refactor(hermes_cli): auth.py inline single-use prefix/logout/base-url wrappers, partial terminal-refresh predicates, compact dict literals

This commit is contained in:
Teknium
2026-09-02 22:20:46 -07:00
parent fd318113e6
commit 9ebacffed0
+64 -127
View File
@@ -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)