fix: place reauthenticated credentials by their saved identity

This commit is contained in:
Teknium
2026-09-07 02:14:51 -07:00
parent 32a59f3bf7
commit de25786dc7
3 changed files with 191 additions and 42 deletions

View File

@@ -258,18 +258,19 @@ def _ask(prompt: str, reader: Callable[[str], str] | None = None) -> str | None:
return None
def _add_nous_oauth_credential(args, provider: str) -> None:
def _add_nous_oauth_credential(args, provider: str) -> PooledCredential:
"""``hermes auth add nous --type oauth``: shared-credential import, else device-code login."""
custom_label = (getattr(args, "label", None) or "").strip() or None
timeout = getattr(args, "timeout", None) or 15.0
def _persist(creds: dict, what: str) -> None:
def _persist(creds: dict, what: str) -> PooledCredential:
# `--label` is embedded into providers.nous so label_from_token doesn't overwrite it on every
# subsequent load_pool("nous").
entry = auth_mod.persist_nous_credentials(creds, label=custom_label)
shown_label = entry.label if entry is not None else label_from_token(
creds.get("access_token", ""), f"{provider}-oauth-1")
print(f'{what} {provider} OAuth {"device-code " if what == "Saved" else ""}credentials: "{shown_label}"')
return entry
# Codex-style auto-import: a shared Nous credential at <hermes-root>/shared/nous_auth.json
# (written by any previous login) makes `hermes --profile <name> auth add nous --type oauth`
@@ -286,8 +287,7 @@ def _add_nous_oauth_credential(args, provider: str) -> None:
print("Rehydrating Nous session from shared credentials...")
rehydrated = auth_mod._try_import_shared_nous_state(timeout_seconds=timeout)
if rehydrated is not None:
_persist(rehydrated, "Imported")
return
return _persist(rehydrated, "Imported")
# Expired refresh_token, portal down, etc. — fall through to device-code.
print("Could not refresh shared credentials — falling back to device-code login.")
@@ -297,7 +297,7 @@ def _add_nous_oauth_credential(args, provider: str) -> None:
client_id=getattr(args, "client_id", None), scope=getattr(args, "scope", None),
open_browser=not getattr(args, "no_browser", False), timeout_seconds=timeout,
insecure=bool(getattr(args, "insecure", False)), ca_bundle=getattr(args, "ca_bundle", None))
_persist(creds, "Saved")
return _persist(creds, "Saved")
def _unsuppress_provider_sources(provider: str) -> None:
@@ -312,7 +312,7 @@ def _unsuppress_provider_sources(provider: str) -> None:
pass
def _add_api_key_credential(args, provider: str, pool) -> None:
def _add_api_key_credential(args, provider: str, pool) -> PooledCredential:
token = ((getattr(args, "api_key", None) or "").strip()
or masked_secret_prompt("Paste your API key: ").strip())
if not token:
@@ -325,8 +325,9 @@ def _add_api_key_credential(args, provider: str, pool) -> None:
entry = PooledCredential(
provider=provider, id=uuid.uuid4().hex[:6], label=label, auth_type=AUTH_TYPE_API_KEY,
priority=0, source=SOURCE_MANUAL, access_token=token, base_url=_provider_base_url(provider))
pool.add_entry(entry)
entry = pool.add_entry(entry)
print(f'Added {provider} credential #{len(pool.entries())}: "{label}"')
return entry
def auth_add_command(args) -> None:
@@ -350,19 +351,18 @@ def auth_add_command(args) -> None:
_unsuppress_provider_sources(provider)
wanted_priority = getattr(args, "priority", None)
before = {entry.id for entry in pool.entries()}
_add_credential(args, provider, pool, requested_type)
entry = _add_credential(args, provider, pool, requested_type)
if wanted_priority is not None:
_place_added_credential(provider, before, int(wanted_priority))
placed_pool = load_pool(provider)
moved = placed_pool.move_entry(entry.id, int(wanted_priority))
_report_priority(provider, placed_pool, moved, int(wanted_priority), "Placed", "at")
def _add_credential(args, provider: str, pool, requested_type: str) -> None:
def _add_credential(args, provider: str, pool, requested_type: str) -> PooledCredential:
if requested_type == AUTH_TYPE_API_KEY:
_add_api_key_credential(args, provider, pool)
return
return _add_api_key_credential(args, provider, pool)
if provider == "nous":
_add_nous_oauth_credential(args, provider)
return
return _add_nous_oauth_credential(args, provider)
spec = _OAUTH_ADD_SPECS.get(provider)
if spec is None:
@@ -379,38 +379,13 @@ def _add_credential(args, provider: str, pool, requested_type: str) -> None:
provider=provider, id=uuid.uuid4().hex[:6], label=label, auth_type=AUTH_TYPE_OAUTH, priority=0,
source=spec.source, access_token=token, **spec.fields(creds, provider))
first_credential = not pool.entries()
pool.add_entry(entry)
entry = pool.add_entry(entry)
# The first Codex/xAI credential becomes the active provider (as the old singleton save path
# did implicitly); subsequent adds leave the active provider as-is.
if spec.activate_first and first_credential:
auth_mod.mark_provider_active_if_unset(provider)
print(f'Added {provider} OAuth credential #{len(pool.entries())}: "{entry.label}"')
def _place_added_credential(provider: str, before_ids: set, priority: int) -> None:
"""Move the credential `auth add` just created to *priority*.
Every add path (api key, OAuth spec, Nous) persists through the pool, so the
new row is the one id that was not there before the add. Reloading rather
than reusing the add's pool object keeps this correct for paths that write
their own store.
"""
pool = load_pool(provider)
entries = pool.entries()
added = [entry for entry in entries if entry.id not in before_ids]
if len(added) == 1:
target = added[0]
elif not added and len(entries) == 1:
# The add updated the sole existing row in place (a repeat Nous login).
target = entries[0]
else:
# The credential was saved; only the placement is unresolved, so do not fail.
print(f"note: could not identify the credential just added to {provider}; set its "
f"priority with `hermes auth priority {provider} <target> {priority}`.",
file=sys.stderr)
return
moved = pool.move_entry(target.id, priority)
_report_priority(provider, pool, moved, priority, "Placed", "at")
return entry
def _report_priority(provider: str, pool, moved, requested: int, verb: str, prep: str) -> None:

