feat: add ModelUsageTracker for per-model token/connection tracking

This commit is contained in:
wangwei
2026-07-02 14:41:21 +08:00
parent 4b451ef97c
commit 74f327c85e
3 changed files with 192 additions and 0 deletions
+113
View File
@@ -0,0 +1,113 @@
"""In-memory registry that tracks per-model call outcomes and token usage.
This module lives in `app/shared` — the same cross-cutting-support tier as
`bootstrap.py` — because it is not business logic: it exists purely so the
System Status page can show which AI models (main LLM, HyDE LLM, embedding,
reranker) are configured, whether their most recent call succeeded, and how
many tokens they have consumed since this process started. Tracking here
must never disrupt a real user-facing call: every public method swallows its
own exceptions and logs a warning instead of raising.
"""
from __future__ import annotations
import threading
from dataclasses import dataclass
from datetime import datetime, timezone
from functools import lru_cache
from loguru import logger
@dataclass
class ModelUsageEntry:
"""Represent accumulated usage/connection state for one provider+model pair."""
provider: str
model: str
total_tokens: int = 0
prompt_tokens: int = 0
completion_tokens: int = 0
call_count_ok: int = 0
call_count_error: int = 0
last_called_at: datetime | None = None
last_latency_ms: int | None = None
last_error: str | None = None
@property
def status(self) -> str:
"""Derive never_called/ok/error from call history.
The "disabled" status (reranker only, when turned off in settings) is
NOT decided here: this dataclass has no access to live settings. The
API route layer (Task 6) applies that override on top of this value,
so config always wins over stale historical data.
"""
if self.last_called_at is None:
return "never_called"
return "error" if self.last_error else "ok"
class ModelUsageTracker:
"""Thread-safe in-memory registry of per-model call/usage stats.
Keyed by "{provider}:{model}" rather than by business role (main LLM /
HyDE / embedding / reranker) so that any future call site is captured
automatically, even before anyone teaches this class about its role.
"""
def __init__(self) -> None:
"""Initialize an empty registry guarded by a single lock."""
self._entries: dict[str, ModelUsageEntry] = {}
# One coarse lock is enough: record() runs at most a few times per
# request, and snapshot() is only read by the low-traffic status page.
self._lock = threading.Lock()
def record(
self,
*,
provider: str,
model: str,
success: bool,
usage: dict | None = None,
latency_ms: int | None = None,
error: str | None = None,
) -> None:
"""Record the outcome of one call to provider/model.
Never raises: any internal failure is logged and swallowed so a bug
in observability code cannot break a real LLM/embedding/reranker call.
"""
try:
key = f"{provider}:{model}"
usage = usage if isinstance(usage, dict) else {}
with self._lock:
entry = self._entries.setdefault(key, ModelUsageEntry(provider=provider, model=model))
entry.total_tokens += int(usage.get("total_tokens", 0) or 0)
entry.prompt_tokens += int(usage.get("prompt_tokens", 0) or 0)
entry.completion_tokens += int(usage.get("completion_tokens", 0) or 0)
if success:
entry.call_count_ok += 1
entry.last_error = None
else:
entry.call_count_error += 1
entry.last_error = error or "unknown error"
entry.last_called_at = datetime.now(timezone.utc)
entry.last_latency_ms = latency_ms
except Exception as exc: # noqa: BLE001 - tracking must never break a real call
logger.warning("ModelUsageTracker.record failed for {}:{} - {}", provider, model, exc)
def snapshot(self) -> dict[str, ModelUsageEntry]:
"""Return a shallow copy of all tracked entries, safe to mutate by the caller."""
with self._lock:
return dict(self._entries)
def get(self, provider: str, model: str) -> ModelUsageEntry | None:
"""Return the entry for one provider/model pair, or None if never recorded."""
return self.snapshot().get(f"{provider}:{model}")
@lru_cache
def get_model_usage_tracker() -> ModelUsageTracker:
"""Return the process-wide singleton tracker (mirrors get_settings()/get_llm_factory())."""
return ModelUsageTracker()
@@ -0,0 +1,79 @@
"""Unit tests for ModelUsageTracker — no mocking needed, pure in-memory state."""
from __future__ import annotations
from app.shared.model_usage_tracker import ModelUsageEntry, ModelUsageTracker, get_model_usage_tracker
def test_never_called_model_has_no_entry():
"""A tracker that has never recorded a call returns None from get()."""
tracker = ModelUsageTracker()
assert tracker.get("deepseek", "deepseek-v4-flash") is None
def test_record_success_accumulates_tokens_and_calls():
"""Two successful calls accumulate tokens and call_count_ok."""
tracker = ModelUsageTracker()
tracker.record(
provider="deepseek", model="deepseek-v4-flash", success=True,
usage={"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, latency_ms=100,
)
tracker.record(
provider="deepseek", model="deepseek-v4-flash", success=True,
usage={"prompt_tokens": 20, "completion_tokens": 8, "total_tokens": 28}, latency_ms=200,
)
entry = tracker.get("deepseek", "deepseek-v4-flash")
assert entry is not None
assert entry.total_tokens == 43
assert entry.prompt_tokens == 30
assert entry.completion_tokens == 13
assert entry.call_count_ok == 2
assert entry.call_count_error == 0
assert entry.status == "ok"
assert entry.last_latency_ms == 200
def test_record_error_sets_error_status_without_losing_prior_tokens():
"""A failed call after successful ones flips status to 'error' but keeps accumulated tokens."""
tracker = ModelUsageTracker()
tracker.record(provider="qwen", model="qwen3.5-flash", success=True, usage={"total_tokens": 50}, latency_ms=50)
tracker.record(provider="qwen", model="qwen3.5-flash", success=False, error="HTTP 500", latency_ms=30)
entry = tracker.get("qwen", "qwen3.5-flash")
assert entry.total_tokens == 50
assert entry.call_count_ok == 1
assert entry.call_count_error == 1
assert entry.status == "error"
assert entry.last_error == "HTTP 500"
def test_record_success_after_error_clears_last_error():
"""A later successful call clears last_error and status returns to 'ok'."""
tracker = ModelUsageTracker()
tracker.record(provider="qwen", model="qwen3.5-flash", success=False, error="timeout", latency_ms=30)
tracker.record(provider="qwen", model="qwen3.5-flash", success=True, usage={"total_tokens": 5}, latency_ms=40)
entry = tracker.get("qwen", "qwen3.5-flash")
assert entry.status == "ok"
assert entry.last_error is None
def test_record_never_raises_on_bad_usage_dict():
"""A malformed usage value (wrong type) is swallowed, not raised, and does not corrupt other entries."""
tracker = ModelUsageTracker()
tracker.record(provider="embedding", model="text-embedding-v3", success=True, usage="not-a-dict", latency_ms=10) # type: ignore[arg-type]
# Must not raise, and must not have created a corrupted entry that breaks snapshot().
snapshot = tracker.snapshot()
assert isinstance(snapshot, dict)
def test_snapshot_returns_independent_copy():
"""snapshot() returns a dict that can be safely mutated without affecting the tracker."""
tracker = ModelUsageTracker()
tracker.record(provider="deepseek", model="deepseek-v4-flash", success=True, usage={"total_tokens": 1}, latency_ms=1)
snap = tracker.snapshot()
snap.clear()
assert tracker.get("deepseek", "deepseek-v4-flash") is not None
def test_get_model_usage_tracker_returns_singleton():
"""get_model_usage_tracker() always returns the same process-wide instance."""
assert get_model_usage_tracker() is get_model_usage_tracker()