fix(auth): bind stale pool writes to token generation

(cherry picked from commit 1f5c7bcd076cd6184fb048ea1a57c4549f5f4991)
This commit is contained in:
JoaoMarcos44
2026-09-23 23:36:57 -03:00
committed by kshitij
parent 299dfd35f3
commit d2f54de2cc
4 changed files with 148 additions and 30 deletions

View File

@@ -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):

View File

@@ -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

View File

@@ -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]]:

View File

@@ -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