refactor(agent): extract post-tool-call compression decision into turn_preflight.compress_after_tool_results

This commit is contained in:
Teknium
2026-09-02 12:10:56 -07:00
parent f0d2355e26
commit 4f53d7bd28
2 changed files with 191 additions and 110 deletions
+26 -110
View File
@@ -18,7 +18,9 @@ import time
from typing import Any, Dict, List, Optional
from agent.codex_responses_adapter import _summarize_user_message_for_log
from agent.conversation_compression import conversation_history_after_compression
from agent.conversation_compression import (
conversation_history_after_compression, # noqa: F401 — resolved lazily by turn_overflow/turn_preflight/turn_recovery (tests patch it here)
)
from agent.display import KawaiiSpinner
from agent.error_classifier import FailoverReason, classify_api_error
from agent.fast_mode import begin_turn as begin_fast_mode_turn
@@ -37,7 +39,7 @@ from agent.turn_empty_response import recover_empty_response
from agent.turn_stop_gates import apply_stop_gates
from agent.turn_tool_validation import validate_tool_calls
from agent.turn_truncation import continue_codex_incomplete, recover_from_truncation
from agent.turn_preflight import run_preflight_compression
from agent.turn_preflight import compress_after_tool_results, run_preflight_compression
from agent.turn_recovery import (
route_classified_error,
describe_invalid_response,
@@ -68,7 +70,7 @@ from agent.model_metadata import (
_estimate_tools_tokens_rough,
anchored_context_tokens,
estimate_messages_tokens_rough,
estimate_request_tokens_rough,
estimate_request_tokens_rough, # noqa: F401 — resolved lazily by turn_overflow/turn_preflight/turn_recovery (tests patch it here)
save_context_length, # noqa: F401 — resolved lazily by agent.turn_overflow (tests patch it here)
)
from agent.process_bootstrap import _install_safe_stdio
@@ -3894,113 +3896,27 @@ def run_conversation(
if _tc_names == {"execute_code"}:
agent.iteration_budget.refund()
# Decide compression from API-reported prompt tokens (tight lower bound;
# tool results get counted on the next call). If last_prompt_tokens is 0
# (disconnect / no usage data) fall back to a rough estimate. (#2153)
_compressor = agent.context_compressor
if _compressor.last_prompt_tokens > 0:
# Only prompt_tokens: thinking models inflate completion_tokens with
# reasoning that uses no context → premature compression. (#12026)
_real_tokens = _compressor.last_prompt_tokens
elif _compressor.last_prompt_tokens == -1:
# Compression just ran, no API prompt count yet: don't treat a rough
# schema-heavy post-compression estimate as real context pressure.
_real_tokens = 0
else:
# Include tool schemas (20-30K tokens the messages-only estimate
# misses) and stay route-aware: on a compacted native-Codex session
# the generic durable-history figure would false-trigger. (#14695)
_real_tokens = _midturn_request_pressure_tokens(
agent,
messages,
active_system_prompt or "",
estimate_request_tokens_rough(
messages, tools=agent.tools or None
),
)
if (
agent.compression_enabled
and compression_attempts < max_compression_attempts
and _compressor.should_compress(_real_tokens)
):
compression_attempts += 1
# Compression is running: reset blocked-overflow warning dedup so a
# future blocked turn can warn again. getattr: test doubles lack it.
_clear_warn = getattr(agent, "_clear_context_overflow_warn", None)
if callable(_clear_warn):
_clear_warn()
agent._safe_print(" ⟳ compacting context…")
_post_tool_input = messages
# Pass overhead-aware _real_tokens, not last_prompt_tokens (0 in
# the no-usage fallback), so the overflow guard sees the true size.
messages, active_system_prompt = agent._compress_context(
messages, system_message,
approx_tokens=_real_tokens,
task_id=effective_task_id,
)
if (
messages is _post_tool_input
and compression_skipped_due_to_lock(agent)
):
# Lock-skip no-op is a temporary defer, not evidence about
# compressibility: refund so a lock-loser loop doesn't burn the
# budget toward compression_exhausted. (#69870)
compression_attempts -= 1
else:
conversation_history = conversation_history_after_compression(
agent, messages, conversation_history
)
if _should_skip_model_call_for_reference_handoff(
messages, user_message
):
logger.info(
"Skipping post-tool compaction model call: "
"reference-only handoff would be the sole "
"active user turn (#80622)"
)
if not final_response:
final_response = _HANDOFF_SKIP_FINAL_RESPONSE
_turn_exit_reason = "compaction_handoff_not_actionable"
break
elif agent.compression_enabled:
# Over threshold but compression blocked (cooldown/anti-thrash):
# deduped warning so context can't silently overflow. (#62625)
_block_reason = None
_info = getattr(_compressor, "should_compress_info", None)
if _info is not None:
try:
_block_reason = _info(_real_tokens)[1]
except Exception:
_block_reason = None
if _block_reason:
agent._warn_context_overflow_blocked(
_block_reason,
_real_tokens,
int(getattr(_compressor, "threshold_tokens", 0) or 0),
)
# Proactive tool-result prune (deterministic, no LLM, keeps tail):
# no-op unless proactive_prune_tokens is exceeded; commits only past
# proactive_prune_min_reclaim_tokens so cache breaks stay episodic.
_prune = getattr(_compressor, "prune_tool_results_only", None)
if callable(_prune):
try:
_pruned_msgs, _pruned_n = _prune(
messages, current_tokens=_real_tokens
)
except Exception:
logger.debug(
"proactive tool-result prune failed; skipping",
exc_info=True,
)
_pruned_msgs, _pruned_n = messages, 0
# Standard no-op caller contract: only commit when the
# engine returned a NEW list object with a non-zero count.
if _pruned_n and _pruned_msgs is not messages:
# Do NOT rebuild conversation_history: rows already carry
# _DB_PERSISTED_MARKER, and on a stale in-place flag the
# helper could seed unpersisted rows into history_ids.
messages = _pruned_msgs
_ptc = compress_after_tool_results(
agent,
messages=messages,
system_message=system_message,
user_message=user_message,
active_system_prompt=active_system_prompt,
conversation_history=conversation_history,
compression_attempts=compression_attempts,
max_compression_attempts=max_compression_attempts,
effective_task_id=effective_task_id,
final_response=final_response,
turn_exit_reason=_turn_exit_reason,
)
messages = _ptc.messages
active_system_prompt = _ptc.active_system_prompt
conversation_history = _ptc.conversation_history
compression_attempts = _ptc.compression_attempts
final_response = _ptc.final_response
_turn_exit_reason = _ptc.turn_exit_reason
if _ptc.end_turn:
break
# Save session log incrementally (so progress is visible even if interrupted)
agent._session_messages = messages
+165
View File
@@ -339,3 +339,168 @@ def run_preflight_compression(
max_compression_attempts,
))
return _verdict("proceed")
@dataclass
class PostToolCompressionVerdict:
"""``end_turn`` True → a reference-only compaction handoff would be the sole active
user turn (#80622): stop without another model call (``final_response`` /
``turn_exit_reason`` set)."""
end_turn: bool
messages: List[Dict[str, Any]]
active_system_prompt: Any
conversation_history: Any
compression_attempts: int
final_response: Any
turn_exit_reason: Any
def compress_after_tool_results(
agent: Any,
*,
messages: List[Dict[str, Any]],
system_message: Any,
user_message: Any,
active_system_prompt: Any,
conversation_history: Any,
compression_attempts: int,
max_compression_attempts: int,
effective_task_id: Any,
final_response: Any,
turn_exit_reason: Any,
) -> PostToolCompressionVerdict:
"""Post-tool-call compression decision. Pressure comes from API-reported
``prompt_tokens`` (a tight lower bound; thinking models inflate completion tokens,
#12026), ``0`` right after compression (no real count yet), else the route-aware
overhead-inclusive estimate (#14695). Over threshold but blocked → deduped warning
(#62625) plus the deterministic tool-result-only prune, committed only when the
engine returns a NEW list (never rebuild ``conversation_history`` for it)."""
from agent.conversation_loop import (
_HANDOFF_SKIP_FINAL_RESPONSE,
_midturn_request_pressure_tokens,
_should_skip_model_call_for_reference_handoff,
estimate_request_tokens_rough,
)
_turn_exit_reason = turn_exit_reason
def _verdict(end_turn: bool) -> PostToolCompressionVerdict:
return PostToolCompressionVerdict(
end_turn=end_turn,
messages=messages,
active_system_prompt=active_system_prompt,
conversation_history=conversation_history,
compression_attempts=compression_attempts,
final_response=final_response,
turn_exit_reason=_turn_exit_reason,
)
# Decide compression from API-reported prompt tokens (tight lower bound;
# tool results get counted on the next call). If last_prompt_tokens is 0
# (disconnect / no usage data) fall back to a rough estimate. (#2153)
_compressor = agent.context_compressor
if _compressor.last_prompt_tokens > 0:
# Only prompt_tokens: thinking models inflate completion_tokens with
# reasoning that uses no context → premature compression. (#12026)
_real_tokens = _compressor.last_prompt_tokens
elif _compressor.last_prompt_tokens == -1:
# Compression just ran, no API prompt count yet: don't treat a rough
# schema-heavy post-compression estimate as real context pressure.
_real_tokens = 0
else:
# Include tool schemas (20-30K tokens the messages-only estimate
# misses) and stay route-aware: on a compacted native-Codex session
# the generic durable-history figure would false-trigger. (#14695)
_real_tokens = _midturn_request_pressure_tokens(
agent,
messages,
active_system_prompt or "",
estimate_request_tokens_rough(
messages, tools=agent.tools or None
),
)
if (
agent.compression_enabled
and compression_attempts < max_compression_attempts
and _compressor.should_compress(_real_tokens)
):
compression_attempts += 1
# Compression is running: reset blocked-overflow warning dedup so a
# future blocked turn can warn again. getattr: test doubles lack it.
_clear_warn = getattr(agent, "_clear_context_overflow_warn", None)
if callable(_clear_warn):
_clear_warn()
agent._safe_print(" ⟳ compacting context…")
_post_tool_input = messages
# Pass overhead-aware _real_tokens, not last_prompt_tokens (0 in
# the no-usage fallback), so the overflow guard sees the true size.
messages, active_system_prompt = agent._compress_context(
messages, system_message,
approx_tokens=_real_tokens,
task_id=effective_task_id,
)
if (
messages is _post_tool_input
and compression_skipped_due_to_lock(agent)
):
# Lock-skip no-op is a temporary defer, not evidence about
# compressibility: refund so a lock-loser loop doesn't burn the
# budget toward compression_exhausted. (#69870)
compression_attempts -= 1
else:
conversation_history = conversation_history_after_compression(
agent, messages, conversation_history
)
if _should_skip_model_call_for_reference_handoff(
messages, user_message
):
logger.info(
"Skipping post-tool compaction model call: "
"reference-only handoff would be the sole "
"active user turn (#80622)"
)
if not final_response:
final_response = _HANDOFF_SKIP_FINAL_RESPONSE
_turn_exit_reason = "compaction_handoff_not_actionable"
return _verdict(True)
elif agent.compression_enabled:
# Over threshold but compression blocked (cooldown/anti-thrash):
# deduped warning so context can't silently overflow. (#62625)
_block_reason = None
_info = getattr(_compressor, "should_compress_info", None)
if _info is not None:
try:
_block_reason = _info(_real_tokens)[1]
except Exception:
_block_reason = None
if _block_reason:
agent._warn_context_overflow_blocked(
_block_reason,
_real_tokens,
int(getattr(_compressor, "threshold_tokens", 0) or 0),
)
# Proactive tool-result prune (deterministic, no LLM, keeps tail):
# no-op unless proactive_prune_tokens is exceeded; commits only past
# proactive_prune_min_reclaim_tokens so cache breaks stay episodic.
_prune = getattr(_compressor, "prune_tool_results_only", None)
if callable(_prune):
try:
_pruned_msgs, _pruned_n = _prune(
messages, current_tokens=_real_tokens
)
except Exception:
logger.debug(
"proactive tool-result prune failed; skipping",
exc_info=True,
)
_pruned_msgs, _pruned_n = messages, 0
# Standard no-op caller contract: only commit when the
# engine returned a NEW list object with a non-zero count.
if _pruned_n and _pruned_msgs is not messages:
# Do NOT rebuild conversation_history: rows already carry
# _DB_PERSISTED_MARKER, and on a stale in-place flag the
# helper could seed unpersisted rows into history_ids.
messages = _pruned_msgs
return _verdict(False)