diff --git a/tests/sync_client.py b/tests/sync_client.py deleted file mode 100644 index ba6d1ea..0000000 --- a/tests/sync_client.py +++ /dev/null @@ -1,85 +0,0 @@ -"""Synchronous test adapter for ASGI apps without Starlette TestClient.""" - -import asyncio -import json - -import httpx -from fastapi import HTTPException -from pydantic import ValidationError - - -class Response: - def __init__(self, status_code, payload=None, headers=None, lines=None): - self.status_code = status_code - self._payload = payload - self.headers = headers or {} - self._lines = lines or [] - - def json(self): - return self._payload - - def iter_lines(self): - return iter(self._lines) - - -def _rag_request(method, path, payload): - from api.routes import rag - - try: - request = rag.QueryRequest(**(payload or {})) - except ValidationError as exc: - return Response(422, {"detail": exc.errors()}) - - endpoint = rag.query if path == "/query" else rag.stream - try: - result = endpoint(request, actor="system") - except HTTPException as exc: - return Response(exc.status_code, {"detail": exc.detail}) - - if path == "/stream": - # StreamingResponse wraps a synchronous iterator in an AnyIO worker - # thread; consuming that wrapper from a fresh event loop can deadlock - # under the current test runner. Recreate the endpoint's finite SSE - # payload from its mocked engine for deterministic client semantics. - try: - engine = rag.get_engine() - context = engine.build_context( - engine.retrieve(request.query, repo=request.repo, bundle=request.bundle) - ) - if hasattr(engine, "stream_llm"): - lines = [ - f"data: {{\"type\":\"token\",\"content\":{json.dumps(token)}}}" - for token in engine.stream_llm(request.query, context) - ] - else: - answer = engine.answer(request.query)["answer"] - lines = [f"data: {{\"type\":\"response\",\"content\":{json.dumps(answer)}}}"] - lines.append('data: {"type":"done"}') - except Exception as exc: - lines = [f"data: {{\"type\":\"error\",\"message\":{json.dumps(str(exc))}}}"] - return Response(200, headers={"content-type": "text/event-stream"}, lines=lines) - return Response(200, result) - - -class SyncASGIClient: - def __init__(self, app): - self.app = app - - def request(self, method, path, **kwargs): - if path in {"/query", "/stream"}: - return _rag_request(method, path, kwargs.get("json")) - - async def send(): - transport = httpx.ASGITransport(app=self.app) - async with httpx.AsyncClient( - transport=transport, base_url="http://testserver" - ) as client: - return await client.request(method, path, **kwargs) - - return asyncio.run(send()) - - def get(self, path, **kwargs): - return self.request("GET", path, **kwargs) - - def post(self, path, **kwargs): - return self.request("POST", path, **kwargs) diff --git a/tests/test_coverage_completion.py b/tests/test_coverage_completion.py index 604445b..ae139f9 100644 --- a/tests/test_coverage_completion.py +++ b/tests/test_coverage_completion.py @@ -1,6 +1,5 @@ """Focused branch coverage for the HTTP auth and RAG route helpers.""" -import asyncio from types import SimpleNamespace import jwt @@ -48,13 +47,9 @@ async def test_require_auth_all_modes_and_identity_fields(monkeypatch): await auth.require_auth(SimpleNamespace(), authorization="bad") -def test_rag_stream_success_fallback_and_error(monkeypatch): - captured = {} - +def test_rag_stream_success_and_error(monkeypatch): class FakeStreamingResponse: def __init__(self, content, media_type): - captured["content"] = content - captured["media_type"] = media_type self.content = content self.media_type = media_type @@ -63,26 +58,20 @@ def __init__(self, content, media_type): engine = SimpleNamespace( retrieve=lambda *args, **kwargs: [{"text": "ctx"}], - build_context=lambda docs: "context", - stream_llm=lambda query, context: ["one", "two"], + answer_from_docs=lambda query, docs: { + "answer": "single", "grounded": True, "answer_status": "ANSWERED", "answer_contract": "v1", + "citations": [], "context_used": len(docs), "llm_invoked": True, + }, ) monkeypatch.setattr(rag, "get_engine", lambda: engine) response = rag.stream(req, actor="system") assert response.media_type == "text/event-stream" events = list(response.content) - assert '"type": "token"' in events[0] + assert '"type": "status"' in events[0] + assert '"type": "response"' in events[1] and '"content": "single"' in events[1] assert '"type": "done"' in events[-1] - fallback = SimpleNamespace( - retrieve=lambda *args, **kwargs: [], - build_context=lambda docs: "", - answer=lambda query: {"answer": "single"}, - ) - monkeypatch.setattr(rag, "get_engine", lambda: fallback) - response = rag.stream(req, actor="system") - events = list(response.content) - assert '"type": "response"' in events[0] - monkeypatch.setattr(rag, "get_engine", lambda: (_ for _ in ()).throw(RuntimeError("boom"))) response = rag.stream(req, actor="system") - assert '"type": "error"' in list(response.content)[0] + events = list(response.content) + assert '"code": "internal_error"' in events[0] and "boom" not in events[0] diff --git a/tests/test_main_api.py b/tests/test_main_api.py index 80d9ec9..b62ece0 100644 --- a/tests/test_main_api.py +++ b/tests/test_main_api.py @@ -11,10 +11,6 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from unittest.mock import MagicMock, patch, AsyncMock -import asyncio -import sys -from starlette.requests import Request from fastapi.testclient import TestClient # Mock heavy dependencies properly @@ -30,32 +26,16 @@ with patch("index.vector_store.VectorStore"), patch("index.graph_store.GraphStore"), patch("index.plugin_index.PluginIndex"): from api.main import app -from api import main as main_module +client = TestClient(app) def test_health_endpoint(): """Report an ok status from /health while the control plane is ready.""" with patch("api.main.CONTROL_PLANE.status", return_value={"status": "READY"}): - response = main_module.health() - assert response["status"] == "ok" + response = client.get("/health") + assert response.status_code == 200 + assert response.json()["status"] == "ok" def test_status_endpoint(): - with patch("api.main.CONTROL_PLANE.status", return_value={"status": "READY"}): - with patch("api.main.graph_store") as mock_gs: - mock_gs.size.return_value = {"nodes": 3, "edges": 4} - response = main_module.status() - assert response["graph_edges"] == 4 - - -@pytest.mark.asyncio -async def test_guard_requests_middleware_ready(): - with patch("api.main.CONTROL_PLANE.status", return_value={"status": "READY"}): - with patch("api.routes.rag.get_engine"): - request = Request({"type": "http", "method": "POST", "path": "/rag/query", "headers": [], "query_string": b""}) - response = await main_module.guard_requests(request, lambda _: asyncio.sleep(0)) - assert response is None - -@pytest.mark.asyncio -async def test_guard_requests_middleware_not_ready(): """Report the control-plane status and the graph edge count from /status.""" with ( patch("api.main.CONTROL_PLANE.status", return_value={"status": "READY"}), @@ -121,10 +101,9 @@ def test_guard_requests_middleware_not_ready(): """Reject RAG requests with a 503 'Control plane not ready' while the control plane is not READY.""" with patch("api.main.CONTROL_PLANE.status", return_value={"status": "INIT"}): - request = Request({"type": "http", "method": "POST", "path": "/rag/query", "headers": [], "query_string": b""}) - response = await main_module.guard_requests(request, lambda _: asyncio.sleep(0)) + response = client.post("/rag/query", json={"query": "q"}) assert response.status_code == 503 - assert response.body == b'{"detail":"Control plane not ready"}' + assert response.json()["detail"] == "Control plane not ready" def test_build_graph_seed(): """Seed the application-level graph store with its four starter edges.""" diff --git a/tests/test_rag_routes.py b/tests/test_rag_routes.py index 0a3fac9..c0abf8c 100644 --- a/tests/test_rag_routes.py +++ b/tests/test_rag_routes.py @@ -9,9 +9,6 @@ import pytest from fastapi import FastAPI -from unittest.mock import MagicMock, patch -import json -from tests.sync_client import SyncASGIClient from fastapi.testclient import TestClient # Import the router and models from the target file @@ -23,7 +20,6 @@ @pytest.fixture def client(): - return SyncASGIClient(app) """Provide a TestClient for a throwaway app that mounts only the RAG router.""" return TestClient(app)