theo-agent-dashboard

Unnamed repository; edit this file 'description' to name the repository.
Log | Files | Refs | README

chat.py (6914B)


      1 """WebSocket chat endpoint — bridges client ↔ Hermes via SSE."""
      2 
      3 from __future__ import annotations
      4 
      5 import json
      6 import logging
      7 from typing import Any
      8 
      9 import aiohttp
     10 from fastapi import APIRouter, WebSocket, WebSocketDisconnect
     11 
     12 from app.config import get_config
     13 from app.hermes_client import _strip_untrusted_tags
     14 from app.routes.attachments import _fetch_metadata
     15 from app.sse_parser import parse_sse_events
     16 from app.ws_manager import ConnectionManager
     17 
     18 logger = logging.getLogger(__name__)
     19 
     20 chat_router = APIRouter()
     21 manager = ConnectionManager()
     22 
     23 
     24 async def _send_error(ws: WebSocket, message: str) -> None:  # pragma: no mutate: block
     25     """Send an error payload to the WebSocket client."""
     26     await ws.send_json({"type": "error", "message": message})
     27 
     28 
     29 def _map_sse_event_name(hermes_event: str) -> str:
     30     """Map Hermes SSE event names to frontend-expected names.
     31 
     32     The frontend's ChatView switch statement uses these exact type strings:
     33       run.started, message.started, assistant.delta, assistant.completed,
     34       tool.started, tool.completed, tool.failed, tool.progress, run.completed
     35 
     36     Hermes uses the same names, so this is a passthrough for known events.
     37     """
     38     return hermes_event
     39 
     40 
     41 def _map_sse_event_data(hermes_event: str, data: dict[str, Any]) -> dict[str, Any]:
     42     """Map Hermes SSE event data fields to what the frontend expects.
     43 
     44     Hermes field names → Frontend field names:
     45       assistant.delta:  message_id→id, delta→content
     46       message.started:  message{id,role}→{id,role} (unwrap nested)
     47       assistant.completed: message_id→id
     48       tool.started:     tool_name→name, args stays, preview→args
     49       tool.completed:   tool_name→name
     50       tool.failed:      tool_name→name
     51       tool.progress:    tool_name→name
     52     """
     53     if hermes_event == "assistant.delta":
     54         return {
     55             "id": data.get("message_id", ""),
     56             "content": data.get("delta", ""),
     57         }
     58     if hermes_event == "message.started":
     59         msg = data.get("message", data)
     60         return {
     61             "id": msg.get("id", ""),
     62             "role": msg.get("role", "assistant"),
     63             "model": msg.get("model"),
     64         }
     65     if hermes_event == "assistant.completed":
     66         return {
     67             "id": data.get("message_id", ""),
     68             "content": data.get("content", ""),
     69             "completed": data.get("completed", True),
     70         }
     71     if hermes_event in ("tool.started", "tool.completed", "tool.failed"):
     72         tool_id = data.get("tool_call_id", "") or data.get("tool_name", "")
     73         return {
     74             "messageId": data.get("message_id", ""),
     75             "toolCall": {
     76                 "id": tool_id,
     77                 "name": data.get("tool_name", ""),
     78                 "args": json.dumps(data.get("args", "")) if isinstance(data.get("args"), (dict, list)) else str(data.get("args", "")),
     79             },
     80             "toolCallId": tool_id,
     81             "result": _strip_untrusted_tags(str(data.get("preview", data.get("result", "")))),
     82             "error": data.get("error", ""),
     83         }
     84     if hermes_event == "tool.progress":
     85         return {
     86             "messageId": data.get("message_id", ""),
     87             "toolCallId": data.get("tool_call_id", data.get("tool_name", "")),
     88             "progress": data.get("delta", data.get("preview", "")),
     89         }
     90     # run.started, run.completed, error, done — pass through
     91     return data
     92 
     93 
     94 async def _stream_hermes_sse(  # pragma: no mutate: block
     95     session: aiohttp.ClientSession,
     96     session_id: str,
     97     content: str,
     98     ws: WebSocket,
     99     base_url: str,
    100     api_key: str,
    101 ) -> None:
    102     """POST to the SSE chat stream endpoint and forward events to the WS client.
    103 
    104     Hermes's session chat stream is a POST endpoint that accepts a message
    105     body and returns SSE events.  This function:
    106       1. POSTs the message to trigger the agent turn
    107       2. Reads the SSE response
    108       3. Forwards each event to the frontend with mapped field names
    109     """
    110     url = f"{base_url}/api/sessions/{session_id}/chat/stream"
    111     headers = {"Authorization": f"Bearer {api_key}"}
    112     payload = {"message": content}
    113     async with session.post(url, json=payload, headers=headers) as resp:
    114         resp.raise_for_status()
    115         text = await resp.text()
    116         for event in parse_sse_events(text):
    117             event_name = _map_sse_event_name(event["event"])
    118             event_data = _map_sse_event_data(event["event"], event["data"])
    119             await ws.send_json({"type": event_name, "data": event_data})
    120 
    121 
    122 async def _load_attachment_content(attachment_ids: list[str]) -> str:
    123     """Resolve attachment IDs to content for the Hermes message.
    124 
    125     Both images and non-image files have their local disk path appended.
    126     Returns a string to append to the message.
    127     """
    128     parts: list[str] = []
    129     for att_id in attachment_ids:
    130         meta = _fetch_metadata(att_id)
    131         if meta is None:
    132             continue
    133         disk_path = meta.get("disk_path", "")
    134         if disk_path:
    135             parts.append(disk_path)
    136     return "\n".join(parts)
    137 
    138 
    139 async def _handle_chat(  # pragma: no mutate: block
    140     ws: WebSocket,
    141     payload: dict,
    142     cfg: Any,
    143 ) -> None:
    144     """Process a single 'chat' message from the client."""
    145     # The WsClient wraps sends as {type, data}, so fields live in payload["data"]
    146     data = payload.get("data", payload)
    147     session_id: str = data["session_id"]
    148     content: str = data["content"]
    149     attachment_ids: list[str] | None = data.get("attachments")
    150 
    151     if attachment_ids:
    152         attachment_content = await _load_attachment_content(attachment_ids)
    153         if attachment_content:
    154             content = f"{content}\n\n{attachment_content}"
    155 
    156     async with aiohttp.ClientSession() as session:
    157         await _stream_hermes_sse(
    158             session, session_id, content, ws,
    159             cfg.hermes_base_url, cfg.hermes_api_key,
    160         )
    161 
    162 
    163 async def ws_chat_handler(ws: WebSocket) -> None:  # pragma: no mutate: block
    164     """Main WebSocket handler for the chat endpoint."""
    165     await ws.accept()
    166     await ws.send_json({"type": "status", "data": {"connected": True}})
    167     await manager.connect("default", ws)
    168     try:
    169         cfg = get_config()
    170         while True:
    171             payload = await ws.receive_json()
    172             msg_type = payload.get("type")
    173             if msg_type == "chat":
    174                 await _handle_chat(ws, payload, cfg)
    175             else:
    176                 await _send_error(ws, f"Unknown type: {msg_type}")
    177     except WebSocketDisconnect:
    178         pass
    179     except Exception:
    180         logger.exception("WebSocket chat error")
    181         await _send_error(ws, "Internal server error")
    182     finally:
    183         manager.disconnect("default")
    184 
    185 
    186 @chat_router.websocket("/ws/chat")
    187 async def ws_chat_route(websocket: WebSocket) -> None:  # pragma: no mutate: block
    188     """FastAPI WebSocket route at ``/ws/chat``."""
    189     await ws_chat_handler(websocket)