"""Terminal session manager for WebSocket connections.""" import asyncio import logging import uuid from datetime import datetime, timezone from fastapi import WebSocket from sqlalchemy.dialects.postgresql import insert as pg_insert from src.database import SessionLocal from src.models import TerminalSessionModel from src.services.terminal.terminal_session import TerminalSession logger = logging.getLogger(__name__) class MaxSessionsExceededError(Exception): """Raised when the maximum number of terminal sessions per instance is reached.""" def __init__(self, instance_id: str, max_sessions: int = 5) -> None: self.instance_id = instance_id self.max_sessions = max_sessions super().__init__( f"Maximum of {max_sessions} terminal sessions reached for instance {instance_id}" ) class TerminalManager: """Manages active terminal sessions with persistence support.""" # Maximum sessions per tool instance MAX_SESSIONS_PER_INSTANCE = 5 def __init__(self) -> None: # Track sessions by (instance_id, session_id) for multi-session support self._sessions: dict[tuple[str, 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 not None and not self._idle_check_task.done(): return try: loop = asyncio.get_running_loop() self._idle_check_task = loop.create_task(self._idle_check_loop()) except RuntimeError: # No event loop running yet, will be started lazily pass 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_keys = [] for (instance_id, session_id), session in list(self._sessions.items()): if session.is_idle(): idle_keys.append((instance_id, session_id)) for key in idle_keys: instance_id, session_id = key logger.info( "Cleaning up idle terminal session %s for instance %s", session_id, instance_id, ) idle_session = self._sessions.get(key) if idle_session is not None: del self._sessions[key] await idle_session.close() # Update DB status fire-and-forget asyncio.create_task(self._mark_closed_in_db(session_id)) async def _insert_db_session_row( self, session_id: str, instance_id: uuid.UUID, name: str, ) -> None: """Insert a TerminalSessionModel row into the database. Uses ON CONFLICT DO NOTHING to handle races when a session is restored from DB and then re-inserted. """ try: async with SessionLocal() as db_session: stmt = ( pg_insert(TerminalSessionModel) .values( id=uuid.UUID(session_id), instance_id=instance_id, name=name, status="active", created_at=datetime.now(timezone.utc), last_activity_at=datetime.now(timezone.utc), ) .on_conflict_do_nothing(index_elements=["id"]) ) await db_session.execute(stmt) await db_session.commit() logger.debug( "Inserted terminal session row %s for instance %s", session_id, instance_id, ) except Exception as exc: logger.error("Failed to insert terminal session row: %s", exc) async def _mark_closed_in_db(self, session_id: str) -> None: """Mark a terminal session as closed in the database.""" try: async with SessionLocal() as db_session: db_row = await db_session.get( TerminalSessionModel, uuid.UUID(session_id) ) if db_row: db_row.status = "closed" db_row.closed_at = datetime.now(timezone.utc) await db_session.commit() logger.debug( "Marked terminal session %s as closed in DB", session_id ) except Exception as exc: logger.error("Failed to mark terminal session as closed in DB: %s", exc) def _count_sessions_for_instance(self, instance_id_str: str) -> int: """Count active in-memory sessions for a given instance.""" return sum(1 for (iid, _sid) in self._sessions if iid == instance_id_str) async def create_session( self, instance_id: uuid.UUID, container_id: str, startup_command: str | None = None, name: str | None = None, session_id: str | None = None, container_user: str | None = None, ) -> TerminalSession: """Create a new terminal session for an instance. Enforces a maximum of MAX_SESSIONS_PER_INSTANCE sessions per instance. Inserts a DB row fire-and-forget. Args: instance_id: UUID of the tool instance. container_id: Docker container ID. startup_command: Optional startup command to run. name: Optional session name (auto-generated if omitted). session_id: Optional explicit session UUID. container_user: Optional container user for docker exec. Returns: The newly created TerminalSession. Raises: MaxSessionsExceededError: If the instance already has max sessions. """ instance_id_str = str(instance_id) if ( self._count_sessions_for_instance(instance_id_str) >= self.MAX_SESSIONS_PER_INSTANCE ): raise MaxSessionsExceededError( instance_id_str, self.MAX_SESSIONS_PER_INSTANCE ) if session_id is None: session_id = str(uuid.uuid4()) session = TerminalSession( session_id=session_id, instance_id=instance_id, container_id=container_id, startup_command=startup_command, name=name, container_user=container_user, ) await session.start(startup_command=startup_command) key = (instance_id_str, session_id) self._sessions[key] = session # Fire-and-forget DB insert (skip if row already exists) asyncio.create_task( self._insert_db_session_row(session_id, instance_id, session.name) ) logger.info( "Created terminal session %s for instance %s (name=%s)", session_id, instance_id, session.name, ) return session async def get_or_create_session( self, instance_id: uuid.UUID, container_id: str, startup_command: str | None = None, container_user: str | None = None, ) -> TerminalSession: """Get existing session or create a new one. Backward-compatible alias that uses 'default' as the session_id. """ # Ensure idle check is running (lazy start) self._start_idle_check() instance_id_str = str(instance_id) key = (instance_id_str, "default") # Check for existing default session if key in self._sessions: session = self._sessions[key] # Check if session is still alive if session.is_alive(): logger.debug( "Reattaching to existing terminal session for instance %s", instance_id, ) return session else: # Session died, clean it up logger.debug( "Existing session for instance %s is dead, cleaning up", instance_id, ) await session.close() del self._sessions[key] # Create new default session logger.info( "Creating new default terminal session for instance %s", instance_id ) session_id = str(uuid.uuid4()) session = TerminalSession( session_id=session_id, instance_id=instance_id, container_id=container_id, startup_command=startup_command, name="Session 1", container_user=container_user, ) await session.start(startup_command=startup_command) self._sessions[key] = session # Fire-and-forget DB insert asyncio.create_task( self._insert_db_session_row(session_id, instance_id, session.name) ) return session def get_session( self, instance_id: str, session_id: str, ) -> TerminalSession | None: """Lookup a session by composite key, or by internal session_id.""" session = self._sessions.get((instance_id, session_id)) if session is not None: return session # Fallback: search by internal TerminalSession.session_id for (iid, _sid), sess in self._sessions.items(): if iid == instance_id and sess.session_id == session_id: return sess return None def _find_key_by_internal_id( self, instance_id: str, internal_session_id: str, ) -> tuple[str, str] | None: """Find the manager dict key for a session by its internal session_id.""" for (iid, sid), session in self._sessions.items(): if iid == instance_id and session.session_id == internal_session_id: return (iid, sid) return None def get_sessions_for_instance( self, instance_id: str, ) -> list[TerminalSession]: """Return all in-memory sessions for a given instance.""" return [ session for (iid, _sid), session in self._sessions.items() if iid == instance_id ] async def close_session( self, instance_id: str, session_id: str, ) -> None: """Close a specific session and update its DB status.""" key = (instance_id, session_id) session = self._sessions.pop(key, None) if session: await session.close() # Fire-and-forget DB update asyncio.create_task(self._mark_closed_in_db(session_id)) logger.info( "Closed terminal session %s for instance %s", session_id, instance_id, ) async def attach_websocket( self, session: TerminalSession, websocket: WebSocket, ) -> None: """Attach a WebSocket to an existing session. Closes existing WebSocket connections only for this specific session. """ # Handle concurrent connections - close existing ones within the same session if session.has_websockets(): logger.debug( "Closing existing WebSocket connections for session %s (instance %s)", session.session_id, session.instance_id, ) for ws in list(session._websockets): try: await ws.close(code=4000, reason="New connection established") except Exception: pass # noqa: S110 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 # noqa: S110 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, startup_command: str | None = None, session_id: str | None = None, name: str | None = None, container_user: str | None = None, ) -> TerminalSession: """Reset a session by killing it and creating a new one. Args: instance_id: UUID of the tool instance. container_id: Docker container ID. startup_command: Optional startup command. session_id: Specific session to reset. If None, resets the default session. name: Optional name to preserve for the new session. container_user: Optional container user for docker exec. Returns: The newly created TerminalSession. """ instance_id_str = str(instance_id) target_session_id = session_id or "default" key = (instance_id_str, target_session_id) # Preserve old name if not provided old_name = name if old_name is None and key in self._sessions: old_name = self._sessions[key].name # Close existing session if any if key in self._sessions: logger.debug( "Resetting terminal session %s for instance %s", target_session_id, instance_id, ) old_session = self._sessions.pop(key) await old_session.close() # Fire-and-forget DB update for old session asyncio.create_task(self._mark_closed_in_db(old_session.session_id)) # Create new session preserving the same session_id slot new_session_id = str(uuid.uuid4()) new_session = TerminalSession( session_id=new_session_id, instance_id=instance_id, container_id=container_id, startup_command=startup_command, name=old_name or ("Session 1" if target_session_id == "default" else None), container_user=container_user, ) await new_session.start(startup_command=startup_command) self._sessions[key] = new_session # Fire-and-forget DB insert asyncio.create_task( self._insert_db_session_row(new_session_id, instance_id, new_session.name) ) return new_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()