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
This commit is contained in:
chengyongru
2026-08-31 18:08:21 +08:00
committed by GitHub
parent 6d6d58d329
commit e111b83af6
9 changed files with 888 additions and 491 deletions
+101 -134
View File
@@ -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])
+136 -58
View File
@@ -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 "",
+42 -1
View File
@@ -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
+3 -3
View File
@@ -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],
)
@@ -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
+30 -29
View File
@@ -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
+526 -264
View File
@@ -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:
+2 -2
View File
@@ -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)
@@ -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"}