View File

@@ -0,0 +1,70 @@
"""Durable pool administration and selection invariants."""
import time
from dataclasses import replace
import pytest
from agent.credential_pool import CredentialPool, PooledCredential
from hermes_cli.auth import read_credential_pool, write_credential_pool
def _pool(provider="openrouter", *, exhausted=False):
rows = [PooledCredential(
provider=provider, id=f"row{i}", label=f"account{i}", source="manual",
auth_type="api_key", access_token=f"fixture-{i}", priority=i,
last_status="exhausted" if exhausted else None,
last_status_at=time.time() if exhausted else None,
last_error_code=429 if exhausted else None,
last_error_reset_at=time.time() + 3600 if exhausted else None,
) for i in range(2)]
write_credential_pool(provider, [e.to_dict() for e in rows])
return CredentialPool(provider, rows)
def test_target_reset_preserves_sibling_cooldown():
pool = _pool(exhausted=True)
before = read_credential_pool(pool.provider)
assert pool.reset_status("missing") is None
assert read_credential_pool(pool.provider) == before
assert pool.reset_status("row1").last_status is None
after = {e["id"]: e for e in read_credential_pool(pool.provider)}
assert after["row0"] == before[0]
assert after["row1"].get("last_error_reset_at") is None
assert pool.reset_statuses() == 1
assert all(e.get("last_status") is None for e in read_credential_pool(pool.provider))
@pytest.mark.parametrize("strategy", ["fill_first", "round_robin", "random", "least_used"])
def test_selection_counts_only_returned_selections(strategy):
pool = _pool()
pool._strategy = strategy
selected = [pool.select().id for _ in range(2)]
assert {e.id: e.request_count for e in pool.entries()} == {
e.id: selected.count(e.id) for e in pool.entries()}
pool.reset_status("row0") # existing persistence boundary, not a selection
assert sum(e.get("request_count", 0) for e in read_credential_pool(pool.provider)) == 2
pool._current_id = None
pool.try_refresh_matching() # API key is not refreshable; lookup must not count
assert sum(e.request_count for e in pool.entries()) == 2
assert pool.peek() is not None
assert sum(e.request_count for e in pool.entries()) == 2
def test_priority_persists_contiguous_order_without_clearing_cooldown():
pool = _pool(exhausted=True)
before = {e.id: e.last_error_reset_at for e in pool.entries()}
assert pool.move_entry("row1", -5).priority == 0
assert [e["id"] for e in read_credential_pool(pool.provider)] == ["row1", "row0"]
assert pool.move_entry("row1", 99).priority == 1
assert [(e.id, e.priority) for e in pool.entries()] == [("row0", 0), ("row1", 1)]
assert {e.id: e.last_error_reset_at for e in pool.entries()} == before
snapshot = read_credential_pool(pool.provider)
assert pool.move_entry("missing", 0) is None
assert read_credential_pool(pool.provider) == snapshot
def test_priority_honors_anthropic_manual_first():
pool = _pool("anthropic")
pool._entries[1] = replace(pool._entries[1], source="env:ANTHROPIC_API_KEY")
assert pool.move_entry("row1", 0).priority == 1
assert [e.id for e in pool.entries()] == ["row0", "row1"]

