fix(auth): bind stale pool writes to token generation
(cherry picked from commit 1f5c7bcd076cd6184fb048ea1a57c4549f5f4991)
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]]:
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user