main-ruqi #1
@@ -1,18 +1,24 @@
|
|||||||
"""Define API routes for status."""
|
"""Define API routes for status."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import time
|
import time
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import APIRouter
|
from fastapi import APIRouter
|
||||||
|
|
||||||
from app.config.settings import settings
|
from app.config.settings import settings
|
||||||
|
from app.domain.retrieval import RetrievedChunk
|
||||||
|
from app.services.llm.llm_factory import get_llm_client
|
||||||
from app.shared.bootstrap import (
|
from app.shared.bootstrap import (
|
||||||
get_bm25_retriever,
|
get_bm25_retriever,
|
||||||
get_binary_store,
|
get_binary_store,
|
||||||
get_conversation_store,
|
get_conversation_store,
|
||||||
get_document_query_service,
|
get_document_query_service,
|
||||||
|
get_embedding_provider,
|
||||||
|
get_reranker,
|
||||||
get_vector_index,
|
get_vector_index,
|
||||||
)
|
)
|
||||||
|
from app.shared.model_usage_tracker import get_model_usage_tracker
|
||||||
|
|
||||||
router = APIRouter(prefix="/status", tags=["系统状态"])
|
router = APIRouter(prefix="/status", tags=["系统状态"])
|
||||||
|
|
||||||
@@ -23,6 +29,16 @@ _stats_cache: dict[str, Any] = {}
|
|||||||
_stats_cache_time: float = 0.0
|
_stats_cache_time: float = 0.0
|
||||||
_STATS_TTL_SECONDS: float = 10.0
|
_STATS_TTL_SECONDS: float = 10.0
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# AI model roles surfaced on the Status page (Task: System Status AI models)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
_MODEL_ROLES: dict[str, str] = {
|
||||||
|
"main_llm": "主问答 LLM",
|
||||||
|
"hyde_llm": "HyDE 查询增强",
|
||||||
|
"embedding": "Embedding",
|
||||||
|
"reranker": "Reranker",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@router.get("/stats")
|
@router.get("/stats")
|
||||||
async def get_stats():
|
async def get_stats():
|
||||||
@@ -111,3 +127,111 @@ async def get_health():
|
|||||||
"max": settings.session_max_sessions,
|
"max": settings.session_max_sessions,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_role_provider_model(role: str) -> tuple[str, str]:
|
||||||
|
"""Return the (provider, model) pair currently configured for one AI model role.
|
||||||
|
|
||||||
|
For "hyde_llm" this mirrors the exact fallback logic already used in
|
||||||
|
hyde_expander.py (settings.hyde_llm_provider or settings.llm_provider, same
|
||||||
|
for model) so tracker lookups here always match what TrackedLLMClient
|
||||||
|
recorded when HyDE actually ran.
|
||||||
|
"""
|
||||||
|
if role == "main_llm":
|
||||||
|
return settings.llm_provider, settings.llm_model
|
||||||
|
if role == "hyde_llm":
|
||||||
|
return (
|
||||||
|
settings.hyde_llm_provider or settings.llm_provider,
|
||||||
|
settings.hyde_llm_model or settings.llm_model,
|
||||||
|
)
|
||||||
|
if role == "embedding":
|
||||||
|
return "embedding", settings.embedding_model
|
||||||
|
if role == "reranker":
|
||||||
|
return "reranker", settings.reranker_model
|
||||||
|
raise ValueError(f"unknown model role: {role}") # pragma: no cover - internal roles are fixed
|
||||||
|
|
||||||
|
|
||||||
|
def _build_model_status(role: str) -> dict[str, Any]:
|
||||||
|
"""Build one /status/models row for the given role from tracker data + live settings."""
|
||||||
|
provider, model = _resolve_role_provider_model(role)
|
||||||
|
entry = get_model_usage_tracker().get(provider, model)
|
||||||
|
|
||||||
|
main_provider, main_model = _resolve_role_provider_model("main_llm")
|
||||||
|
shares_usage_with = (
|
||||||
|
"main_llm" if role != "main_llm" and (provider, model) == (main_provider, main_model) else None
|
||||||
|
)
|
||||||
|
|
||||||
|
enabled = True
|
||||||
|
status = entry.status if entry else "never_called"
|
||||||
|
if role == "reranker":
|
||||||
|
enabled = settings.reranker_enabled
|
||||||
|
if not enabled:
|
||||||
|
# Config always wins: report "disabled" even if the reranker was
|
||||||
|
# enabled and called successfully earlier in this process's life.
|
||||||
|
status = "disabled"
|
||||||
|
|
||||||
|
return {
|
||||||
|
"role": role,
|
||||||
|
"role_label": _MODEL_ROLES[role],
|
||||||
|
"provider": provider,
|
||||||
|
"model": model,
|
||||||
|
"enabled": enabled,
|
||||||
|
"status": status,
|
||||||
|
"total_tokens": entry.total_tokens if entry else 0,
|
||||||
|
"call_count_ok": entry.call_count_ok if entry else 0,
|
||||||
|
"call_count_error": entry.call_count_error if entry else 0,
|
||||||
|
"last_called_at": entry.last_called_at.isoformat() if entry and entry.last_called_at else None,
|
||||||
|
"last_latency_ms": entry.last_latency_ms if entry else None,
|
||||||
|
"last_error": entry.last_error if entry else None,
|
||||||
|
"shares_usage_with": shares_usage_with,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/models")
|
||||||
|
async def get_model_statuses():
|
||||||
|
"""Return connection status + cumulative token usage for all 4 tracked AI model roles.
|
||||||
|
|
||||||
|
Passive: reads tracker state + settings only, makes no outbound network calls.
|
||||||
|
"""
|
||||||
|
return {"models": [_build_model_status(role) for role in _MODEL_ROLES]}
|
||||||
|
|
||||||
|
|
||||||
|
async def _ping_main_or_hyde(role: str) -> None:
|
||||||
|
"""Send one minimal chat completion to the LLM configured for `role`."""
|
||||||
|
provider, model = _resolve_role_provider_model(role)
|
||||||
|
client = get_llm_client(provider=provider, model=model)
|
||||||
|
await asyncio.to_thread(client.chat, [{"role": "user", "content": "ping"}], max_tokens=1)
|
||||||
|
|
||||||
|
|
||||||
|
async def _ping_embedding() -> None:
|
||||||
|
"""Send one minimal embedding request."""
|
||||||
|
await asyncio.to_thread(get_embedding_provider().embed_query, "ping")
|
||||||
|
|
||||||
|
|
||||||
|
async def _ping_reranker() -> None:
|
||||||
|
"""Send one minimal rerank request, only when the reranker is enabled."""
|
||||||
|
reranker = get_reranker()
|
||||||
|
if reranker is None:
|
||||||
|
return
|
||||||
|
# Minimal single-chunk probe — real content doesn't matter, only round-trip success.
|
||||||
|
placeholder = RetrievedChunk(chunk_id="ping", doc_id="ping", doc_title="ping", text="ping", score=0.0)
|
||||||
|
await asyncio.to_thread(reranker.rerank, "ping", [placeholder], 1)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/models/ping")
|
||||||
|
async def ping_model_connections():
|
||||||
|
"""Actively test each configured model with a minimal request, then return fresh statuses.
|
||||||
|
|
||||||
|
Each ping is isolated with return_exceptions=True so one model timing out
|
||||||
|
or erroring does not prevent the other three from completing and being
|
||||||
|
reported. Failures are still visible afterwards via _build_model_status()
|
||||||
|
because the underlying clients record their own outcome into the tracker.
|
||||||
|
"""
|
||||||
|
tasks = [
|
||||||
|
_ping_main_or_hyde("main_llm"),
|
||||||
|
_ping_main_or_hyde("hyde_llm"),
|
||||||
|
_ping_embedding(),
|
||||||
|
_ping_reranker(),
|
||||||
|
]
|
||||||
|
await asyncio.gather(*tasks, return_exceptions=True)
|
||||||
|
return {"models": [_build_model_status(role) for role in _MODEL_ROLES]}
|
||||||
|
|||||||
@@ -0,0 +1,92 @@
|
|||||||
|
"""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_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()
|
||||||
Reference in New Issue
Block a user