"""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