From e111b83af66216ccae5cf7d86350f818fb9059dd Mon Sep 17 00:00:00 2001 From: chengyongru <61816729+chengyongru@users.noreply.github.com> Date: Mon, 31 Aug 2026 18:08:21 +0800 Subject: [PATCH] refactor(agent): unify runner request fitting (#5612) * refactor(agent): unify runner request fitting * fix(agent): fit every runner model request * fix(agent): count resumed state during request fitting * refactor(agent): consolidate request fitting state --- nanobot/agent/context_governance.py | 235 +++--- nanobot/agent/runner.py | 194 +++-- nanobot/providers/conversation_state.py | 43 +- tests/agent/test_dream.py | 6 +- tests/agent/test_loop_consolidation_tokens.py | 3 + tests/agent/test_runner_core.py | 59 +- tests/agent/test_runner_governance.py | 790 ++++++++++++------ tests/agent/test_session_model_runtime.py | 4 +- tests/providers/test_conversation_state.py | 45 + 9 files changed, 888 insertions(+), 491 deletions(-) diff --git a/nanobot/agent/context_governance.py b/nanobot/agent/context_governance.py index 98b1291dc..086f4e723 100644 --- a/nanobot/agent/context_governance.py +++ b/nanobot/agent/context_governance.py @@ -13,6 +13,7 @@ from typing import TYPE_CHECKING, Any, cast from loguru import logger +from nanobot.providers.base import LLMUsage from nanobot.utils.helpers import ( estimate_message_tokens, estimate_prompt_tokens_chain, @@ -27,12 +28,6 @@ if TYPE_CHECKING: from nanobot.providers.base import LLMProvider SNIP_SAFETY_BUFFER = 1024 -MICROCOMPACT_MIN_CHARS = 500 -INFLIGHT_COMPACT_TARGET_RATIO = 0.85 -COMPACTABLE_TOOLS = frozenset({ - "read_file", "exec", "grep", "find_files", - "web_search", "web_fetch", "list_dir", "list_exec_sessions", -}) # read_file is the recovery path for persisted results; exempting it prevents persist->read->persist loops. TOOL_RESULT_OFFLOAD_EXEMPT_TOOLS = frozenset({"read_file"}) BACKFILL_CONTENT = "[Tool result unavailable — call was interrupted or lost]" @@ -41,6 +36,27 @@ PLACEHOLDER_TEXTS = frozenset({ }) +class ContextWindowExceededError(RuntimeError): + """Raised before a locally fitted request that still exceeds its budget.""" + + def __init__( + self, + *, + session_key: str | None, + estimated_tokens: int, + input_budget: int, + source: str, + ) -> None: + self.session_key = session_key + self.estimated_tokens = estimated_tokens + self.input_budget = input_budget + self.source = source + super().__init__( + "Model input still exceeds the local context budget after request fitting " + f"for {session_key or 'default'}: {estimated_tokens}/{input_budget} via {source}" + ) + + def _tool_call_name_is_valid(tool_call: Any) -> bool: """Whether a persisted OpenAI-style tool_call carries a usable name. @@ -67,7 +83,6 @@ class ContextGovernanceConfig: context_window_tokens: int | None = None context_block_limit: int | None = None max_tokens: int | None = None - inflight_start_index: int = 0 class ContextGovernor: @@ -77,17 +92,85 @@ class ContextGovernor: self, config: ContextGovernanceConfig, messages: list[dict[str, Any]], - compacted_tool_call_ids: set[str], ) -> list[dict[str, Any]]: updated = self.strip_placeholder_assistant_messages(messages) updated = self.strip_malformed_tool_calls(updated) updated = self.drop_orphan_tool_results(updated) updated = self.backfill_missing_tool_results(updated) - updated = self.apply_tool_result_budget(config, updated) - updated = self.compact_inflight_overflow(config, updated, compacted_tool_call_ids) - updated = self.snip_history(config, updated) + return self.apply_tool_result_budget(config, updated) + + def fit_to_budget( + self, + config: ContextGovernanceConfig, + messages: list[dict[str, Any]], + *, + tool_definitions: list[dict[str, Any]] | None, + ) -> list[dict[str, Any]]: + """Fit a model-facing copy while keeping the source transcript intact.""" + updated = self.snip_history( + config, + messages, + tool_definitions=tool_definitions, + force=True, + ) updated = self.drop_orphan_tool_results(updated) - return self.backfill_missing_tool_results(updated) + updated = self.backfill_missing_tool_results(updated) + if not config.context_window_tokens: + return updated + budget = self.input_budget(config) + estimated, source = estimate_prompt_tokens_chain( + config.provider, + config.model, + updated, + tool_definitions, + ) + if budget > 0 and estimated <= budget: + return updated + raise ContextWindowExceededError( + session_key=config.session_key, + estimated_tokens=estimated, + input_budget=budget, + source=source, + ) + + def fit_request( + self, + config: ContextGovernanceConfig, + messages: list[dict[str, Any]], + usage: LLMUsage | None, + *, + usage_matches_messages: bool, + tool_definitions: list[dict[str, Any]] | None, + request_context_tokens: int | None = None, + ) -> tuple[list[dict[str, Any]], bool]: + """Fit the request when its measured or estimated input is pressured.""" + if not config.context_window_tokens: + return messages, False + budget = self.input_budget(config) + if ( + request_context_tokens is None + and usage_matches_messages + and usage is not None + and usage.context_tokens is not None + ): + pressured = budget <= 0 or usage.context_tokens >= budget + else: + estimated, _ = estimate_prompt_tokens_chain( + config.provider, + config.model, + messages, + tool_definitions, + ) + if request_context_tokens is not None: + estimated = max(estimated, request_context_tokens) + pressured = budget <= 0 or estimated >= budget + if not pressured: + return messages, False + return self.fit_to_budget( + config, + messages, + tool_definitions=tool_definitions, + ), True @staticmethod def input_budget(config: ContextGovernanceConfig) -> int: @@ -326,71 +409,13 @@ class ContextGovernor: updated[idx]["content"] = normalized return updated - def compact_inflight_overflow( - self, - config: ContextGovernanceConfig, - messages: list[dict[str, Any]], - compacted_tool_call_ids: set[str], - ) -> list[dict[str, Any]]: - """Compact in-flight tool results only when the request would overflow.""" - budget = self.input_budget(config) - if budget <= 0: - return messages - - tools = config.tools.get_definitions() - updated = self._apply_recorded_compactions(messages, compacted_tool_call_ids) - estimate, source = estimate_prompt_tokens_chain( - config.provider, - config.model, - updated, - tools, - ) - if estimate <= budget: - return updated - - target = int(budget * INFLIGHT_COMPACT_TARGET_RATIO) - candidates = self._inflight_compaction_candidates( - config, - updated, - compacted_tool_call_ids, - ) - if not candidates: - return updated - - for candidate_idx, (idx, tool_call_id) in enumerate(candidates): - is_newest_candidate = candidate_idx == len(candidates) - 1 - if is_newest_candidate and estimate <= budget: - break - if tool_call_id in compacted_tool_call_ids: - continue - if updated is messages: - updated = [dict(m) for m in messages] - compacted_tool_call_ids.add(tool_call_id) - self._compact_tool_result_at(updated, idx) - estimate, source = estimate_prompt_tokens_chain( - config.provider, - config.model, - updated, - tools, - ) - if estimate <= target: - break - - logger.debug( - "In-flight context compaction for {}: prompt={} budget={} target={} via {}, ids={}", - config.session_key or "default", - estimate, - budget, - target, - source, - len(compacted_tool_call_ids), - ) - return updated - def snip_history( self, config: ContextGovernanceConfig, messages: list[dict[str, Any]], + *, + tool_definitions: list[dict[str, Any]] | None, + force: bool = False, ) -> list[dict[str, Any]]: if not messages or not config.context_window_tokens: return messages @@ -399,14 +424,13 @@ class ContextGovernor: if budget <= 0: return messages - tools = config.tools.get_definitions() estimate, _ = estimate_prompt_tokens_chain( config.provider, config.model, messages, - tools, + tool_definitions, ) - if estimate <= budget: + if not force and estimate <= budget: return messages system_messages = [dict(msg) for msg in messages if msg.get("role") == "system"] @@ -419,7 +443,7 @@ class ContextGovernor: config.provider, config.model, system_messages, - tools, + tool_definitions, ) remaining_budget = max(0, budget - max(system_tokens, fixed_tokens)) kept: list[dict[str, Any]] = [] @@ -434,16 +458,6 @@ class ContextGovernor: return system_messages + self._legal_history_tail(kept, non_system) - @staticmethod - def _tool_result_compaction_message(message: dict[str, Any]) -> str: - name = message.get("name", "tool") - return ( - f"Error: The previous {name} result was compacted to fit context because it was too " - "large. Do not repeat the same call unchanged. Retry with a narrower path, query, " - "range, or result limit, use another tool, or tell the user the task cannot fit in " - "the available context." - ) - def _legal_history_tail( self, kept: list[dict[str, Any]], @@ -462,50 +476,3 @@ class ContextGovernor: if messages[idx].get("role") == "user": return messages[idx:] return [] - - def _apply_recorded_compactions( - self, - messages: list[dict[str, Any]], - compacted_tool_call_ids: set[str], - ) -> list[dict[str, Any]]: - if not compacted_tool_call_ids: - return messages - updated = messages - for idx, msg in enumerate(messages): - if msg.get("role") != "tool": - continue - tool_call_id = msg.get("tool_call_id") - if not tool_call_id or str(tool_call_id) not in compacted_tool_call_ids: - continue - compaction_message = self._tool_result_compaction_message(msg) - if msg.get("content") == compaction_message: - continue - if updated is messages: - updated = [dict(m) for m in messages] - updated[idx]["content"] = compaction_message - return updated - - def _inflight_compaction_candidates( - self, - config: ContextGovernanceConfig, - messages: list[dict[str, Any]], - compacted_tool_call_ids: set[str], - ) -> list[tuple[int, str]]: - compactable: list[tuple[int, str]] = [] - for idx, msg in enumerate(messages): - if idx < config.inflight_start_index: - continue - if msg.get("role") != "tool" or msg.get("name") not in COMPACTABLE_TOOLS: - continue - tool_call_id = msg.get("tool_call_id") - if not tool_call_id or str(tool_call_id) in compacted_tool_call_ids: - continue - content = msg.get("content") - if not isinstance(content, str) or len(content) < MICROCOMPACT_MIN_CHARS: - continue - compactable.append((idx, str(tool_call_id))) - - return compactable - - def _compact_tool_result_at(self, messages: list[dict[str, Any]], idx: int) -> None: - messages[idx]["content"] = self._tool_result_compaction_message(messages[idx]) diff --git a/nanobot/agent/runner.py b/nanobot/agent/runner.py index 36d5303f7..a19a9d9ae 100644 --- a/nanobot/agent/runner.py +++ b/nanobot/agent/runner.py @@ -139,6 +139,17 @@ class AgentRunResult: provider_state: ProviderConversationState | None = field(default=None, repr=False) +@dataclass(slots=True) +class _ModelRequestState: + """Per-run state used to govern the next provider request.""" + + config: ContextGovernanceConfig + conversation: ProviderConversationStateController + usage: LLMUsage | None = None + messages: list[dict[str, Any]] | None = None + tool_definitions: list[dict[str, Any]] | None = None + + class AgentRunner: """Run a tool-capable LLM loop without product-layer concerns.""" @@ -500,7 +511,6 @@ class AgentRunner: length_recovery_parts: list[str] = [] had_injections = False injection_cycles = 0 - compacted_tool_call_ids: set[str] = set() pending_stream_content: str | None = None conversation_state = ProviderConversationStateController( provider=spec.runtime.provider, @@ -519,39 +529,29 @@ class AgentRunner: context_window_tokens=spec.runtime.context_window_tokens, context_block_limit=spec.context_block_limit, max_tokens=spec.runtime.generation.max_tokens, - inflight_start_index=len(messages), + ) + request_state = _ModelRequestState( + config=governance_config, + conversation=conversation_state, ) for iteration in range(spec.max_iterations): - # Keep the persisted conversation untouched. Context governance - # may repair or compact historical messages for the model, but - # those synthetic edits must not shift the append boundary used - # later when the caller saves only the new turn. A governance - # failure must stop the run instead of sending an ungoverned copy. - messages_for_model = self.context_governor.prepare_for_model( - governance_config, - messages, - compacted_tool_call_ids, - ) context = AgentHookContext( iteration=iteration, messages=messages, session_key=spec.session_key, ) await hook.before_iteration(context) - provider_context = conversation_state.prepare_request( - messages, - context_window_tokens=spec.runtime.context_window_tokens, - model_messages=messages_for_model, - ) response = await self._request_model( spec, - messages_for_model, + messages, hook, context, - conversation_state=conversation_state, - provider_context=provider_context, + request_state=request_state, + transcript=messages, ) + assert request_state.messages is not None + messages_for_model = request_state.messages conversation_state.observe_response(response, messages) context.response = response context.tool_calls = list(response.tool_calls) @@ -563,7 +563,7 @@ class AgentRunner: response.content, ) response.content = cleaned_content - raw_usage = self._usage_or_estimate(spec, messages_for_model, response) + raw_usage = self._record_request_usage(spec, request_state, response) context.usage = raw_usage usage = self._merge_usage(usage, raw_usage) if reasoning_text and not context.streamed_reasoning: @@ -637,7 +637,6 @@ class AgentRunner: self.context_governor.prepare_for_model( governance_config, messages, - compacted_tool_call_ids, ) if response.provider_state is not None else None @@ -703,14 +702,13 @@ class AgentRunner: ) if hook.wants_streaming(): await hook.on_stream_end(context, resuming=False) - retry_messages = self._finalization_retry_messages(messages_for_model) response = await self._request_finalization_retry( spec, messages_for_model, + request_state=request_state, transcript=messages, - conversation_state=conversation_state, ) - retry_usage = self._usage_or_estimate(spec, retry_messages, response) + retry_usage = self._record_request_usage(spec, request_state, response) usage = self._merge_usage(usage, retry_usage) raw_usage = self._merge_usage(raw_usage, retry_usage) context.response = response @@ -897,7 +895,7 @@ class AgentRunner: hook, messages, usage, - conversation_state, + request_state=request_state, ) if terminal_content is None: terminal_content = self._max_iterations_fallback(spec) @@ -944,6 +942,60 @@ class AgentRunner: kwargs["reasoning_effort"] = generation.reasoning_effort return kwargs + def _prepare_model_request( + self, + state: _ModelRequestState, + messages: list[dict[str, Any]], + *, + tool_definitions: list[dict[str, Any]] | None, + transcript: list[dict[str, Any]] | None = None, + ) -> tuple[list[dict[str, Any]], ProviderCallContext | None]: + """Prepare, fit, and record the exact payload sent to a provider.""" + prepared = self.context_governor.prepare_for_model(state.config, messages) + supplemental_messages = ( + [prepared[-1]] if transcript is not None and tool_definitions is None else None + ) + model_messages = None if supplemental_messages is not None else prepared + request_context_tokens = ( + state.conversation.estimate_request_context_tokens( + transcript, + model_messages=model_messages, + supplemental_messages=supplemental_messages, + tool_definitions=tool_definitions, + ) + if transcript is not None + else None + ) + usage_matches_messages = ( + state.messages is not None + and prepared == state.messages + and tool_definitions == state.tool_definitions + ) + prepared, fitted = self.context_governor.fit_request( + state.config, + prepared, + state.usage, + usage_matches_messages=usage_matches_messages, + tool_definitions=tool_definitions, + request_context_tokens=request_context_tokens, + ) + provider_context = ( + state.conversation.prepare_request( + transcript, + context_window_tokens=state.config.context_window_tokens, + model_messages=model_messages, + supplemental_messages=supplemental_messages, + resume_state=not fitted, + ) + if transcript is not None + else state.conversation.independent_request_context( + context_window_tokens=state.config.context_window_tokens, + ) + ) + state.messages = deepcopy(prepared) + state.tool_definitions = deepcopy(tool_definitions) + return prepared, provider_context + async def _request_model( self, spec: AgentRunSpec, @@ -951,16 +1003,23 @@ class AgentRunner: hook: AgentHook, context: AgentHookContext, *, + request_state: _ModelRequestState, malformed_retry: bool = False, - conversation_state: ProviderConversationStateController, - provider_context: ProviderCallContext | None = None, + transcript: list[dict[str, Any]] | None, ) -> LLMResponse: timeout_s = self._resolve_llm_timeout_s(spec) + tool_definitions = spec.tools.get_definitions() + messages, provider_context = self._prepare_model_request( + request_state, + messages, + tool_definitions=tool_definitions, + transcript=transcript, + ) kwargs = self._build_request_kwargs( spec, messages, - tools=spec.tools.get_definitions(), + tools=tool_definitions, ) wants_streaming = hook.wants_streaming() @@ -1138,11 +1197,9 @@ class AgentRunner: ) return await self._request_model( spec, retry_messages, hook, context, + request_state=request_state, malformed_retry=True, - conversation_state=conversation_state, - provider_context=conversation_state.independent_request_context( - context_window_tokens=spec.runtime.context_window_tokens, - ), + transcript=None, ) if ( all_dropped @@ -1158,9 +1215,7 @@ class AgentRunner: return await self._request_no_tools( spec, fallback_messages, - provider_context=conversation_state.independent_request_context( - context_window_tokens=spec.runtime.context_window_tokens, - ), + request_state=request_state, ) return response @@ -1228,21 +1283,17 @@ class AgentRunner: spec: AgentRunSpec, messages: list[dict[str, Any]], *, + request_state: _ModelRequestState, transcript: list[dict[str, Any]], - conversation_state: ProviderConversationStateController, ) -> LLMResponse: retry_messages = self._finalization_retry_messages(messages) - provider_context = conversation_state.prepare_request( - transcript, - context_window_tokens=spec.runtime.context_window_tokens, - supplemental_messages=[retry_messages[-1]], - ) response = await self._request_no_tools( spec, retry_messages, - provider_context=provider_context, + request_state=request_state, + transcript=transcript, ) - conversation_state.observe_response( + request_state.conversation.observe_response( response, transcript, adopt_candidate_state=False, @@ -1261,16 +1312,15 @@ class AgentRunner: hook: AgentHook, messages: list[dict[str, Any]], usage: LLMUsage | None, - conversation_state: ProviderConversationStateController, + *, + request_state: _ModelRequestState, ) -> tuple[str | None, LLMUsage | None]: retry_messages = self._budget_exhausted_finalization_messages(messages) try: response = await self._request_no_tools( spec, retry_messages, - provider_context=conversation_state.independent_request_context( - context_window_tokens=spec.runtime.context_window_tokens, - ), + request_state=request_state, ) except Exception: logger.exception( @@ -1279,7 +1329,7 @@ class AgentRunner: ) return None, usage - raw_usage = self._usage_or_estimate(spec, retry_messages, response) + raw_usage = self._record_request_usage(spec, request_state, response) usage = self._merge_usage(usage, raw_usage) if response.finish_reason == "error" or response.has_tool_calls: logger.warning( @@ -1308,8 +1358,15 @@ class AgentRunner: spec: AgentRunSpec, messages: list[dict[str, Any]], *, - provider_context: ProviderCallContext | None = None, + request_state: _ModelRequestState, + transcript: list[dict[str, Any]] | None = None, ) -> LLMResponse: + messages, provider_context = self._prepare_model_request( + request_state, + messages, + tool_definitions=None, + transcript=transcript, + ) kwargs = self._build_request_kwargs( spec, messages, @@ -1321,17 +1378,18 @@ class AgentRunner: ) timeout_s = self._resolve_llm_timeout_s(spec) try: - return ( + response = ( await coro if timeout_s is None else await asyncio.wait_for(coro, timeout=timeout_s) ) except asyncio.TimeoutError: - return LLMResponse( + response = LLMResponse( content=f"Error calling LLM: timed out after {timeout_s:g}s", finish_reason="error", error_kind="timeout", ) + return response @staticmethod def _resolve_llm_timeout_s(spec: AgentRunSpec) -> float | None: @@ -1373,33 +1431,53 @@ class AgentRunner: spec: AgentRunSpec, messages: list[dict[str, Any]], response: LLMResponse, + *, + tool_definitions: list[dict[str, Any]] | None, ) -> 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) + usage = self._estimate_response_usage( + spec, + messages, + response, + tool_definitions=tool_definitions, + ) return usage.with_timing( generation_ms=response.generation_ms, ttft_ms=response.ttft_ms, ) + def _record_request_usage( + self, + spec: AgentRunSpec, + state: _ModelRequestState, + response: LLMResponse, + ) -> LLMUsage | None: + assert state.messages is not None + state.usage = self._usage_or_estimate( + spec, + state.messages, + response, + tool_definitions=state.tool_definitions, + ) + return state.usage + def _estimate_response_usage( self, spec: AgentRunSpec, messages: list[dict[str, Any]], response: LLMResponse, + *, + tool_definitions: list[dict[str, Any]] | None, ) -> LLMUsage: - try: - tools = spec.tools.get_definitions() - except Exception: - tools = None prompt_tokens, _ = estimate_prompt_tokens_chain( spec.runtime.provider, spec.runtime.model, messages, - tools, + tool_definitions, ) assistant_message = build_assistant_message( response.content or "", diff --git a/nanobot/providers/conversation_state.py b/nanobot/providers/conversation_state.py index 6d891603a..7c70a1848 100644 --- a/nanobot/providers/conversation_state.py +++ b/nanobot/providers/conversation_state.py @@ -11,6 +11,7 @@ from nanobot.providers.base import ( ProviderCallContext, ProviderConversationState, ) +from nanobot.utils.helpers import estimate_prompt_tokens_chain _PROVIDER_STATE_OUTPUT_META = "provider_state_output" _PROVIDER_STATE_BOUNDARY_META = "provider_state_boundary" @@ -69,6 +70,37 @@ class ProviderConversationStateController: session_id=self._session_id, ) + def estimate_request_context_tokens( + self, + messages: list[dict[str, Any]], + *, + model_messages: list[dict[str, Any]] | None = None, + supplemental_messages: list[dict[str, Any]] | None = None, + tool_definitions: list[dict[str, Any]] | None = None, + ) -> int | None: + """Estimate resumed state plus the pending delta for the next request.""" + state = self.checkpoint(messages, model_messages=model_messages) + if state is None: + return None + context_tokens = state.payload.get("context_tokens") + if ( + isinstance(context_tokens, bool) + or not isinstance(context_tokens, int) + or context_tokens < 0 + ): + return None + pending_messages = [ + *state.pending_messages, + *(supplemental_messages or []), + ] + delta_tokens, _ = estimate_prompt_tokens_chain( + self._provider, + self._model, + pending_messages, + tool_definitions, + ) + return context_tokens + max(0, delta_tokens) + def prepare_request( self, messages: list[dict[str, Any]], @@ -76,11 +108,20 @@ class ProviderConversationStateController: context_window_tokens: int | None, model_messages: list[dict[str, Any]] | None = None, supplemental_messages: list[dict[str, Any]] | None = None, + resume_state: bool = True, ) -> ProviderCallContext | None: - """Build typed context for the next request and remember its durable delta.""" + """Build context for the next request and remember its durable delta. + + ``resume_state=False`` abandons opaque history when local request + fitting has produced a new independent model-facing context. + """ independent_context = self.independent_request_context( context_window_tokens=context_window_tokens, ) + if not resume_state: + self._state = None + self._request_messages = [] + return independent_context if self._state is None: self._request_messages = [] return independent_context diff --git a/tests/agent/test_dream.py b/tests/agent/test_dream.py index f14284fc6..4c1378184 100644 --- a/tests/agent/test_dream.py +++ b/tests/agent/test_dream.py @@ -426,7 +426,7 @@ class TestEphemeralDirect: bus=bus, provider=provider, workspace=tmp_path, - context_window_tokens=8000, + context_window_tokens=32_000, ) return loop, store @@ -606,7 +606,7 @@ class TestEphemeralDirect: bus=MessageBus(), provider=provider, workspace=tmp_path, - context_window_tokens=8000, + context_window_tokens=32_000, ) await loop.process_direct( @@ -666,7 +666,7 @@ class TestEphemeralHooks: bus=bus, provider=provider, workspace=tmp_path, - context_window_tokens=8000, + context_window_tokens=32_000, hooks=[spy], ) diff --git a/tests/agent/test_loop_consolidation_tokens.py b/tests/agent/test_loop_consolidation_tokens.py index eab8edc9d..0b304d971 100644 --- a/tests/agent/test_loop_consolidation_tokens.py +++ b/tests/agent/test_loop_consolidation_tokens.py @@ -29,6 +29,9 @@ def _make_loop( workspace=tmp_path, model="test-model", context_window_tokens=context_window_tokens, + # These tests isolate Memory consolidation; Runner request fitting is + # covered separately with realistic context windows. + context_block_limit=10_000, ) loop.tools.get_definitions = MagicMock(return_value=[]) loop.consolidator._SAFETY_BUFFER = 0 diff --git a/tests/agent/test_runner_core.py b/tests/agent/test_runner_core.py index 62f7d0a60..af3d243ce 100644 --- a/tests/agent/test_runner_core.py +++ b/tests/agent/test_runner_core.py @@ -11,6 +11,7 @@ import pytest from agent.runner_helpers import make_run_spec from nanobot.agent.context import TranscriptInput +from nanobot.agent.context_governance import ContextWindowExceededError from nanobot.config.schema import AgentDefaults from nanobot.providers.base import ( LLMProvider, @@ -86,6 +87,7 @@ def test_usage_or_estimate_replaces_reported_zero_for_content(monkeypatch) -> No _make_usage_spec(provider, tools), [{"role": "user", "content": "hello"}], response, + tool_definitions=tools.get_definitions(), ) assert usage == LLMUsage.estimated(input_tokens=12, output_tokens=7).with_timing( @@ -130,6 +132,7 @@ def test_usage_or_estimate_counts_tool_call_output_for_reported_zero(monkeypatch _make_usage_spec(provider, tools), [{"role": "user", "content": "hello"}], response, + tool_definitions=tools.get_definitions(), ) assert usage == LLMUsage.estimated(input_tokens=13, output_tokens=9) @@ -162,6 +165,7 @@ def test_usage_or_estimate_counts_error_without_estimating_tokens( _make_usage_spec(provider, tools), [{"role": "user", "content": "hello"}], response, + tool_definitions=tools.get_definitions(), ) assert usage is not None @@ -197,6 +201,7 @@ def test_usage_or_estimate_trusts_positive_reported_total(monkeypatch) -> None: _make_usage_spec(provider, tools), [{"role": "user", "content": "hello"}], response, + tool_definitions=tools.get_definitions(), ) assert usage is not None @@ -366,14 +371,12 @@ async def test_runner_replays_provider_state_without_chat_projection_duplicates( @pytest.mark.asyncio -async def test_runner_governs_tool_result_before_adding_it_to_provider_state(): +async def test_runner_preserves_tool_result_before_rejecting_unfit_followup(): from nanobot.agent.runner import AgentRunner provider = MagicMock(spec=LLMProvider) provider.can_resume_conversation_state.return_value = True - provider.supports_native_compaction.return_value = False calls = 0 - captured_context: ProviderCallContext | None = None checkpoints: list[dict] = [] state = ProviderConversationState( kind="openai_responses", @@ -384,7 +387,7 @@ async def test_runner_governs_tool_result_before_adding_it_to_provider_state(): ) async def chat_with_retry(**kwargs): - nonlocal calls, captured_context + nonlocal calls calls += 1 if calls == 1: return LLMResponse( @@ -398,7 +401,6 @@ async def test_runner_governs_tool_result_before_adding_it_to_provider_state(): ], provider_state=state, ) - captured_context = kwargs["provider_context"] return LLMResponse(content="done") provider.chat_with_retry = chat_with_retry @@ -409,37 +411,36 @@ async def test_runner_governs_tool_result_before_adding_it_to_provider_state(): async def checkpoint(payload: dict) -> None: checkpoints.append(payload) - await AgentRunner().run(make_run_spec( - provider, - initial_messages=[ - {"role": "system", "content": "system"}, - {"role": "user", "content": "read the file"}, - ], - tools=tools, - model="gpt-5.6", - context_window_tokens=3_000, - context_block_limit=200, - max_tokens=1_000, - max_iterations=3, - max_tool_result_chars=10_000, - checkpoint_callback=checkpoint, - )) + with pytest.raises(ContextWindowExceededError): + await AgentRunner().run(make_run_spec( + provider, + initial_messages=[ + {"role": "system", "content": "system"}, + {"role": "user", "content": "read the file"}, + ], + tools=tools, + model="gpt-5.6", + context_window_tokens=3_000, + context_block_limit=200, + max_tokens=1_000, + max_iterations=3, + max_tool_result_chars=10_000, + checkpoint_callback=checkpoint, + )) - assert captured_context is not None - assert captured_context.conversation_state is not None - pending = captured_context.conversation_state.pending_messages - assert len(pending) == 1 - assert pending[0]["role"] == "tool" - assert "compacted to fit context" in pending[0]["content"] - assert pending[0]["content"] != "x" * 5_000 + assert calls == 1 completed_checkpoint = next( checkpoint for checkpoint in checkpoints if checkpoint["phase"] == "tools_completed" ) checkpoint_pending = completed_checkpoint["provider_state"].pending_messages - assert "compacted to fit context" in checkpoint_pending[0]["content"] - assert checkpoint_pending[0]["content"] != "x" * 5_000 + assert checkpoint_pending == [{ + "role": "tool", + "tool_call_id": "call_1", + "name": "read_file", + "content": "x" * 5_000, + }] @pytest.mark.asyncio diff --git a/tests/agent/test_runner_governance.py b/tests/agent/test_runner_governance.py index 88d36a8c3..fe09bc272 100644 --- a/tests/agent/test_runner_governance.py +++ b/tests/agent/test_runner_governance.py @@ -1,8 +1,7 @@ -"""Tests for AgentRunner context governance: backfill, orphan cleanup, microcompact, snip_history.""" +"""Tests for AgentRunner context governance: repair and request fitting.""" from __future__ import annotations -from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -12,11 +11,14 @@ from nanobot.agent.context_governance import ( BACKFILL_CONTENT, ContextGovernanceConfig, ContextGovernor, + ContextWindowExceededError, ) from nanobot.agent.runner import AgentRunSpec from nanobot.config.schema import AgentDefaults from nanobot.providers.base import ( + LLMProvider, LLMResponse, + LLMUsage, ProviderConversationState, ToolCallRequest, ) @@ -28,8 +30,6 @@ def _governance_config( provider, tools, spec: AgentRunSpec, - *, - inflight_start_index: int = 0, ) -> ContextGovernanceConfig: return ContextGovernanceConfig( provider=provider, @@ -41,7 +41,6 @@ def _governance_config( context_window_tokens=spec.runtime.context_window_tokens, context_block_limit=spec.context_block_limit, max_tokens=spec.runtime.generation.max_tokens, - inflight_start_index=inflight_start_index, ) @@ -89,6 +88,508 @@ async def test_runner_propagates_context_governance_failure(): provider.chat_with_retry.assert_not_awaited() +@pytest.mark.asyncio +async def test_runner_locally_fits_oversized_initial_transcript(monkeypatch): + from nanobot.agent.runner import AgentRunner + + provider = MagicMock(spec=LLMProvider) + provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="done")) + tools = MagicMock() + tools.get_definitions.return_value = [] + old_content = "x" * 20_000 + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_prompt_tokens_chain", + lambda _provider, _model, messages, _tools: ( + (600, "test-counter") + if any(message.get("content") == old_content for message in messages) + else (100, "test-counter") + ), + ) + + result = await AgentRunner().run(make_run_spec( + provider, + initial_messages=[ + {"role": "system", "content": "system"}, + {"role": "user", "content": "old question"}, + {"role": "assistant", "content": old_content}, + {"role": "user", "content": "continue"}, + ], + tools=tools, + model="local-model", + context_window_tokens=2_000, + context_block_limit=500, + max_tokens=100, + max_iterations=1, + max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, + )) + + assert provider.chat_with_retry.await_args.kwargs["messages"] == [ + {"role": "system", "content": "system"}, + {"role": "user", "content": "continue"}, + ] + assert any(message.get("content") == old_content for message in result.messages) + + +@pytest.mark.asyncio +async def test_runner_governs_messages_added_by_before_iteration_hook(monkeypatch): + from nanobot.agent.hook import AgentHook, AgentHookContext + from nanobot.agent.runner import AgentRunner + + provider = MagicMock(spec=LLMProvider) + provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="unexpected")) + tools = MagicMock() + tools.get_definitions.return_value = [] + oversized = "hook-added-oversized-message" + + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_prompt_tokens_chain", + lambda _provider, _model, messages, _tools: ( + (2_000, "test-counter") + if any(message.get("content") == oversized for message in messages) + else (100, "test-counter") + ), + ) + + class MutatingHook(AgentHook): + async def before_iteration(self, context: AgentHookContext) -> None: + context.messages.append({"role": "user", "content": oversized}) + + with pytest.raises(ContextWindowExceededError): + await AgentRunner().run(make_run_spec( + provider, + initial_messages=[{"role": "user", "content": "hello"}], + tools=tools, + model="local-model", + context_window_tokens=2_000, + context_block_limit=500, + max_iterations=1, + max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, + hook=MutatingHook(), + )) + + provider.chat_with_retry.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_runner_drops_resumable_provider_state_when_request_is_fitted(monkeypatch): + from nanobot.agent.runner import AgentRunner + + provider = MagicMock(spec=LLMProvider) + provider.can_resume_conversation_state.return_value = True + captured_contexts = [] + old_content = "old-oversized-history" + candidate = ProviderConversationState( + kind="openai_responses", + provider="openai:test", + model="local-model", + version=1, + payload={"items": [{"type": "message", "content": "fresh state"}]}, + ) + + async def chat_with_retry(*, provider_context=None, **_kwargs): + captured_contexts.append(provider_context) + return LLMResponse( + content="done", + usage=LLMUsage.reported(input_tokens=100, output_tokens=10), + provider_state=candidate, + ) + + provider.chat_with_retry = chat_with_retry + tools = MagicMock() + tools.get_definitions.return_value = [] + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_prompt_tokens_chain", + lambda _provider, _model, messages, _tools: ( + (600, "test-counter") + if any(message.get("content") == old_content for message in messages) + else (100, "test-counter") + ), + ) + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_message_tokens", + lambda message: 450 if message.get("content") == old_content else 50, + ) + saved_state = ProviderConversationState( + kind="openai_responses", + provider="openai:test", + model="local-model", + version=1, + payload={"items": [{"type": "message", "content": "stale state"}]}, + ) + + result = await AgentRunner().run(make_run_spec( + provider, + initial_messages=[ + {"role": "assistant", "content": old_content}, + {"role": "user", "content": "continue"}, + ], + tools=tools, + model="local-model", + context_window_tokens=2_000, + context_block_limit=500, + max_iterations=1, + max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, + provider_state=saved_state, + )) + + assert captured_contexts[0].conversation_state is None + assert result.provider_state is not None + assert result.provider_state.payload == candidate.payload + + +@pytest.mark.asyncio +async def test_runner_fits_each_malformed_retry_with_its_actual_tools(monkeypatch): + from nanobot.agent.runner import AgentRunner + + provider = MagicMock(spec=LLMProvider) + calls: list[dict] = [] + estimated_tools: list[object] = [] + definitions = [{"type": "function", "function": {"name": "read_file"}}] + + async def chat_with_retry(*, messages, tools=None, **_kwargs): + calls.append({"messages": [dict(message) for message in messages], "tools": tools}) + if len(calls) < 3: + return LLMResponse( + content="bad tool request", + tool_calls=[ToolCallRequest(id=f"bad_{len(calls)}", name=None, arguments={})], + finish_reason="tool_calls", + usage=LLMUsage.reported(input_tokens=100, output_tokens=10), + ) + return LLMResponse( + content="recovered", + usage=LLMUsage.reported(input_tokens=100, output_tokens=10), + ) + + def estimate(_provider, _model, messages, _tools): + estimated_tools.append(_tools) + user_count = sum(message.get("role") == "user" for message in messages) + return (600 if user_count > 1 else 100), "test-counter" + + provider.chat_with_retry = chat_with_retry + tools = MagicMock() + tools.get_definitions.return_value = definitions + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_prompt_tokens_chain", + estimate, + ) + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_message_tokens", + lambda _message: 300, + ) + + result = await AgentRunner().run(make_run_spec( + provider, + initial_messages=[{"role": "user", "content": "use a tool"}], + tools=tools, + model="local-model", + context_window_tokens=2_000, + context_block_limit=500, + max_iterations=1, + max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, + )) + + assert [call["tools"] for call in calls] == [definitions, definitions, None] + assert definitions in estimated_tools + assert None in estimated_tools + assert [len(call["messages"]) for call in calls] == [1, 1, 1] + assert result.final_content == "recovered" + assert result.messages == [ + {"role": "user", "content": "use a tool"}, + {"role": "assistant", "content": "recovered"}, + ] + + +@pytest.mark.asyncio +async def test_runner_fits_empty_response_finalization_before_dispatch(monkeypatch): + from nanobot.agent.runner import AgentRunner + + provider = MagicMock(spec=LLMProvider) + calls: list[dict] = [] + + async def chat_with_retry(*, messages, tools=None, **_kwargs): + calls.append({"messages": [dict(message) for message in messages], "tools": tools}) + if len(calls) < 3: + return LLMResponse( + content=None, + usage=LLMUsage.reported(input_tokens=100, output_tokens=1), + ) + return LLMResponse( + content="finalized", + usage=LLMUsage.reported(input_tokens=100, output_tokens=10), + ) + + def estimate(_provider, _model, messages, _tools): + contents = [str(message.get("content") or "") for message in messages] + has_original = "do task" in contents + has_finalization = any("conversation above" in content for content in contents) + return (600 if has_original and has_finalization else 100), "test-counter" + + provider.chat_with_retry = chat_with_retry + tools = MagicMock() + tools.get_definitions.return_value = [] + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_prompt_tokens_chain", + estimate, + ) + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_message_tokens", + lambda _message: 300, + ) + + result = await AgentRunner().run(make_run_spec( + provider, + initial_messages=[{"role": "user", "content": "do task"}], + tools=tools, + model="local-model", + context_window_tokens=2_000, + context_block_limit=500, + max_iterations=3, + max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, + )) + + assert len(calls) == 3 + assert calls[-1]["tools"] is None + assert all(message.get("content") != "do task" for message in calls[-1]["messages"]) + assert result.final_content == "finalized" + + +@pytest.mark.asyncio +async def test_runner_fits_max_iteration_finalization_before_dispatch(monkeypatch): + from nanobot.agent.runner import AgentRunner + + provider = MagicMock(spec=LLMProvider) + calls: list[dict] = [] + oversized_result = "oversized-current-tool-result" + + async def chat_with_retry(*, messages, tools=None, **_kwargs): + calls.append({"messages": [dict(message) for message in messages], "tools": tools}) + if len(calls) == 1: + return LLMResponse( + content="working", + tool_calls=[ToolCallRequest(id="call_1", name="read_file", arguments={})], + finish_reason="tool_calls", + usage=LLMUsage.reported(input_tokens=100, output_tokens=10), + ) + return LLMResponse( + content="safe summary", + usage=LLMUsage.reported(input_tokens=100, output_tokens=10), + ) + + def estimate(_provider, _model, messages, _tools): + has_oversized = any( + message.get("content") == oversized_result for message in messages + ) + return (600 if has_oversized else 100), "test-counter" + + provider.chat_with_retry = chat_with_retry + tools = MagicMock() + tools.get_definitions.return_value = [] + tools.execute = AsyncMock(return_value=oversized_result) + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_prompt_tokens_chain", + estimate, + ) + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_message_tokens", + lambda message: 600 if message.get("content") == oversized_result else 50, + ) + + result = await AgentRunner().run(make_run_spec( + provider, + initial_messages=[{"role": "user", "content": "inspect"}], + tools=tools, + model="local-model", + context_window_tokens=2_000, + context_block_limit=500, + max_iterations=1, + max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, + )) + + assert len(calls) == 2 + assert calls[-1]["tools"] is None + assert all( + message.get("content") != oversized_result + for message in calls[-1]["messages"] + ) + assert any(message.get("content") == oversized_result for message in result.messages) + assert result.final_content == "safe summary" + + +@pytest.mark.parametrize( + ("input_tokens", "expected_fitted"), + [(500, True), (100, False)], +) +def test_matching_reported_provider_usage_avoids_local_estimate( + monkeypatch, + input_tokens, + expected_fitted, +): + provider = MagicMock(spec=LLMProvider) + tools = MagicMock() + tools.get_definitions.return_value = [] + spec = make_run_spec( + provider, + initial_messages=[{"role": "user", "content": "hello"}], + tools=tools, + model="local-model", + context_window_tokens=2_000, + context_block_limit=500, + max_iterations=1, + max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, + ) + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_prompt_tokens_chain", + lambda *_args, **_kwargs: (_ for _ in ()).throw( + AssertionError("matching provider usage must be authoritative") + ), + ) + + governor = ContextGovernor() + monkeypatch.setattr(governor, "fit_to_budget", lambda *_args, **_kwargs: []) + _messages, fitted = governor.fit_request( + _governance_config(provider, tools, spec), + spec.initial_messages, + LLMUsage.reported(input_tokens=input_tokens, output_tokens=10), + usage_matches_messages=True, + tool_definitions=tools.get_definitions(), + ) + + assert fitted is expected_fitted + + +def test_changed_messages_use_local_estimate_after_reported_usage(monkeypatch): + provider = MagicMock(spec=LLMProvider) + tools = MagicMock() + tools.get_definitions.return_value = [] + spec = make_run_spec( + provider, + initial_messages=[{"role": "user", "content": "new tool output"}], + tools=tools, + model="local-model", + context_window_tokens=2_000, + context_block_limit=500, + max_iterations=1, + max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, + ) + estimate = MagicMock(return_value=(600, "test-counter")) + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_prompt_tokens_chain", + estimate, + ) + + governor = ContextGovernor() + monkeypatch.setattr(governor, "fit_to_budget", lambda *_args, **_kwargs: []) + _messages, fitted = governor.fit_request( + _governance_config(provider, tools, spec), + spec.initial_messages, + LLMUsage.reported(input_tokens=900, output_tokens=10), + usage_matches_messages=False, + tool_definitions=tools.get_definitions(), + ) + + assert fitted is True + estimate.assert_called_once() + + +@pytest.mark.asyncio +async def test_runner_counts_resumed_provider_state_before_dispatch(monkeypatch): + from nanobot.agent.runner import AgentRunner + + provider = MagicMock(spec=LLMProvider) + provider.can_resume_conversation_state.return_value = True + captured_contexts = [] + + async def chat_with_retry(*, provider_context=None, **_kwargs): + captured_contexts.append(provider_context) + return LLMResponse( + content="done", + usage=LLMUsage.reported(input_tokens=100, output_tokens=10), + ) + + provider.chat_with_retry = chat_with_retry + tools = MagicMock() + tools.get_definitions.return_value = [] + current_message = {"role": "user", "content": "new delta"} + saved_state = ProviderConversationState( + kind="openai_responses", + provider="openai:test", + model="local-model", + version=1, + payload={ + "items": [{"type": "reasoning", "encrypted_content": "opaque"}], + "context_tokens": 450, + }, + pending_messages=[current_message], + ) + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_prompt_tokens_chain", + lambda *_args, **_kwargs: (100, "test-counter"), + ) + monkeypatch.setattr( + "nanobot.providers.conversation_state.estimate_prompt_tokens_chain", + lambda *_args, **_kwargs: (100, "test-counter"), + ) + + result = await AgentRunner().run(make_run_spec( + provider, + initial_messages=[current_message], + tools=tools, + model="local-model", + context_window_tokens=2_000, + context_block_limit=500, + max_iterations=1, + max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, + provider_state=saved_state, + )) + + assert captured_contexts[0].conversation_state is None + assert result.messages == [ + current_message, + {"role": "assistant", "content": "done"}, + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("context_block_limit", "expected_budget"), + [(500, 500), (None, 0)], +) +async def test_runner_refuses_locally_fitted_request_that_still_cannot_fit( + monkeypatch, + context_block_limit, + expected_budget, +): + from nanobot.agent.runner import AgentRunner + + provider = MagicMock(spec=LLMProvider) + provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="unexpected")) + tools = MagicMock() + tools.get_definitions.return_value = [] + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_prompt_tokens_chain", + lambda *_args, **_kwargs: (2_000, "test-counter"), + ) + + with pytest.raises(ContextWindowExceededError) as exc_info: + await AgentRunner().run(make_run_spec( + provider, + initial_messages=[ + {"role": "system", "content": "oversized system"}, + {"role": "user", "content": "oversized user"}, + ], + tools=tools, + model="local-model", + context_window_tokens=1_000, + context_block_limit=context_block_limit, + max_iterations=1, + max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, + )) + + assert exc_info.value.estimated_tokens == 2_000 + assert exc_info.value.input_budget == expected_budget + provider.chat_with_retry.assert_not_awaited() + + def test_snip_history_drops_orphaned_tool_results_from_trimmed_slice(monkeypatch): provider = MagicMock() tools = MagicMock() @@ -130,7 +631,11 @@ def test_snip_history_drops_orphaned_tool_results_from_trimmed_slice(monkeypatch lambda msg: token_sizes.get(str(msg.get("content")), 40), ) - trimmed = ContextGovernor().snip_history(_governance_config(provider, tools, spec), messages) + trimmed = ContextGovernor().snip_history( + _governance_config(provider, tools, spec), + messages, + tool_definitions=tools.get_definitions(), + ) # After the fix, the user message is recovered so the sequence is valid # for providers that require system → user (e.g. GLM error 1214). @@ -182,7 +687,11 @@ def test_snip_history_reserves_budget_for_tool_definitions(monkeypatch): lambda msg: token_sizes.get(str(msg.get("content")), 40), ) - trimmed = ContextGovernor().snip_history(_governance_config(provider, tools, spec), messages) + trimmed = ContextGovernor().snip_history( + _governance_config(provider, tools, spec), + messages, + tool_definitions=tools.get_definitions(), + ) contents = [message.get("content") for message in trimmed] assert contents == ["system", "recent two"] @@ -465,260 +974,6 @@ async def test_runner_backfill_only_mutates_model_context_not_returned_messages( ] -# --------------------------------------------------------------------------- -# Microcompact (stale tool result compaction) -# --------------------------------------------------------------------------- - - -def _microcompact_messages(*, total: int, tool_name: str, content: str) -> list[dict]: - messages: list[dict] = [{"role": "system", "content": "sys"}] - for i in range(total): - messages.append({ - "role": "assistant", - "content": "", - "tool_calls": [{ - "id": f"c{i}", - "type": "function", - "function": {"name": tool_name, "arguments": "{}"}, - }], - }) - messages.append({ - "role": "tool", - "tool_call_id": f"c{i}", - "name": tool_name, - "content": content, - }) - return messages - - -def test_microcompact_skips_when_prompt_under_hard_budget(monkeypatch): - """Cache-friendly path: in-flight tool results stay stable while prompt fits.""" - provider = MagicMock() - provider.generation = SimpleNamespace(max_tokens=0) - tools = MagicMock() - tools.get_definitions.return_value = [] - - total = 15 - long_content = "x" * 600 - messages = _microcompact_messages(total=total, tool_name="read_file", content=long_content) - spec = make_run_spec(provider, - initial_messages=messages, - tools=tools, - model="test-model", - max_iterations=1, - max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, - max_tokens=0, - context_window_tokens=20_000, - ) - - monkeypatch.setattr( - "nanobot.agent.context_governance.estimate_prompt_tokens_chain", - lambda *_args, **_kwargs: (1000, "test"), - ) - - result = ContextGovernor().compact_inflight_overflow( - _governance_config(provider, tools, spec), - messages, - set(), - ) - - assert result is messages - - -def test_microcompact_overflow_compacts_to_low_watermark(monkeypatch): - """Overflow path: compact in-flight stale results with headroom for later calls.""" - provider = MagicMock() - provider.generation = SimpleNamespace(max_tokens=0) - tools = MagicMock() - tools.get_definitions.return_value = [] - - total = 18 - long_content = "x" * 600 - messages = _microcompact_messages(total=total, tool_name="read_file", content=long_content) - spec = make_run_spec(provider, - initial_messages=messages, - tools=tools, - model="test-model", - max_iterations=1, - max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, - max_tokens=0, - context_window_tokens=2224, # input budget 1200, low target 1020 - ) - - def estimate(_provider, _model, msgs, _tools): - return sum( - 100 if (content := msg.get("content")) == long_content - else 1 if isinstance(content, str) and "compacted to fit context" in content - else 0 - for msg in msgs - if msg.get("role") == "tool" - ), "test" - - monkeypatch.setattr("nanobot.agent.context_governance.estimate_prompt_tokens_chain", estimate) - - result = ContextGovernor().compact_inflight_overflow( - _governance_config(provider, tools, spec), - messages, - set(), - ) - tool_msgs = [m for m in result if m.get("role") == "tool"] - compacted = [m for m in tool_msgs if "compacted to fit context" in str(m.get("content", ""))] - preserved = [m for m in tool_msgs if m.get("content") == long_content] - - assert len(compacted) == 8 - assert len(preserved) == total - 8 - assert [m["tool_call_id"] for m in compacted] == [f"c{i}" for i in range(8)] - - -def test_microcompact_compacts_newest_when_it_alone_overflows(monkeypatch): - """An unfit newest result tells the model to retry narrowly or report the limit.""" - provider = MagicMock() - provider.generation = SimpleNamespace(max_tokens=0) - tools = MagicMock() - tools.get_definitions.return_value = [] - - long_content = "x" * 600 - messages = _microcompact_messages(total=1, tool_name="read_file", content=long_content) - spec = make_run_spec(provider, - initial_messages=messages, - tools=tools, - model="test-model", - max_iterations=1, - max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, - max_tokens=0, - context_window_tokens=2000, - context_block_limit=500, - ) - - def estimate(_provider, _model, msgs, _tools): - return sum( - 1000 if msg.get("content") == long_content else 1 - for msg in msgs - if msg.get("role") == "tool" - ), "test" - - monkeypatch.setattr("nanobot.agent.context_governance.estimate_prompt_tokens_chain", estimate) - - compacted_tool_call_ids: set[str] = set() - result = ContextGovernor().compact_inflight_overflow( - _governance_config(provider, tools, spec), - messages, - compacted_tool_call_ids, - ) - - tool_msg = next(m for m in result if m.get("role") == "tool") - assert "compacted to fit context" in tool_msg["content"] - assert "Do not repeat the same call unchanged" in tool_msg["content"] - assert "Retry with a narrower path, query, range, or result limit" in tool_msg["content"] - assert "tell the user the task cannot fit" in tool_msg["content"] - assert compacted_tool_call_ids == {"c0"} - - -def test_context_governor_keeps_compaction_boundary_stable(monkeypatch): - provider = MagicMock() - provider.generation = SimpleNamespace(max_tokens=0) - tools = MagicMock() - tools.get_definitions.return_value = [] - - total = 18 - long_content = "x" * 600 - messages = _microcompact_messages(total=total, tool_name="read_file", content=long_content) - spec = make_run_spec(provider, - initial_messages=messages, - tools=tools, - model="test-model", - max_iterations=1, - max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, - max_tokens=0, - context_window_tokens=2224, - ) - - def estimate(_provider, _model, msgs, _tools): - return sum( - 100 if msg.get("content") == long_content else 1 - for msg in msgs - if msg.get("role") == "tool" - ), "test" - - monkeypatch.setattr("nanobot.agent.context_governance.estimate_prompt_tokens_chain", estimate) - - governor = ContextGovernor() - compacted_tool_call_ids: set[str] = set() - config = _governance_config(provider, tools, spec, inflight_start_index=0) - first = governor.compact_inflight_overflow(config, messages, compacted_tool_call_ids) - first_ids = set(compacted_tool_call_ids) - - second = governor.compact_inflight_overflow(config, messages, compacted_tool_call_ids) - - assert compacted_tool_call_ids == first_ids - assert [m.get("content") for m in second] == [m.get("content") for m in first] - - -def test_microcompact_preserves_short_results(monkeypatch): - """Short tool results below the compaction threshold should not be replaced.""" - provider = MagicMock() - provider.generation = SimpleNamespace(max_tokens=0) - tools = MagicMock() - tools.get_definitions.return_value = [] - - total = 15 - messages = _microcompact_messages(total=total, tool_name="exec", content="short") - spec = make_run_spec(provider, - initial_messages=messages, - tools=tools, - model="test-model", - max_iterations=1, - max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, - max_tokens=0, - context_window_tokens=2024, - ) - - monkeypatch.setattr( - "nanobot.agent.context_governance.estimate_prompt_tokens_chain", - lambda *_args, **_kwargs: (2000, "test"), - ) - - result = ContextGovernor().compact_inflight_overflow( - _governance_config(provider, tools, spec), - messages, - set(), - ) - assert result is messages # no copy needed — all stale results are short - - -def test_microcompact_skips_non_compactable_tools(monkeypatch): - """Non-compactable tools (e.g. 'message') should never be replaced.""" - provider = MagicMock() - provider.generation = SimpleNamespace(max_tokens=0) - tools = MagicMock() - tools.get_definitions.return_value = [] - - total = 15 - long_content = "y" * 1000 - messages = _microcompact_messages(total=total, tool_name="message", content=long_content) - spec = make_run_spec(provider, - initial_messages=messages, - tools=tools, - model="test-model", - max_iterations=1, - max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, - max_tokens=0, - context_window_tokens=2024, - ) - - monkeypatch.setattr( - "nanobot.agent.context_governance.estimate_prompt_tokens_chain", - lambda *_args, **_kwargs: (2000, "test"), - ) - - result = ContextGovernor().compact_inflight_overflow( - _governance_config(provider, tools, spec), - messages, - set(), - ) - assert result is messages # no compactable tools found - - def test_governance_repairs_orphans_after_snip(): """After snipping clips an assistant+tool_calls, orphan repair cleans up the tail.""" # Simulate snipping that keeps only the tail: drop the assistant with @@ -818,7 +1073,11 @@ def test_snip_history_preserves_user_message_after_truncation(monkeypatch): lambda msg: token_sizes.get(str(msg.get("content")), 100), ) - trimmed = ContextGovernor().snip_history(_governance_config(provider, tools, spec), messages) + trimmed = ContextGovernor().snip_history( + _governance_config(provider, tools, spec), + messages, + tool_definitions=tools.get_definitions(), + ) # The first non-system message MUST be user (not assistant). non_system = [m for m in trimmed if m.get("role") != "system"] @@ -863,7 +1122,11 @@ def test_snip_history_no_user_at_all_falls_back_gracefully(monkeypatch): lambda msg: 100, ) - trimmed = ContextGovernor().snip_history(_governance_config(provider, tools, spec), messages) + trimmed = ContextGovernor().snip_history( + _governance_config(provider, tools, spec), + messages, + tool_definitions=tools.get_definitions(), + ) # Should not crash. The result should still be a valid list. assert isinstance(trimmed, list) @@ -871,7 +1134,6 @@ def test_snip_history_no_user_at_all_falls_back_gracefully(monkeypatch): assert any(m.get("role") == "system" for m in trimmed) # The _enforce_role_alternation safety net must be able to fix whatever # _snip_history returns here — verify it produces a valid sequence. - from nanobot.providers.base import LLMProvider fixed = LLMProvider._enforce_role_alternation(trimmed) non_system = [m for m in fixed if m["role"] != "system"] if non_system: diff --git a/tests/agent/test_session_model_runtime.py b/tests/agent/test_session_model_runtime.py index dac43632a..dd1dee5a6 100644 --- a/tests/agent/test_session_model_runtime.py +++ b/tests/agent/test_session_model_runtime.py @@ -114,7 +114,7 @@ async def test_removed_session_model_preset_falls_back_and_clears_metadata(tmp_p provider=base, workspace=tmp_path, model="base-model", - context_window_tokens=8_000, + context_window_tokens=16_000, ) loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign] session_key = "sdk:removed-preset" @@ -196,7 +196,7 @@ async def test_sdk_custom_model_preset_metadata_does_not_select_runtime( provider=base, workspace=tmp_path, model="base-model", - context_window_tokens=8_000, + context_window_tokens=16_000, ) loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign] bot = Nanobot(loop) diff --git a/tests/providers/test_conversation_state.py b/tests/providers/test_conversation_state.py index 7c62159e8..157acb14d 100644 --- a/tests/providers/test_conversation_state.py +++ b/tests/providers/test_conversation_state.py @@ -146,6 +146,51 @@ def test_controller_uses_governed_messages_for_provider_state_delta() -> None: assert governed_checkpoint.pending_messages[-1]["content"] == "compacted result" +def test_controller_estimates_active_state_plus_pending_delta(monkeypatch) -> None: + provider = _provider() + current_message = {"role": "user", "content": "new delta"} + state = ProviderConversationState( + kind="openai_responses", + provider="openai:test", + model="gpt-5.6", + version=1, + payload={ + "items": [{"type": "reasoning", "encrypted_content": "opaque"}], + "context_tokens": 450, + }, + pending_messages=[current_message], + ) + controller = ProviderConversationStateController( + provider=provider, + model="gpt-5.6", + messages=[current_message], + state=state, + ) + seen = {} + + def estimate(_provider, _model, messages, tools): + seen["messages"] = messages + seen["tools"] = tools + return 100, "test-counter" + + monkeypatch.setattr( + "nanobot.providers.conversation_state.estimate_prompt_tokens_chain", + estimate, + ) + + tokens = controller.estimate_request_context_tokens( + [current_message], + model_messages=[current_message], + tool_definitions=[{"type": "web_search"}], + ) + + assert tokens == 550 + assert seen == { + "messages": [current_message], + "tools": [{"type": "web_search"}], + } + + def test_transient_response_preserves_only_durable_request_messages() -> None: provider = _provider() current_message = {"role": "user", "content": "continue"}