mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-09-04 02:01:48 +03:00
refactor(providers): define typed usage contract
This commit is contained in:
@@ -380,7 +380,8 @@ async def test_chat_success():
|
||||
assert isinstance(result, LLMResponse)
|
||||
assert result.content == "Hello!"
|
||||
assert result.finish_reason == "stop"
|
||||
assert result.usage["prompt_tokens"] == 10
|
||||
assert result.usage is not None
|
||||
assert result.usage.input_tokens == 10
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -229,8 +229,11 @@ def test_parse_response_maps_text_tools_reasoning_usage_and_stop_reason() -> Non
|
||||
|
||||
assert result.content == "hello"
|
||||
assert result.finish_reason == "tool_calls"
|
||||
assert result.usage["prompt_tokens"] == 10
|
||||
assert result.usage["cached_tokens"] == 2
|
||||
assert result.usage is not None
|
||||
assert result.usage.input_tokens == 12
|
||||
assert result.usage.output_tokens == 5
|
||||
assert result.usage.cache_read_tokens == 2
|
||||
assert result.usage.cache_write_tokens is None
|
||||
assert result.reasoning_content == "think"
|
||||
assert result.thinking_blocks == [{"type": "thinking", "thinking": "think", "signature": "sig"}]
|
||||
assert result.tool_calls[0].id == "t1"
|
||||
@@ -276,7 +279,10 @@ async def test_chat_stream_aggregates_text_tool_use_and_usage() -> None:
|
||||
assert deltas == ["he", "llo"]
|
||||
assert result.content == "hello"
|
||||
assert result.finish_reason == "tool_calls"
|
||||
assert result.usage == {"prompt_tokens": 3, "completion_tokens": 4, "total_tokens": 7}
|
||||
assert result.usage is not None
|
||||
assert result.usage.input_tokens == 3
|
||||
assert result.usage.output_tokens == 4
|
||||
assert result.usage.total_tokens == 7
|
||||
assert result.tool_calls[0].name == "search"
|
||||
assert result.tool_calls[0].arguments == {"q": "x"}
|
||||
|
||||
@@ -285,6 +291,48 @@ async def _append_delta(deltas: list[str], text: str) -> None:
|
||||
deltas.append(text)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("wire_usage", "expected_read", "expected_write", "expected_input"),
|
||||
[
|
||||
({"inputTokens": 5, "outputTokens": 1}, None, None, 5),
|
||||
(
|
||||
{
|
||||
"inputTokens": 5,
|
||||
"outputTokens": 1,
|
||||
"cacheReadInputTokens": 0,
|
||||
"cacheWriteInputTokens": 0,
|
||||
},
|
||||
0,
|
||||
0,
|
||||
5,
|
||||
),
|
||||
(
|
||||
{
|
||||
"inputTokens": 5,
|
||||
"outputTokens": 1,
|
||||
"cacheReadInputTokens": 7,
|
||||
"cacheWriteInputTokens": 3,
|
||||
},
|
||||
7,
|
||||
3,
|
||||
15,
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_bedrock_usage_preserves_cache_reporting_and_logical_input(
|
||||
wire_usage: dict[str, int],
|
||||
expected_read: int | None,
|
||||
expected_write: int | None,
|
||||
expected_input: int,
|
||||
) -> None:
|
||||
usage = BedrockProvider._usage(wire_usage)
|
||||
|
||||
assert usage is not None
|
||||
assert usage.cache_read_tokens == expected_read
|
||||
assert usage.cache_write_tokens == expected_write
|
||||
assert usage.input_tokens == expected_input
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_error_maps_retry_metadata() -> None:
|
||||
provider = BedrockProvider(region="us-east-1", client=FakeClient(error=FakeBedrockError()))
|
||||
|
||||
@@ -14,8 +14,9 @@ class FakeUsage:
|
||||
|
||||
class FakePromptDetails:
|
||||
"""Mimics prompt_tokens_details sub-object."""
|
||||
def __init__(self, cached_tokens=0):
|
||||
def __init__(self, cached_tokens=0, cache_write_tokens=None):
|
||||
self.cached_tokens = cached_tokens
|
||||
self.cache_write_tokens = cache_write_tokens
|
||||
|
||||
|
||||
class _FakeSpec:
|
||||
@@ -62,8 +63,9 @@ def test_extract_usage_openai_cached_tokens_dict():
|
||||
}
|
||||
}
|
||||
result = p._parse(response)
|
||||
assert result.usage["cached_tokens"] == 1200
|
||||
assert result.usage["prompt_tokens"] == 2000
|
||||
assert result.usage is not None
|
||||
assert result.usage.cache_read_tokens == 1200
|
||||
assert result.usage.input_tokens == 2000
|
||||
|
||||
|
||||
def test_extract_usage_deepseek_cached_tokens_dict():
|
||||
@@ -80,11 +82,12 @@ def test_extract_usage_deepseek_cached_tokens_dict():
|
||||
}
|
||||
}
|
||||
result = p._parse(response)
|
||||
assert result.usage["cached_tokens"] == 1200
|
||||
assert result.usage is not None
|
||||
assert result.usage.cache_read_tokens == 1200
|
||||
|
||||
|
||||
def test_extract_usage_no_cached_tokens_dict():
|
||||
"""Response without any cache fields -> no cached_tokens key."""
|
||||
"""Response without any cache fields preserves an unreported cache count."""
|
||||
p = _provider()
|
||||
response = {
|
||||
"choices": [_DICT_CHOICE],
|
||||
@@ -95,11 +98,13 @@ def test_extract_usage_no_cached_tokens_dict():
|
||||
}
|
||||
}
|
||||
result = p._parse(response)
|
||||
assert "cached_tokens" not in result.usage
|
||||
assert result.usage is not None
|
||||
assert result.usage.cache_read_tokens is None
|
||||
assert result.usage.cache_write_tokens is None
|
||||
|
||||
|
||||
def test_extract_usage_openai_cached_zero_dict():
|
||||
"""cached_tokens=0 should NOT be included (same as existing fields)."""
|
||||
"""cached_tokens=0 remains distinct from an unreported cache count."""
|
||||
p = _provider()
|
||||
response = {
|
||||
"choices": [_DICT_CHOICE],
|
||||
@@ -107,11 +112,42 @@ def test_extract_usage_openai_cached_zero_dict():
|
||||
"prompt_tokens": 2000,
|
||||
"completion_tokens": 300,
|
||||
"total_tokens": 2300,
|
||||
"prompt_tokens_details": {"cached_tokens": 0},
|
||||
"prompt_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 0},
|
||||
}
|
||||
}
|
||||
result = p._parse(response)
|
||||
assert "cached_tokens" not in result.usage
|
||||
assert result.usage is not None
|
||||
assert result.usage.cache_read_tokens == 0
|
||||
assert result.usage.cache_write_tokens == 0
|
||||
|
||||
|
||||
def test_extract_usage_preserves_reported_total_and_cache_write_dict():
|
||||
response = {
|
||||
"choices": [_DICT_CHOICE],
|
||||
"usage": {
|
||||
"prompt_tokens": 15,
|
||||
"completion_tokens": 18,
|
||||
"total_tokens": 175,
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 0,
|
||||
"cache_write_tokens": 7,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result = _provider()._parse(response)
|
||||
|
||||
assert result.usage is not None
|
||||
assert result.usage.total_tokens == 175
|
||||
assert result.usage.reported_tokens == 175
|
||||
assert result.usage.cache_read_tokens == 0
|
||||
assert result.usage.cache_write_tokens == 7
|
||||
|
||||
|
||||
def test_extract_usage_missing_is_none():
|
||||
result = _provider()._parse({"choices": [_DICT_CHOICE]})
|
||||
|
||||
assert result.usage is None
|
||||
|
||||
|
||||
# --- object-based response (OpenAI SDK Pydantic model) ---
|
||||
@@ -127,7 +163,29 @@ def test_extract_usage_openai_cached_tokens_obj():
|
||||
)
|
||||
response = FakeUsage(choices=[_FakeChoice()], usage=usage_obj)
|
||||
result = p._parse(response)
|
||||
assert result.usage["cached_tokens"] == 1200
|
||||
assert result.usage is not None
|
||||
assert result.usage.cache_read_tokens == 1200
|
||||
|
||||
|
||||
def test_extract_usage_preserves_reported_total_and_cache_write_obj():
|
||||
usage_obj = FakeUsage(
|
||||
prompt_tokens=15,
|
||||
completion_tokens=18,
|
||||
total_tokens=175,
|
||||
prompt_tokens_details=FakePromptDetails(
|
||||
cached_tokens=0,
|
||||
cache_write_tokens=7,
|
||||
),
|
||||
)
|
||||
response = FakeUsage(choices=[_FakeChoice()], usage=usage_obj)
|
||||
|
||||
result = _provider()._parse(response)
|
||||
|
||||
assert result.usage is not None
|
||||
assert result.usage.total_tokens == 175
|
||||
assert result.usage.reported_tokens == 175
|
||||
assert result.usage.cache_read_tokens == 0
|
||||
assert result.usage.cache_write_tokens == 7
|
||||
|
||||
|
||||
def test_extract_usage_deepseek_cached_tokens_obj():
|
||||
@@ -141,7 +199,8 @@ def test_extract_usage_deepseek_cached_tokens_obj():
|
||||
)
|
||||
response = FakeUsage(choices=[_FakeChoice()], usage=usage_obj)
|
||||
result = p._parse(response)
|
||||
assert result.usage["cached_tokens"] == 1200
|
||||
assert result.usage is not None
|
||||
assert result.usage.cache_read_tokens == 1200
|
||||
|
||||
|
||||
def test_extract_usage_stepfun_top_level_cached_tokens_dict():
|
||||
@@ -157,7 +216,8 @@ def test_extract_usage_stepfun_top_level_cached_tokens_dict():
|
||||
}
|
||||
}
|
||||
result = p._parse(response)
|
||||
assert result.usage["cached_tokens"] == 512
|
||||
assert result.usage is not None
|
||||
assert result.usage.cache_read_tokens == 512
|
||||
|
||||
|
||||
def test_extract_usage_stepfun_top_level_cached_tokens_obj():
|
||||
@@ -171,7 +231,8 @@ def test_extract_usage_stepfun_top_level_cached_tokens_obj():
|
||||
)
|
||||
response = FakeUsage(choices=[_FakeChoice()], usage=usage_obj)
|
||||
result = p._parse(response)
|
||||
assert result.usage["cached_tokens"] == 512
|
||||
assert result.usage is not None
|
||||
assert result.usage.cache_read_tokens == 512
|
||||
|
||||
|
||||
def test_extract_usage_priority_nested_over_top_level_dict():
|
||||
@@ -188,11 +249,12 @@ def test_extract_usage_priority_nested_over_top_level_dict():
|
||||
}
|
||||
}
|
||||
result = p._parse(response)
|
||||
assert result.usage["cached_tokens"] == 100
|
||||
assert result.usage is not None
|
||||
assert result.usage.cache_read_tokens == 100
|
||||
|
||||
|
||||
def test_anthropic_maps_cache_fields_to_cached_tokens():
|
||||
"""Anthropic's cache_read_input_tokens should map to cached_tokens."""
|
||||
def test_anthropic_adds_native_cache_fields_to_logical_input():
|
||||
"""Anthropic excludes cache reads/writes from its native input_tokens."""
|
||||
from nanobot.providers.anthropic_provider import AnthropicProvider
|
||||
|
||||
usage_obj = FakeUsage(
|
||||
@@ -210,14 +272,15 @@ def test_anthropic_maps_cache_fields_to_cached_tokens():
|
||||
usage=usage_obj,
|
||||
)
|
||||
result = AnthropicProvider._parse_response(response)
|
||||
assert result.usage["cached_tokens"] == 1200
|
||||
assert result.usage["prompt_tokens"] == 2300
|
||||
assert result.usage["total_tokens"] == 2500
|
||||
assert result.usage["cache_creation_input_tokens"] == 300
|
||||
assert result.usage is not None
|
||||
assert result.usage.cache_read_tokens == 1200
|
||||
assert result.usage.cache_write_tokens == 300
|
||||
assert result.usage.input_tokens == 2300
|
||||
assert result.usage.total_tokens == 2500
|
||||
|
||||
|
||||
def test_anthropic_no_cache_fields():
|
||||
"""Anthropic response without cache fields should not have cached_tokens."""
|
||||
"""Anthropic response without cache fields preserves unreported counts."""
|
||||
from nanobot.providers.anthropic_provider import AnthropicProvider
|
||||
|
||||
usage_obj = FakeUsage(input_tokens=800, output_tokens=200)
|
||||
@@ -230,4 +293,7 @@ def test_anthropic_no_cache_fields():
|
||||
usage=usage_obj,
|
||||
)
|
||||
result = AnthropicProvider._parse_response(response)
|
||||
assert "cached_tokens" not in result.usage
|
||||
assert result.usage is not None
|
||||
assert result.usage.input_tokens == 800
|
||||
assert result.usage.cache_read_tokens is None
|
||||
assert result.usage.cache_write_tokens is None
|
||||
|
||||
@@ -46,7 +46,8 @@ def test_custom_provider_parse_accepts_dict_response() -> None:
|
||||
|
||||
assert result.finish_reason == "stop"
|
||||
assert result.content == "hello from dict"
|
||||
assert result.usage["total_tokens"] == 3
|
||||
assert result.usage is not None
|
||||
assert result.usage.total_tokens == 3
|
||||
|
||||
|
||||
def test_custom_provider_parse_normalizes_text_tool_call() -> None:
|
||||
|
||||
@@ -732,11 +732,7 @@ async def test_codex_compacts_state_at_ninety_percent_before_next_request(
|
||||
"content": [{"type": "output_text", "text": "old answer"}],
|
||||
},
|
||||
],
|
||||
usage={
|
||||
"prompt_tokens": 90,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 95,
|
||||
},
|
||||
usage=provider_base.LLMUsage.reported(input_tokens=90, output_tokens=5),
|
||||
)
|
||||
bodies: list[dict[str, Any]] = []
|
||||
|
||||
@@ -772,11 +768,10 @@ async def test_codex_compacts_state_at_ninety_percent_before_next_request(
|
||||
model="gpt-5.6-sol",
|
||||
input_items=body["input"],
|
||||
output_items=[compact_item],
|
||||
usage={
|
||||
"prompt_tokens": 95,
|
||||
"completion_tokens": 2,
|
||||
"total_tokens": 97,
|
||||
},
|
||||
usage=provider_base.LLMUsage.reported(
|
||||
input_tokens=95,
|
||||
output_tokens=2,
|
||||
),
|
||||
),
|
||||
)
|
||||
return provider_base.LLMResponse(content="done")
|
||||
@@ -830,7 +825,7 @@ async def test_codex_disables_unsupported_native_compaction_and_continues(
|
||||
model="gpt-5.6-sol",
|
||||
input_items=[{"type": "message", "role": "user", "content": "old"}],
|
||||
output_items=[{"type": "reasoning", "encrypted_content": "opaque"}],
|
||||
usage={"prompt_tokens": 90, "completion_tokens": 5, "total_tokens": 95},
|
||||
usage=provider_base.LLMUsage.reported(input_tokens=90, output_tokens=5),
|
||||
)
|
||||
bodies: list[dict[str, Any]] = []
|
||||
|
||||
@@ -914,7 +909,7 @@ async def test_codex_stream_surfaces_reasoning_summary(monkeypatch) -> None:
|
||||
return provider_base.LLMResponse(
|
||||
content="answer",
|
||||
finish_reason="stop",
|
||||
usage={"prompt_tokens": 10, "completion_tokens": 5},
|
||||
usage=provider_base.LLMUsage.reported(input_tokens=10, output_tokens=5),
|
||||
reasoning_content="summary",
|
||||
)
|
||||
|
||||
@@ -934,7 +929,7 @@ async def test_codex_stream_surfaces_reasoning_summary(monkeypatch) -> None:
|
||||
assert content_deltas == ["answer"]
|
||||
assert thinking_deltas == ["summary"]
|
||||
assert response.content == "answer"
|
||||
assert response.usage == {"prompt_tokens": 10, "completion_tokens": 5}
|
||||
assert response.usage == provider_base.LLMUsage.reported(input_tokens=10, output_tokens=5)
|
||||
assert response.reasoning_content == "summary"
|
||||
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ from unittest.mock import MagicMock, patch
|
||||
import pytest
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.providers.base import LLMUsage
|
||||
from nanobot.providers.openai_responses.converters import (
|
||||
convert_messages,
|
||||
convert_tools,
|
||||
@@ -484,7 +485,7 @@ class TestParseResponseOutput:
|
||||
result = parse_response_output(resp)
|
||||
assert result.content == "Hello!"
|
||||
assert result.finish_reason == "stop"
|
||||
assert result.usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
|
||||
assert result.usage == LLMUsage.reported(input_tokens=10, output_tokens=5)
|
||||
assert result.tool_calls == []
|
||||
|
||||
def test_refusal_response_surfaces_text_without_advancing_state(self):
|
||||
@@ -652,7 +653,8 @@ class TestParseResponseOutput:
|
||||
}
|
||||
result = parse_response_output(mock)
|
||||
assert result.content == "sdk"
|
||||
assert result.usage["prompt_tokens"] == 1
|
||||
assert result.usage is not None
|
||||
assert result.usage.input_tokens == 1
|
||||
|
||||
def test_usage_maps_responses_api_keys(self):
|
||||
"""Responses API uses input_tokens/output_tokens, not prompt_tokens/completion_tokens."""
|
||||
@@ -662,9 +664,20 @@ class TestParseResponseOutput:
|
||||
"usage": {"input_tokens": 100, "output_tokens": 50, "total_tokens": 150},
|
||||
}
|
||||
result = parse_response_output(resp)
|
||||
assert result.usage["prompt_tokens"] == 100
|
||||
assert result.usage["completion_tokens"] == 50
|
||||
assert result.usage["total_tokens"] == 150
|
||||
assert result.usage == LLMUsage.reported(input_tokens=100, output_tokens=50)
|
||||
|
||||
def test_non_stream_preserves_provider_reported_total(self):
|
||||
result = parse_response_output({
|
||||
"output": [],
|
||||
"status": "completed",
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 999},
|
||||
})
|
||||
|
||||
assert result.usage == LLMUsage.reported(
|
||||
input_tokens=10,
|
||||
output_tokens=5,
|
||||
total_tokens=999,
|
||||
)
|
||||
|
||||
def test_preserves_every_output_item_as_opaque_state(self):
|
||||
input_items = [{"role": "user", "content": "inspect the repo"}]
|
||||
@@ -713,18 +726,18 @@ class TestResponsesConversationState:
|
||||
{"type": "compaction", "encrypted_content": "compact"},
|
||||
{"type": "message", "role": "assistant", "content": "new"},
|
||||
],
|
||||
usage={
|
||||
"prompt_tokens": 90,
|
||||
"completion_tokens": 10,
|
||||
"total_tokens": 100,
|
||||
},
|
||||
usage=LLMUsage.reported(
|
||||
input_tokens=90,
|
||||
output_tokens=10,
|
||||
total_tokens=175,
|
||||
),
|
||||
)
|
||||
|
||||
assert responses_state_items(state) == [
|
||||
{"type": "compaction", "encrypted_content": "compact"},
|
||||
{"type": "message", "role": "assistant", "content": "new"},
|
||||
]
|
||||
assert responses_state_context_tokens(state) == 100
|
||||
assert responses_state_context_tokens(state) == 175
|
||||
|
||||
def test_existing_compaction_keeps_canonical_retained_prefix(self):
|
||||
canonical_input = [
|
||||
@@ -1090,7 +1103,7 @@ class TestConsumeSse:
|
||||
assert content == "answer"
|
||||
assert tool_calls == []
|
||||
assert finish_reason == "stop"
|
||||
assert usage == {}
|
||||
assert usage is None
|
||||
assert reasoning == "thinking briefly\nChecking result"
|
||||
assert deltas == ["thinking ", "briefly", "\nChecking result"]
|
||||
|
||||
@@ -1224,7 +1237,7 @@ class TestConsumeSse:
|
||||
|
||||
assert content == "partial"
|
||||
assert finish_reason == expected_finish_reason
|
||||
assert usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
|
||||
assert usage == LLMUsage.reported(input_tokens=10, output_tokens=5)
|
||||
assert capture.completed is True
|
||||
assert capture.response == terminal_response
|
||||
assert capture.output_items == output
|
||||
@@ -1296,7 +1309,10 @@ class TestConsumeSse:
|
||||
"status": "completed",
|
||||
"usage": {
|
||||
"input_tokens": 10,
|
||||
"input_tokens_details": {"cached_tokens": 8},
|
||||
"input_tokens_details": {
|
||||
"cached_tokens": 8,
|
||||
"cache_write_tokens": 0,
|
||||
},
|
||||
"output_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
},
|
||||
@@ -1306,12 +1322,68 @@ class TestConsumeSse:
|
||||
|
||||
_, _, _, usage, _ = await consume_sse_with_reasoning(response)
|
||||
|
||||
assert usage == {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
"cached_tokens": 8,
|
||||
assert usage == LLMUsage.reported(
|
||||
input_tokens=10,
|
||||
output_tokens=5,
|
||||
cache_read_tokens=8,
|
||||
cache_write_tokens=0,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_and_non_stream_share_usage_normalization(self):
|
||||
terminal = {
|
||||
"status": "completed",
|
||||
"output": [],
|
||||
"usage": {
|
||||
"input_tokens": 15,
|
||||
"input_tokens_details": {
|
||||
"cached_tokens": 0,
|
||||
"cache_write_tokens": 7,
|
||||
},
|
||||
"output_tokens": 18,
|
||||
"total_tokens": 175,
|
||||
},
|
||||
}
|
||||
non_stream = parse_response_output(terminal).usage
|
||||
sse = _SseResponse([
|
||||
{"type": "response.completed", "response": terminal},
|
||||
])
|
||||
_, _, _, streamed, _ = await consume_sse_with_reasoning(sse)
|
||||
|
||||
sdk_response = SimpleNamespace(**terminal)
|
||||
sdk_response.usage = SimpleNamespace(
|
||||
input_tokens=15,
|
||||
input_tokens_details=SimpleNamespace(
|
||||
cached_tokens=0,
|
||||
cache_write_tokens=7,
|
||||
),
|
||||
output_tokens=18,
|
||||
total_tokens=175,
|
||||
)
|
||||
|
||||
async def sdk_stream():
|
||||
yield SimpleNamespace(type="response.completed", response=sdk_response)
|
||||
|
||||
_, _, _, sdk_streamed, _ = await consume_sdk_stream(sdk_stream())
|
||||
expected = LLMUsage.reported(
|
||||
input_tokens=15,
|
||||
output_tokens=18,
|
||||
total_tokens=175,
|
||||
cache_read_tokens=0,
|
||||
cache_write_tokens=7,
|
||||
)
|
||||
assert non_stream == streamed == sdk_streamed == expected
|
||||
|
||||
def test_missing_usage_is_not_explicit_zero_usage(self):
|
||||
missing = parse_response_output({"status": "completed", "output": []})
|
||||
explicit_zero = parse_response_output({
|
||||
"status": "completed",
|
||||
"output": [],
|
||||
"usage": {"input_tokens": 0, "output_tokens": 0},
|
||||
})
|
||||
|
||||
assert missing.usage is None
|
||||
assert explicit_zero.usage == LLMUsage.reported(input_tokens=0, output_tokens=0)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_done_arguments_callback(self):
|
||||
@@ -1778,25 +1850,24 @@ class TestConsumeSdkStream:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usage_extracted(self):
|
||||
usage_obj = MagicMock(
|
||||
usage_obj = SimpleNamespace(
|
||||
input_tokens=10,
|
||||
input_tokens_details=MagicMock(cached_tokens=8),
|
||||
input_tokens_details=SimpleNamespace(cached_tokens=8),
|
||||
output_tokens=5,
|
||||
total_tokens=15,
|
||||
)
|
||||
resp_obj = MagicMock(status="completed", usage=usage_obj, output=[])
|
||||
ev = MagicMock(type="response.completed", response=resp_obj)
|
||||
resp_obj = SimpleNamespace(status="completed", usage=usage_obj, output=[])
|
||||
ev = SimpleNamespace(type="response.completed", response=resp_obj)
|
||||
|
||||
async def stream():
|
||||
yield ev
|
||||
|
||||
_, _, _, usage, _ = await consume_sdk_stream(stream())
|
||||
assert usage == {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
"cached_tokens": 8,
|
||||
}
|
||||
assert usage == LLMUsage.reported(
|
||||
input_tokens=10,
|
||||
output_tokens=5,
|
||||
cache_read_tokens=8,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
@@ -1851,7 +1922,7 @@ class TestConsumeSdkStream:
|
||||
|
||||
assert content == "partial"
|
||||
assert finish_reason == expected_finish_reason
|
||||
assert usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
|
||||
assert usage == LLMUsage.reported(input_tokens=10, output_tokens=5)
|
||||
assert capture.completed is True
|
||||
assert capture.response == terminal_response
|
||||
assert capture.output_items == output
|
||||
|
||||
@@ -15,7 +15,7 @@ from nanobot.providers.base import (
|
||||
|
||||
class ScriptedProvider(LLMProvider):
|
||||
def __init__(self, responses):
|
||||
super().__init__()
|
||||
super().__init__(provider_name="scripted")
|
||||
self._responses = list(responses)
|
||||
self.calls = 0
|
||||
self.last_kwargs: dict = {}
|
||||
|
||||
@@ -32,6 +32,7 @@ def test_importing_providers_package_is_lazy(monkeypatch) -> None:
|
||||
assert providers.__all__ == [
|
||||
"LLMProvider",
|
||||
"LLMResponse",
|
||||
"LLMUsage",
|
||||
"AnthropicProvider",
|
||||
"OpenAICompatProvider",
|
||||
"OpenAICodexProvider",
|
||||
|
||||
@@ -0,0 +1,122 @@
|
||||
import pytest
|
||||
|
||||
from nanobot.providers.base import LLMUsage
|
||||
|
||||
|
||||
def test_reported_usage_derives_total_and_preserves_unreported_cache() -> None:
|
||||
usage = LLMUsage.reported(input_tokens=12, output_tokens=3)
|
||||
|
||||
assert usage.total_tokens == 15
|
||||
assert usage.cache_read_tokens is None
|
||||
assert usage.cache_write_tokens is None
|
||||
assert usage.source == "reported"
|
||||
|
||||
|
||||
def test_reported_usage_preserves_explicit_total_across_contract_operations() -> None:
|
||||
usage = LLMUsage.reported(input_tokens=15, output_tokens=18, total_tokens=175)
|
||||
|
||||
assert usage.total_tokens == 175
|
||||
assert usage.reported_tokens == 175
|
||||
assert usage.estimated_tokens == 0
|
||||
assert LLMUsage.from_dict(usage.to_dict()) == usage
|
||||
assert usage.with_timing(generation_ms=25, ttft_ms=5).total_tokens == 175
|
||||
|
||||
combined = usage + LLMUsage.estimated(input_tokens=2, output_tokens=1)
|
||||
assert combined.total_tokens == 178
|
||||
assert combined.reported_tokens == 175
|
||||
assert combined.estimated_tokens == 3
|
||||
|
||||
|
||||
def test_reported_usage_normalizes_missing_or_underreported_total() -> None:
|
||||
missing = LLMUsage.reported(input_tokens=15, output_tokens=18)
|
||||
underreported = LLMUsage.reported(
|
||||
input_tokens=15,
|
||||
output_tokens=18,
|
||||
total_tokens=12,
|
||||
)
|
||||
|
||||
assert missing.total_tokens == 33
|
||||
assert underreported.total_tokens == 33
|
||||
assert underreported.reported_tokens == 33
|
||||
|
||||
|
||||
def test_reported_usage_preserves_explicit_zero_cache() -> None:
|
||||
usage = LLMUsage.reported(
|
||||
input_tokens=12,
|
||||
output_tokens=3,
|
||||
cache_read_tokens=0,
|
||||
cache_write_tokens=0,
|
||||
)
|
||||
|
||||
assert usage.cache_read_tokens == 0
|
||||
assert usage.cache_write_tokens == 0
|
||||
|
||||
|
||||
def test_usage_rejects_inconsistent_token_partitions_and_cache_totals() -> None:
|
||||
with pytest.raises(ValueError, match="must equal"):
|
||||
LLMUsage(input_tokens=10, output_tokens=2, total_tokens=12, reported_tokens=11)
|
||||
|
||||
with pytest.raises(ValueError, match="at least"):
|
||||
LLMUsage(input_tokens=10, output_tokens=2, total_tokens=11, reported_tokens=11)
|
||||
|
||||
with pytest.raises(ValueError, match="cache token counts"):
|
||||
LLMUsage.reported(input_tokens=10, output_tokens=2, cache_read_tokens=11)
|
||||
|
||||
|
||||
def test_usage_serialization_is_strict_and_rejects_legacy_or_tampered_data() -> None:
|
||||
usage = LLMUsage.estimated(input_tokens=10, output_tokens=2)
|
||||
serialized = usage.to_dict()
|
||||
|
||||
assert LLMUsage.from_dict(serialized) == usage
|
||||
assert LLMUsage.from_dict({"prompt_tokens": 10, "completion_tokens": 2}) is None
|
||||
assert LLMUsage.from_dict({**serialized, "total_tokens": 99}) is None
|
||||
assert LLMUsage.from_dict({**serialized, "source": "reported"}) is None
|
||||
assert LLMUsage.from_dict({**serialized, "legacy_alias": 12}) is None
|
||||
|
||||
|
||||
def test_usage_aggregation_keeps_reported_estimated_split_and_unknown_cache() -> None:
|
||||
reported = LLMUsage.reported(
|
||||
input_tokens=10,
|
||||
output_tokens=2,
|
||||
total_tokens=20,
|
||||
cache_read_tokens=4,
|
||||
)
|
||||
estimated = LLMUsage.estimated(input_tokens=5, output_tokens=1)
|
||||
|
||||
combined = reported + estimated
|
||||
|
||||
assert combined.input_tokens == 15
|
||||
assert combined.output_tokens == 3
|
||||
assert combined.total_tokens == 26
|
||||
assert combined.reported_tokens == 20
|
||||
assert combined.estimated_tokens == 6
|
||||
assert combined.source == "mixed"
|
||||
assert combined.cache_read_tokens is None
|
||||
assert combined.context_tokens == 5
|
||||
assert combined.request_count == 2
|
||||
|
||||
|
||||
def test_usage_projects_compact_turn_observability_shape() -> None:
|
||||
usage = LLMUsage.reported(
|
||||
input_tokens=12,
|
||||
output_tokens=3,
|
||||
total_tokens=20,
|
||||
cache_read_tokens=4,
|
||||
) + LLMUsage.estimated(input_tokens=18, output_tokens=2)
|
||||
|
||||
assert usage.to_turn_dict() == {
|
||||
"prompt_tokens": 30,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 40,
|
||||
"context_tokens": 18,
|
||||
"request_count": 2,
|
||||
"estimated_tokens": 20,
|
||||
}
|
||||
|
||||
|
||||
def test_empty_request_counts_without_replacing_last_context() -> None:
|
||||
usage = LLMUsage.reported(input_tokens=12, output_tokens=3) + LLMUsage.empty_request()
|
||||
|
||||
assert usage.total_tokens == 15
|
||||
assert usage.context_tokens == 12
|
||||
assert usage.request_count == 2
|
||||
@@ -10,6 +10,7 @@ import httpx
|
||||
import pytest
|
||||
|
||||
from nanobot.config.schema import Config
|
||||
from nanobot.providers.base import LLMUsage
|
||||
from nanobot.providers.factory import make_provider
|
||||
from nanobot.providers.registry import find_by_name
|
||||
from nanobot.providers.xai_grok_provider import (
|
||||
@@ -454,7 +455,7 @@ async def test_raw_response_request_streams_text_usage_and_inline_citations(monk
|
||||
|
||||
assert result[0] == "Live result [[1]](https://x.com/example/status/1)"
|
||||
assert result[2] == "stop"
|
||||
assert result[3] == {"prompt_tokens": 8, "completion_tokens": 4, "total_tokens": 12}
|
||||
assert result[3] == LLMUsage.reported(input_tokens=8, output_tokens=4)
|
||||
assert deltas == ["Live result ", "[[1]](https://x.com/example/status/1)"]
|
||||
assert captured["json"]["tools"] == [{"type": "x_search"}]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user