From 37ea27fcbe5d55f71b08b683e96011a23a65b336 Mon Sep 17 00:00:00 2001 From: wangwei Date: Thu, 2 Jul 2026 14:49:06 +0800 Subject: [PATCH] feat: add TrackedLLMClient decorator for transparent usage recording Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- backend/app/services/llm/tracked_client.py | 86 +++++++++++++++++++ .../observability/test_tracked_client.py | 83 ++++++++++++++++++ 2 files changed, 169 insertions(+) create mode 100644 backend/app/services/llm/tracked_client.py create mode 100644 backend/tests/observability/test_tracked_client.py diff --git a/backend/app/services/llm/tracked_client.py b/backend/app/services/llm/tracked_client.py new file mode 100644 index 0000000..1cc249f --- /dev/null +++ b/backend/app/services/llm/tracked_client.py @@ -0,0 +1,86 @@ +"""Transparent decorator around BaseLLMClient implementations. + +Records per-call token usage, latency, and success/failure into a +ModelUsageTracker without changing any caller-visible behavior. +""" + +from __future__ import annotations + +import time +from typing import Any, Dict, List, Optional + +from app.shared.model_usage_tracker import ModelUsageTracker + +from .base_client import BaseLLMClient, LLMResponse +from .tool_types import Tool + + +class TrackedLLMClient: + """Wrap any BaseLLMClient and record its usage into a ModelUsageTracker. + + Deliberately does NOT subclass BaseLLMClient: that ABC declares abstract + methods (_init_client, get_available_models) with no meaningful override + here, and subclassing would make Python refuse to instantiate this class + ("Can't instantiate abstract class") before __getattr__ ever got a chance + to forward the call. Plain composition + __getattr__ delegation works + because every caller in this codebase only ever uses duck-typed access: + .chat(), .stream_chat(), .get_available_models(), .close(), .config. + """ + + def __init__(self, inner: BaseLLMClient, tracker: ModelUsageTracker) -> None: + """Store the wrapped client and the tracker to report into.""" + self._inner = inner + self._tracker = tracker + + def chat( + self, + messages: List[Dict[str, str]], + max_tokens: Optional[int] = None, + temperature: Optional[float] = None, + tools: Optional[List[Tool]] = None, + **kwargs: Any, + ) -> LLMResponse: + """Delegate to the wrapped client's chat(), then record the outcome.""" + start = time.time() + response = self._inner.chat(messages, max_tokens, temperature, tools, **kwargs) + # Key by the *configured* model, not response.model, so lookups driven + # by settings (llm_model / hyde_llm_model) always match what we recorded. + self._tracker.record( + provider=self._inner.config.provider.value, + model=self._inner.config.model, + success=response.is_success, + usage=response.usage, + latency_ms=int((time.time() - start) * 1000), + error=response.error, + ) + return response + + def stream_chat(self, messages: List[Dict[str, str]], *args: Any, **kwargs: Any): + """Delegate to the wrapped client's stream_chat(), recording call outcome only. + + Token usage is NOT recorded here: none of the current provider + stream_chat() implementations parse a trailing usage chunk from the + gateway (see the design doc's Known Limitations), so accumulating a + token count here would silently be wrong. Only call success/failure + and latency are tracked for streaming calls. + """ + start = time.time() + error: Optional[str] = None + try: + for chunk in self._inner.stream_chat(messages, *args, **kwargs): + yield chunk + except Exception as exc: # noqa: BLE001 - report, then re-raise unchanged + error = str(exc) + raise + finally: + self._tracker.record( + provider=self._inner.config.provider.value, + model=self._inner.config.model, + success=error is None, + latency_ms=int((time.time() - start) * 1000), + error=error, + ) + + def __getattr__(self, name: str) -> Any: + """Forward any other attribute/method access to the wrapped client.""" + return getattr(self._inner, name) diff --git a/backend/tests/observability/test_tracked_client.py b/backend/tests/observability/test_tracked_client.py new file mode 100644 index 0000000..5741f6b --- /dev/null +++ b/backend/tests/observability/test_tracked_client.py @@ -0,0 +1,83 @@ +"""Unit tests for TrackedLLMClient — verifies transparent delegation + recording.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +from app.services.llm.base_client import LLMConfig, LLMProvider, LLMResponse +from app.services.llm.tracked_client import TrackedLLMClient +from app.shared.model_usage_tracker import ModelUsageTracker + + +def _make_inner(model: str = "deepseek-v4-flash") -> MagicMock: + """Build a MagicMock standing in for a concrete BaseLLMClient subclass.""" + inner = MagicMock() + inner.config = LLMConfig( + provider=LLMProvider.DEEPSEEK, model=model, api_key="test-key", base_url="http://example.test/v1", + ) + return inner + + +def test_chat_delegates_and_returns_unchanged_response(): + """chat() must return exactly what the wrapped client returned.""" + inner = _make_inner() + expected = LLMResponse(content="hello", model="deepseek-v4-flash", usage={"total_tokens": 12}) + inner.chat.return_value = expected + tracker = ModelUsageTracker() + + tracked = TrackedLLMClient(inner, tracker) + result = tracked.chat([{"role": "user", "content": "hi"}]) + + assert result is expected + inner.chat.assert_called_once_with([{"role": "user", "content": "hi"}], None, None, None) + + +def test_chat_records_success_and_tokens(): + """A successful chat() call must be recorded under 'deepseek:deepseek-v4-flash'.""" + inner = _make_inner() + inner.chat.return_value = LLMResponse(content="hi", model="deepseek-v4-flash", usage={"total_tokens": 42}) + tracker = ModelUsageTracker() + + TrackedLLMClient(inner, tracker).chat([{"role": "user", "content": "hi"}]) + + entry = tracker.get("deepseek", "deepseek-v4-flash") + assert entry is not None + assert entry.total_tokens == 42 + assert entry.status == "ok" + + +def test_chat_records_error_from_response(): + """A chat() call that returns an error-carrying LLMResponse is recorded as a failure.""" + inner = _make_inner() + inner.chat.return_value = LLMResponse(content="", model="deepseek-v4-flash", error="API error: 500") + tracker = ModelUsageTracker() + + TrackedLLMClient(inner, tracker).chat([{"role": "user", "content": "hi"}]) + + entry = tracker.get("deepseek", "deepseek-v4-flash") + assert entry.status == "error" + assert entry.last_error == "API error: 500" + + +def test_getattr_forwards_to_inner_client(): + """Attributes not defined on TrackedLLMClient must forward to the wrapped client.""" + inner = _make_inner() + inner.get_available_models.return_value = ["deepseek-v4-flash"] + tracked = TrackedLLMClient(inner, ModelUsageTracker()) + + assert tracked.get_available_models() == ["deepseek-v4-flash"] + assert tracked.config is inner.config + + +def test_stream_chat_records_call_without_token_usage(): + """stream_chat() must record a call (latency/success) but not fabricate token counts.""" + inner = _make_inner() + inner.stream_chat.return_value = iter(["chunk-1", "chunk-2"]) + tracker = ModelUsageTracker() + + chunks = list(TrackedLLMClient(inner, tracker).stream_chat([{"role": "user", "content": "hi"}])) + + assert chunks == ["chunk-1", "chunk-2"] + entry = tracker.get("deepseek", "deepseek-v4-flash") + assert entry.call_count_ok == 1 + assert entry.total_tokens == 0