Complete local hybrid search and improve agent reliability
This commit is contained in:
@@ -0,0 +1,96 @@
|
||||
import asyncio
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
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
|
||||
|
||||
from app.main import AgentRequest, run_agent_stream
|
||||
from app.model_provider import OllamaProvider
|
||||
|
||||
|
||||
class TimeoutAndCancellationTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_provider_timeout_is_bounded_and_reported_as_504(self) -> None:
|
||||
observed_timeouts: list[float] = []
|
||||
|
||||
class TimeoutClient:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *_args):
|
||||
return None
|
||||
|
||||
async def post(self, *_args, **_kwargs):
|
||||
request = httpx.Request("POST", "http://127.0.0.1/chat/completions")
|
||||
raise httpx.ReadTimeout("test timeout", request=request)
|
||||
|
||||
def client_factory(*, timeout):
|
||||
observed_timeouts.append(timeout)
|
||||
return TimeoutClient()
|
||||
|
||||
provider = OllamaProvider("http://127.0.0.1:11434/v1", "test-model")
|
||||
with patch("app.model_provider.httpx.AsyncClient", side_effect=client_factory):
|
||||
with self.assertRaises(HTTPException) as caught:
|
||||
await provider.complete({"messages": []}, timeout_seconds=0.025)
|
||||
|
||||
self.assertEqual(observed_timeouts, [0.025])
|
||||
self.assertEqual(caught.exception.status_code, 504)
|
||||
|
||||
async def test_provider_cancellation_propagates_and_closes_client(self) -> None:
|
||||
started = asyncio.Event()
|
||||
closed = asyncio.Event()
|
||||
|
||||
class WaitingClient:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *_args):
|
||||
closed.set()
|
||||
return None
|
||||
|
||||
async def post(self, *_args, **_kwargs):
|
||||
started.set()
|
||||
await asyncio.Future()
|
||||
|
||||
provider = OllamaProvider("http://127.0.0.1:11434/v1", "test-model")
|
||||
with patch("app.model_provider.httpx.AsyncClient", return_value=WaitingClient()):
|
||||
task = asyncio.create_task(provider.complete({"messages": []}))
|
||||
await asyncio.wait_for(started.wait(), timeout=1)
|
||||
task.cancel()
|
||||
with self.assertRaises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
self.assertTrue(closed.is_set())
|
||||
|
||||
async def test_disconnecting_agent_stream_cancels_agent_task(self) -> None:
|
||||
started = asyncio.Event()
|
||||
cancelled = asyncio.Event()
|
||||
|
||||
async def waiting_agent(_request, *, report_progress):
|
||||
await report_progress("بدأ الاختبار")
|
||||
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):
|
||||
response = await run_agent_stream(request)
|
||||
stream = response.body_iterator
|
||||
first_event = await asyncio.wait_for(anext(stream), timeout=1)
|
||||
self.assertIn("event: progress", first_event)
|
||||
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