fix(agent): honor provider reset for fallback cooldown

This commit is contained in:
fangliquanflq
2026-09-21 04:22:46 +08:00
committed by Teknium
parent 348568ddaf
commit 125bdff601
5 changed files with 128 additions and 14 deletions

View File

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

View File

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

View File

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

View File

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

View File

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