From d2f54de2cc5b1b2d2e430301a273940940f466f7 Mon Sep 17 00:00:00 2001 From: JoaoMarcos44 Date: Wed, 23 Sep 2026 23:36:57 -0300 Subject: [PATCH] fix(auth): bind stale pool writes to token generation (cherry picked from commit 1f5c7bcd076cd6184fb048ea1a57c4549f5f4991) --- agent/credential_pool.py | 56 ++++++++++++------ agent/credential_pool_admin.py | 8 +-- hermes_cli/auth.py | 59 +++++++++++++++++-- ...test_credential_pool_profile_oauth_fork.py | 55 +++++++++++++++++ 4 files changed, 148 insertions(+), 30 deletions(-) diff --git a/agent/credential_pool.py b/agent/credential_pool.py index 4b671068eb..de9e868b3a 100644 --- a/agent/credential_pool.py +++ b/agent/credential_pool.py @@ -857,7 +857,8 @@ def _borrowed_single_use_pool_root() -> Optional[Path]: def _update_root_pool_rows( provider: str, payloads: List[Dict[str, Any]], global_path: Path, *, status_cleared_ids: Optional[Iterable[str]] = None, -) -> None: + token_bases: Optional[Dict[str, Tuple[Any, Any]]] = None, +) -> List[Dict[str, Any]]: """UPDATE-ONLY merge of *payloads* into the root store's rows for *provider*. A borrower may refresh the root's rows (rotation, cooldown state) but @@ -875,6 +876,7 @@ def _update_root_pool_rows( existing_list = existing if isinstance(existing, list) else [] incoming_by_id = {p.get("id"): p for p in payloads if isinstance(p, dict) and p.get("id")} cleared = {cid for cid in (status_cleared_ids or ()) if cid} + bases = token_bases or {} merged: List[Dict[str, Any]] = [] changed = False for disk_entry in existing_list: @@ -883,9 +885,9 @@ def _update_root_pool_rows( if incoming is None: merged.append(disk_entry) continue - # A deliberately cleared entry has no disk cooldown worth keeping. - updated = auth_mod._merge_disk_cooldown_state( - incoming, None if did in cleared else disk_entry, provider, + updated = auth_mod._merge_pool_row_generation( + incoming, disk_entry, provider, + base_pair=bases.get(did), status_cleared=did in cleared, ) if updated != disk_entry: changed = True @@ -893,6 +895,7 @@ def _update_root_pool_rows( if changed: pool[provider] = merged _save_auth_store(store, target_path=global_path) + return merged def persist_pool_entries( @@ -901,7 +904,8 @@ def persist_pool_entries( *, removed_ids: Optional[Iterable[str]] = None, status_cleared_ids: Optional[Iterable[str]] = None, -) -> None: + token_bases: Optional[Dict[str, Tuple[Any, Any]]] = None, +) -> Optional[List[Dict[str, Any]]]: """Persist a provider's pool rows to the store that OWNS them. A named profile that sees a single-use-refresh provider (Anthropic, @@ -916,9 +920,9 @@ def persist_pool_entries( global_path = _borrowed_single_use_pool_root() if global_path is not None: try: - _update_root_pool_rows( + return _update_root_pool_rows( provider, payloads, global_path, - status_cleared_ids=status_cleared_ids, + status_cleared_ids=status_cleared_ids, token_bases=token_bases, ) except Exception as exc: # Fail closed on the FORK, not on the save: never fall back to @@ -929,9 +933,10 @@ def persist_pool_entries( "not materializing a profile-local copy", provider, exc, ) - return - write_credential_pool( + return None + return write_credential_pool( provider, payloads, removed_ids=removed_ids, status_cleared_ids=status_cleared_ids, + token_bases=token_bases, ) @@ -987,6 +992,7 @@ class CredentialPool(CredentialPoolAdminMixin, CredentialPoolModelCooldownMixin) # Ids of rows read via the global-root fallback (single-use OAuth # providers only); set by load_pool(), consumed by add_entry(). self._borrowed_root_ids: Set[str] = set() + self._persisted_token_pairs: Dict[str, Tuple[Any, Any]] = {} self._strategy = get_pool_strategy(provider) # RLock: _replace_entry/_persist self-acquire it so the DEFERRED # single-use-token refresh path (network I/O outside the lock by @@ -1115,12 +1121,24 @@ class CredentialPool(CredentialPoolAdminMixin, CredentialPoolModelCooldownMixin) ) -> None: # Self-locking: snapshotting self._entries must not race a rotation. with self._lock: - persist_pool_entries( - self.provider, - [entry.to_dict() for entry in self._entries], + payloads = [entry.to_dict() for entry in self._entries] + written = persist_pool_entries( + self.provider, payloads, removed_ids=removed_ids, status_cleared_ids=status_cleared_ids, + token_bases=getattr(self, "_persisted_token_pairs", {}), ) + if written is None: + return + rows = {row.get("id"): row for row in written if isinstance(row, dict) and row.get("id")} + pairs = {row_id: auth_mod._credential_token_pair(row) for row_id, row in rows.items()} + self._persisted_token_pairs = {row_id: pair for row_id, pair in pairs.items() if any(pair)} + for index, entry in enumerate(self._entries): + row = rows.get(entry.id) + if row is not None and any(pairs[entry.id]): + # Reference-only rows are intentionally secret-free on disk; never dehydrate + # their live in-memory credential while adopting a concurrent generation. + self._entries[index] = PooledCredential.from_dict(self.provider, row) def _adopt(self, entry: PooledCredential, *, persist: bool = True, **updates: Any) -> PooledCredential: """``replace(entry, **updates)``, swap it into the pool, optionally persist.""" @@ -1271,6 +1289,7 @@ class CredentialPool(CredentialPoolAdminMixin, CredentialPoolModelCooldownMixin) # 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 + self._persisted_token_pairs[entry.id] = auth_mod._credential_token_pair(persisted) 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", @@ -3042,14 +3061,13 @@ def load_pool(provider: str) -> CredentialPool: ) changed |= _normalize_pool_priorities(provider, entries) - if changed: - new_ids = {entry.id for entry in entries} - persist_pool_entries( - provider, - [entry.to_dict() for entry in sorted(entries, key=lambda item: item.priority)], - removed_ids=disk_ids - new_ids, - ) pool = CredentialPool(provider, entries) + pool._persisted_token_pairs = { + payload["id"]: auth_mod._credential_token_pair(payload) + for payload in raw_entries if isinstance(payload, dict) and payload.get("id") + } + if changed: + pool._persist(removed_ids=sorted(disk_ids - {entry.id for entry in entries})) # Remember the root's borrowed rows so a later ``add_entry`` in this # profile leaves them out of the profile's own store (#100339). if provider in SINGLE_USE_REFRESH_POOL_PROVIDERS and not _profile_owns_pool_provider(provider): diff --git a/agent/credential_pool_admin.py b/agent/credential_pool_admin.py index f5b89f7e08..a4ebad859f 100644 --- a/agent/credential_pool_admin.py +++ b/agent/credential_pool_admin.py @@ -54,18 +54,12 @@ class CredentialPoolAdminMixin: return len(stale) def remove_index(self, index: int) -> Optional[PooledCredential]: - from agent.credential_pool import persist_pool_entries - with self._lock: if index < 1 or index > len(self._entries): return None removed = self._entries.pop(index - 1) self._entries = [replace(e, priority=p) for p, e in enumerate(self._entries)] - persist_pool_entries( - self.provider, - [entry.to_dict() for entry in self._entries], - removed_ids=[removed.id], - ) + self._persist(removed_ids=[removed.id]) if self._current_id == removed.id: self._current_id = None return removed diff --git a/hermes_cli/auth.py b/hermes_cli/auth.py index ee50be10c4..2057722bcf 100644 --- a/hermes_cli/auth.py +++ b/hermes_cli/auth.py @@ -897,6 +897,52 @@ def read_credential_pool(provider_id: Optional[str] = None) -> Dict[str, Any]: _POOL_STATUS_FIELDS = ( "last_status", "last_status_at", "last_error_code", "last_error_reason", "last_error_message", "last_error_reset_at", "status_cleared_at") +_POOL_TOKEN_GENERATION_FIELDS = ( + "access_token", "refresh_token", "expires_at", "expires_at_ms", "expires_in", "obtained_at", + "last_refresh", "agent_key", "agent_key_expires_at", "agent_key_expires_in", "agent_key_id", + "agent_key_obtained_at", "agent_key_reused", +) + + +def _credential_token_pair(row: Any) -> Tuple[Any, Any]: + if not isinstance(row, dict): + return (None, None) + return row.get("access_token"), row.get("refresh_token") + + +def _merge_pool_row_generation( + entry: Dict[str, Any], + disk_entry: Optional[Dict[str, Any]], + provider_id: str, + *, + base_pair: Optional[Tuple[Any, Any]] = None, + status_cleared: bool = False, +) -> Dict[str, Any]: + """Keep a newer on-disk token generation authoritative during stale writes.""" + merge_disk = None if status_cleared else disk_entry + if not isinstance(disk_entry, dict) or base_pair is None: + return _merge_disk_cooldown_state(entry, merge_disk, provider_id) + disk_pair = _credential_token_pair(disk_entry) + if not any(disk_pair) or disk_pair == base_pair: + return _merge_disk_cooldown_state(entry, merge_disk, provider_id) + + merged = dict(entry) + for field in _POOL_TOKEN_GENERATION_FIELDS: + if field in disk_entry: + merged[field] = disk_entry[field] + else: + merged.pop(field, None) + if not status_cleared: + for field in _POOL_STATUS_FIELDS: + if field in disk_entry: + merged[field] = disk_entry[field] + else: + merged.pop(field, None) + if "failure_reason" in disk_entry: + merged["failure_reason"] = disk_entry["failure_reason"] + else: + merged.pop("failure_reason", None) + return _merge_disk_cooldown_state(merged, merge_disk, provider_id) def _merge_disk_cooldown_state( @@ -957,7 +1003,8 @@ def write_credential_pool( provider_id: str, entries: List[Dict[str, Any]], *, removed_ids: Optional[Iterable[str]] = None, status_cleared_ids: Optional[Iterable[str]] = None, -) -> Path: + token_bases: Optional[Dict[str, Tuple[Any, Any]]] = None, +) -> List[Dict[str, Any]]: """Persist one provider's credential pool under auth.json. Final disk-boundary sanitizer for borrowed credentials (callers may pass raw dicts). Entries on @@ -967,6 +1014,7 @@ def write_credential_pool( recency merge, which would otherwise read their cleared ``last_status_at`` (None -> epoch 0) as a stale snapshot and copy a still-binding cooldown back.""" removed = {rid for rid in (removed_ids or ()) if rid} + bases = token_bases or {} with _auth_store_lock(): auth_store = _load_auth_store() pool = _store_section(auth_store, "credential_pool") @@ -979,8 +1027,10 @@ def write_credential_pool( new_ids = set(_entry_ids(sanitized)) status_cleared = {cid for cid in (status_cleared_ids or ()) if cid} merged: List[Dict[str, Any]] = [ - _merge_disk_cooldown_state( - e, None if e.get("id") in status_cleared else existing_by_id.get(e.get("id")), provider_id, + _merge_pool_row_generation( + e, existing_by_id.get(e.get("id")), provider_id, + base_pair=bases.get(e.get("id")), + status_cleared=e.get("id") in status_cleared, ) if isinstance(e, dict) else e for e in sanitized] @@ -989,7 +1039,8 @@ def write_credential_pool( if disk_id and disk_id not in new_ids and disk_id not in removed: merged.append(sanitize_borrowed_credential_payload(disk_entry, provider_id)) pool[provider_id] = merged - return _save_auth_store(auth_store) + _save_auth_store(auth_store) + return merged def _suppressed_source_list(suppressed: Dict[str, Any], provider_id: str) -> Optional[List[str]]: diff --git a/tests/agent/test_credential_pool_profile_oauth_fork.py b/tests/agent/test_credential_pool_profile_oauth_fork.py index 156f0c0762..694ec47d54 100644 --- a/tests/agent/test_credential_pool_profile_oauth_fork.py +++ b/tests/agent/test_credential_pool_profile_oauth_fork.py @@ -756,3 +756,58 @@ def test_persisted_mark_still_re_heals_when_the_root_store_gains_a_grant(fleet): _new_process(auth_mod) assert auth_mod.heal_forked_single_use_oauth_grants("anthropic") is None assert grants._oauth_heal_clean_mark_path().read_text() != marked + + + +# ── G. token generation is the authority boundary (#120815) ───────────── + +@pytest.mark.parametrize("borrowed", [False, True], ids=["owned-rows", "borrowed-root-rows"]) +def test_stale_pool_cannot_roll_back_peer_token_generation(fleet, borrowed): + from agent.credential_pool import load_pool + + fleet["use"](_profile(fleet, "stale-write") if borrowed else fleet["root"]) + stale, peer = load_pool("anthropic"), load_pool("anthropic") + assert peer.try_refresh_matching(credential_id="abc123").refresh_token == "sk-ant-ort-RT1" + + stale.mark_exhausted_and_rotate(status_code=429, credential_id="abc123") + + row = fleet["rows"](fleet["root"])[0] + assert row["refresh_token"] == "sk-ant-ort-RT1" + selected = stale.select() + assert selected is not None and selected.refresh_token == "sk-ant-ort-RT1" + + +@pytest.mark.parametrize("borrowed", [False, True], ids=["owned-rows", "borrowed-root-rows"]) +def test_stale_terminal_verdict_cannot_kill_peer_token_generation(fleet, borrowed): + from agent.credential_pool import STATUS_DEAD, load_pool + + fleet["use"](_profile(fleet, "stale-dead") if borrowed else fleet["root"]) + stale, peer = load_pool("anthropic"), load_pool("anthropic") + assert peer.try_refresh_matching(credential_id="abc123").refresh_token == "sk-ant-ort-RT1" + + stale.mark_exhausted_and_rotate( + status_code=401, credential_id="abc123", + error_context={"reason": "invalid_grant", "message": "refresh token already used"}, + ) + + row = fleet["rows"](fleet["root"])[0] + assert row["refresh_token"] == "sk-ant-ort-RT1" + assert row.get("last_status") != STATUS_DEAD + selected = stale.select() + assert selected is not None and selected.refresh_token == "sk-ant-ort-RT1" + + +def test_terminal_verdict_for_current_token_generation_still_persists(fleet): + from agent.credential_pool import STATUS_DEAD, load_pool + + fleet["use"](fleet["root"]) + pool = load_pool("anthropic") + pool.mark_exhausted_and_rotate( + status_code=401, credential_id="abc123", + error_context={"reason": "invalid_grant", "message": "refresh token already used"}, + ) + + row = fleet["rows"](fleet["root"])[0] + assert row["refresh_token"] == "sk-ant-ort-RT0" + assert row["last_status"] == STATUS_DEAD + assert pool.select() is None