mirror of
https://github.com/hansjone/netx.git
synced 2026-10-10 12:40:44 +08:00
Reuse SSH for WebCRT SFTP and fix folder navigation UX.
Open SFTP on the live session transport after SSH connect, pass the target path on every list refresh, and use folder/file icons in the browser panel. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
d87bdf4be1
commit
499b4cc2a1
10 changed files with 525 additions and 131 deletions
|
|
@ -298,6 +298,7 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None:
|
|||
"device_type": sess.device_type,
|
||||
"vendor": sess.vendor,
|
||||
"cli_hop": bool(sess.cli_hop_guard),
|
||||
"sftp_ready": bool(sess.sftp_ready),
|
||||
}
|
||||
)
|
||||
|
||||
|
|
@ -372,6 +373,7 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None:
|
|||
"device_type": sess.device_type,
|
||||
"vendor": sess.vendor,
|
||||
"cli_hop": bool(sess.cli_hop_guard),
|
||||
"sftp_ready": bool(sess.sftp_ready),
|
||||
"connect_ms": (
|
||||
int((sess.connect_finished_at - sess.connect_started_at) * 1000)
|
||||
if sess.connect_finished_at
|
||||
|
|
|
|||
|
|
@ -499,10 +499,85 @@ class WebcrtSession:
|
|||
_cli_hop_seen_other_prompt: bool = field(default=False, repr=False)
|
||||
_log_fh: Any = field(default=None, repr=False)
|
||||
_ready_event: threading.Event = field(default_factory=threading.Event, repr=False)
|
||||
# SFTP channel on the same SSH transport as the interactive shell (direct SSH only).
|
||||
sftp_ready: bool = False
|
||||
_sftp: Any = field(default=None, repr=False)
|
||||
_sftp_lock: threading.RLock = field(default_factory=threading.RLock, repr=False)
|
||||
|
||||
def touch(self) -> None:
|
||||
self.last_activity = time.time()
|
||||
|
||||
def close_sftp(self) -> None:
|
||||
with self._sftp_lock:
|
||||
sftp = self._sftp
|
||||
self._sftp = None
|
||||
self.sftp_ready = False
|
||||
if sftp is None:
|
||||
return
|
||||
try:
|
||||
sftp.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _ensure_sftp_unlocked(self) -> Any:
|
||||
"""Caller must hold ``_sftp_lock``."""
|
||||
import paramiko
|
||||
|
||||
if self.closed or self.conn is None:
|
||||
raise RuntimeError("session_closed")
|
||||
if str(self.protocol or "ssh").lower() != "ssh":
|
||||
raise RuntimeError("sftp_requires_ssh")
|
||||
if self.cli_hop_guard:
|
||||
raise RuntimeError("sftp_hop_not_supported")
|
||||
if self._sftp is not None:
|
||||
sock = getattr(self._sftp, "sock", None)
|
||||
if sock is not None and not bool(getattr(sock, "closed", False)):
|
||||
return self._sftp
|
||||
try:
|
||||
self._sftp.close()
|
||||
except Exception:
|
||||
pass
|
||||
self._sftp = None
|
||||
channel = getattr(self.conn, "remote_conn", None)
|
||||
transport = None
|
||||
if channel is not None and hasattr(channel, "get_transport"):
|
||||
try:
|
||||
transport = channel.get_transport()
|
||||
except Exception:
|
||||
transport = None
|
||||
if transport is None or not bool(getattr(transport, "is_active", lambda: False)()):
|
||||
raise RuntimeError("ssh_transport_unavailable")
|
||||
self._sftp = paramiko.SFTPClient.from_transport(transport)
|
||||
if self._sftp is None:
|
||||
raise RuntimeError("sftp_open_failed")
|
||||
self.sftp_ready = True
|
||||
return self._sftp
|
||||
|
||||
def open_sftp(self) -> Any:
|
||||
"""Open/reuse an SFTP client on this session's SSH transport."""
|
||||
with self._sftp_lock:
|
||||
return self._ensure_sftp_unlocked()
|
||||
|
||||
def run_sftp(self, fn: Any) -> Any:
|
||||
"""Run ``fn(sftp)`` while holding the session SFTP lock."""
|
||||
with self._sftp_lock:
|
||||
sftp = self._ensure_sftp_unlocked()
|
||||
return fn(sftp)
|
||||
|
||||
def try_attach_sftp(self) -> bool:
|
||||
"""Best-effort SFTP channel open after SSH login (does not fail the shell)."""
|
||||
if str(self.protocol or "ssh").lower() != "ssh" or self.cli_hop_guard:
|
||||
self.sftp_ready = False
|
||||
return False
|
||||
try:
|
||||
self.open_sftp()
|
||||
self.sftp_ready = True
|
||||
return True
|
||||
except Exception:
|
||||
self.sftp_ready = False
|
||||
_log.debug("webcrt sftp attach skipped session=%s", self.session_id, exc_info=True)
|
||||
return False
|
||||
|
||||
def open_session_log(self) -> None:
|
||||
if not bool(getattr(settings, "webcrt_session_log_enabled", True)):
|
||||
return
|
||||
|
|
@ -790,6 +865,7 @@ class WebcrtSession:
|
|||
self.state = "closed"
|
||||
self.close_reason = reason or "closed"
|
||||
self._ready_event.set()
|
||||
self.close_sftp()
|
||||
try:
|
||||
close_netmiko_connection(self.conn)
|
||||
except Exception:
|
||||
|
|
@ -892,6 +968,29 @@ def get_session(session_id: str) -> WebcrtSession | None:
|
|||
return sess
|
||||
|
||||
|
||||
def find_ssh_session_for_ne(ne_id: str) -> WebcrtSession | None:
|
||||
"""Prefer a ready/attached interactive SSH session for SFTP channel reuse."""
|
||||
nid = str(ne_id or "").strip()
|
||||
if not nid:
|
||||
return None
|
||||
with _sessions_lock:
|
||||
candidates = [
|
||||
s
|
||||
for s in _sessions.values()
|
||||
if (not s.closed)
|
||||
and str(s.ne_id) == nid
|
||||
and str(s.protocol or "ssh").lower() == "ssh"
|
||||
and s.conn is not None
|
||||
and s.state in ("ready", "connecting")
|
||||
and not s.cli_hop_guard
|
||||
]
|
||||
if not candidates:
|
||||
return None
|
||||
# Prefer attached + ready sessions.
|
||||
candidates.sort(key=lambda s: (0 if s.attached and s.state == "ready" else 1, -s.last_activity))
|
||||
return candidates[0]
|
||||
|
||||
|
||||
def wait_session_ready(session_id: str, *, timeout: float = 120.0) -> WebcrtSession:
|
||||
"""Block until async connect finishes (ready or error). Used by tests and WS."""
|
||||
deadline = time.time() + max(1.0, float(timeout))
|
||||
|
|
@ -1067,6 +1166,8 @@ def _finish_connect(
|
|||
sess.run_post_login_commands()
|
||||
except Exception:
|
||||
_log.debug("post_login failed session=%s", sess.session_id, exc_info=True)
|
||||
# Same SSH transport: open SFTP channel when the device supports it.
|
||||
sftp_ok = sess.try_attach_sftp()
|
||||
sess.state = "ready"
|
||||
sess.connect_finished_at = time.time()
|
||||
sess._ready_event.set()
|
||||
|
|
@ -1084,6 +1185,7 @@ def _finish_connect(
|
|||
hop_vendor=str(creds.get("hop_vendor") or "") if creds.get("hop_enabled") else "",
|
||||
cli_hop_guard=bool(hop_guard),
|
||||
cli_hop_prompt=str((hop_guard or {}).get("hop_prompt") or ""),
|
||||
sftp_ready=bool(sftp_ok),
|
||||
client=client or "",
|
||||
connect_ms=elapsed_ms,
|
||||
active=active_session_count(),
|
||||
|
|
@ -1226,6 +1328,7 @@ def create_session(
|
|||
"state": sess.state,
|
||||
"ws_path": f"/v1/webcrt/sessions/{session_id}/ws",
|
||||
"cli_hop": bool(sess.cli_hop_guard),
|
||||
"sftp_ready": bool(sess.sftp_ready),
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,10 +1,15 @@
|
|||
"""Lightweight SFTP helpers for WebCRT (SSH targets only; separate from interactive PTY)."""
|
||||
"""WebCRT SFTP helpers — prefer the live SSH session channel; pool only as fallback."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import posixpath
|
||||
from typing import Any
|
||||
import stat as statmod
|
||||
import threading
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Iterator
|
||||
|
||||
import paramiko
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -13,21 +18,64 @@ from sqlalchemy.orm import Session
|
|||
from .cli_resolve import resolve_cli_target
|
||||
from .config import settings
|
||||
from .ne_crypto import CredentialCryptoError
|
||||
from .webcrt_service import _webcrt_creds_ready
|
||||
from .webcrt_service import (
|
||||
_webcrt_creds_ready,
|
||||
find_ssh_session_for_ne,
|
||||
)
|
||||
|
||||
_log = logging.getLogger("netx.webcrt.sftp")
|
||||
|
||||
_pool_lock = threading.Lock()
|
||||
_pool: dict[str, "_PooledSftp"] = {}
|
||||
_POOL_IDLE_SEC = 180
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PooledSftp:
|
||||
key: str
|
||||
client: paramiko.SSHClient
|
||||
sftp: paramiko.SFTPClient
|
||||
last_used: float = field(default_factory=time.time)
|
||||
lock: threading.RLock = field(default_factory=threading.RLock)
|
||||
|
||||
|
||||
def _require_ssh_direct(creds: dict[str, Any], device: dict[str, Any]) -> None:
|
||||
protocol = str(device.get("protocol") or creds.get("protocol") or "ssh").lower()
|
||||
if protocol != "ssh":
|
||||
raise HTTPException(status_code=400, detail="sftp_requires_ssh")
|
||||
if creds.get("hop_enabled"):
|
||||
# Keep v1 simple: SFTP only for direct SSH (no hop/proxy jump).
|
||||
raise HTTPException(status_code=400, detail="sftp_hop_not_supported")
|
||||
|
||||
|
||||
def _open_sftp(creds: dict[str, Any]) -> tuple[paramiko.SSHClient, paramiko.SFTPClient]:
|
||||
def _pool_key(*, managed_ne_id: str | None, ume_ne_id: str | None) -> str:
|
||||
mid = str(managed_ne_id or "").strip()
|
||||
uid = str(ume_ne_id or "").strip()
|
||||
if mid:
|
||||
return f"m:{mid}"
|
||||
return f"u:{uid}"
|
||||
|
||||
|
||||
def _close_pooled(entry: _PooledSftp) -> None:
|
||||
try:
|
||||
entry.sftp.close()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
entry.client.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _reap_pool_unlocked(now: float | None = None) -> None:
|
||||
ts = float(now or time.time())
|
||||
dead = [k for k, e in _pool.items() if (ts - e.last_used) > _POOL_IDLE_SEC]
|
||||
for k in dead:
|
||||
entry = _pool.pop(k, None)
|
||||
if entry is not None:
|
||||
_close_pooled(entry)
|
||||
|
||||
|
||||
def _open_pooled_sftp(creds: dict[str, Any]) -> tuple[paramiko.SSHClient, paramiko.SFTPClient]:
|
||||
timeout = int(settings.ne_connect_timeout_sec or 30)
|
||||
client = paramiko.SSHClient()
|
||||
client.set_missing_host_key_policy(paramiko.AutoAddPolicy())
|
||||
|
|
@ -45,12 +93,6 @@ def _open_sftp(creds: dict[str, Any]) -> tuple[paramiko.SSHClient, paramiko.SFTP
|
|||
)
|
||||
sftp = client.open_sftp()
|
||||
return client, sftp
|
||||
except HTTPException:
|
||||
try:
|
||||
client.close()
|
||||
except Exception:
|
||||
pass
|
||||
raise
|
||||
except Exception as exc:
|
||||
try:
|
||||
client.close()
|
||||
|
|
@ -59,6 +101,23 @@ def _open_sftp(creds: dict[str, Any]) -> tuple[paramiko.SSHClient, paramiko.SFTP
|
|||
raise HTTPException(status_code=502, detail=f"sftp_connect_failed:{exc}") from exc
|
||||
|
||||
|
||||
def _get_pooled(key: str, creds: dict[str, Any]) -> _PooledSftp:
|
||||
with _pool_lock:
|
||||
_reap_pool_unlocked()
|
||||
entry = _pool.get(key)
|
||||
if entry is not None:
|
||||
sock = getattr(entry.sftp, "sock", None)
|
||||
if sock is not None and not bool(getattr(sock, "closed", False)):
|
||||
entry.last_used = time.time()
|
||||
return entry
|
||||
_pool.pop(key, None)
|
||||
_close_pooled(entry)
|
||||
client, sftp = _open_pooled_sftp(creds)
|
||||
entry = _PooledSftp(key=key, client=client, sftp=sftp)
|
||||
_pool[key] = entry
|
||||
return entry
|
||||
|
||||
|
||||
def _resolve(db: Session, *, managed_ne_id: str | None, ume_ne_id: str | None) -> tuple[dict[str, Any], dict[str, Any]]:
|
||||
try:
|
||||
creds, device = resolve_cli_target(db, managed_ne_id=managed_ne_id, ume_ne_id=ume_ne_id)
|
||||
|
|
@ -74,6 +133,57 @@ def _resolve(db: Session, *, managed_ne_id: str | None, ume_ne_id: str | None) -
|
|||
return creds, device
|
||||
|
||||
|
||||
def _normalize_remote(path: str, *, allow_dot: bool = True) -> str:
|
||||
remote = str(path or "").strip() or ("." if allow_dot else "")
|
||||
if not remote:
|
||||
return remote
|
||||
if remote not in (".", "/"):
|
||||
remote = posixpath.normpath(remote.replace("\\", "/")) or ("." if allow_dot else "")
|
||||
return remote
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _sftp_client(
|
||||
db: Session,
|
||||
*,
|
||||
managed_ne_id: str | None,
|
||||
ume_ne_id: str | None,
|
||||
) -> Iterator[tuple[Any, dict[str, Any]]]:
|
||||
"""Yield ``(sftp, device)`` — prefers live WebCRT SSH session channel."""
|
||||
creds, device = _resolve(db, managed_ne_id=managed_ne_id, ume_ne_id=ume_ne_id)
|
||||
ne_key = str(device.get("id") or managed_ne_id or ume_ne_id or "").strip()
|
||||
sess = find_ssh_session_for_ne(ne_key) if ne_key else None
|
||||
if sess is not None:
|
||||
opened = False
|
||||
try:
|
||||
with sess._sftp_lock:
|
||||
sftp = sess._ensure_sftp_unlocked()
|
||||
opened = True
|
||||
yield sftp, device
|
||||
return
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
if opened:
|
||||
# Operation failed on an already-open session channel — don't double-yield.
|
||||
raise HTTPException(status_code=502, detail=f"sftp_failed:{exc}") from exc
|
||||
_log.debug("session sftp open failed ne=%s: %s — pool fallback", ne_key, exc)
|
||||
|
||||
key = _pool_key(managed_ne_id=managed_ne_id, ume_ne_id=ume_ne_id)
|
||||
entry = _get_pooled(key, creds)
|
||||
with entry.lock:
|
||||
entry.last_used = time.time()
|
||||
try:
|
||||
yield entry.sftp, device
|
||||
except Exception:
|
||||
# Drop broken pooled socket so the next call reconnects once.
|
||||
with _pool_lock:
|
||||
cur = _pool.pop(key, None)
|
||||
if cur is not None:
|
||||
_close_pooled(cur)
|
||||
raise
|
||||
|
||||
|
||||
def sftp_list(
|
||||
db: Session,
|
||||
*,
|
||||
|
|
@ -81,42 +191,34 @@ def sftp_list(
|
|||
ume_ne_id: str | None,
|
||||
path: str = ".",
|
||||
) -> dict[str, Any]:
|
||||
creds, device = _resolve(db, managed_ne_id=managed_ne_id, ume_ne_id=ume_ne_id)
|
||||
remote = str(path or ".").strip() or "."
|
||||
client, sftp = _open_sftp(creds)
|
||||
remote = _normalize_remote(path, allow_dot=True) or "."
|
||||
try:
|
||||
entries = []
|
||||
for attr in sftp.listdir_attr(remote):
|
||||
mode = int(getattr(attr, "st_mode", 0) or 0)
|
||||
is_dir = bool(mode & 0o40000)
|
||||
entries.append(
|
||||
{
|
||||
"name": attr.filename,
|
||||
"size": int(getattr(attr, "st_size", 0) or 0),
|
||||
"mtime": int(getattr(attr, "st_mtime", 0) or 0),
|
||||
"is_dir": is_dir,
|
||||
}
|
||||
)
|
||||
entries.sort(key=lambda x: (not x["is_dir"], str(x["name"]).lower()))
|
||||
return {
|
||||
"ne_id": str(device.get("id") or ""),
|
||||
"ne_name": str(device.get("name") or ""),
|
||||
"path": remote,
|
||||
"items": entries,
|
||||
}
|
||||
with _sftp_client(db, managed_ne_id=managed_ne_id, ume_ne_id=ume_ne_id) as (sftp, device):
|
||||
entries = []
|
||||
for attr in sftp.listdir_attr(remote):
|
||||
mode = int(getattr(attr, "st_mode", 0) or 0)
|
||||
name = str(attr.filename or "")
|
||||
if not name or name in (".", ".."):
|
||||
continue
|
||||
entries.append(
|
||||
{
|
||||
"name": name,
|
||||
"size": int(getattr(attr, "st_size", 0) or 0),
|
||||
"mtime": int(getattr(attr, "st_mtime", 0) or 0),
|
||||
"is_dir": bool(statmod.S_ISDIR(mode)),
|
||||
}
|
||||
)
|
||||
entries.sort(key=lambda x: (not x["is_dir"], str(x["name"]).lower()))
|
||||
return {
|
||||
"ne_id": str(device.get("id") or ""),
|
||||
"ne_name": str(device.get("name") or ""),
|
||||
"path": remote,
|
||||
"items": entries,
|
||||
}
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=502, detail=f"sftp_list_failed:{exc}") from exc
|
||||
finally:
|
||||
try:
|
||||
sftp.close()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
client.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def sftp_download(
|
||||
|
|
@ -126,14 +228,13 @@ def sftp_download(
|
|||
ume_ne_id: str | None,
|
||||
path: str,
|
||||
) -> tuple[bytes, str]:
|
||||
creds, _device = _resolve(db, managed_ne_id=managed_ne_id, ume_ne_id=ume_ne_id)
|
||||
remote = str(path or "").strip()
|
||||
remote = _normalize_remote(path, allow_dot=False)
|
||||
if not remote or remote.endswith("/"):
|
||||
raise HTTPException(status_code=400, detail="sftp_path_required")
|
||||
client, sftp = _open_sftp(creds)
|
||||
try:
|
||||
with sftp.open(remote, "rb") as fh:
|
||||
data = fh.read(8 * 1024 * 1024 + 1)
|
||||
with _sftp_client(db, managed_ne_id=managed_ne_id, ume_ne_id=ume_ne_id) as (sftp, _device):
|
||||
with sftp.open(remote, "rb") as fh:
|
||||
data = fh.read(8 * 1024 * 1024 + 1)
|
||||
if len(data) > 8 * 1024 * 1024:
|
||||
raise HTTPException(status_code=413, detail="sftp_file_too_large")
|
||||
return data, posixpath.basename(remote) or "download.bin"
|
||||
|
|
@ -141,15 +242,6 @@ def sftp_download(
|
|||
raise
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=502, detail=f"sftp_download_failed:{exc}") from exc
|
||||
finally:
|
||||
try:
|
||||
sftp.close()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
client.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def sftp_upload(
|
||||
|
|
@ -160,30 +252,20 @@ def sftp_upload(
|
|||
remote_path: str,
|
||||
data: bytes,
|
||||
) -> dict[str, Any]:
|
||||
creds, device = _resolve(db, managed_ne_id=managed_ne_id, ume_ne_id=ume_ne_id)
|
||||
remote = str(remote_path or "").strip()
|
||||
remote = _normalize_remote(remote_path, allow_dot=False)
|
||||
if not remote:
|
||||
raise HTTPException(status_code=400, detail="sftp_path_required")
|
||||
client, sftp = _open_sftp(creds)
|
||||
try:
|
||||
with sftp.open(remote, "wb") as fh:
|
||||
fh.write(data)
|
||||
return {
|
||||
"ok": True,
|
||||
"ne_id": str(device.get("id") or ""),
|
||||
"path": remote,
|
||||
"size": len(data),
|
||||
}
|
||||
with _sftp_client(db, managed_ne_id=managed_ne_id, ume_ne_id=ume_ne_id) as (sftp, device):
|
||||
with sftp.open(remote, "wb") as fh:
|
||||
fh.write(data)
|
||||
return {
|
||||
"ok": True,
|
||||
"ne_id": str(device.get("id") or ""),
|
||||
"path": remote,
|
||||
"size": len(data),
|
||||
}
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=502, detail=f"sftp_upload_failed:{exc}") from exc
|
||||
finally:
|
||||
try:
|
||||
sftp.close()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
client.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue