fix agent stream audit lifecycle
This commit is contained in:
@@ -4,6 +4,9 @@ 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-")
|
||||
@@ -13,6 +16,77 @@ 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("بدأ تحليل المهمة.")
|
||||
|
||||
Reference in New Issue
Block a user