diff --git a/backend/app/shared/model_usage_tracker.py b/backend/app/shared/model_usage_tracker.py new file mode 100644 index 0000000..9b4ce9c --- /dev/null +++ b/backend/app/shared/model_usage_tracker.py @@ -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() diff --git a/backend/tests/observability/__init__.py b/backend/tests/observability/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/backend/tests/observability/test_model_usage_tracker.py b/backend/tests/observability/test_model_usage_tracker.py new file mode 100644 index 0000000..d09db11 --- /dev/null +++ b/backend/tests/observability/test_model_usage_tracker.py @@ -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()