Merge branch 'feat/tool-definition-manifest' into dev
Conflicts resolved: - models/__init__.py: kept both TerminalSessionModel (from dev) and ToolDefinitionManifest (from feature branch) - alembic migration: kept full migration (already applied to DB) - openspec/config.yaml: kept full config with SDD settings
This commit is contained in:
+470
-72
@@ -1,16 +1,20 @@
|
||||
"""WebSocket terminal endpoint for tool instances."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import uuid
|
||||
from contextlib import suppress
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, WebSocket, status
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.dependencies import get_db_session
|
||||
from src.auth.dependencies import get_current_user_id, get_db_session
|
||||
from src.models.terminal_session import TerminalSessionModel
|
||||
from src.models.tool_instance import ToolInstance
|
||||
from src.models.tool_type import ToolType
|
||||
from src.services.terminal_manager import terminal_manager
|
||||
from src.services.terminal_manager import MaxSessionsExceededError, terminal_manager
|
||||
|
||||
router = APIRouter()
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -19,32 +23,58 @@ logger = logging.getLogger(__name__)
|
||||
class SessionRef:
|
||||
"""Mutable reference to a terminal session, allowing updates during reset."""
|
||||
|
||||
def __init__(self, session):
|
||||
def __init__(self, session, slot_session_id: str | None = None):
|
||||
self.session = session
|
||||
self.slot_session_id = slot_session_id or session.session_id
|
||||
|
||||
|
||||
@router.websocket(
|
||||
"/ws/tool-instances/{instance_id}/terminal",
|
||||
)
|
||||
async def terminal_websocket(
|
||||
async def terminal_websocket_default(
|
||||
websocket: WebSocket,
|
||||
instance_id: str,
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> None:
|
||||
"""WebSocket endpoint for terminal access to a tool instance.
|
||||
"""WebSocket endpoint for terminal access (default session alias).
|
||||
|
||||
Provides an interactive terminal session inside a running tool instance container.
|
||||
Sessions persist across WebSocket disconnections.
|
||||
Backward-compatible route that maps to the default session.
|
||||
"""
|
||||
await _handle_terminal_websocket(websocket, instance_id, None, db_session)
|
||||
|
||||
|
||||
@router.websocket(
|
||||
"/ws/tool-instances/{instance_id}/terminal/{session_id}",
|
||||
)
|
||||
async def terminal_websocket_specific(
|
||||
websocket: WebSocket,
|
||||
instance_id: str,
|
||||
session_id: str,
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> None:
|
||||
"""WebSocket endpoint for a specific terminal session."""
|
||||
await _handle_terminal_websocket(websocket, instance_id, session_id, db_session)
|
||||
|
||||
|
||||
async def _handle_terminal_websocket(
|
||||
websocket: WebSocket,
|
||||
instance_id: str,
|
||||
target_session_id: str | None,
|
||||
db_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Shared WebSocket handler for terminal sessions.
|
||||
|
||||
Args:
|
||||
websocket: The WebSocket connection.
|
||||
instance_id: UUID string of the tool instance.
|
||||
target_session_id: Specific session ID (slot key). None means default session.
|
||||
db_session: Database session.
|
||||
|
||||
Returns:
|
||||
None. Communicates via WebSocket messages.
|
||||
"""
|
||||
logger.debug("Terminal WebSocket connection attempt for instance %s", instance_id)
|
||||
logger.debug(
|
||||
"Terminal WebSocket connection attempt for instance %s (session=%s)",
|
||||
instance_id,
|
||||
target_session_id or "default",
|
||||
)
|
||||
await websocket.accept()
|
||||
logger.debug("Terminal WebSocket accepted for instance %s", instance_id)
|
||||
|
||||
@@ -59,7 +89,9 @@ async def terminal_websocket(
|
||||
# Authenticate user from session cookie
|
||||
user_id = await _get_user_from_websocket(websocket, db_session)
|
||||
if user_id is None:
|
||||
logger.warning("Unauthorized terminal access attempt for instance %s", instance_id)
|
||||
logger.warning(
|
||||
"Unauthorized terminal access attempt for instance %s", instance_id
|
||||
)
|
||||
await websocket.close(code=4003, reason="Unauthorized")
|
||||
return
|
||||
|
||||
@@ -71,12 +103,21 @@ async def terminal_websocket(
|
||||
return
|
||||
|
||||
if instance.owner_id != user_id:
|
||||
logger.warning("Forbidden terminal access for instance %s by user %s", instance_id, user_id)
|
||||
logger.warning(
|
||||
"Forbidden terminal access for instance %s by user %s",
|
||||
instance_id,
|
||||
user_id,
|
||||
)
|
||||
await websocket.close(code=4003, reason="Forbidden")
|
||||
return
|
||||
|
||||
if instance.status != "running" or not instance.container_id:
|
||||
logger.warning("Instance %s not running (status=%s, container_id=%s)", instance_id, instance.status, instance.container_id)
|
||||
logger.warning(
|
||||
"Instance %s not running (status=%s, container_id=%s)",
|
||||
instance_id,
|
||||
instance.status,
|
||||
instance.container_id,
|
||||
)
|
||||
await websocket.close(code=4004, reason="Instance not running")
|
||||
return
|
||||
|
||||
@@ -86,16 +127,50 @@ async def terminal_websocket(
|
||||
tool_type = await db_session.get(ToolType, instance.tool_type_id)
|
||||
startup_command = tool_type.startup_command if tool_type else None
|
||||
if startup_command:
|
||||
logger.debug("Using startup command for instance %s: %s", instance_id, startup_command)
|
||||
logger.debug(
|
||||
"Using startup command for instance %s: %s",
|
||||
instance_id,
|
||||
startup_command,
|
||||
)
|
||||
|
||||
session = None
|
||||
|
||||
# Get or create terminal session
|
||||
try:
|
||||
session = await terminal_manager.get_or_create_session(
|
||||
instance_uuid,
|
||||
instance.container_id,
|
||||
startup_command=startup_command,
|
||||
if target_session_id is None:
|
||||
# Default session alias
|
||||
session = await terminal_manager.get_or_create_session(
|
||||
instance_uuid,
|
||||
instance.container_id,
|
||||
startup_command=startup_command,
|
||||
)
|
||||
slot_session_id = "default"
|
||||
else:
|
||||
# Specific session
|
||||
session = terminal_manager.get_session(
|
||||
instance_id,
|
||||
target_session_id,
|
||||
)
|
||||
if session is None:
|
||||
logger.warning(
|
||||
"Session %s not found for instance %s",
|
||||
target_session_id,
|
||||
instance_id,
|
||||
)
|
||||
await websocket.close(code=4004, reason="Session not found")
|
||||
return
|
||||
# Determine slot key for reset scoping
|
||||
key = terminal_manager._find_key_by_internal_id(
|
||||
instance_id, session.session_id
|
||||
)
|
||||
slot_session_id = key[1] if key else target_session_id
|
||||
|
||||
logger.debug(
|
||||
"Terminal session ready for instance %s (session_id=%s, slot=%s)",
|
||||
instance_id,
|
||||
session.session_id,
|
||||
slot_session_id,
|
||||
)
|
||||
logger.debug("Terminal session ready for instance %s (session_id=%s)", instance_id, session.session_id)
|
||||
|
||||
# Attach WebSocket to session
|
||||
await terminal_manager.attach_websocket(session, websocket)
|
||||
@@ -106,11 +181,13 @@ async def terminal_websocket(
|
||||
logger.debug("Sent connected status for instance %s", instance_id)
|
||||
|
||||
# Use mutable session reference so loops can survive reset
|
||||
session_ref = SessionRef(session)
|
||||
session_ref = SessionRef(session, slot_session_id)
|
||||
|
||||
# Start I/O loops and heartbeat
|
||||
read_task = asyncio.create_task(_read_loop(session_ref, websocket))
|
||||
write_task = asyncio.create_task(_write_loop(session_ref, websocket, instance_id))
|
||||
write_task = asyncio.create_task(
|
||||
_write_loop(session_ref, websocket, instance_id)
|
||||
)
|
||||
heartbeat_task = asyncio.create_task(_heartbeat_loop(websocket))
|
||||
logger.debug("Started terminal loops for instance %s", instance_id)
|
||||
|
||||
@@ -119,24 +196,33 @@ async def terminal_websocket(
|
||||
[read_task, write_task, heartbeat_task],
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
|
||||
logger.debug("Terminal loop completed for instance %s, done=%s", instance_id, len(done))
|
||||
|
||||
|
||||
logger.debug(
|
||||
"Terminal loop completed for instance %s, done=%s",
|
||||
instance_id,
|
||||
len(done),
|
||||
)
|
||||
|
||||
# Cancel remaining tasks
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
|
||||
except Exception as exc:
|
||||
logger.error("Terminal session error for instance %s: %s", instance_id, str(exc), exc_info=True)
|
||||
logger.error(
|
||||
"Terminal session error for instance %s: %s",
|
||||
instance_id,
|
||||
str(exc),
|
||||
exc_info=True,
|
||||
)
|
||||
await websocket.close(code=4000, reason=f"Error: {exc}")
|
||||
finally:
|
||||
# Detach WebSocket, don't kill session
|
||||
try:
|
||||
if 'session' in locals():
|
||||
with suppress(Exception):
|
||||
if session is not None:
|
||||
await terminal_manager.detach_websocket(session, websocket)
|
||||
logger.debug("WebSocket detached from session for instance %s", instance_id)
|
||||
except Exception:
|
||||
pass
|
||||
logger.debug(
|
||||
"WebSocket detached from session for instance %s", instance_id
|
||||
)
|
||||
|
||||
|
||||
async def _read_loop(session_ref: SessionRef, websocket) -> None:
|
||||
@@ -175,38 +261,54 @@ async def _write_loop(session_ref: SessionRef, websocket, instance_id: str) -> N
|
||||
text = message["text"]
|
||||
if text.startswith("{"):
|
||||
# Control message (JSON)
|
||||
import json
|
||||
try:
|
||||
ctrl = json.loads(text)
|
||||
msg_type = ctrl.get("type")
|
||||
|
||||
|
||||
if msg_type == "resize":
|
||||
cols = ctrl.get("cols", 80)
|
||||
rows = ctrl.get("rows", 24)
|
||||
logger.debug(f"Received resize message for instance {instance_id}: {cols}x{rows}")
|
||||
logger.debug(
|
||||
"Received resize message for instance %s: %sx%s",
|
||||
instance_id,
|
||||
cols,
|
||||
rows,
|
||||
)
|
||||
await session.resize(cols, rows)
|
||||
elif msg_type == "reset":
|
||||
# Reset terminal session
|
||||
logger.debug("Resetting terminal session for instance %s", session.instance_id)
|
||||
await websocket.send_json({"type": "status", "status": "resetting"})
|
||||
|
||||
# Reset the session
|
||||
# Reset terminal session (scoped to current slot)
|
||||
logger.debug(
|
||||
"Resetting terminal session for instance %s (slot=%s)",
|
||||
session.instance_id,
|
||||
session_ref.slot_session_id,
|
||||
)
|
||||
await websocket.send_json(
|
||||
{"type": "status", "status": "resetting"}
|
||||
)
|
||||
|
||||
# Reset the session scoped to its slot
|
||||
new_session = await terminal_manager.reset_session(
|
||||
session.instance_id,
|
||||
session.container_id,
|
||||
startup_command=session.startup_command,
|
||||
session_id=session_ref.slot_session_id,
|
||||
name=session.name,
|
||||
)
|
||||
|
||||
# Update the mutable session reference so read_loop uses the new session
|
||||
|
||||
# Update the mutable session reference
|
||||
session_ref.session = new_session
|
||||
|
||||
|
||||
# Attach to new session
|
||||
await terminal_manager.attach_websocket(new_session, websocket)
|
||||
await websocket.send_json({"type": "status", "status": "connected"})
|
||||
|
||||
await terminal_manager.attach_websocket(
|
||||
new_session, websocket
|
||||
)
|
||||
await websocket.send_json(
|
||||
{"type": "status", "status": "connected"}
|
||||
)
|
||||
|
||||
# Continue the loop with the new session
|
||||
continue
|
||||
|
||||
|
||||
except json.JSONDecodeError:
|
||||
# Not a valid JSON control message, treat as regular input
|
||||
await session.write_input(text.encode("utf-8"))
|
||||
@@ -232,56 +334,347 @@ async def _heartbeat_loop(websocket: WebSocket) -> None:
|
||||
pass
|
||||
|
||||
|
||||
@router.post(
|
||||
"/projects/{project_id}/repositories/{repo_id}/instances/{instance_id}/terminal/reset",
|
||||
summary="Reset terminal session",
|
||||
description="Reset the terminal session for a tool instance, killing the current shell and starting fresh.",
|
||||
)
|
||||
async def reset_terminal_session(
|
||||
project_id: uuid.UUID,
|
||||
repo_id: uuid.UUID,
|
||||
async def _get_terminal_instance(
|
||||
instance_id: uuid.UUID,
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Reset the terminal session for an instance.
|
||||
user_id: uuid.UUID,
|
||||
db_session: AsyncSession,
|
||||
) -> ToolInstance:
|
||||
"""Fetch instance and validate auth, ownership, and running status.
|
||||
|
||||
Args:
|
||||
project_id: UUID of the project.
|
||||
repo_id: UUID of the repository.
|
||||
instance_id: UUID of the tool instance.
|
||||
user_id: ID of the authenticated user.
|
||||
db_session: Database session.
|
||||
|
||||
Returns:
|
||||
Dictionary with status message.
|
||||
The validated ToolInstance.
|
||||
|
||||
Raises:
|
||||
HTTPException: If instance not found, not owned, or not running.
|
||||
"""
|
||||
# Get instance and verify it exists and is running
|
||||
instance = await db_session.get(ToolInstance, instance_id)
|
||||
if instance is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="Instance not found"
|
||||
status_code=status.HTTP_404_NOT_FOUND, detail="Instance not found"
|
||||
)
|
||||
|
||||
if instance.owner_id != user_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Not authorized to access this instance",
|
||||
)
|
||||
|
||||
if instance.status != "running" or not instance.container_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Instance is not running"
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail="Instance is not running"
|
||||
)
|
||||
|
||||
return instance
|
||||
|
||||
|
||||
@router.get(
|
||||
"/instances/{instance_id}/terminal/sessions",
|
||||
summary="List terminal sessions",
|
||||
description="List terminal sessions for a tool instance with live WebSocket state.",
|
||||
)
|
||||
async def list_terminal_sessions(
|
||||
instance_id: uuid.UUID,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""List terminal sessions for an instance.
|
||||
|
||||
Args:
|
||||
instance_id: UUID of the tool instance.
|
||||
user_id: ID of the authenticated user.
|
||||
db_session: Database session.
|
||||
|
||||
Returns:
|
||||
Dictionary with sessions list.
|
||||
"""
|
||||
await _get_terminal_instance(instance_id, user_id, db_session)
|
||||
|
||||
# Query active DB rows for this instance
|
||||
result = await db_session.execute(
|
||||
select(TerminalSessionModel)
|
||||
.where(TerminalSessionModel.instance_id == instance_id)
|
||||
.where(TerminalSessionModel.status != "closed")
|
||||
.order_by(TerminalSessionModel.created_at.asc())
|
||||
)
|
||||
db_rows = result.scalars().all()
|
||||
|
||||
# Build response with live has_websockets flag
|
||||
sessions = []
|
||||
for row in db_rows:
|
||||
live_session = terminal_manager.get_session(str(instance_id), str(row.id))
|
||||
sessions.append(
|
||||
{
|
||||
"id": str(row.id),
|
||||
"name": row.name,
|
||||
"status": row.status,
|
||||
"has_websockets": live_session.has_websockets()
|
||||
if live_session
|
||||
else False,
|
||||
"created_at": row.created_at.isoformat() if row.created_at else None,
|
||||
"last_activity_at": row.last_activity_at.isoformat()
|
||||
if row.last_activity_at
|
||||
else None,
|
||||
}
|
||||
)
|
||||
|
||||
return {"sessions": sessions}
|
||||
|
||||
|
||||
@router.post(
|
||||
"/instances/{instance_id}/terminal/sessions",
|
||||
summary="Create terminal session",
|
||||
description="Create a new terminal session for a running tool instance.",
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
)
|
||||
async def create_terminal_session(
|
||||
instance_id: uuid.UUID,
|
||||
data: dict,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Create a new terminal session.
|
||||
|
||||
Args:
|
||||
instance_id: UUID of the tool instance.
|
||||
data: Request body with optional name.
|
||||
user_id: ID of the authenticated user.
|
||||
db_session: Database session.
|
||||
|
||||
Returns:
|
||||
Dictionary with new session details.
|
||||
|
||||
Raises:
|
||||
HTTPException: 409 if max sessions reached.
|
||||
"""
|
||||
instance = await _get_terminal_instance(instance_id, user_id, db_session)
|
||||
assert instance.container_id is not None
|
||||
|
||||
# Fetch tool type to get startup_command
|
||||
tool_type = await db_session.get(ToolType, instance.tool_type_id)
|
||||
startup_command = tool_type.startup_command if tool_type else None
|
||||
|
||||
name = data.get("name")
|
||||
|
||||
try:
|
||||
session = await terminal_manager.create_session(
|
||||
instance_id,
|
||||
instance.container_id,
|
||||
startup_command=startup_command,
|
||||
name=name,
|
||||
)
|
||||
except MaxSessionsExceededError:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail="Maximum of 5 terminal sessions reached for this instance",
|
||||
) from None
|
||||
|
||||
return {
|
||||
"id": session.session_id,
|
||||
"name": session.name,
|
||||
"status": session.status,
|
||||
"created_at": session.last_activity,
|
||||
}
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/instances/{instance_id}/terminal/sessions/{session_id}",
|
||||
summary="Close terminal session",
|
||||
description="Close a specific terminal session.",
|
||||
)
|
||||
async def close_terminal_session(
|
||||
instance_id: uuid.UUID,
|
||||
session_id: str,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Close a terminal session.
|
||||
|
||||
Args:
|
||||
instance_id: UUID of the tool instance.
|
||||
session_id: ID of the session to close.
|
||||
user_id: ID of the authenticated user.
|
||||
db_session: Database session.
|
||||
|
||||
Returns:
|
||||
Dictionary with closure status.
|
||||
"""
|
||||
await _get_terminal_instance(instance_id, user_id, db_session)
|
||||
|
||||
# Find the session by internal ID to determine its slot key
|
||||
key = terminal_manager._find_key_by_internal_id(str(instance_id), session_id)
|
||||
if (
|
||||
key is None
|
||||
and terminal_manager.get_session(str(instance_id), session_id) is not None
|
||||
):
|
||||
key = (str(instance_id), session_id)
|
||||
|
||||
if key is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND, detail="Session not found"
|
||||
)
|
||||
|
||||
await terminal_manager.close_session(key[0], key[1])
|
||||
|
||||
return {"status": "closed", "session_id": session_id}
|
||||
|
||||
|
||||
@router.post(
|
||||
"/instances/{instance_id}/terminal/sessions/{session_id}/reset",
|
||||
summary="Reset terminal session",
|
||||
description="Reset a specific terminal session, killing the current shell and starting fresh.",
|
||||
)
|
||||
async def reset_specific_terminal_session(
|
||||
instance_id: uuid.UUID,
|
||||
session_id: str,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Reset a specific terminal session.
|
||||
|
||||
Args:
|
||||
instance_id: UUID of the tool instance.
|
||||
session_id: ID of the session to reset.
|
||||
user_id: ID of the authenticated user.
|
||||
db_session: Database session.
|
||||
|
||||
Returns:
|
||||
Dictionary with reset session details.
|
||||
"""
|
||||
instance = await _get_terminal_instance(instance_id, user_id, db_session)
|
||||
assert instance.container_id is not None
|
||||
|
||||
# Determine slot key for reset
|
||||
key = terminal_manager._find_key_by_internal_id(str(instance_id), session_id)
|
||||
if (
|
||||
key is None
|
||||
and terminal_manager.get_session(str(instance_id), session_id) is not None
|
||||
):
|
||||
key = (str(instance_id), session_id)
|
||||
|
||||
if key is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND, detail="Session not found"
|
||||
)
|
||||
|
||||
# Fetch tool type to get startup_command
|
||||
tool_type = await db_session.get(ToolType, instance.tool_type_id)
|
||||
startup_command = tool_type.startup_command if tool_type else None
|
||||
|
||||
# Preserve name if possible
|
||||
live_session = terminal_manager.get_session(str(instance_id), session_id)
|
||||
name = live_session.name if live_session else None
|
||||
|
||||
new_session = await terminal_manager.reset_session(
|
||||
instance_id,
|
||||
instance.container_id,
|
||||
startup_command=startup_command,
|
||||
session_id=key[1],
|
||||
name=name,
|
||||
)
|
||||
|
||||
return {
|
||||
"id": new_session.session_id,
|
||||
"name": new_session.name,
|
||||
"status": new_session.status,
|
||||
}
|
||||
|
||||
|
||||
@router.post(
|
||||
"/instances/{instance_id}/terminal/sessions/{session_id}/rename",
|
||||
summary="Rename terminal session",
|
||||
description="Rename a specific terminal session.",
|
||||
)
|
||||
async def rename_terminal_session(
|
||||
instance_id: uuid.UUID,
|
||||
session_id: str,
|
||||
data: dict,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Rename a terminal session.
|
||||
|
||||
Args:
|
||||
instance_id: UUID of the tool instance.
|
||||
session_id: ID of the session to rename.
|
||||
data: Request body with new name.
|
||||
user_id: ID of the authenticated user.
|
||||
db_session: Database session.
|
||||
|
||||
Returns:
|
||||
Dictionary with updated session details.
|
||||
"""
|
||||
await _get_terminal_instance(instance_id, user_id, db_session)
|
||||
|
||||
new_name = data.get("name")
|
||||
if not new_name or not isinstance(new_name, str):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail="Name is required"
|
||||
)
|
||||
|
||||
# Update in-memory session name if live
|
||||
live_session = terminal_manager.get_session(str(instance_id), session_id)
|
||||
if live_session:
|
||||
live_session.name = new_name
|
||||
|
||||
# Update DB row
|
||||
db_row = await db_session.get(TerminalSessionModel, uuid.UUID(session_id))
|
||||
if db_row is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND, detail="Session not found"
|
||||
)
|
||||
|
||||
db_row.name = new_name
|
||||
await db_session.commit()
|
||||
|
||||
return {"id": str(db_row.id), "name": new_name}
|
||||
|
||||
|
||||
@router.post(
|
||||
"/instances/{instance_id}/terminal/reset",
|
||||
summary="Reset terminal session (legacy alias)",
|
||||
description="Reset the default terminal session for a tool instance. Preserved for backward compatibility.",
|
||||
)
|
||||
async def reset_terminal_session(
|
||||
instance_id: uuid.UUID,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Reset the default terminal session for an instance (legacy alias).
|
||||
|
||||
Args:
|
||||
instance_id: UUID of the tool instance.
|
||||
user_id: ID of the authenticated user.
|
||||
db_session: Database session.
|
||||
|
||||
Returns:
|
||||
Dictionary with status message.
|
||||
"""
|
||||
instance = await _get_terminal_instance(instance_id, user_id, db_session)
|
||||
assert instance.container_id is not None
|
||||
|
||||
# Fetch tool type to get startup_command
|
||||
tool_type = await db_session.get(ToolType, instance.tool_type_id)
|
||||
startup_command = tool_type.startup_command if tool_type else None
|
||||
|
||||
try:
|
||||
# Reset the session
|
||||
# Reset the default session
|
||||
new_session = await terminal_manager.reset_session(
|
||||
instance_id,
|
||||
instance.container_id,
|
||||
startup_command=startup_command,
|
||||
)
|
||||
|
||||
logger.info("Terminal session reset for instance %s (new session_id=%s)", instance_id, new_session.session_id)
|
||||
|
||||
|
||||
logger.info(
|
||||
"Terminal session reset for instance %s (new session_id=%s)",
|
||||
instance_id,
|
||||
new_session.session_id,
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "success",
|
||||
"message": "Terminal session reset successfully",
|
||||
@@ -289,11 +682,16 @@ async def reset_terminal_session(
|
||||
"session_id": new_session.session_id,
|
||||
}
|
||||
except Exception as exc:
|
||||
logger.error("Failed to reset terminal session for instance %s: %s", instance_id, str(exc), exc_info=True)
|
||||
logger.error(
|
||||
"Failed to reset terminal session for instance %s: %s",
|
||||
instance_id,
|
||||
str(exc),
|
||||
exc_info=True,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to reset terminal session: {exc}"
|
||||
)
|
||||
detail=f"Failed to reset terminal session: {exc}",
|
||||
) from exc
|
||||
|
||||
|
||||
async def _get_user_from_websocket(
|
||||
|
||||
@@ -25,6 +25,7 @@ from src.api.tool_types import router as tool_types_router
|
||||
from src.api.user_config import router as user_config_router
|
||||
from src.api.users import router as users_router
|
||||
from src.config import Settings
|
||||
from src.models.terminal_session import TerminalSessionModel # noqa: F401 – Alembic model discovery
|
||||
from src.database import init_database
|
||||
from src.logging_config import (
|
||||
ExceptionLoggingMiddleware,
|
||||
|
||||
@@ -4,6 +4,7 @@ from src.models.config_profile import ConfigProfile, ConfigProfileInclude
|
||||
from src.models.git_repository import GitRepository
|
||||
from src.models.project import Project
|
||||
from src.models.ssh_key import SSHKey
|
||||
from src.models.terminal_session import TerminalSessionModel
|
||||
from src.models.tool_definition_manifest import ToolDefinitionManifest
|
||||
from src.models.tool_instance import ToolInstance
|
||||
from src.models.tool_type import ToolType
|
||||
@@ -18,6 +19,7 @@ __all__ = [
|
||||
"GitRepository",
|
||||
"Project",
|
||||
"SSHKey",
|
||||
"TerminalSessionModel",
|
||||
"ToolDefinitionManifest",
|
||||
"ToolInstance",
|
||||
"ToolType",
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
"""Terminal session database model."""
|
||||
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, ForeignKey, String
|
||||
from sqlalchemy import Uuid as UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
|
||||
class TerminalSessionModel(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
"""Database model for terminal session metadata."""
|
||||
|
||||
__tablename__ = "terminal_sessions"
|
||||
|
||||
instance_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(),
|
||||
ForeignKey("tool_instances.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
index=True,
|
||||
)
|
||||
name: Mapped[str | None] = mapped_column(String(255), nullable=True)
|
||||
status: Mapped[str] = mapped_column(
|
||||
String(50),
|
||||
nullable=False,
|
||||
default="active",
|
||||
)
|
||||
last_activity_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
nullable=True,
|
||||
)
|
||||
closed_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
nullable=True,
|
||||
)
|
||||
@@ -3,20 +3,37 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from fastapi import WebSocket
|
||||
|
||||
from src.database import SessionLocal
|
||||
from src.models.terminal_session import TerminalSessionModel
|
||||
from src.services.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 for persistence
|
||||
self._sessions: dict[str, TerminalSession] = {}
|
||||
# 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()
|
||||
|
||||
@@ -42,16 +59,131 @@ class TerminalManager:
|
||||
|
||||
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()):
|
||||
idle_keys = []
|
||||
for (instance_id, session_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)
|
||||
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,
|
||||
)
|
||||
session = self._sessions.pop(key, None)
|
||||
if session:
|
||||
await 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."""
|
||||
try:
|
||||
async with SessionLocal() as db_session:
|
||||
db_row = TerminalSessionModel(
|
||||
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),
|
||||
)
|
||||
db_session.add(db_row)
|
||||
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,
|
||||
) -> 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).
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
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,
|
||||
)
|
||||
await session.start(startup_command=startup_command)
|
||||
|
||||
key = (instance_id_str, session_id)
|
||||
self._sessions[key] = session
|
||||
|
||||
# Fire-and-forget DB insert
|
||||
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,
|
||||
@@ -59,61 +191,146 @@ class TerminalManager:
|
||||
container_id: str,
|
||||
startup_command: str | None = None,
|
||||
) -> TerminalSession:
|
||||
"""Get existing session or create a new one."""
|
||||
"""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)
|
||||
|
||||
# Check for existing session
|
||||
if instance_id_str in self._sessions:
|
||||
session = self._sessions[instance_id_str]
|
||||
|
||||
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)
|
||||
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)
|
||||
logger.debug(
|
||||
"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)
|
||||
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, instance_id, container_id, startup_command=startup_command)
|
||||
session = TerminalSession(
|
||||
session_id=session_id,
|
||||
instance_id=instance_id,
|
||||
container_id=container_id,
|
||||
startup_command=startup_command,
|
||||
name="Session 1",
|
||||
)
|
||||
await session.start(startup_command=startup_command)
|
||||
self._sessions[instance_id_str] = session
|
||||
|
||||
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."""
|
||||
# Handle concurrent connections - close existing ones
|
||||
"""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 instance %s", session.instance_id)
|
||||
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
|
||||
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
|
||||
pass # noqa: S110
|
||||
|
||||
async def detach_websocket(
|
||||
self,
|
||||
@@ -128,23 +345,60 @@ class TerminalManager:
|
||||
instance_id: uuid.UUID,
|
||||
container_id: str,
|
||||
startup_command: str | None = None,
|
||||
session_id: str | None = None,
|
||||
name: str | None = None,
|
||||
) -> TerminalSession:
|
||||
"""Reset a session by killing it and creating a new one."""
|
||||
"""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.
|
||||
|
||||
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 instance_id_str in self._sessions:
|
||||
logger.debug("Resetting terminal session for instance %s", instance_id)
|
||||
old_session = self._sessions.pop(instance_id_str)
|
||||
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()
|
||||
|
||||
# Create new session
|
||||
session_id = str(uuid.uuid4())
|
||||
session = TerminalSession(session_id, instance_id, container_id, startup_command=startup_command)
|
||||
await session.start(startup_command=startup_command)
|
||||
self._sessions[instance_id_str] = session
|
||||
|
||||
return session
|
||||
# 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),
|
||||
)
|
||||
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."""
|
||||
@@ -152,7 +406,7 @@ class TerminalManager:
|
||||
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()
|
||||
|
||||
|
||||
@@ -18,18 +18,28 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
class TerminalSession:
|
||||
"""Manages a single terminal session connected to a docker container.
|
||||
|
||||
|
||||
Supports persistent sessions that survive WebSocket disconnections.
|
||||
Multiple WebSocket connections can attach/detach from the same session.
|
||||
"""
|
||||
|
||||
# Circular buffer size (10KB)
|
||||
BUFFER_SIZE = 10 * 1024
|
||||
|
||||
|
||||
# Idle timeout in seconds (30 minutes)
|
||||
IDLE_TIMEOUT = 30 * 60
|
||||
|
||||
def __init__(self, session_id: str, instance_id: uuid.UUID, container_id: str, startup_command: str | None = None) -> None:
|
||||
# Session number counter per instance_id for auto-naming
|
||||
_instance_counters: dict[str, int] = {}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
session_id: str,
|
||||
instance_id: uuid.UUID,
|
||||
container_id: str,
|
||||
startup_command: str | None = None,
|
||||
name: str | None = None,
|
||||
) -> None:
|
||||
self.session_id = session_id
|
||||
self.instance_id = instance_id
|
||||
self.container_id = container_id
|
||||
@@ -38,37 +48,52 @@ class TerminalSession:
|
||||
self._closed = False
|
||||
self._master_fd: int | None = None
|
||||
self._slave_fd: int | None = None
|
||||
|
||||
|
||||
# Circular buffer for output replay
|
||||
self._output_buffer: deque[bytes] = deque(maxlen=self.BUFFER_SIZE)
|
||||
self._buffer_size = 0
|
||||
|
||||
|
||||
# WebSocket connections
|
||||
self._websockets: set[Any] = set()
|
||||
|
||||
|
||||
# Activity tracking
|
||||
self.last_activity = time.time()
|
||||
|
||||
|
||||
# Terminal size
|
||||
self._cols = 80
|
||||
self._rows = 24
|
||||
|
||||
# Session metadata
|
||||
self.name = name or self._generate_name(str(instance_id))
|
||||
self.status: str = "active"
|
||||
|
||||
@classmethod
|
||||
def _generate_name(cls, instance_id: str) -> str:
|
||||
"""Generate an auto-incremented session name for the instance."""
|
||||
count = cls._instance_counters.get(instance_id, 0) + 1
|
||||
cls._instance_counters[instance_id] = count
|
||||
return f"Session {count}"
|
||||
|
||||
async def start(self, startup_command: str | None = None) -> None:
|
||||
"""Start the docker exec process with a shell using a PTY."""
|
||||
# Create a pseudo-terminal on the host
|
||||
self._master_fd, self._slave_fd = pty.openpty()
|
||||
|
||||
|
||||
# Set the terminal size initially
|
||||
self._set_terminal_size(self._cols, self._rows)
|
||||
logger.debug(f"Starting terminal session {self.session_id} for container {self.container_id} with initial size {self._cols}x{self._rows}")
|
||||
|
||||
logger.debug(
|
||||
f"Starting terminal session {self.session_id} for container {self.container_id} with initial size {self._cols}x{self._rows}"
|
||||
)
|
||||
|
||||
# Build the shell command
|
||||
if startup_command:
|
||||
shell_cmd = f'bash -c "{startup_command}" || true; exec bash -il'
|
||||
logger.debug(f"Using startup command for session {self.session_id}: {startup_command}")
|
||||
logger.debug(
|
||||
f"Using startup command for session {self.session_id}: {startup_command}"
|
||||
)
|
||||
else:
|
||||
shell_cmd = "bash -il"
|
||||
|
||||
|
||||
# Start docker exec with the slave fd as stdin/stdout/stderr
|
||||
# Using -it because the slave fd IS a TTY
|
||||
self.process = await asyncio.create_subprocess_exec(
|
||||
@@ -85,11 +110,11 @@ class TerminalSession:
|
||||
stdout=self._slave_fd,
|
||||
stderr=self._slave_fd,
|
||||
)
|
||||
|
||||
|
||||
# Close slave fd in parent process
|
||||
os.close(self._slave_fd)
|
||||
self._slave_fd = None
|
||||
|
||||
|
||||
self.last_activity = time.time()
|
||||
|
||||
def _set_terminal_size(self, cols: int, rows: int) -> None:
|
||||
@@ -99,7 +124,7 @@ class TerminalSession:
|
||||
return
|
||||
# TIOCSWINSZ = 0x5414 on Linux
|
||||
TIOCSWINSZ = 0x5414
|
||||
size = struct.pack('HHHH', rows, cols, 0, 0)
|
||||
size = struct.pack("HHHH", rows, cols, 0, 0)
|
||||
try:
|
||||
fcntl.ioctl(self._master_fd, TIOCSWINSZ, size)
|
||||
logger.debug(f"Resized PTY to {cols}x{rows} (fd={self._master_fd})")
|
||||
@@ -127,7 +152,7 @@ class TerminalSession:
|
||||
"""Add data to circular buffer, maintaining size limit."""
|
||||
self._output_buffer.append(data)
|
||||
self._buffer_size += len(data)
|
||||
|
||||
|
||||
# Trim if exceeds max size
|
||||
while self._buffer_size > self.BUFFER_SIZE and self._output_buffer:
|
||||
removed = self._output_buffer.popleft()
|
||||
@@ -152,16 +177,16 @@ class TerminalSession:
|
||||
if self._closed:
|
||||
logger.warning("Cannot resize: session is closed")
|
||||
return
|
||||
|
||||
|
||||
# Only resize if dimensions actually changed
|
||||
if cols == self._cols and rows == self._rows:
|
||||
return
|
||||
|
||||
|
||||
self._cols = cols
|
||||
self._rows = rows
|
||||
logger.debug(f"resize() called for session {self.session_id}: {cols}x{rows}")
|
||||
self._set_terminal_size(cols, rows)
|
||||
|
||||
|
||||
# Docker exec -it creates its own PTY inside the container,
|
||||
# so host PTY resize doesn't propagate to the container shell.
|
||||
# Send SIGWINCH to the docker exec process on the host.
|
||||
@@ -170,14 +195,19 @@ class TerminalSession:
|
||||
if self.process and self.process.pid:
|
||||
try:
|
||||
os.kill(self.process.pid, signal.SIGWINCH)
|
||||
logger.debug(f"Sent SIGWINCH to docker exec process {self.process.pid} for session {self.session_id}")
|
||||
logger.debug(
|
||||
f"Sent SIGWINCH to docker exec process {self.process.pid} for session {self.session_id}"
|
||||
)
|
||||
except ProcessLookupError:
|
||||
logger.warning(f"docker exec process {self.process.pid} not found for session {self.session_id}")
|
||||
logger.warning(
|
||||
f"docker exec process {self.process.pid} not found for session {self.session_id}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to send SIGWINCH: {e}")
|
||||
|
||||
async def reset(self) -> None:
|
||||
"""Reset the session by killing the process and clearing state."""
|
||||
self.status = "resetting"
|
||||
await self.close()
|
||||
self._closed = False
|
||||
self._output_buffer.clear()
|
||||
@@ -186,18 +216,20 @@ class TerminalSession:
|
||||
self.process = None
|
||||
self._master_fd = None
|
||||
self._slave_fd = None
|
||||
self.status = "active"
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close the session and cleanup."""
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
self.status = "closed"
|
||||
|
||||
if self._master_fd is not None:
|
||||
try:
|
||||
os.close(self._master_fd)
|
||||
except OSError:
|
||||
pass
|
||||
pass # noqa: S110
|
||||
self._master_fd = None
|
||||
|
||||
if self.process is not None:
|
||||
@@ -240,7 +272,7 @@ class TerminalSession:
|
||||
await ws.send_bytes(data)
|
||||
except Exception:
|
||||
dead_sockets.add(ws)
|
||||
|
||||
|
||||
# Clean up dead sockets
|
||||
for ws in dead_sockets:
|
||||
self._websockets.discard(ws)
|
||||
|
||||
Reference in New Issue
Block a user