Files
AIRegulation-DocAnalysis/tests/test_status_models_routes.py
T

134 lines
5.7 KiB
Python
Raw Normal View History

"""Integration tests for the /status/models routes.
Uses FastAPI TestClient with mocked LLM/embedding/reranker clients so no
external gateway or database is required.
"""
from __future__ import annotations
from unittest.mock import MagicMock, patch
import pytest
from fastapi.testclient import TestClient
from app.services.llm.base_client import LLMResponse
from app.shared.model_usage_tracker import get_model_usage_tracker
@pytest.fixture(autouse=True)
def _reset_tracker():
"""Clear the process-wide tracker before and after each test in this file."""
get_model_usage_tracker()._entries.clear()
yield
get_model_usage_tracker()._entries.clear()
@pytest.fixture
def client():
"""Return a TestClient for the real app (status routes require no auth)."""
from app.api.main import app
with TestClient(app, raise_server_exceptions=False) as c:
yield c
def test_get_models_returns_four_roles_never_called_by_default(client):
"""With no calls made yet, all 4 roles are returned with status 'never_called' or 'disabled'."""
resp = client.get("/api/v1/status/models")
assert resp.status_code == 200
body = resp.json()
roles = {m["role"] for m in body["models"]}
assert roles == {"main_llm", "hyde_llm", "embedding", "reranker"}
reranker_row = next(m for m in body["models"] if m["role"] == "reranker")
# Default .env.example ships RERANKER_ENABLED=false.
from app.config.settings import settings
assert reranker_row["enabled"] == settings.reranker_enabled
if not settings.reranker_enabled:
assert reranker_row["status"] == "disabled"
def test_get_models_reflects_recorded_usage(client):
"""A previously recorded call must show up in total_tokens/status."""
from app.config.settings import settings
get_model_usage_tracker().record(
provider=settings.llm_provider, model=settings.llm_model, success=True, usage={"total_tokens": 99},
)
resp = client.get("/api/v1/status/models")
main_row = next(m for m in resp.json()["models"] if m["role"] == "main_llm")
assert main_row["total_tokens"] == 99
assert main_row["status"] == "ok"
def test_get_models_hyde_llm_disabled_forces_disabled_status(client):
"""settings.hyde_enabled=False must force hyde_llm to enabled=False/status='disabled',
mirroring the reranker override, even if HyDE previously ran successfully."""
from app.config.settings import settings
get_model_usage_tracker().record(
provider=settings.hyde_llm_provider or settings.llm_provider,
model=settings.hyde_llm_model or settings.llm_model,
success=True,
)
with patch.object(settings, "hyde_enabled", False):
resp = client.get("/api/v1/status/models")
assert resp.status_code == 200
hyde_row = next(m for m in resp.json()["models"] if m["role"] == "hyde_llm")
assert hyde_row["enabled"] is False
assert hyde_row["status"] == "disabled"
def test_ping_models_calls_each_enabled_model_once(client):
"""POST /status/models/ping must invoke chat()/embed_query() and return fresh statuses."""
mock_llm_response = LLMResponse(content="pong", model="test-model", usage={"total_tokens": 1})
mock_llm_client = MagicMock()
mock_llm_client.chat.return_value = mock_llm_response
mock_embedding = MagicMock()
mock_embedding.embed_query.return_value = [0.1]
with patch("app.api.routes.status.get_llm_client", return_value=mock_llm_client), \
patch("app.api.routes.status.get_embedding_provider", return_value=mock_embedding), \
patch("app.api.routes.status.get_reranker", return_value=None):
resp = client.post("/api/v1/status/models/ping")
assert resp.status_code == 200
body = resp.json()
assert len(body["models"]) == 4
assert mock_llm_client.chat.call_count >= 1
mock_embedding.embed_query.assert_called_once()
def test_ping_models_survives_one_model_failing(client):
"""If the LLM ping raises, embedding/reranker pings must still be attempted and a 200 returned."""
mock_embedding = MagicMock()
mock_embedding.embed_query.return_value = [0.1]
with patch("app.api.routes.status.get_llm_client", side_effect=RuntimeError("gateway down")), \
patch("app.api.routes.status.get_embedding_provider", return_value=mock_embedding), \
patch("app.api.routes.status.get_reranker", return_value=None):
resp = client.post("/api/v1/status/models/ping")
assert resp.status_code == 200
mock_embedding.embed_query.assert_called_once()
def test_ping_records_get_llm_client_failure_instead_of_dropping_it(client):
"""A get_llm_client() failure (raised before any TrackedLLMClient exists) must still
be recorded into the tracker, so it is visible afterwards via _build_model_status()
instead of being silently discarded by asyncio.gather(return_exceptions=True)."""
mock_embedding = MagicMock()
mock_embedding.embed_query.return_value = [0.1]
with patch("app.api.routes.status.get_llm_client", side_effect=RuntimeError("missing api key")), \
patch("app.api.routes.status.get_embedding_provider", return_value=mock_embedding), \
patch("app.api.routes.status.get_reranker", return_value=None):
resp = client.post("/api/v1/status/models/ping")
assert resp.status_code == 200
body = resp.json()
main_row = next(m for m in body["models"] if m["role"] == "main_llm")
hyde_row = next(m for m in body["models"] if m["role"] == "hyde_llm")
assert main_row["status"] == "error"
assert main_row["call_count_error"] == 1
assert main_row["last_error"] == "missing api key"
assert hyde_row["status"] == "error"
assert hyde_row["call_count_error"] == 1
assert hyde_row["last_error"] == "missing api key"