main-ruqi #1
@@ -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()
|
||||||
Reference in New Issue
Block a user