mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-09-01 08:42:20 +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 loguru import logger
|
||||||
|
|
||||||
|
from nanobot.providers.base import LLMUsage
|
||||||
from nanobot.utils.helpers import (
|
from nanobot.utils.helpers import (
|
||||||
estimate_message_tokens,
|
estimate_message_tokens,
|
||||||
estimate_prompt_tokens_chain,
|
estimate_prompt_tokens_chain,
|
||||||
@@ -27,12 +28,6 @@ if TYPE_CHECKING:
|
|||||||
from nanobot.providers.base import LLMProvider
|
from nanobot.providers.base import LLMProvider
|
||||||
|
|
||||||
SNIP_SAFETY_BUFFER = 1024
|
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.
|
# read_file is the recovery path for persisted results; exempting it prevents persist->read->persist loops.
|
||||||
TOOL_RESULT_OFFLOAD_EXEMPT_TOOLS = frozenset({"read_file"})
|
TOOL_RESULT_OFFLOAD_EXEMPT_TOOLS = frozenset({"read_file"})
|
||||||
BACKFILL_CONTENT = "[Tool result unavailable — call was interrupted or lost]"
|
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:
|
def _tool_call_name_is_valid(tool_call: Any) -> bool:
|
||||||
"""Whether a persisted OpenAI-style tool_call carries a usable name.
|
"""Whether a persisted OpenAI-style tool_call carries a usable name.
|
||||||
|
|
||||||
@@ -67,7 +83,6 @@ class ContextGovernanceConfig:
|
|||||||
context_window_tokens: int | None = None
|
context_window_tokens: int | None = None
|
||||||
context_block_limit: int | None = None
|
context_block_limit: int | None = None
|
||||||
max_tokens: int | None = None
|
max_tokens: int | None = None
|
||||||
inflight_start_index: int = 0
|
|
||||||
|
|
||||||
|
|
||||||
class ContextGovernor:
|
class ContextGovernor:
|
||||||
@@ -77,17 +92,85 @@ class ContextGovernor:
|
|||||||
self,
|
self,
|
||||||
config: ContextGovernanceConfig,
|
config: ContextGovernanceConfig,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
compacted_tool_call_ids: set[str],
|
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
updated = self.strip_placeholder_assistant_messages(messages)
|
updated = self.strip_placeholder_assistant_messages(messages)
|
||||||
updated = self.strip_malformed_tool_calls(updated)
|
updated = self.strip_malformed_tool_calls(updated)
|
||||||
updated = self.drop_orphan_tool_results(updated)
|
updated = self.drop_orphan_tool_results(updated)
|
||||||
updated = self.backfill_missing_tool_results(updated)
|
updated = self.backfill_missing_tool_results(updated)
|
||||||
updated = self.apply_tool_result_budget(config, updated)
|
return self.apply_tool_result_budget(config, updated)
|
||||||
updated = self.compact_inflight_overflow(config, updated, compacted_tool_call_ids)
|
|
||||||
updated = self.snip_history(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)
|
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
|
@staticmethod
|
||||||
def input_budget(config: ContextGovernanceConfig) -> int:
|
def input_budget(config: ContextGovernanceConfig) -> int:
|
||||||
@@ -326,71 +409,13 @@ class ContextGovernor:
|
|||||||
updated[idx]["content"] = normalized
|
updated[idx]["content"] = normalized
|
||||||
return updated
|
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(
|
def snip_history(
|
||||||
self,
|
self,
|
||||||
config: ContextGovernanceConfig,
|
config: ContextGovernanceConfig,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
|
*,
|
||||||
|
tool_definitions: list[dict[str, Any]] | None,
|
||||||
|
force: bool = False,
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
if not messages or not config.context_window_tokens:
|
if not messages or not config.context_window_tokens:
|
||||||
return messages
|
return messages
|
||||||
@@ -399,14 +424,13 @@ class ContextGovernor:
|
|||||||
if budget <= 0:
|
if budget <= 0:
|
||||||
return messages
|
return messages
|
||||||
|
|
||||||
tools = config.tools.get_definitions()
|
|
||||||
estimate, _ = estimate_prompt_tokens_chain(
|
estimate, _ = estimate_prompt_tokens_chain(
|
||||||
config.provider,
|
config.provider,
|
||||||
config.model,
|
config.model,
|
||||||
messages,
|
messages,
|
||||||
tools,
|
tool_definitions,
|
||||||
)
|
)
|
||||||
if estimate <= budget:
|
if not force and estimate <= budget:
|
||||||
return messages
|
return messages
|
||||||
|
|
||||||
system_messages = [dict(msg) for msg in messages if msg.get("role") == "system"]
|
system_messages = [dict(msg) for msg in messages if msg.get("role") == "system"]
|
||||||
@@ -419,7 +443,7 @@ class ContextGovernor:
|
|||||||
config.provider,
|
config.provider,
|
||||||
config.model,
|
config.model,
|
||||||
system_messages,
|
system_messages,
|
||||||
tools,
|
tool_definitions,
|
||||||
)
|
)
|
||||||
remaining_budget = max(0, budget - max(system_tokens, fixed_tokens))
|
remaining_budget = max(0, budget - max(system_tokens, fixed_tokens))
|
||||||
kept: list[dict[str, Any]] = []
|
kept: list[dict[str, Any]] = []
|
||||||
@@ -434,16 +458,6 @@ class ContextGovernor:
|
|||||||
|
|
||||||
return system_messages + self._legal_history_tail(kept, non_system)
|
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(
|
def _legal_history_tail(
|
||||||
self,
|
self,
|
||||||
kept: list[dict[str, Any]],
|
kept: list[dict[str, Any]],
|
||||||
@@ -462,50 +476,3 @@ class ContextGovernor:
|
|||||||
if messages[idx].get("role") == "user":
|
if messages[idx].get("role") == "user":
|
||||||
return messages[idx:]
|
return messages[idx:]
|
||||||
return []
|
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)
|
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:
|
class AgentRunner:
|
||||||
"""Run a tool-capable LLM loop without product-layer concerns."""
|
"""Run a tool-capable LLM loop without product-layer concerns."""
|
||||||
|
|
||||||
@@ -500,7 +511,6 @@ class AgentRunner:
|
|||||||
length_recovery_parts: list[str] = []
|
length_recovery_parts: list[str] = []
|
||||||
had_injections = False
|
had_injections = False
|
||||||
injection_cycles = 0
|
injection_cycles = 0
|
||||||
compacted_tool_call_ids: set[str] = set()
|
|
||||||
pending_stream_content: str | None = None
|
pending_stream_content: str | None = None
|
||||||
conversation_state = ProviderConversationStateController(
|
conversation_state = ProviderConversationStateController(
|
||||||
provider=spec.runtime.provider,
|
provider=spec.runtime.provider,
|
||||||
@@ -519,39 +529,29 @@ class AgentRunner:
|
|||||||
context_window_tokens=spec.runtime.context_window_tokens,
|
context_window_tokens=spec.runtime.context_window_tokens,
|
||||||
context_block_limit=spec.context_block_limit,
|
context_block_limit=spec.context_block_limit,
|
||||||
max_tokens=spec.runtime.generation.max_tokens,
|
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):
|
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(
|
context = AgentHookContext(
|
||||||
iteration=iteration,
|
iteration=iteration,
|
||||||
messages=messages,
|
messages=messages,
|
||||||
session_key=spec.session_key,
|
session_key=spec.session_key,
|
||||||
)
|
)
|
||||||
await hook.before_iteration(context)
|
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(
|
response = await self._request_model(
|
||||||
spec,
|
spec,
|
||||||
messages_for_model,
|
messages,
|
||||||
hook,
|
hook,
|
||||||
context,
|
context,
|
||||||
conversation_state=conversation_state,
|
request_state=request_state,
|
||||||
provider_context=provider_context,
|
transcript=messages,
|
||||||
)
|
)
|
||||||
|
assert request_state.messages is not None
|
||||||
|
messages_for_model = request_state.messages
|
||||||
conversation_state.observe_response(response, messages)
|
conversation_state.observe_response(response, messages)
|
||||||
context.response = response
|
context.response = response
|
||||||
context.tool_calls = list(response.tool_calls)
|
context.tool_calls = list(response.tool_calls)
|
||||||
@@ -563,7 +563,7 @@ class AgentRunner:
|
|||||||
response.content,
|
response.content,
|
||||||
)
|
)
|
||||||
response.content = cleaned_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
|
context.usage = raw_usage
|
||||||
usage = self._merge_usage(usage, raw_usage)
|
usage = self._merge_usage(usage, raw_usage)
|
||||||
if reasoning_text and not context.streamed_reasoning:
|
if reasoning_text and not context.streamed_reasoning:
|
||||||
@@ -637,7 +637,6 @@ class AgentRunner:
|
|||||||
self.context_governor.prepare_for_model(
|
self.context_governor.prepare_for_model(
|
||||||
governance_config,
|
governance_config,
|
||||||
messages,
|
messages,
|
||||||
compacted_tool_call_ids,
|
|
||||||
)
|
)
|
||||||
if response.provider_state is not None
|
if response.provider_state is not None
|
||||||
else None
|
else None
|
||||||
@@ -703,14 +702,13 @@ class AgentRunner:
|
|||||||
)
|
)
|
||||||
if hook.wants_streaming():
|
if hook.wants_streaming():
|
||||||
await hook.on_stream_end(context, resuming=False)
|
await hook.on_stream_end(context, resuming=False)
|
||||||
retry_messages = self._finalization_retry_messages(messages_for_model)
|
|
||||||
response = await self._request_finalization_retry(
|
response = await self._request_finalization_retry(
|
||||||
spec,
|
spec,
|
||||||
messages_for_model,
|
messages_for_model,
|
||||||
|
request_state=request_state,
|
||||||
transcript=messages,
|
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)
|
usage = self._merge_usage(usage, retry_usage)
|
||||||
raw_usage = self._merge_usage(raw_usage, retry_usage)
|
raw_usage = self._merge_usage(raw_usage, retry_usage)
|
||||||
context.response = response
|
context.response = response
|
||||||
@@ -897,7 +895,7 @@ class AgentRunner:
|
|||||||
hook,
|
hook,
|
||||||
messages,
|
messages,
|
||||||
usage,
|
usage,
|
||||||
conversation_state,
|
request_state=request_state,
|
||||||
)
|
)
|
||||||
if terminal_content is None:
|
if terminal_content is None:
|
||||||
terminal_content = self._max_iterations_fallback(spec)
|
terminal_content = self._max_iterations_fallback(spec)
|
||||||
@@ -944,6 +942,60 @@ class AgentRunner:
|
|||||||
kwargs["reasoning_effort"] = generation.reasoning_effort
|
kwargs["reasoning_effort"] = generation.reasoning_effort
|
||||||
return kwargs
|
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(
|
async def _request_model(
|
||||||
self,
|
self,
|
||||||
spec: AgentRunSpec,
|
spec: AgentRunSpec,
|
||||||
@@ -951,16 +1003,23 @@ class AgentRunner:
|
|||||||
hook: AgentHook,
|
hook: AgentHook,
|
||||||
context: AgentHookContext,
|
context: AgentHookContext,
|
||||||
*,
|
*,
|
||||||
|
request_state: _ModelRequestState,
|
||||||
malformed_retry: bool = False,
|
malformed_retry: bool = False,
|
||||||
conversation_state: ProviderConversationStateController,
|
transcript: list[dict[str, Any]] | None,
|
||||||
provider_context: ProviderCallContext | None = None,
|
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
timeout_s = self._resolve_llm_timeout_s(spec)
|
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(
|
kwargs = self._build_request_kwargs(
|
||||||
spec,
|
spec,
|
||||||
messages,
|
messages,
|
||||||
tools=spec.tools.get_definitions(),
|
tools=tool_definitions,
|
||||||
)
|
)
|
||||||
wants_streaming = hook.wants_streaming()
|
wants_streaming = hook.wants_streaming()
|
||||||
|
|
||||||
@@ -1138,11 +1197,9 @@ class AgentRunner:
|
|||||||
)
|
)
|
||||||
return await self._request_model(
|
return await self._request_model(
|
||||||
spec, retry_messages, hook, context,
|
spec, retry_messages, hook, context,
|
||||||
|
request_state=request_state,
|
||||||
malformed_retry=True,
|
malformed_retry=True,
|
||||||
conversation_state=conversation_state,
|
transcript=None,
|
||||||
provider_context=conversation_state.independent_request_context(
|
|
||||||
context_window_tokens=spec.runtime.context_window_tokens,
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
all_dropped
|
all_dropped
|
||||||
@@ -1158,9 +1215,7 @@ class AgentRunner:
|
|||||||
return await self._request_no_tools(
|
return await self._request_no_tools(
|
||||||
spec,
|
spec,
|
||||||
fallback_messages,
|
fallback_messages,
|
||||||
provider_context=conversation_state.independent_request_context(
|
request_state=request_state,
|
||||||
context_window_tokens=spec.runtime.context_window_tokens,
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
return response
|
return response
|
||||||
|
|
||||||
@@ -1228,21 +1283,17 @@ class AgentRunner:
|
|||||||
spec: AgentRunSpec,
|
spec: AgentRunSpec,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
*,
|
*,
|
||||||
|
request_state: _ModelRequestState,
|
||||||
transcript: list[dict[str, Any]],
|
transcript: list[dict[str, Any]],
|
||||||
conversation_state: ProviderConversationStateController,
|
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
retry_messages = self._finalization_retry_messages(messages)
|
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(
|
response = await self._request_no_tools(
|
||||||
spec,
|
spec,
|
||||||
retry_messages,
|
retry_messages,
|
||||||
provider_context=provider_context,
|
request_state=request_state,
|
||||||
|
transcript=transcript,
|
||||||
)
|
)
|
||||||
conversation_state.observe_response(
|
request_state.conversation.observe_response(
|
||||||
response,
|
response,
|
||||||
transcript,
|
transcript,
|
||||||
adopt_candidate_state=False,
|
adopt_candidate_state=False,
|
||||||
@@ -1261,16 +1312,15 @@ class AgentRunner:
|
|||||||
hook: AgentHook,
|
hook: AgentHook,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
usage: LLMUsage | None,
|
usage: LLMUsage | None,
|
||||||
conversation_state: ProviderConversationStateController,
|
*,
|
||||||
|
request_state: _ModelRequestState,
|
||||||
) -> tuple[str | None, LLMUsage | None]:
|
) -> tuple[str | None, LLMUsage | None]:
|
||||||
retry_messages = self._budget_exhausted_finalization_messages(messages)
|
retry_messages = self._budget_exhausted_finalization_messages(messages)
|
||||||
try:
|
try:
|
||||||
response = await self._request_no_tools(
|
response = await self._request_no_tools(
|
||||||
spec,
|
spec,
|
||||||
retry_messages,
|
retry_messages,
|
||||||
provider_context=conversation_state.independent_request_context(
|
request_state=request_state,
|
||||||
context_window_tokens=spec.runtime.context_window_tokens,
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception(
|
logger.exception(
|
||||||
@@ -1279,7 +1329,7 @@ class AgentRunner:
|
|||||||
)
|
)
|
||||||
return None, usage
|
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)
|
usage = self._merge_usage(usage, raw_usage)
|
||||||
if response.finish_reason == "error" or response.has_tool_calls:
|
if response.finish_reason == "error" or response.has_tool_calls:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@@ -1308,8 +1358,15 @@ class AgentRunner:
|
|||||||
spec: AgentRunSpec,
|
spec: AgentRunSpec,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
*,
|
*,
|
||||||
provider_context: ProviderCallContext | None = None,
|
request_state: _ModelRequestState,
|
||||||
|
transcript: list[dict[str, Any]] | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
|
messages, provider_context = self._prepare_model_request(
|
||||||
|
request_state,
|
||||||
|
messages,
|
||||||
|
tool_definitions=None,
|
||||||
|
transcript=transcript,
|
||||||
|
)
|
||||||
kwargs = self._build_request_kwargs(
|
kwargs = self._build_request_kwargs(
|
||||||
spec,
|
spec,
|
||||||
messages,
|
messages,
|
||||||
@@ -1321,17 +1378,18 @@ class AgentRunner:
|
|||||||
)
|
)
|
||||||
timeout_s = self._resolve_llm_timeout_s(spec)
|
timeout_s = self._resolve_llm_timeout_s(spec)
|
||||||
try:
|
try:
|
||||||
return (
|
response = (
|
||||||
await coro
|
await coro
|
||||||
if timeout_s is None
|
if timeout_s is None
|
||||||
else await asyncio.wait_for(coro, timeout=timeout_s)
|
else await asyncio.wait_for(coro, timeout=timeout_s)
|
||||||
)
|
)
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
return LLMResponse(
|
response = LLMResponse(
|
||||||
content=f"Error calling LLM: timed out after {timeout_s:g}s",
|
content=f"Error calling LLM: timed out after {timeout_s:g}s",
|
||||||
finish_reason="error",
|
finish_reason="error",
|
||||||
error_kind="timeout",
|
error_kind="timeout",
|
||||||
)
|
)
|
||||||
|
return response
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _resolve_llm_timeout_s(spec: AgentRunSpec) -> float | None:
|
def _resolve_llm_timeout_s(spec: AgentRunSpec) -> float | None:
|
||||||
@@ -1373,33 +1431,53 @@ class AgentRunner:
|
|||||||
spec: AgentRunSpec,
|
spec: AgentRunSpec,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
response: LLMResponse,
|
response: LLMResponse,
|
||||||
|
*,
|
||||||
|
tool_definitions: list[dict[str, Any]] | None,
|
||||||
) -> LLMUsage | None:
|
) -> LLMUsage | None:
|
||||||
usage = response.usage
|
usage = response.usage
|
||||||
if response.finish_reason == "error":
|
if response.finish_reason == "error":
|
||||||
if usage is None or usage.total_tokens == 0:
|
if usage is None or usage.total_tokens == 0:
|
||||||
usage = LLMUsage.empty_request()
|
usage = LLMUsage.empty_request()
|
||||||
elif usage is None or usage.total_tokens == 0:
|
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(
|
return usage.with_timing(
|
||||||
generation_ms=response.generation_ms,
|
generation_ms=response.generation_ms,
|
||||||
ttft_ms=response.ttft_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(
|
def _estimate_response_usage(
|
||||||
self,
|
self,
|
||||||
spec: AgentRunSpec,
|
spec: AgentRunSpec,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
response: LLMResponse,
|
response: LLMResponse,
|
||||||
|
*,
|
||||||
|
tool_definitions: list[dict[str, Any]] | None,
|
||||||
) -> LLMUsage:
|
) -> LLMUsage:
|
||||||
try:
|
|
||||||
tools = spec.tools.get_definitions()
|
|
||||||
except Exception:
|
|
||||||
tools = None
|
|
||||||
prompt_tokens, _ = estimate_prompt_tokens_chain(
|
prompt_tokens, _ = estimate_prompt_tokens_chain(
|
||||||
spec.runtime.provider,
|
spec.runtime.provider,
|
||||||
spec.runtime.model,
|
spec.runtime.model,
|
||||||
messages,
|
messages,
|
||||||
tools,
|
tool_definitions,
|
||||||
)
|
)
|
||||||
assistant_message = build_assistant_message(
|
assistant_message = build_assistant_message(
|
||||||
response.content or "",
|
response.content or "",
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ from nanobot.providers.base import (
|
|||||||
ProviderCallContext,
|
ProviderCallContext,
|
||||||
ProviderConversationState,
|
ProviderConversationState,
|
||||||
)
|
)
|
||||||
|
from nanobot.utils.helpers import estimate_prompt_tokens_chain
|
||||||
|
|
||||||
_PROVIDER_STATE_OUTPUT_META = "provider_state_output"
|
_PROVIDER_STATE_OUTPUT_META = "provider_state_output"
|
||||||
_PROVIDER_STATE_BOUNDARY_META = "provider_state_boundary"
|
_PROVIDER_STATE_BOUNDARY_META = "provider_state_boundary"
|
||||||
@@ -69,6 +70,37 @@ class ProviderConversationStateController:
|
|||||||
session_id=self._session_id,
|
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(
|
def prepare_request(
|
||||||
self,
|
self,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
@@ -76,11 +108,20 @@ class ProviderConversationStateController:
|
|||||||
context_window_tokens: int | None,
|
context_window_tokens: int | None,
|
||||||
model_messages: list[dict[str, Any]] | None = None,
|
model_messages: list[dict[str, Any]] | None = None,
|
||||||
supplemental_messages: list[dict[str, Any]] | None = None,
|
supplemental_messages: list[dict[str, Any]] | None = None,
|
||||||
|
resume_state: bool = True,
|
||||||
) -> ProviderCallContext | None:
|
) -> 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(
|
independent_context = self.independent_request_context(
|
||||||
context_window_tokens=context_window_tokens,
|
context_window_tokens=context_window_tokens,
|
||||||
)
|
)
|
||||||
|
if not resume_state:
|
||||||
|
self._state = None
|
||||||
|
self._request_messages = []
|
||||||
|
return independent_context
|
||||||
if self._state is None:
|
if self._state is None:
|
||||||
self._request_messages = []
|
self._request_messages = []
|
||||||
return independent_context
|
return independent_context
|
||||||
|
|||||||
@@ -426,7 +426,7 @@ class TestEphemeralDirect:
|
|||||||
bus=bus,
|
bus=bus,
|
||||||
provider=provider,
|
provider=provider,
|
||||||
workspace=tmp_path,
|
workspace=tmp_path,
|
||||||
context_window_tokens=8000,
|
context_window_tokens=32_000,
|
||||||
)
|
)
|
||||||
|
|
||||||
return loop, store
|
return loop, store
|
||||||
@@ -606,7 +606,7 @@ class TestEphemeralDirect:
|
|||||||
bus=MessageBus(),
|
bus=MessageBus(),
|
||||||
provider=provider,
|
provider=provider,
|
||||||
workspace=tmp_path,
|
workspace=tmp_path,
|
||||||
context_window_tokens=8000,
|
context_window_tokens=32_000,
|
||||||
)
|
)
|
||||||
|
|
||||||
await loop.process_direct(
|
await loop.process_direct(
|
||||||
@@ -666,7 +666,7 @@ class TestEphemeralHooks:
|
|||||||
bus=bus,
|
bus=bus,
|
||||||
provider=provider,
|
provider=provider,
|
||||||
workspace=tmp_path,
|
workspace=tmp_path,
|
||||||
context_window_tokens=8000,
|
context_window_tokens=32_000,
|
||||||
hooks=[spy],
|
hooks=[spy],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -29,6 +29,9 @@ def _make_loop(
|
|||||||
workspace=tmp_path,
|
workspace=tmp_path,
|
||||||
model="test-model",
|
model="test-model",
|
||||||
context_window_tokens=context_window_tokens,
|
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.tools.get_definitions = MagicMock(return_value=[])
|
||||||
loop.consolidator._SAFETY_BUFFER = 0
|
loop.consolidator._SAFETY_BUFFER = 0
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import pytest
|
|||||||
|
|
||||||
from agent.runner_helpers import make_run_spec
|
from agent.runner_helpers import make_run_spec
|
||||||
from nanobot.agent.context import TranscriptInput
|
from nanobot.agent.context import TranscriptInput
|
||||||
|
from nanobot.agent.context_governance import ContextWindowExceededError
|
||||||
from nanobot.config.schema import AgentDefaults
|
from nanobot.config.schema import AgentDefaults
|
||||||
from nanobot.providers.base import (
|
from nanobot.providers.base import (
|
||||||
LLMProvider,
|
LLMProvider,
|
||||||
@@ -86,6 +87,7 @@ def test_usage_or_estimate_replaces_reported_zero_for_content(monkeypatch) -> No
|
|||||||
_make_usage_spec(provider, tools),
|
_make_usage_spec(provider, tools),
|
||||||
[{"role": "user", "content": "hello"}],
|
[{"role": "user", "content": "hello"}],
|
||||||
response,
|
response,
|
||||||
|
tool_definitions=tools.get_definitions(),
|
||||||
)
|
)
|
||||||
|
|
||||||
assert usage == LLMUsage.estimated(input_tokens=12, output_tokens=7).with_timing(
|
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),
|
_make_usage_spec(provider, tools),
|
||||||
[{"role": "user", "content": "hello"}],
|
[{"role": "user", "content": "hello"}],
|
||||||
response,
|
response,
|
||||||
|
tool_definitions=tools.get_definitions(),
|
||||||
)
|
)
|
||||||
|
|
||||||
assert usage == LLMUsage.estimated(input_tokens=13, output_tokens=9)
|
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),
|
_make_usage_spec(provider, tools),
|
||||||
[{"role": "user", "content": "hello"}],
|
[{"role": "user", "content": "hello"}],
|
||||||
response,
|
response,
|
||||||
|
tool_definitions=tools.get_definitions(),
|
||||||
)
|
)
|
||||||
|
|
||||||
assert usage is not None
|
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),
|
_make_usage_spec(provider, tools),
|
||||||
[{"role": "user", "content": "hello"}],
|
[{"role": "user", "content": "hello"}],
|
||||||
response,
|
response,
|
||||||
|
tool_definitions=tools.get_definitions(),
|
||||||
)
|
)
|
||||||
|
|
||||||
assert usage is not None
|
assert usage is not None
|
||||||
@@ -366,14 +371,12 @@ async def test_runner_replays_provider_state_without_chat_projection_duplicates(
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@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
|
from nanobot.agent.runner import AgentRunner
|
||||||
|
|
||||||
provider = MagicMock(spec=LLMProvider)
|
provider = MagicMock(spec=LLMProvider)
|
||||||
provider.can_resume_conversation_state.return_value = True
|
provider.can_resume_conversation_state.return_value = True
|
||||||
provider.supports_native_compaction.return_value = False
|
|
||||||
calls = 0
|
calls = 0
|
||||||
captured_context: ProviderCallContext | None = None
|
|
||||||
checkpoints: list[dict] = []
|
checkpoints: list[dict] = []
|
||||||
state = ProviderConversationState(
|
state = ProviderConversationState(
|
||||||
kind="openai_responses",
|
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):
|
async def chat_with_retry(**kwargs):
|
||||||
nonlocal calls, captured_context
|
nonlocal calls
|
||||||
calls += 1
|
calls += 1
|
||||||
if calls == 1:
|
if calls == 1:
|
||||||
return LLMResponse(
|
return LLMResponse(
|
||||||
@@ -398,7 +401,6 @@ async def test_runner_governs_tool_result_before_adding_it_to_provider_state():
|
|||||||
],
|
],
|
||||||
provider_state=state,
|
provider_state=state,
|
||||||
)
|
)
|
||||||
captured_context = kwargs["provider_context"]
|
|
||||||
return LLMResponse(content="done")
|
return LLMResponse(content="done")
|
||||||
|
|
||||||
provider.chat_with_retry = chat_with_retry
|
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:
|
async def checkpoint(payload: dict) -> None:
|
||||||
checkpoints.append(payload)
|
checkpoints.append(payload)
|
||||||
|
|
||||||
await AgentRunner().run(make_run_spec(
|
with pytest.raises(ContextWindowExceededError):
|
||||||
provider,
|
await AgentRunner().run(make_run_spec(
|
||||||
initial_messages=[
|
provider,
|
||||||
{"role": "system", "content": "system"},
|
initial_messages=[
|
||||||
{"role": "user", "content": "read the file"},
|
{"role": "system", "content": "system"},
|
||||||
],
|
{"role": "user", "content": "read the file"},
|
||||||
tools=tools,
|
],
|
||||||
model="gpt-5.6",
|
tools=tools,
|
||||||
context_window_tokens=3_000,
|
model="gpt-5.6",
|
||||||
context_block_limit=200,
|
context_window_tokens=3_000,
|
||||||
max_tokens=1_000,
|
context_block_limit=200,
|
||||||
max_iterations=3,
|
max_tokens=1_000,
|
||||||
max_tool_result_chars=10_000,
|
max_iterations=3,
|
||||||
checkpoint_callback=checkpoint,
|
max_tool_result_chars=10_000,
|
||||||
))
|
checkpoint_callback=checkpoint,
|
||||||
|
))
|
||||||
|
|
||||||
assert captured_context is not None
|
assert calls == 1
|
||||||
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
|
|
||||||
completed_checkpoint = next(
|
completed_checkpoint = next(
|
||||||
checkpoint
|
checkpoint
|
||||||
for checkpoint in checkpoints
|
for checkpoint in checkpoints
|
||||||
if checkpoint["phase"] == "tools_completed"
|
if checkpoint["phase"] == "tools_completed"
|
||||||
)
|
)
|
||||||
checkpoint_pending = completed_checkpoint["provider_state"].pending_messages
|
checkpoint_pending = completed_checkpoint["provider_state"].pending_messages
|
||||||
assert "compacted to fit context" in checkpoint_pending[0]["content"]
|
assert checkpoint_pending == [{
|
||||||
assert checkpoint_pending[0]["content"] != "x" * 5_000
|
"role": "tool",
|
||||||
|
"tool_call_id": "call_1",
|
||||||
|
"name": "read_file",
|
||||||
|
"content": "x" * 5_000,
|
||||||
|
}]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@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 __future__ import annotations
|
||||||
|
|
||||||
from types import SimpleNamespace
|
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -12,11 +11,14 @@ from nanobot.agent.context_governance import (
|
|||||||
BACKFILL_CONTENT,
|
BACKFILL_CONTENT,
|
||||||
ContextGovernanceConfig,
|
ContextGovernanceConfig,
|
||||||
ContextGovernor,
|
ContextGovernor,
|
||||||
|
ContextWindowExceededError,
|
||||||
)
|
)
|
||||||
from nanobot.agent.runner import AgentRunSpec
|
from nanobot.agent.runner import AgentRunSpec
|
||||||
from nanobot.config.schema import AgentDefaults
|
from nanobot.config.schema import AgentDefaults
|
||||||
from nanobot.providers.base import (
|
from nanobot.providers.base import (
|
||||||
|
LLMProvider,
|
||||||
LLMResponse,
|
LLMResponse,
|
||||||
|
LLMUsage,
|
||||||
ProviderConversationState,
|
ProviderConversationState,
|
||||||
ToolCallRequest,
|
ToolCallRequest,
|
||||||
)
|
)
|
||||||
@@ -28,8 +30,6 @@ def _governance_config(
|
|||||||
provider,
|
provider,
|
||||||
tools,
|
tools,
|
||||||
spec: AgentRunSpec,
|
spec: AgentRunSpec,
|
||||||
*,
|
|
||||||
inflight_start_index: int = 0,
|
|
||||||
) -> ContextGovernanceConfig:
|
) -> ContextGovernanceConfig:
|
||||||
return ContextGovernanceConfig(
|
return ContextGovernanceConfig(
|
||||||
provider=provider,
|
provider=provider,
|
||||||
@@ -41,7 +41,6 @@ def _governance_config(
|
|||||||
context_window_tokens=spec.runtime.context_window_tokens,
|
context_window_tokens=spec.runtime.context_window_tokens,
|
||||||
context_block_limit=spec.context_block_limit,
|
context_block_limit=spec.context_block_limit,
|
||||||
max_tokens=spec.runtime.generation.max_tokens,
|
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()
|
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):
|
def test_snip_history_drops_orphaned_tool_results_from_trimmed_slice(monkeypatch):
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
tools = 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),
|
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
|
# After the fix, the user message is recovered so the sequence is valid
|
||||||
# for providers that require system → user (e.g. GLM error 1214).
|
# 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),
|
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]
|
contents = [message.get("content") for message in trimmed]
|
||||||
assert contents == ["system", "recent two"]
|
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():
|
def test_governance_repairs_orphans_after_snip():
|
||||||
"""After snipping clips an assistant+tool_calls, orphan repair cleans up the tail."""
|
"""After snipping clips an assistant+tool_calls, orphan repair cleans up the tail."""
|
||||||
# Simulate snipping that keeps only the tail: drop the assistant with
|
# 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),
|
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).
|
# The first non-system message MUST be user (not assistant).
|
||||||
non_system = [m for m in trimmed if m.get("role") != "system"]
|
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,
|
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.
|
# Should not crash. The result should still be a valid list.
|
||||||
assert isinstance(trimmed, 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)
|
assert any(m.get("role") == "system" for m in trimmed)
|
||||||
# The _enforce_role_alternation safety net must be able to fix whatever
|
# The _enforce_role_alternation safety net must be able to fix whatever
|
||||||
# _snip_history returns here — verify it produces a valid sequence.
|
# _snip_history returns here — verify it produces a valid sequence.
|
||||||
from nanobot.providers.base import LLMProvider
|
|
||||||
fixed = LLMProvider._enforce_role_alternation(trimmed)
|
fixed = LLMProvider._enforce_role_alternation(trimmed)
|
||||||
non_system = [m for m in fixed if m["role"] != "system"]
|
non_system = [m for m in fixed if m["role"] != "system"]
|
||||||
if non_system:
|
if non_system:
|
||||||
|
|||||||
@@ -114,7 +114,7 @@ async def test_removed_session_model_preset_falls_back_and_clears_metadata(tmp_p
|
|||||||
provider=base,
|
provider=base,
|
||||||
workspace=tmp_path,
|
workspace=tmp_path,
|
||||||
model="base-model",
|
model="base-model",
|
||||||
context_window_tokens=8_000,
|
context_window_tokens=16_000,
|
||||||
)
|
)
|
||||||
loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
|
loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
|
||||||
session_key = "sdk:removed-preset"
|
session_key = "sdk:removed-preset"
|
||||||
@@ -196,7 +196,7 @@ async def test_sdk_custom_model_preset_metadata_does_not_select_runtime(
|
|||||||
provider=base,
|
provider=base,
|
||||||
workspace=tmp_path,
|
workspace=tmp_path,
|
||||||
model="base-model",
|
model="base-model",
|
||||||
context_window_tokens=8_000,
|
context_window_tokens=16_000,
|
||||||
)
|
)
|
||||||
loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
|
loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
|
||||||
bot = Nanobot(loop)
|
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"
|
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:
|
def test_transient_response_preserves_only_durable_request_messages() -> None:
|
||||||
provider = _provider()
|
provider = _provider()
|
||||||
current_message = {"role": "user", "content": "continue"}
|
current_message = {"role": "user", "content": "continue"}
|
||||||
|
|||||||
Reference in New Issue
Block a user