From 9895c23cb533aa6269c1a2ef9d6b80ce5056db21 Mon Sep 17 00:00:00 2001 From: chengyongru <2755839590@qq.com> Date: Sat, 22 Aug 2026 00:53:16 +0800 Subject: [PATCH] refactor(providers): define typed usage contract --- nanobot/__init__.py | 3 + nanobot/agent/hook.py | 12 +- nanobot/agent/loop.py | 14 +- nanobot/agent/progress_hook.py | 12 +- nanobot/agent/runner.py | 119 +++---- nanobot/agent/subagent.py | 6 +- nanobot/agent/tools/runtime_control.py | 9 +- nanobot/agent/tools/self.py | 2 +- nanobot/agent/turn_delivery.py | 5 +- nanobot/api/server.py | 9 +- nanobot/bus/outbound_events.py | 8 +- nanobot/bus/runtime_events.py | 18 +- nanobot/channels/websocket/runtime.py | 22 +- .../websocket/tests/test_websocket_channel.py | 24 +- nanobot/command/builtin.py | 3 +- nanobot/nanobot.py | 4 +- nanobot/providers/__init__.py | 3 +- nanobot/providers/anthropic_provider.py | 40 +-- nanobot/providers/azure_openai_provider.py | 4 +- nanobot/providers/base.py | 295 +++++++++++++++++- nanobot/providers/bedrock_provider.py | 40 +-- nanobot/providers/factory.py | 8 +- nanobot/providers/fallback_provider.py | 3 + nanobot/providers/github_copilot_provider.py | 8 +- nanobot/providers/openai_codex_provider.py | 4 +- nanobot/providers/openai_compat_provider.py | 65 ++-- nanobot/providers/openai_responses/parsing.py | 68 ++-- nanobot/providers/openai_responses/state.py | 16 +- nanobot/providers/unconfigured_provider.py | 2 +- nanobot/providers/xai_grok_provider.py | 7 +- nanobot/sdk/types.py | 5 +- nanobot/session/webui_turns.py | 6 +- nanobot/skills/my/references/examples.md | 4 +- nanobot/utils/helpers.py | 15 +- nanobot/webui/session_context.py | 15 +- nanobot/webui/token_usage.py | 96 +++--- tests/agent/test_attachment_references.py | 2 +- tests/agent/test_dream.py | 8 +- tests/agent/test_evaluator.py | 2 +- tests/agent/test_history_replay.py | 4 +- tests/agent/test_hook_composite.py | 8 +- tests/agent/test_loop_progress.py | 4 +- tests/agent/test_loop_runner_integration.py | 44 +-- tests/agent/test_loop_save_turn.py | 11 +- tests/agent/test_loop_session_policy.py | 2 +- tests/agent/test_runner_core.py | 221 +++++++++++-- tests/agent/test_runner_errors.py | 16 +- tests/agent/test_runner_fallback.py | 37 ++- tests/agent/test_runner_goal_continue.py | 16 +- tests/agent/test_runner_governance.py | 6 +- tests/agent/test_runner_hooks.py | 61 ++-- tests/agent/test_runner_injections.py | 56 ++-- tests/agent/test_runner_persistence.py | 14 +- tests/agent/test_runner_progress_deltas.py | 20 +- tests/agent/test_runner_reasoning.py | 24 +- tests/agent/test_runner_tool_execution.py | 16 +- tests/agent/test_session_inputs.py | 2 +- tests/agent/test_session_model_runtime.py | 2 +- tests/agent/test_subagent_lifecycle.py | 10 +- tests/agent/tools/test_self_tool.py | 19 +- tests/agent/tools/test_subagent_tools.py | 8 +- tests/bus/test_runtime_events.py | 7 +- tests/cli/test_restart_command.py | 17 +- tests/providers/test_azure_openai_provider.py | 3 +- tests/providers/test_bedrock_provider.py | 54 +++- tests/providers/test_cached_tokens.py | 110 +++++-- tests/providers/test_custom_provider.py | 3 +- tests/providers/test_openai_codex_provider.py | 21 +- tests/providers/test_openai_responses.py | 131 ++++++-- tests/providers/test_provider_retry.py | 2 +- tests/providers/test_providers_init.py | 1 + tests/providers/test_usage_contract.py | 122 ++++++++ tests/providers/test_xai_grok_provider.py | 3 +- tests/test_api_attachment.py | 2 +- tests/test_api_stream.py | 16 +- tests/test_build_status.py | 19 +- tests/test_nanobot_facade.py | 20 +- tests/test_openai_api.py | 21 +- tests/utils/test_searchusage.py | 3 +- tests/webui/test_session_context.py | 14 +- tests/webui/test_settings_api.py | 5 +- tests/webui/test_token_usage.py | 164 ++++++++-- webui/src/lib/types.ts | 26 +- webui/src/tests/settings-overview.test.tsx | 18 +- 84 files changed, 1643 insertions(+), 726 deletions(-) create mode 100644 tests/providers/test_usage_contract.py diff --git a/nanobot/__init__.py b/nanobot/__init__.py index 20986b2bb..215dd80c8 100644 --- a/nanobot/__init__.py +++ b/nanobot/__init__.py @@ -23,6 +23,7 @@ if TYPE_CHECKING: STREAM_EVENT_TOOL_FAILED, STREAM_EVENT_TOOL_STARTED, STREAM_EVENT_TYPES, + LLMUsage, Nanobot, RunResult, RunStream, @@ -56,6 +57,7 @@ __logo__ = "🐈" _LAZY_EXPORTS = { "Nanobot": ".nanobot", + "LLMUsage": ".nanobot", "RunStream": ".nanobot", "RunResult": ".nanobot", "RequestContext": ".agent.tools.context", @@ -93,6 +95,7 @@ def __getattr__(name: str) -> Any: __all__ = [ "Nanobot", + "LLMUsage", "RunResult", "RequestContext", "RuntimeContextBlock", diff --git a/nanobot/agent/hook.py b/nanobot/agent/hook.py index ff5b1639a..2455c4b56 100644 --- a/nanobot/agent/hook.py +++ b/nanobot/agent/hook.py @@ -9,7 +9,7 @@ from typing import Any from loguru import logger -from nanobot.providers.base import LLMResponse, ToolCallRequest +from nanobot.providers.base import LLMResponse, LLMUsage, ToolCallRequest @dataclass(slots=True) @@ -19,7 +19,7 @@ class AgentHookContext: iteration: int messages: list[dict[str, Any]] response: LLMResponse | None = None - usage: dict[str, int] = field(default_factory=dict) + usage: LLMUsage | None = None tool_calls: list[ToolCallRequest] = field(default_factory=list) tool_results: list[Any] = field(default_factory=list) tool_events: list[dict[str, str]] = field(default_factory=list) @@ -39,7 +39,7 @@ class AgentRunHookContext: messages: list[dict[str, Any]] final_content: str | None = None tools_used: list[str] = field(default_factory=list) - usage: dict[str, int] = field(default_factory=dict) + usage: LLMUsage | None = None stop_reason: str | None = None error: str | None = None tool_events: list[dict[str, str]] = field(default_factory=list) @@ -284,7 +284,7 @@ class SDKCaptureHook(AgentHook): super().__init__() self.tools_used: list[str] = [] self.messages: list[dict[str, Any]] = [] - self.usage: dict[str, int] = {} + self.usage: LLMUsage | None = None self.stop_reason: str | None = None self.error: str | None = None self.tool_events: list[dict[str, str]] = [] @@ -294,7 +294,7 @@ class SDKCaptureHook(AgentHook): for call in context.tool_calls: self.tools_used.append(call.name) self.messages = list(context.messages) - self.usage = dict(context.usage) + self.usage = context.usage self.stop_reason = context.stop_reason self.error = context.error self.tool_events = list(context.tool_events) @@ -302,7 +302,7 @@ class SDKCaptureHook(AgentHook): async def after_run(self, context: AgentRunHookContext) -> None: self.tools_used = list(context.tools_used) self.messages = list(context.messages) - self.usage = dict(context.usage) + self.usage = context.usage self.stop_reason = context.stop_reason self.error = context.error self.tool_events = list(context.tool_events) diff --git a/nanobot/agent/loop.py b/nanobot/agent/loop.py index ae1ce06cf..2f34664f2 100644 --- a/nanobot/agent/loop.py +++ b/nanobot/agent/loop.py @@ -49,7 +49,7 @@ from nanobot.bus.queue import MessageBus from nanobot.bus.runtime_events import RuntimeEventBus from nanobot.command import CommandContext, CommandRouter, register_builtin_commands from nanobot.config.schema import AgentDefaults, ModelPresetConfig -from nanobot.providers.base import LLMProvider, ProviderConversationState +from nanobot.providers.base import LLMProvider, LLMUsage, ProviderConversationState from nanobot.providers.factory import ProviderSnapshot from nanobot.runtime_context import ( RUNTIME_CONTEXT_HISTORY_META, @@ -167,7 +167,7 @@ class TurnContext: turn_wall_started_at: float = field(default_factory=time.time) visible_run_started_at: float | None = None turn_latency_ms: int | None = None - usage: dict[str, int] = field(default_factory=dict) + usage: LLMUsage | None = None def require_runtime(self) -> LLMRuntime: """Return the runtime established by the BUILD stage.""" @@ -203,7 +203,7 @@ class AgentLoop: return self.tools.tool_names @property - def last_usage(self) -> Mapping[str, int]: + def last_usage(self) -> LLMUsage | None: """Latest aggregate usage exposed through the runtime-control snapshot.""" return self._last_usage @@ -378,7 +378,7 @@ class AgentLoop: default_restrict_to_workspace=restrict_to_workspace, ) self._start_time = time.time() - self._last_usage: dict[str, int] = {} + self._last_usage: LLMUsage | None = None self._extra_hooks: list[AgentHook] = hooks or [] self._hook_factories: list[AgentTurnHookFactory] = hook_factories or [] @@ -2031,7 +2031,7 @@ class AgentLoop: ctx.all_messages = all_msgs ctx.stop_reason = stop_reason ctx.had_injections = had_injections - ctx.usage = dict(self._last_usage) + ctx.usage = self._last_usage ctx.delivery.record_usage(ctx.usage) if ctx.kind is TurnKind.USER: await turn_continuation.maybe_continue_turn(ctx) @@ -2058,8 +2058,8 @@ class AgentLoop: else ctx.turn_wall_started_at ) ctx.turn_latency_ms = max(0, int((time.time() - latency_started_at) * 1000)) - if ctx.usage and not ctx.ephemeral: - session.metadata["_last_usage"] = dict(ctx.usage) + if ctx.usage is not None and not ctx.ephemeral: + session.metadata["_last_usage"] = ctx.usage.to_dict() self._save_turn( session, ctx.all_messages, ctx.save_skip, turn_latency_ms=ctx.turn_latency_ms, diff --git a/nanobot/agent/progress_hook.py b/nanobot/agent/progress_hook.py index 82b493cd4..0fa04748d 100644 --- a/nanobot/agent/progress_hook.py +++ b/nanobot/agent/progress_hook.py @@ -210,12 +210,14 @@ class AgentProgressHook(AgentHook): tool_hint=False, tool_events=tool_events, ) - u = context.usage or {} + u = context.usage logger.debug( - "LLM usage: prompt={} completion={} cached={}", - u.get("prompt_tokens", 0), - u.get("completion_tokens", 0), - u.get("cached_tokens", 0), + "LLM usage: input={} output={} cache_read={} cache_write={} source={}", + u.input_tokens if u else 0, + u.output_tokens if u else 0, + u.cache_read_tokens if u else None, + u.cache_write_tokens if u else None, + u.source if u else "missing", ) def finalize_content(self, context: AgentHookContext, content: str | None) -> str | None: diff --git a/nanobot/agent/runner.py b/nanobot/agent/runner.py index f66f01409..14131d2d5 100644 --- a/nanobot/agent/runner.py +++ b/nanobot/agent/runner.py @@ -23,6 +23,7 @@ from nanobot.agent.tools.registry import ToolRegistry, is_tool_error_result from nanobot.providers.base import ( LLMProvider, LLMResponse, + LLMUsage, ProviderCallContext, ProviderConversationState, ToolCallRequest, @@ -126,7 +127,7 @@ class AgentRunResult: final_content: str | None messages: list[dict[str, Any]] tools_used: list[str] = field(default_factory=list) - usage: dict[str, int] = field(default_factory=dict) + usage: LLMUsage | None = None stop_reason: str = "completed" error: str | None = None tool_events: list[dict[str, str]] = field(default_factory=list) @@ -412,7 +413,7 @@ class AgentRunner: context.messages = deepcopy(result.messages) context.final_content = result.final_content context.tools_used = list(result.tools_used) - context.usage = dict(result.usage) + context.usage = result.usage context.stop_reason = result.stop_reason context.error = result.error context.tool_events = deepcopy(result.tool_events) @@ -443,7 +444,7 @@ class AgentRunner: ) -> AgentRunResult: final_content: str | None = None tools_used: list[str] = [] - usage = {"prompt_tokens": 0, "completion_tokens": 0} + usage: LLMUsage | None = None error: str | None = None stop_reason = "completed" tool_events: list[dict[str, str]] = [] @@ -519,8 +520,8 @@ class AgentRunner: ) response.content = cleaned_content raw_usage = self._usage_or_estimate(spec, messages_for_model, response) - context.usage = dict(raw_usage) - self._accumulate_usage(usage, raw_usage) + context.usage = raw_usage + usage = self._merge_usage(usage, raw_usage) if reasoning_text and not context.streamed_reasoning: await hook.emit_reasoning(reasoning_text) await hook.emit_reasoning_end() @@ -683,10 +684,10 @@ class AgentRunner: conversation_state=conversation_state, ) retry_usage = self._usage_or_estimate(spec, retry_messages, response) - self._accumulate_usage(usage, retry_usage) + usage = self._merge_usage(usage, retry_usage) raw_usage = self._merge_usage(raw_usage, retry_usage) context.response = response - context.usage = dict(raw_usage) + context.usage = raw_usage context.tool_calls = list(response.tool_calls) original_content = response.content clean = hook.finalize_content(context, response.content) @@ -859,7 +860,7 @@ class AgentRunner: had_injections = True terminal_content = None if spec.finalize_on_max_iterations: - terminal_content = await self._try_finalize_after_max_iterations( + terminal_content, usage = await self._try_finalize_after_max_iterations( spec, hook, messages, @@ -1236,9 +1237,9 @@ class AgentRunner: spec: AgentRunSpec, hook: AgentHook, messages: list[dict[str, Any]], - usage: dict[str, int], + usage: LLMUsage | None, conversation_state: ProviderConversationStateController, - ) -> str | None: + ) -> tuple[str | None, LLMUsage | None]: retry_messages = self._budget_exhausted_finalization_messages(messages) try: response = await self._request_no_tools( @@ -1253,10 +1254,10 @@ class AgentRunner: "Budget-exhausted finalization failed for {}; using fallback", spec.session_key or "default", ) - return None + return None, usage raw_usage = self._usage_or_estimate(spec, retry_messages, response) - self._accumulate_usage(usage, raw_usage) + usage = self._merge_usage(usage, raw_usage) if response.finish_reason == "error" or response.has_tool_calls: logger.warning( "Budget-exhausted finalization returned finish_reason='{}' " @@ -1265,19 +1266,19 @@ class AgentRunner: len(response.tool_calls), spec.session_key or "default", ) - return None + return None, usage context = AgentHookContext( iteration=spec.max_iterations, messages=messages, response=response, - usage=dict(raw_usage), + usage=raw_usage, session_key=spec.session_key, ) clean = hook.finalize_content(context, response.content) if is_blank_text(clean): - return None - return clean + return None, usage + return clean, usage async def _request_no_tools( self, @@ -1349,31 +1350,24 @@ class AgentRunner: spec: AgentRunSpec, messages: list[dict[str, Any]], response: LLMResponse, - ) -> dict[str, int]: - usage = self._usage_dict(response.usage) - total = self._usage_total(usage) - if total > 0: - usage["total_tokens"] = total - usage.setdefault("provider_tokens", total) - elif response.finish_reason == "error": - return {} - else: + ) -> LLMUsage | None: + usage = response.usage + if response.finish_reason == "error": + if usage is None or usage.total_tokens == 0: + usage = LLMUsage.empty_request() + elif usage is None or usage.total_tokens == 0: usage = self._estimate_response_usage(spec, messages, response) - completion = usage.get("completion_tokens", 0) - if response.generation_ms is not None and completion > 0: - usage["generation_ms"] = response.generation_ms - usage["measured_completion_tokens"] = completion - if response.ttft_ms is not None: - usage["ttft_ms"] = response.ttft_ms - usage["timed_requests"] = 1 - return usage + return usage.with_timing( + generation_ms=response.generation_ms, + ttft_ms=response.ttft_ms, + ) def _estimate_response_usage( self, spec: AgentRunSpec, messages: list[dict[str, Any]], response: LLMResponse, - ) -> dict[str, int]: + ) -> LLMUsage: try: tools = spec.tools.get_definitions() except Exception: @@ -1391,52 +1385,21 @@ class AgentRunner: thinking_blocks=response.thinking_blocks, ) completion_tokens = estimate_message_tokens(assistant_message) - total_tokens = max(0, prompt_tokens) + max(0, completion_tokens) - if total_tokens <= 0: - return {} - return { - "prompt_tokens": max(0, prompt_tokens), - "completion_tokens": max(0, completion_tokens), - "total_tokens": total_tokens, - "estimated_tokens": total_tokens, - } + return LLMUsage.estimated( + input_tokens=max(0, prompt_tokens), + output_tokens=max(0, completion_tokens), + ) @staticmethod - def _usage_dict(usage: dict[str, Any] | None) -> dict[str, int]: - if not usage: - return {} - result: dict[str, int] = {} - for key, value in usage.items(): - try: - result[key] = int(value or 0) - except (TypeError, ValueError): - continue - return result - - @staticmethod - def _usage_total(usage: dict[str, int]) -> int: - return max(0, usage.get("total_tokens", 0) or ( - usage.get("prompt_tokens", 0) + usage.get("completion_tokens", 0) - )) - - @staticmethod - def _merge_usage(left: dict[str, int], right: dict[str, int]) -> dict[str, int]: - merged = dict(left) - for key, value in right.items(): - merged[key] = merged.get(key, 0) + value - return merged - - @staticmethod - def _accumulate_usage(total: dict[str, int], request: dict[str, int]) -> None: - """Fold one model request into the current turn's usage.""" - total["request_count"] = total.get("request_count", 0) + 1 - prompt_tokens = request.get("prompt_tokens") - if prompt_tokens is not None and prompt_tokens >= 0: - total["context_tokens"] = prompt_tokens - for key, value in request.items(): - if key in {"context_tokens", "request_count"} or value < 0: - continue - total[key] = total.get(key, 0) + value + def _merge_usage( + left: LLMUsage | None, + right: LLMUsage | None, + ) -> LLMUsage | None: + if left is None: + return right + if right is None: + return left + return left + right async def _execute_tools( self, diff --git a/nanobot/agent/subagent.py b/nanobot/agent/subagent.py index a3ea0a367..e88e43d38 100644 --- a/nanobot/agent/subagent.py +++ b/nanobot/agent/subagent.py @@ -28,7 +28,7 @@ from nanobot.agent.tools.registry import ToolRegistry from nanobot.bus.events import InboundMessage from nanobot.bus.queue import MessageBus from nanobot.config.schema import AgentDefaults, ToolsConfig -from nanobot.providers.base import LLMProvider +from nanobot.providers.base import LLMProvider, LLMUsage from nanobot.security.workspace_access import ( WorkspaceScope, bind_workspace_scope, @@ -56,7 +56,7 @@ class SubagentStatus: phase: str = "initializing" # initializing | awaiting_tools | tools_completed | final_response | done | error iteration: int = 0 tool_events: list[dict[str, str]] = field(default_factory=list) - usage: dict[str, int] = field(default_factory=dict) + usage: LLMUsage | None = None stop_reason: str | None = None error: str | None = None @@ -82,7 +82,7 @@ class _SubagentHook(AgentHook): return self._status.iteration = context.iteration self._status.tool_events = list(context.tool_events) - self._status.usage = dict(context.usage) + self._status.usage = context.usage if context.error: self._status.error = str(context.error) diff --git a/nanobot/agent/tools/runtime_control.py b/nanobot/agent/tools/runtime_control.py index 7118edc39..85227b4de 100644 --- a/nanobot/agent/tools/runtime_control.py +++ b/nanobot/agent/tools/runtime_control.py @@ -12,6 +12,7 @@ if TYPE_CHECKING: from nanobot.agent.tools.shell import ExecToolConfig from nanobot.agent.tools.web import WebToolsConfig from nanobot.config.schema import ModelPresetConfig + from nanobot.providers.base import LLMUsage from nanobot.utils.llm_runtime import LLMRuntime @@ -65,7 +66,7 @@ class RuntimeSnapshot: web_config: dict[str, object] exec_config: dict[str, object] subagent_statuses: dict[str, dict[str, object]] - last_usage: dict[str, int] + last_usage: Mapping[str, JsonScalar] scratchpad: dict[str, JsonValue] def as_mapping(self) -> Mapping[str, object]: @@ -151,7 +152,7 @@ class _RuntimeControlTarget(Protocol): def tool_names(self) -> list[str]: ... @property - def last_usage(self) -> Mapping[str, int]: ... + def last_usage(self) -> LLMUsage | None: ... def set_runtime_model(self, model: str) -> LLMRuntime: ... @@ -190,7 +191,7 @@ class AgentRuntimeControl: web_config=_snapshot_web_config(target.web_config), exec_config=_snapshot_exec_config(target.exec_config), subagent_statuses=_snapshot_subagent_statuses(target.subagents), - last_usage=dict(target.last_usage), + last_usage=target.last_usage.to_dict() if target.last_usage is not None else {}, scratchpad=_snapshot_json_mapping(self.__scratchpad), ) @@ -297,7 +298,7 @@ def _snapshot_subagent_status(status: SubagentStatus) -> dict[str, object]: "phase": status.phase, "iteration": status.iteration, "tool_events": [dict(event) for event in status.tool_events], - "usage": dict(status.usage), + "usage": status.usage.to_dict() if status.usage is not None else None, "stop_reason": status.stop_reason, "error": status.error, } diff --git a/nanobot/agent/tools/self.py b/nanobot/agent/tools/self.py index 45ed72dd6..1bc5000d0 100644 --- a/nanobot/agent/tools/self.py +++ b/nanobot/agent/tools/self.py @@ -150,7 +150,7 @@ class MyTool(Tool): "Actions: check, set.\n" "- check (no key): full config overview — start here.\n" "- check (key): drill into a value. Dot-paths allowed " - "(e.g. '_last_usage.prompt_tokens', 'web_config.enable').\n" + "(e.g. '_last_usage.input_tokens', 'web_config.enable').\n" "- set (key, value): change config or store notes in your scratchpad. " "Scratchpad keys persist across turns but not restarts.\n" "Key values: _current_iteration (current progress), " diff --git a/nanobot/agent/turn_delivery.py b/nanobot/agent/turn_delivery.py index c9563b76c..c1de57f01 100644 --- a/nanobot/agent/turn_delivery.py +++ b/nanobot/agent/turn_delivery.py @@ -4,7 +4,7 @@ from __future__ import annotations import dataclasses import time -from collections.abc import Awaitable, Callable, Mapping +from collections.abc import Awaitable, Callable from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, cast @@ -19,6 +19,7 @@ from nanobot.bus.outbound_events import ( from nanobot.bus.progress import build_bus_progress_callback from nanobot.bus.queue import MessageBus from nanobot.bus.runtime_events import RuntimeEventBus, RuntimeEventPublisher +from nanobot.providers.base import LLMUsage if TYPE_CHECKING: from nanobot.utils.llm_runtime import LLMRuntime @@ -203,7 +204,7 @@ class TurnDelivery: def record_latency(self, latency_ms: int | None) -> None: self.runtime_event_publisher.record_turn_latency(self.session_key, latency_ms) - def record_usage(self, usage: Mapping[str, int]) -> None: + def record_usage(self, usage: LLMUsage | None) -> None: self.runtime_event_publisher.record_turn_usage(self.session_key, usage) def background_response( diff --git a/nanobot/api/server.py b/nanobot/api/server.py index 59da093fe..2933c311f 100644 --- a/nanobot/api/server.py +++ b/nanobot/api/server.py @@ -18,6 +18,7 @@ from aiohttp import web from loguru import logger from nanobot.config.paths import get_media_dir +from nanobot.providers.base import LLMUsage from nanobot.utils.helpers import safe_filename from nanobot.utils.media_decode import ( MAX_FILE_SIZE, @@ -93,11 +94,11 @@ def _error_json(status: int, message: str, err_type: str = "invalid_request_erro def _chat_completion_response( content: str, model: str, - usage: dict[str, int] | None = None, + usage: LLMUsage | None = None, ) -> dict[str, Any]: - prompt = (usage or {}).get("prompt_tokens", 0) - completion = (usage or {}).get("completion_tokens", 0) - total = (usage or {}).get("total_tokens", 0) or prompt + completion + prompt = usage.input_tokens if usage else 0 + completion = usage.output_tokens if usage else 0 + total = usage.total_tokens if usage else 0 return { "id": f"chatcmpl-{uuid.uuid4().hex[:12]}", "object": "chat.completion", diff --git a/nanobot/bus/outbound_events.py b/nanobot/bus/outbound_events.py index 82ab3ed9f..68e7e1026 100644 --- a/nanobot/bus/outbound_events.py +++ b/nanobot/bus/outbound_events.py @@ -12,6 +12,7 @@ from dataclasses import dataclass, replace from typing import Any, cast from nanobot.bus.events import OutboundMessage +from nanobot.providers.base import LLMUsage class OutboundEvent: @@ -58,7 +59,7 @@ class StreamedResponseEvent(OutboundEvent): class TurnEndEvent(OutboundEvent): latency_ms: int | None = None goal_state: dict[str, Any] | None = None - usage: dict[str, int] | None = None + usage: LLMUsage | None = None context_window_tokens: int | None = None @@ -197,11 +198,6 @@ def _legacy_event_from_metadata(msg: OutboundMessage) -> OutboundEvent | None: return TurnEndEvent( latency_ms=_metadata_int(meta, "latency_ms"), goal_state=cast(dict[str, Any], goal_state) if isinstance(goal_state, dict) else None, - usage=( - cast(dict[str, int], meta.get("usage")) - if isinstance(meta.get("usage"), dict) - else None - ), context_window_tokens=_metadata_int(meta, "context_window_tokens"), ) if meta.get("_session_updated"): diff --git a/nanobot/bus/runtime_events.py b/nanobot/bus/runtime_events.py index d2cde53c3..c92890878 100644 --- a/nanobot/bus/runtime_events.py +++ b/nanobot/bus/runtime_events.py @@ -10,13 +10,14 @@ from __future__ import annotations import asyncio import contextlib import inspect -from collections.abc import Awaitable, Callable, Mapping +from collections.abc import Awaitable, Callable from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any from loguru import logger from nanobot.bus.events import InboundMessage +from nanobot.providers.base import LLMUsage if TYPE_CHECKING: from nanobot.utils.llm_runtime import LLMRuntime @@ -72,7 +73,7 @@ class TurnCompleted: context: RuntimeEventContext latency_ms: int | None = None runtime: LLMRuntime | None = None - usage: dict[str, int] = field(default_factory=dict) + usage: LLMUsage | None = None @dataclass(frozen=True) @@ -180,7 +181,7 @@ class RuntimeEventPublisher: self.bus = bus or RuntimeEventBus() self._turn_latency_ms: dict[str, int] = {} self._turn_runtime: dict[str, LLMRuntime] = {} - self._turn_usage: dict[str, dict[str, int]] = {} + self._turn_usage: dict[str, LLMUsage] = {} @staticmethod def _context( @@ -206,12 +207,9 @@ class RuntimeEventPublisher: if latency_ms is not None: self._turn_latency_ms[session_key] = int(latency_ms) - def record_turn_usage(self, session_key: str, usage: Mapping[str, int]) -> None: - self._turn_usage[session_key] = { - key: int(value) - for key, value in usage.items() - if type(value) is int and value >= 0 - } + def record_turn_usage(self, session_key: str, usage: LLMUsage | None) -> None: + if usage is not None: + self._turn_usage[session_key] = usage def clear_turn(self, session_key: str) -> None: self._turn_latency_ms.pop(session_key, None) @@ -332,7 +330,7 @@ class RuntimeEventPublisher: ), latency_ms=self._turn_latency_ms.pop(session_key, None), runtime=self._turn_runtime.pop(session_key, None), - usage=self._turn_usage.pop(session_key, {}), + usage=self._turn_usage.pop(session_key, None), ) ) diff --git a/nanobot/channels/websocket/runtime.py b/nanobot/channels/websocket/runtime.py index 1d8ee8d55..fc8e8bf20 100644 --- a/nanobot/channels/websocket/runtime.py +++ b/nanobot/channels/websocket/runtime.py @@ -44,6 +44,7 @@ from nanobot.bus.queue import MessageBus from nanobot.channels.base import BaseChannel from nanobot.command.builtin import USER_SHELL_COMMAND, builtin_command_starts_agent_turn from nanobot.config.schema import Base +from nanobot.providers.base import LLMUsage from nanobot.runtime_context import ( RUNTIME_CONTEXT_INPUT_META, WEBUI_QUOTE_METADATA, @@ -458,18 +459,9 @@ class WebSocketChannel(BaseChannel): recovery_state = recovery_state_from_metadata(metadata) if recovery_state is not None: fields["recovery_state"] = recovery_state - usage = metadata.get("_last_usage") - if isinstance(usage, dict): - sanitized_usage: dict[str, int | float] = {} - for key, value in cast(dict[object, object], usage).items(): - if ( - isinstance(key, str) - and isinstance(value, (int, float)) - and not isinstance(value, bool) - and value >= 0 - ): - sanitized_usage[key] = value - fields["usage"] = sanitized_usage + usage = LLMUsage.from_dict(metadata.get("_last_usage")) + if usage is not None: + fields["usage"] = usage.to_turn_dict() return fields def _detach(self, connection: ServerConnection, chat_id: str) -> None: @@ -2019,7 +2011,7 @@ class WebSocketChannel(BaseChannel): latency_ms: int | None = None, *, goal_state: dict[str, Any] | None = None, - usage: dict[str, int] | None = None, + usage: LLMUsage | None = None, context_window_tokens: int | None = None, metadata: dict[str, Any] | None = None, turn_owner: str | None = None, @@ -2034,8 +2026,8 @@ class WebSocketChannel(BaseChannel): body["latency_ms"] = int(latency_ms) if goal_state is not None: body["goal_state"] = goal_state - if usage: - body["usage"] = usage + if usage is not None: + body["usage"] = usage.to_turn_dict() if context_window_tokens is not None: body["context_window_tokens"] = int(context_window_tokens) canonical_webui_turn = (metadata or {}).get("webui") is True diff --git a/nanobot/channels/websocket/tests/test_websocket_channel.py b/nanobot/channels/websocket/tests/test_websocket_channel.py index c914e6620..96f08068d 100644 --- a/nanobot/channels/websocket/tests/test_websocket_channel.py +++ b/nanobot/channels/websocket/tests/test_websocket_channel.py @@ -44,6 +44,7 @@ from nanobot.channels.websocket.runtime import ( ) from nanobot.config.loader import load_config, save_config from nanobot.config.schema import Config, ModelPresetConfig +from nanobot.providers.base import LLMUsage from nanobot.runtime_context import RUNTIME_CONTEXT_INPUT_META, WEBUI_QUOTE_SOURCE from nanobot.security.workspace_access import WORKSPACE_SCOPE_METADATA_KEY from nanobot.session import webui_turns as wth @@ -2190,16 +2191,12 @@ async def test_send_scopes_turn_model_updates_to_the_subscribed_chat() -> None: def test_attach_fields_restore_the_session_model_and_latest_usage() -> None: + usage = LLMUsage.reported(input_tokens=120, output_tokens=8, total_tokens=175) manager = MagicMock() manager.read_session_metadata.return_value = { "metadata": { SESSION_MODEL_PRESET_METADATA_KEY: "Deep Research", - "_last_usage": { - "prompt_tokens": 120, - "completion_tokens": 8, - "negative": -1, - "boolean": True, - }, + "_last_usage": usage.to_dict(), } } bus = MagicMock() @@ -2211,7 +2208,7 @@ def test_attach_fields_restore_the_session_model_and_latest_usage() -> None: assert channel._attached_model_fields("chat-1") == { "model_preset": "Deep Research", - "usage": {"prompt_tokens": 120, "completion_tokens": 8}, + "usage": usage.to_turn_dict(), } @@ -3329,7 +3326,7 @@ async def test_send_turn_end_includes_latency_ms_when_present() -> None: content="", event=TurnEndEvent( latency_ms=1500, - usage={"prompt_tokens": 80, "completion_tokens": 20, "cached_tokens": 40}, + usage=LLMUsage.reported(input_tokens=80, output_tokens=20, cache_read_tokens=40), context_window_tokens=128_000, ), )) @@ -3339,7 +3336,11 @@ async def test_send_turn_end_includes_latency_ms_when_present() -> None: "event": "turn_end", "chat_id": "chat-1", "latency_ms": 1500, - "usage": {"prompt_tokens": 80, "completion_tokens": 20, "cached_tokens": 40}, + "usage": LLMUsage.reported( + input_tokens=80, + output_tokens=20, + cache_read_tokens=40, + ).to_turn_dict(), "context_window_tokens": 128_000, }, {"event": "session_updated", "chat_id": "chat-1", "scope": "thread"}, @@ -5306,10 +5307,11 @@ async def test_handle_session_context_get_reads_detached_session() -> None: from nanobot.session import Session + usage = LLMUsage.reported(input_tokens=12, output_tokens=3, total_tokens=175) session = Session( key="websocket:context-route", messages=[{"role": "user", "content": "hello"}], - metadata={"_last_usage": {"prompt_tokens": 12, "completion_tokens": 3}}, + metadata={"_last_usage": usage.to_dict()}, ) manager = MagicMock() manager.read_session_snapshot.return_value = session @@ -5326,7 +5328,7 @@ async def test_handle_session_context_get_reads_detached_session() -> None: assert response.status_code == 200 body = json.loads(response.body.decode()) assert body["replay_messages"] == 1 - assert body["last_usage"] == {"prompt_tokens": 12, "completion_tokens": 3} + assert body["last_usage"] == usage.to_dict() manager.read_session_snapshot.assert_called_once_with(session.key) diff --git a/nanobot/command/builtin.py b/nanobot/command/builtin.py index cff9b76c5..0d3e1c2cb 100644 --- a/nanobot/command/builtin.py +++ b/nanobot/command/builtin.py @@ -266,7 +266,8 @@ async def cmd_status(ctx: CommandContext) -> OutboundMessage: runtime=runtime, ) if ctx_est <= 0: - ctx_est = loop._last_usage.get("prompt_tokens", 0) # pyright: ignore[reportPrivateUsage] + last_usage = loop._last_usage # pyright: ignore[reportPrivateUsage] + ctx_est = last_usage.input_tokens if last_usage is not None else 0 # Fetch web search provider usage (best-effort, never blocks the response) search_usage_text: str | None = None diff --git a/nanobot/nanobot.py b/nanobot/nanobot.py index 272043838..6c9951ed8 100644 --- a/nanobot/nanobot.py +++ b/nanobot/nanobot.py @@ -13,6 +13,7 @@ from nanobot.agent.loop import AgentLoop from nanobot.agent.tools.mcp import MCPProvider from nanobot.agent.tools.registry import ToolRegistry from nanobot.config.schema import Config +from nanobot.providers.base import LLMUsage from nanobot.providers.image_generation import image_gen_provider_configs from nanobot.sdk.clients import MemoryClient, RuntimeClient, SessionClient from nanobot.sdk.runtime import ( @@ -43,6 +44,7 @@ from nanobot.utils.llm_runtime import LLMRuntime __all__ = [ "Nanobot", + "LLMUsage", "RunResult", "RunStream", "SessionInfo", @@ -287,7 +289,7 @@ class Nanobot: type=STREAM_EVENT_RUN_COMPLETED, content=result.content, result=result, - usage=dict(result.usage), + usage=result.usage, metadata=dict(result.metadata), )) return result diff --git a/nanobot/providers/__init__.py b/nanobot/providers/__init__.py index 4eea5d3d1..0d12ace8e 100644 --- a/nanobot/providers/__init__.py +++ b/nanobot/providers/__init__.py @@ -5,11 +5,12 @@ from __future__ import annotations from importlib import import_module from typing import TYPE_CHECKING -from nanobot.providers.base import LLMProvider, LLMResponse +from nanobot.providers.base import LLMProvider, LLMResponse, LLMUsage __all__ = [ "LLMProvider", "LLMResponse", + "LLMUsage", "AnthropicProvider", "OpenAICompatProvider", "OpenAICodexProvider", diff --git a/nanobot/providers/anthropic_provider.py b/nanobot/providers/anthropic_provider.py index 841387a83..201172510 100644 --- a/nanobot/providers/anthropic_provider.py +++ b/nanobot/providers/anthropic_provider.py @@ -17,6 +17,7 @@ from loguru import logger from nanobot.providers.base import ( LLMProvider, LLMResponse, + LLMUsage, ToolCallRequest, resolve_stream_idle_timeout_s, tool_arguments_object_for_replay, @@ -90,8 +91,10 @@ class AnthropicProvider(LLMProvider): api_base: str | None = None, default_model: str = "claude-sonnet-4-6", extra_headers: dict[str, str] | None = None, + *, + provider_name: str = "anthropic", ): - super().__init__(api_key, api_base) + super().__init__(api_key, api_base, provider_name=provider_name) self.default_model = default_model self.extra_headers = extra_headers or {} @@ -689,24 +692,25 @@ class AnthropicProvider(LLMProvider): stop_map = {"tool_use": "tool_calls", "end_turn": "stop", "max_tokens": "length"} finish_reason = stop_map.get(response.stop_reason or "", response.stop_reason or "stop") - usage: dict[str, int] = {} + usage: LLMUsage | None = None if response.usage: - input_tokens = response.usage.input_tokens - cache_creation = getattr(response.usage, "cache_creation_input_tokens", 0) or 0 - cache_read = getattr(response.usage, "cache_read_input_tokens", 0) or 0 - total_prompt_tokens = input_tokens + cache_creation + cache_read - usage = { - "prompt_tokens": total_prompt_tokens, - "completion_tokens": response.usage.output_tokens, - "total_tokens": total_prompt_tokens + response.usage.output_tokens, - } - for attr in ("cache_creation_input_tokens", "cache_read_input_tokens"): - val = getattr(response.usage, attr, 0) - if val: - usage[attr] = val - # Normalize to cached_tokens for downstream consistency. - if cache_read: - usage["cached_tokens"] = cache_read + cache_write_raw = getattr( + response.usage, + "cache_creation_input_tokens", + None, + ) + cache_read_raw = getattr(response.usage, "cache_read_input_tokens", None) + cache_write = int(cache_write_raw) if cache_write_raw is not None else None + cache_read = int(cache_read_raw) if cache_read_raw is not None else None + logical_input = int(response.usage.input_tokens) + (cache_write or 0) + ( + cache_read or 0 + ) + usage = LLMUsage.reported( + input_tokens=logical_input, + output_tokens=int(response.usage.output_tokens), + cache_read_tokens=cache_read, + cache_write_tokens=cache_write, + ) return LLMResponse( content="".join(content_parts) or None, diff --git a/nanobot/providers/azure_openai_provider.py b/nanobot/providers/azure_openai_provider.py index e7d718133..cf851fb05 100644 --- a/nanobot/providers/azure_openai_provider.py +++ b/nanobot/providers/azure_openai_provider.py @@ -106,8 +106,10 @@ class AzureOpenAIProvider(LLMProvider): api_key: str = "", api_base: str = "", default_model: str = "gpt-5.2-chat", + *, + provider_name: str = "azure_openai", ): - super().__init__(api_key, api_base) + super().__init__(api_key, api_base, provider_name=provider_name) self.default_model = default_model self._native_compaction_available = True diff --git a/nanobot/providers/base.py b/nanobot/providers/base.py index f29cd2c82..749d9963d 100644 --- a/nanobot/providers/base.py +++ b/nanobot/providers/base.py @@ -13,7 +13,7 @@ from copy import deepcopy from dataclasses import dataclass, field from datetime import datetime, timezone from email.utils import parsedate_to_datetime -from typing import Any, cast +from typing import Any, Literal, cast import json_repair from loguru import logger @@ -253,13 +253,292 @@ class ProviderCallContext: context_window_tokens: int | None = None +@dataclass(frozen=True, slots=True) +class LLMUsage: + """Canonical token usage reported by, or estimated for, one or more LLM calls. + + ``input_tokens`` is the logical input total and therefore includes cache reads + and writes. ``None`` cache counts mean the wire protocol did not report that + metric, while zero means it explicitly reported no cache activity. + + ``total_tokens`` preserves a provider-reported total when it exceeds the + visible input plus output (for example, hidden reasoning or tool usage). It + must be at least ``input_tokens + output_tokens``. The reported and estimated + totals partition it exactly, including after multi-call aggregation. + """ + + input_tokens: int + output_tokens: int + total_tokens: int + cache_read_tokens: int | None = None + cache_write_tokens: int | None = None + reported_tokens: int = 0 + estimated_tokens: int = 0 + generation_ms: int = 0 + measured_output_tokens: int = 0 + ttft_ms: int = 0 + timed_requests: int = 0 + context_tokens: int | None = None + request_count: int = 0 + + def __post_init__(self) -> None: + token_fields = { + "input_tokens": self.input_tokens, + "output_tokens": self.output_tokens, + "total_tokens": self.total_tokens, + "reported_tokens": self.reported_tokens, + "estimated_tokens": self.estimated_tokens, + "generation_ms": self.generation_ms, + "measured_output_tokens": self.measured_output_tokens, + "ttft_ms": self.ttft_ms, + "timed_requests": self.timed_requests, + "request_count": self.request_count, + } + for name, value in token_fields.items(): + runtime_value = cast(object, value) + if ( + not isinstance(runtime_value, int) + or isinstance(runtime_value, bool) + or runtime_value < 0 + ): + raise ValueError(f"{name} must be a non-negative integer") + for name, value in ( + ("cache_read_tokens", self.cache_read_tokens), + ("cache_write_tokens", self.cache_write_tokens), + ("context_tokens", self.context_tokens), + ): + runtime_value = cast(object, value) + if runtime_value is not None and ( + not isinstance(runtime_value, int) + or isinstance(runtime_value, bool) + or runtime_value < 0 + ): + raise ValueError(f"{name} must be None or a non-negative integer") + + visible_total = self.input_tokens + self.output_tokens + if self.total_tokens < visible_total: + raise ValueError("total_tokens must be at least input_tokens + output_tokens") + if self.reported_tokens + self.estimated_tokens != self.total_tokens: + raise ValueError("reported_tokens + estimated_tokens must equal total_tokens") + cache_total = (self.cache_read_tokens or 0) + (self.cache_write_tokens or 0) + if cache_total > self.input_tokens: + raise ValueError("cache token counts cannot exceed logical input_tokens") + + @classmethod + def reported( + cls, + *, + input_tokens: int, + output_tokens: int, + total_tokens: int | None = None, + cache_read_tokens: int | None = None, + cache_write_tokens: int | None = None, + ) -> LLMUsage: + """Build usage normalized from a provider response.""" + visible_total = input_tokens + output_tokens + normalized_total = ( + visible_total if total_tokens is None else max(visible_total, total_tokens) + ) + return cls( + input_tokens=input_tokens, + output_tokens=output_tokens, + total_tokens=normalized_total, + cache_read_tokens=cache_read_tokens, + cache_write_tokens=cache_write_tokens, + reported_tokens=normalized_total, + context_tokens=input_tokens, + request_count=1, + ) + + @classmethod + def estimated(cls, *, input_tokens: int, output_tokens: int) -> LLMUsage: + """Build usage estimated locally because the provider omitted it.""" + return cls( + input_tokens=input_tokens, + output_tokens=output_tokens, + total_tokens=input_tokens + output_tokens, + estimated_tokens=input_tokens + output_tokens, + context_tokens=input_tokens, + request_count=1, + ) + + @classmethod + def empty_request(cls) -> LLMUsage: + """Represent a completed model request with no measurable token usage.""" + return cls( + input_tokens=0, + output_tokens=0, + total_tokens=0, + request_count=1, + ) + + @property + def source(self) -> Literal["reported", "estimated", "mixed"]: + if self.estimated_tokens == 0: + return "reported" + if self.reported_tokens == 0: + return "estimated" + return "mixed" + + def with_timing( + self, + *, + generation_ms: int | None, + ttft_ms: int | None, + ) -> LLMUsage: + """Attach locally measured streaming telemetry to this usage value.""" + return LLMUsage( + input_tokens=self.input_tokens, + output_tokens=self.output_tokens, + total_tokens=self.total_tokens, + cache_read_tokens=self.cache_read_tokens, + cache_write_tokens=self.cache_write_tokens, + reported_tokens=self.reported_tokens, + estimated_tokens=self.estimated_tokens, + generation_ms=max(0, generation_ms or 0), + measured_output_tokens=self.output_tokens if generation_ms is not None else 0, + ttft_ms=max(0, ttft_ms or 0), + timed_requests=1 if ttft_ms is not None else 0, + context_tokens=self.context_tokens, + request_count=self.request_count, + ) + + def __add__(self, other: LLMUsage) -> LLMUsage: + """Aggregate calls without turning partially reported cache data into a count.""" + + def _sum_cache(left: int | None, right: int | None) -> int | None: + return left + right if left is not None and right is not None else None + + return LLMUsage( + input_tokens=self.input_tokens + other.input_tokens, + output_tokens=self.output_tokens + other.output_tokens, + total_tokens=self.total_tokens + other.total_tokens, + cache_read_tokens=_sum_cache(self.cache_read_tokens, other.cache_read_tokens), + cache_write_tokens=_sum_cache(self.cache_write_tokens, other.cache_write_tokens), + reported_tokens=self.reported_tokens + other.reported_tokens, + estimated_tokens=self.estimated_tokens + other.estimated_tokens, + generation_ms=self.generation_ms + other.generation_ms, + measured_output_tokens=( + self.measured_output_tokens + other.measured_output_tokens + ), + ttft_ms=self.ttft_ms + other.ttft_ms, + timed_requests=self.timed_requests + other.timed_requests, + context_tokens=( + other.context_tokens + if other.context_tokens is not None + else self.context_tokens + ), + request_count=self.request_count + other.request_count, + ) + + def to_dict(self) -> dict[str, int | str | None]: + """Serialize the canonical contract at JSON/persistence boundaries.""" + return { + "input_tokens": self.input_tokens, + "output_tokens": self.output_tokens, + "total_tokens": self.total_tokens, + "cache_read_tokens": self.cache_read_tokens, + "cache_write_tokens": self.cache_write_tokens, + "reported_tokens": self.reported_tokens, + "estimated_tokens": self.estimated_tokens, + "source": self.source, + "generation_ms": self.generation_ms, + "measured_output_tokens": self.measured_output_tokens, + "ttft_ms": self.ttft_ms, + "timed_requests": self.timed_requests, + "context_tokens": self.context_tokens, + "request_count": self.request_count, + } + + def to_turn_dict(self) -> dict[str, int]: + """Project canonical usage into the WebUI's compact per-turn shape.""" + result: dict[str, int] = { + "prompt_tokens": self.input_tokens, + "completion_tokens": self.output_tokens, + "total_tokens": self.total_tokens, + "request_count": self.request_count, + "estimated_tokens": self.estimated_tokens, + } + if self.context_tokens is not None: + result["context_tokens"] = self.context_tokens + if self.cache_read_tokens is not None: + result["cached_tokens"] = self.cache_read_tokens + if self.cache_write_tokens is not None: + result["cache_write_tokens"] = self.cache_write_tokens + return result + + @classmethod + def from_dict(cls, value: object) -> LLMUsage | None: + """Validate the exact first-party serialized contract.""" + if not isinstance(value, dict): + return None + data = cast(dict[object, object], value) + integer_fields = ( + "input_tokens", + "output_tokens", + "reported_tokens", + "estimated_tokens", + "generation_ms", + "measured_output_tokens", + "ttft_ms", + "timed_requests", + "request_count", + ) + serialized_fields = { + *integer_fields, + "total_tokens", + "cache_read_tokens", + "cache_write_tokens", + "context_tokens", + "source", + } + if set(data) != serialized_fields: + return None + if any( + not isinstance(item := data.get(name), int) or isinstance(item, bool) + for name in integer_fields + ): + return None + cache_read = data.get("cache_read_tokens") + cache_write = data.get("cache_write_tokens") + context_tokens = data.get("context_tokens") + total = data.get("total_tokens") + source = data.get("source") + if any( + item is not None and (not isinstance(item, int) or isinstance(item, bool)) + for item in (cache_read, cache_write, context_tokens) + ) or not isinstance(total, int) or isinstance(total, bool): + return None + try: + usage = cls( + input_tokens=cast(int, data["input_tokens"]), + output_tokens=cast(int, data["output_tokens"]), + total_tokens=total, + cache_read_tokens=cast(int | None, cache_read), + cache_write_tokens=cast(int | None, cache_write), + reported_tokens=cast(int, data["reported_tokens"]), + estimated_tokens=cast(int, data["estimated_tokens"]), + generation_ms=cast(int, data["generation_ms"]), + measured_output_tokens=cast(int, data["measured_output_tokens"]), + ttft_ms=cast(int, data["ttft_ms"]), + timed_requests=cast(int, data["timed_requests"]), + context_tokens=cast(int | None, context_tokens), + request_count=cast(int, data["request_count"]), + ) + except (KeyError, ValueError): + return None + if source != usage.source: + return None + return usage + + @dataclass class LLMResponse: """Response from an LLM provider.""" content: str | None tool_calls: list[ToolCallRequest] = field(default_factory=list) finish_reason: str = "stop" - usage: dict[str, int] = field(default_factory=dict) + usage: LLMUsage | None = None # Locally measured streaming telemetry. ``generation_ms`` excludes time to # first token and provider retry gaps; ``ttft_ms`` measures the first # streamed reasoning/content delta from request start. They stay separate @@ -383,9 +662,19 @@ class LLMProvider(ABC): _SENTINEL = object() - def __init__(self, api_key: str | None = None, api_base: str | None = None): + def __init__( + self, + api_key: str | None = None, + api_base: str | None = None, + *, + provider_name: str, + ): + runtime_provider_name = cast(object, provider_name) + if not isinstance(runtime_provider_name, str) or not runtime_provider_name.strip(): + raise ValueError("provider_name must be a non-empty configured identity") self.api_key = api_key self.api_base = api_base + self.provider_name = provider_name self.generation: GenerationSettings = GenerationSettings() def can_resume_conversation_state( diff --git a/nanobot/providers/bedrock_provider.py b/nanobot/providers/bedrock_provider.py index 7a4728302..c461545cb 100644 --- a/nanobot/providers/bedrock_provider.py +++ b/nanobot/providers/bedrock_provider.py @@ -14,6 +14,7 @@ from typing import Any, cast from nanobot.providers.base import ( LLMProvider, LLMResponse, + LLMUsage, ToolCallRequest, parse_tool_arguments, resolve_stream_idle_timeout_s, @@ -60,8 +61,9 @@ class BedrockProvider(LLMProvider): profile: str | None = None, extra_body: dict[str, Any] | None = None, client: Any | None = None, + provider_name: str = "bedrock", ): - super().__init__(api_key, api_base) + super().__init__(api_key, api_base, provider_name=provider_name) self.default_model = default_model self.region = region or os.environ.get("AWS_REGION") or os.environ.get("AWS_DEFAULT_REGION") self.profile = profile @@ -453,25 +455,25 @@ class BedrockProvider(LLMProvider): }.get(stop_reason or "", stop_reason or "stop") @staticmethod - def _usage(usage: dict[str, Any] | None) -> dict[str, int]: + def _usage(usage: dict[str, Any] | None) -> LLMUsage | None: if not usage: - return {} - prompt = int(usage.get("inputTokens") or 0) - completion = int(usage.get("outputTokens") or 0) - total = int(usage.get("totalTokens") or prompt + completion) - result = { - "prompt_tokens": prompt, - "completion_tokens": completion, - "total_tokens": total, - } - cache_read = int(usage.get("cacheReadInputTokens") or 0) - cache_write = int(usage.get("cacheWriteInputTokens") or 0) - if cache_read: - result["cached_tokens"] = cache_read - result["cache_read_input_tokens"] = cache_read - if cache_write: - result["cache_creation_input_tokens"] = cache_write - return result + return None + + def _optional_count(key: str) -> int | None: + raw = usage.get(key) + return int(raw) if raw is not None else None + + cache_read = _optional_count("cacheReadInputTokens") + cache_write = _optional_count("cacheWriteInputTokens") + logical_input = int(usage.get("inputTokens") or 0) + (cache_read or 0) + ( + cache_write or 0 + ) + return LLMUsage.reported( + input_tokens=logical_input, + output_tokens=int(usage.get("outputTokens") or 0), + cache_read_tokens=cache_read, + cache_write_tokens=cache_write, + ) @staticmethod def _parse_reasoning(block: dict[str, Any]) -> tuple[str | None, dict[str, Any] | None]: diff --git a/nanobot/providers/factory.py b/nanobot/providers/factory.py index c9346a137..d6f0483b6 100644 --- a/nanobot/providers/factory.py +++ b/nanobot/providers/factory.py @@ -172,6 +172,7 @@ def _make_provider_core( default_model=model, proxy=getattr(p, "proxy", None) if p else None, extra_body=p.extra_body if p else None, + provider_name=provider_name, ) elif backend == "xai_grok": from nanobot.providers.xai_grok_provider import XAIGrokProvider @@ -180,6 +181,7 @@ def _make_provider_core( default_model=model, proxy=getattr(p, "proxy", None) if p else None, extra_body=p.extra_body if p else None, + provider_name=provider_name, ) elif backend == "azure_openai": from nanobot.providers.azure_openai_provider import AzureOpenAIProvider @@ -190,11 +192,12 @@ def _make_provider_core( api_key=p.api_key or "", api_base=p.api_base, default_model=model, + provider_name=provider_name, ) elif backend == "github_copilot": from nanobot.providers.github_copilot_provider import GitHubCopilotProvider - provider = GitHubCopilotProvider(default_model=model) + provider = GitHubCopilotProvider(default_model=model, provider_name=provider_name) elif backend == "anthropic": from nanobot.providers.anthropic_provider import AnthropicProvider @@ -203,6 +206,7 @@ def _make_provider_core( api_base=config.get_api_base(model, preset=preset), default_model=model, extra_headers=_provider_extra_headers(spec, p), + provider_name=provider_name, ) elif backend == "bedrock": from nanobot.providers.bedrock_provider import BedrockProvider @@ -214,6 +218,7 @@ def _make_provider_core( region=getattr(p, "region", None) if p else None, profile=getattr(p, "profile", None) if p else None, extra_body=p.extra_body if p else None, + provider_name=provider_name, ) else: from nanobot.providers.openai_compat_provider import OpenAICompatProvider @@ -228,6 +233,7 @@ def _make_provider_core( api_type=p.api_type if p and provider_name == "openai" else "auto", extra_query=p.extra_query if p else None, proxy=p.proxy if p else None, + provider_name=provider_name, ) provider.generation = preset.to_generation_settings() diff --git a/nanobot/providers/fallback_provider.py b/nanobot/providers/fallback_provider.py index 141bc496f..0920a724f 100644 --- a/nanobot/providers/fallback_provider.py +++ b/nanobot/providers/fallback_provider.py @@ -124,7 +124,10 @@ class FallbackProvider(LLMProvider): fallback_model_observer: FallbackModelObserver | None = None, primary_context_window_tokens: int | None = None, ): + primary_generation = primary.generation self._primary = primary + super().__init__(provider_name=primary.provider_name) + self._primary.generation = primary_generation self._fallback_presets = list(fallback_presets) self._provider_factory = provider_factory self._fallback_model_observer = fallback_model_observer diff --git a/nanobot/providers/github_copilot_provider.py b/nanobot/providers/github_copilot_provider.py index 4ae3d9f39..65f40b892 100644 --- a/nanobot/providers/github_copilot_provider.py +++ b/nanobot/providers/github_copilot_provider.py @@ -174,7 +174,12 @@ def login_github_copilot( class GitHubCopilotProvider(OpenAICompatProvider): """Provider that exchanges a stored GitHub OAuth token for Copilot access tokens.""" - def __init__(self, default_model: str = "github-copilot/gpt-4.1"): + def __init__( + self, + default_model: str = "github-copilot/gpt-4.1", + *, + provider_name: str = "github_copilot", + ): from nanobot.providers.registry import find_by_name self._copilot_access_token: str | None = None @@ -190,6 +195,7 @@ class GitHubCopilotProvider(OpenAICompatProvider): "User-Agent": USER_AGENT, }, spec=find_by_name("github_copilot"), + provider_name=provider_name, ) async def _get_copilot_access_token(self) -> str: diff --git a/nanobot/providers/openai_codex_provider.py b/nanobot/providers/openai_codex_provider.py index 1a3881dd0..d8cd4cd82 100644 --- a/nanobot/providers/openai_codex_provider.py +++ b/nanobot/providers/openai_codex_provider.py @@ -51,8 +51,10 @@ class OpenAICodexProvider(LLMProvider): default_model: str = "openai-codex/gpt-5.6-sol", proxy: str | None = None, extra_body: dict[str, Any] | None = None, + *, + provider_name: str = "openai_codex", ): - super().__init__(api_key=None, api_base=None) + super().__init__(api_key=None, api_base=None, provider_name=provider_name) self.default_model = default_model self.proxy = proxy or None self._extra_body = dict(extra_body or {}) diff --git a/nanobot/providers/openai_compat_provider.py b/nanobot/providers/openai_compat_provider.py index 61649d9ac..54a9bb0b9 100644 --- a/nanobot/providers/openai_compat_provider.py +++ b/nanobot/providers/openai_compat_provider.py @@ -26,6 +26,7 @@ from pydantic.alias_generators import to_snake from nanobot.providers.base import ( LLMProvider, LLMResponse, + LLMUsage, ProviderCallContext, ProviderConversationState, ToolCallRequest, @@ -517,8 +518,9 @@ class OpenAICompatProvider(LLMProvider): api_type: str = "auto", extra_query: dict[str, str] | None = None, proxy: str | None = None, + provider_name: str = "openai", ): - super().__init__(api_key, api_base) + super().__init__(api_key, api_base, provider_name=provider_name) self.default_model = default_model self.extra_headers = extra_headers or {} self._spec = spec @@ -1428,12 +1430,12 @@ class OpenAICompatProvider(LLMProvider): return "".join(parts) or None @classmethod - def _extract_usage(cls, response: Any) -> dict[str, int]: + def _extract_usage(cls, response: Any) -> LLMUsage | None: """Extract token usage from an OpenAI-compatible response. Handles both dict-based (raw JSON) and object-based (SDK Pydantic) - responses. Provider-specific ``cached_tokens`` fields are normalised - under a single key; see the priority chain inside for details. + responses. Provider-specific cache fields are normalized once at + this Chat Completions wire boundary. """ # --- resolve usage object --- usage_obj = None @@ -1445,21 +1447,18 @@ class OpenAICompatProvider(LLMProvider): usage_map = cls._maybe_mapping(usage_obj) if usage_map is not None: - result = { - "prompt_tokens": int(usage_map.get("prompt_tokens") or 0), - "completion_tokens": int(usage_map.get("completion_tokens") or 0), - "total_tokens": int(usage_map.get("total_tokens") or 0), - } + input_tokens = int(usage_map.get("prompt_tokens") or 0) + output_tokens = int(usage_map.get("completion_tokens") or 0) elif usage_obj: - result = { - "prompt_tokens": getattr(usage_obj, "prompt_tokens", 0) or 0, - "completion_tokens": getattr(usage_obj, "completion_tokens", 0) or 0, - "total_tokens": getattr(usage_obj, "total_tokens", 0) or 0, - } + input_tokens = int(getattr(usage_obj, "prompt_tokens", 0) or 0) + output_tokens = int(getattr(usage_obj, "completion_tokens", 0) or 0) else: - return {} + return None - # --- cached_tokens (normalised across providers) --- + wire_total = cls._get_nested_int(usage_obj, ("total_tokens",)) + + cache_read: int | None = None + # --- cached_tokens (normalised across Chat-compatible providers) --- # Try nested paths first (dict), fall back to attribute (SDK object). # Priority order ensures the most specific field wins. for path in ( @@ -1468,17 +1467,28 @@ class OpenAICompatProvider(LLMProvider): ("prompt_cache_hit_tokens",), # DeepSeek/SiliconFlow ): cached = cls._get_nested_int(usage_map, path) - if not cached and usage_obj: + if cached is None and usage_obj: cached = cls._get_nested_int(usage_obj, path) - if cached: - result["cached_tokens"] = cached + if cached is not None: + cache_read = cached break - return result + cache_write = cls._get_nested_int( + usage_obj, + ("prompt_tokens_details", "cache_write_tokens"), + ) + + return LLMUsage.reported( + input_tokens=input_tokens, + output_tokens=output_tokens, + total_tokens=wire_total, + cache_read_tokens=cache_read, + cache_write_tokens=cache_write, + ) @staticmethod - def _get_nested_int(obj: object, path: tuple[str, ...]) -> int: - """Drill into *obj* by *path* segments and return an ``int`` value. + def _get_nested_int(obj: object, path: tuple[str, ...]) -> int | None: + """Return a present usage count while preserving explicit zero. Supports both dict-key access and attribute access so it works uniformly with raw JSON dicts **and** SDK Pydantic models. @@ -1486,12 +1496,17 @@ class OpenAICompatProvider(LLMProvider): current: object = obj for segment in path: if current is None: - return 0 + return None if isinstance(current, dict): current = cast(dict[str, Any], current).get(segment) else: current = getattr(current, segment, None) - return int(cast(Any, current) or 0) if current is not None else 0 + if current is None or isinstance(current, bool): + return None + try: + return int(cast(Any, current)) + except (TypeError, ValueError): + return None def _parse(self, response: Any) -> LLMResponse: if isinstance(response, str): @@ -1645,7 +1660,7 @@ class OpenAICompatProvider(LLMProvider): reasoning_parts: list[str] = [] tc_bufs: dict[int, dict[str, Any]] = {} finish_reason = "stop" - usage: dict[str, int] = {} + usage: LLMUsage | None = None def _accum_tc(tc: Any, idx_hint: int) -> None: """Accumulate one streaming tool-call delta into *tc_bufs*.""" diff --git a/nanobot/providers/openai_responses/parsing.py b/nanobot/providers/openai_responses/parsing.py index b06c93211..b5676a70a 100644 --- a/nanobot/providers/openai_responses/parsing.py +++ b/nanobot/providers/openai_responses/parsing.py @@ -10,7 +10,7 @@ from typing import Any, AsyncGenerator, cast import httpx from loguru import logger -from nanobot.providers.base import LLMResponse, ToolCallRequest, parse_tool_arguments +from nanobot.providers.base import LLMResponse, LLMUsage, ToolCallRequest, parse_tool_arguments from nanobot.providers.openai_responses.state import build_responses_state FINISH_REASON_MAP = { @@ -186,33 +186,40 @@ def _response_finish_reason( return map_finish_reason(terminal_status) -def _usage_from_response_obj(response: object) -> dict[str, int]: +def _usage_from_response_obj(response: object) -> LLMUsage | None: response_object = _response_object(response) usage_raw: object = ( response_object.get("usage") if response_object is not None else getattr(response, "usage", None) ) - if not usage_raw: - return {} + if usage_raw is None: + return None usage = _response_object(usage_raw) if usage is None: - return {} - prompt_tokens = int(usage.get("input_tokens") or usage.get("prompt_tokens") or 0) - completion_tokens = int( - usage.get("output_tokens") or usage.get("completion_tokens") or 0 - ) - total_tokens = int(usage.get("total_tokens") or prompt_tokens + completion_tokens) - result = { - "prompt_tokens": prompt_tokens, - "completion_tokens": completion_tokens, - "total_tokens": total_tokens, - } + return None + + def _usage_int(container: dict[str, Any] | None, key: str) -> int | None: + if container is None: + return None + raw = container.get(key) + if raw is None or isinstance(raw, bool): + return None + try: + return int(raw) + except (TypeError, ValueError): + return None + + input_tokens = _usage_int(usage, "input_tokens") or 0 + output_tokens = _usage_int(usage, "output_tokens") or 0 input_details = _response_object(usage.get("input_tokens_details")) - cached_tokens = int(input_details.get("cached_tokens") or 0) if input_details else 0 - if cached_tokens > 0: - result["cached_tokens"] = cached_tokens - return result + return LLMUsage.reported( + input_tokens=input_tokens, + output_tokens=output_tokens, + total_tokens=_usage_int(usage, "total_tokens"), + cache_read_tokens=_usage_int(input_details, "cached_tokens"), + cache_write_tokens=_usage_int(input_details, "cache_write_tokens"), + ) def _parse_tool_call_arguments(args_raw: Any, name: str | None) -> Any: @@ -352,14 +359,14 @@ async def consume_sse_with_reasoning( on_reasoning_delta: Callable[[str], Awaitable[None]] | None = None, on_response_event: Callable[[dict[str, Any]], Awaitable[None]] | None = None, capture: ResponsesStreamCapture | None = None, -) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]: +) -> tuple[str, list[ToolCallRequest], str, LLMUsage | None, str | None]: """Consume a Responses API SSE stream, including visible reasoning summaries.""" content = "" tool_calls: list[ToolCallRequest] = [] tool_call_buffers: dict[str, dict[str, Any]] = {} tool_call_args_emitted: set[str] = set() finish_reason = "stop" - usage: dict[str, int] = {} + usage: LLMUsage | None = None reasoning_content: str | None = None streamed_reasoning = False reasoning_summary_key: tuple[str | None, int] | None = None @@ -657,14 +664,14 @@ async def consume_sdk_stream( on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None, on_reasoning_delta: Callable[[str], Awaitable[None]] | None = None, capture: ResponsesStreamCapture | None = None, -) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]: +) -> tuple[str, list[ToolCallRequest], str, LLMUsage | None, str | None]: """Consume an SDK async stream from ``client.responses.create(stream=True)``.""" content = "" tool_calls: list[ToolCallRequest] = [] tool_call_buffers: dict[str, dict[str, Any]] = {} tool_call_args_emitted: set[str] = set() finish_reason = "stop" - usage: dict[str, int] = {} + usage: LLMUsage | None = None reasoning_content: str | None = None streamed_reasoning = False refusal_seen = False @@ -823,20 +830,7 @@ async def consume_sdk_stream( if on_content_delta and remaining_text: await on_content_delta(remaining_text) if resp: - usage_obj = getattr(resp, "usage", None) - if usage_obj: - usage = { - "prompt_tokens": int(getattr(usage_obj, "input_tokens", 0) or 0), - "completion_tokens": int(getattr(usage_obj, "output_tokens", 0) or 0), - "total_tokens": int(getattr(usage_obj, "total_tokens", 0) or 0), - } - usage_data = _response_object(usage_obj) or {} - input_details = _response_object(usage_data.get("input_tokens_details")) - cached_tokens = ( - int(input_details.get("cached_tokens") or 0) if input_details else 0 - ) - if cached_tokens > 0: - usage["cached_tokens"] = cached_tokens + usage = _usage_from_response_obj(resp) or usage if not reasoning_content: reasoning_content = _extract_reasoning_summary_from_output( getattr(resp, "output", None) diff --git a/nanobot/providers/openai_responses/state.py b/nanobot/providers/openai_responses/state.py index a2a95d70f..a424c66d8 100644 --- a/nanobot/providers/openai_responses/state.py +++ b/nanobot/providers/openai_responses/state.py @@ -7,7 +7,7 @@ from typing import Any, cast from loguru import logger -from nanobot.providers.base import ProviderConversationState +from nanobot.providers.base import LLMUsage, ProviderConversationState from nanobot.providers.openai_responses.converters import convert_messages RESPONSES_STATE_KIND = "openai_responses" @@ -84,7 +84,7 @@ def build_responses_state( model: str, input_items: list[dict[str, Any]], output_items: list[dict[str, Any]], - usage: dict[str, int] | None = None, + usage: LLMUsage | None = None, ) -> ProviderConversationState: """Create the canonical next state from request input and every output item.""" unpruned_items = [*input_items, *output_items] @@ -178,16 +178,8 @@ def _prune_before_latest_output_compaction( return output_items[latest:] -def _context_tokens_from_usage(usage: dict[str, int] | None) -> int: - if not usage: - return 0 - prompt_tokens = usage.get("prompt_tokens", 0) - completion_tokens = usage.get("completion_tokens", 0) - total_tokens = usage.get("total_tokens", 0) - values = (prompt_tokens, completion_tokens, total_tokens) - if any(isinstance(value, bool) for value in values): - return 0 - return max(0, total_tokens or prompt_tokens + completion_tokens) +def _context_tokens_from_usage(usage: LLMUsage | None) -> int: + return usage.total_tokens if usage is not None else 0 def _state_items( diff --git a/nanobot/providers/unconfigured_provider.py b/nanobot/providers/unconfigured_provider.py index 98d7b69ec..9a3f6420b 100644 --- a/nanobot/providers/unconfigured_provider.py +++ b/nanobot/providers/unconfigured_provider.py @@ -11,7 +11,7 @@ class UnconfiguredProvider(LLMProvider): """Keep the gateway available for settings before a model is configured.""" def __init__(self, default_model: str) -> None: - super().__init__() + super().__init__(provider_name="unconfigured") self._default_model = default_model async def chat( diff --git a/nanobot/providers/xai_grok_provider.py b/nanobot/providers/xai_grok_provider.py index 86d0cf945..fc6864acc 100644 --- a/nanobot/providers/xai_grok_provider.py +++ b/nanobot/providers/xai_grok_provider.py @@ -18,6 +18,7 @@ from nanobot import __version__ from nanobot.providers.base import ( LLMProvider, LLMResponse, + LLMUsage, ToolCallRequest, resolve_stream_idle_timeout_s, ) @@ -69,8 +70,10 @@ class XAIGrokProvider(LLMProvider): default_model: str = DEFAULT_XAI_GROK_MODEL, proxy: str | None = None, extra_body: dict[str, Any] | None = None, + *, + provider_name: str = "xai_grok", ): - super().__init__(api_key=None, api_base=None) + super().__init__(api_key=None, api_base=None, provider_name=provider_name) self.default_model = default_model self.proxy = proxy or None self._extra_body = dict(extra_body or {}) @@ -436,7 +439,7 @@ async def _request_xai( on_content_delta: Callable[[str], Awaitable[None]] | None = None, on_thinking_delta: Callable[[str], Awaitable[None]] | None = None, on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None, -) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]: +) -> tuple[str, list[ToolCallRequest], str, LLMUsage | None, str | None]: async def _on_response_event(event: dict[str, Any]) -> None: hosted_event = _xai_hosted_tool_event(event) if hosted_event is not None and on_tool_call_delta is not None: diff --git a/nanobot/sdk/types.py b/nanobot/sdk/types.py index e28feef99..358cf8aaf 100644 --- a/nanobot/sdk/types.py +++ b/nanobot/sdk/types.py @@ -6,6 +6,7 @@ from copy import deepcopy from dataclasses import dataclass, field from typing import Any, Literal, Mapping, TypeAlias, cast +from nanobot.providers.base import LLMUsage from nanobot.runtime_context import public_history_messages StreamEventType: TypeAlias = Literal[ @@ -53,7 +54,7 @@ class RunResult: content: str tools_used: list[str] = field(default_factory=list) messages: list[dict[str, Any]] = field(default_factory=list) - usage: dict[str, int] = field(default_factory=dict) + usage: LLMUsage | None = None stop_reason: str | None = None error: str | None = None metadata: dict[str, Any] = field(default_factory=dict) @@ -72,7 +73,7 @@ class StreamEvent: arguments: dict[str, Any] | None = None iteration: int | None = None resuming: bool | None = None - usage: dict[str, int] = field(default_factory=dict) + usage: LLMUsage | None = None error: str | None = None metadata: dict[str, Any] = field(default_factory=dict) diff --git a/nanobot/session/webui_turns.py b/nanobot/session/webui_turns.py index a6cf76b50..a4fc5d0b7 100644 --- a/nanobot/session/webui_turns.py +++ b/nanobot/session/webui_turns.py @@ -37,7 +37,7 @@ from nanobot.bus.runtime_events import ( TurnRuntimeAdmitted, UserInputAccepted, ) -from nanobot.providers.base import LLMProvider +from nanobot.providers.base import LLMProvider, LLMUsage from nanobot.providers.fallback_provider import FallbackModelObserver from nanobot.runtime_context import public_history_message from nanobot.session.goal_state import goal_state_ws_blob @@ -695,7 +695,7 @@ class WebuiTurnCoordinator: *, session_key: str, latency_ms: int | None, - usage: dict[str, int] | None = None, + usage: LLMUsage | None = None, context_window_tokens: int | None = None, ) -> None: if msg.channel != "websocket": @@ -709,7 +709,7 @@ class WebuiTurnCoordinator: event=TurnEndEvent( latency_ms=latency_ms, goal_state=goal_state_ws_blob(session.metadata), - usage=usage or None, + usage=usage, context_window_tokens=context_window_tokens, ), metadata=msg.metadata, diff --git a/nanobot/skills/my/references/examples.md b/nanobot/skills/my/references/examples.md index 961813d97..1fa59408e 100644 --- a/nanobot/skills/my/references/examples.md +++ b/nanobot/skills/my/references/examples.md @@ -16,7 +16,7 @@ Concrete scenarios showing when and how to use the my tool effectively. → my(action="check", key="max_iterations") → 40 → my(action="check", key="_last_usage") - → {"prompt_tokens": 62000, "completion_tokens": 3000} + → {"input_tokens": 62000, "output_tokens": 3000} → "I hit the iteration limit (40). The task was complex. I can ask the user if they want to increase it." ``` @@ -72,6 +72,6 @@ Concrete scenarios showing when and how to use the my tool effectively. ### Token-conscious behavior ``` → my(action="check", key="_last_usage") - → {"prompt_tokens": 58000, "completion_tokens": 12000} + → {"input_tokens": 58000, "output_tokens": 12000} → "I've consumed ~70k tokens. I'll keep my remaining responses focused." ``` diff --git a/nanobot/utils/helpers.py b/nanobot/utils/helpers.py index 1f1720dd3..de6f89862 100644 --- a/nanobot/utils/helpers.py +++ b/nanobot/utils/helpers.py @@ -1,5 +1,7 @@ """Utility functions for nanobot.""" +from __future__ import annotations + import base64 import json import os @@ -12,11 +14,14 @@ from contextlib import suppress from datetime import datetime from functools import lru_cache from pathlib import Path -from typing import Any, TypeVar, cast, overload +from typing import TYPE_CHECKING, Any, TypeVar, cast, overload import tiktoken from loguru import logger +if TYPE_CHECKING: + from nanobot.providers.base import LLMUsage + _TOOLS_TOKEN_CACHE_MAX_ENTRIES = 64 _TOOLS_TOKEN_CACHE: dict[int, tuple[tuple[int, ...], dict[bool, int]]] = {} _T = TypeVar("_T") @@ -793,7 +798,7 @@ def build_status_content( version: str, model: str, start_time: float, - last_usage: dict[str, int], + last_usage: LLMUsage | None, context_window_tokens: int, session_msg_count: int, context_tokens_estimate: int, @@ -814,9 +819,9 @@ def build_status_content( if uptime_s >= 3600 else f"{uptime_s // 60}m {uptime_s % 60}s" ) - last_in = last_usage.get("prompt_tokens", 0) - last_out = last_usage.get("completion_tokens", 0) - cached = last_usage.get("cached_tokens", 0) + last_in = last_usage.input_tokens if last_usage else 0 + last_out = last_usage.output_tokens if last_usage else 0 + cached = last_usage.cache_read_tokens if last_usage else None ctx_total = max(context_window_tokens, 0) # Budget mirrors Consolidator formula: ctx_window - max_completion - _SAFETY_BUFFER ctx_budget = max(ctx_total - int(max_completion_tokens) - 1024, 1) diff --git a/nanobot/webui/session_context.py b/nanobot/webui/session_context.py index 54f5b57f9..5ec3b18ab 100644 --- a/nanobot/webui/session_context.py +++ b/nanobot/webui/session_context.py @@ -4,6 +4,7 @@ from __future__ import annotations from typing import Any, cast +from nanobot.providers.base import LLMUsage from nanobot.session.manager import Session from nanobot.utils.helpers import estimate_message_tokens, truncate_text @@ -36,18 +37,8 @@ def session_context_payload(session: Session) -> dict[str, Any]: summary_tokens = ( estimate_message_tokens({"role": "system", "content": summary}) if summary else 0 ) - raw_usage = session.metadata.get("_last_usage") - last_usage = ( - { - key: value - for key, value in cast(dict[object, object], raw_usage).items() - if isinstance(key, str) - and type(value) is int - and value >= 0 - } - if isinstance(raw_usage, dict) - else None - ) + stored_usage = LLMUsage.from_dict(session.metadata.get("_last_usage")) + last_usage = stored_usage.to_dict() if stored_usage is not None else None return { "schema_version": 1, diff --git a/nanobot/webui/token_usage.py b/nanobot/webui/token_usage.py index f06e9421f..59264e620 100644 --- a/nanobot/webui/token_usage.py +++ b/nanobot/webui/token_usage.py @@ -15,19 +15,23 @@ from loguru import logger from nanobot.agent.hook import AgentHook, AgentHookContext from nanobot.config.paths import get_webui_dir +from nanobot.providers.base import LLMUsage -TOKEN_USAGE_SCHEMA_VERSION = 1 +TOKEN_USAGE_SCHEMA_VERSION = 2 _MAX_STATE_FILE_BYTES = 512 * 1024 _MAX_DAYS_RETAINED = 400 _USAGE_KEYS = ( - "prompt_tokens", - "completion_tokens", - "cached_tokens", + "input_tokens", + "output_tokens", + "cache_read_tokens", + "cache_write_tokens", + "cache_read_observed_input_tokens", + "cache_write_observed_input_tokens", "total_tokens", - "provider_tokens", + "reported_tokens", "estimated_tokens", ) -_REQUEST_KEYS = ("requests", "provider_requests", "estimated_requests") +_REQUEST_KEYS = ("requests", "reported_requests", "estimated_requests") _SOURCE_KEYS = ("user", "api", "cron", "dream", "system") _WRITE_LOCK = threading.Lock() @@ -88,38 +92,43 @@ def _source_from_session_key(session_key: str | None) -> str: return "user" -def _normalize_usage(raw: dict[str, Any] | None) -> dict[str, int]: - if not isinstance(raw, dict): +def _normalize_usage(raw: LLMUsage | None) -> dict[str, int]: + if raw is None: return {} - usage = {key: _clean_int(raw.get(key)) for key in _USAGE_KEYS} - fallback_total = usage["prompt_tokens"] + usage["completion_tokens"] - if usage["total_tokens"] <= 0: - usage["total_tokens"] = fallback_total - if usage["estimated_tokens"] <= 0 and usage["provider_tokens"] <= 0: - usage["provider_tokens"] = usage["total_tokens"] - elif usage["estimated_tokens"] > 0 and usage["provider_tokens"] <= 0: - usage["estimated_tokens"] = min(usage["estimated_tokens"], usage["total_tokens"]) - elif usage["provider_tokens"] > 0 and usage["estimated_tokens"] <= 0: - usage["provider_tokens"] = min(usage["provider_tokens"], usage["total_tokens"]) + usage = { + "input_tokens": raw.input_tokens, + "output_tokens": raw.output_tokens, + "cache_read_tokens": raw.cache_read_tokens or 0, + "cache_write_tokens": raw.cache_write_tokens or 0, + "cache_read_observed_input_tokens": ( + raw.input_tokens if raw.cache_read_tokens is not None else 0 + ), + "cache_write_observed_input_tokens": ( + raw.input_tokens if raw.cache_write_tokens is not None else 0 + ), + "total_tokens": raw.total_tokens, + "reported_tokens": raw.reported_tokens, + "estimated_tokens": raw.estimated_tokens, + } return usage if usage["total_tokens"] > 0 else {} def _normalize_usage_row(row: dict[str, Any]) -> dict[str, int]: cleaned = {key: _clean_int(row.get(key)) for key in _USAGE_KEYS} if cleaned["total_tokens"] <= 0: - cleaned["total_tokens"] = cleaned["prompt_tokens"] + cleaned["completion_tokens"] - if cleaned["provider_tokens"] <= 0 and cleaned["estimated_tokens"] <= 0: - cleaned["provider_tokens"] = cleaned["total_tokens"] + cleaned["total_tokens"] = cleaned["input_tokens"] + cleaned["output_tokens"] + if cleaned["reported_tokens"] <= 0 and cleaned["estimated_tokens"] <= 0: + cleaned["reported_tokens"] = cleaned["total_tokens"] requests = {key: _clean_int(row.get(key)) for key in _REQUEST_KEYS} if ( requests["requests"] > 0 - and requests["provider_requests"] <= 0 + and requests["reported_requests"] <= 0 and requests["estimated_requests"] <= 0 ): - if cleaned["estimated_tokens"] > 0 and cleaned["provider_tokens"] <= 0: + if cleaned["estimated_tokens"] > 0 and cleaned["reported_tokens"] <= 0: requests["estimated_requests"] = requests["requests"] else: - requests["provider_requests"] = requests["requests"] + requests["reported_requests"] = requests["requests"] return {**cleaned, **requests} @@ -150,6 +159,8 @@ def normalize_token_usage_state(raw: Any) -> dict[str, Any]: if not isinstance(raw, dict): return state raw = cast(dict[str, Any], raw) + if raw.get("schema_version") != TOKEN_USAGE_SCHEMA_VERSION: + return state days_raw = raw.get("days") if not isinstance(days_raw, dict): return state @@ -197,24 +208,35 @@ def read_token_usage_state() -> dict[str, Any]: return normalize_token_usage_state(raw) -def write_token_usage_state(raw: dict[str, Any]) -> dict[str, Any]: - state = normalize_token_usage_state(raw) - state["updated_at"] = _utc_now_iso() - encoded = json.dumps( +def _encode_token_usage_state(state: dict[str, Any]) -> bytes: + """Encode the persisted state compactly, including its trailing newline.""" + payload = json.dumps( state, ensure_ascii=False, - indent=2, + separators=(",", ":"), sort_keys=True, - ).encode("utf-8") + ) + return f"{payload}\n".encode("utf-8") + + +def write_token_usage_state(raw: dict[str, Any]) -> dict[str, Any]: + # Day-count retention is applied by normalization first. The byte budget + # then trims only the oldest remaining days, preserving a contiguous suffix. + state = normalize_token_usage_state(raw) + state["updated_at"] = _utc_now_iso() + days = cast(dict[str, dict[str, Any]], state["days"]) + encoded = _encode_token_usage_state(state) + while len(encoded) > _MAX_STATE_FILE_BYTES and len(days) > 1: + del days[min(days)] + encoded = _encode_token_usage_state(state) if len(encoded) > _MAX_STATE_FILE_BYTES: - raise ValueError("token usage state is too large") + raise ValueError("latest token usage day exceeds the state byte limit") path = token_usage_state_path() path.parent.mkdir(parents=True, exist_ok=True) tmp = path.with_suffix(".json.tmp") with open(tmp, "wb") as f: f.write(encoded) - f.write(b"\n") f.flush() os.fsync(f.fileno()) os.replace(tmp, path) @@ -230,7 +252,7 @@ def write_token_usage_state(raw: dict[str, Any]) -> dict[str, Any]: def record_token_usage( - usage: dict[str, Any] | None, + usage: LLMUsage | None, *, source: str = "user", timezone_name: str | None = None, @@ -248,10 +270,10 @@ def record_token_usage( for key in _USAGE_KEYS: row[key] = _clean_int(row.get(key)) + normalized.get(key, 0) row["requests"] = _clean_int(row.get("requests")) + 1 - if normalized.get("estimated_tokens", 0) > 0 and normalized.get("provider_tokens", 0) <= 0: + if normalized.get("estimated_tokens", 0) > 0 and normalized.get("reported_tokens", 0) <= 0: row["estimated_requests"] = _clean_int(row.get("estimated_requests")) + 1 else: - row["provider_requests"] = _clean_int(row.get("provider_requests")) + 1 + row["reported_requests"] = _clean_int(row.get("reported_requests")) + 1 source_key = _clean_source(source) sources: dict[str, dict[str, Any]] = dict( @@ -261,10 +283,10 @@ def record_token_usage( for key in _USAGE_KEYS: source_row[key] = _clean_int(source_row.get(key)) + normalized.get(key, 0) source_row["requests"] = _clean_int(source_row.get("requests")) + 1 - if normalized.get("estimated_tokens", 0) > 0 and normalized.get("provider_tokens", 0) <= 0: + if normalized.get("estimated_tokens", 0) > 0 and normalized.get("reported_tokens", 0) <= 0: source_row["estimated_requests"] = _clean_int(source_row.get("estimated_requests")) + 1 else: - source_row["provider_requests"] = _clean_int(source_row.get("provider_requests")) + 1 + source_row["reported_requests"] = _clean_int(source_row.get("reported_requests")) + 1 sources[source_key] = source_row row["sources"] = sources diff --git a/tests/agent/test_attachment_references.py b/tests/agent/test_attachment_references.py index 75b9bc2d8..875f80877 100644 --- a/tests/agent/test_attachment_references.py +++ b/tests/agent/test_attachment_references.py @@ -128,7 +128,7 @@ async def test_pending_document_attachment_keeps_body_out_of_prompt( nonlocal call_count call_count += 1 captured_messages.append([dict(message) for message in messages]) - return LLMResponse(content=f"answer-{call_count}", tool_calls=[], usage={}) + return LLMResponse(content=f"answer-{call_count}", tool_calls=[], usage=None) loop = _make_loop(workspace) loop.provider.chat_with_retry = chat_with_retry diff --git a/tests/agent/test_dream.py b/tests/agent/test_dream.py index 772898111..f14284fc6 100644 --- a/tests/agent/test_dream.py +++ b/tests/agent/test_dream.py @@ -412,7 +412,7 @@ class TestEphemeralDirect: provider.supports_tools = True provider.generation = MagicMock(max_tokens=4096) provider.chat_with_retry = AsyncMock( - return_value=LLMResponse(content="done", tool_calls=[], finish_reason="stop", usage={}) + return_value=LLMResponse(content="done", tool_calls=[], finish_reason="stop", usage=None) ) with ( @@ -556,9 +556,9 @@ class TestEphemeralDirect: "new_text": "replacement", }, )], - usage={}, + usage=None, ), - LLMResponse(content="done", finish_reason="stop", tool_calls=[], usage={}), + LLMResponse(content="done", finish_reason="stop", tool_calls=[], usage=None), ]) resp = await loop.process_direct( @@ -646,7 +646,7 @@ class TestEphemeralHooks: provider.generation = MagicMock(max_tokens=4096) provider.chat_with_retry = AsyncMock( return_value=LLMResponse( - content="done", finish_reason="stop", tool_calls=[], usage={}, + content="done", finish_reason="stop", tool_calls=[], usage=None, ) ) diff --git a/tests/agent/test_evaluator.py b/tests/agent/test_evaluator.py index f5f2dff7a..22c568e73 100644 --- a/tests/agent/test_evaluator.py +++ b/tests/agent/test_evaluator.py @@ -12,7 +12,7 @@ from nanobot.utils.evaluator import ( class DummyProvider(LLMProvider): def __init__(self, responses: list[LLMResponse]): - super().__init__() + super().__init__(provider_name="dummy") self._responses = list(responses) async def chat(self, *args, **kwargs) -> LLMResponse: diff --git a/tests/agent/test_history_replay.py b/tests/agent/test_history_replay.py index 04372c3ca..e9e7ce634 100644 --- a/tests/agent/test_history_replay.py +++ b/tests/agent/test_history_replay.py @@ -69,7 +69,7 @@ def test_explicit_message_limit_still_starts_at_user_turn() -> None: async def test_process_message_replays_with_token_budget_only(tmp_path: Path) -> None: loop = _make_loop(tmp_path, context_window_tokens=32_768) loop.provider.chat_with_retry = AsyncMock( - return_value=LLMResponse(content="ok", tool_calls=[], usage={}) + return_value=LLMResponse(content="ok", tool_calls=[], usage=None) ) loop.tools.get_definitions = MagicMock(return_value=[]) loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign] @@ -91,7 +91,7 @@ async def test_process_message_replays_with_token_budget_only(tmp_path: Path) -> async def test_token_budget_keeps_current_user_as_replay_boundary(tmp_path: Path) -> None: loop = _make_loop(tmp_path, context_window_tokens=8_000) loop.provider.chat_with_retry = AsyncMock( - return_value=LLMResponse(content="ok", tool_calls=[], usage={}) + return_value=LLMResponse(content="ok", tool_calls=[], usage=None) ) loop.tools.get_definitions = MagicMock(return_value=[]) loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign] diff --git a/tests/agent/test_hook_composite.py b/tests/agent/test_hook_composite.py index 51bca6973..229029ee5 100644 --- a/tests/agent/test_hook_composite.py +++ b/tests/agent/test_hook_composite.py @@ -453,7 +453,7 @@ async def test_agent_loop_extra_hook_receives_calls(tmp_path): loop = _make_loop(tmp_path, hooks=[TrackingHook()]) loop.provider.chat_with_retry = AsyncMock( - return_value=LLMResponse(content="done", tool_calls=[], usage={}) + return_value=LLMResponse(content="done", tool_calls=[], usage=None) ) loop.tools.get_definitions = MagicMock(return_value=[]) @@ -494,7 +494,7 @@ async def test_agent_loop_turn_hook_factories_receive_context(tmp_path): loop = _make_loop(tmp_path, hook_factories=[factory("registered")]) loop.provider.chat_with_retry = AsyncMock( - return_value=LLMResponse(content="done", tool_calls=[], usage={}) + return_value=LLMResponse(content="done", tool_calls=[], usage=None) ) loop.tools.get_definitions = MagicMock(return_value=[]) @@ -541,7 +541,7 @@ async def test_agent_loop_extra_hook_error_isolation(tmp_path): loop = _make_loop(tmp_path, hooks=[BadHook()]) loop.provider.chat_with_retry = AsyncMock( - return_value=LLMResponse(content="still works", tool_calls=[], usage={}) + return_value=LLMResponse(content="still works", tool_calls=[], usage=None) ) loop.tools.get_definitions = MagicMock(return_value=[]) @@ -562,7 +562,7 @@ async def test_agent_loop_extra_hooks_do_not_swallow_loop_hook_errors(tmp_path): loop.provider.chat_with_retry = AsyncMock(return_value=LLMResponse( content="working", tool_calls=[ToolCallRequest(id="c1", name="list_dir", arguments={"path": "."})], - usage={}, + usage=None, )) loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.execute = AsyncMock(return_value="ok") diff --git a/tests/agent/test_loop_progress.py b/tests/agent/test_loop_progress.py index cb80be1cb..92e9b14c5 100644 --- a/tests/agent/test_loop_progress.py +++ b/tests/agent/test_loop_progress.py @@ -393,9 +393,9 @@ class TestToolEventProgress: }, ) ], - usage={}, + usage=None, ) - return LLMResponse(content="Done", tool_calls=[], usage={}) + return LLMResponse(content="Done", tool_calls=[], usage=None) provider.chat_stream_with_retry = chat_stream_with_retry provider.chat_with_retry = AsyncMock() diff --git a/tests/agent/test_loop_runner_integration.py b/tests/agent/test_loop_runner_integration.py index 8f5231363..62aacce21 100644 --- a/tests/agent/test_loop_runner_integration.py +++ b/tests/agent/test_loop_runner_integration.py @@ -48,7 +48,7 @@ async def test_ephemeral_runner_enters_and_restores_turn_scopes(tmp_path): async def chat_with_retry(**_kwargs): assert goal_mutation_allowed() is True - return LLMResponse(content="done", tool_calls=[], usage={}) + return LLMResponse(content="done", tool_calls=[], usage=None) loop.provider.chat_with_retry = AsyncMock(side_effect=chat_with_retry) loop.tools.get_definitions = MagicMock(return_value=[]) @@ -83,7 +83,7 @@ async def test_goal_command_can_implement_plan_from_prior_discussion(tmp_path): }, ) ], - usage={}, + usage=None, ), LLMResponse( content="closing goal", @@ -94,7 +94,7 @@ async def test_goal_command_can_implement_plan_from_prior_discussion(tmp_path): arguments={"action": "complete", "recap": "Implemented and tested."}, ) ], - usage={}, + usage=None, ), LLMResponse( content="trying to start another goal", @@ -105,9 +105,9 @@ async def test_goal_command_can_implement_plan_from_prior_discussion(tmp_path): arguments={"objective": "Start an unrelated follow-up."}, ) ], - usage={}, + usage=None, ), - LLMResponse(content="done", tool_calls=[], usage={}), + LLMResponse(content="done", tool_calls=[], usage=None), ]) loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=None) @@ -160,8 +160,8 @@ async def test_runtime_context_is_persisted_as_next_turn_prompt_prefix(tmp_path) provider.get_default_model.return_value = "test-model" provider.generation = GenerationSettings() provider.chat_with_retry = AsyncMock(side_effect=[ - LLMResponse(content="first answer", usage={}), - LLMResponse(content="second answer", usage={}), + LLMResponse(content="first answer", usage=None), + LLMResponse(content="second answer", usage=None), ]) loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=None) @@ -216,7 +216,7 @@ async def test_webui_quote_reaches_model_without_leaking_into_public_history(tmp provider = MagicMock() provider.get_default_model.return_value = "test-model" provider.generation = GenerationSettings() - provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="answer", usage={})) + provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="answer", usage=None)) loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=None) session = loop.sessions.get_or_create("websocket:chat") @@ -258,9 +258,9 @@ async def test_runtime_context_provider_runs_once_across_tool_iterations(tmp_pat name="read_file", arguments={"path": "note.txt"}, )], - usage={}, + usage=None, ), - LLMResponse(content="done", usage={}), + LLMResponse(content="done", usage=None), ]) loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=None) @@ -303,9 +303,9 @@ async def test_non_goal_direct_turn_cannot_reuse_prior_goal_command(tmp_path): arguments={"objective": "Unauthorized persistent objective."}, ) ], - usage={}, + usage=None, ), - LLMResponse(content="handled as a one-time task", tool_calls=[], usage={}), + LLMResponse(content="handled as a one-time task", tool_calls=[], usage=None), ]) loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=None) @@ -383,7 +383,7 @@ async def test_loop_stream_filter_handles_think_only_prefix_without_crashing(tmp async def chat_stream_with_retry(*, on_content_delta, **kwargs): await on_content_delta("hidden") await on_content_delta("Hello") - return LLMResponse(content="hiddenHello", tool_calls=[], usage={}) + return LLMResponse(content="hiddenHello", tool_calls=[], usage=None) loop.provider.chat_stream_with_retry = chat_stream_with_retry @@ -413,7 +413,7 @@ async def test_loop_stream_filter_hides_partial_trailing_think_prefix(tmp_path): async def chat_stream_with_retry(*, on_content_delta, **kwargs): await on_content_delta("Hello hiddenWorld") - return LLMResponse(content="Hello hiddenWorld", tool_calls=[], usage={}) + return LLMResponse(content="Hello hiddenWorld", tool_calls=[], usage=None) loop.provider.chat_stream_with_retry = chat_stream_with_retry @@ -436,7 +436,7 @@ async def test_loop_stream_filter_hides_complete_trailing_think_tag(tmp_path): async def chat_stream_with_retry(*, on_content_delta, **kwargs): await on_content_delta("Hello ") await on_content_delta("hiddenWorld") - return LLMResponse(content="Hello hiddenWorld", tool_calls=[], usage={}) + return LLMResponse(content="Hello hiddenWorld", tool_calls=[], usage=None) loop.provider.chat_stream_with_retry = chat_stream_with_retry @@ -459,8 +459,8 @@ async def test_loop_retries_think_only_final_response(tmp_path): async def chat_with_retry(**kwargs): call_count["n"] += 1 if call_count["n"] == 1: - return LLMResponse(content="hidden", tool_calls=[], usage={}) - return LLMResponse(content="Recovered answer", tool_calls=[], usage={}) + return LLMResponse(content="hidden", tool_calls=[], usage=None) + return LLMResponse(content="Recovered answer", tool_calls=[], usage=None) loop.provider.chat_with_retry = chat_with_retry @@ -485,7 +485,7 @@ async def test_streamed_flag_not_set_on_llm_error(tmp_path): provider.get_default_model.return_value = "test-model" loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") error_resp = LLMResponse( - content="503 service unavailable", finish_reason="error", tool_calls=[], usage={}, + content="503 service unavailable", finish_reason="error", tool_calls=[], usage=None, ) loop.provider.chat_with_retry = AsyncMock(return_value=error_resp) loop.provider.chat_stream_with_retry = AsyncMock(return_value=error_resp) @@ -523,14 +523,14 @@ async def test_ssrf_soft_block_can_finalize_after_streamed_tool_call(tmp_path): name="exec", arguments={"command": "curl http://169.254.169.254/latest/meta-data/"}, )], - usage={}, + usage=None, ) responses = iter([ tool_call_resp, LLMResponse( content="I cannot access private URLs. Please share the local file.", tool_calls=[], - usage={}, + usage=None, ), ]) @@ -569,8 +569,8 @@ async def test_next_turn_after_llm_error_keeps_turn_boundary(tmp_path): provider = MagicMock() provider.get_default_model.return_value = "test-model" provider.chat_with_retry = AsyncMock(side_effect=[ - LLMResponse(content="429 rate limit exceeded", finish_reason="error", tool_calls=[], usage={}), - LLMResponse(content="Recovered answer", tool_calls=[], usage={}), + LLMResponse(content="429 rate limit exceeded", finish_reason="error", tool_calls=[], usage=None), + LLMResponse(content="Recovered answer", tool_calls=[], usage=None), ]) loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") diff --git a/tests/agent/test_loop_save_turn.py b/tests/agent/test_loop_save_turn.py index 25a96a8ce..82953b893 100644 --- a/tests/agent/test_loop_save_turn.py +++ b/tests/agent/test_loop_save_turn.py @@ -20,7 +20,7 @@ from nanobot.bus.outbound_events import ( ) from nanobot.bus.queue import MessageBus from nanobot.cron.session_turns import CRON_HISTORY_META, CRON_TRIGGER_META -from nanobot.providers.base import LLMProvider, LLMResponse, ProviderConversationState +from nanobot.providers.base import LLMProvider, LLMResponse, LLMUsage, ProviderConversationState from nanobot.providers.factory import ProviderSnapshot from nanobot.runtime_context import ( RUNTIME_CONTEXT_HISTORY_META, @@ -1883,7 +1883,7 @@ async def test_turn_usage_is_persisted_with_the_saved_session(tmp_path: Path) -> loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign] async def fake_run_agent_loop(initial_messages, **_kwargs): - loop._last_usage = {"prompt_tokens": 64, "completion_tokens": 9} + loop._last_usage = LLMUsage.reported(input_tokens=64, output_tokens=9) return ( "done", [], @@ -1898,10 +1898,9 @@ async def test_turn_usage_is_persisted_with_the_saved_session(tmp_path: Path) -> ) loop.sessions.invalidate("cli:usage") - assert loop.sessions.get_or_create("cli:usage").metadata["_last_usage"] == { - "prompt_tokens": 64, - "completion_tokens": 9, - } + assert loop.sessions.get_or_create("cli:usage").metadata["_last_usage"] == ( + LLMUsage.reported(input_tokens=64, output_tokens=9).to_dict() + ) @pytest.mark.asyncio diff --git a/tests/agent/test_loop_session_policy.py b/tests/agent/test_loop_session_policy.py index adee6817b..59c372e42 100644 --- a/tests/agent/test_loop_session_policy.py +++ b/tests/agent/test_loop_session_policy.py @@ -30,7 +30,7 @@ def _loop(tmp_path, responses: list[str], **kwargs) -> AgentLoop: provider.get_default_model.return_value = "test-model" provider.generation = GenerationSettings() provider.chat_with_retry = AsyncMock( - side_effect=[LLMResponse(content=response, usage={}) for response in responses] + side_effect=[LLMResponse(content=response, usage=None) for response in responses] ) return AgentLoop( bus=MessageBus(), diff --git a/tests/agent/test_runner_core.py b/tests/agent/test_runner_core.py index e1adc5d76..b16af1a42 100644 --- a/tests/agent/test_runner_core.py +++ b/tests/agent/test_runner_core.py @@ -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 diff --git a/tests/agent/test_runner_errors.py b/tests/agent/test_runner_errors.py index 6125fb06b..11332c093 100644 --- a/tests/agent/test_runner_errors.py +++ b/tests/agent/test_runner_errors.py @@ -90,7 +90,7 @@ async def test_llm_error_not_appended_to_session_messages(): provider = MagicMock(spec=LLMProvider) provider.chat_with_retry = AsyncMock(return_value=LLMResponse( - content="429 rate limit exceeded", finish_reason="error", tool_calls=[], usage={}, + content="429 rate limit exceeded", finish_reason="error", tool_calls=[], usage=None, )) tools = MagicMock() tools.get_definitions.return_value = [] @@ -158,7 +158,7 @@ async def test_runner_ignores_tool_calls_when_finish_reason_blocks_execution( content="Request blocked by provider policy.", finish_reason=finish_reason, tool_calls=[ToolCallRequest(id="call_1", name="exec", arguments={"command": "echo nope"})], - usage={}, + usage=None, )) tools = MagicMock() tools.get_definitions.return_value = [] @@ -189,7 +189,7 @@ async def test_runner_tool_error_sets_final_content(): return LLMResponse( content="working", tool_calls=[ToolCallRequest(id="call_1", name="read_file", arguments={"path": "x"})], - usage={}, + usage=None, ) provider.chat_with_retry = chat_with_retry @@ -224,9 +224,9 @@ async def test_runner_preserves_successful_exec_output_that_starts_with_error(): tool_calls=[ ToolCallRequest(id="call_1", name="exec", arguments={"command": "report"}) ], - usage={}, + usage=None, ) - return LLMResponse(content="done", usage={}) + return LLMResponse(content="done", usage=None) provider.chat_with_retry = chat_with_retry output = "Error: generated report successfully\n\nExit code: 0" @@ -266,7 +266,7 @@ async def test_runner_tool_error_preserves_tool_results_in_messages(): ToolCallRequest(id="tc1", name="read_file", arguments={"path": "a"}), ToolCallRequest(id="tc2", name="exec", arguments={"cmd": "bad"}), ], - usage={}, + usage=None, ) provider.chat_with_retry = chat_with_retry @@ -333,9 +333,9 @@ async def test_length_finish_with_blank_content_routes_to_length_recovery(): content="", finish_reason="length", tool_calls=[ToolCallRequest(id="call_1", name="exec", arguments={})], - usage={}, + usage=None, ), - LLMResponse(content="done", finish_reason="stop", tool_calls=[], usage={}), + LLMResponse(content="done", finish_reason="stop", tool_calls=[], usage=None), ]) tools = MagicMock() tools.get_definitions.return_value = [] diff --git a/tests/agent/test_runner_fallback.py b/tests/agent/test_runner_fallback.py index 649f15811..fd4fcdaa3 100644 --- a/tests/agent/test_runner_fallback.py +++ b/tests/agent/test_runner_fallback.py @@ -78,7 +78,7 @@ class _FakeProvider(LLMProvider): *, responses: list[LLMResponse] | None = None, ): - super().__init__() + super().__init__(provider_name=name) self.name = name self._response = response or _make_response() self._responses = iter(responses) if responses is not None else None @@ -260,6 +260,41 @@ def test_provider_snapshot_uses_smallest_fallback_context_window() -> None: assert snapshot.provider._primary_context_window_tokens == 128000 +def test_factory_injects_configured_identity_into_primary_and_fallback_leaves() -> None: + from nanobot.config.schema import Config + from nanobot.providers.factory import build_provider_snapshot + + config = Config.model_validate({ + "agents": { + "defaults": { + "modelPreset": "primary", + "fallbackModels": ["backup"], + } + }, + "modelPresets": { + "primary": {"model": "primary-model", "provider": "primary_edge"}, + "backup": {"model": "backup-model", "provider": "backup_edge"}, + }, + "providers": { + "primary_edge": { + "apiKey": "primary-key", + "apiBase": "https://primary.example/v1", + }, + "backup_edge": { + "apiKey": "backup-key", + "apiBase": "https://backup.example/v1", + }, + }, + }) + + snapshot = build_provider_snapshot(config) + + assert isinstance(snapshot.provider, FallbackProvider) + assert snapshot.provider._primary.provider_name == "primary_edge" + fallback = snapshot.provider._provider_factory(snapshot.provider._fallback_presets[0]) + assert fallback.provider_name == "backup_edge" + + def test_inline_fallback_reasoning_effort_does_not_inherit_primary() -> None: from nanobot.config.schema import Config from nanobot.providers.factory import provider_signature diff --git a/tests/agent/test_runner_goal_continue.py b/tests/agent/test_runner_goal_continue.py index db86f24dc..0ebcf022b 100644 --- a/tests/agent/test_runner_goal_continue.py +++ b/tests/agent/test_runner_goal_continue.py @@ -25,7 +25,7 @@ async def test_runner_exits_normally_without_predicate(): provider = MagicMock(spec=LLMProvider) provider.chat_with_retry = AsyncMock(return_value=LLMResponse( - content="all done", tool_calls=[], usage={}, + content="all done", tool_calls=[], usage=None, )) tools = MagicMock() tools.get_definitions.return_value = [] @@ -50,7 +50,7 @@ async def test_runner_exits_normally_with_inactive_goal(): provider = MagicMock(spec=LLMProvider) provider.chat_with_retry = AsyncMock(return_value=LLMResponse( - content="all done", tool_calls=[], usage={}, + content="all done", tool_calls=[], usage=None, )) tools = MagicMock() tools.get_definitions.return_value = [] @@ -82,7 +82,7 @@ async def test_runner_forces_continue_when_goal_active(): provider = MagicMock(spec=LLMProvider) provider.chat_with_retry = AsyncMock(return_value=LLMResponse( - content="still working", tool_calls=[], usage={}, + content="still working", tool_calls=[], usage=None, )) tools = MagicMock() tools.get_definitions.return_value = [] @@ -112,7 +112,7 @@ async def test_runner_respects_max_iterations_even_with_active_goal(): provider = MagicMock(spec=LLMProvider) provider.chat_with_retry = AsyncMock(return_value=LLMResponse( - content="still working", tool_calls=[], usage={}, + content="still working", tool_calls=[], usage=None, )) tools = MagicMock() tools.get_definitions.return_value = [] @@ -137,7 +137,7 @@ async def test_runner_goal_continue_not_limited_by_injection_cycle_cap(): provider = MagicMock(spec=LLMProvider) provider.chat_with_retry = AsyncMock(return_value=LLMResponse( - content="still working", tool_calls=[], usage={}, + content="still working", tool_calls=[], usage=None, )) tools = MagicMock() tools.get_definitions.return_value = [] @@ -165,7 +165,7 @@ async def test_runner_does_not_force_continue_on_error(): provider = MagicMock(spec=LLMProvider) provider.chat_with_retry = AsyncMock(return_value=LLMResponse( - content=None, tool_calls=[], usage={}, + content=None, tool_calls=[], usage=None, finish_reason="error", )) tools = MagicMock() @@ -191,7 +191,7 @@ async def test_runner_uses_custom_goal_continue_message(): provider = MagicMock(spec=LLMProvider) provider.chat_with_retry = AsyncMock(return_value=LLMResponse( - content="still working", tool_calls=[], usage={}, + content="still working", tool_calls=[], usage=None, )) tools = MagicMock() tools.get_definitions.return_value = [] @@ -220,7 +220,7 @@ async def test_runner_resolves_goal_continue_message_lazily(): provider = MagicMock(spec=LLMProvider) provider.chat_with_retry = AsyncMock(return_value=LLMResponse( - content="still working", tool_calls=[], usage={}, + content="still working", tool_calls=[], usage=None, )) tools = MagicMock() tools.get_definitions.return_value = [] diff --git a/tests/agent/test_runner_governance.py b/tests/agent/test_runner_governance.py index 29379c99d..88d36a8c3 100644 --- a/tests/agent/test_runner_governance.py +++ b/tests/agent/test_runner_governance.py @@ -273,7 +273,7 @@ async def test_runner_drops_orphan_tool_results_before_model_request(): async def chat_with_retry(*, messages, **kwargs): captured_messages[:] = 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() @@ -312,7 +312,7 @@ async def test_backfill_repairs_model_context_without_shifting_save_turn_boundar provider = MagicMock() provider.get_default_model.return_value = "test-model" - response = LLMResponse(content="new answer", tool_calls=[], usage={}) + response = LLMResponse(content="new answer", tool_calls=[], usage=None) provider.chat_with_retry = AsyncMock(return_value=response) provider.chat_stream_with_retry = AsyncMock(return_value=response) @@ -397,7 +397,7 @@ async def test_runner_backfill_only_mutates_model_context_not_returned_messages( async def chat_with_retry(*, messages, **kwargs): captured_messages[:] = 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() diff --git a/tests/agent/test_runner_hooks.py b/tests/agent/test_runner_hooks.py index 71c4e12e6..288c32bcf 100644 --- a/tests/agent/test_runner_hooks.py +++ b/tests/agent/test_runner_hooks.py @@ -9,7 +9,7 @@ import pytest from agent.runner_helpers import make_run_spec from nanobot.config.schema import AgentDefaults -from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest +from nanobot.providers.base import LLMProvider, LLMResponse, LLMUsage, ToolCallRequest _MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars @@ -30,7 +30,7 @@ async def test_runner_calls_hooks_in_order(): content="thinking", tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})], ) - return LLMResponse(content="done", tool_calls=[], usage={}) + return LLMResponse(content="done", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry tools = MagicMock() @@ -110,7 +110,7 @@ async def test_runner_streaming_hook_receives_deltas_and_end_signal(): async def chat_stream_with_retry(*, on_content_delta, **kwargs): await on_content_delta("he") await on_content_delta("llo") - return LLMResponse(content="hello", tool_calls=[], usage={}) + return LLMResponse(content="hello", tool_calls=[], usage=None) provider.chat_stream_with_retry = chat_stream_with_retry provider.chat_with_retry = AsyncMock() @@ -155,7 +155,7 @@ async def test_runner_measures_stream_generation_without_time_to_first_token(): await on_content_delta("llo") return LLMResponse( content="hello", - usage={"prompt_tokens": 100, "completion_tokens": 12}, + usage=LLMUsage.reported(input_tokens=100, output_tokens=12), ) provider.chat_stream_with_retry = chat_stream_with_retry @@ -181,10 +181,11 @@ async def test_runner_measures_stream_generation_without_time_to_first_token(): hook=StreamingHook(), )) - assert result.usage["generation_ms"] == 600 - assert result.usage["measured_completion_tokens"] == 12 - assert result.usage["ttft_ms"] == 200 - assert result.usage["timed_requests"] == 1 + assert result.usage is not None + assert result.usage.generation_ms == 600 + assert result.usage.measured_output_tokens == 12 + assert result.usage.ttft_ms == 200 + assert result.usage.timed_requests == 1 @pytest.mark.asyncio @@ -240,23 +241,24 @@ async def test_runner_length_recovery_streams_segments_once_and_returns_all_cont @pytest.mark.asyncio -async def test_runner_passes_cached_tokens_to_hook_context(): - """Hook context.usage should contain cached_tokens.""" +async def test_runner_passes_cache_read_tokens_to_hook_context(): + """Hook context usage preserves a reported cache-read count.""" from nanobot.agent.hook import AgentHook, AgentHookContext from nanobot.agent.runner import AgentRunner provider = MagicMock(spec=LLMProvider) - captured_usage: list[dict] = [] + captured_usage: list[LLMUsage] = [] class UsageHook(AgentHook): async def after_iteration(self, context: AgentHookContext) -> None: - captured_usage.append(dict(context.usage)) + assert context.usage is not None + captured_usage.append(context.usage) async def chat_with_retry(**kwargs): 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 @@ -274,8 +276,8 @@ async def test_runner_passes_cached_tokens_to_hook_context(): )) assert len(captured_usage) == 1 - assert captured_usage[0]["cached_tokens"] == 150 - assert captured_usage[0]["provider_tokens"] == 220 + assert captured_usage[0].cache_read_tokens == 150 + assert captured_usage[0].reported_tokens == 220 @pytest.mark.asyncio @@ -284,14 +286,15 @@ async def test_runner_estimates_usage_when_provider_omits_usage(monkeypatch): from nanobot.agent.runner import AgentRunner provider = MagicMock(spec=LLMProvider) - captured_usage: list[dict] = [] + captured_usage: list[LLMUsage] = [] class UsageHook(AgentHook): async def after_iteration(self, context: AgentHookContext) -> None: - captured_usage.append(dict(context.usage)) + assert context.usage is not None + captured_usage.append(context.usage) async def chat_with_retry(**kwargs): - return LLMResponse(content="done", tool_calls=[], usage={}) + return LLMResponse(content="done", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry tools = MagicMock() @@ -312,11 +315,8 @@ async def test_runner_estimates_usage_when_provider_omits_usage(monkeypatch): hook=UsageHook(), )) - assert result.usage["prompt_tokens"] == 123 - assert result.usage["completion_tokens"] == 7 - assert result.usage["total_tokens"] == 130 - assert result.usage["estimated_tokens"] == 130 - assert captured_usage[0]["estimated_tokens"] == 130 + assert result.usage == LLMUsage.estimated(input_tokens=123, output_tokens=7) + assert captured_usage[0].estimated_tokens == 130 @pytest.mark.asyncio @@ -332,7 +332,7 @@ async def test_runner_calls_run_level_hooks_on_success(): return LLMResponse( content="done", tool_calls=[], - usage={"prompt_tokens": 3, "completion_tokens": 2}, + usage=LLMUsage.reported(input_tokens=3, output_tokens=2), ) provider.chat_with_retry = chat_with_retry @@ -350,7 +350,7 @@ async def test_runner_calls_run_level_hooks_on_success(): context.final_content, context.stop_reason, context.error, - dict(context.usage), + context.usage, [msg["role"] for msg in context.messages], )) @@ -379,14 +379,7 @@ async def test_runner_calls_run_level_hooks_on_success(): "done", "completed", None, - { - "prompt_tokens": 3, - "completion_tokens": 2, - "total_tokens": 5, - "provider_tokens": 5, - "request_count": 1, - "context_tokens": 3, - }, + LLMUsage.reported(input_tokens=3, output_tokens=2), ["user", "assistant"], ), ("on_finally", "completed", None), @@ -410,7 +403,7 @@ async def test_runner_run_level_context_is_detached_snapshot(): content="thinking", tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})], ) - return LLMResponse(content="done", tool_calls=[], usage={}) + return LLMResponse(content="done", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry tools = MagicMock() diff --git a/tests/agent/test_runner_injections.py b/tests/agent/test_runner_injections.py index 38a4a977f..0d17c2e5a 100644 --- a/tests/agent/test_runner_injections.py +++ b/tests/agent/test_runner_injections.py @@ -263,9 +263,9 @@ async def test_checkpoint1_injects_after_tool_execution(): return LLMResponse( content="using tool", tool_calls=[ToolCallRequest(id="c1", name="read_file", arguments={"path": "x"})], - usage={}, + usage=None, ) - return LLMResponse(content="final answer", tool_calls=[], usage={}) + return LLMResponse(content="final answer", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry tools = MagicMock() @@ -323,8 +323,8 @@ async def test_checkpoint2_injects_after_final_response_with_resuming_stream(): async def chat_stream_with_retry(*, messages, on_content_delta=None, **kwargs): call_count["n"] += 1 if call_count["n"] == 1: - return LLMResponse(content="first answer", tool_calls=[], usage={}) - return LLMResponse(content="second answer", tool_calls=[], usage={}) + return LLMResponse(content="first answer", tool_calls=[], usage=None) + return LLMResponse(content="second answer", tool_calls=[], usage=None) provider.chat_stream_with_retry = chat_stream_with_retry tools = MagicMock() @@ -411,8 +411,8 @@ async def test_checkpoint2_preserves_final_response_in_history_before_followup() call_count["n"] += 1 captured_messages.append([dict(message) for message in messages]) if call_count["n"] == 1: - return LLMResponse(content="first answer", tool_calls=[], usage={}) - return LLMResponse(content="second answer", tool_calls=[], usage={}) + return LLMResponse(content="first answer", tool_calls=[], usage=None) + return LLMResponse(content="second answer", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry tools = MagicMock() @@ -474,8 +474,8 @@ async def test_loop_injected_followup_preserves_image_media(tmp_path): call_count["n"] += 1 captured_messages.append(list(messages)) if call_count["n"] == 1: - return LLMResponse(content="first answer", tool_calls=[], usage={}) - return LLMResponse(content="second answer", tool_calls=[], usage={}) + return LLMResponse(content="first answer", tool_calls=[], usage=None) + return LLMResponse(content="second answer", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") @@ -528,8 +528,8 @@ async def test_pending_injection_resolves_its_own_runtime_context(tmp_path): provider = MagicMock() provider.get_default_model.return_value = "test-model" provider.chat_with_retry = AsyncMock(side_effect=[ - LLMResponse(content="first answer", tool_calls=[], usage={}), - LLMResponse(content="second answer", tool_calls=[], usage={}), + LLMResponse(content="first answer", tool_calls=[], usage=None), + LLMResponse(content="second answer", tool_calls=[], usage=None), ]) loop = AgentLoop( bus=MessageBus(), @@ -653,8 +653,8 @@ async def test_subagent_pending_injection_is_hidden_history_and_not_merged(tmp_p async def chat_with_retry(*, messages, **kwargs): call_count["n"] += 1 if call_count["n"] == 1: - return LLMResponse(content="first answer", tool_calls=[], usage={}) - return LLMResponse(content="second answer", tool_calls=[], usage={}) + return LLMResponse(content="first answer", tool_calls=[], usage=None) + return LLMResponse(content="second answer", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") @@ -714,8 +714,8 @@ async def test_runner_merges_multiple_injected_user_messages_without_losing_medi call_count["n"] += 1 captured_messages.append([dict(message) for message in messages]) if call_count["n"] == 1: - return LLMResponse(content="first answer", tool_calls=[], usage={}) - return LLMResponse(content="second answer", tool_calls=[], usage={}) + return LLMResponse(content="first answer", tool_calls=[], usage=None) + return LLMResponse(content="second answer", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry tools = MagicMock() @@ -841,7 +841,7 @@ async def test_injection_cycles_capped_at_max(): async def chat_with_retry(*, messages, **kwargs): call_count["n"] += 1 - return LLMResponse(content=f"answer-{call_count['n']}", tool_calls=[], usage={}) + return LLMResponse(content=f"answer-{call_count['n']}", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry tools = MagicMock() @@ -879,7 +879,7 @@ async def test_no_injections_flag_is_false_by_default(): provider = MagicMock() async def chat_with_retry(**kwargs): - return LLMResponse(content="done", tool_calls=[], usage={}) + return LLMResponse(content="done", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry tools = MagicMock() @@ -903,7 +903,7 @@ async def test_pending_queue_cleanup_on_dispatch(tmp_path): loop = _make_loop(tmp_path) async def chat_with_retry(**kwargs): - return LLMResponse(content="done", tool_calls=[], usage={}) + return LLMResponse(content="done", tool_calls=[], usage=None) loop.provider.chat_with_retry = chat_with_retry @@ -1329,7 +1329,7 @@ async def test_pending_queue_preserves_overflow_for_next_injection_cycle(tmp_pat async def chat_with_retry(*, messages, **kwargs): call_count["n"] += 1 captured_messages.append([dict(message) for message in messages]) - return LLMResponse(content=f"answer-{call_count['n']}", tool_calls=[], usage={}) + return LLMResponse(content=f"answer-{call_count['n']}", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") @@ -1502,16 +1502,16 @@ async def test_drain_injections_on_fatal_tool_error(): return LLMResponse( content="stale prefix ", finish_reason="length", - usage={}, + usage=None, ) if call_count["n"] == 2: return LLMResponse( content="", tool_calls=[ToolCallRequest(id="c1", name="exec", arguments={"cmd": "bad"})], - usage={}, + usage=None, ) # Third call: respond normally to the injected follow-up. - return LLMResponse(content="reply to follow-up", tool_calls=[], usage={}) + return LLMResponse(content="reply to follow-up", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry tools = MagicMock() @@ -1563,10 +1563,10 @@ async def test_drain_injections_on_llm_error(): content=None, tool_calls=[], finish_reason="error", - usage={}, + usage=None, ) # Second call: respond normally to the injected follow-up - return LLMResponse(content="recovered answer", tool_calls=[], usage={}) + return LLMResponse(content="recovered answer", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry tools = MagicMock() @@ -1614,9 +1614,9 @@ async def test_drain_injections_on_empty_final_response(): async def chat_with_retry(*, messages, **kwargs): call_count["n"] += 1 if call_count["n"] <= _MAX_EMPTY_RETRIES + 1: - return LLMResponse(content="", tool_calls=[], usage={}) + return LLMResponse(content="", tool_calls=[], usage=None) # After retries exhausted + injection drain, respond normally - return LLMResponse(content="answer after empty", tool_calls=[], usage={}) + return LLMResponse(content="answer after empty", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry tools = MagicMock() @@ -1671,7 +1671,7 @@ async def test_drain_injections_on_max_iterations(): return LLMResponse( content="", tool_calls=[ToolCallRequest(id=f"c{call_count['n']}", name="read_file", arguments={"path": "x"})], - usage={}, + usage=None, ) provider.chat_with_retry = chat_with_retry @@ -1723,7 +1723,7 @@ async def test_drain_injections_set_flag_when_followup_arrives_after_last_iterat return LLMResponse( content="", tool_calls=[ToolCallRequest(id=f"c{call_count['n']}", name="read_file", arguments={"path": "x"})], - usage={}, + usage=None, ) provider.chat_with_retry = chat_with_retry @@ -1786,7 +1786,7 @@ async def test_injection_cycle_cap_on_error_path(): content=None, tool_calls=[], finish_reason="error", - usage={}, + usage=None, ) provider.chat_with_retry = chat_with_retry diff --git a/tests/agent/test_runner_persistence.py b/tests/agent/test_runner_persistence.py index 975c787a4..1b7aa6ae8 100644 --- a/tests/agent/test_runner_persistence.py +++ b/tests/agent/test_runner_persistence.py @@ -8,7 +8,7 @@ from unittest.mock import AsyncMock, MagicMock, patch from agent.runner_helpers import make_run_spec from nanobot.config.schema import AgentDefaults -from nanobot.providers.base import LLMResponse, ToolCallRequest +from nanobot.providers.base import LLMResponse, LLMUsage, ToolCallRequest _MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars @@ -25,10 +25,10 @@ async def test_runner_persists_large_tool_results_for_follow_up_calls(tmp_path): return LLMResponse( content="working", tool_calls=[ToolCallRequest(id="call_big", name="list_dir", arguments={"path": "."})], - 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() @@ -138,10 +138,10 @@ async def test_read_file_result_is_not_offloaded(tmp_path): return LLMResponse( content="reading", tool_calls=[ToolCallRequest(id="call_rf", name="read_file", arguments={"path": "big.txt"})], - 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() @@ -183,10 +183,10 @@ async def test_runner_keeps_going_when_tool_result_persistence_fails(): return LLMResponse( content="working", tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})], - 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() diff --git a/tests/agent/test_runner_progress_deltas.py b/tests/agent/test_runner_progress_deltas.py index 7996e256a..cde3b8308 100644 --- a/tests/agent/test_runner_progress_deltas.py +++ b/tests/agent/test_runner_progress_deltas.py @@ -23,7 +23,7 @@ async def test_runner_can_disable_provider_progress_delta_streaming(): provider = MagicMock() provider.supports_progress_deltas = True provider.chat_with_retry = AsyncMock( - return_value=LLMResponse(content="done", tool_calls=[], usage={}) + return_value=LLMResponse(content="done", tool_calls=[], usage=None) ) provider.chat_stream_with_retry = AsyncMock() tools = MagicMock() @@ -59,7 +59,7 @@ async def test_runner_streams_provider_progress_deltas_by_default(): async def chat_stream_with_retry(*, on_content_delta, **kwargs): await on_content_delta("he") await on_content_delta("llo") - return LLMResponse(content="hello", tool_calls=[], usage={}) + return LLMResponse(content="hello", tool_calls=[], usage=None) provider.chat_stream_with_retry = chat_stream_with_retry provider.chat_with_retry = AsyncMock() @@ -113,7 +113,7 @@ async def test_runner_routes_hosted_tool_events_to_structured_progress(): "result": {"name": "x_semantic_search"}, }) await on_content_delta("done") - return LLMResponse(content="done", tool_calls=[], usage={}) + return LLMResponse(content="done", tool_calls=[], usage=None) provider.chat_stream_with_retry = chat_stream_with_retry provider.chat_with_retry = AsyncMock() @@ -264,9 +264,9 @@ async def test_runner_emits_write_file_diff_from_tool_execution_snapshots(tmp_pa arguments={"path": "big.txt", "content": "line\n" * 24}, ) ], - usage={}, + usage=None, ) - return LLMResponse(content="done", tool_calls=[], usage={}) + return LLMResponse(content="done", tool_calls=[], usage=None) provider.chat_stream_with_retry = chat_stream_with_retry provider.chat_with_retry = AsyncMock() @@ -338,9 +338,9 @@ async def test_runner_emits_edit_file_diff_from_tool_execution_snapshots(tmp_pat }, ) ], - usage={}, + usage=None, ) - return LLMResponse(content="done", tool_calls=[], usage={}) + return LLMResponse(content="done", tool_calls=[], usage=None) provider.chat_stream_with_retry = chat_stream_with_retry provider.chat_with_retry = AsyncMock() @@ -404,9 +404,9 @@ async def test_runner_marks_file_edit_activity_failed_when_tool_errors(tmp_path) arguments={"path": "aborted.txt"}, ) ], - usage={}, + usage=None, ) - return LLMResponse(content="done", tool_calls=[], usage={}) + return LLMResponse(content="done", tool_calls=[], usage=None) provider.chat_stream_with_retry = chat_stream_with_retry provider.chat_with_retry = AsyncMock() @@ -469,7 +469,7 @@ async def test_runner_marks_file_edit_activity_failed_when_cancelled(tmp_path): arguments={"path": "cancelled.txt", "content": "new\n"}, ) ], - usage={}, + usage=None, ) provider.chat_stream_with_retry = chat_stream_with_retry diff --git a/tests/agent/test_runner_reasoning.py b/tests/agent/test_runner_reasoning.py index 048d715e5..3e8a49617 100644 --- a/tests/agent/test_runner_reasoning.py +++ b/tests/agent/test_runner_reasoning.py @@ -16,7 +16,7 @@ import pytest from agent.runner_helpers import make_run_spec from nanobot.agent.hook import AgentHook, AgentHookContext from nanobot.config.schema import AgentDefaults -from nanobot.providers.base import LLMResponse, ToolCallRequest +from nanobot.providers.base import LLMResponse, LLMUsage, ToolCallRequest _MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars @@ -53,10 +53,10 @@ async def test_runner_preserves_reasoning_fields_in_assistant_history(): 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() @@ -99,7 +99,7 @@ async def test_runner_emits_anthropic_thinking_blocks(): {"type": "thinking", "thinking": "After careful consideration.", "signature": "sig2"}, ], tool_calls=[], - usage={"prompt_tokens": 5, "completion_tokens": 3}, + usage=LLMUsage.reported(input_tokens=5, output_tokens=3), ) provider.chat_with_retry = chat_with_retry @@ -135,7 +135,7 @@ async def test_runner_emits_inline_think_content_as_reasoning(): return LLMResponse( content="Let me think about this...\nThe answer is 42.The answer is 42.", tool_calls=[], - usage={"prompt_tokens": 5, "completion_tokens": 3}, + usage=LLMUsage.reported(input_tokens=5, output_tokens=3), ) provider.chat_with_retry = chat_with_retry @@ -171,7 +171,7 @@ async def test_runner_prefers_reasoning_content_over_inline_think(): content="inline thinkingThe answer.", reasoning_content="dedicated reasoning field", tool_calls=[], - usage={"prompt_tokens": 5, "completion_tokens": 3}, + usage=LLMUsage.reported(input_tokens=5, output_tokens=3), ) provider.chat_with_retry = chat_with_retry @@ -211,7 +211,7 @@ async def test_runner_emits_reasoning_content_even_when_answer_was_streamed(): content="The answer.", reasoning_content="step-by-step deduction", tool_calls=[], - usage={"prompt_tokens": 5, "completion_tokens": 3}, + usage=LLMUsage.reported(input_tokens=5, output_tokens=3), ) provider.chat_stream_with_retry = chat_stream_with_retry @@ -257,7 +257,7 @@ async def test_runner_does_not_double_emit_when_inline_think_already_streamed(): return LLMResponse( content="working...The answer.", tool_calls=[], - usage={"prompt_tokens": 5, "completion_tokens": 3}, + usage=LLMUsage.reported(input_tokens=5, output_tokens=3), ) provider.chat_stream_with_retry = chat_stream_with_retry @@ -299,7 +299,7 @@ async def test_runner_closes_reasoning_stream_after_one_shot_response(): content="answer", reasoning_content="hidden thought", tool_calls=[], - usage={"prompt_tokens": 5, "completion_tokens": 3}, + usage=LLMUsage.reported(input_tokens=5, output_tokens=3), ) provider.chat_with_retry = chat_with_retry @@ -350,7 +350,7 @@ async def test_runner_streams_native_thinking_deltas_without_post_hoc_dup(): content="done", tool_calls=[], thinking_blocks=[{"type": "thinking", "thinking": "part1part2"}], - usage={"prompt_tokens": 1, "completion_tokens": 2}, + usage=LLMUsage.reported(input_tokens=1, output_tokens=2), ) provider.chat_stream_with_retry = chat_stream_with_retry @@ -387,7 +387,7 @@ async def test_runner_strips_thinking_tags_from_native_thinking_deltas(): await on_thinking_delta("") if on_content_delta: await on_content_delta("done") - return LLMResponse(content="done", tool_calls=[], usage={}) + return LLMResponse(content="done", tool_calls=[], usage=None) provider.chat_stream_with_retry = chat_stream_with_retry tools = MagicMock() @@ -425,7 +425,7 @@ async def test_runner_ignores_empty_thinking_marker_before_final_reasoning(): content="done", reasoning_content="Preparing final response", tool_calls=[], - usage={}, + usage=None, ) provider.chat_stream_with_retry = chat_stream_with_retry diff --git a/tests/agent/test_runner_tool_execution.py b/tests/agent/test_runner_tool_execution.py index 3d2a85d42..f16127bd6 100644 --- a/tests/agent/test_runner_tool_execution.py +++ b/tests/agent/test_runner_tool_execution.py @@ -106,7 +106,7 @@ async def _run_optional_tool_response(response: LLMResponse): calls["n"] += 1 if calls["n"] == 1: return response - return LLMResponse(content="done", tool_calls=[], usage={}) + return LLMResponse(content="done", tool_calls=[], usage=None) provider.chat_with_retry = chat_with_retry tools = ToolRegistry() @@ -272,10 +272,10 @@ async def test_runner_rejects_near_miss_tool_name_without_executing(): ) ], finish_reason="tool_calls", - 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 = ToolRegistry() @@ -403,7 +403,7 @@ async def test_runner_treats_legacy_entry_point_error_prefix_as_tool_error(tmp_p provider.chat_with_retry = AsyncMock(return_value=LLMResponse( content="working", tool_calls=[ToolCallRequest(id="call_1", name="legacy_plugin", arguments={})], - usage={}, + usage=None, )) result = await AgentRunner().run(make_run_spec(provider, @@ -430,9 +430,9 @@ async def test_runner_preserves_structured_plugin_success_that_starts_with_error tool_calls=[ ToolCallRequest(id="call_1", name="structured_success_plugin", arguments={}) ], - usage={}, + usage=None, ), - LLMResponse(content="done", tool_calls=[], usage={}), + LLMResponse(content="done", tool_calls=[], usage=None), ]) result = await AgentRunner().run(make_run_spec(provider, @@ -466,10 +466,10 @@ async def test_runner_blocks_repeated_external_fetches(): return LLMResponse( content="working", tool_calls=[ToolCallRequest(id=f"call_{call_count['n']}", name="web_fetch", arguments={"url": "https://example.com"})], - usage={}, + usage=None, ) captured_final_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() diff --git a/tests/agent/test_session_inputs.py b/tests/agent/test_session_inputs.py index 7606f54d1..e2f34ef4f 100644 --- a/tests/agent/test_session_inputs.py +++ b/tests/agent/test_session_inputs.py @@ -18,7 +18,7 @@ def _loop(tmp_path: Path) -> AgentLoop: provider.get_default_model.return_value = "test-model" provider.generation = SimpleNamespace(max_tokens=4096) provider.chat_with_retry = AsyncMock( - return_value=LLMResponse(content="Reviewed", tool_calls=[], usage={}) + return_value=LLMResponse(content="Reviewed", tool_calls=[], usage=None) ) return AgentLoop( bus=MessageBus(), diff --git a/tests/agent/test_session_model_runtime.py b/tests/agent/test_session_model_runtime.py index 276067ea2..dac43632a 100644 --- a/tests/agent/test_session_model_runtime.py +++ b/tests/agent/test_session_model_runtime.py @@ -18,7 +18,7 @@ from nanobot.utils.llm_runtime import LLMRuntime class RecordingProvider(LLMProvider): def __init__(self, name: str) -> None: - super().__init__() + super().__init__(provider_name=name) self.name = name self.generation = GenerationSettings(max_tokens=256, temperature=0.1) self.calls: list[str | None] = [] diff --git a/tests/agent/test_subagent_lifecycle.py b/tests/agent/test_subagent_lifecycle.py index c85f37f16..fc406becd 100644 --- a/tests/agent/test_subagent_lifecycle.py +++ b/tests/agent/test_subagent_lifecycle.py @@ -16,7 +16,7 @@ from nanobot.agent.subagent import ( ) from nanobot.agent.tools.context import current_request_context from nanobot.bus.queue import MessageBus -from nanobot.providers.base import GenerationSettings, LLMProvider +from nanobot.providers.base import GenerationSettings, LLMProvider, LLMUsage from nanobot.utils.llm_runtime import LLMRuntime # --------------------------------------------------------------------------- @@ -49,7 +49,7 @@ def _make_hook_context(**overrides) -> AgentHookContext: tool_calls=[], tool_events=[], messages=[], - usage={}, + usage=None, error=None, stop_reason="completed", final_content="ok", @@ -97,7 +97,7 @@ class TestSubagentStatus: assert s.phase == "initializing" assert s.iteration == 0 assert s.tool_events == [] - assert s.usage == {} + assert s.usage is None assert s.stop_reason is None assert s.error is None @@ -677,12 +677,12 @@ class TestSubagentHook: ctx = _make_hook_context( iteration=3, tool_events=[{"name": "read_file", "status": "ok", "detail": ""}], - usage={"prompt_tokens": 100}, + usage=LLMUsage.reported(input_tokens=100, output_tokens=0), ) await hook.after_iteration(ctx) assert status.iteration == 3 assert len(status.tool_events) == 1 - assert status.usage == {"prompt_tokens": 100} + assert status.usage == LLMUsage.reported(input_tokens=100, output_tokens=0) @pytest.mark.asyncio async def test_after_iteration_no_status_noop(self): diff --git a/tests/agent/tools/test_self_tool.py b/tests/agent/tools/test_self_tool.py index 42727f9e7..2a59f9a15 100644 --- a/tests/agent/tools/test_self_tool.py +++ b/tests/agent/tools/test_self_tool.py @@ -15,6 +15,7 @@ from nanobot.agent.tools.self import MyTool from nanobot.agent.tools.shell import ExecToolConfig from nanobot.agent.tools.web import WebSearchConfig, WebToolsConfig from nanobot.config.schema import ModelPresetConfig +from nanobot.providers.base import LLMUsage # --------------------------------------------------------------------------- # Helpers @@ -31,7 +32,7 @@ def _make_mock_loop(**overrides): loop._start_time = 1000.0 loop.exec_config = ExecToolConfig() loop.channels_config = MagicMock() - loop._last_usage = {"prompt_tokens": 100, "completion_tokens": 50} + loop._last_usage = LLMUsage.reported(input_tokens=100, output_tokens=50) loop.last_usage = loop._last_usage loop._current_iteration = 0 loop.current_iteration = loop._current_iteration @@ -163,9 +164,9 @@ class TestInspectPathNavigation: @pytest.mark.asyncio async def test_inspect_dict_key_via_dotpath(self): loop = _make_mock_loop() - loop._last_usage = {"prompt_tokens": 100, "completion_tokens": 50} + loop._last_usage = LLMUsage.reported(input_tokens=100, output_tokens=50) tool = _make_tool(loop=loop) - result = await tool.execute(action="check", key="_last_usage.prompt_tokens") + result = await tool.execute(action="check", key="_last_usage.input_tokens") assert "100" in result @pytest.mark.asyncio @@ -624,7 +625,7 @@ class TestSubagentStatusFormatting: {"name": "grep", "status": "ok", "detail": "searched ERROR"}, {"name": "exec", "status": "error", "detail": "timeout"}, ], - usage={"prompt_tokens": 4500, "completion_tokens": 1200}, + usage=LLMUsage.reported(input_tokens=4500, output_tokens=1200), ) result = MyTool._format_value(status) assert "abc12345" in result @@ -698,14 +699,14 @@ class TestSubagentHookStatus: iteration=5, messages=[], tool_events=[{"name": "read_file", "status": "ok", "detail": "ok"}], - usage={"prompt_tokens": 100, "completion_tokens": 50}, + usage=LLMUsage.reported(input_tokens=100, output_tokens=50), ) await hook.after_iteration(context) assert status.iteration == 5 assert len(status.tool_events) == 1 assert status.tool_events[0]["name"] == "read_file" - assert status.usage == {"prompt_tokens": 100, "completion_tokens": 50} + assert status.usage == LLMUsage.reported(input_tokens=100, output_tokens=50) @pytest.mark.asyncio async def test_after_iteration_with_error(self): @@ -821,7 +822,7 @@ class TestInspectTaskStatuses: phase="awaiting_tools", iteration=2, tool_events=[{"name": "read_file", "status": "ok", "detail": "ok"}], - usage={"prompt_tokens": 500, "completion_tokens": 100}, + usage=LLMUsage.reported(input_tokens=500, output_tokens=100), ), } tool = _make_tool(loop=loop) @@ -1127,12 +1128,12 @@ class TestLastUsageInSummary: tool = _make_tool() result = await tool.execute(action="check") assert "_last_usage" in result - assert "prompt_tokens" in result + assert "input_tokens" in result @pytest.mark.asyncio async def test_last_usage_not_shown_when_empty(self): loop = _make_mock_loop() - loop._last_usage = {} + loop._last_usage = None loop.last_usage = loop._last_usage tool = _make_tool(loop=loop) result = await tool.execute(action="check") diff --git a/tests/agent/tools/test_subagent_tools.py b/tests/agent/tools/test_subagent_tools.py index 84a941459..0a90b1ff6 100644 --- a/tests/agent/tools/test_subagent_tools.py +++ b/tests/agent/tools/test_subagent_tools.py @@ -456,7 +456,7 @@ async def test_agent_loop_syncs_updated_max_iterations_before_run(tmp_path): error=None, tool_events=[], messages=[], - usage={}, + usage=None, had_injections=False, tools_used=[], ) @@ -501,7 +501,7 @@ async def test_drain_pending_blocks_while_subagents_running(tmp_path): error=None, tool_events=[], messages=[], - usage={}, + usage=None, had_injections=False, tools_used=[], provider_state=None, @@ -587,7 +587,7 @@ async def test_drain_pending_no_block_when_no_subagents(tmp_path): error=None, tool_events=[], messages=[], - usage={}, + usage=None, had_injections=False, tools_used=[], provider_state=None, @@ -637,7 +637,7 @@ async def test_drain_pending_timeout(tmp_path): error=None, tool_events=[], messages=[], - usage={}, + usage=None, had_injections=False, tools_used=[], provider_state=None, diff --git a/tests/bus/test_runtime_events.py b/tests/bus/test_runtime_events.py index c2441d31b..44dd2440a 100644 --- a/tests/bus/test_runtime_events.py +++ b/tests/bus/test_runtime_events.py @@ -12,6 +12,7 @@ from nanobot.bus.runtime_events import ( TurnRunStatusChanged, TurnRuntimeAdmitted, ) +from nanobot.providers.base import LLMUsage @pytest.mark.asyncio @@ -99,7 +100,7 @@ async def test_runtime_event_publisher_consumes_turn_metadata_on_complete() -> N bus.subscribe(seen.append) publisher.record_turn_runtime("cli:direct", "runtime") publisher.record_turn_latency("cli:direct", 123) - publisher.record_turn_usage("cli:direct", {"prompt_tokens": 40, "completion_tokens": 2}) + publisher.record_turn_usage("cli:direct", LLMUsage.reported(input_tokens=40, output_tokens=2)) await publisher.turn_completed( channel="cli", @@ -120,11 +121,11 @@ async def test_runtime_event_publisher_consumes_turn_metadata_on_complete() -> N assert first.context.metadata == {"source": "test"} assert first.latency_ms == 123 assert first.runtime == "runtime" - assert first.usage == {"prompt_tokens": 40, "completion_tokens": 2} + assert first.usage == LLMUsage.reported(input_tokens=40, output_tokens=2) assert isinstance(second, TurnCompleted) assert second.latency_ms is None assert second.runtime is None - assert second.usage == {} + assert second.usage is None @pytest.mark.asyncio diff --git a/tests/cli/test_restart_command.py b/tests/cli/test_restart_command.py index 0409489fe..0152f1081 100644 --- a/tests/cli/test_restart_command.py +++ b/tests/cli/test_restart_command.py @@ -12,7 +12,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from nanobot.bus.events import InboundMessage -from nanobot.providers.base import LLMResponse +from nanobot.providers.base import LLMResponse, LLMUsage def _make_loop(): @@ -238,7 +238,7 @@ class TestRestartCommand: session.get_history.return_value = [{"role": "user"}] * 3 loop.sessions.get_or_create.return_value = session loop._start_time = time.time() - 125 - loop._last_usage = {"prompt_tokens": 0, "completion_tokens": 0} + loop._last_usage = LLMUsage.reported(input_tokens=0, output_tokens=0) loop.consolidator.estimate_session_prompt_tokens = MagicMock( return_value=(20500, "tiktoken") ) @@ -305,18 +305,15 @@ class TestRestartCommand: lambda _message: 7, ) loop.provider.chat_with_retry = AsyncMock(side_effect=[ - LLMResponse(content="first", usage={"prompt_tokens": 9, "completion_tokens": 4}), - LLMResponse(content="second", usage={}), + LLMResponse(content="first", usage=LLMUsage.reported(input_tokens=9, output_tokens=4)), + LLMResponse(content="second", usage=None), ]) await loop._run_agent_loop([], runtime=loop.llm_runtime()) - assert loop._last_usage["prompt_tokens"] == 9 - assert loop._last_usage["completion_tokens"] == 4 + assert loop._last_usage == LLMUsage.reported(input_tokens=9, output_tokens=4) await loop._run_agent_loop([], runtime=loop.llm_runtime()) - assert loop._last_usage["prompt_tokens"] == 123 - assert loop._last_usage["completion_tokens"] == 7 - assert loop._last_usage["estimated_tokens"] == 130 + assert loop._last_usage == LLMUsage.estimated(input_tokens=123, output_tokens=7) @pytest.mark.asyncio async def test_status_falls_back_to_last_usage_when_context_estimate_missing(self): @@ -324,7 +321,7 @@ class TestRestartCommand: session = MagicMock() session.get_history.return_value = [{"role": "user"}] loop.sessions.get_or_create.return_value = session - loop._last_usage = {"prompt_tokens": 1200, "completion_tokens": 34} + loop._last_usage = LLMUsage.reported(input_tokens=1200, output_tokens=34) loop.consolidator.estimate_session_prompt_tokens = MagicMock( return_value=(0, "none") ) diff --git a/tests/providers/test_azure_openai_provider.py b/tests/providers/test_azure_openai_provider.py index 10d2770c2..5978f0270 100644 --- a/tests/providers/test_azure_openai_provider.py +++ b/tests/providers/test_azure_openai_provider.py @@ -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 diff --git a/tests/providers/test_bedrock_provider.py b/tests/providers/test_bedrock_provider.py index a1c175245..668e6b0fa 100644 --- a/tests/providers/test_bedrock_provider.py +++ b/tests/providers/test_bedrock_provider.py @@ -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())) diff --git a/tests/providers/test_cached_tokens.py b/tests/providers/test_cached_tokens.py index 1b01408a4..acd63420a 100644 --- a/tests/providers/test_cached_tokens.py +++ b/tests/providers/test_cached_tokens.py @@ -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 diff --git a/tests/providers/test_custom_provider.py b/tests/providers/test_custom_provider.py index 023888976..551492257 100644 --- a/tests/providers/test_custom_provider.py +++ b/tests/providers/test_custom_provider.py @@ -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: diff --git a/tests/providers/test_openai_codex_provider.py b/tests/providers/test_openai_codex_provider.py index 1bbad4ba1..57f37a094 100644 --- a/tests/providers/test_openai_codex_provider.py +++ b/tests/providers/test_openai_codex_provider.py @@ -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" diff --git a/tests/providers/test_openai_responses.py b/tests/providers/test_openai_responses.py index b25e51592..abb4005f8 100644 --- a/tests/providers/test_openai_responses.py +++ b/tests/providers/test_openai_responses.py @@ -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 diff --git a/tests/providers/test_provider_retry.py b/tests/providers/test_provider_retry.py index e3e32c28f..5f603de91 100644 --- a/tests/providers/test_provider_retry.py +++ b/tests/providers/test_provider_retry.py @@ -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 = {} diff --git a/tests/providers/test_providers_init.py b/tests/providers/test_providers_init.py index f3c3b51ca..71d16bd06 100644 --- a/tests/providers/test_providers_init.py +++ b/tests/providers/test_providers_init.py @@ -32,6 +32,7 @@ def test_importing_providers_package_is_lazy(monkeypatch) -> None: assert providers.__all__ == [ "LLMProvider", "LLMResponse", + "LLMUsage", "AnthropicProvider", "OpenAICompatProvider", "OpenAICodexProvider", diff --git a/tests/providers/test_usage_contract.py b/tests/providers/test_usage_contract.py new file mode 100644 index 000000000..9b9cbb733 --- /dev/null +++ b/tests/providers/test_usage_contract.py @@ -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 diff --git a/tests/providers/test_xai_grok_provider.py b/tests/providers/test_xai_grok_provider.py index 8ba95da26..5d9629a2f 100644 --- a/tests/providers/test_xai_grok_provider.py +++ b/tests/providers/test_xai_grok_provider.py @@ -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"}] diff --git a/tests/test_api_attachment.py b/tests/test_api_attachment.py index cb23f669f..1fadef636 100644 --- a/tests/test_api_attachment.py +++ b/tests/test_api_attachment.py @@ -33,7 +33,7 @@ def _make_mock_agent(response_text: str = "mock response") -> MagicMock: agent = MagicMock() agent.process_direct = AsyncMock(return_value=response_text) agent.aclose = AsyncMock() - agent._last_usage = {} + agent._last_usage = None return agent diff --git a/tests/test_api_stream.py b/tests/test_api_stream.py index 35847f3ce..0b6f54e21 100644 --- a/tests/test_api_stream.py +++ b/tests/test_api_stream.py @@ -77,7 +77,7 @@ def _make_streaming_agent(tokens: list[str]) -> MagicMock: return " ".join(tokens) agent.process_direct = fake_process_direct - agent._last_usage = {} + agent._last_usage = None return agent @@ -136,7 +136,7 @@ async def test_stream_false_returns_json(aiohttp_client) -> None: agent = MagicMock() agent.process_direct = AsyncMock(return_value="normal reply") agent.aclose = AsyncMock() - agent._last_usage = {} + agent._last_usage = None app = create_app(agent, model_name="m", api_key=API_KEY) client = await aiohttp_client(app) @@ -159,7 +159,7 @@ async def test_stream_default_is_false(aiohttp_client) -> None: agent = MagicMock() agent.process_direct = AsyncMock(return_value="default reply") agent.aclose = AsyncMock() - agent._last_usage = {} + agent._last_usage = None app = create_app(agent, model_name="m", api_key=API_KEY) client = await aiohttp_client(app) @@ -215,7 +215,7 @@ async def test_stream_passes_on_stream_callbacks(aiohttp_client) -> None: agent = MagicMock() agent.process_direct = fake_process_direct agent.aclose = AsyncMock() - agent._last_usage = {} + agent._last_usage = None app = create_app(agent, model_name="m", api_key=API_KEY) client = await aiohttp_client(app) @@ -248,7 +248,7 @@ async def test_stream_segment_end_does_not_close_sse(aiohttp_client) -> None: agent.process_direct = fake_process_direct agent.aclose = AsyncMock() - agent._last_usage = {} + agent._last_usage = None app = create_app(agent, model_name="m", api_key=API_KEY) client = await aiohttp_client(app) @@ -287,7 +287,7 @@ async def test_stream_uses_final_response_when_no_deltas(aiohttp_client) -> None agent.process_direct = fake_process_direct agent.aclose = AsyncMock() - agent._last_usage = {} + agent._last_usage = None app = create_app(agent, model_name="m", api_key=API_KEY) client = await aiohttp_client(app) @@ -329,7 +329,7 @@ async def test_stream_with_session_id(aiohttp_client) -> None: agent = MagicMock() agent.process_direct = fake_process_direct agent.aclose = AsyncMock() - agent._last_usage = {} + agent._last_usage = None app = create_app(agent, model_name="m", api_key=API_KEY) client = await aiohttp_client(app) @@ -358,7 +358,7 @@ async def test_streaming_backend_failure_does_not_emit_success_terminator(aiohtt agent.process_direct = boom agent.aclose = AsyncMock() - agent._last_usage = {} + agent._last_usage = None app = create_app(agent, model_name="m", api_key=API_KEY) client = await aiohttp_client(app) diff --git a/tests/test_build_status.py b/tests/test_build_status.py index 922243d56..00cfc6ff1 100644 --- a/tests/test_build_status.py +++ b/tests/test_build_status.py @@ -1,5 +1,6 @@ """Tests for build_status_content cache hit rate display.""" +from nanobot.providers.base import LLMUsage from nanobot.utils.helpers import build_status_content @@ -8,7 +9,7 @@ def test_status_shows_cache_hit_rate(): version="0.1.0", model="glm-4-plus", start_time=1000000.0, - last_usage={"prompt_tokens": 2000, "completion_tokens": 300, "cached_tokens": 1200}, + last_usage=LLMUsage.reported(input_tokens=2000, output_tokens=300, cache_read_tokens=1200), context_window_tokens=128000, session_msg_count=10, context_tokens_estimate=5000, @@ -19,12 +20,12 @@ def test_status_shows_cache_hit_rate(): def test_status_no_cache_info(): - """Without cached_tokens, display should not show cache percentage.""" + """Without a reported cache-read count, omit the cache percentage.""" content = build_status_content( version="0.1.0", model="glm-4-plus", start_time=1000000.0, - last_usage={"prompt_tokens": 2000, "completion_tokens": 300}, + last_usage=LLMUsage.reported(input_tokens=2000, output_tokens=300), context_window_tokens=128000, session_msg_count=10, context_tokens_estimate=5000, @@ -34,13 +35,13 @@ def test_status_no_cache_info(): assert "Tasks: 0 active" in content -def test_status_zero_cached_tokens(): - """cached_tokens=0 should not show cache percentage.""" +def test_status_zero_cache_read_tokens(): + """An explicit zero cache-read count should not show cache percentage.""" content = build_status_content( version="0.1.0", model="glm-4-plus", start_time=1000000.0, - last_usage={"prompt_tokens": 2000, "completion_tokens": 300, "cached_tokens": 0}, + last_usage=LLMUsage.reported(input_tokens=2000, output_tokens=300, cache_read_tokens=0), context_window_tokens=128000, session_msg_count=10, context_tokens_estimate=5000, @@ -53,7 +54,7 @@ def test_status_100_percent_cached(): version="0.1.0", model="glm-4-plus", start_time=1000000.0, - last_usage={"prompt_tokens": 1000, "completion_tokens": 100, "cached_tokens": 1000}, + last_usage=LLMUsage.reported(input_tokens=1000, output_tokens=100, cache_read_tokens=1000), context_window_tokens=128000, session_msg_count=5, context_tokens_estimate=3000, @@ -67,7 +68,7 @@ def test_status_context_pct_uses_budget_not_total(): version="0.1.0", model="test", start_time=1000000.0, - last_usage={"prompt_tokens": 2000, "completion_tokens": 300}, + last_usage=LLMUsage.reported(input_tokens=2000, output_tokens=300), context_window_tokens=128000, session_msg_count=10, context_tokens_estimate=120000, @@ -83,7 +84,7 @@ def test_status_context_pct_capped_at_999(): version="0.1.0", model="test", start_time=1000000.0, - last_usage={"prompt_tokens": 2000, "completion_tokens": 300}, + last_usage=LLMUsage.reported(input_tokens=2000, output_tokens=300), context_window_tokens=10000, session_msg_count=10, context_tokens_estimate=100000, diff --git a/tests/test_nanobot_facade.py b/tests/test_nanobot_facade.py index d168849c0..be6dc4a46 100644 --- a/tests/test_nanobot_facade.py +++ b/tests/test_nanobot_facade.py @@ -30,6 +30,7 @@ from nanobot.nanobot import ( StreamEvent, StreamEventType, ) +from nanobot.providers.base import LLMUsage from nanobot.runtime_context import ( RUNTIME_CONTEXT_HISTORY_META, RuntimeContextBlock, @@ -601,7 +602,7 @@ async def test_run_no_iterations_leaves_defaults_empty(tmp_path): result = await bot.run("hi") assert result.tools_used == [] assert result.messages == [] - assert result.usage == {} + assert result.usage is None assert result.stop_reason is None assert result.error is None @@ -622,7 +623,7 @@ async def test_run_populates_observability_fields(tmp_path): ], final_content="done", tools_used=["read_file"], - usage={"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12}, + usage=LLMUsage.reported(input_tokens=10, output_tokens=2), stop_reason="completed", error=None, tool_events=[{"tool": "read_file", "status": "ok"}], @@ -641,7 +642,7 @@ async def test_run_populates_observability_fields(tmp_path): assert result.content == "done" assert result.tools_used == ["read_file"] - assert result.usage == {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12} + assert result.usage == LLMUsage.reported(input_tokens=10, output_tokens=2) assert result.stop_reason == "completed" assert result.error is None assert result.metadata == {"latency_ms": 42} @@ -658,7 +659,7 @@ async def test_run_ephemeral_still_captures_runner_observability(tmp_path): provider.chat_with_retry = AsyncMock(return_value=LLMResponse( content="done", tool_calls=[], - usage={"total_tokens": 3}, + usage=LLMUsage.reported(input_tokens=3, output_tokens=0), )) bot = Nanobot(AgentLoop( bus=MessageBus(), @@ -670,8 +671,7 @@ async def test_run_ephemeral_still_captures_runner_observability(tmp_path): result = await bot.run("hi", ephemeral=True) assert result.content == "done" - assert result.usage["total_tokens"] == 3 - assert result.usage["provider_tokens"] == 3 + assert result.usage == LLMUsage.reported(input_tokens=3, output_tokens=0) @pytest.mark.asyncio @@ -1053,7 +1053,7 @@ async def test_run_streamed_wait_returns_full_result_without_consuming_events(tm ], final_content="done", tools_used=["read_file"], - usage={"total_tokens": 9}, + usage=LLMUsage.reported(input_tokens=9, output_tokens=0), stop_reason="completed", ) for hook in hooks: @@ -1073,7 +1073,7 @@ async def test_run_streamed_wait_returns_full_result_without_consuming_events(tm assert result.content == "done" assert result.tools_used == ["read_file"] - assert result.usage == {"total_tokens": 9} + assert result.usage == LLMUsage.reported(input_tokens=9, output_tokens=0) assert result.stop_reason == "completed" assert result.metadata == {"latency_ms": 5} @@ -1397,13 +1397,13 @@ async def test_sdk_capture_prefers_run_level_snapshot(): await hook.after_run(AgentRunHookContext( messages=final_messages, tools_used=["read_file"], - usage={"total_tokens": 3}, + usage=LLMUsage.reported(input_tokens=3, output_tokens=0), stop_reason="completed", )) assert hook.tools_used == ["read_file"] assert hook.messages == final_messages - assert hook.usage == {"total_tokens": 3} + assert hook.usage == LLMUsage.reported(input_tokens=3, output_tokens=0) assert hook.stop_reason == "completed" diff --git a/tests/test_openai_api.py b/tests/test_openai_api.py index 013e84907..00ff24ff0 100644 --- a/tests/test_openai_api.py +++ b/tests/test_openai_api.py @@ -17,6 +17,7 @@ from nanobot.api.server import ( create_app, handle_chat_completions, ) +from nanobot.providers.base import LLMUsage try: from aiohttp.test_utils import TestClient, TestServer @@ -35,7 +36,7 @@ def _make_mock_agent(response_text: str = "mock response") -> MagicMock: agent = MagicMock() agent.process_direct = AsyncMock(return_value=response_text) agent.aclose = AsyncMock() - agent._last_usage = {"prompt_tokens": 100, "completion_tokens": 50} + agent._last_usage = LLMUsage.reported(input_tokens=100, output_tokens=50) return agent @@ -87,19 +88,19 @@ def test_chat_completion_response() -> None: def test_chat_completion_response_with_usage() -> None: - usage = {"prompt_tokens": 150, "completion_tokens": 42} + usage = LLMUsage.reported(input_tokens=150, output_tokens=42) result = _chat_completion_response("hello world", "test-model", usage) assert result["usage"]["prompt_tokens"] == 150 assert result["usage"]["completion_tokens"] == 42 assert result["usage"]["total_tokens"] == 192 -def test_chat_completion_response_preserves_provider_total_usage() -> None: - usage = {"total_tokens": 77} +def test_chat_completion_response_preserves_explicit_total_usage() -> None: + usage = LLMUsage.reported(input_tokens=70, output_tokens=7, total_tokens=175) result = _chat_completion_response("hello world", "test-model", usage) - assert result["usage"]["prompt_tokens"] == 0 - assert result["usage"]["completion_tokens"] == 0 - assert result["usage"]["total_tokens"] == 77 + assert result["usage"]["prompt_tokens"] == 70 + assert result["usage"]["completion_tokens"] == 7 + assert result["usage"]["total_tokens"] == 175 @pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed") @@ -328,7 +329,7 @@ async def test_followup_requests_share_same_session_key(aiohttp_client) -> None: agent = MagicMock() agent.process_direct = fake_process agent.aclose = AsyncMock() - agent._last_usage = {} + agent._last_usage = None app = create_app(agent, model_name="m", api_key=API_KEY) client = await aiohttp_client(app) @@ -367,7 +368,7 @@ async def test_fixed_session_requests_are_serialized(aiohttp_client) -> None: agent = MagicMock() agent.process_direct = slow_process agent.aclose = AsyncMock() - agent._last_usage = {} + agent._last_usage = None app = create_app(agent, model_name="m", api_key=API_KEY) client = await aiohttp_client(app) @@ -484,7 +485,7 @@ async def test_empty_response_falls_back_without_retry(aiohttp_client) -> None: agent = MagicMock() agent.process_direct = always_empty agent.aclose = AsyncMock() - agent._last_usage = {} + agent._last_usage = None app = create_app(agent, model_name="m", api_key=API_KEY) client = await aiohttp_client(app) diff --git a/tests/utils/test_searchusage.py b/tests/utils/test_searchusage.py index 0bb71f5f8..ea4bb1ea0 100644 --- a/tests/utils/test_searchusage.py +++ b/tests/utils/test_searchusage.py @@ -6,6 +6,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from nanobot.providers.base import LLMUsage from nanobot.utils.helpers import build_status_content from nanobot.utils.searchusage import ( SearchUsageInfo, @@ -273,7 +274,7 @@ class TestBuildStatusContentWithSearchUsage: version="0.1.0", model="claude-opus-4-5", start_time=1_000_000.0, - last_usage={"prompt_tokens": 1000, "completion_tokens": 200}, + last_usage=LLMUsage.reported(input_tokens=1000, output_tokens=200), context_window_tokens=65536, session_msg_count=5, context_tokens_estimate=3000, diff --git a/tests/webui/test_session_context.py b/tests/webui/test_session_context.py index 9cc90c3df..b125b1b5f 100644 --- a/tests/webui/test_session_context.py +++ b/tests/webui/test_session_context.py @@ -1,3 +1,4 @@ +from nanobot.providers.base import LLMUsage from nanobot.session import Session from nanobot.utils.helpers import estimate_message_tokens from nanobot.webui.session_context import session_context_payload @@ -58,19 +59,12 @@ def test_session_context_tolerates_untrusted_summary_metadata() -> None: def test_session_context_sanitizes_usage_metadata() -> None: + usage = LLMUsage.reported(input_tokens=120, output_tokens=8, total_tokens=175) session = Session( key="websocket:context", - metadata={ - "_last_usage": { - "prompt_tokens": 120, - "completion_tokens": 8, - "negative": -1, - "boolean": True, - "text": "invalid", - } - }, + metadata={"_last_usage": usage.to_dict()}, ) payload = session_context_payload(session) - assert payload["last_usage"] == {"prompt_tokens": 120, "completion_tokens": 8} + assert payload["last_usage"] == usage.to_dict() diff --git a/tests/webui/test_settings_api.py b/tests/webui/test_settings_api.py index a4b471863..caab2b864 100644 --- a/tests/webui/test_settings_api.py +++ b/tests/webui/test_settings_api.py @@ -9,6 +9,7 @@ import pytest from nanobot.config.loader import load_config, save_config from nanobot.config.schema import Config, InlineFallbackConfig, ModelPresetConfig +from nanobot.providers.base import LLMUsage from nanobot.providers.registry import find_by_name from nanobot.session.manager import SessionManager from nanobot.session.model_selection import SESSION_MODEL_PRESET_METADATA_KEY @@ -1467,7 +1468,7 @@ def test_settings_payload_includes_token_usage_summary( from nanobot.webui.token_usage import record_token_usage record_token_usage( - {"prompt_tokens": 10, "completion_tokens": 5}, + LLMUsage.reported(input_tokens=10, output_tokens=5), timezone_name=config.agents.defaults.timezone, ) @@ -1495,7 +1496,7 @@ def test_settings_usage_payload_returns_lightweight_token_usage( from nanobot.webui.token_usage import record_token_usage record_token_usage( - {"prompt_tokens": 20, "completion_tokens": 2}, + LLMUsage.reported(input_tokens=20, output_tokens=2), timezone_name=config.agents.defaults.timezone, ) diff --git a/tests/webui/test_token_usage.py b/tests/webui/test_token_usage.py index 6514909fa..07f2cc202 100644 --- a/tests/webui/test_token_usage.py +++ b/tests/webui/test_token_usage.py @@ -1,17 +1,20 @@ from __future__ import annotations import json -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from types import SimpleNamespace import pytest from nanobot.agent.hook import AgentHookContext +from nanobot.providers.base import LLMUsage from nanobot.webui.token_usage import ( TokenUsageHook, + read_token_usage_state, record_response_token_usage, record_token_usage, token_usage_payload, + write_token_usage_state, ) @@ -19,7 +22,7 @@ def _write_state(tmp_path, days: dict) -> None: state_dir = tmp_path / "webui" state_dir.mkdir(parents=True, exist_ok=True) (state_dir / "token-usage.json").write_text( - json.dumps({"days": days}), encoding="utf-8" + json.dumps({"schema_version": 2, "days": days}), encoding="utf-8" ) @@ -58,7 +61,7 @@ def test_record_scrubs_malformed_day_keys(tmp_path, monkeypatch) -> None: }) record_token_usage( - {"prompt_tokens": 1, "completion_tokens": 1}, + LLMUsage.reported(input_tokens=1, output_tokens=1), timezone_name="UTC", now=datetime(2026, 6, 3, 12, 0, tzinfo=timezone.utc), ) @@ -73,12 +76,16 @@ def test_record_token_usage_aggregates_by_local_day(tmp_path, monkeypatch) -> No monkeypatch.setattr("nanobot.webui.token_usage.get_webui_dir", lambda: tmp_path / "webui") record_token_usage( - {"prompt_tokens": 100, "completion_tokens": 40, "cached_tokens": 20}, + LLMUsage.reported( + input_tokens=100, + output_tokens=40, + cache_read_tokens=20, + ), timezone_name="Asia/Shanghai", now=datetime(2026, 6, 2, 18, 0, tzinfo=timezone.utc), ) record_token_usage( - {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + LLMUsage.reported(input_tokens=10, output_tokens=5), timezone_name="Asia/Shanghai", now=datetime(2026, 6, 2, 19, 0, tzinfo=timezone.utc), ) @@ -94,25 +101,31 @@ def test_record_token_usage_aggregates_by_local_day(tmp_path, monkeypatch) -> No assert payload["days"] == [ { "date": "2026-06-03", - "prompt_tokens": 110, - "completion_tokens": 45, - "cached_tokens": 20, + "input_tokens": 110, + "output_tokens": 45, + "cache_read_tokens": 20, + "cache_write_tokens": 0, + "cache_read_observed_input_tokens": 100, + "cache_write_observed_input_tokens": 0, "total_tokens": 155, - "provider_tokens": 155, + "reported_tokens": 155, "estimated_tokens": 0, "requests": 2, - "provider_requests": 2, + "reported_requests": 2, "estimated_requests": 0, "sources": { "user": { - "prompt_tokens": 110, - "completion_tokens": 45, - "cached_tokens": 20, + "input_tokens": 110, + "output_tokens": 45, + "cache_read_tokens": 20, + "cache_write_tokens": 0, + "cache_read_observed_input_tokens": 100, + "cache_write_observed_input_tokens": 0, "total_tokens": 155, - "provider_tokens": 155, + "reported_tokens": 155, "estimated_tokens": 0, "requests": 2, - "provider_requests": 2, + "reported_requests": 2, "estimated_requests": 0, } }, @@ -120,10 +133,113 @@ def test_record_token_usage_aggregates_by_local_day(tmp_path, monkeypatch) -> No ] +def test_cache_observation_denominators_distinguish_missing_from_zero( + tmp_path, + monkeypatch, +) -> None: + monkeypatch.setattr("nanobot.webui.token_usage.get_webui_dir", lambda: tmp_path / "webui") + now = datetime(2026, 6, 3, tzinfo=timezone.utc) + + record_token_usage( + LLMUsage.reported(input_tokens=100, output_tokens=10), + source="user", + now=now, + ) + record_token_usage( + LLMUsage.reported( + input_tokens=40, + output_tokens=5, + cache_read_tokens=0, + cache_write_tokens=0, + ), + source="dream", + now=now, + ) + + row = token_usage_payload(now=now)["days"][0] + + assert row["cache_read_tokens"] == 0 + assert row["cache_write_tokens"] == 0 + assert row["cache_read_observed_input_tokens"] == 40 + assert row["cache_write_observed_input_tokens"] == 40 + assert row["sources"]["user"]["cache_read_observed_input_tokens"] == 0 + assert row["sources"]["user"]["cache_write_observed_input_tokens"] == 0 + assert row["sources"]["dream"]["cache_read_observed_input_tokens"] == 40 + assert row["sources"]["dream"]["cache_write_observed_input_tokens"] == 40 + + +def _retention_state(sources: tuple[str, ...], *, day_count: int = 400) -> dict: + start = datetime(2025, 1, 1, tzinfo=timezone.utc) + source_usage = { + "input_tokens": 100, + "output_tokens": 10, + "total_tokens": 110, + "reported_tokens": 110, + "requests": 1, + "reported_requests": 1, + } + days = {} + for offset in range(day_count): + day = (start + timedelta(days=offset)).date().isoformat() + days[day] = { + "input_tokens": 100 * len(sources), + "output_tokens": 10 * len(sources), + "total_tokens": 110 * len(sources), + "reported_tokens": 110 * len(sources), + "requests": len(sources), + "reported_requests": len(sources), + "sources": {source: dict(source_usage) for source in sources}, + } + return {"schema_version": 2, "days": days} + + +def test_write_compact_state_keeps_400_days_with_two_sources(tmp_path, monkeypatch) -> None: + monkeypatch.setattr("nanobot.webui.token_usage.get_webui_dir", lambda: tmp_path / "webui") + + written = write_token_usage_state(_retention_state(("user", "api"))) + persisted = (tmp_path / "webui" / "token-usage.json").read_bytes() + + assert len(written["days"]) == 400 + assert len(persisted) <= 512 * 1024 + assert persisted.endswith(b"\n") + assert json.loads(persisted) == written + + +def test_write_prunes_only_oldest_days_to_fit_byte_budget(tmp_path, monkeypatch) -> None: + monkeypatch.setattr("nanobot.webui.token_usage.get_webui_dir", lambda: tmp_path / "webui") + sources = ("user", "api", "cron", "dream", "system") + raw = _retention_state(sources) + all_dates = list(raw["days"]) + + written = write_token_usage_state(raw) + retained_dates = list(written["days"]) + persisted = (tmp_path / "webui" / "token-usage.json").read_bytes() + + assert 1 <= len(retained_dates) < len(all_dates) + assert retained_dates == all_dates[-len(retained_dates) :] + assert retained_dates[-1] == all_dates[-1] + assert all(set(row["sources"]) == set(sources) for row in written["days"].values()) + assert len(persisted) <= 512 * 1024 + assert read_token_usage_state() == written + + +def test_write_raises_when_latest_day_alone_exceeds_byte_budget( + tmp_path, + monkeypatch, +) -> None: + monkeypatch.setattr("nanobot.webui.token_usage.get_webui_dir", lambda: tmp_path / "webui") + monkeypatch.setattr("nanobot.webui.token_usage._MAX_STATE_FILE_BYTES", 256) + + with pytest.raises(ValueError, match="latest token usage day exceeds"): + write_token_usage_state(_retention_state(("user", "api"), day_count=1)) + + assert not (tmp_path / "webui" / "token-usage.json").exists() + + def test_record_token_usage_skips_empty_usage(tmp_path, monkeypatch) -> None: monkeypatch.setattr("nanobot.webui.token_usage.get_webui_dir", lambda: tmp_path / "webui") - record_token_usage({"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}) + record_token_usage(LLMUsage.reported(input_tokens=0, output_tokens=0)) payload = token_usage_payload(now=datetime(2026, 6, 3, tzinfo=timezone.utc)) assert payload["days"] == [] @@ -134,14 +250,14 @@ def test_record_token_usage_keeps_estimated_split(tmp_path, monkeypatch) -> None monkeypatch.setattr("nanobot.webui.token_usage.get_webui_dir", lambda: tmp_path / "webui") record_token_usage( - {"prompt_tokens": 100, "completion_tokens": 25, "estimated_tokens": 125}, + LLMUsage.estimated(input_tokens=100, output_tokens=25), now=datetime(2026, 6, 3, tzinfo=timezone.utc), ) payload = token_usage_payload(now=datetime(2026, 6, 3, tzinfo=timezone.utc)) assert payload["days"][0]["total_tokens"] == 125 - assert payload["days"][0]["provider_tokens"] == 0 + assert payload["days"][0]["reported_tokens"] == 0 assert payload["days"][0]["estimated_tokens"] == 125 assert payload["days"][0]["estimated_requests"] == 1 @@ -150,12 +266,12 @@ def test_record_token_usage_keeps_source_breakdown(tmp_path, monkeypatch) -> Non monkeypatch.setattr("nanobot.webui.token_usage.get_webui_dir", lambda: tmp_path / "webui") record_token_usage( - {"prompt_tokens": 100, "completion_tokens": 25}, + LLMUsage.reported(input_tokens=100, output_tokens=25, total_tokens=175), source="user", now=datetime(2026, 6, 3, tzinfo=timezone.utc), ) record_token_usage( - {"prompt_tokens": 20, "completion_tokens": 5}, + LLMUsage.reported(input_tokens=20, output_tokens=5), source="dream", now=datetime(2026, 6, 3, tzinfo=timezone.utc), ) @@ -163,8 +279,8 @@ def test_record_token_usage_keeps_source_breakdown(tmp_path, monkeypatch) -> Non payload = token_usage_payload(now=datetime(2026, 6, 3, tzinfo=timezone.utc)) row = payload["days"][0] - assert row["total_tokens"] == 150 - assert row["sources"]["user"]["total_tokens"] == 125 + assert row["total_tokens"] == 200 + assert row["sources"]["user"]["total_tokens"] == 175 assert row["sources"]["user"]["requests"] == 1 assert row["sources"]["dream"]["total_tokens"] == 25 assert row["sources"]["dream"]["requests"] == 1 @@ -175,7 +291,7 @@ def test_record_response_token_usage_uses_response_usage(tmp_path, monkeypatch) monkeypatch.setattr("nanobot.webui.token_usage._local_day", lambda *_, **__: "2026-06-03") record_response_token_usage( - SimpleNamespace(usage={"prompt_tokens": 20, "completion_tokens": 5}), + SimpleNamespace(usage=LLMUsage.reported(input_tokens=20, output_tokens=5)), source="dream", ) @@ -194,7 +310,7 @@ async def test_token_usage_hook_classifies_source_from_session_key(tmp_path, mon iteration=0, messages=[], session_key="cron:drink-water", - usage={"prompt_tokens": 10, "completion_tokens": 5}, + usage=LLMUsage.reported(input_tokens=10, output_tokens=5), ) ) diff --git a/webui/src/lib/types.ts b/webui/src/lib/types.ts index d404d511e..9e110f6e8 100644 --- a/webui/src/lib/types.ts +++ b/webui/src/lib/types.ts @@ -712,26 +712,32 @@ export interface SettingsPayload { usage?: { days: Array<{ date: string; - prompt_tokens: number; - completion_tokens: number; - cached_tokens: number; + input_tokens: number; + output_tokens: number; + cache_read_tokens: number; + cache_write_tokens: number; + cache_read_observed_input_tokens: number; + cache_write_observed_input_tokens: number; total_tokens: number; - provider_tokens?: number; + reported_tokens?: number; estimated_tokens?: number; requests: number; - provider_requests?: number; + reported_requests?: number; estimated_requests?: number; sources?: Record< "user" | "api" | "cron" | "dream" | "system" | string, { - prompt_tokens: number; - completion_tokens: number; - cached_tokens: number; + input_tokens: number; + output_tokens: number; + cache_read_tokens: number; + cache_write_tokens: number; + cache_read_observed_input_tokens: number; + cache_write_observed_input_tokens: number; total_tokens: number; - provider_tokens?: number; + reported_tokens?: number; estimated_tokens?: number; requests: number; - provider_requests?: number; + reported_requests?: number; estimated_requests?: number; } >; diff --git a/webui/src/tests/settings-overview.test.tsx b/webui/src/tests/settings-overview.test.tsx index a9dfdc970..49871c504 100644 --- a/webui/src/tests/settings-overview.test.tsx +++ b/webui/src/tests/settings-overview.test.tsx @@ -111,9 +111,12 @@ describe("Settings overview and appearance", () => { days: [ { date: "2026-06-03", - prompt_tokens: 1200, - completion_tokens: 300, - cached_tokens: 500, + input_tokens: 1200, + output_tokens: 300, + cache_read_tokens: 500, + cache_write_tokens: 0, + cache_read_observed_input_tokens: 1200, + cache_write_observed_input_tokens: 1200, total_tokens: 1500, requests: 2, }, @@ -214,9 +217,12 @@ describe("Settings overview and appearance", () => { days: [ { date: "2026-06-03", - prompt_tokens: 1200, - completion_tokens: 300, - cached_tokens: 500, + input_tokens: 1200, + output_tokens: 300, + cache_read_tokens: 500, + cache_write_tokens: 0, + cache_read_observed_input_tokens: 1200, + cache_write_observed_input_tokens: 1200, total_tokens: 1500, requests: 2, },