From 125bdff6014b0186c29e9f4e69ee6740acddaff5 Mon Sep 17 00:00:00 2001 From: fangliquanflq Date: Mon, 21 Sep 2026 04:22:46 +0800 Subject: [PATCH] fix(agent): honor provider reset for fallback cooldown --- agent/chat_completion_helpers.py | 4 +- agent/fallback_cooldown.py | 32 +++++-- agent/turn_api_error.py | 11 ++- agent/turn_recovery.py | 3 +- tests/agent/test_provider_reset_cooldown.py | 92 +++++++++++++++++++++ 5 files changed, 128 insertions(+), 14 deletions(-) create mode 100644 tests/agent/test_provider_reset_cooldown.py diff --git a/agent/chat_completion_helpers.py b/agent/chat_completion_helpers.py index 98b5019d7f..327e734b99 100644 --- a/agent/chat_completion_helpers.py +++ b/agent/chat_completion_helpers.py @@ -1997,12 +1997,12 @@ def _buffer_fallback_notice(agent, notice: str) -> None: agent._pending_fallback_notice = [str(pending), notice] if pending else [notice] -def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool: +def try_activate_fallback(agent, reason: "FailoverReason | None" = None, reset_at=None) -> bool: """Switch to the next fallback model/provider in the chain; False when exhausted. Swaps client, model slug and provider in place so the retry loop continues on the new backend; client construction goes through resolve_provider_client (no duplicated provider→key mappings).""" from agent.fallback_cooldown import _arm_rate_limit_cooldown - cooldown_seconds = _arm_rate_limit_cooldown(agent, reason) + cooldown_seconds = _arm_rate_limit_cooldown(agent, reason, reset_at=reset_at) while True: if agent._fallback_index >= len(agent._fallback_chain): return _fallback_chain_exhausted(agent, reason) diff --git a/agent/fallback_cooldown.py b/agent/fallback_cooldown.py index fcd62fbed0..785353dbcd 100644 --- a/agent/fallback_cooldown.py +++ b/agent/fallback_cooldown.py @@ -1,6 +1,7 @@ """Primary rate-limit cooldown arming and per-session model rejection markers, shared by the fallback walk (chat_completion_helpers) and restore_primary_runtime (agent_runtime_helpers).""" import logging +import math import time from agent.error_classifier import FailoverReason @@ -10,11 +11,17 @@ logger = logging.getLogger(__name__) _RATE_LIMIT_FAILOVER_REASONS = frozenset({FailoverReason.rate_limit, FailoverReason.billing, FailoverReason.upstream_rate_limit}) -def _arm_rate_limit_cooldown(agent, reason: "FailoverReason | None") -> int | None: - """Arm the primary's exponential cooldown (60s → 2m → ... → 4h cap) on CONSECUTIVE rate-limits; - restore_primary_runtime resets the counter. Only when leaving the primary: chain-switching from - an active fallback means the primary was not the 429 source, so its cooldown is left alone. - Return the armed cooldown in seconds, or None when no cooldown was armed.""" +def _arm_rate_limit_cooldown( + agent, reason: "FailoverReason | None", reset_at=None, +) -> int | None: + """Arm the primary cooldown until the provider reset, or use exponential backoff. + + ``reset_at`` is an absolute wall-clock timestamp while ``_rate_limited_until`` is monotonic; + convert through a duration so wall-clock epoch values never enter the monotonic comparison. + Missing, invalid, or expired provider resets retain the 60s → 2m → ... → 4h fallback. + Only arm when leaving the primary: chain-switching from an active fallback means the primary + was not the failing source. Return the armed cooldown in seconds, or None when not armed. + """ if reason not in _RATE_LIMIT_FAILOVER_REASONS: return None current_provider = (getattr(agent, "provider", "") or "").strip().lower() @@ -23,9 +30,20 @@ def _arm_rate_limit_cooldown(agent, reason: "FailoverReason | None") -> int | No return None backoff_count = getattr(agent, "_rate_limit_backoff_count", 0) agent._rate_limit_backoff_count = backoff_count + 1 - backoff_seconds = min(60 * (2 ** backoff_count), 14400) + from agent.credential_pool import _parse_absolute_timestamp + parsed_reset_at = _parse_absolute_timestamp(reset_at) + provider_delay = parsed_reset_at - time.time() if parsed_reset_at is not None else None + if provider_delay is not None and math.isfinite(provider_delay) and provider_delay > 0: + backoff_seconds = math.ceil(provider_delay) + source = "provider reset" + else: + backoff_seconds = min(60 * (2 ** backoff_count), 14400) + source = "exponential fallback" agent._rate_limited_until = time.monotonic() + backoff_seconds - logging.info("Rate-limit backoff level %d: cooldown %d s (%.1f min, backoff#%d)", backoff_count, backoff_seconds, backoff_seconds / 60, backoff_count + 1) + logging.info( + "Rate-limit backoff level %d: cooldown %d s (%.1f min, backoff#%d, %s)", + backoff_count, backoff_seconds, backoff_seconds / 60, backoff_count + 1, source, + ) return backoff_seconds diff --git a/agent/turn_api_error.py b/agent/turn_api_error.py index 9b2440b647..376098d44f 100644 --- a/agent/turn_api_error.py +++ b/agent/turn_api_error.py @@ -204,7 +204,8 @@ def handle_api_error( _ue = settle_unrecovered_error( agent, api_error=api_error, classified=classified, _retry=_retry, status_code=status_code, - error_msg=error_msg, is_context_length_error=is_context_length_error, + error_msg=error_msg, error_context=error_context, + is_context_length_error=is_context_length_error, is_rate_limited=is_rate_limited, _is_zai_coding_overload=_is_zai_coding_overload, _provider=_provider, _base=_base, _model=_model, messages=messages, api_messages=api_messages, api_kwargs=api_kwargs, active_system_prompt=active_system_prompt, @@ -250,7 +251,7 @@ def settle_unrecovered_error( is_context_length_error: Any, is_rate_limited: Any, _is_zai_coding_overload: Any, _provider: Any, _base: Any, _model: Any, messages: Any, api_messages: Any, api_kwargs: Any, active_system_prompt: Any, conversation_history: Any, approx_tokens: Any, retry_count: Any, - max_retries: Any, compression_attempts: Any, api_call_count: Any, + max_retries: Any, compression_attempts: Any, api_call_count: Any, error_context: Any = None, ) -> UnrecoveredErrorVerdict: """Decide the fate of an API error that every recovery chain declined: local validation / non-retryable client errors (Copilot stale-credential self-heal first, then fallback, then a @@ -331,7 +332,8 @@ def settle_unrecovered_error( if agent._has_pending_fallback(): _label = _NONRETRYABLE_LABELS.get(classified.reason, f"Non-retryable error (HTTP {status_code})") agent._buffer_diagnostic_status(f"⚠️ {_label} — trying fallback...") - if agent._try_activate_fallback(): + reset_at = error_context.get("reset_at") if isinstance(error_context, dict) else None + if agent._try_activate_fallback(reason=classified.reason, reset_at=reset_at): # Direct ``return _verdict("break")`` is load-bearing: the restart handler # re-runs the pre-API preflight against the fallback's context window. active_system_prompt = _arm_fallback_restart(agent, api_messages, active_system_prompt, _retry) @@ -360,7 +362,8 @@ def settle_unrecovered_error( return _verdict("continue") if agent._has_pending_fallback(): agent._buffer_diagnostic_status(f"⚠️ Max retries ({max_retries}) exhausted — trying fallback...") - if agent._try_activate_fallback(): + reset_at = error_context.get("reset_at") if isinstance(error_context, dict) else None + if agent._try_activate_fallback(reason=classified.reason, reset_at=reset_at): # Direct ``return _verdict("break")`` is load-bearing: the restart handler # re-runs the pre-API preflight against the fallback's context window. active_system_prompt = _arm_fallback_restart(agent, api_messages, active_system_prompt, _retry) diff --git a/agent/turn_recovery.py b/agent/turn_recovery.py index 2f86176540..a8c95540f8 100644 --- a/agent/turn_recovery.py +++ b/agent/turn_recovery.py @@ -1786,7 +1786,8 @@ def route_classified_error( ) if not pool_may_recover: agent._buffer_diagnostic_status(_eager_fallback_status(classified, _is_upstream, _is_transport_failure)) - if agent._try_activate_fallback(reason=classified.reason): + reset_at = error_context.get("reset_at") if isinstance(error_context, dict) else None + if agent._try_activate_fallback(reason=classified.reason, reset_at=reset_at): return _fallback_break() # A 401/403 surviving credential refresh means a broken credential or endpoint: diff --git a/tests/agent/test_provider_reset_cooldown.py b/tests/agent/test_provider_reset_cooldown.py new file mode 100644 index 0000000000..2f9cd05086 --- /dev/null +++ b/tests/agent/test_provider_reset_cooldown.py @@ -0,0 +1,92 @@ +"""Provider-declared reset windows govern primary fallback cooldowns (#117484).""" + +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest + +from agent.error_classifier import FailoverReason +from agent.fallback_cooldown import _arm_rate_limit_cooldown +from agent.turn_recovery import route_classified_error + + +@pytest.mark.parametrize( + ("reset_at", "expected_seconds"), + [ + (1_700_000_090.2, 91), + (1_700_274_291, 274_291), + (None, 60), + ("not-a-timestamp", 60), + (1_699_999_999, 60), + ], +) +def test_rate_limit_cooldown_prefers_only_valid_future_provider_resets( + reset_at, expected_seconds, +): + agent = SimpleNamespace( + provider="openrouter", + _primary_runtime={"provider": "openrouter"}, + _fallback_activated=False, + _rate_limit_backoff_count=0, + ) + with ( + patch("agent.fallback_cooldown.time.time", return_value=1_700_000_000), + patch("agent.fallback_cooldown.time.monotonic", return_value=500), + ): + armed = _arm_rate_limit_cooldown( + agent, FailoverReason.rate_limit, reset_at=reset_at, + ) + + assert armed == expected_seconds + assert agent._rate_limited_until == 500 + expected_seconds + + +def test_eager_rate_limit_fallback_forwards_extracted_reset_time(): + reset_at = 1_900_000_000 + agent = MagicMock() + agent.compression_enabled = True + agent.provider = "openrouter" + agent._fallback_index = 0 + agent._fallback_chain = [{"provider": "anthropic", "model": "claude"}] + agent._credential_pool = None + agent._try_activate_fallback.return_value = True + classified = SimpleNamespace(reason=FailoverReason.rate_limit) + retry = SimpleNamespace() + api_error = SimpleNamespace(status_code=429) + + with ( + patch( + "agent.conversation_loop._ra", + return_value=SimpleNamespace( + _pool_may_recover_from_rate_limit=lambda _pool: False, + ), + ), + patch("agent.conversation_loop._arm_fallback_restart", return_value="fallback prompt"), + ): + verdict = route_classified_error( + agent, + api_error, + classified, + retry, + error_msg="rate limited", + error_context={"reset_at": reset_at}, + recovered_with_pool=False, + base_url="https://openrouter.ai/api/v1", + model="primary", + messages=[], + api_messages=[], + system_message="system", + active_system_prompt="system", + conversation_history=[], + retry_count=1, + max_retries=3, + compression_attempts=0, + max_compression_attempts=2, + api_call_count=1, + effective_task_id=None, + ) + + assert verdict.action == "break" + agent._try_activate_fallback.assert_called_once_with( + reason=FailoverReason.rate_limit, reset_at=reset_at, + ) \ No newline at end of file