Files
sovereign_ai/SovereignAI-Starter/tests/test_timeouts_and_cancellation.py
T

127 lines
4.8 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 = 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
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)
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()