diff --git a/hermes_cli/auth_commands.py b/hermes_cli/auth_commands.py index f77ffbec68..04d66e76f7 100644 --- a/hermes_cli/auth_commands.py +++ b/hermes_cli/auth_commands.py @@ -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 /shared/nous_auth.json # (written by any previous login) makes `hermes --profile 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} {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: diff --git a/tests/agent/test_credential_pool_operations.py b/tests/agent/test_credential_pool_operations.py new file mode 100644 index 0000000000..760cda64b4 --- /dev/null +++ b/tests/agent/test_credential_pool_operations.py @@ -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"] diff --git a/tests/hermes_cli/test_auth_pool_operations.py b/tests/hermes_cli/test_auth_pool_operations.py new file mode 100644 index 0000000000..92c226346b --- /dev/null +++ b/tests/hermes_cli/test_auth_pool_operations.py @@ -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))