fix(auth): stop stale-key credential recovery loops

Track the selected credential by stable pool entry ID so token refreshes and shared cursor movement cannot detach failures from the entry that issued them. Stop unmatched single-entry pools from reporting a no-op rotation as successful recovery.

Co-authored-by: Maxim Esipov <maksesipov@gmail.com>
This commit is contained in:
Gille 2026-07-23 15:54:00 -06:00 committed by kshitij
parent d51bd5fdc6
commit 73c4b5a045
6 changed files with 234 additions and 46 deletions

View file

@ -1341,6 +1341,21 @@ def init_agent(
print("⚠️ Warning: API key appears invalid or missing")
except Exception as e:
raise RuntimeError(f"Failed to initialize OpenAI client: {e}")
# Keep a stable identity for the pool entry that supplied this runtime.
# OAuth refreshes can replace the runtime token before a failed request is
# recovered, so the mutable API-key value alone cannot reliably attribute
# the failure to its source entry.
agent._credential_pool_entry_id = None
if agent._credential_pool is not None:
try:
agent._credential_pool_entry_id = (
agent._credential_pool.entry_id_for_api_key(
getattr(agent, "api_key", None)
)
)
except Exception:
agent._credential_pool_entry_id = None
# Provider fallback chain — ordered list of backup providers tried
# when the primary is exhausted (rate-limit, overload, connection

View file

@ -934,10 +934,30 @@ def recover_with_credential_pool(
# failing entry exactly; fall back to current()'s key only when the agent
# carries no key at all.
_api_key_hint = getattr(agent, "api_key", None) or None
_raw_credential_id = getattr(agent, "_credential_pool_entry_id", None)
_credential_id = (
_raw_credential_id
if isinstance(_raw_credential_id, str) and _raw_credential_id
else None
)
if not _api_key_hint:
_cur = pool.current()
if _cur:
_api_key_hint = getattr(_cur, "runtime_api_key", None)
if not _credential_id:
_current_id = getattr(_cur, "id", None)
if isinstance(_current_id, str) and _current_id:
_credential_id = _current_id
def _rotate_failed_credential(rotate_status: int):
kwargs = {
"status_code": rotate_status,
"error_context": error_context,
"api_key_hint": _api_key_hint,
}
if _credential_id:
kwargs["credential_id"] = _credential_id
return pool.mark_exhausted_and_rotate(**kwargs)
effective_reason = classified_reason
if effective_reason is None:
@ -972,11 +992,7 @@ def recover_with_credential_pool(
# Runtime credentials can be resolved by a separate pool instance,
# leaving this recovery pool without ``current_id``. Match the key
# that actually failed instead of quarantining a different account.
next_entry = pool.mark_exhausted_and_rotate(
status_code=rotate_status,
error_context=error_context,
api_key_hint=_api_key_hint,
)
next_entry = _rotate_failed_credential(rotate_status)
if next_entry is not None:
_ra().logger.info(
"Credential %s (billing) — rotated to pool entry %s",
@ -995,8 +1011,13 @@ def recover_with_credential_pool(
# Prefer the entry matching the failing key over the shared current()
# pointer, for the same attribution reason as above.
current_entry = None
if _api_key_hint:
if _credential_id:
current_entry = next(
(e for e in pool.entries() if e.id == _credential_id),
None,
)
if _api_key_hint:
current_entry = current_entry or next(
(e for e in pool.entries() if e.runtime_api_key == _api_key_hint),
None,
)
@ -1009,11 +1030,7 @@ def recover_with_credential_pool(
current_last_status,
)
rotate_status = status_code if status_code is not None else 429
next_entry = pool.mark_exhausted_and_rotate(
status_code=rotate_status,
error_context=error_context,
api_key_hint=_api_key_hint,
)
next_entry = _rotate_failed_credential(rotate_status)
if next_entry is not None:
_ra().logger.info(
"Credential %s (rate limit, pre-exhausted) — rotated to pool entry %s",
@ -1037,11 +1054,7 @@ def recover_with_credential_pool(
if not has_retried_429 and not usage_limit_reached:
return False, True
rotate_status = status_code if status_code is not None else 429
next_entry = pool.mark_exhausted_and_rotate(
status_code=rotate_status,
error_context=error_context,
api_key_hint=_api_key_hint,
)
next_entry = _rotate_failed_credential(rotate_status)
if next_entry is not None:
_ra().logger.info(
"Credential %s (rate limit) — rotated to pool entry %s",
@ -1113,7 +1126,10 @@ def recover_with_credential_pool(
# the shared pointer can reference a different, healthy entry, and
# refreshing it would consume that entry's single-use refresh token
# (or mark it exhausted on failure) for a failure it never had.
refreshed = pool.try_refresh_matching(api_key_hint=_api_key_hint)
refresh_kwargs = {"api_key_hint": _api_key_hint}
if _credential_id:
refresh_kwargs["credential_id"] = _credential_id
refreshed = pool.try_refresh_matching(**refresh_kwargs)
if refreshed is not None:
# ``try_refresh_matching()`` re-mints a fresh OAuth token and reports
# success even when the upstream keeps rejecting it — a single-entry
@ -1145,11 +1161,7 @@ def recover_with_credential_pool(
# Refresh failed — rotate to next credential instead of giving up.
# The failed entry is already marked exhausted by the refresh attempt.
rotate_status = status_code if status_code is not None else 401
next_entry = pool.mark_exhausted_and_rotate(
status_code=rotate_status,
error_context=error_context,
api_key_hint=_api_key_hint,
)
next_entry = _rotate_failed_credential(rotate_status)
if next_entry is not None:
_ra().logger.info(
"Credential %s (auth refresh failed) — rotated to pool entry %s",
@ -1460,6 +1472,7 @@ def restore_primary_runtime(agent) -> bool:
pool_matches_primary = False
if pool is not None and pool_provider and not pool_matches_primary:
agent._credential_pool = None
agent._credential_pool_entry_id = None
try:
from agent.credential_pool import load_pool
@ -1479,6 +1492,7 @@ def restore_primary_runtime(agent) -> bool:
# the pool for its current best entry and swap the live credential in.
# When the pool is absent, empty, or the entry has no usable key, we
# keep the snapshot key (the existing behavior). Fixes #25205.
agent._credential_pool_entry_id = None
pool = getattr(agent, "_credential_pool", None)
if pool is not None and pool.has_available():
entry = pool.select()
@ -2075,6 +2089,9 @@ def switch_model(agent, new_model, new_provider, api_key='', base_url='', api_mo
# restore the original pool (issue #52727: pool reload is part of this
# switch and must be reversible on rollback).
_snapshot["_credential_pool"] = getattr(agent, "_credential_pool", _MISSING)
_snapshot["_credential_pool_entry_id"] = getattr(
agent, "_credential_pool_entry_id", _MISSING
)
try:
# Clear the per-config context_length override so the new model's
@ -2131,6 +2148,7 @@ def switch_model(agent, new_model, new_provider, api_key='', base_url='', api_mo
# A pool bound to the old provider is worse than no pool: the
# recovery guard rejects it and every later 401/429 skips rotation.
agent._credential_pool = None
agent._credential_pool_entry_id = None
try:
from agent.credential_pool import load_pool
agent._credential_pool = load_pool(new_provider)
@ -2140,7 +2158,6 @@ def switch_model(agent, new_model, new_provider, api_key='', base_url='', api_mo
"continuing without pool rotation this turn",
new_provider, _pool_exc,
)
# ── Build new client ──
if (new_provider or "").strip().lower() == "moa":
from agent.moa_loop import build_moa_facade
@ -2236,6 +2253,17 @@ def switch_model(agent, new_model, new_provider, api_key='', base_url='', api_mo
reason="switch_model",
shared=True,
)
pool = getattr(agent, "_credential_pool", None)
if pool is not None:
try:
agent._credential_pool_entry_id = pool.entry_id_for_api_key(
getattr(agent, "api_key", None)
)
except Exception:
agent._credential_pool_entry_id = None
else:
agent._credential_pool_entry_id = None
except Exception:
# Rollback every mutated field to the pre-swap snapshot so the agent
# is left consistent (old model + old provider + old client) and the

View file

@ -1739,6 +1739,7 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool
fb_provider, fb_model, _pool_provider,
)
agent._credential_pool = None
agent._credential_pool_entry_id = None
if getattr(agent, "_credential_pool", None) is None:
try:
from agent.credential_pool import load_pool
@ -1801,6 +1802,17 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool
# not only after a later credential-rotation rebuild.
agent._replace_primary_openai_client(reason="fallback_timeout_apply")
fallback_pool = getattr(agent, "_credential_pool", None)
if fallback_pool is not None:
try:
agent._credential_pool_entry_id = (
fallback_pool.entry_id_for_api_key(agent.api_key)
)
except Exception:
agent._credential_pool_entry_id = None
else:
agent._credential_pool_entry_id = None
# Re-evaluate prompt caching for the new provider/model
agent._use_prompt_caching, agent._use_native_cache_layout = (
agent._anthropic_prompt_cache_policy(

View file

@ -622,6 +622,28 @@ class CredentialPool:
with self._lock:
return self._current_unlocked()
def entry_id_for_api_key(self, api_key_hint: Any = None) -> Optional[str]:
"""Return the stable id for the runtime credential in use.
Prefer the current selection when it still supplies ``api_key_hint``.
If the cursor was cleared, fall back to an unambiguous key match.
"""
with self._lock:
current = self._current_unlocked()
if current is not None and (
api_key_hint is None
or current.runtime_api_key == api_key_hint
):
return current.id
if api_key_hint is None:
return None
matches = [
entry
for entry in self._entries
if entry.runtime_api_key == api_key_hint
]
return matches[0].id if len(matches) == 1 else None
def _replace_entry(self, old: PooledCredential, new: PooledCredential) -> None:
"""Swap an entry in-place by id, preserving sort order."""
for idx, entry in enumerate(self._entries):
@ -1763,10 +1785,17 @@ class CredentialPool:
status_code: Optional[int],
error_context: Optional[Dict[str, Any]] = None,
api_key_hint: Optional[str] = None,
credential_id: Optional[str] = None,
) -> Optional[PooledCredential]:
with self._lock:
entry = None
if api_key_hint:
identity_supplied = bool(credential_id or api_key_hint)
if credential_id:
entry = next(
(e for e in self._entries if e.id == credential_id),
None,
)
if entry is None and api_key_hint:
# Prefer the specific entry whose API key matches the one that
# actually failed. When this pool was freshly loaded from disk
# (another process already rotated), current() is None and
@ -1775,20 +1804,26 @@ class CredentialPool:
(e for e in self._entries if e.runtime_api_key == api_key_hint),
None,
)
if entry is None:
# The failed key is identifiable but matches no entry
# (rotated away, or a wrapper whose runtime key differs).
# Falling through to current()/_select_unlocked() would
# mark an INNOCENT healthy key exhausted for the full
# cooldown TTL. Don't guess — just hand back a fresh
# selection so the caller can retry.
logger.info(
"credential pool: failed key hint matched no %s entry; "
"rotating without marking any credential exhausted",
self.provider,
)
if entry is None and identity_supplied:
# The failed credential is identifiable but matches no entry
# (rotated away, or a wrapper whose runtime key differs).
# Falling through to current()/_select_unlocked() would mark an
# innocent healthy key exhausted for the full cooldown TTL.
logger.info(
"credential pool: failed credential identity matched no %s "
"entry; rotating without marking any credential exhausted",
self.provider,
)
self._current_id = None
next_entry = self._select_unlocked()
if next_entry is not None and len(self._available_entries()) == 1:
# A single-entry pool cannot rotate. Returning its only
# entry reports a successful recovery without changing
# the credential, so the caller retries the same 401
# indefinitely. Let fallback/error propagation proceed.
self._current_id = None
return self._select_unlocked()
return None
return next_entry
if entry is None:
entry = self._current_unlocked() or self._select_unlocked()
if entry is None:
@ -1806,12 +1841,13 @@ class CredentialPool:
# disconnects (a ~2.5min hang with no error surfaced to the user).
# Mark every entry sharing the failed key so the pool can reach the
# "no available entries" state and let the error propagate.
if api_key_hint:
failed_runtime_key = getattr(entry, "runtime_api_key", None)
if identity_supplied and failed_runtime_key:
siblings_marked = False
for sibling in self._entries:
if sibling.id == entry.id:
continue
if sibling.runtime_api_key == api_key_hint:
if sibling.runtime_api_key == failed_runtime_key:
self._mark_exhausted(
sibling, status_code, error_context, persist=False
)
@ -1885,9 +1921,11 @@ class CredentialPool:
return self._try_refresh_current_unlocked()
def try_refresh_matching(
self, api_key_hint: Optional[str] = None
self,
api_key_hint: Optional[str] = None,
credential_id: Optional[str] = None,
) -> Optional[PooledCredential]:
"""Force-refresh the entry that supplied ``api_key_hint``.
"""Force-refresh the entry that supplied the failed request.
Direct provider integrations may reload the pool after a request has
already failed, so they cannot rely on ``current_id`` identifying the
@ -1897,17 +1935,29 @@ class CredentialPool:
"""
with self._lock:
entry = None
if api_key_hint:
if credential_id:
entry = next(
(
candidate
for candidate in self._entries
if candidate.runtime_api_key == api_key_hint
if candidate.id == credential_id
),
None,
)
else:
entry = self._current_unlocked() or self._select_unlocked(refresh=False)
if entry is None:
if api_key_hint:
entry = next(
(
candidate
for candidate in self._entries
if candidate.runtime_api_key == api_key_hint
),
None,
)
else:
entry = self._current_unlocked() or self._select_unlocked(
refresh=False
)
if entry is None:
return None
self._current_id = entry.id

View file

@ -4963,6 +4963,7 @@ class AIAgent:
def _swap_credential(self, entry) -> None:
runtime_key = getattr(entry, "runtime_api_key", None) or getattr(entry, "access_token", "")
runtime_base = getattr(entry, "runtime_base_url", None) or getattr(entry, "base_url", None) or self.base_url
self._credential_pool_entry_id = getattr(entry, "id", None)
from hermes_cli.route_identity import normalize_route_base_url
route_changed = normalize_route_base_url(self.base_url) != normalize_route_base_url(

View file

@ -405,11 +405,12 @@ class TestFailureAttribution:
entry.update(overrides)
return entry
def _agent(self, pool, failing_key):
def _agent(self, pool, failing_key, credential_id=None):
return SimpleNamespace(
provider="anthropic",
api_key=failing_key,
_credential_pool=pool,
_credential_pool_entry_id=credential_id,
_swap_credential=MagicMock(),
)
@ -521,3 +522,84 @@ class TestFailureAttribution:
assert statuses["cred-0"] != "exhausted"
swapped = agent._swap_credential.call_args[0][0]
assert swapped.id == "cred-0"
def test_auth_refresh_uses_stable_id_after_runtime_key_changes(
self, tmp_path, monkeypatch
):
"""A refreshed pool token must not detach the failed request from the
entry that supplied its now-stale runtime key."""
pool = self._make_pool(
tmp_path, monkeypatch,
[self._entry(0, "new-runtime-key")],
)
selected = pool.select()
assert selected.id == "cred-0"
assert pool.entry_id_for_api_key("new-runtime-key") == "cred-0"
agent = self._agent(
pool,
failing_key="stale-runtime-key",
credential_id="cred-0",
)
agent._is_entitlement_failure = MagicMock(return_value=False)
from agent.agent_runtime_helpers import recover_with_credential_pool
recovered, _ = recover_with_credential_pool(
agent, status_code=401, has_retried_429=False
)
assert recovered is False
assert self._statuses(pool)["cred-0"] == "exhausted"
agent._swap_credential.assert_not_called()
def test_unmatched_key_does_not_retry_only_pool_entry(
self, tmp_path, monkeypatch
):
"""Legacy agents without a stable id must stop when an unmatched key
has no different credential to rotate to."""
pool = self._make_pool(
tmp_path, monkeypatch,
[self._entry(0, "pool-runtime-key")],
)
agent = self._agent(pool, failing_key="wrapper-runtime-key")
agent._is_entitlement_failure = MagicMock(return_value=False)
from agent.agent_runtime_helpers import recover_with_credential_pool
recovered, _ = recover_with_credential_pool(
agent, status_code=401, has_retried_429=False
)
assert recovered is False
assert self._statuses(pool)["cred-0"] != "exhausted"
agent._swap_credential.assert_not_called()
def test_stable_id_rotates_from_failed_entry_when_cursor_points_elsewhere(
self, tmp_path, monkeypatch
):
"""Stable identity wins over both a stale key and the shared cursor."""
pool = self._make_pool(
tmp_path, monkeypatch,
[self._entry(0, "key-a"), self._entry(1, "key-b-new")],
)
assert pool.select().id == "cred-0"
agent = self._agent(
pool,
failing_key="key-b-old",
credential_id="cred-1",
)
agent._is_entitlement_failure = MagicMock(return_value=False)
from agent.agent_runtime_helpers import recover_with_credential_pool
recovered, _ = recover_with_credential_pool(
agent, status_code=401, has_retried_429=False
)
assert recovered is True
statuses = self._statuses(pool)
assert statuses["cred-1"] == "exhausted"
assert statuses["cred-0"] != "exhausted"
swapped = agent._swap_credential.call_args[0][0]
assert swapped.id == "cred-0"