main-ruqi #1

Merged
wangwei merged 17 commits from main-ruqi into main 2026-07-02 22:05:17 +08:00
2 changed files with 216 additions and 0 deletions
Showing only changes of commit 169911ab46 - Show all commits
+124
View File
@@ -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]}
+92
View File
@@ -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()