refactor(agent): unify runner request fitting (#5612)

* refactor(agent): unify runner request fitting

* fix(agent): fit every runner model request

* fix(agent): count resumed state during request fitting

* refactor(agent): consolidate request fitting state
This commit is contained in:
chengyongru
2026-08-31 18:08:21 +08:00
committed by GitHub
parent 6d6d58d329
commit e111b83af6
9 changed files with 888 additions and 491 deletions
+101 -134
View File
@@ -13,6 +13,7 @@ from typing import TYPE_CHECKING, Any, cast
from loguru import logger
from nanobot.providers.base import LLMUsage
from nanobot.utils.helpers import (
estimate_message_tokens,
estimate_prompt_tokens_chain,
@@ -27,12 +28,6 @@ if TYPE_CHECKING:
from nanobot.providers.base import LLMProvider
SNIP_SAFETY_BUFFER = 1024
MICROCOMPACT_MIN_CHARS = 500
INFLIGHT_COMPACT_TARGET_RATIO = 0.85
COMPACTABLE_TOOLS = frozenset({
"read_file", "exec", "grep", "find_files",
"web_search", "web_fetch", "list_dir", "list_exec_sessions",
})
# read_file is the recovery path for persisted results; exempting it prevents persist->read->persist loops.
TOOL_RESULT_OFFLOAD_EXEMPT_TOOLS = frozenset({"read_file"})
BACKFILL_CONTENT = "[Tool result unavailable — call was interrupted or lost]"
@@ -41,6 +36,27 @@ PLACEHOLDER_TEXTS = frozenset({
})
class ContextWindowExceededError(RuntimeError):
"""Raised before a locally fitted request that still exceeds its budget."""
def __init__(
self,
*,
session_key: str | None,
estimated_tokens: int,
input_budget: int,
source: str,
) -> None:
self.session_key = session_key
self.estimated_tokens = estimated_tokens
self.input_budget = input_budget
self.source = source
super().__init__(
"Model input still exceeds the local context budget after request fitting "
f"for {session_key or 'default'}: {estimated_tokens}/{input_budget} via {source}"
)
def _tool_call_name_is_valid(tool_call: Any) -> bool:
"""Whether a persisted OpenAI-style tool_call carries a usable name.
@@ -67,7 +83,6 @@ class ContextGovernanceConfig:
context_window_tokens: int | None = None
context_block_limit: int | None = None
max_tokens: int | None = None
inflight_start_index: int = 0
class ContextGovernor:
@@ -77,17 +92,85 @@ class ContextGovernor:
self,
config: ContextGovernanceConfig,
messages: list[dict[str, Any]],
compacted_tool_call_ids: set[str],
) -> list[dict[str, Any]]:
updated = self.strip_placeholder_assistant_messages(messages)
updated = self.strip_malformed_tool_calls(updated)
updated = self.drop_orphan_tool_results(updated)
updated = self.backfill_missing_tool_results(updated)
updated = self.apply_tool_result_budget(config, updated)
updated = self.compact_inflight_overflow(config, updated, compacted_tool_call_ids)
updated = self.snip_history(config, updated)
return self.apply_tool_result_budget(config, updated)
def fit_to_budget(
self,
config: ContextGovernanceConfig,
messages: list[dict[str, Any]],
*,
tool_definitions: list[dict[str, Any]] | None,
) -> list[dict[str, Any]]:
"""Fit a model-facing copy while keeping the source transcript intact."""
updated = self.snip_history(
config,
messages,
tool_definitions=tool_definitions,
force=True,
)
updated = self.drop_orphan_tool_results(updated)
return self.backfill_missing_tool_results(updated)
updated = self.backfill_missing_tool_results(updated)
if not config.context_window_tokens:
return updated
budget = self.input_budget(config)
estimated, source = estimate_prompt_tokens_chain(
config.provider,
config.model,
updated,
tool_definitions,
)
if budget > 0 and estimated <= budget:
return updated
raise ContextWindowExceededError(
session_key=config.session_key,
estimated_tokens=estimated,
input_budget=budget,
source=source,
)
def fit_request(
self,
config: ContextGovernanceConfig,
messages: list[dict[str, Any]],
usage: LLMUsage | None,
*,
usage_matches_messages: bool,
tool_definitions: list[dict[str, Any]] | None,
request_context_tokens: int | None = None,
) -> tuple[list[dict[str, Any]], bool]:
"""Fit the request when its measured or estimated input is pressured."""
if not config.context_window_tokens:
return messages, False
budget = self.input_budget(config)
if (
request_context_tokens is None
and usage_matches_messages
and usage is not None
and usage.context_tokens is not None
):
pressured = budget <= 0 or usage.context_tokens >= budget
else:
estimated, _ = estimate_prompt_tokens_chain(
config.provider,
config.model,
messages,
tool_definitions,
)
if request_context_tokens is not None:
estimated = max(estimated, request_context_tokens)
pressured = budget <= 0 or estimated >= budget
if not pressured:
return messages, False
return self.fit_to_budget(
config,
messages,
tool_definitions=tool_definitions,
), True
@staticmethod
def input_budget(config: ContextGovernanceConfig) -> int:
@@ -326,71 +409,13 @@ class ContextGovernor:
updated[idx]["content"] = normalized
return updated
def compact_inflight_overflow(
self,
config: ContextGovernanceConfig,
messages: list[dict[str, Any]],
compacted_tool_call_ids: set[str],
) -> list[dict[str, Any]]:
"""Compact in-flight tool results only when the request would overflow."""
budget = self.input_budget(config)
if budget <= 0:
return messages
tools = config.tools.get_definitions()
updated = self._apply_recorded_compactions(messages, compacted_tool_call_ids)
estimate, source = estimate_prompt_tokens_chain(
config.provider,
config.model,
updated,
tools,
)
if estimate <= budget:
return updated
target = int(budget * INFLIGHT_COMPACT_TARGET_RATIO)
candidates = self._inflight_compaction_candidates(
config,
updated,
compacted_tool_call_ids,
)
if not candidates:
return updated
for candidate_idx, (idx, tool_call_id) in enumerate(candidates):
is_newest_candidate = candidate_idx == len(candidates) - 1
if is_newest_candidate and estimate <= budget:
break
if tool_call_id in compacted_tool_call_ids:
continue
if updated is messages:
updated = [dict(m) for m in messages]
compacted_tool_call_ids.add(tool_call_id)
self._compact_tool_result_at(updated, idx)
estimate, source = estimate_prompt_tokens_chain(
config.provider,
config.model,
updated,
tools,
)
if estimate <= target:
break
logger.debug(
"In-flight context compaction for {}: prompt={} budget={} target={} via {}, ids={}",
config.session_key or "default",
estimate,
budget,
target,
source,
len(compacted_tool_call_ids),
)
return updated
def snip_history(
self,
config: ContextGovernanceConfig,
messages: list[dict[str, Any]],
*,
tool_definitions: list[dict[str, Any]] | None,
force: bool = False,
) -> list[dict[str, Any]]:
if not messages or not config.context_window_tokens:
return messages
@@ -399,14 +424,13 @@ class ContextGovernor:
if budget <= 0:
return messages
tools = config.tools.get_definitions()
estimate, _ = estimate_prompt_tokens_chain(
config.provider,
config.model,
messages,
tools,
tool_definitions,
)
if estimate <= budget:
if not force and estimate <= budget:
return messages
system_messages = [dict(msg) for msg in messages if msg.get("role") == "system"]
@@ -419,7 +443,7 @@ class ContextGovernor:
config.provider,
config.model,
system_messages,
tools,
tool_definitions,
)
remaining_budget = max(0, budget - max(system_tokens, fixed_tokens))
kept: list[dict[str, Any]] = []
@@ -434,16 +458,6 @@ class ContextGovernor:
return system_messages + self._legal_history_tail(kept, non_system)
@staticmethod
def _tool_result_compaction_message(message: dict[str, Any]) -> str:
name = message.get("name", "tool")
return (
f"Error: The previous {name} result was compacted to fit context because it was too "
"large. Do not repeat the same call unchanged. Retry with a narrower path, query, "
"range, or result limit, use another tool, or tell the user the task cannot fit in "
"the available context."
)
def _legal_history_tail(
self,
kept: list[dict[str, Any]],
@@ -462,50 +476,3 @@ class ContextGovernor:
if messages[idx].get("role") == "user":
return messages[idx:]
return []
def _apply_recorded_compactions(
self,
messages: list[dict[str, Any]],
compacted_tool_call_ids: set[str],
) -> list[dict[str, Any]]:
if not compacted_tool_call_ids:
return messages
updated = messages
for idx, msg in enumerate(messages):
if msg.get("role") != "tool":
continue
tool_call_id = msg.get("tool_call_id")
if not tool_call_id or str(tool_call_id) not in compacted_tool_call_ids:
continue
compaction_message = self._tool_result_compaction_message(msg)
if msg.get("content") == compaction_message:
continue
if updated is messages:
updated = [dict(m) for m in messages]
updated[idx]["content"] = compaction_message
return updated
def _inflight_compaction_candidates(
self,
config: ContextGovernanceConfig,
messages: list[dict[str, Any]],
compacted_tool_call_ids: set[str],
) -> list[tuple[int, str]]:
compactable: list[tuple[int, str]] = []
for idx, msg in enumerate(messages):
if idx < config.inflight_start_index:
continue
if msg.get("role") != "tool" or msg.get("name") not in COMPACTABLE_TOOLS:
continue
tool_call_id = msg.get("tool_call_id")
if not tool_call_id or str(tool_call_id) in compacted_tool_call_ids:
continue
content = msg.get("content")
if not isinstance(content, str) or len(content) < MICROCOMPACT_MIN_CHARS:
continue
compactable.append((idx, str(tool_call_id)))
return compactable
def _compact_tool_result_at(self, messages: list[dict[str, Any]], idx: int) -> None:
messages[idx]["content"] = self._tool_result_compaction_message(messages[idx])
+136 -58
View File
@@ -139,6 +139,17 @@ class AgentRunResult:
provider_state: ProviderConversationState | None = field(default=None, repr=False)
@dataclass(slots=True)
class _ModelRequestState:
"""Per-run state used to govern the next provider request."""
config: ContextGovernanceConfig
conversation: ProviderConversationStateController
usage: LLMUsage | None = None
messages: list[dict[str, Any]] | None = None
tool_definitions: list[dict[str, Any]] | None = None
class AgentRunner:
"""Run a tool-capable LLM loop without product-layer concerns."""
@@ -500,7 +511,6 @@ class AgentRunner:
length_recovery_parts: list[str] = []
had_injections = False
injection_cycles = 0
compacted_tool_call_ids: set[str] = set()
pending_stream_content: str | None = None
conversation_state = ProviderConversationStateController(
provider=spec.runtime.provider,
@@ -519,39 +529,29 @@ class AgentRunner:
context_window_tokens=spec.runtime.context_window_tokens,
context_block_limit=spec.context_block_limit,
max_tokens=spec.runtime.generation.max_tokens,
inflight_start_index=len(messages),
)
request_state = _ModelRequestState(
config=governance_config,
conversation=conversation_state,
)
for iteration in range(spec.max_iterations):
# Keep the persisted conversation untouched. Context governance
# may repair or compact historical messages for the model, but
# those synthetic edits must not shift the append boundary used
# later when the caller saves only the new turn. A governance
# failure must stop the run instead of sending an ungoverned copy.
messages_for_model = self.context_governor.prepare_for_model(
governance_config,
messages,
compacted_tool_call_ids,
)
context = AgentHookContext(
iteration=iteration,
messages=messages,
session_key=spec.session_key,
)
await hook.before_iteration(context)
provider_context = conversation_state.prepare_request(
messages,
context_window_tokens=spec.runtime.context_window_tokens,
model_messages=messages_for_model,
)
response = await self._request_model(
spec,
messages_for_model,
messages,
hook,
context,
conversation_state=conversation_state,
provider_context=provider_context,
request_state=request_state,
transcript=messages,
)
assert request_state.messages is not None
messages_for_model = request_state.messages
conversation_state.observe_response(response, messages)
context.response = response
context.tool_calls = list(response.tool_calls)
@@ -563,7 +563,7 @@ class AgentRunner:
response.content,
)
response.content = cleaned_content
raw_usage = self._usage_or_estimate(spec, messages_for_model, response)
raw_usage = self._record_request_usage(spec, request_state, response)
context.usage = raw_usage
usage = self._merge_usage(usage, raw_usage)
if reasoning_text and not context.streamed_reasoning:
@@ -637,7 +637,6 @@ class AgentRunner:
self.context_governor.prepare_for_model(
governance_config,
messages,
compacted_tool_call_ids,
)
if response.provider_state is not None
else None
@@ -703,14 +702,13 @@ class AgentRunner:
)
if hook.wants_streaming():
await hook.on_stream_end(context, resuming=False)
retry_messages = self._finalization_retry_messages(messages_for_model)
response = await self._request_finalization_retry(
spec,
messages_for_model,
request_state=request_state,
transcript=messages,
conversation_state=conversation_state,
)
retry_usage = self._usage_or_estimate(spec, retry_messages, response)
retry_usage = self._record_request_usage(spec, request_state, response)
usage = self._merge_usage(usage, retry_usage)
raw_usage = self._merge_usage(raw_usage, retry_usage)
context.response = response
@@ -897,7 +895,7 @@ class AgentRunner:
hook,
messages,
usage,
conversation_state,
request_state=request_state,
)
if terminal_content is None:
terminal_content = self._max_iterations_fallback(spec)
@@ -944,6 +942,60 @@ class AgentRunner:
kwargs["reasoning_effort"] = generation.reasoning_effort
return kwargs
def _prepare_model_request(
self,
state: _ModelRequestState,
messages: list[dict[str, Any]],
*,
tool_definitions: list[dict[str, Any]] | None,
transcript: list[dict[str, Any]] | None = None,
) -> tuple[list[dict[str, Any]], ProviderCallContext | None]:
"""Prepare, fit, and record the exact payload sent to a provider."""
prepared = self.context_governor.prepare_for_model(state.config, messages)
supplemental_messages = (
[prepared[-1]] if transcript is not None and tool_definitions is None else None
)
model_messages = None if supplemental_messages is not None else prepared
request_context_tokens = (
state.conversation.estimate_request_context_tokens(
transcript,
model_messages=model_messages,
supplemental_messages=supplemental_messages,
tool_definitions=tool_definitions,
)
if transcript is not None
else None
)
usage_matches_messages = (
state.messages is not None
and prepared == state.messages
and tool_definitions == state.tool_definitions
)
prepared, fitted = self.context_governor.fit_request(
state.config,
prepared,
state.usage,
usage_matches_messages=usage_matches_messages,
tool_definitions=tool_definitions,
request_context_tokens=request_context_tokens,
)
provider_context = (
state.conversation.prepare_request(
transcript,
context_window_tokens=state.config.context_window_tokens,
model_messages=model_messages,
supplemental_messages=supplemental_messages,
resume_state=not fitted,
)
if transcript is not None
else state.conversation.independent_request_context(
context_window_tokens=state.config.context_window_tokens,
)
)
state.messages = deepcopy(prepared)
state.tool_definitions = deepcopy(tool_definitions)
return prepared, provider_context
async def _request_model(
self,
spec: AgentRunSpec,
@@ -951,16 +1003,23 @@ class AgentRunner:
hook: AgentHook,
context: AgentHookContext,
*,
request_state: _ModelRequestState,
malformed_retry: bool = False,
conversation_state: ProviderConversationStateController,
provider_context: ProviderCallContext | None = None,
transcript: list[dict[str, Any]] | None,
) -> LLMResponse:
timeout_s = self._resolve_llm_timeout_s(spec)
tool_definitions = spec.tools.get_definitions()
messages, provider_context = self._prepare_model_request(
request_state,
messages,
tool_definitions=tool_definitions,
transcript=transcript,
)
kwargs = self._build_request_kwargs(
spec,
messages,
tools=spec.tools.get_definitions(),
tools=tool_definitions,
)
wants_streaming = hook.wants_streaming()
@@ -1138,11 +1197,9 @@ class AgentRunner:
)
return await self._request_model(
spec, retry_messages, hook, context,
request_state=request_state,
malformed_retry=True,
conversation_state=conversation_state,
provider_context=conversation_state.independent_request_context(
context_window_tokens=spec.runtime.context_window_tokens,
),
transcript=None,
)
if (
all_dropped
@@ -1158,9 +1215,7 @@ class AgentRunner:
return await self._request_no_tools(
spec,
fallback_messages,
provider_context=conversation_state.independent_request_context(
context_window_tokens=spec.runtime.context_window_tokens,
),
request_state=request_state,
)
return response
@@ -1228,21 +1283,17 @@ class AgentRunner:
spec: AgentRunSpec,
messages: list[dict[str, Any]],
*,
request_state: _ModelRequestState,
transcript: list[dict[str, Any]],
conversation_state: ProviderConversationStateController,
) -> LLMResponse:
retry_messages = self._finalization_retry_messages(messages)
provider_context = conversation_state.prepare_request(
transcript,
context_window_tokens=spec.runtime.context_window_tokens,
supplemental_messages=[retry_messages[-1]],
)
response = await self._request_no_tools(
spec,
retry_messages,
provider_context=provider_context,
request_state=request_state,
transcript=transcript,
)
conversation_state.observe_response(
request_state.conversation.observe_response(
response,
transcript,
adopt_candidate_state=False,
@@ -1261,16 +1312,15 @@ class AgentRunner:
hook: AgentHook,
messages: list[dict[str, Any]],
usage: LLMUsage | None,
conversation_state: ProviderConversationStateController,
*,
request_state: _ModelRequestState,
) -> tuple[str | None, LLMUsage | None]:
retry_messages = self._budget_exhausted_finalization_messages(messages)
try:
response = await self._request_no_tools(
spec,
retry_messages,
provider_context=conversation_state.independent_request_context(
context_window_tokens=spec.runtime.context_window_tokens,
),
request_state=request_state,
)
except Exception:
logger.exception(
@@ -1279,7 +1329,7 @@ class AgentRunner:
)
return None, usage
raw_usage = self._usage_or_estimate(spec, retry_messages, response)
raw_usage = self._record_request_usage(spec, request_state, response)
usage = self._merge_usage(usage, raw_usage)
if response.finish_reason == "error" or response.has_tool_calls:
logger.warning(
@@ -1308,8 +1358,15 @@ class AgentRunner:
spec: AgentRunSpec,
messages: list[dict[str, Any]],
*,
provider_context: ProviderCallContext | None = None,
request_state: _ModelRequestState,
transcript: list[dict[str, Any]] | None = None,
) -> LLMResponse:
messages, provider_context = self._prepare_model_request(
request_state,
messages,
tool_definitions=None,
transcript=transcript,
)
kwargs = self._build_request_kwargs(
spec,
messages,
@@ -1321,17 +1378,18 @@ class AgentRunner:
)
timeout_s = self._resolve_llm_timeout_s(spec)
try:
return (
response = (
await coro
if timeout_s is None
else await asyncio.wait_for(coro, timeout=timeout_s)
)
except asyncio.TimeoutError:
return LLMResponse(
response = LLMResponse(
content=f"Error calling LLM: timed out after {timeout_s:g}s",
finish_reason="error",
error_kind="timeout",
)
return response
@staticmethod
def _resolve_llm_timeout_s(spec: AgentRunSpec) -> float | None:
@@ -1373,33 +1431,53 @@ class AgentRunner:
spec: AgentRunSpec,
messages: list[dict[str, Any]],
response: LLMResponse,
*,
tool_definitions: list[dict[str, Any]] | None,
) -> LLMUsage | None:
usage = response.usage
if response.finish_reason == "error":
if usage is None or usage.total_tokens == 0:
usage = LLMUsage.empty_request()
elif usage is None or usage.total_tokens == 0:
usage = self._estimate_response_usage(spec, messages, response)
usage = self._estimate_response_usage(
spec,
messages,
response,
tool_definitions=tool_definitions,
)
return usage.with_timing(
generation_ms=response.generation_ms,
ttft_ms=response.ttft_ms,
)
def _record_request_usage(
self,
spec: AgentRunSpec,
state: _ModelRequestState,
response: LLMResponse,
) -> LLMUsage | None:
assert state.messages is not None
state.usage = self._usage_or_estimate(
spec,
state.messages,
response,
tool_definitions=state.tool_definitions,
)
return state.usage
def _estimate_response_usage(
self,
spec: AgentRunSpec,
messages: list[dict[str, Any]],
response: LLMResponse,
*,
tool_definitions: list[dict[str, Any]] | None,
) -> LLMUsage:
try:
tools = spec.tools.get_definitions()
except Exception:
tools = None
prompt_tokens, _ = estimate_prompt_tokens_chain(
spec.runtime.provider,
spec.runtime.model,
messages,
tools,
tool_definitions,
)
assistant_message = build_assistant_message(
response.content or "",
+42 -1
View File
@@ -11,6 +11,7 @@ from nanobot.providers.base import (
ProviderCallContext,
ProviderConversationState,
)
from nanobot.utils.helpers import estimate_prompt_tokens_chain
_PROVIDER_STATE_OUTPUT_META = "provider_state_output"
_PROVIDER_STATE_BOUNDARY_META = "provider_state_boundary"
@@ -69,6 +70,37 @@ class ProviderConversationStateController:
session_id=self._session_id,
)
def estimate_request_context_tokens(
self,
messages: list[dict[str, Any]],
*,
model_messages: list[dict[str, Any]] | None = None,
supplemental_messages: list[dict[str, Any]] | None = None,
tool_definitions: list[dict[str, Any]] | None = None,
) -> int | None:
"""Estimate resumed state plus the pending delta for the next request."""
state = self.checkpoint(messages, model_messages=model_messages)
if state is None:
return None
context_tokens = state.payload.get("context_tokens")
if (
isinstance(context_tokens, bool)
or not isinstance(context_tokens, int)
or context_tokens < 0
):
return None
pending_messages = [
*state.pending_messages,
*(supplemental_messages or []),
]
delta_tokens, _ = estimate_prompt_tokens_chain(
self._provider,
self._model,
pending_messages,
tool_definitions,
)
return context_tokens + max(0, delta_tokens)
def prepare_request(
self,
messages: list[dict[str, Any]],
@@ -76,11 +108,20 @@ class ProviderConversationStateController:
context_window_tokens: int | None,
model_messages: list[dict[str, Any]] | None = None,
supplemental_messages: list[dict[str, Any]] | None = None,
resume_state: bool = True,
) -> ProviderCallContext | None:
"""Build typed context for the next request and remember its durable delta."""
"""Build context for the next request and remember its durable delta.
``resume_state=False`` abandons opaque history when local request
fitting has produced a new independent model-facing context.
"""
independent_context = self.independent_request_context(
context_window_tokens=context_window_tokens,
)
if not resume_state:
self._state = None
self._request_messages = []
return independent_context
if self._state is None:
self._request_messages = []
return independent_context