test_chat.py (15697B)
1 """T024 — Tests for the chat WebSocket endpoint.""" 2 3 from __future__ import annotations 4 5 import json 6 from typing import Any 7 from unittest.mock import AsyncMock, MagicMock, patch 8 import pytest 9 from fastapi import FastAPI 10 from fastapi.testclient import TestClient 11 from starlette.websockets import WebSocketDisconnect 12 13 from app.routes.chat import chat_router 14 from app.ws_manager import ConnectionManager 15 16 # ── helpers ────────────────────────────────────────────────────────────── 17 18 19 class FakeWebSocket: 20 """Tracks sent JSON and simulates receive / lifecycle.""" 21 22 def __init__(self, receive_side_effects: list[Any]) -> None: 23 self.sent: list[dict] = [] 24 self._receive_queue = receive_side_effects 25 self._idx = 0 26 27 async def accept(self) -> None: 28 pass 29 30 async def send_json(self, data: dict) -> None: 31 self.sent.append(data) 32 33 async def receive_json(self) -> dict: 34 item = self._receive_queue[self._idx] 35 self._idx += 1 36 if isinstance(item, Exception): 37 raise item 38 return item 39 40 async def close(self, code: int = 1000) -> None: 41 pass 42 43 class _AsyncCtx: 44 """Async context manager wrapper for a mock response.""" 45 46 def __init__(self, response: Any) -> None: 47 self._response = response 48 49 async def __aenter__(self) -> Any: 50 return self._response 51 52 async def __aexit__(self, *args: Any) -> None: 53 pass 54 55 56 def _make_sse_stream(events: list[dict]) -> str: 57 """Build an SSE text block from a list of event dicts.""" 58 parts: list[str] = [] 59 for ev in events: 60 parts.append(f"event: {ev['event']}") 61 parts.append(f"data: {json.dumps(ev['data'])}") 62 parts.append("") 63 parts.append("") 64 return "\n".join(parts) 65 66 # ── tests ──────────────────────────────────────────────────────────────── 67 68 69 @pytest.mark.asyncio 70 async def test_connect_sends_status() -> None: 71 """On connect, the endpoint sends a status event with connected=true.""" 72 mgr = ConnectionManager() 73 with patch("app.routes.chat.manager", mgr): 74 ws = FakeWebSocket([WebSocketDisconnect()]) 75 from app.routes.chat import ws_chat_handler 76 await ws_chat_handler(ws) 77 78 assert ws.sent[0] == {"type": "status", "data": {"connected": True}} 79 80 81 @pytest.mark.asyncio 82 async def test_chat_message_forwards_to_hermes() -> None: 83 """Sending a 'chat' message triggers Hermes SSE POST and streams back.""" 84 fake_sse = _make_sse_stream([ 85 {"event": "run.started", "data": {"run_id": "r1"}}, 86 {"event": "assistant.delta", "data": {"message_id": "m1", "delta": "Hello"}}, 87 {"event": "run.completed", "data": {"run_id": "r1"}}, 88 ]) 89 90 post_resp = AsyncMock() 91 post_resp.status = 200 92 post_resp.text = AsyncMock(return_value=fake_sse) 93 post_resp.raise_for_status = AsyncMock() 94 95 session = MagicMock() 96 session.post = MagicMock(return_value=_AsyncCtx(post_resp)) 97 98 with patch("app.routes.chat.aiohttp.ClientSession", return_value=_AsyncCtx(session)): 99 mgr = ConnectionManager() 100 with patch("app.routes.chat.manager", mgr): 101 with patch("app.routes.chat.get_config") as mock_cfg: 102 mock_cfg.return_value.hermes_base_url = "http://fake" 103 mock_cfg.return_value.hermes_api_key = "test-key" 104 105 ws = FakeWebSocket([ 106 {"type": "chat", "session_id": "s1", "content": "Hi"}, 107 WebSocketDisconnect(), 108 ]) 109 110 from app.routes.chat import ws_chat_handler 111 await ws_chat_handler(ws) 112 113 assert ws.sent[0] == {"type": "status", "data": {"connected": True}} 114 # Events are forwarded with type = event name (not "event") 115 events = [e for e in ws.sent if e.get("type") not in ("status", "error")] 116 assert len(events) == 3 117 event_types = [e["type"] for e in events] 118 assert "run.started" in event_types 119 assert "assistant.delta" in event_types 120 assert "run.completed" in event_types 121 122 # Verify assistant.delta data is mapped: message_id→id, delta→content 123 delta = [e for e in events if e["type"] == "assistant.delta"][0] 124 assert delta["data"]["id"] == "m1" 125 assert delta["data"]["content"] == "Hello" 126 127 128 @pytest.mark.asyncio 129 async def test_chat_posts_to_stream_endpoint() -> None: 130 """Verify the POST goes to /chat/stream with correct auth and body.""" 131 fake_sse = 'event: done\ndata: {}\n\n' 132 133 post_resp = AsyncMock() 134 post_resp.status = 200 135 post_resp.text = AsyncMock(return_value=fake_sse) 136 post_resp.raise_for_status = AsyncMock() 137 138 session = MagicMock() 139 post_mock = MagicMock(return_value=_AsyncCtx(post_resp)) 140 session.post = post_mock 141 142 with patch("app.routes.chat.aiohttp.ClientSession", return_value=_AsyncCtx(session)): 143 mgr = ConnectionManager() 144 with patch("app.routes.chat.manager", mgr): 145 with patch("app.routes.chat.get_config") as mock_cfg: 146 mock_cfg.return_value.hermes_base_url = "http://hermes:8642" 147 mock_cfg.return_value.hermes_api_key = "sk-test-key" 148 149 ws = FakeWebSocket([ 150 {"type": "chat", "session_id": "s1", "content": "Hello"}, 151 WebSocketDisconnect(), 152 ]) 153 154 from app.routes.chat import ws_chat_handler 155 await ws_chat_handler(ws) 156 157 # Verify POST was called to the stream endpoint 158 call_args = post_mock.call_args 159 assert call_args[0][0] == "http://hermes:8642/api/sessions/s1/chat/stream" 160 assert call_args[1]["json"] == {"message": "Hello"} 161 assert call_args[1]["headers"]["Authorization"] == "Bearer sk-test-key" 162 163 164 @pytest.mark.asyncio 165 async def test_handles_hermes_connection_error() -> None: 166 """Hermes API error sends an error event to the client.""" 167 import aiohttp 168 169 with patch("app.routes.chat.aiohttp.ClientSession") as mock_session_cls: 170 mock_session = AsyncMock() 171 mock_session.post.side_effect = aiohttp.ClientError("connection refused") 172 mock_session.__aenter__ = AsyncMock(return_value=mock_session) 173 mock_session.__aexit__ = AsyncMock(return_value=False) 174 mock_session_cls.return_value = mock_session 175 176 mgr = ConnectionManager() 177 with patch("app.routes.chat.manager", mgr): 178 with patch("app.routes.chat.get_config") as mock_cfg: 179 mock_cfg.return_value.hermes_base_url = "http://fake" 180 mock_cfg.return_value.hermes_api_key = "test-key" 181 182 ws = FakeWebSocket([ 183 {"type": "chat", "session_id": "s1", "content": "Hi"}, 184 WebSocketDisconnect(), 185 ]) 186 187 from app.routes.chat import ws_chat_handler 188 await ws_chat_handler(ws) 189 190 assert ws.sent[0] == {"type": "status", "data": {"connected": True}} 191 error_msgs = [e for e in ws.sent if e.get("type") == "error"] 192 assert len(error_msgs) == 1 193 assert error_msgs[0]["message"] == "Internal server error" 194 195 196 @pytest.mark.asyncio 197 async def test_unknown_message_type_sends_error() -> None: 198 """An unrecognized message type returns an error event.""" 199 mgr = ConnectionManager() 200 with patch("app.routes.chat.manager", mgr): 201 ws = FakeWebSocket([ 202 {"type": "ping", "session_id": "s1"}, 203 WebSocketDisconnect(), 204 ]) 205 206 from app.routes.chat import ws_chat_handler 207 await ws_chat_handler(ws) 208 209 assert ws.sent[0] == {"type": "status", "data": {"connected": True}} 210 error_msgs = [e for e in ws.sent if e.get("type") == "error"] 211 assert len(error_msgs) == 1 212 assert error_msgs[0]["type"] == "error" 213 assert "Unknown type" in error_msgs[0]["message"] 214 assert "ping" in error_msgs[0]["message"] 215 216 217 class TestSendError: 218 """_send_error sends error JSON to WebSocket.""" 219 220 async def test_sends_error_payload(self) -> None: 221 from unittest.mock import AsyncMock 222 from app.routes.chat import _send_error 223 ws = AsyncMock() 224 await _send_error(ws, "Something failed") 225 ws.send_json.assert_awaited_once_with({"type": "error", "message": "Something failed"}) 226 227 228 class TestStreamHermesSSE: 229 """_stream_hermes_sse POSTs to the stream endpoint and forwards events.""" 230 231 async def test_forwards_mapped_events(self) -> None: 232 """Events are forwarded with type = event name and mapped data fields.""" 233 from unittest.mock import AsyncMock, MagicMock 234 from app.routes.chat import _stream_hermes_sse 235 236 ws = AsyncMock() 237 session = AsyncMock() 238 resp = AsyncMock() 239 resp.text = AsyncMock( 240 return_value='event: assistant.delta\ndata: {"message_id": "m1", "delta": "Hi"}\n\n' 241 ) 242 resp.raise_for_status = MagicMock() 243 resp.__aenter__ = AsyncMock(return_value=resp) 244 resp.__aexit__ = AsyncMock(return_value=False) 245 session.post = MagicMock(return_value=resp) 246 247 await _stream_hermes_sse(session, "s1", "Hello", ws, "http://localhost:8642", "sk-test") 248 ws.send_json.assert_awaited() 249 call_args = ws.send_json.call_args[0][0] 250 # type should be the event name, not "event" 251 assert call_args["type"] == "assistant.delta" 252 # data fields should be mapped 253 assert call_args["data"]["id"] == "m1" 254 assert call_args["data"]["content"] == "Hi" 255 256 async def test_sends_post_with_auth(self) -> None: 257 """Verify POST is used with correct auth header and message body.""" 258 from unittest.mock import AsyncMock, MagicMock 259 from app.routes.chat import _stream_hermes_sse 260 261 ws = AsyncMock() 262 session = AsyncMock() 263 resp = AsyncMock() 264 resp.text = AsyncMock(return_value='event: done\ndata: {}\n\n') 265 resp.raise_for_status = MagicMock() 266 resp.__aenter__ = AsyncMock(return_value=resp) 267 resp.__aexit__ = AsyncMock(return_value=False) 268 post_mock = MagicMock(return_value=resp) 269 session.post = post_mock 270 271 await _stream_hermes_sse(session, "sess-42", "test msg", ws, "http://hermes:8642", "sk-abc") 272 post_mock.assert_called_once_with( 273 "http://hermes:8642/api/sessions/sess-42/chat/stream", 274 json={"message": "test msg"}, 275 headers={"Authorization": "Bearer sk-abc"}, 276 ) 277 278 279 class TestMapSSEEventName: 280 """Event name mapping passthrough test.""" 281 282 def test_known_events_pass_through(self): 283 from app.routes.chat import _map_sse_event_name 284 for name in ["run.started", "message.started", "assistant.delta", 285 "assistant.completed", "tool.started", "tool.completed", 286 "tool.failed", "tool.progress", "run.completed"]: 287 assert _map_sse_event_name(name) == name 288 289 def test_unknown_events_pass_through(self): 290 from app.routes.chat import _map_sse_event_name 291 assert _map_sse_event_name("custom.event") == "custom.event" 292 293 294 class TestMapSSEEventData: 295 """Event data field mapping tests.""" 296 297 def test_assistant_delta_maps_fields(self): 298 from app.routes.chat import _map_sse_event_data 299 data = _map_sse_event_data("assistant.delta", {"message_id": "m1", "delta": "Hello"}) 300 assert data == {"id": "m1", "content": "Hello"} 301 302 def test_message_started_unwraps_nested(self): 303 from app.routes.chat import _map_sse_event_data 304 data = _map_sse_event_data("message.started", {"message": {"id": "m1", "role": "assistant"}}) 305 assert data["id"] == "m1" 306 assert data["role"] == "assistant" 307 308 def test_assistant_completed_maps_fields(self): 309 from app.routes.chat import _map_sse_event_data 310 data = _map_sse_event_data("assistant.completed", {"message_id": "m1", "content": "Final"}) 311 assert data["id"] == "m1" 312 assert data["content"] == "Final" 313 314 def test_tool_started_maps_fields(self): 315 from app.routes.chat import _map_sse_event_data 316 data = _map_sse_event_data("tool.started", { 317 "message_id": "m1", "tool_name": "web_search", "args": {"q": "test"} 318 }) 319 assert data["messageId"] == "m1" 320 assert data["toolCall"]["name"] == "web_search" 321 322 def test_run_started_passthrough(self): 323 from app.routes.chat import _map_sse_event_data 324 data = _map_sse_event_data("run.started", {"run_id": "r1"}) 325 assert data == {"run_id": "r1"} 326 327 328 class TestAttachmentForwarding: 329 """Attachments sent with chat messages are resolved and forwarded to Hermes.""" 330 331 @pytest.mark.asyncio 332 async def test_chat_with_image_attachment_forwards_data_url(self) -> None: 333 """An image attachment forwards its disk path (not base64) to Hermes.""" 334 from app.routes.chat import _load_attachment_content 335 336 mock_meta = { 337 "id": "att-123", 338 "filename": "screenshot.png", 339 "content_type": "image/png", 340 "size_bytes": 1024, 341 "disk_path": "/data/attachments/s1/att-123.png", 342 "created_at": 1000, 343 } 344 345 with patch("app.routes.chat._fetch_metadata", return_value=mock_meta): 346 result = await _load_attachment_content(["att-123"]) 347 348 # Should contain the file path, NOT a base64 data URL 349 assert "/data/attachments/s1/att-123.png" in result 350 assert "data:image" not in result 351 352 @pytest.mark.asyncio 353 async def test_chat_with_non_image_attachment_forwards_path(self) -> None: 354 """A non-image attachment has its disk path appended to the message.""" 355 from app.routes.chat import _load_attachment_content 356 357 mock_meta = { 358 "id": "att-456", 359 "filename": "document.pdf", 360 "content_type": "application/pdf", 361 "size_bytes": 1024, 362 "disk_path": "/data/attachments/s1/att-456.pdf", 363 "created_at": 1000, 364 } 365 366 with patch("app.routes.chat._fetch_metadata", return_value=mock_meta): 367 result = await _load_attachment_content(["att-456"]) 368 369 assert "/data/attachments/s1/att-456.pdf" in result 370 371 @pytest.mark.asyncio 372 async def test_chat_without_attachments_unchanged(self) -> None: 373 """Sending a message without attachments does not change the payload.""" 374 fake_sse = 'event: done\ndata: {}\n\n' 375 376 post_resp = AsyncMock() 377 post_resp.status = 200 378 post_resp.text = AsyncMock(return_value=fake_sse) 379 post_resp.raise_for_status = AsyncMock() 380 381 session = MagicMock() 382 post_mock = MagicMock(return_value=_AsyncCtx(post_resp)) 383 session.post = post_mock 384 385 with patch("app.routes.chat.aiohttp.ClientSession", return_value=_AsyncCtx(session)): 386 mgr = ConnectionManager() 387 with patch("app.routes.chat.manager", mgr): 388 with patch("app.routes.chat.get_config") as mock_cfg: 389 mock_cfg.return_value.hermes_base_url = "http://hermes:8642" 390 mock_cfg.return_value.hermes_api_key = "sk-test" 391 392 ws = FakeWebSocket([ 393 {"type": "chat", "session_id": "s1", "content": "Hello"}, 394 WebSocketDisconnect(), 395 ]) 396 397 from app.routes.chat import ws_chat_handler 398 await ws_chat_handler(ws) 399 400 call_args = post_mock.call_args 401 assert call_args[1]["json"] == {"message": "Hello"}