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:
+64
-127
@@ -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)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user