Files
sovereign_ai/SovereignAI-Starter/tests/test_agent_stream.py
T

117 lines
4.2 KiB
Python

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()