fix: place reauthenticated credentials by their saved identity
This commit is contained in:
@@ -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:
|
||||
|
||||
70
tests/agent/test_credential_pool_operations.py
Normal file
70
tests/agent/test_credential_pool_operations.py
Normal 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"]
|
||||
104
tests/hermes_cli/test_auth_pool_operations.py
Normal file
104
tests/hermes_cli/test_auth_pool_operations.py
Normal 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))
|
||||
Reference in New Issue
Block a user