refactor(agent/turn): working-state carriers subclass their Verdict (drop duplicate field copies)
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user