refactor(agent/turn): working-state carriers subclass their Verdict (drop duplicate field copies)

This commit is contained in:
Teknium
2026-09-02 18:49:52 -07:00
parent 498c7a1c08
commit 761cff3221
2 changed files with 30 additions and 47 deletions

View File

@@ -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")

View File

@@ -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(