fix(agent): honor provider reset for fallback cooldown
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
92
tests/agent/test_provider_reset_cooldown.py
Normal file
92
tests/agent/test_provider_reset_cooldown.py
Normal 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,
|
||||
)
|
||||
Reference in New Issue
Block a user