fix: keep agent streams alive during long tasks

This commit is contained in:
Hamza Ayed
2026-10-04 13:41:22 +03:00
parent ad221e24aa
commit d1010b79d7
11 changed files with 334 additions and 37 deletions
@@ -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()