Implement manual generator driving using next()/StopIteration to capture the return value (trailing usage dict) from inner stream_chat() implementations, enabling token tracking for streaming LLM calls. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
106 lines
4.1 KiB
Python
106 lines
4.1 KiB
Python
"""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."""
|
|
# Use MagicMock to avoid requiring a real LLM provider implementation (e.g., DeepseekClient);
|
|
# tests focus on TrackedLLMClient's delegation and recording behavior, not provider logic.
|
|
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
|
|
|
|
|
|
def test_stream_chat_records_usage_from_generator_return_value():
|
|
"""stream_chat() must forward the inner generator's returned usage dict to record()."""
|
|
inner = _make_inner()
|
|
|
|
def fake_stream(*args, **kwargs):
|
|
yield "chunk-1"
|
|
yield "chunk-2"
|
|
return {"prompt_tokens": 6, "completion_tokens": 2, "total_tokens": 8}
|
|
|
|
inner.stream_chat.side_effect = fake_stream
|
|
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.total_tokens == 8
|
|
assert entry.call_count_ok == 1
|