From 9ebacffed0f0b6e476d8bd7c24e758209230daaf Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 22:20:46 -0700 Subject: [PATCH] refactor(hermes_cli): auth.py inline single-use prefix/logout/base-url wrappers, partial terminal-refresh predicates, compact dict literals --- hermes_cli/auth.py | 191 +++++++++++++++------------------------------ 1 file changed, 64 insertions(+), 127 deletions(-) diff --git a/hermes_cli/auth.py b/hermes_cli/auth.py index 30f3624d23..1cc03394ff 100644 --- a/hermes_cli/auth.py +++ b/hermes_cli/auth.py @@ -24,6 +24,7 @@ import webbrowser # noqa: F401 (tests patch auth_mod.webbrowser.open; same mod from contextlib import ExitStack, contextmanager from dataclasses import dataclass, field +from functools import partial from datetime import datetime, timezone from pathlib import Path from typing import Any, Callable, Dict, FrozenSet, Iterable, List, Optional, Tuple @@ -327,28 +328,20 @@ KNOWN_PROVIDER_KEY_PREFIXES: Dict[str, tuple] = { } -def _secret_matches_declared_prefix(provider_id: str, value: str) -> bool: - """False only when the provider declares key prefixes and none match (fail-open otherwise).""" - prefixes = KNOWN_PROVIDER_KEY_PREFIXES.get(provider_id) - return not prefixes or any(value.startswith(p) for p in prefixes) - - -def _warn_malformed_secret(provider_id: str, source: str) -> None: - logger.warning( - "Ignoring %s for provider %r: value does not match the expected key " - "prefix (%s). Falling back to the next credential source. Fix or " - "remove the malformed key to silence this warning.", - source, provider_id, " or ".join(KNOWN_PROVIDER_KEY_PREFIXES.get(provider_id, ()))) - - def _usable_declared_secret(provider_id: str, value: Any, source: str) -> Optional[str]: """*value* stripped when it is a usable, prefix-valid secret; None (after warning on a provable - prefix mismatch, so it never shadows a later credential source) otherwise.""" + prefix mismatch, so it never shadows a later credential source) otherwise. Providers without a + declared prefix are fail-open.""" val = str(value or "").strip() if not has_usable_secret(val): return None - if not _secret_matches_declared_prefix(provider_id, val): - _warn_malformed_secret(provider_id, source) + prefixes = KNOWN_PROVIDER_KEY_PREFIXES.get(provider_id) + if prefixes and not any(val.startswith(p) for p in prefixes): + logger.warning( + "Ignoring %s for provider %r: value does not match the expected key " + "prefix (%s). Falling back to the next credential source. Fix or " + "remove the malformed key to silence this warning.", + source, provider_id, " or ".join(prefixes)) return None return val @@ -384,10 +377,8 @@ def _resolve_api_key_provider_secret(provider_id: str, pconfig: ProviderConfig) from agent.credential_pool import load_pool pool = load_pool(provider_id) if pool and pool.has_credentials(): - candidates = [] entry = pool.peek() - if entry is not None: - candidates.append(entry) + candidates = [entry] if entry is not None else [] try: for extra in pool.entries(): if extra is not None and all(extra is not c for c in candidates): @@ -410,9 +401,8 @@ def _resolve_api_key_provider_secret(provider_id: str, pconfig: ProviderConfig) def is_rate_limited_auth_error(error: Exception) -> bool: """True when an :class:`AuthError` is upstream rate-limiting / quota: transient, and re-authenticating cannot fix it, so callers should say "retry later", not ``hermes auth``.""" - return ( - isinstance(error, AuthError) and not error.relogin_required and error.code == CODEX_RATE_LIMITED_CODE - ) + return (isinstance(error, AuthError) and not error.relogin_required + and error.code == CODEX_RATE_LIMITED_CODE) # Entitlement failures: Nous gets a Portal-aware message; other providers a fixed generic one (or @@ -498,16 +488,14 @@ def _load_global_auth_store() -> Dict[str, Any]: cached_path, cached_mtime, cached_store = _global_auth_store_cache if (cached_path, cached_mtime) == cache_key: return cached_store - if os.environ.get("PYTEST_CURRENT_TEST"): - real_home_env = os.environ.get("HOME", "") - if real_home_env: - real_root = Path(real_home_env) / ".hermes" / "auth.json" - try: - if global_path.resolve(strict=False) == real_root.resolve(strict=False): - _global_auth_store_cache = None - return {} - except Exception: - pass + if os.environ.get("PYTEST_CURRENT_TEST") and os.environ.get("HOME"): + real_root = Path(os.environ["HOME"]) / ".hermes" / "auth.json" + try: + if global_path.resolve(strict=False) == real_root.resolve(strict=False): + _global_auth_store_cache = None + return {} + except Exception: + pass try: store = _load_auth_store(global_path) except Exception: @@ -648,9 +636,8 @@ def _load_auth_store(auth_file: Optional[Path] = None) -> Dict[str, Any]: preserved = True except Exception: preserved = False - logger.debug( - "auth: could not preserve a copy of the corrupt store at %s", corrupt_path, exc_info=True, - ) + logger.debug("auth: could not preserve a copy of the corrupt store at %s", corrupt_path, + exc_info=True) logger.warning( "auth: failed to parse %s (%s), starting with empty store. %s %s", auth_file, exc, @@ -668,18 +655,14 @@ def _load_auth_store(auth_file: Optional[Path] = None) -> Dict[str, Any]: if isinstance(raw, dict) and isinstance(raw.get("systems"), dict): # legacy "systems" format systems = raw["systems"] providers = {"nous": systems["nous_portal"]} if "nous_portal" in systems else {} - return { - **_empty_auth_store(), "providers": providers, "active_provider": "nous" if providers else None, - } + return {**_empty_auth_store(), "providers": providers, + "active_provider": "nous" if providers else None} return _empty_auth_store() def _write_private_file_atomic( - target: Path, - payload: str, - *, - replace: Optional[Callable[[Any, Any], Any]] = None, + target: Path, payload: str, *, replace: Optional[Callable[[Any, Any], Any]] = None, fsync_dir: bool = False) -> None: """Write *payload* to *target* via a 0o600 temp file + atomic rename. @@ -700,8 +683,8 @@ def _write_private_file_atomic( try: dir_fd = os.open(str(target.parent), os.O_RDONLY) except OSError: - dir_fd = None - if dir_fd is not None: + pass + else: try: os.fsync(dir_fd) finally: @@ -1223,8 +1206,7 @@ def _refuse_env_adoption_if_config_corrupt() -> None: different one. Fires ONLY on the auto path and clears itself as soon as the file parses again. """ try: - from hermes_cli.config import get_active_config_parse_failure, get_config_path - + from hermes_cli.config import get_active_config_parse_failure err = get_active_config_parse_failure() if not err: return @@ -1402,13 +1384,11 @@ def resolve_provider( if normalized in ("openrouter", "custom") or normalized in PROVIDER_REGISTRY: return normalized if normalized != "auto": - _config_hint = _get_config_hint_for_unknown_provider(normalized) - msg = f"Unknown provider '{normalized}'." - if _config_hint: - msg += f"\n\n{_config_hint}" - else: - msg += " Check 'hermes model' for available providers, or run 'hermes doctor' to diagnose config issues." - raise AuthError(msg, code="invalid_provider") + hint = _get_config_hint_for_unknown_provider(normalized) + raise AuthError( + f"Unknown provider '{normalized}'." + (f"\n\n{hint}" if hint else ( + " Check 'hermes model' for available providers, or run 'hermes doctor' to diagnose config issues.")), + code="invalid_provider") if explicit_api_key or explicit_base_url: # one-off CLI creds always mean openrouter/custom return "openrouter" @@ -1493,11 +1473,8 @@ def _last_auth_error_marker( ) -> Dict[str, Any]: """The ``last_auth_error`` record persisted when dead OAuth material is quarantined.""" return { - "provider": provider, - "code": error.code if default_code is None else (error.code or default_code), - "message": str(error), - "reason": reason, - "relogin_required": True, + "provider": provider, "code": error.code if default_code is None else (error.code or default_code), + "message": str(error), "reason": reason, "relogin_required": True, "at": datetime.now(timezone.utc).isoformat()} @@ -1609,10 +1586,8 @@ def resolve_nous_access_token( if not isinstance(refresh_token, str) or not refresh_token: raise _nous_err("Session expired and no refresh token is available.", relogin=True) - timeout = httpx.Timeout(timeout_seconds if timeout_seconds else 15.0) - with httpx.Client( - timeout=timeout, headers={"Accept": "application/json"}, verify=verify, - ) as client: + with httpx.Client(timeout=httpx.Timeout(timeout_seconds or 15.0), + headers={"Accept": "application/json"}, verify=verify) as client: refreshed = _refresh_nous_or_quarantine( client=client, auth_store=auth_store, state=state, portal_base_url=portal_base_url, client_id=client_id, refresh_token=refresh_token, @@ -1735,16 +1710,9 @@ def _is_terminal_refresh_error(exc: Exception, provider: str) -> bool: return OAUTH_PROVIDER_FLOWS[provider].is_terminal_refresh_error(exc) -def _is_terminal_nous_refresh_error(exc: Exception) -> bool: - return _is_terminal_refresh_error(exc, "nous") - - -def _is_terminal_xai_oauth_refresh_error(exc: Exception) -> bool: - return _is_terminal_refresh_error(exc, "xai-oauth") - - -def _is_terminal_codex_oauth_refresh_error(exc: Exception) -> bool: - return _is_terminal_refresh_error(exc, "openai-codex") +_is_terminal_nous_refresh_error = partial(_is_terminal_refresh_error, provider="nous") +_is_terminal_xai_oauth_refresh_error = partial(_is_terminal_refresh_error, provider="xai-oauth") +_is_terminal_codex_oauth_refresh_error = partial(_is_terminal_refresh_error, provider="openai-codex") def _codex_pool_rate_limited_status() -> Optional[Dict[str, Any]]: @@ -1752,16 +1720,11 @@ def _codex_pool_rate_limited_status() -> Optional[Dict[str, Any]]: if not rate_limit: return None return { - "logged_in": True, - "auth_store": str(_auth_file_path()), - "last_refresh": rate_limit.get("last_refresh"), - "auth_mode": "chatgpt", - "source": f"pool:{rate_limit.get('label') or 'unknown'}", - "rate_limited": True, + "logged_in": True, "auth_store": str(_auth_file_path()), "last_refresh": rate_limit.get("last_refresh"), + "auth_mode": "chatgpt", "source": f"pool:{rate_limit.get('label') or 'unknown'}", "rate_limited": True, "error_code": CODEX_RATE_LIMITED_CODE, - "error": ( - rate_limit.get("message") - or "Codex provider quota exhausted; retry after the usage limit resets."), + "error": (rate_limit.get("message") + or "Codex provider quota exhausted; retry after the usage limit resets."), "reset_at": rate_limit.get("reset_at")} @@ -1817,11 +1780,9 @@ def get_api_key_provider_status(provider_id: str) -> Dict[str, Any]: base_url = normalize_actual_base_url(base_url) actual_local_noauth = not api_key and is_actual_local_base_url(base_url) configured = bool(api_key) or actual_local_noauth - status.update( - configured=configured, base_url=base_url, - key_source=key_source or ("local-offline" if actual_local_noauth else ""), - logged_in=configured, # compat with the OAuth status shape - ) + status.update( # logged_in mirrors configured for compat with the OAuth status shape + configured=configured, base_url=base_url, logged_in=configured, + key_source=key_source or ("local-offline" if actual_local_noauth else "")) return status @@ -1967,22 +1928,18 @@ def _get_azure_foundry_auth_status() -> Dict[str, Any]: cfg = {} model_cfg = cfg.get("model") if isinstance(cfg, dict) else None - auth_mode = "api_key" - base_url = "" - if isinstance(model_cfg, dict): - auth_mode = str(model_cfg.get("auth_mode") or "api_key").strip().lower() or "api_key" - base_url = str(model_cfg.get("base_url") or "").strip() + if not isinstance(model_cfg, dict): + model_cfg = {} + auth_mode = str(model_cfg.get("auth_mode") or "api_key").strip().lower() or "api_key" info["auth_mode"] = auth_mode - info["base_url"] = base_url + info["base_url"] = str(model_cfg.get("base_url") or "").strip() if auth_mode == "entra_id": try: from agent.azure_identity_adapter import ( EntraIdentityConfig, SCOPE_AI_AZURE_DEFAULT, has_azure_identity_installed) installed = has_azure_identity_installed() - entra_cfg = {} - if isinstance(model_cfg, dict) and isinstance(model_cfg.get("entra"), dict): - entra_cfg = model_cfg["entra"] + entra_cfg = model_cfg["entra"] if isinstance(model_cfg.get("entra"), dict) else {} identity_config = EntraIdentityConfig.from_dict(entra_cfg, default_scope=SCOPE_AI_AZURE_DEFAULT) info.update( azure_identity_installed=installed, scope=identity_config.scope, credential_probe="not_run", @@ -2027,14 +1984,6 @@ def _copilot_runtime_base_url(api_key: str, default: str, env_url: str) -> str: return base_url -def _lmstudio_runtime_base_url(api_key: str, default: str, env_url: str) -> str: - return _normalize_lmstudio_runtime_base_url(_default_api_key_base_url(api_key, default, env_url)) - - -def _actual_runtime_base_url(api_key: str, default: str, env_url: str) -> str: - return normalize_actual_base_url(_default_api_key_base_url(api_key, default, env_url)) - - # Providers whose runtime base URL is not simply env-override-or-registry-default: # ``(api_key, registry_default, env_override) -> base_url``. _API_KEY_BASE_URL_RESOLVERS: Dict[str, Callable[[str, str, str], str]] = { @@ -2042,8 +1991,8 @@ _API_KEY_BASE_URL_RESOLVERS: Dict[str, Callable[[str, str, str], str]] = { "kimi-coding-cn": _resolve_kimi_base_url, "zai": _resolve_zai_base_url, "copilot": _copilot_runtime_base_url, - "lmstudio": _lmstudio_runtime_base_url, - "actual": _actual_runtime_base_url} + "lmstudio": lambda k, d, e: _normalize_lmstudio_runtime_base_url(_default_api_key_base_url(k, d, e)), + "actual": lambda k, d, e: normalize_actual_base_url(_default_api_key_base_url(k, d, e))} def resolve_api_key_provider_credentials(provider_id: str) -> Dict[str, Any]: @@ -2094,15 +2043,11 @@ def resolve_external_process_provider_credentials(provider_id: str) -> Dict[str, f"'{command or '(none configured)'}'. Install it{_hint}.", provider=provider_id, code="missing_external_process_cli") + # api_key is a placeholder: the subprocess owns real auth. Keyed on the provider id so each + # external-process provider gets a distinct value. return { - "provider": provider_id, - # Placeholder credential: the subprocess owns real auth. Keyed on the provider id so each - # external-process provider gets a distinct value. - "api_key": pconfig.id or provider_id, - "base_url": base_url.rstrip("/"), - "command": resolved_command or command, - "args": args, - "source": "process"} + "provider": provider_id, "api_key": pconfig.id or provider_id, "base_url": base_url.rstrip("/"), + "command": resolved_command or command, "args": args, "source": "process"} # ── CLI Commands — login / logout ─────────────────────────────────────────────────────────────────── @@ -2166,17 +2111,10 @@ def _get_config_provider() -> Optional[str]: return (provider.strip().lower() or None) if isinstance(provider, str) else None -def _config_provider_matches(provider_id: Optional[str]) -> bool: - """Return True when config.yaml currently selects *provider_id*.""" - return bool(provider_id) and _get_config_provider() == provider_id.strip().lower() - - def _should_reset_config_provider_on_logout(provider_id: Optional[str]) -> bool: - """Return True when logout should reset the model provider config.""" - if not provider_id: - return False - normalized = provider_id.strip().lower() - return normalized in PROVIDER_REGISTRY and _config_provider_matches(normalized) + """True when logout should reset model.provider: a registry provider that config.yaml selects.""" + normalized = (provider_id or "").strip().lower() + return normalized in PROVIDER_REGISTRY and _get_config_provider() == normalized def _logout_default_provider_from_config() -> Optional[str]: @@ -2208,9 +2146,8 @@ def _reset_config_provider() -> Path: def login_command(args) -> None: """Deprecated: use 'hermes model' or 'hermes setup' instead.""" - print("The 'hermes login' command has been removed.") - print("Use 'hermes auth' to manage credentials,") - print("'hermes model' to select a provider, or 'hermes setup' for full setup.") + print("The 'hermes login' command has been removed.\nUse 'hermes auth' to manage credentials,\n" + "'hermes model' to select a provider, or 'hermes setup' for full setup.") raise SystemExit(0)