refactor(agent/turn_overflow): unify attempt counting and token-scored compression across the overflow handlers

This commit is contained in:
Teknium
2026-09-02 18:30:28 -07:00
parent 5db64a48e8
commit 4cf65d7d89
+62 -63
View File
@@ -1,11 +1,10 @@
"""Overflow recovery for the conversation turn loop: 413 payload-too-large and
context-length errors after ``classify_api_error``.
Extracted from ``run_conversation``'s ``except`` branch. Each path either compresses and
signals a restart, defers softly (compression lock / transient block), or ends the turn
with a typed result. Nothing here imports ``agent.conversation_loop`` at module level
(cycle); loop-internal helpers and the token estimators that tests patch on the loop
module are imported lazily inside the handlers so they keep resolving through the loop.
Each path either compresses and signals a restart, defers softly (compression lock /
transient block), or ends the turn with a typed result. Nothing here imports
``agent.conversation_loop`` at module level (cycle); loop-internal helpers and the token
estimators that tests patch on the loop module are imported lazily inside the handlers.
"""
from __future__ import annotations
@@ -13,7 +12,7 @@ from __future__ import annotations
import logging
import time
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
from typing import Any, Dict, List, Optional, Tuple
from agent.conversation_compression import (
COMPRESSION_RETRY_MESSAGES_STATUS_TEMPLATE, COMPRESSION_RETRY_TOKENS_STATUS_TEMPLATE,
@@ -41,6 +40,8 @@ _GITHUB_MODELS_HINT = (
" setup` → GitHub Copilot), or pick any other provider.",
)
_MINIMAX_ANTHROPIC_URLS = ("https://api.minimax.io/anthropic", "https://api.minimaxi.com/anthropic")
@dataclass
class OverflowVerdict:
@@ -120,9 +121,12 @@ class _Recovery:
result.update(extra)
return self.verdict("return", result)
def exhausted(self, *, payload_too_large: bool = False) -> OverflowVerdict:
"""Terminal: ``compression_attempts`` exceeded ``max_compression_attempts``."""
def count_attempt(self, *, payload_too_large: bool = False) -> Optional[OverflowVerdict]:
"""Bump ``compression_attempts``; the terminal verdict once the cap is exceeded."""
self.compression_attempts += 1
cap = self.max_compression_attempts
if self.compression_attempts <= cap:
return None
if payload_too_large:
return self.fail_turn(
f"Request payload too large: max compression attempts ({cap}) reached.",
@@ -140,13 +144,13 @@ class _Recovery:
def compress(self, request_tokens: int, *, fail_on_timeout: bool = False) -> Optional[OverflowVerdict]:
"""One compression pass with the summary-failure cooldown bypassed (the
provider proved the request doesn't fit, #100661). Returns ``None`` when
history was compressed, or a soft-defer verdict when another path holds the
compression lock (#69870) or a timed guard no-oped the pass (#97488): the
attempt is refunded and the turn ends WITHOUT ``compression_exhausted`` so the
gateway does not auto-reset. With ``fail_on_timeout`` a host timeout (recovery
spent its wait budget with no committed summary) ends the turn via the typed
contract, since re-sending would hit the same overflow (#98722)."""
provider proved the request doesn't fit). Returns ``None`` when history was
compressed, or a soft-defer verdict when another path holds the compression
lock or a timed guard no-oped the pass: the attempt is refunded and the turn
ends WITHOUT ``compression_exhausted`` so the gateway does not auto-reset. With
``fail_on_timeout`` a host timeout (recovery spent its wait budget with no
committed summary) ends the turn via the typed contract, since re-sending would
hit the same overflow."""
from agent.conversation_loop import (
_COMPRESSION_TIMEOUT_FINAL_RESPONSE, _compression_deferred_result,
conversation_history_after_compression,
@@ -179,6 +183,29 @@ class _Recovery:
)
return None
def compress_scored_by_tokens(
self, request_tokens: int, *, fail_on_timeout: bool = False,
) -> Tuple[Optional[OverflowVerdict], bool, int]:
"""``compress`` scored in message count / tokens (context-overflow errors ARE
token-budget errors). Same-message-count compression (tool-result pruning,
in-place summarization) can shrink the request, so re-estimate rather than trust
the array length. Returns ``(deferred_verdict, shrank, new_tokens)``."""
from agent.conversation_loop import estimate_messages_tokens_rough
original_len = len(self.messages)
original_tokens = estimate_messages_tokens_rough(self.messages)
deferred = self.compress(request_tokens, fail_on_timeout=fail_on_timeout)
if deferred is not None:
return deferred, False, original_tokens
messages = self.messages
new_tokens = estimate_messages_tokens_rough(messages)
shrank_tokens = new_tokens > 0 and new_tokens < original_tokens * 0.95
if len(messages) < original_len:
self.agent._buffer_status(COMPRESSION_RETRY_MESSAGES_STATUS_TEMPLATE.format(before=original_len, after=len(messages)))
elif shrank_tokens:
self.agent._buffer_status(COMPRESSION_RETRY_TOKENS_STATUS_TEMPLATE.format(before=original_tokens, after=new_tokens))
return None, len(messages) < original_len or shrank_tokens, new_tokens
def request_tokens(self) -> int:
"""Overhead-aware request size (msgs + tools + system) so LCM forced-overflow
recovery arms on the TRUE request, not the tool-blind message count."""
@@ -190,13 +217,13 @@ class _Recovery:
def _recover_payload_too_large(st: _Recovery, _retry: TurnRetryState) -> OverflowVerdict:
"""413: compress and retry. A 413 is a BYTE-size error, so progress is scored in
payload bytes — never the token estimate, which is deliberately byte-blind to images
and wedged sessions on "no progress" (#88960 / #47339)."""
and wedged sessions on "no progress"."""
from agent.conversation_loop import estimate_messages_tokens_rough
agent = st.agent
st.compression_attempts += 1
if st.compression_attempts > st.max_compression_attempts:
return st.exhausted(payload_too_large=True)
exhausted = st.count_attempt(payload_too_large=True)
if exhausted is not None:
return exhausted
agent._buffer_status(
f"⚠️ Request payload too large (413) — compression attempt "
f"{st.compression_attempts}/{st.max_compression_attempts}..."
@@ -243,8 +270,6 @@ def _clamp_output_cap(st: _Recovery, _retry: TurnRetryState, available_out: int,
"""Output-cap error ("max_tokens too large": input fits but input + max_tokens >
window). The provider's available_tokens is the authoritative bound; also estimate
the real request shape (API-only content) and use the smaller minus a margin."""
from agent.conversation_loop import estimate_messages_tokens_rough
agent = st.agent
request_input_estimate = st.request_tokens()
local_available_out = old_ctx - request_input_estimate
@@ -262,24 +287,16 @@ def _clamp_output_cap(st: _Recovery, _retry: TurnRetryState, available_out: int,
f"context_length unchanged at {old_ctx:,})"
)
# Still count against compression_attempts so a recurring error can't loop forever.
st.compression_attempts += 1
if st.compression_attempts > st.max_compression_attempts:
return st.exhausted()
exhausted = st.count_attempt()
if exhausted is not None:
return exhausted
# Also compress history so the retry doesn't spin on max_tokens alone; dropping the
# middle window makes the total fit (#55546). Compression must never turn an
# output-cap error fatal — on error, fall through and retry on max_tokens alone.
# middle window makes the total fit. Compression must never turn an output-cap error
# fatal — on error, fall through and retry on max_tokens alone.
try:
original_len = len(st.messages)
original_tokens = estimate_messages_tokens_rough(st.messages)
deferred = st.compress(request_input_estimate)
deferred, _shrank, _new_tokens = st.compress_scored_by_tokens(request_input_estimate)
if deferred is not None:
return deferred
messages = st.messages
new_tokens = estimate_messages_tokens_rough(messages)
if len(messages) < original_len:
agent._buffer_status(COMPRESSION_RETRY_MESSAGES_STATUS_TEMPLATE.format(before=original_len, after=len(messages)))
elif new_tokens > 0 and new_tokens < original_tokens * 0.95:
agent._buffer_status(COMPRESSION_RETRY_TOKENS_STATUS_TEMPLATE.format(before=original_tokens, after=new_tokens))
except Exception:
logger.warning(
"%sOutput-cap compression hit an error; retrying on max_tokens only.", agent.log_prefix
@@ -315,13 +332,9 @@ def _adopt_provider_context_limit(st: _Recovery, error_msg: str, old_ctx: int) -
agent._buffer_vprint(f"⚠️ Context length exceeded — using provider limit: {old_ctx:,} → {new_ctx:,} tokens")
return new_ctx
_provider_lower = (getattr(agent, "provider", "") or "").lower()
_base_lower = (getattr(agent, "base_url", "") or "").rstrip("/").lower()
is_minimax_provider = (
_provider_lower in {"minimax", "minimax-cn"}
or _base_lower.startswith((
"https://api.minimax.io/anthropic", "https://api.minimaxi.com/anthropic"
))
(getattr(agent, "provider", "") or "").lower() in {"minimax", "minimax-cn"}
or (getattr(agent, "base_url", "") or "").rstrip("/").lower().startswith(_MINIMAX_ANTHROPIC_URLS)
)
if is_minimax_provider and "context window exceeds limit (" in error_msg:
agent._buffer_vprint(
@@ -340,8 +353,6 @@ def _recover_context_length(st: _Recovery, _retry: TurnRetryState, error_msg: st
"""Context-length error. Two shapes: "prompt too long" = INPUT overflows the window
(shrink context_length + compress); "max_tokens too large" = input fits but
input + max_tokens > window (shrink the OUTPUT cap only)."""
from agent.conversation_loop import estimate_messages_tokens_rough
agent = st.agent
old_ctx = agent.context_compressor.context_length
@@ -350,7 +361,7 @@ def _recover_context_length(st: _Recovery, _retry: TurnRetryState, error_msg: st
return _clamp_output_cap(st, _retry, available_out, old_ctx)
# Output-cap error with unparseable budget: compression can't help (input already
# fits) and would death-loop on the same 400. Fail fast (#55546).
# fits) and would death-loop on the same 400. Fail fast.
if is_output_cap_error(error_msg):
return st.fail_turn(
"max_tokens exceeds the provider's output cap for this model. "
@@ -372,30 +383,18 @@ def _recover_context_length(st: _Recovery, _retry: TurnRetryState, error_msg: st
new_ctx = _adopt_provider_context_limit(st, error_msg, old_ctx)
st.compression_attempts += 1
if st.compression_attempts > st.max_compression_attempts:
return st.exhausted()
exhausted = st.count_attempt()
if exhausted is not None:
return exhausted
agent._buffer_status(COMPRESSION_RETRY_TOO_LARGE_STATUS_TEMPLATE.format(
tokens=st.approx_tokens, attempt=st.compression_attempts, cap=st.max_compression_attempts,
))
original_len = len(st.messages)
original_tokens = estimate_messages_tokens_rough(st.messages)
deferred = st.compress(st.request_tokens(), fail_on_timeout=True)
deferred, shrank, new_tokens = st.compress_scored_by_tokens(st.request_tokens(), fail_on_timeout=True)
if deferred is not None:
return deferred
# Re-estimate: same-message-count compression (tool-result pruning, in-place
# summarization) can shrink the request (#39550).
messages = st.messages
new_tokens = estimate_messages_tokens_rough(messages)
st.approx_tokens = new_tokens
shrank_tokens = new_tokens > 0 and new_tokens < original_tokens * 0.95
if len(messages) < original_len or shrank_tokens or (new_ctx and new_ctx < old_ctx):
if len(messages) < original_len:
agent._buffer_status(COMPRESSION_RETRY_MESSAGES_STATUS_TEMPLATE.format(before=original_len, after=len(messages)))
elif shrank_tokens:
agent._buffer_status(COMPRESSION_RETRY_TOKENS_STATUS_TEMPLATE.format(before=original_tokens, after=new_tokens))
if shrank or (new_ctx and new_ctx < old_ctx):
time.sleep(2) # Brief pause between compression retries
# Rebuild the full request and force normal preflight to honor it; message
# count alone doesn't prove system/tool-inclusive pressure fell.
@@ -450,8 +449,8 @@ def recover_from_overflow(
return _recover_payload_too_large(st, _retry)
# Relay-wrapped output-cap 429s (parsed by the caller) go to the clamp, not
# failover or generic retries (#72281). The classifier also covers
# 400/disconnect + large-session heuristics.
# failover or generic retries. The classifier also covers 400/disconnect +
# large-session heuristics.
st.is_context_length_error = (
classified.reason == FailoverReason.context_overflow
or wrapped_output_cap_budget is not None