diff --git a/agent/turn_overflow.py b/agent/turn_overflow.py index 3ee123647a..1e059194c7 100644 --- a/agent/turn_overflow.py +++ b/agent/turn_overflow.py @@ -64,10 +64,11 @@ class OverflowVerdict: is_context_length_error: bool -@dataclass -class _Recovery: - """Mutable working state shared by the overflow sub-handlers; the loop locals - they may rebind live here and are handed back through ``verdict()``.""" +@dataclass(kw_only=True) +class _Recovery(OverflowVerdict): + """Working state shared by the overflow sub-handlers — the verdict itself, plus the + read-only call context; handlers mutate the loop-local fields and ``done()`` stamps + the action.""" agent: Any api_messages: Any @@ -75,23 +76,14 @@ class _Recovery: effective_task_id: Any api_call_count: int max_compression_attempts: int - messages: List[Dict[str, Any]] - active_system_prompt: Any - conversation_history: Any - approx_tokens: int - compression_attempts: int + action: str = "fallthrough" + result: Optional[Dict[str, Any]] = None provider_overflow_recovery_pending: bool = False is_context_length_error: bool = False - def verdict(self, action: str, result: Optional[Dict[str, Any]] = None) -> OverflowVerdict: - return OverflowVerdict( - action=action, result=result, messages=self.messages, - active_system_prompt=self.active_system_prompt, - conversation_history=self.conversation_history, approx_tokens=self.approx_tokens, - compression_attempts=self.compression_attempts, - provider_overflow_recovery_pending=self.provider_overflow_recovery_pending, - is_context_length_error=self.is_context_length_error, - ) + def done(self, action: str, result: Optional[Dict[str, Any]] = None) -> OverflowVerdict: + self.action, self.result = action, result + return self def fail_turn( self, final_response: str, *, notices: tuple = (), log: Optional[tuple] = None, @@ -119,7 +111,7 @@ class _Recovery: if compression_exhausted: result["compression_exhausted"] = True result.update(extra) - return self.verdict("return", result) + return self.done("return", result) def count_attempt(self, *, payload_too_large: bool = False) -> Optional[OverflowVerdict]: """Bump ``compression_attempts``; the terminal verdict once the cap is exceeded.""" @@ -173,7 +165,7 @@ class _Recovery: if deferred is not None: self.compression_attempts -= 1 agent._persist_session(self.messages, self.conversation_history) - return self.verdict("return", deferred) + return self.done("return", deferred) if fail_on_timeout and context_compression_timed_out(agent): return self.fail_turn( _COMPRESSION_TIMEOUT_FINAL_RESPONSE, turn_exit_reason="context_compression_timeout" @@ -250,14 +242,14 @@ def _recover_payload_too_large(st: _Recovery, _retry: TurnRetryState) -> Overflo ) time.sleep(2) # Brief pause between compression retries _retry.restart_with_compressed_messages = True - return st.verdict("break") + return st.done("break") if agent._try_strip_image_parts_from_tool_messages(st.api_messages, remember_model=False): agent._buffer_status( "📐 Compression could not reduce the request further — " "removed retained vision payloads and retrying..." ) - return st.verdict("continue") + return st.done("continue") return st.fail_turn( "Request payload too large (413). Cannot compress further.", @@ -302,7 +294,7 @@ def _clamp_output_cap(st: _Recovery, _retry: TurnRetryState, available_out: int, "%sOutput-cap compression hit an error; retrying on max_tokens only.", agent.log_prefix ) _retry.restart_with_compressed_messages = True - return st.verdict("break") + return st.done("break") def _adopt_provider_context_limit(st: _Recovery, error_msg: str, old_ctx: int) -> Optional[int]: @@ -397,7 +389,7 @@ def _recover_context_length(st: _Recovery, _retry: TurnRetryState, error_msg: st # count alone doesn't prove system/tool-inclusive pressure fell. st.provider_overflow_recovery_pending = True _retry.restart_with_compressed_messages = True - return st.verdict("break") + return st.done("break") # Can't compress further and already at minimum tier. return st.fail_turn( @@ -454,4 +446,4 @@ def recover_from_overflow( ) if st.is_context_length_error: return _recover_context_length(st, _retry, error_msg) - return st.verdict("fallthrough") + return st.done("fallthrough") diff --git a/agent/turn_truncation.py b/agent/turn_truncation.py index b9eadaaa80..38ff9a7b90 100644 --- a/agent/turn_truncation.py +++ b/agent/turn_truncation.py @@ -99,10 +99,10 @@ class TruncationVerdict: compression_attempts: int -@dataclass -class _Trunc: - """Mutable working state for the truncation phases; rebound loop locals are handed - back through ``verdict()``.""" +@dataclass(kw_only=True) +class _Trunc(TruncationVerdict): + """Working state for the truncation phases — the verdict itself, plus the read-only + call context; phases mutate the loop-local fields and ``done()`` stamps the action.""" agent: Any response: Any @@ -111,21 +111,12 @@ class _Trunc: api_call_count: int effective_task_id: Any current_turn_user_idx: Any - messages: List[Dict[str, Any]] - length_continue_retries: int - truncated_response_parts: List[str] - truncated_tool_call_retries: int - retry_count: int - compression_attempts: int + action: str = "fallthrough" + result: Optional[Dict[str, Any]] = None - def verdict(self, action: str, result: Optional[Dict[str, Any]] = None) -> TruncationVerdict: - return TruncationVerdict( - action=action, result=result, messages=self.messages, - length_continue_retries=self.length_continue_retries, - truncated_response_parts=self.truncated_response_parts, - truncated_tool_call_retries=self.truncated_tool_call_retries, - retry_count=self.retry_count, compression_attempts=self.compression_attempts, - ) + def done(self, action: str, result: Optional[Dict[str, Any]] = None) -> TruncationVerdict: + self.action, self.result = action, result + return self def end_turn( self, final_response: str, error: Optional[str] = None, *, @@ -137,7 +128,7 @@ class _Trunc: if cleanup: agent._cleanup_task_resources(self.effective_task_id) agent._persist_session(self.messages, self.conversation_history) - return self.verdict("return", partial_result( + return self.done("return", partial_result( self.messages if result_messages is None else result_messages, self.api_call_count, final_response, error, failed=failed, )) @@ -194,7 +185,7 @@ def _content_filter_fallback(st: _Trunc, _retry: TurnRetryState) -> Optional[Tru st.compression_attempts = 0 _retry.primary_recovery_attempted = False _retry.restart_with_rebuilt_messages = True - return st.verdict("break") + return st.done("break") agent._vprint( f"{agent.log_prefix}⚠️ No fallback provider " f"configured — retrying with same provider " @@ -244,7 +235,7 @@ def _continue_text(st: _Trunc, _retry: TurnRetryState, assistant_message: Any) - }) agent._session_messages = messages _retry.restart_with_length_continuation = True - return st.verdict("break") + return st.done("break") partial_response = agent._strip_think_blocks(_join_truncated_parts(st.truncated_response_parts)).strip() # The one-shot reasoning-off override must not leak into the next turn. @@ -293,7 +284,7 @@ def _retry_truncated_tool_call(st: _Trunc, api_kwargs: Any) -> TruncationVerdict if _tc_requested_cap is not None: _tc_boost = max(_tc_boost, _tc_requested_cap) agent._ephemeral_max_output_tokens = min(_tc_boost, max(32768, _tc_requested_cap or 0)) - return st.verdict("continue") # don't append the broken response + return st.done("continue") # don't append the broken response agent._flush_status_buffer() if st.is_stub: agent._vprint(