fix(agent): refund the API call when an empty response activates fallback (#77305)

Fallback activation after empty-response retries re-entered the outer loop
without refunding the provisional call, so a provider hop consumed an
iteration. Refund via _refund_api_call and carry the count back through
EmptyResponseVerdict -> FinalResponseVerdict -> s.api_call_count.

Salvaged from #88867.

Co-authored-by: RelaxJonh <92573950+RelaxJonh@users.noreply.github.com>
This commit is contained in:
Axl Ibiza, MBA
2026-09-24 16:57:40 +05:30
committed by kshitij
parent b7f91258b9
commit bc6808fbf5
3 changed files with 57 additions and 0 deletions

View File

@@ -40,6 +40,7 @@ class EmptyResponseVerdict:
turn_exit_reason: Any
active_system_prompt: Any
preflight_compression_blocked: bool
api_call_count: int
def _retry_empty(
@@ -154,6 +155,7 @@ def recover_empty_response(
action=action, result=result, final_response=final_response,
turn_exit_reason=_turn_exit_reason, active_system_prompt=active_system_prompt,
preflight_compression_blocked=_preflight_compression_blocked,
api_call_count=api_call_count,
)
# Partial stream recovery: content streamed before the connection died becomes the
@@ -279,6 +281,10 @@ def recover_empty_response(
# OUTER loop: `continue` re-runs preflight against the fallback's window;
# `break` would end the turn without calling the fallback.
_preflight_compression_blocked = False
# The fallback hop is a provider switch, not a model turn: refund the empty
# call so a mid-turn fallback doesn't eat the iteration budget (#77305).
from agent.turn_context_compaction import _refund_api_call
api_call_count = _refund_api_call(agent, api_call_count)
return _verdict("continue")
_turn_exit_reason = "empty_response_exhausted"

View File

@@ -40,6 +40,7 @@ class FinalResponseVerdict:
length_continue_retries: Any
_pending_verification_response: Any
_pending_verification_response_previewed: Any
api_call_count: Any
result: Optional[Dict[str, Any]] = None
@@ -70,6 +71,7 @@ def finish_text_response(
length_continue_retries=length_continue_retries,
_pending_verification_response=_pending_verification_response,
_pending_verification_response_previewed=_pending_verification_response_previewed,
api_call_count=api_call_count,
result=result,
)
@@ -119,6 +121,7 @@ def finish_text_response(
_turn_exit_reason = _ev.turn_exit_reason
active_system_prompt = _ev.active_system_prompt
_preflight_compression_blocked = _ev.preflight_compression_blocked
api_call_count = _ev.api_call_count
if _ev.action == "return":
return _verdict("return", _ev.result)
if _ev.action == "break":

View File

@@ -0,0 +1,48 @@
"""Empty-response fallback hop refunds its API call (#77305)."""
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from agent import turn_empty_response as ter
def _agent():
a = MagicMock()
a._current_streamed_assistant_text = ""
a._has_content_after_think_block.return_value = False
a._strip_think_blocks.side_effect = lambda t: t or ""
a._last_content_with_tools = None
a._thinking_prefill_retries = 2
a._fallback_chain = ["fallback"]
a._api_call_count = 5
return a
def _recover(agent):
msg = SimpleNamespace(reasoning=None, reasoning_content=None, reasoning_details=None)
with patch.object(ter, "_retry_empty", return_value=(None, None, False)), \
patch("agent.conversation_loop._sync_failover_system_message", return_value="sys"):
return ter.recover_empty_response(
agent, msg, None, "stop", final_response="", messages=[{"role": "user", "content": "hi"}],
api_messages=[], conversation_history=[], active_system_prompt="sys", api_call_count=5,
turn_exit_reason=None, preflight_compression_blocked=False,
)
def test_fallback_activation_refunds_call_and_budget():
agent = _agent()
agent._try_activate_fallback.return_value = True
v = _recover(agent)
assert v.action == "continue"
assert v.api_call_count == 4
assert agent._api_call_count == 4
agent.iteration_budget.refund.assert_called_once()
def test_exhausted_without_fallback_keeps_count():
agent = _agent()
agent._try_activate_fallback.return_value = False
with patch.object(ter, "_terminal_empty", return_value="(empty)"):
v = _recover(agent)
assert v.action == "break"
assert v.api_call_count == 5
agent.iteration_budget.refund.assert_not_called()