117 lines
4.2 KiB
Python
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()
|