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)