diff --git a/agent/credential_pool.py b/agent/credential_pool.py index c51c1ff86a..556d8c5322 100644 --- a/agent/credential_pool.py +++ b/agent/credential_pool.py @@ -22,6 +22,7 @@ from hermes_cli.config import load_env from agent.secret_scope import get_secret as _get_secret, get_secret_str from agent.retry_utils import reset_delay_from_message from hermes_cli.auth_plugin_providers import plugin_refresh_hook +from agent.credential_pool_plugin import apply_plugin_refresh_result, recover_failed_plugin_refresh from agent.credential_persistence import ( fingerprint_secret_value, is_borrowed_credential_source, @@ -254,8 +255,12 @@ class PooledCredential: data["last_status_at"] = _parse_absolute_timestamp(data["last_status_at"]) # Every non-field key rides in ``extra`` (to_dict writes them all back), so metadata a plugin # stores on its own rows survives load -> save -> load. ``_EXTRA_KEYS`` stays the attribute - # surface for core logic; unknown keys are opaque payload. - data["extra"] = {k: v for k, v in payload.items() if k not in field_names and v is not None} + # surface for core logic; unknown keys are opaque payload. ``provider`` is the row's owner + # (excluded from ``field_names`` above), never metadata — sweeping it in would write a + # stray provider name back over the row on to_dict(). + data["extra"] = { + k: v for k, v in payload.items() if k not in field_names and k != "provider" and v is not None + } data.setdefault("id", uuid.uuid4().hex[:6]) data.setdefault("label", payload.get("source", provider)) data.setdefault("auth_type", AUTH_TYPE_API_KEY) @@ -1067,9 +1072,11 @@ class CredentialPool(CredentialPoolAdminMixin, CredentialPoolModelCooldownMixin) the pool store, is token authority for those sources; a row with no token material at all is refused for the same reason. """ - if self.provider not in ("anthropic", "xai-oauth"): + if self.provider not in ("anthropic", "xai-oauth") and plugin_refresh_hook(self.provider) is None: return entry is_anthropic = self.provider == "anthropic" + is_xai = self.provider == "xai-oauth" + display = {"anthropic": "Anthropic", "xai-oauth": "xAI"}.get(self.provider, self.provider) if is_anthropic and is_borrowed_credential_source(entry.source, self.provider): return entry try: @@ -1080,20 +1087,19 @@ class CredentialPool(CredentialPoolAdminMixin, CredentialPoolModelCooldownMixin) if not isinstance(persisted, dict): return entry stored = PooledCredential.from_dict(self.provider, persisted) - if is_anthropic and not (stored.access_token or "").strip() and not (stored.refresh_token or "").strip(): + # No token material at all is never a "rotation" (anthropic borrowed rows, a plugin row a + # peer blanked mid-write): adopting it would replace a usable credential with nothing. + if not is_xai and not (stored.access_token or "").strip() and not (stored.refresh_token or "").strip(): return entry if stored.access_token != entry.access_token or stored.refresh_token != entry.refresh_token: logger.debug( "Pool entry %s: adopting %s OAuth tokens rotated by another pool instance", - entry.id, "Anthropic" if is_anthropic else "xAI", + entry.id, display, ) self._replace_entry(entry, stored) return stored except Exception as exc: - logger.debug( - "Failed to sync %s OAuth entry from credential pool: %s", - "Anthropic" if is_anthropic else "xAI", exc, - ) + logger.debug("Failed to sync %s OAuth entry from credential pool: %s", display, exc) return entry _sync_anthropic_entry_from_pool_store = _sync_entry_from_pool_store @@ -1272,7 +1278,10 @@ class CredentialPool(CredentialPoolAdminMixin, CredentialPoolModelCooldownMixin) if force: self._mark_exhausted(entry, None) return None - if self.provider not in _SINGLE_USE_REFRESH_PROVIDERS: + # Plugin providers with a ``refresh_credential`` hook are treated as single-use by default: + # the pool cannot know their grant semantics, and a needless in-lock re-read is cheaper than + # a ``refresh_token_reused`` login loss. Eligibility comes from the hook, never a name set. + if self.provider not in _SINGLE_USE_REFRESH_PROVIDERS and plugin_refresh_hook(self.provider) is None: return self._refresh_entry_impl(entry, force=force) # Single-use refresh tokens: sync -> POST -> write-back must be atomic @@ -1463,7 +1472,7 @@ class CredentialPool(CredentialPoolAdminMixin, CredentialPoolModelCooldownMixin) entry = self._sync_entry_from_auth_store(entry) updated = self._post_tokens_refresh(entry) elif (plugin_refresh := plugin_refresh_hook(self.provider)) is not None: - updated = replace(entry, **dict(plugin_refresh(entry) or {})) + updated = apply_plugin_refresh_result(entry, plugin_refresh(entry)) elif self.provider == "nous": stale_key = entry.runtime_api_key or entry.agent_key or entry.access_token synced = self._sync_nous_entry_from_auth_store(entry) @@ -1595,6 +1604,10 @@ class CredentialPool(CredentialPoolAdminMixin, CredentialPoolModelCooldownMixin) ) self._mark_dead_refresh_grant(entry, exc) return None + elif plugin_refresh_hook(self.provider) is not None: + handled, result = recover_failed_plugin_refresh(self, entry, exc) + if handled: + return result self._mark_exhausted(entry, None) return None diff --git a/agent/credential_pool_plugin.py b/agent/credential_pool_plugin.py new file mode 100644 index 0000000000..3a4f238521 --- /dev/null +++ b/agent/credential_pool_plugin.py @@ -0,0 +1,92 @@ +"""Plugin-provider refresh support for the credential pool (#116408). + +A model-provider plugin makes its pooled OAuth rows refreshable by shipping +``ProviderProfile.refresh_credential(entry) -> Mapping | None``. The pool owns +the locking, persistence and failure classification around that hook; the +plugin owns only the token POST. This sibling module keeps that logic out of +``agent/credential_pool.py`` (near the size cap). + +Contract (documented in website/docs/developer-guide/model-provider-plugin.md): + +* the hook returns a mapping of rotated values — dataclass field names + (``access_token``, ``refresh_token``, ``expires_at_ms`` …) replace the row's + fields, every other key (``expires_in``, ``token_type``, ``scope`` — the + natural token-endpoint shape) lands in ``entry.extra``; ``None`` = no rotation; +* raising ``AuthError(..., relogin_required=True)`` (or a grant-dead OAuth code) + is terminal: the row goes DEAD with a WARNING naming ``hermes auth add``; + any other exception is transient and only benches the row. +""" + +from __future__ import annotations + +import logging +from dataclasses import fields, replace +from typing import TYPE_CHECKING, Any, Mapping, Optional, Tuple + +from hermes_cli.auth import _OAUTH_GRANT_DEAD_CODES +from hermes_cli.auth_constants import AuthError + +if TYPE_CHECKING: # pragma: no cover + from agent.credential_pool import CredentialPool, PooledCredential + +logger = logging.getLogger(__name__) + + +def apply_plugin_refresh_result(entry: "PooledCredential", result: Any) -> "PooledCredential": + """Merge a ``refresh_credential`` return value into *entry*. + + Field names go through ``dataclasses.replace``; everything else is merged into ``extra`` + (mirroring ``PooledCredential.from_dict``). Before this split a token-endpoint-shaped mapping + with ``expires_in`` raised ``TypeError`` inside ``replace`` and the pool benched a row whose + single-use refresh token the server had already rotated — the login was lost. + """ + if not result: + return entry + mapping: Mapping[str, Any] = dict(result) + field_names = {f.name for f in fields(type(entry))} - {"provider", "extra"} + field_updates = {k: v for k, v in mapping.items() if k in field_names} + extra_updates = {k: v for k, v in mapping.items() if k not in field_names and k != "provider"} + if extra_updates: + field_updates["extra"] = {**entry.extra, **extra_updates} + return replace(entry, **field_updates) if field_updates else entry + + +def is_terminal_plugin_refresh_error(exc: BaseException) -> bool: + """True when retrying the same plugin refresh token cannot succeed. + + Plugins have no per-provider code table, so the predicate is the structural one every built-in + flow shares: a structured ``AuthError`` that asks for a re-login, or one carrying a grant-dead + OAuth code. Transport errors, 429/5xx-shaped ``RuntimeError`` and plain ``AuthError`` without + that signal stay transient (EXHAUSTED for one cooldown). + """ + if not isinstance(exc, AuthError): + return False + return bool(exc.relogin_required) or (exc.code or "") in _OAUTH_GRANT_DEAD_CODES + + +def recover_failed_plugin_refresh( + pool: "CredentialPool", entry: "PooledCredential", exc: Exception, +) -> Tuple[bool, Optional["PooledCredential"]]: + """Recovery for a plugin hook that raised: adopt a peer's rotation, or quarantine a dead grant. + + Returns ``(handled, result)``; ``handled=False`` means the caller should bench the row as a + transient failure. A peer process may have rotated the single-use token between our in-lock + sync and the hook's POST — adopt that pair before classifying the failure. + """ + from agent.credential_pool import _MARK_OK + + synced = pool._sync_entry_from_pool_store(entry) + if synced.refresh_token != entry.refresh_token and (synced.access_token or "").strip(): + logger.debug("%s refresh failed but the pool store has newer tokens — adopting", pool.provider) + return True, pool._adopt(synced, **_MARK_OK) + if is_terminal_plugin_refresh_error(exc): + # WARNING, not debug: this is the moment a login is lost. Benching for a TTL would replay + # the dead token every cooldown at DEBUG with no trace for the user. + logger.warning( + "%s refresh token for %s is terminally invalid (%s); the credential leaves rotation. " + "Re-run 'hermes auth add %s' to sign in again.", + pool.provider, entry.label or entry.id[:8], exc, pool.provider, + ) + pool._mark_dead_refresh_grant(entry, exc) + return True, None + return False, None diff --git a/tests/agent/test_credential_pool_plugin_seam.py b/tests/agent/test_credential_pool_plugin_seam.py index ccac5eb780..f43022c43c 100644 --- a/tests/agent/test_credential_pool_plugin_seam.py +++ b/tests/agent/test_credential_pool_plugin_seam.py @@ -3,6 +3,7 @@ through the profile's ``refresh_credential`` hook — eligibility derives from t from __future__ import annotations +import logging from dataclasses import replace import pytest @@ -10,7 +11,16 @@ import pytest import providers from providers.base import ProviderProfile -from agent.credential_pool import AUTH_TYPE_OAUTH, CredentialPool, PooledCredential +from agent.credential_pool import ( + AUTH_TYPE_OAUTH, + STATUS_DEAD, + STATUS_EXHAUSTED, + STATUS_OK, + CredentialPool, + PooledCredential, +) +from hermes_cli.auth import read_credential_pool +from hermes_cli.auth_constants import AuthError from hermes_cli.auth_plugin_providers import is_refreshable_oauth_provider @@ -28,6 +38,10 @@ def test_plugin_metadata_survives_load_save_load(): assert PooledCredential.from_dict("example-oauth", again).extra == {"tenant": "acme", "region": "eu"} # Core-known extra keys keep their attribute surface; unknown ones stay opaque payload. assert PooledCredential.from_dict("nous", {"access_token": "t", "org_id": "o1"}).org_id == "o1" + # A stray row-level ``provider`` is pool bookkeeping, not plugin metadata: it must not be swept + # into ``extra`` and written back over the owning provider's row. + assert "provider" not in PooledCredential.from_dict( + "example-oauth", {**payload, "provider": "other"}).to_dict() @pytest.fixture @@ -64,3 +78,66 @@ def test_pool_refresh_dispatches_to_profile_hook(plugin_profiles, monkeypatch): # Without the hook the pool must not pretend it refreshed anything. nohook = replace(entry, provider="example-oauth-nohook") assert CredentialPool("example-oauth-nohook", [nohook])._refresh_entry_impl(nohook, force=True) is nohook + + +def _register_hook(hook): + providers.register_provider(ProviderProfile(name="example-oauth", auth_type="oauth_external", + base_url="https://example.invalid/v1", + refresh_credential=hook)) + + +def _token_endpoint_shape(entry): + # The natural token-endpoint response: field keys plus keys that are not dataclass fields. + return {"access_token": "tok-2", "refresh_token": "rt-2", "expires_in": 3600, + "expires_at_ms": 4102444800000, "token_type": "Bearer"} + + +def _raise_transient(entry): + raise RuntimeError("token endpoint said 503") + + +def _raise_terminal(entry): + raise AuthError("invalid_grant", provider="example-oauth", code="example_refresh_failed", + relogin_required=True) + + +@pytest.mark.parametrize("hook, status, tokens, extra_key", [ + (_token_endpoint_shape, STATUS_OK, ("tok-2", "rt-2"), "expires_in"), + (_raise_transient, STATUS_EXHAUSTED, ("tok-1", "rt-1"), None), + (_raise_terminal, STATUS_DEAD, ("tok-1", "rt-1"), None), +], ids=["token-endpoint-shape-rotates", "transient-error-benches", "relogin-required-is-terminal"]) +def test_plugin_refresh_outcome(plugin_profiles, caplog, hook, status, tokens, extra_key): + _register_hook(hook) + entry = _entry(expires_at_ms=1) + pool = CredentialPool("example-oauth", [entry]) + pool._persist() + with caplog.at_level(logging.WARNING, logger="agent.credential_pool"): + pool._refresh_entry(entry, force=True) + row = PooledCredential.from_dict("example-oauth", read_credential_pool("example-oauth")[0]) + assert (row.last_status, row.access_token, row.refresh_token) == (status, *tokens) + if extra_key: + assert row.extra[extra_key] == 3600 and row.expires_at_ms == 4102444800000 + if status == STATUS_DEAD: + assert "hermes auth add example-oauth" in caplog.text + else: + assert "hermes auth add" not in caplog.text + + +def test_plugin_refresh_adopts_peer_rotation_without_spending_token(plugin_profiles): + """Two Hermes processes share one auth.json: the second refresh adopts the first's rotated pair + instead of POSTing the same single-use refresh token again.""" + calls = [] + + def hook(entry): + calls.append(entry.refresh_token) + return {"access_token": f"tok-{len(calls) + 1}", "refresh_token": f"rt-{len(calls) + 1}"} + + _register_hook(hook) + first = CredentialPool("example-oauth", [_entry(expires_at_ms=1)]) + first._persist() + second = CredentialPool("example-oauth", [_entry(expires_at_ms=1)]) # stale copy, same on-disk row + + assert first._refresh_entry(first.entries()[0], force=True).access_token == "tok-2" + adopted = second._refresh_entry(second.entries()[0], force=True) + assert (adopted.access_token, adopted.refresh_token) == ("tok-2", "rt-2") + assert calls == ["rt-1"] diff --git a/website/docs/developer-guide/model-provider-plugin.md b/website/docs/developer-guide/model-provider-plugin.md index a0b6fc948a..c2d1f1d7ee 100644 --- a/website/docs/developer-guide/model-provider-plugin.md +++ b/website/docs/developer-guide/model-provider-plugin.md @@ -267,10 +267,10 @@ def example_auth(action: str, args) -> bool: def example_refresh(entry): - """Called by the credential pool with the pooled row; return the rotated fields or raise.""" - tokens = post_refresh(entry.refresh_token) + """Called by the credential pool with the pooled row; return the rotated values, None, or raise.""" + tokens = post_refresh(entry.refresh_token) # the raw token-endpoint response is fine as-is return {"access_token": tokens["access_token"], "refresh_token": tokens["refresh_token"], - "expires_at_ms": tokens["expires_at_ms"]} + "expires_at_ms": tokens["expires_at_ms"], "expires_in": tokens["expires_in"]} register_provider(ProviderProfile( @@ -281,7 +281,9 @@ register_provider(ProviderProfile( | Contract | | |---|---| | `auth_handler(action, args)` | `args` is the parsed `hermes auth` namespace. Truthy = handled (Hermes prints nothing more, exit 0); falsy = fall back to the built-in path **for that action**. An exception becomes `SystemExit(" auth handler failed for `<action>`: …")`. | -| `refresh_credential(entry)` | Receives the `PooledCredential`; returns a mapping of rotated fields (`access_token`, `refresh_token`, `expires_at_ms`, …) applied to the row, or raises (the pool benches the row). Its presence is what makes the provider *refreshable* — `hermes auth refresh ` and the 401 recovery paths (main loop and auxiliary client) call it; no core name list is involved. | +| `refresh_credential(entry)` | Receives the `PooledCredential`; returns a mapping of rotated values or `None`. Keys that are `PooledCredential` fields (`access_token`, `refresh_token`, `expires_at_ms`, …) replace the row's fields; every other key (`expires_in`, `token_type`, `scope` — the raw token-endpoint shape) lands in `entry.extra` and round-trips through `auth.json`. `None` = no rotation happened, the row is marked ok. Its presence is what makes the provider *refreshable* — `hermes auth refresh ` and the 401 recovery paths (main loop and auxiliary client) call it; no core name list is involved. | +| Refresh failures | Raise `hermes_cli.auth_constants.AuthError(..., relogin_required=True)` (or with `code` `invalid_grant` / `invalid_token` / `refresh_token_reused`) when the grant is dead: the row goes **DEAD**, leaves rotation and Hermes logs a WARNING naming `hermes auth add `. Any other exception (network, 429, 5xx) is transient — the row is benched for one cooldown and retried. | +| Concurrency | The hook runs under the shared `auth.json` lock. Before calling it the pool re-reads the row; if another Hermes process (gateway + CLI, two profiles) already rotated the pair, that pair is adopted and your hook is **not** called — safe for single-use refresh tokens. After the hook returns, the rotated row is written through to `auth.json`. | | No hooks | `api_key` profiles behave exactly as before. Any other `auth_type` without `auth_handler` fails loud on `hermes auth add`. | `hermes auth add|status|logout|refresh ` consults the handler **first** — before the built-in