import asyncio import os import tempfile import unittest from unittest.mock import patch from fastapi import Request from fastapi.responses import StreamingResponse _TEST_DATA_DIR = None if "SOVEREIGNAI_DATA_DIR" not in os.environ: _TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-stream-tests-") os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name from app import main class AgentStreamTests(unittest.IsolatedAsyncioTestCase): def _agent_request(self) -> Request: return Request( { "type": "http", "asgi": {"version": "3.0"}, "http_version": "1.1", "method": "POST", "scheme": "http", "path": "/v1/agent/run/stream", "raw_path": b"/v1/agent/run/stream", "query_string": b"", "headers": [], "client": ("127.0.0.1", 12345), "server": ("127.0.0.1", 8000), "root_path": "", } ) async def test_audit_duration_includes_complete_stream_lifetime(self) -> None: recorded = [] async def body(): yield b"event: progress\n\n" await asyncio.sleep(0.04) yield b"event: done\n\n" async def call_next(_request): return StreamingResponse(body(), media_type="text/event-stream") with patch.object( main.database, "record_agent_audit_event", side_effect=lambda **event: recorded.append(event), ): response = await main.audit_agent_routes(self._agent_request(), call_next) chunks = [chunk async for chunk in response.body_iterator] self.assertEqual(chunks, [b"event: progress\n\n", b"event: done\n\n"]) self.assertEqual(len(recorded), 1) self.assertEqual(recorded[0]["status_code"], 200) self.assertGreaterEqual(recorded[0]["duration_ms"], 35) async def test_audit_marks_client_closed_stream_as_499(self) -> None: recorded = [] started = asyncio.Event() async def body(): yield b"event: progress\n\n" started.set() await asyncio.Future() async def call_next(_request): return StreamingResponse(body(), media_type="text/event-stream") with patch.object( main.database, "record_agent_audit_event", side_effect=lambda **event: recorded.append(event), ): response = await main.audit_agent_routes(self._agent_request(), call_next) stream = response.body_iterator self.assertEqual(await anext(stream), b"event: progress\n\n") waiting_read = asyncio.create_task(anext(stream)) await asyncio.wait_for(started.wait(), timeout=1) waiting_read.cancel() with self.assertRaises(asyncio.CancelledError): await waiting_read self.assertEqual(len(recorded), 1) self.assertEqual(recorded[0]["status_code"], 499) async def test_idle_model_wait_sends_heartbeat_then_final_result(self) -> None: async def slow_agent(request, report_progress=None, *, user_id): await report_progress("بدأ تحليل المهمة.") await asyncio.sleep(0.05) return {"task": request.task, "result": "اكتمل التحليل."} with ( patch.object(main, "AGENT_STREAM_HEARTBEAT_SECONDS", 0.01), patch.object(main, "_execute_agent", side_effect=slow_agent), ): response = await main.run_agent_stream( main.AgentRequest(task="اختبار انتظار النموذج"), user_id="test-user" ) chunks = [chunk async for chunk in response.body_iterator] body = b"".join( chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in chunks ).decode("utf-8") self.assertIn("event: progress", body) self.assertIn("event: heartbeat\ndata: {}", body) self.assertIn("event: done", body) self.assertIn("اكتمل التحليل", body) if __name__ == "__main__": unittest.main()