View File

@@ -0,0 +1,104 @@
"""OAuth control commands against a loopback token endpoint."""
import json
import threading
import time
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from types import SimpleNamespace
from urllib.parse import parse_qs
import pytest
from hermes_cli import auth_commands
from hermes_cli.auth import read_credential_pool, write_credential_pool
def _rows():
return [dict(id=f"row{i}", label=f"account{i}", source="manual:device_code",
auth_type="oauth", access_token=f"fixture-access-{i}",
refresh_token=f"fixture-refresh-{i}", priority=i,
last_status="exhausted", last_status_at=time.time(),
last_error_code=429, last_error_reset_at=time.time()+3600)
for i in range(2)]
@pytest.mark.parametrize("status", [200, 503, 401])
def test_refresh_uses_target_grant_and_preserves_sibling(monkeypatch, status):
from hermes_cli import auth_codex
requests = []
class Endpoint(BaseHTTPRequestHandler):
def do_POST(self):
requests.append(parse_qs(self.rfile.read(int(self.headers["Content-Length"])).decode()))
body = ({"access_token": "fixture-new-access", "refresh_token": "fixture-new-refresh"}
if status == 200 else {"error": "invalid_grant" if status == 401 else "unavailable"})
self.send_response(status)
self.send_header("Content-Type", "application/json")
self.end_headers()
self.wfile.write(json.dumps(body).encode())
def log_message(self, *_args):
pass
server = ThreadingHTTPServer(("127.0.0.1", 0), Endpoint)
worker = threading.Thread(target=server.serve_forever, daemon=True)
worker.start()
monkeypatch.setattr(auth_codex, "CODEX_OAUTH_TOKEN_URL", f"http://127.0.0.1:{server.server_port}/token")
from agent.credential_pool import PooledCredential
rows = [PooledCredential.from_dict("openai-codex", row).to_dict() for row in _rows()]
write_credential_pool("openai-codex", rows)
before = read_credential_pool("openai-codex")
try:
args = SimpleNamespace(provider="openai-codex", target="row1")
if status == 200:
auth_commands.auth_refresh_command(args)
else:
with pytest.raises(SystemExit, match="Refresh failed"):
auth_commands.auth_refresh_command(args)
after = {e["id"]: e for e in read_credential_pool("openai-codex")}
assert requests == [{"grant_type": ["refresh_token"], "refresh_token": ["fixture-refresh-1"],
"client_id": [auth_codex.CODEX_OAUTH_CLIENT_ID]}]
assert after["row0"] == before[0], (after["row0"], before[0])
target = after["row1"]
if status == 200:
assert target["access_token"] == "fixture-new-access"
assert target["refresh_token"] == "fixture-new-refresh"
assert target.get("last_error_reset_at") is None
assert target["last_status"] == "ok"
else:
# Manual grants remain in the pool on terminal failure; only
# singleton-seeded grants are removed by the existing quarantine.
assert target["last_status"] == "exhausted"
assert target["access_token"] == before[1]["access_token"]
finally:
server.shutdown()
worker.join(timeout=5)
server.server_close()
def test_add_priority_places_reauthenticated_row_in_multi_entry_pool(monkeypatch):
rows = _rows()
rows[1]["source"] = "device_code"
write_credential_pool("nous", rows)
monkeypatch.setattr(auth_commands.auth_mod, "_read_shared_nous_state", lambda: None)
monkeypatch.setattr(auth_commands.auth_mod, "_nous_device_code_login", lambda **_kwargs: {
"access_token": "fixture-renewed", "refresh_token": "fixture-renewed-refresh",
"agent_key": "fixture-agent-key", "expires_at": time.time() + 3600,
})
auth_commands.auth_add_command(SimpleNamespace(
provider="nous", auth_type="oauth", priority=0, label="reauthenticated"))
entries = read_credential_pool("nous")
assert [e["id"] for e in entries] == ["row1", "row0"]
assert entries[0]["priority"] == 0
def test_refresh_rejects_ambiguous_and_non_oauth_targets():
rows = _rows()
write_credential_pool("openai-codex", rows)
with pytest.raises(SystemExit, match="pass an index"):
auth_commands.auth_refresh_command(SimpleNamespace(provider="openai-codex", target=None))
with pytest.raises(SystemExit, match="No credential matching"):
auth_commands.auth_refresh_command(SimpleNamespace(provider="openai-codex", target="missing"))
rows[0].update(auth_type="api_key", source="manual")
write_credential_pool("openrouter", rows[:1])
with pytest.raises(SystemExit, match="not a refreshable"):
auth_commands.auth_refresh_command(SimpleNamespace(provider="openrouter", target=None))