"""Terminal session manager for WebSocket connections.""" import asyncio import logging import uuid from typing import Any from fastapi import WebSocket from src.services.terminal_session import TerminalSession logger = logging.getLogger(__name__) class TerminalManager: """Manages active terminal sessions with persistence support.""" def __init__(self) -> None: # Track sessions by instance_id for persistence self._sessions: dict[str, TerminalSession] = {} self._idle_check_task: asyncio.Task | None = None self._start_idle_check() def _start_idle_check(self) -> None: """Start the idle timeout background task.""" if self._idle_check_task is None or self._idle_check_task.done(): self._idle_check_task = asyncio.create_task(self._idle_check_loop()) async def _idle_check_loop(self) -> None: """Periodically check for idle sessions and clean them up.""" while True: try: await asyncio.sleep(60) # Check every minute await self._cleanup_idle_sessions() except Exception as exc: logger.error("Error in idle check loop: %s", exc) async def _cleanup_idle_sessions(self) -> None: """Clean up sessions that have been idle for too long.""" idle_sessions = [] for instance_id, session in list(self._sessions.items()): if session.is_idle(): idle_sessions.append(instance_id) for instance_id in idle_sessions: logger.info("Cleaning up idle terminal session for instance %s", instance_id) session = self._sessions.pop(instance_id, None) if session: await session.close() async def get_or_create_session( self, instance_id: uuid.UUID, container_id: str, ) -> TerminalSession: """Get existing session or create a new one.""" instance_id_str = str(instance_id) # Check for existing session if instance_id_str in self._sessions: session = self._sessions[instance_id_str] # Check if session is still alive if session.is_alive(): logger.info("Reattaching to existing terminal session for instance %s", instance_id) return session else: # Session died, clean it up logger.info("Existing session for instance %s is dead, cleaning up", instance_id) await session.close() del self._sessions[instance_id_str] # Create new session logger.info("Creating new terminal session for instance %s", instance_id) session_id = str(uuid.uuid4()) session = TerminalSession(session_id, instance_id, container_id) await session.start() self._sessions[instance_id_str] = session return session async def attach_websocket( self, session: TerminalSession, websocket: WebSocket, ) -> None: """Attach a WebSocket to an existing session.""" # Handle concurrent connections - close existing ones if session.has_websockets(): logger.info("Closing existing WebSocket connections for instance %s", session.instance_id) for ws in list(session._websockets): try: await ws.close(code=4000, reason="New connection established") except Exception: pass session._websockets.clear() # Attach new WebSocket session.attach_websocket(websocket) # Replay buffer buffer = session.get_buffer() if buffer: try: await websocket.send_bytes(buffer) except Exception: pass async def detach_websocket( self, session: TerminalSession, websocket: WebSocket, ) -> None: """Detach a WebSocket from a session.""" session.detach_websocket(websocket) async def reset_session( self, instance_id: uuid.UUID, container_id: str, ) -> TerminalSession: """Reset a session by killing it and creating a new one.""" instance_id_str = str(instance_id) # Close existing session if any if instance_id_str in self._sessions: logger.info("Resetting terminal session for instance %s", instance_id) old_session = self._sessions.pop(instance_id_str) await old_session.close() # Create new session session_id = str(uuid.uuid4()) session = TerminalSession(session_id, instance_id, container_id) await session.start() self._sessions[instance_id_str] = session return session async def close_all(self) -> None: """Close all active sessions.""" sessions = list(self._sessions.values()) self._sessions.clear() for session in sessions: await session.close() if self._idle_check_task and not self._idle_check_task.done(): self._idle_check_task.cancel() # Global terminal manager instance terminal_manager = TerminalManager()