theo-agent-dashboard

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

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