"""SSH local forward with Nova watchdog. Ctrl+C and Ctrl+D (EOF) both run ``stop`` (disks kept). """ from __future__ import annotations import os import sys import time import webbrowser from collections.abc import Callable from dataclasses import dataclass from gpu_rent.cloud import ( ensure_floating_ip, pick_existing_server, server_status, unshelve, ) from gpu_rent.config import Config from gpu_rent.errors import CloudError, GpuRentError from gpu_rent.llm_runtime import normalize_runtime from gpu_rent.os_client import connect from gpu_rent.ssh_ops import wait_ssh from gpu_rent.state import load_state, save_state, utc_now Log = Callable[[str], None] EXIT_STATUSES = frozenset( {"ERROR", "DELETED", "SOFT_DELETED", "UNKNOWN", "BUILD_FAILED"} ) SHELVED_STATUSES = frozenset({"EXPIRED", "SHELVED", "SHELVED_OFFLOADED"}) def _patch_paramiko_for_sshtunnel() -> None: """sshtunnel 0.4 still refs paramiko.DSSKey; Paramiko 4+ removed it.""" import paramiko if hasattr(paramiko, "DSSKey"): return class _DSSKeyRemoved(paramiko.PKey): def __init__(self, *args, **kwargs): raise paramiko.SSHException("DSA keys unsupported (paramiko>=4)") paramiko.DSSKey = _DSSKeyRemoved # type: ignore[attr-defined, assignment] def _ssh_tunnel_forwarder(): _patch_paramiko_for_sshtunnel() from sshtunnel import SSHTunnelForwarder return SSHTunnelForwarder @dataclass class WatchDecision: kind: str # ok | reconnect | unshelve | exit detail: str = "" def _poll_nova(cfg: Config, log: Log) -> tuple[str | None, str]: """Return (status, detail). Refreshes IAM token via connect(). On soft auth/API failure return SOFT_FAIL (not fake ACTIVE) so we do not mask DELETED/ERROR forever. """ try: conn = connect(cfg) server = pick_existing_server(conn) if not server: return None, "нет сервера" return server_status(server), server.id except GpuRentError as exc: log(f"watch: OpenStack временно недоступен ({exc})") return "SOFT_FAIL", "auth-soft-fail" def decide_watch(status: str | None, tunnel_alive: bool) -> WatchDecision: """Pure policy for tunnel watchdog (unit-tested).""" if not status: return WatchDecision("exit", "нет сервера gpu-rent") st = status.upper() if st == "SOFT_FAIL": # Transient OpenStack blip — keep tunnel, do not pretend ACTIVE forever. return WatchDecision("ok", "openstack soft-fail") if st in EXIT_STATUSES: return WatchDecision("exit", f"Nova {st}") if st in SHELVED_STATUSES: return WatchDecision("unshelve", f"Nova {st}") if st == "ACTIVE" and not tunnel_alive: return WatchDecision("reconnect", "туннель мёртв, сервер ACTIVE") if st == "ACTIVE": return WatchDecision("ok", "ACTIVE") return WatchDecision("ok", f"ждём {st}") def tunnel_forwards(cfg: Config) -> list[tuple[int, int]]: """List of (local_port, remote_port). SwarmUI if enabled; LLM if configured.""" pairs: list[tuple[int, int]] = [] if bool(getattr(cfg, "enable_swarmui", True)): pairs.append((cfg.swarmui_local_port, 7801)) runtime = normalize_runtime(cfg.llm_runtime) if runtime == "ollama": pairs.append((cfg.ollama_local_port, 11434)) if not pairs: # Failsafe: at least SwarmUI port so tunnel isn't empty. pairs.append((cfg.swarmui_local_port, 7801)) return pairs def _start_forwarder(cfg: Config, host: str, forwards: list[tuple[int, int]] | None = None): SSHTunnelForwarder = _ssh_tunnel_forwarder() pairs = forwards or tunnel_forwards(cfg) local_binds = [("127.0.0.1", loc) for loc, _ in pairs] remote_binds = [("127.0.0.1", rem) for _, rem in pairs] server = SSHTunnelForwarder( (host, 22), ssh_username=cfg.ssh_user, ssh_pkey=str(cfg.ssh_private_key_path), remote_bind_addresses=remote_binds, local_bind_addresses=local_binds, set_keepalive=30, ) try: server.start() except Exception as exc: ports = ",".join(str(p[0]) for p in pairs) raise CloudError( f"не открыть туннель на {ports}: {exc}. Порт занят?" ) from exc return server def _stop_forwarder(server) -> None: if server is None: return try: server.stop() except Exception: pass def poll_ctrl_d(timeout: float = 1.0) -> bool: """True if the user sent Ctrl+D / EOF. Ctrl+C stays KeyboardInterrupt. Windows console delivers Ctrl+D as ``\\x04`` (and Ctrl+Z as ``\\x1a``). Those keys are ignored unless we read them — the old sleep-loop never did. """ try: if not sys.stdin.isatty(): if timeout > 0: time.sleep(timeout) return False except Exception: if timeout > 0: time.sleep(timeout) return False if os.name == "nt": try: import msvcrt except ImportError: if timeout > 0: time.sleep(timeout) return False deadline = time.time() + max(timeout, 0.0) while True: if msvcrt.kbhit(): ch = msvcrt.getch() if ch in (b"\x00", b"\xe0") and msvcrt.kbhit(): msvcrt.getch() continue if ch in (b"\x04", b"\x1a"): return True if ch == b"\x03": raise KeyboardInterrupt continue if time.time() >= deadline: return False time.sleep(0.05) import select r, _, _ = select.select([sys.stdin], [], [], max(timeout, 0.0)) if not r: return False try: data = os.read(sys.stdin.fileno(), 64) except OSError: return False return (not data) or (b"\x04" in data) def _recover_unshelve(cfg: Config, log: Log) -> str: """Unshelve EXPIRED VM, rebind FIP, wait SSH. Returns new host.""" conn = connect(cfg) server = pick_existing_server(conn) if not server: raise CloudError("сервер gpu-rent исчез во время EXPIRED") status = server_status(server) if status in SHELVED_STATUSES: server = unshelve(conn, server, log) elif status != "ACTIVE": raise CloudError(f"после preempt статус {status} — не unshelve") state = load_state() ip, fip_id = ensure_floating_ip( conn, server, state.floating_ip_id, state.floating_ip, log ) state.floating_ip = ip if fip_id: state.floating_ip_id = fip_id state.server_id = server.id state.unshelved_at = utc_now() state.phase = "ready_tunneled" save_state(state) log(f"жду SSH на {ip}…") wait_ssh(cfg, ip, timeout=420, log=log) return ip def run_tunnel( cfg: Config, host: str, *, open_browser: bool = False, log: Log = print, wait: Callable[[], None] | None = None, poll_seconds: float = 30.0, stop_gpu: Callable[[], None] | None = None, session_end_poll: Callable[[float], bool] | None = None, ) -> None: try: _ssh_tunnel_forwarder() except ImportError as exc: raise CloudError("Нет sshtunnel. Переустанови пакет: pip install -e .") from exc forwards = tunnel_forwards(cfg) current_host = host for loc, rem in forwards: log(f"туннель 127.0.0.1:{loc} -> {current_host}:{rem}") log("Ctrl+C / Ctrl+D — stop GPU (диски остаются).") log("watchdog: EXPIRED → unshelve + reconnect") server = _start_forwarder(cfg, current_host, forwards) swarm_on = bool(getattr(cfg, "enable_swarmui", True)) runtime = normalize_runtime(cfg.llm_runtime) if swarm_on: open_url = f"http://127.0.0.1:{cfg.swarmui_local_port}" elif runtime == "ollama": open_url = f"http://127.0.0.1:{cfg.ollama_local_port}" else: open_url = f"http://127.0.0.1:{cfg.swarmui_local_port}" from gpu_rent.local_watchdog import ( clear_lease, start_heartbeat_thread, stop_heartbeat_thread, watchdog_installed, ) def _halt_gpu(reason: str) -> None: nonlocal server log(f"{reason} — гашу GPU (диски остаются)") _stop_forwarder(server) server = None stop_heartbeat_thread() clear_lease() if stop_gpu is not None: stop_gpu() else: from gpu_rent.session import cmd_stop cmd_stop(cfg, log=log) log("туннель закрыт. GPU остановлен.") try: if watchdog_installed(): start_heartbeat_thread() log( "local-watchdog: heartbeat активен — kill/reboot без stop " "→ stop после grace" ) from gpu_rent.ready import verify_stack_local try: local_checks = verify_stack_local(cfg, log, timeout=90.0) state = load_state() state.notes = dict(state.notes or {}) state.notes["stack_local"] = [ {"name": c.name, "ok": c.ok, "detail": c.detail} for c in local_checks ] save_state(state) except CloudError as exc: log(f"проверка туннеля: {exc}") raise from gpu_rent.llm_runtime import maybe_warmup_ollama_local maybe_warmup_ollama_local(cfg, log) from gpu_rent.access_card import print_access_card from gpu_rent.term import console as term_console from gpu_rent.vm_logs import print_log_digest print_access_card(cfg, tunneled=True, host=current_host) print_log_digest(cfg, current_host, console=term_console) if open_browser: webbrowser.open(open_url) state = load_state() state.phase = "ready_tunneled" save_state(state) if wait is not None: wait() return end_poll = session_end_poll or poll_ctrl_d next_poll = time.time() + poll_seconds while True: if end_poll(1.0): _halt_gpu("Ctrl+D") return if not server.is_active: next_poll = 0 if time.time() < next_poll: continue next_poll = time.time() + poll_seconds tunnel_alive = bool(server.is_active) status, detail = _poll_nova(cfg, log) decision = decide_watch(status, tunnel_alive) if decision.kind == "ok": continue if decision.kind == "exit": log(f"watchdog стоп: {decision.detail} ({detail})") return if decision.kind == "reconnect": log(f"reconnect: {decision.detail}") _stop_forwarder(server) try: server = _start_forwarder(cfg, current_host, forwards) log(f"туннель снова на {current_host}") except CloudError as exc: log(f"reconnect не вышел: {exc}") time.sleep(10) continue if decision.kind == "unshelve": log(f"watchdog: {decision.detail} — unshelve…") _stop_forwarder(server) try: current_host = _recover_unshelve(cfg, log) server = _start_forwarder(cfg, current_host, forwards) log(f"туннель после unshelve → {current_host}") except (CloudError, GpuRentError) as exc: log(f"unshelve/reconnect fail: {exc}") return except KeyboardInterrupt: _halt_gpu("Ctrl+C") finally: stop_heartbeat_thread() _stop_forwarder(server)