"""Terminal session manager for WebSocket connections.""" import asyncio import contextlib import json import logging import time import uuid from collections.abc import Coroutine from typing import Any from fastapi import WebSocket from src.services.terminal_session import TerminalSession logger = logging.getLogger(__name__) _READ_BATCH_INTERVAL_S = 0.016 # 16ms max batching delay _READ_POLL_TIMEOUT_S = 0.005 _READ_POLL_SLEEP_S = 0.001 _HEARTBEAT_INTERVAL_S = 15.0 _IDLE_TIMEOUT_S = 60.0 class TerminalManager: """Manages active terminal sessions.""" def __init__(self) -> None: """Initialise the terminal manager.""" self._sessions: dict[str, TerminalSession] = {} self._last_client_message: dict[str, float] = {} self._background_tasks: set[asyncio.Task[Any]] = set() async def create_session( self, instance_id: uuid.UUID, container_id: str, websocket: WebSocket, ) -> TerminalSession: """Create a new terminal session.""" session_id = str(uuid.uuid4()) session = TerminalSession(session_id, instance_id, container_id) await session.start() self._sessions[session_id] = session self._last_client_message[session_id] = time.monotonic() # Start background tasks for I/O streaming self._start_task(self._read_loop(session, websocket)) self._start_task(self._write_loop(session, websocket)) self._start_task(self._heartbeat_loop(session, websocket)) return session def _start_task(self, coro: Coroutine[Any, Any, None]) -> None: """Start a background task and store a reference to prevent GC.""" task = asyncio.create_task(coro) self._background_tasks.add(task) task.add_done_callback(self._background_tasks.discard) async def _read_loop( self, session: TerminalSession, websocket: WebSocket, ) -> None: """Read output from the container and send to WebSocket with batching.""" try: buffer = bytearray() last_flush = time.monotonic() while session.is_alive() and not session.closed: data = await session.read_output(select_timeout=_READ_POLL_TIMEOUT_S) if data: buffer.extend(data) now = time.monotonic() flush_due = buffer and ( now - last_flush >= _READ_BATCH_INTERVAL_S or not data ) if flush_due: await websocket.send_bytes(bytes(buffer)) buffer.clear() last_flush = now elif not data: await asyncio.sleep(_READ_POLL_SLEEP_S) # Flush any remaining data if buffer: with contextlib.suppress(Exception): await websocket.send_bytes(bytes(buffer)) except Exception: logger.exception("Read loop error for session %s", session.session_id) finally: await self._cleanup_session(session) async def _write_loop( self, session: TerminalSession, websocket: WebSocket, ) -> None: """Read input from WebSocket and send to container.""" try: while session.is_alive() and not session.closed: message = await websocket.receive() self._last_client_message[session.session_id] = time.monotonic() if message["type"] == "websocket.receive": if "bytes" in message: await session.write_input(message["bytes"]) elif "text" in message: text = message["text"] if text.startswith("{"): try: ctrl = json.loads(text) await self._handle_control_message( session, websocket, ctrl, ) except json.JSONDecodeError: logger.debug("Invalid JSON control message: %s", text) else: await session.write_input(text.encode("utf-8")) elif message["type"] == "websocket.disconnect": break except Exception: logger.exception("Write loop error for session %s", session.session_id) finally: await self._cleanup_session(session) async def _handle_control_message( self, session: TerminalSession, websocket: WebSocket, ctrl: dict[str, Any], ) -> None: """Handle a JSON control message from the client.""" msg_type = ctrl.get("type") if msg_type == "resize": await session.resize( ctrl.get("cols", 80), ctrl.get("rows", 24), ) elif msg_type == "ping": await websocket.send_json( {"type": "pong", "id": ctrl.get("id")}, ) async def _heartbeat_loop( self, session: TerminalSession, websocket: WebSocket, ) -> None: """Monitor client activity and close idle connections.""" try: while session.is_alive() and not session.closed: await asyncio.sleep(_HEARTBEAT_INTERVAL_S) last_msg = self._last_client_message.get(session.session_id, 0) if time.monotonic() - last_msg > _IDLE_TIMEOUT_S: # Client has been silent for 60s — close connection with contextlib.suppress(Exception): await websocket.close( code=1000, reason="Idle timeout", ) break except Exception: logger.exception( "Heartbeat loop error for session %s", session.session_id, ) finally: await self._cleanup_session(session) async def _cleanup_session(self, session: TerminalSession) -> None: """Clean up a session.""" if session.session_id in self._sessions: del self._sessions[session.session_id] self._last_client_message.pop(session.session_id, None) await session.close() async def close_all(self) -> None: """Close all active sessions.""" sessions = list(self._sessions.values()) self._sessions.clear() self._last_client_message.clear() for session in sessions: await session.close() # Global terminal manager instance terminal_manager = TerminalManager()