refactor(hermes_cli): auth.py collapse resolve/cache/config defensive ladders, shared _resolved_key helper

This commit is contained in:
Teknium
2026-09-02 22:25:15 -07:00
parent 9ebacffed0
commit 83cdfdcaf5
+70 -88
View File
@@ -443,17 +443,12 @@ def _auth_file_path() -> Path:
# Seat belt: under pytest, refuse to touch the real user's auth store (catches tests that forgot
# to monkeypatch HERMES_HOME or escaped the hermetic conftest). In production this is one dict
# lookup.
if os.environ.get("PYTEST_CURRENT_TEST"):
real_home_auth = (Path.home() / ".hermes" / "auth.json").resolve(strict=False)
try:
resolved = path.resolve(strict=False)
except Exception:
resolved = path
if resolved == real_home_auth:
raise RuntimeError(
f"Refusing to touch real user auth store during test run: {path}. "
"Set HERMES_HOME to a tmp_path in your test fixture, or run "
"via scripts/run_tests.sh for hermetic CI-parity env.")
if (os.environ.get("PYTEST_CURRENT_TEST")
and _same_path(path, Path.home() / ".hermes" / "auth.json")):
raise RuntimeError(
f"Refusing to touch real user auth store during test run: {path}. "
"Set HERMES_HOME to a tmp_path in your test fixture, or run "
"via scripts/run_tests.sh for hermetic CI-parity env.")
return path
@@ -484,10 +479,9 @@ def _load_global_auth_store() -> Dict[str, Any]:
str(global_path.resolve(strict=False)), global_path.stat().st_mtime_ns)
except Exception:
cache_key = None
if cache_key is not None and _global_auth_store_cache is not None:
cached_path, cached_mtime, cached_store = _global_auth_store_cache
if (cached_path, cached_mtime) == cache_key:
return cached_store
cached = _global_auth_store_cache
if cache_key is not None and cached is not None and cached[:2] == cache_key:
return cached[2]
if os.environ.get("PYTEST_CURRENT_TEST") and os.environ.get("HOME"):
real_root = Path(os.environ["HOME"]) / ".hermes" / "auth.json"
try:
@@ -502,7 +496,7 @@ def _load_global_auth_store() -> Dict[str, Any]:
_global_auth_store_cache = None
return {}
if cache_key is not None:
_global_auth_store_cache = (cache_key[0], cache_key[1], store)
_global_auth_store_cache = (*cache_key, store)
return store
@@ -517,14 +511,18 @@ def _same_path(left: Path, right: Path) -> bool:
return left == right
def _resolved_key(path: Path) -> str:
"""Canonical string for *path* (resolved when possible) used as a cache / lock-holder key."""
try:
return str(path.resolve(strict=False))
except Exception:
return str(path)
def _auth_lock_holder_for(target_path: Path) -> threading.local:
"""Return a reentrancy tracker keyed to one canonical auth-store path."""
try:
key = str(target_path.resolve(strict=False))
except Exception:
key = str(target_path)
with _auth_target_lock_holders_guard:
return _auth_target_lock_holders.setdefault(key, threading.local())
return _auth_target_lock_holders.setdefault(_resolved_key(target_path), threading.local())
def _kernel_lock(lock_file: Any, acquire: bool) -> None:
@@ -853,19 +851,17 @@ def read_credential_pool(provider_id: Optional[str] = None) -> Dict[str, Any]:
the profile fully shadows global on the next read."""
pool = _load_auth_store().get("credential_pool")
pool = pool if isinstance(pool, dict) else {}
global_store = _load_global_auth_store()
global_pool = global_store.get("credential_pool") if global_store else None
global_pool = _load_global_auth_store().get("credential_pool")
global_pool = global_pool if isinstance(global_pool, dict) else {}
if provider_id is None:
merged = dict(pool)
for gp_key, gp_entries in global_pool.items():
existing = merged.get(gp_key)
if not isinstance(gp_entries, list) or not gp_entries:
if not (isinstance(gp_entries, list) and gp_entries):
continue
if isinstance(existing, list) and existing: # profile wins whenever it has ANY entries
continue
merged[gp_key] = list(gp_entries)
if not (isinstance(existing, list) and existing): # profile wins when it has ANY entries
merged[gp_key] = list(gp_entries)
return merged
provider_entries = pool.get(provider_id)
@@ -939,8 +935,8 @@ def write_credential_pool(
sanitized = [
sanitize_borrowed_credential_payload(e, provider_id) if isinstance(e, dict) else e
for e in entries]
existing = pool.get(provider_id)
existing_list = existing if isinstance(existing, list) else []
existing_list = pool.get(provider_id)
existing_list = existing_list if isinstance(existing_list, list) else []
existing_by_id = _entry_ids(existing_list)
new_ids = set(_entry_ids(sanitized))
merged: List[Dict[str, Any]] = [
@@ -1047,8 +1043,8 @@ def _config_selects_provider(normalized: str) -> bool:
if not isinstance(moa_cfg, dict):
return False
presets = moa_cfg.get("presets")
return _moa_block_matches(moa_cfg) or (
isinstance(presets, dict) and any(_moa_block_matches(p) for p in presets.values()))
presets = presets.values() if isinstance(presets, dict) else ()
return _moa_block_matches(moa_cfg) or any(_moa_block_matches(p) for p in presets)
def _explicit_pool_entry_present(normalized: str) -> bool:
@@ -1079,8 +1075,8 @@ def _explicit_env_credentials_present(normalized: str) -> bool:
if pconfig is None:
from hermes_cli.providers import get_provider
pconfig = get_provider(normalized)
if not pconfig:
return False
if not pconfig:
return False
if pconfig.auth_type == "api_key":
return any(_env_secret(v) for v in pconfig.api_key_env_vars if v not in _IMPLICIT_ENV_VARS)
if pconfig.auth_type == "aws_sdk":
@@ -1113,8 +1109,8 @@ def _keyless_provider_has_explicit_config(normalized: str) -> bool:
from agent.vertex_adapter import has_explicit_vertex_config
return bool(has_explicit_vertex_config())
if normalized == "bedrock":
from hermes_cli.config import load_config as _load_cfg
bedrock_cfg = _load_cfg().get("bedrock")
from hermes_cli.config import load_config
bedrock_cfg = load_config().get("bedrock")
return isinstance(bedrock_cfg, dict) and bool(str(bedrock_cfg.get("region") or "").strip())
return False
@@ -1189,9 +1185,8 @@ def _get_config_hint_for_unknown_provider(provider_name: str) -> str:
lines = ["Config issue detected — run 'hermes doctor' for full diagnostics:"]
for ci in issues:
lines.append(f" [{'ERROR' if ci.severity == 'error' else 'WARNING'}] {ci.message}")
first_hint = ci.hint.splitlines()[0] if ci.hint else ""
if first_hint:
lines.append(f" → {first_hint}")
if ci.hint and ci.hint.splitlines()[0]:
lines.append(f" → {ci.hint.splitlines()[0]}")
return "\n".join(lines)
except Exception:
return ""
@@ -1329,11 +1324,9 @@ def _config_model_provider() -> Tuple[Any, Optional[str]]:
from hermes_cli.config import load_config
model_cfg = (load_config() or {}).get("model")
if isinstance(model_cfg, dict):
_cfg_provider = model_cfg.get("provider")
if isinstance(_cfg_provider, str) and _cfg_provider.strip().lower() in PROVIDER_REGISTRY:
return model_cfg, _cfg_provider.strip().lower()
return model_cfg, None
provider = model_cfg.get("provider") if isinstance(model_cfg, dict) else None
provider = provider.strip().lower() if isinstance(provider, str) else ""
return model_cfg, (provider if provider in PROVIDER_REGISTRY else None)
except Exception as e:
logger.debug("Could not read config.yaml model.provider for auto-resolution: %s", e)
return None, None
@@ -1385,10 +1378,9 @@ def resolve_provider(
return normalized
if normalized != "auto":
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")
tail = (f"\n\n{hint}" if hint else " Check 'hermes model' for available providers, "
"or run 'hermes doctor' to diagnose config issues.")
raise AuthError(f"Unknown provider '{normalized}'." + tail, code="invalid_provider")
if explicit_api_key or explicit_base_url: # one-off CLI creds always mean openrouter/custom
return "openrouter"
@@ -1473,8 +1465,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, "message": str(error), "reason": reason, "relogin_required": True,
"code": error.code if default_code is None else (error.code or default_code),
"at": datetime.now(timezone.utc).isoformat()}
@@ -1546,10 +1538,9 @@ def resolve_nous_access_token(
memoable = not insecure and ca_bundle is None
if memoable:
with _RESOLVE_TOKEN_CACHE_LOCK:
if _RESOLVE_TOKEN_CACHE is not None:
cached_at, cached_token = _RESOLVE_TOKEN_CACHE
if (time.monotonic() - cached_at) < _RESOLVE_TOKEN_CACHE_TTL_S:
return cached_token
cached = _RESOLVE_TOKEN_CACHE
if cached is not None and (time.monotonic() - cached[0]) < _RESOLVE_TOKEN_CACHE_TTL_S:
return cached[1]
def _memo(token: str) -> str:
global _RESOLVE_TOKEN_CACHE
@@ -1620,13 +1611,9 @@ _global_auth_store_cache: Optional[Tuple[str, int, Dict[str, Any]]] = None
def _auth_file_cache_key() -> Tuple[str, Optional[float]]:
auth_file = _auth_file_path()
try:
auth_file_key = str(auth_file.resolve(strict=False))
except Exception:
auth_file_key = str(auth_file)
try:
return auth_file_key, auth_file.stat().st_mtime
return _resolved_key(auth_file), auth_file.stat().st_mtime
except Exception: # missing file included: key without an mtime
return auth_file_key, None
return _resolved_key(auth_file), None
def invalidate_nous_auth_status_cache() -> None:
@@ -1645,12 +1632,9 @@ def get_nous_auth_status() -> Dict[str, Any]:
now = time.monotonic()
auth_file_key, mtime = _auth_file_cache_key()
cached = _nous_auth_status_cache
if cached is not None:
cached_at, cached_auth_file_key, cached_mtime, cached_status = cached
if (
(cached_auth_file_key, cached_mtime) == (auth_file_key, mtime)
and (now - cached_at) < _NOUS_AUTH_STATUS_CACHE_TTL):
return dict(cached_status)
if (cached is not None and cached[1:3] == (auth_file_key, mtime)
and (now - cached[0]) < _NOUS_AUTH_STATUS_CACHE_TTL):
return dict(cached[3])
status = _compute_nous_auth_status()
_nous_auth_status_cache = (now, auth_file_key, mtime, dict(status))
return status
@@ -1712,7 +1696,8 @@ def _is_terminal_refresh_error(exc: Exception, provider: str) -> bool:
_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")
_is_terminal_codex_oauth_refresh_error = partial(
_is_terminal_refresh_error, provider="openai-codex")
def _codex_pool_rate_limited_status() -> Optional[Dict[str, Any]]:
@@ -1720,8 +1705,9 @@ 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."),
@@ -1848,9 +1834,8 @@ def _external_process_spec(
command_env_vars = tuple(getattr(profile, "process_command_env_vars", ()) or ())
args_env_var = str(getattr(profile, "process_args_env_var", "") or "")
command = next((v for v in (os.getenv(var, "").strip() for var in command_env_vars) if v), "")
if not command:
command = str(getattr(profile, "process_command", "") or "")
command = (next((v for v in (os.getenv(var, "").strip() for var in command_env_vars) if v), "")
or str(getattr(profile, "process_command", "") or ""))
raw_args = os.getenv(args_env_var, "").strip() if args_env_var else ""
args = shlex.split(raw_args) if raw_args else list(getattr(profile, "process_args", ()) or [])
return command, args, base_url, shutil.which(command) if command else None, command_env_vars
@@ -1991,8 +1976,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": 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))}
"lmstudio": lambda *a: _normalize_lmstudio_runtime_base_url(_default_api_key_base_url(*a)),
"actual": lambda *a: normalize_actual_base_url(_default_api_key_base_url(*a))}
def resolve_api_key_provider_credentials(provider_id: str) -> Dict[str, Any]:
@@ -2046,8 +2031,9 @@ def resolve_external_process_provider_credentials(provider_id: str) -> Dict[str,
# 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, "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 ───────────────────────────────────────────────────────────────────
@@ -2072,11 +2058,8 @@ def _update_config_for_provider(
current_model = config.get("model")
if isinstance(current_model, dict):
model_cfg = dict(current_model)
elif _nonempty_str(current_model):
model_cfg = {"default": current_model.strip()}
else:
model_cfg = {}
model_cfg = {"default": current_model.strip()} if _nonempty_str(current_model) else {}
model_cfg["provider"] = provider_id
if inference_base_url and inference_base_url.strip():
model_cfg["base_url"] = inference_base_url.rstrip("/")
@@ -2112,7 +2095,7 @@ def _get_config_provider() -> Optional[str]:
def _should_reset_config_provider_on_logout(provider_id: Optional[str]) -> bool:
"""True when logout should reset model.provider: a registry provider that config.yaml selects."""
"""True when logout should reset model.provider (a registry provider config.yaml selects)."""
normalized = (provider_id or "").strip().lower()
return normalized in PROVIDER_REGISTRY and _get_config_provider() == normalized
@@ -2157,10 +2140,9 @@ def get_minimax_oauth_auth_status() -> Dict[str, Any]:
if not state or not state.get("access_token"):
return {"logged_in": False, "provider": "minimax-oauth"}
try:
expires_at = datetime.fromisoformat(state.get("expires_at", "")).timestamp()
token_valid = (expires_at - time.time()) > 0
token_valid = datetime.fromisoformat(state.get("expires_at", "")).timestamp() > time.time()
except Exception:
token_valid = bool(state.get("access_token"))
token_valid = True # access_token is known non-empty here
return {
"logged_in": token_valid, "provider": "minimax-oauth",
"region": state.get("region", "global"), "expires_at": state.get("expires_at")}
@@ -2186,9 +2168,9 @@ def logout_command(args) -> None:
if should_reset_config:
_reset_config_provider()
print(f"Logged out of {provider_name}.")
if should_reset_config and os.getenv("OPENROUTER_API_KEY"):
print("Hermes will use OpenRouter for inference.")
elif should_reset_config:
print("Run `hermes model` or configure an API key to use Hermes.")
else:
if not should_reset_config:
print("Model provider configuration was unchanged.")
elif os.getenv("OPENROUTER_API_KEY"):
print("Hermes will use OpenRouter for inference.")
else:
print("Run `hermes model` or configure an API key to use Hermes.")