fix: keep agent streams alive during long tasks
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
import os
|
||||
import tempfile
|
||||
import asyncio
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, patch
|
||||
@@ -50,7 +51,18 @@ class MixedPdfKnowledgeIntegrationTests(unittest.TestCase):
|
||||
database.initialize_database()
|
||||
knowledge.initialize()
|
||||
client = authenticated_client(app)
|
||||
with patch("app.main.recognize_image_text", return_value=ocr):
|
||||
ocr_execution_contexts: list[str] = []
|
||||
|
||||
def fake_ocr(*_args, **_kwargs):
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
ocr_execution_contexts.append("worker")
|
||||
else:
|
||||
ocr_execution_contexts.append("event-loop")
|
||||
return ocr
|
||||
|
||||
with patch("app.main.recognize_image_text", side_effect=fake_ocr):
|
||||
indexed = client.post(
|
||||
"/v1/agent/knowledge/index",
|
||||
json={"workspace_path": str(workspace), "files": ["mixed.pdf"]},
|
||||
@@ -75,6 +87,7 @@ class MixedPdfKnowledgeIntegrationTests(unittest.TestCase):
|
||||
|
||||
self.assertEqual(indexed.status_code, 200, indexed.text)
|
||||
self.assertEqual(indexed.json()["indexed"][0]["ocr_pages"], 1)
|
||||
self.assertEqual(ocr_execution_contexts, ["worker"])
|
||||
self.assertTrue(indexed.json()["indexed"][0]["semantic_indexed"])
|
||||
self.assertEqual(digital.status_code, 200, digital.text)
|
||||
self.assertIn("IndexMarker", digital.json()["results"][0]["text"])
|
||||
|
||||
@@ -8,8 +8,10 @@ import httpx
|
||||
from fastapi import HTTPException
|
||||
|
||||
# Importing the API initializes SQLite; keep this test process away from user data.
|
||||
_TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-timeout-tests-")
|
||||
os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name
|
||||
_TEST_DATA_DIR = None
|
||||
if not os.environ.get("SOVEREIGNAI_DATA_DIR"):
|
||||
_TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-timeout-tests-")
|
||||
os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name
|
||||
|
||||
from app import database
|
||||
from app.main import AgentRequest, run_agent_stream
|
||||
@@ -93,6 +95,32 @@ class TimeoutAndCancellationTests(unittest.IsolatedAsyncioTestCase):
|
||||
await stream.aclose()
|
||||
await asyncio.wait_for(cancelled.wait(), timeout=1)
|
||||
|
||||
async def test_idle_agent_stream_emits_keepalive_until_model_finishes(self) -> None:
|
||||
started = asyncio.Event()
|
||||
cancelled = asyncio.Event()
|
||||
|
||||
async def waiting_agent(_request, *, report_progress, user_id):
|
||||
self.assertEqual(user_id, database.LOCAL_USER_ID)
|
||||
started.set()
|
||||
try:
|
||||
await asyncio.Future()
|
||||
except asyncio.CancelledError:
|
||||
cancelled.set()
|
||||
raise
|
||||
|
||||
request = AgentRequest(task="اختبار نبضة اتصال الوكيل")
|
||||
with (
|
||||
patch("app.main._execute_agent", side_effect=waiting_agent),
|
||||
patch("app.main.AGENT_STREAM_HEARTBEAT_SECONDS", 0.01),
|
||||
):
|
||||
response = await run_agent_stream(request, user_id=database.LOCAL_USER_ID)
|
||||
stream = response.body_iterator
|
||||
keepalive = await asyncio.wait_for(anext(stream), timeout=1)
|
||||
self.assertEqual(keepalive, ": keep-alive\n\n")
|
||||
self.assertTrue(started.is_set())
|
||||
await stream.aclose()
|
||||
await asyncio.wait_for(cancelled.wait(), timeout=1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user