130 lines
5.0 KiB
Python
130 lines
5.0 KiB
Python
|
|
"""WebSocket connection manager for AI UI Control.
|
||
|
|
|
||
|
|
Manages WebSocket connections from frontend clients. When an AI agent sends
|
||
|
|
a UI command via REST API, the command is forwarded to the frontend via
|
||
|
|
these WebSocket connections. The frontend executes the command and sends
|
||
|
|
feedback back through the same WebSocket.
|
||
|
|
"""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import json
|
||
|
|
import logging
|
||
|
|
import time
|
||
|
|
import uuid
|
||
|
|
from typing import Any
|
||
|
|
|
||
|
|
from fastapi import WebSocket
|
||
|
|
|
||
|
|
logger = logging.getLogger(__name__)
|
||
|
|
|
||
|
|
|
||
|
|
class AIUIControlWSManager:
|
||
|
|
"""Manages WebSocket connections for AI UI control.
|
||
|
|
|
||
|
|
Connections are per-user: each authenticated user can have one or more
|
||
|
|
frontend tabs connected. Commands are delivered to all tabs of the target
|
||
|
|
user. Feedback from any tab is accepted and stored.
|
||
|
|
"""
|
||
|
|
|
||
|
|
def __init__(self) -> None:
|
||
|
|
# user_id → list of WebSocket connections
|
||
|
|
self._connections: dict[str, list[WebSocket]] = {}
|
||
|
|
# command_id → feedback dict (stored when frontend responds)
|
||
|
|
self._feedback: dict[str, dict[str, Any]] = {}
|
||
|
|
# command_id → timestamp when delivered (for timeout tracking)
|
||
|
|
self._delivered_at: dict[str, float] = {}
|
||
|
|
# command_id → user_id (to route feedback)
|
||
|
|
self._command_user: dict[str, str] = {}
|
||
|
|
|
||
|
|
async def connect(self, websocket: WebSocket, user_id: str) -> None:
|
||
|
|
"""Accept and register a new WebSocket connection."""
|
||
|
|
await websocket.accept()
|
||
|
|
if user_id not in self._connections:
|
||
|
|
self._connections[user_id] = []
|
||
|
|
self._connections[user_id].append(websocket)
|
||
|
|
logger.debug(f"AI UI Control WS connected: user={user_id}, total={len(self._connections[user_id])}")
|
||
|
|
|
||
|
|
async def disconnect(self, websocket: WebSocket, user_id: str) -> None:
|
||
|
|
"""Remove a WebSocket connection."""
|
||
|
|
conns = self._connections.get(user_id, [])
|
||
|
|
if websocket in conns:
|
||
|
|
conns.remove(websocket)
|
||
|
|
if not conns:
|
||
|
|
self._connections.pop(user_id, None)
|
||
|
|
logger.debug(f"AI UI Control WS disconnected: user={user_id}, remaining={len(conns)}")
|
||
|
|
|
||
|
|
async def send_command(self, user_id: str, command: dict[str, Any]) -> str | None:
|
||
|
|
"""Send a UI command to all frontend connections of a user.
|
||
|
|
|
||
|
|
Returns the command_id if delivered, None if user has no active connections.
|
||
|
|
"""
|
||
|
|
command_id = command.get("command_id") or str(uuid.uuid4())
|
||
|
|
command["command_id"] = command_id
|
||
|
|
command.setdefault("timestamp", time.time())
|
||
|
|
|
||
|
|
conns = self._connections.get(user_id, [])
|
||
|
|
if not conns:
|
||
|
|
logger.debug(f"AI UI Control: no active connections for user {user_id}")
|
||
|
|
return None
|
||
|
|
|
||
|
|
self._command_user[command_id] = user_id
|
||
|
|
self._delivered_at[command_id] = time.time()
|
||
|
|
|
||
|
|
text = json.dumps(command, default=str)
|
||
|
|
delivered = False
|
||
|
|
for ws in conns:
|
||
|
|
try:
|
||
|
|
await ws.send_text(text)
|
||
|
|
delivered = True
|
||
|
|
except Exception:
|
||
|
|
logger.warning(f"Failed to send command to user {user_id}, removing connection")
|
||
|
|
await self.disconnect(ws, user_id)
|
||
|
|
|
||
|
|
if delivered:
|
||
|
|
# Initialize as pending feedback
|
||
|
|
self._feedback.setdefault(command_id, {
|
||
|
|
"command_id": command_id,
|
||
|
|
"status": "delivered",
|
||
|
|
"action": command.get("action"),
|
||
|
|
})
|
||
|
|
return command_id
|
||
|
|
return None
|
||
|
|
|
||
|
|
def store_feedback(self, feedback: dict[str, Any]) -> None:
|
||
|
|
"""Store feedback from frontend after command execution."""
|
||
|
|
command_id = feedback.get("command_id")
|
||
|
|
if command_id:
|
||
|
|
self._feedback[command_id] = feedback
|
||
|
|
logger.debug(f"AI UI Control: feedback stored for command {command_id}: {feedback.get('status')}")
|
||
|
|
|
||
|
|
def get_feedback(self, command_id: str) -> dict[str, Any] | None:
|
||
|
|
"""Get stored feedback for a command."""
|
||
|
|
return self._feedback.get(command_id)
|
||
|
|
|
||
|
|
def is_user_online(self, user_id: str) -> bool:
|
||
|
|
"""Check if a user has any active frontend connections."""
|
||
|
|
return user_id in self._connections and len(self._connections[user_id]) > 0
|
||
|
|
|
||
|
|
def get_online_users(self) -> list[str]:
|
||
|
|
"""Get list of currently connected user IDs."""
|
||
|
|
return list(self._connections.keys())
|
||
|
|
|
||
|
|
def cleanup_stale(self, timeout_seconds: int = 60) -> None:
|
||
|
|
"""Remove stale command tracking entries older than timeout."""
|
||
|
|
now = time.time()
|
||
|
|
stale_ids = [
|
||
|
|
cid for cid, ts in self._delivered_at.items()
|
||
|
|
if now - ts > timeout_seconds
|
||
|
|
]
|
||
|
|
for cid in stale_ids:
|
||
|
|
if cid not in self._feedback or self._feedback[cid].get("status") == "delivered":
|
||
|
|
self._feedback[cid] = {
|
||
|
|
"command_id": cid,
|
||
|
|
"status": "timeout",
|
||
|
|
"action": None,
|
||
|
|
"message": "Command timed out waiting for frontend response",
|
||
|
|
}
|
||
|
|
self._delivered_at.pop(cid, None)
|
||
|
|
self._command_user.pop(cid, None)
|