feat: add TrackedLLMClient decorator for transparent usage recording
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
@@ -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)
|
||||||
@@ -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
|
||||||
Reference in New Issue
Block a user