Files
leocrm/app/plugins/builtins/ai_ui_control/websocket_manager.py
T

130 lines
5.0 KiB
Python
Raw Normal View History

"""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)