fix(auth): rotate credentials for named custom providers after 401/429

Salvage of #93214 (5 commits squashed onto current main; agent_runtime_helpers.py
diverged since the PR base and was 3-way reapplied). The credential-rotation
guard in recover_with_credential_pool and both restore_primary_runtime paths
only tolerated the custom-naming split when the agent carried the literal label
'custom', so a named custom provider (agent.provider='gemini-no-filter', pool
'custom:gemini-no-filter') tripped the mismatch guard and skipped rotation on
every 401/429. Now all three guard sites use the canonical
credential_pool_matches_provider boundary predicate + resolve_runtime_pool_key,
which recognizes configured named-custom aliases and validates endpoints.

Fixes #93188.
This commit is contained in:
fangliquanflq
2026-08-23 19:28:14 -07:00
committed by Teknium
parent 030edf9774
commit 37411f349a
5 changed files with 452 additions and 93 deletions
+66 -67
View File
@@ -42,7 +42,11 @@ from agent.message_sanitization import (
from agent.prompt_builder import format_steer_marker
from agent.tool_dispatch_helpers import _trajectory_normalize_msg, make_tool_result_message
from agent.trajectory import convert_scratchpad_to_think
from agent.credential_pool import STATUS_EXHAUSTED, credential_pool_matches_provider
from agent.credential_pool import (
STATUS_EXHAUSTED,
credential_pool_matches_provider,
resolve_runtime_pool_key,
)
from agent.error_classifier import FailoverReason
from agent.turn_context import drop_stale_api_content
from utils import base_url_host_matches, base_url_hostname, env_var_enabled, atomic_json_write
@@ -1001,28 +1005,15 @@ def recover_with_credential_pool(
# because swapping the pool's credentials would set base_url/api_key
# without fixing the empty provider field, leaving the agent in a
# corrupted state (provider="" model="").
if pool_provider and current_provider != pool_provider:
# Custom endpoints use two naming conventions for the SAME provider:
# the agent carries the generic ``custom`` label while the pool is
# keyed ``custom:<name>`` (see CUSTOM_POOL_PREFIX). A literal string
# compare treats them as a mismatch and skips recovery for every
# custom-provider user — 401s/429s then burn the full retry cycle
# with no rotation or refresh. Accept the pair as matching only when
# the agent's CURRENT base_url actually resolves to this pool key,
# so a fallback provider (or a different custom endpoint) still
# triggers the guard.
_custom_match = False
if current_provider == "custom" and pool_provider.startswith("custom:"):
try:
from agent.credential_pool import get_custom_provider_pool_key
_agent_base = (getattr(agent, "base_url", "") or "").strip()
_custom_match = bool(_agent_base) and (
(get_custom_provider_pool_key(_agent_base) or "").strip().lower()
== pool_provider
)
except Exception:
_custom_match = False
if not _custom_match:
if pool_provider:
# Use the same fail-closed boundary predicate as runtime binding. This
# recognizes configured named-custom aliases, validates endpoints even
# for exact custom:* identities, and preserves fallback isolation.
if not credential_pool_matches_provider(
pool,
current_provider,
base_url=getattr(agent, "base_url", None),
):
_ra().logger.warning(
"Credential pool provider mismatch: pool=%s, agent=%s — "
"skipping pool mutation to avoid cross-provider contamination",
@@ -1545,22 +1536,39 @@ def restore_primary_runtime(agent) -> bool:
# pool-rebind block below via ``prefetched_primary_pool`` so the load
# happens at most once per restore.
prefetched_primary_pool = None
primary_pool_prefetched = False
try:
primary_provider = str(
(agent._primary_runtime or {}).get("provider") or ""
).strip().lower()
primary_runtime_base_url = str(
(agent._primary_runtime or {}).get("base_url") or ""
)
primary_pool_key = resolve_runtime_pool_key(
primary_provider,
primary_runtime_base_url,
)
pool = getattr(agent, "_credential_pool", None)
if not credential_pool_matches_provider(
pool,
primary_provider,
base_url=str((agent._primary_runtime or {}).get("base_url") or ""),
base_url=primary_runtime_base_url,
):
from agent.credential_pool import load_pool
prefetched_primary_pool = (
load_pool(primary_provider) if primary_provider else None
load_pool(primary_pool_key) if primary_pool_key else None
)
pool = prefetched_primary_pool
primary_pool_prefetched = True
if prefetched_primary_pool is not None and credential_pool_matches_provider(
prefetched_primary_pool,
primary_provider,
base_url=primary_runtime_base_url,
):
pool = prefetched_primary_pool
else:
prefetched_primary_pool = None
pool = None
next_at = getattr(pool, "next_available_at", lambda: None)()
if next_at is not None and next_at > time.time():
if not getattr(agent, "_restore_wait_logged", False):
@@ -1655,34 +1663,44 @@ def restore_primary_runtime(agent) -> bool:
# and disables credential rotation. Reload the primary pool first; if
# auth storage is temporarily unreadable, clear the mismatched pool.
primary_provider = str(rt.get("provider") or "").strip().lower()
primary_runtime_base_url = str(rt.get("base_url") or "")
primary_pool_key = resolve_runtime_pool_key(
primary_provider,
primary_runtime_base_url,
)
pool = getattr(agent, "_credential_pool", None)
pool_provider = str(getattr(pool, "provider", "") or "").strip().lower()
pool_matches_primary = pool_provider == primary_provider
if (
primary_provider == "custom"
and pool_provider.startswith("custom:")
):
try:
from agent.credential_pool import get_custom_provider_pool_key
primary_key = (
get_custom_provider_pool_key(str(rt.get("base_url") or "")) or ""
).strip().lower()
pool_matches_primary = bool(primary_key) and primary_key == pool_provider
except Exception:
pool_matches_primary = False
pool_matches_primary = credential_pool_matches_provider(
pool,
primary_provider,
base_url=primary_runtime_base_url,
)
if pool is not None and pool_provider and not pool_matches_primary:
agent._credential_pool = None
agent._credential_pool_entry_id = None
try:
if prefetched_primary_pool is not None:
if primary_pool_prefetched:
# Reuse the pool the reset-aware gate already loaded for
# this restore — avoids a second disk read of auth.json.
agent._credential_pool = prefetched_primary_pool
if (
prefetched_primary_pool is not None
and credential_pool_matches_provider(
prefetched_primary_pool,
primary_provider,
base_url=primary_runtime_base_url,
)
):
agent._credential_pool = prefetched_primary_pool
else:
from agent.credential_pool import load_pool
agent._credential_pool = load_pool(primary_provider)
loaded_pool = load_pool(primary_pool_key)
if loaded_pool is not None and credential_pool_matches_provider(
loaded_pool,
primary_provider,
base_url=primary_runtime_base_url,
):
agent._credential_pool = loaded_pool
except Exception as exc:
logger.warning(
"Restore could not reload primary credential pool for %s: %s",
@@ -1704,30 +1722,11 @@ def restore_primary_runtime(agent) -> bool:
entry = pool.select()
if entry is not None:
entry_provider = str(getattr(entry, "provider", "") or "").strip().lower()
entry_matches_primary = entry_provider == primary_provider
# Custom endpoints all carry the generic ``custom`` provider on
# the agent while the pool entry is keyed ``custom:<name>`` (see
# CUSTOM_POOL_PREFIX). Resolve the primary's base_url to its
# ``custom:<name>`` key via the canonical helper and compare
# against the entry's key — this mirrors the sibling guard in
# ``recover_with_credential_pool`` (see above) and correctly
# disambiguates multiple custom providers that share one gateway
# base_url. Fixes #56885.
from agent.credential_pool import CUSTOM_POOL_PREFIX
if (
primary_provider == "custom"
and entry_provider.startswith(CUSTOM_POOL_PREFIX)
):
entry_matches_primary = False
try:
from agent.credential_pool import get_custom_provider_pool_key
primary_base_url = str(rt.get("base_url") or "").strip()
primary_key = (
get_custom_provider_pool_key(primary_base_url) or ""
).strip().lower()
entry_matches_primary = bool(primary_key) and primary_key == entry_provider
except Exception:
entry_matches_primary = False
entry_matches_primary = credential_pool_matches_provider(
entry,
primary_provider,
base_url=primary_runtime_base_url,
)
entry_key = (
getattr(entry, "runtime_api_key", None)
+81 -18
View File
@@ -458,20 +458,17 @@ def _normalize_custom_pool_name(name: str) -> str:
def _iter_custom_providers(config: Optional[dict] = None):
"""Yield (normalized_name, entry_dict) for each valid custom_providers entry."""
"""Yield normalized entries from the merged custom-provider config view."""
if config is None:
config = _load_config_safe()
if config is None:
return
custom_providers = config.get("custom_providers")
if not isinstance(custom_providers, list):
# Fall back to the v12+ providers dict via the compatibility layer
try:
from hermes_cli.config import get_compatible_custom_providers
try:
from hermes_cli.config import get_compatible_custom_providers
custom_providers = get_compatible_custom_providers(config)
except Exception:
return
custom_providers = get_compatible_custom_providers(config)
except Exception:
return
if not custom_providers:
return
for entry in custom_providers:
@@ -559,10 +556,11 @@ def credential_pool_matches_provider(
) -> bool:
"""Return whether a pool belongs to the requested runtime provider.
Named custom endpoints intentionally use two identities: the live agent is
``custom`` while its pool is keyed ``custom:<name>``. Accept that pair only
when the runtime base URL resolves to the exact same custom pool key.
Empty string identities fail closed. Legacy pool adapters without a
Named custom endpoints may use three identities: the live agent can retain
the configured name/provider key, newer runtime paths normalize it to
``custom``, and the pool is keyed ``custom:<name>``. Accept those aliases
only when the runtime endpoint belongs to the same configured custom
provider. Empty identities fail closed. Legacy pool adapters without a
``provider`` attribute remain compatible; production pools are scoped.
"""
raw_pool_provider = getattr(pool_or_provider, "provider", None)
@@ -578,15 +576,80 @@ def credential_pool_matches_provider(
provider_norm = str(provider or "").strip().lower()
if not pool_provider or not provider_norm:
return False
if pool_provider == provider_norm:
return True
if provider_norm != "custom" or not pool_provider.startswith(CUSTOM_POOL_PREFIX):
if not pool_provider.startswith(CUSTOM_POOL_PREFIX):
return pool_provider == provider_norm
if provider_norm == "custom":
try:
matched_pool = get_custom_provider_pool_key(base_url or "")
except Exception:
return False
return str(matched_pool or "").strip().lower() == pool_provider
runtime_url = str(base_url or "").strip().rstrip("/")
if not runtime_url:
return False
try:
matched_pool = get_custom_provider_pool_key(base_url or "")
for normalized_name, entry in _iter_custom_providers():
if f"{CUSTOM_POOL_PREFIX}{normalized_name}" != pool_provider:
continue
aliases = {normalized_name}
for value in (entry.get("name"), entry.get("provider_key")):
alias = _normalize_custom_pool_name(str(value or ""))
if alias:
aliases.add(alias)
if alias.startswith(CUSTOM_POOL_PREFIX):
aliases.add(alias[len(CUSTOM_POOL_PREFIX):])
configured_url = str(entry.get("base_url") or "").strip().rstrip("/")
runtime_aliases = {_normalize_custom_pool_name(provider_norm)}
if provider_norm.startswith(CUSTOM_POOL_PREFIX):
runtime_aliases.add(
_normalize_custom_pool_name(
provider_norm[len(CUSTOM_POOL_PREFIX):]
)
)
return bool(runtime_aliases & aliases) and runtime_url == configured_url
except Exception:
return False
return str(matched_pool or "").strip().lower() == pool_provider
return False
def resolve_runtime_pool_key(provider: Optional[str], base_url: Optional[str]) -> str:
"""Resolve the credential-pool key for a runtime provider identity.
Named custom runtimes retain their configured alias while their pool is
stored under ``custom:<name>``. Return that scoped key only when the
canonical provider/endpoint boundary accepts it; otherwise preserve the
normalized runtime identity so callers fail closed.
"""
provider_norm = str(provider or "").strip().lower()
if not provider_norm:
return ""
try:
if provider_norm == "custom":
candidate = get_custom_provider_pool_key(base_url)
if candidate and credential_pool_matches_provider(
candidate,
provider_norm,
base_url=base_url,
):
return str(candidate).strip().lower()
else:
# Named and exact custom runtimes are keyed by provider identity,
# while auth storage remains keyed by display name. Search the
# configured candidates by identity before considering endpoint;
# this prevents a sibling sharing the URL from lending its pool.
for normalized_name, _entry in _iter_custom_providers():
candidate = f"{CUSTOM_POOL_PREFIX}{normalized_name}"
if credential_pool_matches_provider(
candidate,
provider_norm,
base_url=base_url,
):
return candidate
except Exception:
pass
return provider_norm
DEFAULT_MAX_CONCURRENT_PER_CREDENTIAL = 1
@@ -3,7 +3,10 @@
from types import SimpleNamespace
from unittest.mock import patch
from agent.credential_pool import credential_pool_matches_provider
from agent.credential_pool import (
credential_pool_matches_provider,
resolve_runtime_pool_key,
)
from hermes_cli import runtime_provider as rp
@@ -26,6 +29,127 @@ def test_custom_pool_match_is_scoped_by_endpoint():
)
def test_named_custom_pool_match_requires_configured_identity_and_endpoint():
configured = [
(
"gemini-display",
{
"name": "Gemini Display",
"provider_key": "gemini-no-filter",
"base_url": "https://generativelanguage.googleapis.com/v1beta/",
},
)
]
with patch("agent.credential_pool._iter_custom_providers", return_value=configured):
assert credential_pool_matches_provider(
"custom:gemini-display",
"gemini-no-filter",
base_url="https://generativelanguage.googleapis.com/v1beta",
)
assert credential_pool_matches_provider(
"custom:gemini-display",
"custom:gemini-no-filter",
base_url="https://generativelanguage.googleapis.com/v1beta",
)
assert not credential_pool_matches_provider(
"custom:gemini-display",
"gemini-no-filter",
base_url="https://fallback.example/v1",
)
assert not credential_pool_matches_provider(
"custom:gemini-display",
"custom:gemini-no-filter",
base_url="https://fallback.example/v1",
)
assert not credential_pool_matches_provider(
"custom:gemini-display",
"other-provider",
base_url="https://generativelanguage.googleapis.com/v1beta",
)
def test_runtime_pool_key_resolves_all_custom_runtime_identities():
endpoint = "https://generativelanguage.googleapis.com/v1beta"
configured = [
(
"sibling-display",
{
"name": "Sibling Display",
"provider_key": "sibling-provider",
"base_url": endpoint,
},
),
(
"gemini-display",
{
"name": "Gemini Display",
"provider_key": "gemini-no-filter",
"base_url": endpoint,
},
)
]
with patch("agent.credential_pool._iter_custom_providers", return_value=configured):
assert resolve_runtime_pool_key("custom", endpoint) == "custom:sibling-display"
assert (
resolve_runtime_pool_key("gemini-no-filter", endpoint)
== "custom:gemini-display"
)
assert (
resolve_runtime_pool_key("custom:gemini-no-filter", endpoint)
== "custom:gemini-display"
)
assert (
resolve_runtime_pool_key(
"gemini-no-filter",
"https://fallback.example/v1",
)
== "gemini-no-filter"
)
def test_runtime_pool_key_resolves_modern_provider_in_mixed_config():
endpoint = "https://generativelanguage.googleapis.com/v1beta"
config = {
"custom_providers": [
{
"name": "Legacy Provider",
"base_url": "https://legacy.example/v1",
}
],
"providers": {
"gemini-no-filter": {
"name": "Gemini Display",
"api": endpoint,
}
},
}
with patch("agent.credential_pool._load_config_safe", return_value=config):
assert (
resolve_runtime_pool_key("gemini-no-filter", endpoint)
== "custom:gemini-display"
)
assert (
resolve_runtime_pool_key("custom:gemini-no-filter", endpoint)
== "custom:gemini-display"
)
assert (
resolve_runtime_pool_key(
"custom:gemini-no-filter",
"https://fallback.example/v1",
)
== "custom:gemini-no-filter"
)
def test_runtime_pool_key_preserves_non_custom_identity():
with patch("agent.credential_pool._iter_custom_providers", return_value=[]):
assert (
resolve_runtime_pool_key("openai-codex", "https://chatgpt.com/backend-api")
== "openai-codex"
)
def test_runtime_ignores_pool_loaded_for_different_provider(monkeypatch):
entry = SimpleNamespace(
provider="openai-codex",
+100 -7
View File
@@ -1,18 +1,17 @@
"""Regression tests for the credential-pool provider-mismatch guard with
custom providers (Bernard's Fireworks report, June 2026).
Custom endpoints carry two naming conventions for the same provider: the
agent's ``provider`` attribute is the generic ``"custom"`` label while the
pool is keyed ``custom:<normalized-name>`` (``CUSTOM_POOL_PREFIX``). The
defensive guard in ``recover_with_credential_pool`` compared the two
literally, logged "Credential pool provider mismatch: pool=custom:<name>,
agent=custom", and skipped recovery — so 401/429 recovery (refresh,
rotation) never ran for ANY custom-provider user.
Custom endpoints can carry a generic ``"custom"`` label or retain their
configured name/provider key while the pool is keyed
``custom:<normalized-name>`` (``CUSTOM_POOL_PREFIX``). The defensive guard in
``recover_with_credential_pool`` must recognize each identity without letting
a different endpoint or fallback provider mutate the pool.
The fix accepts the pair only when the agent's current base_url resolves to
the same pool key, preserving the guard's original purpose (#33088/#33163:
never mutate the primary's pool while a fallback provider is active).
"""
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
@@ -36,6 +35,100 @@ def _agent(provider, base_url, pool_provider):
class TestCustomPoolMismatchGuard:
@staticmethod
def _gemini_config():
return [
(
"gemini-no-filter",
{
"name": "Gemini No Filter",
"provider_key": "gemini-no-filter",
"base_url": "https://generativelanguage.googleapis.com/v1beta",
},
)
]
def test_named_custom_provider_rotates_its_matching_pool(self):
agent, pool = _agent(
"gemini-no-filter",
"https://generativelanguage.googleapis.com/v1beta",
"custom:gemini-no-filter",
)
agent.api_key = "key-a"
agent._credential_pool_entry_id = None
agent._swap_credential = MagicMock()
pool.entries.return_value = []
pool.current.return_value = None
next_entry = SimpleNamespace(id="key-b", runtime_api_key="key-b")
pool.mark_exhausted_and_rotate.return_value = next_entry
with patch(
"agent.credential_pool._iter_custom_providers",
return_value=self._gemini_config(),
):
recovered, retried = recover_with_credential_pool(
agent,
status_code=429,
has_retried_429=True,
classified_reason=FailoverReason.rate_limit,
)
assert recovered is True
assert retried is False
pool.mark_exhausted_and_rotate.assert_called_once()
agent._swap_credential.assert_called_once_with(next_entry)
def test_exact_custom_identity_requires_matching_endpoint(self):
agent, pool = _agent(
"custom:gemini-no-filter",
"https://fallback.example/v1",
"custom:gemini-no-filter",
)
with patch(
"agent.credential_pool._iter_custom_providers",
return_value=self._gemini_config(),
):
recovered, retried = recover_with_credential_pool(
agent,
status_code=429,
has_retried_429=True,
classified_reason=FailoverReason.rate_limit,
)
assert recovered is False
assert retried is True
assert not pool.method_calls
def test_exact_custom_identity_rotates_at_matching_endpoint(self):
agent, pool = _agent(
"custom:gemini-no-filter",
"https://generativelanguage.googleapis.com/v1beta",
"custom:gemini-no-filter",
)
agent.api_key = "key-a"
agent._credential_pool_entry_id = None
agent._swap_credential = MagicMock()
pool.entries.return_value = []
pool.current.return_value = None
next_entry = SimpleNamespace(id="key-b", runtime_api_key="key-b")
pool.mark_exhausted_and_rotate.return_value = next_entry
with patch(
"agent.credential_pool._iter_custom_providers",
return_value=self._gemini_config(),
):
recovered, retried = recover_with_credential_pool(
agent,
status_code=429,
has_retried_429=True,
classified_reason=FailoverReason.rate_limit,
)
assert recovered is True
assert retried is False
pool.mark_exhausted_and_rotate.assert_called_once()
agent._swap_credential.assert_called_once_with(next_entry)
def test_unrelated_custom_pool_still_guarded(self):
"""agent=custom pointed at a DIFFERENT endpoint than the pool's
custom provider must still skip pool mutation."""
@@ -374,6 +374,86 @@ class TestRestorePrimaryRuntime:
assert result is True
agent._swap_credential.assert_called_once()
def test_restore_reloads_named_custom_pool_by_scoped_key(self):
class _Entry:
provider = "custom:gemini-display"
id = "gemini-key"
label = "gemini"
runtime_api_key = "gemini-key"
access_token = "gemini-key"
primary_pool = MagicMock()
primary_pool.provider = "custom:gemini-display"
primary_pool.has_available.return_value = True
primary_pool.select.return_value = _Entry()
fallback_pool = MagicMock()
fallback_pool.provider = "openrouter"
agent = _make_agent(
provider="custom:gemini-no-filter",
base_url="https://generativelanguage.googleapis.com/v1beta",
)
agent._fallback_activated = True
agent._credential_pool = fallback_pool
agent._swap_credential = MagicMock()
config = {
"custom_providers": [
{
"name": "Legacy Provider",
"base_url": "https://legacy.example/v1",
}
],
"providers": {
"gemini-no-filter": {
"name": "Gemini Display",
"api": "https://generativelanguage.googleapis.com/v1beta",
}
},
}
with (
patch("agent.credential_pool._load_config_safe", return_value=config),
patch("agent.credential_pool.load_pool", return_value=primary_pool) as load_pool,
patch("run_agent.OpenAI", return_value=MagicMock()),
):
result = agent._restore_primary_runtime()
assert result is True
assert agent._credential_pool is primary_pool
load_pool.assert_called_once_with("custom:gemini-display")
agent._swap_credential.assert_called_once_with(primary_pool.select.return_value)
def test_restore_named_custom_pool_wrong_endpoint_fails_closed(self):
pool = MagicMock()
pool.provider = "custom:gemini-no-filter"
agent = _make_agent(
provider="gemini-no-filter",
base_url="https://fallback.example/v1",
)
agent._fallback_activated = True
agent._credential_pool = pool
agent._swap_credential = MagicMock()
configured = [(
"gemini-no-filter",
{
"name": "Gemini No Filter",
"provider_key": "gemini-no-filter",
"base_url": "https://generativelanguage.googleapis.com/v1beta",
},
)]
with (
patch("agent.credential_pool._iter_custom_providers", return_value=configured),
patch("agent.credential_pool.load_pool", return_value=None) as load_pool,
patch("run_agent.OpenAI", return_value=MagicMock()),
):
result = agent._restore_primary_runtime()
assert result is True
assert agent._credential_pool is None
load_pool.assert_called_once_with("gemini-no-filter")
agent._swap_credential.assert_not_called()