99 lines
3.6 KiB
Python
99 lines
3.6 KiB
Python
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 import database
|
|
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, user_id):
|
|
self.assertEqual(user_id, database.LOCAL_USER_ID)
|
|
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, user_id=database.LOCAL_USER_ID)
|
|
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()
|