Files
nanobot/tests/providers/test_llm_usage_observer.py
T

181 lines
5.5 KiB
Python

from __future__ import annotations
import asyncio
from collections.abc import Iterator
from types import SimpleNamespace
import pytest
from nanobot.llm_usage.context import llm_usage_source
from nanobot.llm_usage.models import LLMCallRecord
from nanobot.providers.base import LLMProvider, LLMResponse, LLMUsage
from nanobot.providers.fallback_provider import FallbackProvider
class _SequenceProvider(LLMProvider):
_CHAT_RETRY_DELAYS = (0,)
def __init__(self, responses: Iterator[LLMResponse]) -> None:
super().__init__(provider_name="test-provider")
self._responses = responses
async def chat(self, **_kwargs: object) -> LLMResponse:
return next(self._responses)
def get_default_model(self) -> str:
return "test-model"
class _NoRetryProvider(_SequenceProvider):
_CHAT_RETRY_DELAYS = ()
class _BlockingProvider(LLMProvider):
async def chat(self, **_kwargs: object) -> LLMResponse:
await asyncio.Event().wait()
raise AssertionError("unreachable")
async def chat_stream(self, **_kwargs: object) -> LLMResponse:
await asyncio.Event().wait()
raise AssertionError("unreachable")
def get_default_model(self) -> str:
return "test-model"
@pytest.mark.asyncio
async def test_observer_receives_every_retry_attempt() -> None:
provider = _SequenceProvider(
iter(
[
LLMResponse(
content="temporary failure",
finish_reason="error",
error_kind="timeout",
),
LLMResponse(
content="ok",
usage=LLMUsage.reported(
input_tokens=100,
output_tokens=20,
cache_read_tokens=60,
),
),
]
)
)
events: list[LLMCallRecord] = []
provider.set_llm_call_observer(events.append)
with llm_usage_source("api"):
response = await provider.chat_with_retry(
messages=[{"role": "user", "content": "hello"}],
model="selected-model",
)
assert response.finish_reason == "stop"
assert len(events) == 2
assert [event.finish_reason for event in events] == ["error", "stop"]
assert all(event.provider == "test-provider" for event in events)
assert all(event.model == "selected-model" for event in events)
assert all(event.source == "api" for event in events)
assert events[1].usage is not None
assert events[1].usage.cache_read_tokens == 60
@pytest.mark.asyncio
async def test_observer_estimates_missing_success_usage_without_storing_content() -> None:
provider = _SequenceProvider(iter([LLMResponse(content="hello")]))
events: list[LLMCallRecord] = []
provider.set_llm_call_observer(events.append)
response = await provider.chat_with_retry(
messages=[{"role": "user", "content": "hello"}],
)
assert response.usage is not None
assert response.usage.source == "estimated"
assert events[0].usage == response.usage
assert "content" not in LLMCallRecord.__dataclass_fields__
@pytest.mark.asyncio
async def test_observer_failure_never_breaks_provider_call() -> None:
provider = _SequenceProvider(iter([LLMResponse(content="ok")]))
def _fail(_event: LLMCallRecord) -> None:
raise RuntimeError("disk unavailable")
provider.set_llm_call_observer(_fail)
response = await provider.chat_with_retry(
messages=[{"role": "user", "content": "hello"}],
)
assert response.content == "ok"
@pytest.mark.asyncio
@pytest.mark.parametrize("stream", [False, True])
async def test_observer_records_cancelled_provider_attempt(stream: bool) -> None:
provider = _BlockingProvider(provider_name="blocking-provider")
events: list[LLMCallRecord] = []
provider.set_llm_call_observer(events.append)
call = provider.chat_stream_with_retry if stream else provider.chat_with_retry
with pytest.raises(asyncio.TimeoutError):
await asyncio.wait_for(
call(messages=[{"role": "user", "content": "hello"}]),
timeout=0.01,
)
assert len(events) == 1
assert events[0].finish_reason == "cancelled"
assert events[0].error_kind == "cancelled"
assert events[0].usage is None
@pytest.mark.asyncio
async def test_fallback_provider_propagates_observer_to_every_leaf() -> None:
primary = _NoRetryProvider(
iter(
[
LLMResponse(
content="primary unavailable",
finish_reason="error",
error_kind="timeout",
)
]
)
)
fallback = _SequenceProvider(
iter(
[
LLMResponse(
content="fallback ok",
usage=LLMUsage.reported(input_tokens=12, output_tokens=3),
)
]
)
)
preset = SimpleNamespace(
model="fallback-model",
max_tokens=256,
temperature=0.2,
reasoning_effort=None,
context_window_tokens=4_096,
)
provider = FallbackProvider(primary, [preset], lambda _preset: fallback)
events: list[LLMCallRecord] = []
provider.set_llm_call_observer(events.append)
response = await provider.chat_with_retry(
messages=[{"role": "user", "content": "hello"}],
model="primary-model",
)
assert response.content == "fallback ok"
assert [(event.model, event.finish_reason) for event in events] == [
("primary-model", "error"),
("fallback-model", "stop"),
]