refactor(providers): define typed usage contract

This commit is contained in:
chengyongru
2026-08-25 01:04:25 +08:00
committed by chengyongru
parent 89c94d8744
commit 9895c23cb5
84 changed files with 1643 additions and 726 deletions
+190 -31
View File
@@ -14,6 +14,7 @@ from nanobot.config.schema import AgentDefaults
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
LLMUsage,
ProviderCallContext,
ProviderConversationState,
ToolCallRequest,
@@ -22,6 +23,163 @@ from nanobot.providers.base import (
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
def _make_usage_spec(provider, tools):
return make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "hello"}],
tools=tools,
model="test-model",
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
)
def test_usage_or_estimate_replaces_reported_zero_for_content(monkeypatch) -> None:
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
tools = MagicMock()
tools.get_definitions.return_value = []
monkeypatch.setattr(
"nanobot.agent.runner.estimate_prompt_tokens_chain",
lambda provider, model, messages, definitions: (12, "test"),
)
monkeypatch.setattr("nanobot.agent.runner.estimate_message_tokens", lambda message: 7)
response = LLMResponse(
content="answer",
usage=LLMUsage.reported(input_tokens=0, output_tokens=0),
generation_ms=25,
ttft_ms=5,
)
usage = AgentRunner()._usage_or_estimate(
_make_usage_spec(provider, tools),
[{"role": "user", "content": "hello"}],
response,
)
assert usage == LLMUsage.estimated(input_tokens=12, output_tokens=7).with_timing(
generation_ms=25,
ttft_ms=5,
)
assert usage.source == "estimated"
assert usage.total_tokens == 19
def test_usage_or_estimate_counts_tool_call_output_for_reported_zero(monkeypatch) -> None:
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
tools = MagicMock()
tools.get_definitions.return_value = []
captured_message: dict = {}
monkeypatch.setattr(
"nanobot.agent.runner.estimate_prompt_tokens_chain",
lambda provider, model, messages, definitions: (13, "test"),
)
def estimate_output(message):
captured_message.update(message)
return 9
monkeypatch.setattr("nanobot.agent.runner.estimate_message_tokens", estimate_output)
response = LLMResponse(
content=None,
tool_calls=[
ToolCallRequest(
id="call_1",
name="lookup",
arguments={"query": "nanobot"},
)
],
finish_reason="tool_calls",
usage=LLMUsage.reported(input_tokens=0, output_tokens=0),
)
usage = AgentRunner()._usage_or_estimate(
_make_usage_spec(provider, tools),
[{"role": "user", "content": "hello"}],
response,
)
assert usage == LLMUsage.estimated(input_tokens=13, output_tokens=9)
assert usage.total_tokens == 22
assert captured_message["tool_calls"][0]["function"]["name"] == "lookup"
@pytest.mark.parametrize(
"provider_usage",
[None, LLMUsage.reported(input_tokens=0, output_tokens=0)],
)
def test_usage_or_estimate_counts_error_without_estimating_tokens(
monkeypatch,
provider_usage: LLMUsage | None,
) -> None:
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
tools = MagicMock()
estimate = MagicMock()
runner = AgentRunner()
monkeypatch.setattr(runner, "_estimate_response_usage", estimate)
response = LLMResponse(
content="upstream failed",
finish_reason="error",
usage=provider_usage,
)
usage = runner._usage_or_estimate(
_make_usage_spec(provider, tools),
[{"role": "user", "content": "hello"}],
response,
)
assert usage is not None
assert usage.total_tokens == 0
assert usage.request_count == 1
assert usage.context_tokens is None
aggregate = LLMUsage.reported(input_tokens=12, output_tokens=3) + usage
assert aggregate.context_tokens == 12
assert aggregate.request_count == 2
estimate.assert_not_called()
def test_usage_or_estimate_trusts_positive_reported_total(monkeypatch) -> None:
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
tools = MagicMock()
estimate = MagicMock()
runner = AgentRunner()
monkeypatch.setattr(runner, "_estimate_response_usage", estimate)
response = LLMResponse(
content="answer",
usage=LLMUsage.reported(
input_tokens=15,
output_tokens=18,
total_tokens=175,
),
generation_ms=30,
ttft_ms=6,
)
usage = runner._usage_or_estimate(
_make_usage_spec(provider, tools),
[{"role": "user", "content": "hello"}],
response,
)
assert usage is not None
assert usage.source == "reported"
assert usage.input_tokens == 15
assert usage.output_tokens == 18
assert usage.total_tokens == 175
assert usage.reported_tokens == 175
assert usage.generation_ms == 30
assert usage.ttft_ms == 6
estimate.assert_not_called()
@pytest.mark.asyncio
async def test_runner_preserves_reasoning_fields_and_tool_results():
from nanobot.agent.runner import AgentRunner
@@ -38,10 +196,10 @@ async def test_runner_preserves_reasoning_fields_and_tool_results():
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})],
reasoning_content="hidden reasoning",
thinking_blocks=[{"type": "thinking", "thinking": "step"}],
usage={"prompt_tokens": 5, "completion_tokens": 3},
usage=LLMUsage.reported(input_tokens=5, output_tokens=3),
)
captured_second_call[:] = messages
return LLMResponse(content="done", tool_calls=[], usage={})
return LLMResponse(content="done", tool_calls=[], usage=None)
provider.chat_with_retry = chat_with_retry
tools = MagicMock()
@@ -441,7 +599,7 @@ async def test_runner_uses_no_tools_finalization_after_max_iterations():
return LLMResponse(
content="Read the directory twice. More investigation remains.",
tool_calls=[],
usage={"prompt_tokens": 10, "completion_tokens": 7},
usage=LLMUsage.reported(input_tokens=10, output_tokens=7),
)
provider.chat_with_retry = chat_with_retry
@@ -713,10 +871,10 @@ async def test_runner_replaces_empty_tool_result_with_marker():
return LLMResponse(
content="working",
tool_calls=[ToolCallRequest(id="call_1", name="noop", arguments={})],
usage={},
usage=None,
)
captured_second_call[:] = messages
return LLMResponse(content="done", tool_calls=[], usage={})
return LLMResponse(content="done", tool_calls=[], usage=None)
provider.chat_with_retry = chat_with_retry
tools = MagicMock()
@@ -751,12 +909,12 @@ async def test_runner_retries_empty_final_response_with_summary_prompt():
return LLMResponse(
content=None,
tool_calls=[],
usage={"prompt_tokens": 5, "completion_tokens": 1},
usage=LLMUsage.reported(input_tokens=5, output_tokens=1),
)
return LLMResponse(
content="final answer",
tool_calls=[],
usage={"prompt_tokens": 3, "completion_tokens": 7},
usage=LLMUsage.reported(input_tokens=3, output_tokens=7),
)
provider.chat_with_retry = chat_with_retry
@@ -778,8 +936,9 @@ async def test_runner_retries_empty_final_response_with_summary_prompt():
assert calls[0]["tools"] is not None
assert calls[1]["tools"] is not None
assert calls[2]["tools"] is None
assert result.usage["prompt_tokens"] == 13
assert result.usage["completion_tokens"] == 9
assert result.usage is not None
assert result.usage.input_tokens == 13
assert result.usage.output_tokens == 9
@pytest.mark.asyncio
@@ -851,7 +1010,7 @@ async def test_runner_uses_specific_message_after_empty_finalization_retry():
provider = MagicMock(spec=LLMProvider)
async def chat_with_retry(*, messages, **kwargs):
return LLMResponse(content=None, tool_calls=[], usage={})
return LLMResponse(content=None, tool_calls=[], usage=None)
provider.chat_with_retry = chat_with_retry
tools = MagicMock()
@@ -891,14 +1050,14 @@ async def test_empty_finalization_retry_discards_candidate_provider_state():
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.chat_with_retry = AsyncMock(side_effect=[
LLMResponse(content=None, tool_calls=[], usage={}),
LLMResponse(content=None, tool_calls=[], usage={}),
LLMResponse(content=None, tool_calls=[], usage=None),
LLMResponse(content=None, tool_calls=[], usage=None),
LLMResponse(
content="finalized without tools",
tool_calls=[ToolCallRequest(id="call_1", name="exec", arguments={})],
finish_reason="stop",
provider_state=candidate,
usage={},
usage=None,
),
])
tools = MagicMock()
@@ -1037,20 +1196,20 @@ async def test_runner_empty_response_does_not_break_tool_chain():
return LLMResponse(
content=None,
tool_calls=[ToolCallRequest(id="tc1", name="read_file", arguments={"path": "a.txt"})],
usage={"prompt_tokens": 10, "completion_tokens": 5},
usage=LLMUsage.reported(input_tokens=10, output_tokens=5),
)
if call_count == 2:
return LLMResponse(content=None, tool_calls=[], usage={"prompt_tokens": 10, "completion_tokens": 1})
return LLMResponse(content=None, tool_calls=[], usage=LLMUsage.reported(input_tokens=10, output_tokens=1))
if call_count == 3:
return LLMResponse(
content=None,
tool_calls=[ToolCallRequest(id="tc2", name="read_file", arguments={"path": "b.txt"})],
usage={"prompt_tokens": 10, "completion_tokens": 5},
usage=LLMUsage.reported(input_tokens=10, output_tokens=5),
)
return LLMResponse(
content="Here are the results.",
tool_calls=[],
usage={"prompt_tokens": 10, "completion_tokens": 10},
usage=LLMUsage.reported(input_tokens=10, output_tokens=10),
)
provider.chat_with_retry = chat_with_retry
@@ -1079,9 +1238,8 @@ async def test_runner_empty_response_does_not_break_tool_chain():
@pytest.mark.asyncio
async def test_runner_accumulates_usage_and_preserves_cached_tokens():
"""Runner should accumulate prompt/completion tokens across iterations
and preserve cached_tokens from provider responses."""
async def test_runner_accumulates_usage_and_preserves_cache_reads():
"""Runner accumulates usage across iterations, including cache reads."""
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
@@ -1093,12 +1251,12 @@ async def test_runner_accumulates_usage_and_preserves_cached_tokens():
return LLMResponse(
content="thinking",
tool_calls=[ToolCallRequest(id="call_1", name="read_file", arguments={"path": "x"})],
usage={"prompt_tokens": 100, "completion_tokens": 10, "cached_tokens": 80},
usage=LLMUsage.reported(input_tokens=100, output_tokens=10, cache_read_tokens=80),
)
return LLMResponse(
content="done",
tool_calls=[],
usage={"prompt_tokens": 200, "completion_tokens": 20, "cached_tokens": 150},
usage=LLMUsage.reported(input_tokens=200, output_tokens=20, cache_read_tokens=150),
)
provider.chat_with_retry = chat_with_retry
@@ -1116,11 +1274,12 @@ async def test_runner_accumulates_usage_and_preserves_cached_tokens():
))
# Usage should be accumulated across iterations
assert result.usage["prompt_tokens"] == 300 # 100 + 200
assert result.usage["completion_tokens"] == 30 # 10 + 20
assert result.usage["cached_tokens"] == 230 # 80 + 150
assert result.usage["context_tokens"] == 200
assert result.usage["request_count"] == 2
assert result.usage is not None
assert result.usage.input_tokens == 300 # 100 + 200
assert result.usage.output_tokens == 30 # 10 + 20
assert result.usage.cache_read_tokens == 230 # 80 + 150
assert result.usage.context_tokens == 200
assert result.usage.request_count == 2
@pytest.mark.asyncio
@@ -1137,7 +1296,7 @@ async def test_runner_binds_on_retry_wait_to_retry_callback_not_progress():
async def chat_with_retry(**kwargs):
captured.update(kwargs)
return LLMResponse(content="done", tool_calls=[], usage={})
return LLMResponse(content="done", tool_calls=[], usage=None)
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = chat_with_retry
@@ -1179,7 +1338,7 @@ async def test_runner_passes_temperature_to_provider():
async def chat_with_retry(**kwargs):
captured.update(kwargs)
return LLMResponse(content="done", tool_calls=[], usage={})
return LLMResponse(content="done", tool_calls=[], usage=None)
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = chat_with_retry
@@ -1208,7 +1367,7 @@ async def test_runner_passes_max_tokens_to_provider():
async def chat_with_retry(**kwargs):
captured.update(kwargs)
return LLMResponse(content="done", tool_calls=[], usage={})
return LLMResponse(content="done", tool_calls=[], usage=None)
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = chat_with_retry
@@ -1237,7 +1396,7 @@ async def test_runner_passes_reasoning_effort_to_provider():
async def chat_with_retry(**kwargs):
captured.update(kwargs)
return LLMResponse(content="done", tool_calls=[], usage={})
return LLMResponse(content="done", tool_calls=[], usage=None)
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = chat_with_retry