mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-31 16:21:50 +03:00
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:
+101
-134
@@ -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
@@ -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 "",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"}
|
||||
|
||||
Reference in New Issue
Block a user