refactor(auth): split OAuth-grant hygiene, device-flow helpers, model picker and Kimi/Z.AI detection into modules

This commit is contained in:
Teknium
2026-09-02 15:48:08 -07:00
parent 53a24170ab
commit d32014a4a6
5 changed files with 1587 additions and 1436 deletions

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,415 @@
"""Shared device-code / browser / TLS helpers for interactive OAuth logins.
Split out of ``hermes_cli/auth.py``; every moved name is re-imported there, so
``hermes_cli.auth.<name>`` keeps resolving (and monkeypatching) as before. Origin-internal
helpers are imported lazily inside each function (no import cycle; patches on
``hermes_cli.auth.<helper>`` still intercept).
"""
from __future__ import annotations
import logging
import os
import ssl
import sys
import time
import webbrowser
from pathlib import Path
from typing import Any, Callable, Dict, Optional
from urllib.parse import urlparse
from hermes_cli.auth_constants import (
AuthError,
DEFAULT_NOUS_PORTAL_URL,
DEVICE_AUTH_POLL_INTERVAL_CAP_SECONDS,
DEVICE_CODE_GRANT_TYPE,
OAUTH_OVER_SSH_DOCS_URL,
httpx,
)
from utils import is_truthy_value
# Log-record parity with the origin module (caplog tests pin "hermes_cli.auth").
logger = logging.getLogger("hermes_cli.auth")
def _is_remote_session() -> bool:
"""Detect environments where loopback OAuth can't reach the local browser.
These environments typically don't set ``SSH_CLIENT`` / ``SSH_TTY``, so the SSH-only check left
them with no guidance and no fallback.
"""
if os.getenv("SSH_CLIENT") or os.getenv("SSH_TTY"):
return True
# Browser-only remote IDEs / cloud shells. Keep this list narrow
# (well-known, documented env vars set by the host platform) so
# we don't falsely trip on a developer's local shell.
for var in (
"CLOUD_SHELL", # GCP Cloud Shell
"CODESPACES", # GitHub Codespaces
"CODESPACE_NAME", # GitHub Codespaces (alt)
"GITPOD_WORKSPACE_ID", # Gitpod
"REPL_ID", # Replit
"STACKBLITZ", # StackBlitz
):
if os.getenv(var):
return True
return False
def _can_open_graphical_browser() -> bool:
"""Return True only when a *graphical* browser is likely to open.
``webbrowser.open()`` resolves to whatever the platform offers, and on a headless / CLI-only
Linux box with no GUI browser installed that is often a text-mode browser (w3m/lynx/links) which
launches inside the terminal and takes over the user's session.
Heuristics: * Respect ``$BROWSER`` — if it names a known console browser, refuse. * On Linux,
require a display server (``$DISPLAY`` / ``$WAYLAND_DISPLAY``) unless ``$BROWSER`` points at
something graphical; no display server almost always means no GUI browser.
"""
from hermes_cli.auth import _CONSOLE_BROWSER_NAMES
import webbrowser as _webbrowser
def _names_console_browser(value: str) -> bool:
token = value.strip().split()[0] if value.strip() else ""
base = os.path.basename(token).lower()
return base in _CONSOLE_BROWSER_NAMES
browser_env = os.environ.get("BROWSER", "")
if browser_env and _names_console_browser(browser_env):
return False
if sys.platform.startswith("linux"):
has_display = bool(
os.environ.get("DISPLAY") or os.environ.get("WAYLAND_DISPLAY")
)
# An explicit graphical $BROWSER can work without $DISPLAY in odd
# setups, but a console $BROWSER already returned False above, so the
# only way to reach here with a $BROWSER set is a graphical one.
if not has_display and not browser_env:
return False
try:
controller = _webbrowser.get()
except Exception:
# No browser resolvable at all → definitely don't auto-open.
return False
candidate = (
getattr(controller, "name", "")
or getattr(controller, "basename", "")
or ""
)
return not (candidate and _names_console_browser(candidate))
def _ssh_user_at_host() -> str:
"""Return best-effort 'user@hostname' for the SSH tunnel hint command.
Falls back to placeholder tokens when the values cannot be determined so the hint is always
syntactically valid even if not copy-pasteable.
"""
try:
import socket as _socket
hostname = _socket.gethostname() or "<this-host>"
except OSError:
hostname = "<this-host>"
user = os.getenv("USER") or os.getenv("LOGNAME") or "<user>"
return f"{user}@{hostname}"
def _print_loopback_ssh_hint(redirect_uri: str, *, docs_url: str | None = None) -> None:
"""Print an SSH tunnel hint when running a loopback-redirect OAuth flow on a remote host. The auth
server (Spotify, MCP servers, ...) will redirect the user's browser to
``127.0.0.1:<port>/callback``. If the browser is on a different machine than the loopback
listener (the usual SSH case), the redirect can't reach the listener without a local port
forward.
"""
from hermes_cli.auth import _is_remote_session
if not _is_remote_session():
return
try:
parsed = urlparse(redirect_uri)
except Exception:
return
host = parsed.hostname or ""
port = parsed.port
if host not in {"127.0.0.1", "::1", "localhost"} or not port:
return
divider = "-" * 60
print()
print(divider)
print("Remote session detected — SSH tunnel required")
print(divider)
print(f"Hermes is waiting for the OAuth callback on {redirect_uri}")
print("but your browser is on a different machine. Run this command")
print("in a NEW terminal on your local machine BEFORE opening the URL:")
print()
print(f" ssh -N -L {port}:127.0.0.1:{port} {_ssh_user_at_host()}")
print()
print("Then open the authorize URL above in your local browser.")
if docs_url:
print(f"Provider docs: {docs_url}")
print(f"SSH/jump-box guide: {OAUTH_OVER_SSH_DOCS_URL}")
print(divider)
print()
def _default_verify() -> bool | ssl.SSLContext:
"""Platform-aware default SSL verify for httpx clients.
On macOS with Homebrew Python, the system OpenSSL cannot locate the system trust store and valid
public certs fail verification. When certifi is importable we pin its bundle explicitly;
elsewhere we defer to httpx's built-in default (certifi via its own dependency). Mirrors the
weixin fix in 3a0ec1d93.
"""
if sys.platform == "darwin":
try:
import certifi
return ssl.create_default_context(cafile=certifi.where())
except ImportError:
pass
return True
def _resolve_verify(
*,
insecure: Optional[bool] = None,
ca_bundle: Optional[str] = None,
auth_state: Optional[Dict[str, Any]] = None,
) -> bool | ssl.SSLContext:
from hermes_cli.auth import _default_verify
tls_state = auth_state.get("tls") if isinstance(auth_state, dict) else {}
tls_state = tls_state if isinstance(tls_state, dict) else {}
effective_insecure = (
is_truthy_value(insecure, default=False) if insecure is not None
else is_truthy_value(tls_state.get("insecure", False), default=False)
)
effective_ca = (
ca_bundle
or tls_state.get("ca_bundle")
or os.getenv("HERMES_CA_BUNDLE")
or os.getenv("SSL_CERT_FILE")
or os.getenv("REQUESTS_CA_BUNDLE")
)
if effective_insecure:
return False
if effective_ca:
ca_path = str(effective_ca)
if not os.path.isfile(ca_path):
logger.warning(
"CA bundle path does not exist: %s — falling back to default certificates",
ca_path,
)
return _default_verify()
return ssl.create_default_context(cafile=ca_path)
return _default_verify()
def _request_device_code(
client: httpx.Client,
portal_base_url: str,
client_id: str,
scope: Optional[str],
) -> Dict[str, Any]:
"""POST to the device code endpoint. Returns device_code, user_code, etc."""
response = client.post(
f"{portal_base_url}/api/oauth/device/code",
data={
"client_id": client_id,
**({"scope": scope} if scope else {}),
},
)
response.raise_for_status()
data = response.json()
required_fields = [
"device_code", "user_code", "verification_uri",
"verification_uri_complete", "expires_in", "interval",
]
missing = [f for f in required_fields if f not in data]
if missing:
raise ValueError(f"Device code response missing fields: {', '.join(missing)}")
return data
def _nous_device_auth_timeout_message(portal_base_url: str) -> str:
"""Actionable timeout text for Nous device-code login failures.
A bare "timed out" gives the user nothing to act on; the usual cause is Portal sign-in failing
in the opened browser tab, so point at the Portal login page and the retry command.
"""
portal = (portal_base_url or DEFAULT_NOUS_PORTAL_URL).rstrip("/")
return (
"Timed out waiting for device authorization.\n"
" Portal sign-in is required before the device code can be approved.\n"
" If the browser showed a CAPTCHA / 'You did not pass CAPTCHA' error,\n"
" finish signing in at the Portal in a normal browser tab, then retry:\n"
" hermes portal\n"
f" Portal login: {portal}/login"
)
def _print_device_code_instructions(
verification_url: str,
user_code: str,
*,
open_browser: bool,
failure_dash: str = "--",
swallow_open_errors: bool = False,
) -> None:
"""Print the shared "To continue" device-code block and optionally open the browser.
Callers decide *whether* to open (remote-session / graphical-browser gating differs per
provider); the wording of the fallback hint is parameterized so each provider keeps its
historical dash style.
"""
print()
print("To continue:")
print(f" 1. Open: {verification_url}")
print(f" 2. If prompted, enter code: {user_code}")
if not open_browser:
return
if swallow_open_errors:
try:
opened = webbrowser.open(verification_url)
except Exception:
opened = False
else:
opened = webbrowser.open(verification_url)
if opened:
print(" (Opened browser for verification)")
else:
print(f" Could not open browser automatically {failure_dash} use the URL above.")
def _poll_device_token_generic(
post: Callable[[], "httpx.Response"],
*,
expires_in: int,
poll_interval: int,
validate_success: Callable[[Dict[str, Any]], None],
on_non_json_error: Callable[["httpx.Response"], Exception],
on_error: Callable[["httpx.Response", Dict[str, Any]], Exception],
on_timeout: Callable[[], Exception],
) -> Dict[str, Any]:
"""RFC 8628 device-code polling loop shared by the Nous and xAI flows.
``authorization_pending`` sleeps and retries; ``slow_down`` grows the interval by 1s (cap 30s).
Every other error, a non-JSON error body, and the deadline are turned into provider-specific
exceptions by the supplied factories so each caller keeps its exact error contract.
"""
deadline = time.monotonic() + max(1, expires_in)
current_interval = poll_interval
while time.monotonic() < deadline:
response = post()
if response.status_code == 200:
payload = response.json()
validate_success(payload)
return payload
try:
error_payload = response.json()
except Exception:
response.raise_for_status()
raise on_non_json_error(response)
error_code = str(error_payload.get("error") or "")
if error_code == "authorization_pending":
time.sleep(current_interval)
continue
if error_code == "slow_down":
current_interval = min(current_interval + 1, 30)
time.sleep(current_interval)
continue
raise on_error(response, error_payload)
raise on_timeout()
def _poll_for_token(
client: httpx.Client,
portal_base_url: str,
client_id: str,
device_code: str,
expires_in: int,
poll_interval: int,
) -> Dict[str, Any]:
"""Poll the Nous token endpoint until the user approves or the code expires."""
def _validate(payload: Dict[str, Any]) -> None:
if "access_token" not in payload:
raise ValueError("Token response did not include access_token")
def _error(_response, error_payload) -> Exception:
error_code = error_payload.get("error", "")
description = error_payload.get("error_description") or "Unknown authentication error"
return RuntimeError(f"{error_code}: {description}")
return _poll_device_token_generic(
lambda: client.post(
f"{portal_base_url}/api/oauth/token",
data={
"grant_type": DEVICE_CODE_GRANT_TYPE,
"client_id": client_id,
"device_code": device_code,
},
),
expires_in=expires_in,
poll_interval=max(1, min(poll_interval, DEVICE_AUTH_POLL_INTERVAL_CAP_SECONDS)),
validate_success=_validate,
on_non_json_error=lambda _r: RuntimeError("Token endpoint returned a non-JSON error response"),
on_error=_error,
# Enriched at the SOURCE so every caller inherits the guidance:
# the CLI login (_nous_device_code_login) and the dashboard/desktop
# poller (web_server._nous_poller, which surfaces str(e) to the UI).
on_timeout=lambda: TimeoutError(_nous_device_auth_timeout_message(portal_base_url)),
)
def _prompt_yes_no(prompt: str, *, default: str) -> bool:
"""``input()`` a [Y/n]-style question; EOF/Ctrl-C count as *default*."""
try:
answer = input(prompt).strip().lower()
except (EOFError, KeyboardInterrupt):
answer = default
return answer in {"", "y", "yes"} if default == "y" else answer in {"y", "yes"}
def _print_login_success(provider_id: str, config_path: Path, *, show_auth_state: bool = False) -> None:
print()
print("Login successful!")
if show_auth_state:
from hermes_constants import display_hermes_home as _dhh
print(f" Auth state: {_dhh()}/auth.json")
print(f" Config updated: {config_path} (model.provider={provider_id})")
def _offer_existing_oauth_credentials(
provider_id: str,
*,
resolve: Callable[[], Dict[str, Any]],
is_expiring: Callable[[str, int], bool],
display_name: str,
default_base_url: str,
expired_notice: Optional[str] = None,
) -> bool:
"""Offer to reuse still-valid stored OAuth credentials. Returns True when the user accepted.
*resolve* attempts a refresh, so a resolved token should be valid — but double-check the
expiry before telling the user "Login successful!".
"""
from hermes_cli.auth import _update_config_for_provider
try:
existing = resolve()
api_key = existing.get("api_key", "")
if isinstance(api_key, str) and api_key and not is_expiring(api_key, 60):
print(f"Existing {display_name} credentials found in Hermes auth store.")
if _prompt_yes_no("Use existing credentials? [Y/n]: ", default="y"):
config_path = _update_config_for_provider(
provider_id, existing.get("base_url", default_base_url),
)
_print_login_success(provider_id, config_path)
return True
elif expired_notice:
print(expired_notice)
except AuthError:
pass
return False

View File

@@ -0,0 +1,387 @@
"""Interactive model picker used after OAuth login.
Split out of ``hermes_cli/auth.py``; every moved name is re-imported there, so
``hermes_cli.auth.<name>`` keeps resolving (and monkeypatching) as before. Origin-internal
helpers are imported lazily inside each function (no import cycle; patches on
``hermes_cli.auth.<helper>`` still intercept).
"""
from __future__ import annotations
import logging
import subprocess
from typing import Dict, List, Optional
from hermes_cli.auth_constants import DEFAULT_NOUS_PORTAL_URL
# Log-record parity with the origin module (caplog tests pin "hermes_cli.auth").
logger = logging.getLogger("hermes_cli.auth")
def _confirm_selection_guards(
model_id: str,
*,
provider: str = "",
base_url: str = "",
api_key: str = "",
include_kinds: Optional[List[str]] = None,
) -> bool:
"""Prompt before saving a model that trips any selection guard.
Runs the unified guard registry (cost, data-policy, future guards) and shows one [y/N] confirm
listing every warning that fired. Returns True to proceed, False to cancel.
"""
try:
from hermes_cli.model_selection_guards import (
combined_message,
selection_warnings,
)
warnings = selection_warnings(
model_id,
provider=provider,
base_url=base_url,
api_key=api_key,
include_kinds=include_kinds,
)
except Exception:
warnings = []
if not warnings:
return True
print()
print("=" * 72)
print(combined_message(warnings))
print("=" * 72)
try:
response = input("Switch anyway? [y/N]: ").strip().lower()
except (KeyboardInterrupt, EOFError):
print()
return False
return response in {"y", "yes"}
class _ModelPickerRows:
"""Column-aligned model rows (name + $/Mtok prices + Nous sale chrome) for the model picker.
Sale chrome (★ / -N% / was) is drawn as curses/ANSI segments (yellow % / dim "was"), not baked
into one plain string — curses addnstr would otherwise render escape bytes literally.
"""
def __init__(
self,
all_models: List[str],
pricing: Optional[Dict[str, Dict[str, str]]],
*,
current_model: str,
sale_chrome: bool,
) -> None:
from hermes_cli.models import _format_price_per_mtok, compute_sale_discount
self.current_model = current_model
self.has_pricing = bool(pricing and any(pricing.get(m) for m in all_models))
# Leave room for a leading "★ " on sale rows (Nous only).
name_pad = 3 if sale_chrome else 2
self.name_col = (
max((len(m) for m in all_models), default=0) + name_pad
if self.has_pricing
else 0
)
# (inp, out, cache, pct|None, was_inp, was_out)
self._price_cache: dict[str, tuple[str, str, str, int | None, str, str]] = {}
self.price_col = 3 # minimum width
self.cache_col = 0 # only set if any model has cache pricing
self.has_cache = False
self.any_on_sale = False
if not self.has_pricing:
return
for mid in all_models:
p = pricing.get(mid) # type: ignore[union-attr]
pct: int | None = None
was_inp = was_out = ""
if p:
inp = _format_price_per_mtok(p.get("prompt", ""))
out = _format_price_per_mtok(p.get("completion", ""))
cache_read = p.get("input_cache_read", "")
cache = _format_price_per_mtok(cache_read) if cache_read else ""
if cache:
self.has_cache = True
if sale_chrome:
sale = compute_sale_discount(
p.get("prompt", ""),
p.get("completion", ""),
p.get("original"),
)
if sale is not None:
self.any_on_sale = True
pct, was_prompt_raw, was_out_raw = sale
# Natively-free models (no gateway original) carry
# empty was_* raws — leave them empty so the row
# shows bare "-100%" with no "was ?/?" suffix.
if was_prompt_raw == "" and was_out_raw == "":
was_inp = was_out = ""
else:
was_inp = (
_format_price_per_mtok(was_prompt_raw)
if was_prompt_raw != ""
else "?"
)
was_out = (
_format_price_per_mtok(was_out_raw)
if was_out_raw != ""
else "?"
)
else:
inp, out, cache = "", "", ""
self._price_cache[mid] = (inp, out, cache, pct, was_inp, was_out)
self.price_col = max(self.price_col, len(inp), len(out))
self.cache_col = max(self.cache_col, len(cache))
if self.has_cache:
self.cache_col = max(self.cache_col, 5) # minimum: "Cache" header
def segments(self, mid: str) -> list[tuple[str, str | None]]:
"""Build a rich radiolist row: yellow ★/% , dim was, plain prices."""
if not self.has_pricing:
segs: list[tuple[str, str | None]] = [(mid, None)]
if mid == self.current_model:
segs.append((" ← currently in use", None))
return segs
inp, out, cache, pct, was_inp, was_out = self._price_cache.get(
mid, ("", "", "", None, "", "")
)
on_sale = pct is not None
# Reserve 2 columns for "★ " so sale and non-sale names share alignment.
star_w = 2
if on_sale:
name_segs: list[tuple[str, str | None]] = [
("★ ", "yellow"),
(f"{mid:<{self.name_col - star_w}}", None),
]
else:
name_segs = [(f"{mid:<{self.name_col}}", None)]
price_part = f" {inp:>{self.price_col}} {out:>{self.price_col}}"
if self.has_cache:
price_part += f" {cache:>{self.cache_col}}"
segs = [*name_segs, (price_part, None)]
if on_sale:
segs.append((f" -{pct}%", "yellow"))
if was_inp or was_out:
segs.append((f" was {was_inp}/{was_out}", "dim"))
if mid == self.current_model:
segs.append((" ← currently in use", None))
return segs
def label(self, mid: str) -> str:
return "".join(text for text, _style in self.segments(mid))
def menu_title(self) -> str:
"""``Select default model:`` plus an aligned pricing header hint when priced."""
title = "Select default model:"
if self.has_pricing:
# Align the header with the model column.
# Each choice is " {label}" (2 spaces) and we prepend
# a 3-char cursor region ("-> " or " "), so content starts at col 5.
pad = " " * 5
header = f"\n{pad}{'':>{self.name_col}} {'In':>{self.price_col}} {'Out':>{self.price_col}}"
if self.has_cache:
header += f" {'Cache':>{self.cache_col}}"
# Legend lives on the column-header line so it reads as a key
# (★ = on sale), not a fake menu row.
title += header + " $/Mtok"
if self.any_on_sale:
title += " ★ = on sale"
return title
def _prompt_model_selection(
model_ids: List[str],
current_model: str = "",
pricing: Optional[Dict[str, Dict[str, str]]] = None,
unavailable_models: Optional[List[str]] = None,
portal_url: str = "",
unavailable_message: str = "",
confirm_provider: str = "",
confirm_base_url: str = "",
confirm_api_key: str = "",
) -> Optional[str]:
"""Interactive model picker; current_model listed first. Returns the chosen model ID or None.
With *pricing* (``{model_id: {prompt, completion}}``) a compact price column is shown; models in
*unavailable_models* render grayed out and unselectable with an upgrade link to *portal_url*.
"""
from hermes_cli.cli_output import line_input
_unavailable = unavailable_models or []
# Sale chrome (★ / -N% / was) is Nous Portal-only — never for OpenRouter
# or other providers even if pricing.original is somehow present.
sale_chrome = (confirm_provider or "").strip().lower() == "nous"
def _confirmed_selection(mid: str) -> Optional[str]:
if not mid:
return None
# Unified guard registry (hermes_cli.model_selection_guards): the cost
# guard only runs when a provider is known (pricing lookups need one);
# id-keyed guards like the data-policy guard always run — they must
# fire even via a custom endpoint or gateway.
_kinds = None if confirm_provider else ["data_policy"]
if not _confirm_selection_guards(
mid,
provider=confirm_provider,
base_url=confirm_base_url,
api_key=confirm_api_key,
include_kinds=_kinds,
):
return None
return mid
# Reorder: current model first, then the rest (deduplicated)
ordered = []
if current_model and current_model in model_ids:
ordered.append(current_model)
for mid in model_ids:
if mid not in ordered:
ordered.append(mid)
# All models for column-width computation (selectable + unavailable)
rows = _ModelPickerRows(
list(ordered) + list(_unavailable), pricing,
current_model=current_model, sale_chrome=sale_chrome,
)
_DIM = "\033[2m"
_RESET = "\033[0m"
# Default cursor on the current model (index 0 if it was reordered to top)
default_idx = 0
menu_title = rows.menu_title()
_upgrade_url = (portal_url or DEFAULT_NOUS_PORTAL_URL).rstrip("/")
# Try arrow-key menu first, fall back to number input.
try:
from hermes_cli.curses_ui import curses_radiolist
choices = [rows.segments(mid) for mid in ordered]
choices.append("Enter custom model name")
choices.append("Skip (keep current)")
unavailable_footer = unavailable_message.strip()
if not unavailable_footer and _unavailable:
unavailable_footer = f"Upgrade at {_upgrade_url} for paid models"
# The pricing column header (and any unavailable-models block) is shown
# as a multi-line description above the list so it survives the curses
# screen clear. menu_title already embeds the aligned price header.
desc_lines: list[str] = []
if rows.has_pricing:
# menu_title is "Select default model:\n<pad><header> $/Mtok\n…"
# Keep only the header/legend portion for the description.
header_part = menu_title.split("\n", 1)
if len(header_part) > 1:
desc_lines.extend(header_part[1].splitlines())
if _unavailable:
for mid in _unavailable:
desc_lines.append(f" {rows.label(mid)}")
desc_lines.append(f" ── {unavailable_footer} ──")
description = "\n".join(desc_lines) if desc_lines else None
# Search haystacks keep pricing labels visible while adding aliases
# for brand-less wire ids (e.g. Kimi Coding `k3` ↔ query "kimi").
from hermes_cli.model_search import model_search_text
model_search_labels = []
for mid in ordered:
label = rows.label(mid)
haystack = model_search_text(mid)
# model_search_text always starts with the wire id; only append when
# aliases add tokens beyond the bare id already in the label.
model_search_labels.append(
label if haystack == mid else f"{label} {haystack}"
)
model_search_labels.append("Enter custom model name")
model_search_labels.append("Skip (keep current)")
idx = curses_radiolist(
"Select default model:",
choices,
selected=default_idx,
cancel_returns=-1,
description=description,
searchable=True,
search_labels=model_search_labels,
)
if idx < 0:
return None
print()
if idx < len(ordered):
return _confirmed_selection(ordered[idx])
elif idx == len(ordered):
try:
custom = line_input("Enter model name: ").strip()
except (EOFError, KeyboardInterrupt):
return None
return _confirmed_selection(custom) if custom else None
return None
except (ImportError, NotImplementedError, OSError, subprocess.SubprocessError):
pass
# Fallback: numbered list (ANSI colors for sale chrome)
from hermes_cli.curses_ui import format_radio_item_ansi
from hermes_cli.colors import Colors, color
for line in menu_title.splitlines():
if "★" in line:
print(line.replace("★", color("★", Colors.YELLOW), 1))
else:
print(line)
num_width = len(str(len(ordered) + 2))
for i, mid in enumerate(ordered, 1):
print(f" {i:>{num_width}}. {format_radio_item_ansi(rows.segments(mid))}")
n = len(ordered)
print(f" {n + 1:>{num_width}}. Enter custom model name")
print(f" {n + 2:>{num_width}}. Skip (keep current)")
if _unavailable:
unavailable_footer = unavailable_message.strip() or (
f"Unavailable models (requires paid tier — upgrade at {_upgrade_url})"
)
print()
print(f" {_DIM}── {unavailable_footer} ──{_RESET}")
for mid in _unavailable:
print(f" {'':>{num_width}} {_DIM}{rows.label(mid)}{_RESET}")
print()
while True:
try:
choice = input(f"Choice [1-{n + 2}] (default: skip): ").strip()
if not choice:
return None
idx = int(choice)
if 1 <= idx <= n:
return _confirmed_selection(ordered[idx - 1])
elif idx == n + 1:
custom = line_input("Enter model name: ").strip()
return _confirmed_selection(custom) if custom else None
elif idx == n + 2:
return None
print(f"Please enter 1-{n + 2}")
except ValueError:
print("Please enter a number")
except (KeyboardInterrupt, EOFError):
return None
def _save_model_choice(model_id: str) -> None:
"""Save the selected model to config.yaml (single source of truth).
The model is stored in config.yaml only — NOT in .env. This avoids conflicts in multi-agent
setups where env vars would stomp each other.
"""
from hermes_cli.config import save_config, load_config
config = load_config()
# Always use dict format so provider/base_url can be stored alongside
if isinstance(config.get("model"), dict):
config["model"]["default"] = model_id
else:
config["model"] = {"default": model_id}
save_config(config)

View File

@@ -0,0 +1,513 @@
"""Single-use OAuth grant hygiene: strip cloned grants from profiles, heal forked grants.
Split out of ``hermes_cli/auth.py``; every moved name is re-imported there, so
``hermes_cli.auth.<name>`` keeps resolving (and monkeypatching) as before. Origin-internal
helpers are imported lazily inside each function (no import cycle; patches on
``hermes_cli.auth.<helper>`` still intercept).
"""
from __future__ import annotations
import logging
import json
import os
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
from hermes_cli.auth_nous import _decode_jwt_claims
# Log-record parity with the origin module (caplog tests pin "hermes_cli.auth").
logger = logging.getLogger("hermes_cli.auth")
# Pool providers whose OAuth refresh tokens are SINGLE-USE: redeeming the
# refresh token rotates the pair and revokes the old one. A grant forked into
# two auth.json files is therefore not two credentials but one credential with
# two owners — the first owner to refresh strands the other with
# ``invalid_grant`` / ``refresh_token_reused`` (#100339; same class as the
# ``providers.<id>`` write-through hazard in #48415 / #43589). Profiles must
# never receive a copy of these grants: ONE grant lives at the global root and
# named profiles read it through the ``read_credential_pool`` root fallback.
SINGLE_USE_REFRESH_POOL_PROVIDERS = frozenset({
"anthropic",
"openai-codex",
"xai-oauth",
})
# Singleton credential files that hold the same single-use grants outside
# ``auth.json``. Copying one into a profile re-seeds a forked pool row on the
# profile's next ``load_pool()``.
SINGLE_USE_OAUTH_SINGLETON_FILES = (".anthropic_oauth.json",)
def _is_oauth_pool_payload(entry: Any) -> bool:
if not isinstance(entry, dict):
return False
auth_type = str(entry.get("auth_type") or "").strip().lower()
if auth_type == "oauth":
return True
# Legacy rows predating ``auth_type``: an Anthropic OAuth access token or
# any row carrying a refresh token is an OAuth grant.
if str(entry.get("refresh_token") or "").strip():
return True
return str(entry.get("access_token") or "").startswith("sk-ant-oat")
def strip_cloned_single_use_oauth_grants(profile_dir: Path) -> Dict[str, Any]:
"""Remove forked single-use OAuth grants from a freshly cloned profile.
Called after any code path that copies credential files from one profile into another (``hermes
profile create --clone-all``, the dashboard/TUI ``mirror_credentials`` flow). API-key pool rows
are kept — a static key is safe to duplicate.
Returns a summary ``{"pool": [...provider ids], "providers": [...], "files": [...]}`` of what
was stripped (empty lists when nothing was). Never raises: a clone must not fail because
credential hygiene could not run — the caller logs the summary.
"""
from hermes_cli.auth import _save_auth_store
stripped: Dict[str, Any] = {"pool": [], "providers": [], "files": []}
profile_dir = Path(profile_dir)
for name in SINGLE_USE_OAUTH_SINGLETON_FILES:
try:
target = profile_dir / name
if target.is_file() or target.is_symlink():
target.unlink()
stripped["files"].append(name)
except OSError:
logger.debug("Could not remove cloned %s from %s", name, profile_dir, exc_info=True)
auth_path = profile_dir / "auth.json"
if not auth_path.is_file():
return stripped
try:
store = json.loads(auth_path.read_text(encoding="utf-8-sig"))
except (OSError, json.JSONDecodeError):
return stripped
if not isinstance(store, dict):
return stripped
changed = False
pool = store.get("credential_pool")
if isinstance(pool, dict):
for provider_id in list(pool):
if provider_id not in SINGLE_USE_REFRESH_POOL_PROVIDERS:
continue
entries = pool.get(provider_id)
if not isinstance(entries, list):
continue
kept = [e for e in entries if not _is_oauth_pool_payload(e)]
if len(kept) != len(entries):
changed = True
stripped["pool"].append(provider_id)
if kept:
pool[provider_id] = kept
else:
# No local rows at all → read_credential_pool falls back
# to the root slice for this provider.
del pool[provider_id]
providers = store.get("providers")
if isinstance(providers, dict):
# Device-code grants for these providers live under providers.<id>;
# _load_provider_state has the same root fallback, so dropping the
# copy keeps the profile working while removing the fork.
for provider_id in ("openai-codex", "xai-oauth"):
block = providers.get(provider_id)
if isinstance(block, dict) and block:
del providers[provider_id]
stripped["providers"].append(provider_id)
changed = True
if not changed:
return stripped
try:
_save_auth_store(store, target_path=auth_path)
except Exception:
logger.debug(
"Failed to strip cloned single-use OAuth grants from %s",
auth_path,
exc_info=True,
)
return stripped
_OAUTH_TOKEN_FIELDS = (
"access_token",
"refresh_token",
"expires_at",
"expires_at_ms",
"last_refresh",
)
_oauth_heal_notices: List[str] = []
# provider -> (profile auth.json path, auth.json mtime_ns, singleton mtime_ns)
# of the last store verified fork-free; lets load_pool() skip the locked scan.
_oauth_heal_clean_marks: Dict[str, Tuple[str, Optional[int], Optional[int]]] = {}
def consume_oauth_heal_notices() -> List[str]:
"""Return (and clear) human-readable notes about heals run in this process.
``hermes auth list`` / ``hermes auth status`` print them so the user sees that a forked grant
was consolidated rather than only finding it in logs.
"""
from hermes_cli.auth import _oauth_heal_notices
notes = list(_oauth_heal_notices)
_oauth_heal_notices.clear()
return notes
def _oauth_identity(entry: Dict[str, Any]) -> Optional[str]:
"""Stable account identity for an OAuth row when the token carries one.
Codex / xAI access tokens are JWTs with ``sub`` / ``email`` / ``chatgpt_account_id`` claims;
Anthropic ``sk-ant-oat`` tokens carry none (returns None, so lineage rests on id / token
material).
"""
from hermes_cli.auth import _nonempty_str
if not isinstance(entry, dict):
return None
for token in (entry.get("access_token"), entry.get("id_token")):
claims = _decode_jwt_claims(token)
if not claims:
continue
nested = claims.get("https://api.openai.com/auth")
account = nested.get("chatgpt_account_id") if isinstance(nested, dict) else None
for value in (account, claims.get("sub"), claims.get("email")):
if _nonempty_str(value):
return value.strip()
return None
def _oauth_freshness(entry: Dict[str, Any]) -> float:
"""Best-effort 'how recently was this pair issued' score (epoch seconds).
A rotation always issues a later-expiring access token, so ``expires_at`` ordering identifies
the live copy; ``last_refresh`` and the JWT ``exp`` claim are fallbacks for rows that do not
persist expiry.
"""
from agent.credential_pool import _parse_absolute_timestamp
best = 0.0
for key in ("expires_at_ms", "expires_at", "last_refresh"):
ts = _parse_absolute_timestamp(entry.get(key))
if ts and ts > best:
best = ts
if best == 0.0:
exp = _decode_jwt_claims(entry.get("access_token")).get("exp")
ts = _parse_absolute_timestamp(exp)
if ts:
best = ts
return best
def _find_root_counterpart(
profile_row: Dict[str, Any], root_rows: List[Dict[str, Any]]
) -> Optional[int]:
"""Index of the root OAuth row that shares a grant lineage with *profile_row*.
Fallback per the one-grant-at-root rule: same provider + same OAuth client — every Anthropic
``hermes_pkce`` grant uses one client id and carries no claims, so two Anthropic OAuth rows with
no contrary identity are one lineage.
"""
from hermes_cli.auth import _nonempty_str
candidates = [i for i, r in enumerate(root_rows) if _is_oauth_pool_payload(r)]
if not candidates:
return None
pid = profile_row.get("id")
for i in candidates:
if pid and root_rows[i].get("id") == pid:
return i
p_ident = _oauth_identity(profile_row)
for i in candidates:
r_ident = _oauth_identity(root_rows[i])
if p_ident and r_ident and p_ident == r_ident:
return i
for key in ("refresh_token", "access_token"):
p_val = profile_row.get(key)
if not _nonempty_str(p_val):
continue
for i in candidates:
if root_rows[i].get(key) == p_val:
return i
# Fallback: same provider + same client. Only a contradicting identity
# (both sides carry claims and they differ from every root row) blocks it.
if p_ident:
for i in candidates:
if not _oauth_identity(root_rows[i]):
return i
return None
return candidates[0]
def _adopt_oauth_material(target: Dict[str, Any], winner: Dict[str, Any]) -> Dict[str, Any]:
"""Return *target* carrying *winner*'s token pair, status markers cleared."""
from hermes_cli.auth import _POOL_STATUS_FIELDS
merged = dict(target)
for key in _OAUTH_TOKEN_FIELDS:
if winner.get(key) is not None:
merged[key] = winner[key]
else:
merged.pop(key, None)
for status_field in _POOL_STATUS_FIELDS:
merged[status_field] = None
return merged
def _singleton_as_row(path: Path) -> Optional[Dict[str, Any]]:
"""Read a ``.anthropic_oauth.json`` as a pool-row-shaped dict, or None."""
try:
data = json.loads(path.read_text(encoding="utf-8"))
except (OSError, ValueError):
return None
if not isinstance(data, dict) or not str(data.get("accessToken") or "").strip():
return None
return {
"access_token": data.get("accessToken"),
"refresh_token": data.get("refreshToken"),
"expires_at_ms": data.get("expiresAt"),
}
def heal_forked_single_use_oauth_grants(provider_id: str) -> Optional[Dict[str, Any]]:
"""Consolidate a profile's forked copy of a single-use OAuth grant into root.
Runs only in profile mode for ``SINGLE_USE_REFRESH_POOL_PROVIDERS``. Returns a summary
``{"adopted": bool, "stripped_ids": [...], "files": [...], "providers_block": bool}`` when
something was healed, else ``None``. Never raises.
"""
if provider_id not in SINGLE_USE_REFRESH_POOL_PROVIDERS:
return None
try:
return _heal_forked_single_use_oauth_grants(provider_id)
except Exception:
logger.debug("%s: forked-OAuth heal skipped", provider_id, exc_info=True)
return None
def _heal_forked_provider_block(
profile_store: Dict[str, Any], root_store: Dict[str, Any], provider_id: str,
) -> Optional[bool]:
"""Consolidate a forked ``providers.<id>`` device-code block into root.
Returns None when nothing matched, False when the profile copy was dropped (root already
newest), True when the profile copy was fresher and was adopted into root.
"""
p_providers = profile_store.get("providers")
r_providers = root_store.get("providers")
if not (isinstance(p_providers, dict) and isinstance(r_providers, dict)):
return None
p_block = p_providers.get(provider_id)
r_block = r_providers.get(provider_id)
if not (isinstance(p_block, dict) and p_block and isinstance(r_block, dict) and r_block):
return None
def _flat(block: Dict[str, Any]) -> Dict[str, Any]:
tokens = block.get("tokens") if isinstance(block.get("tokens"), dict) else {}
return {**tokens, "last_refresh": block.get("last_refresh")}
p_flat, r_flat = _flat(p_block), _flat(r_block)
p_ident, r_ident = _oauth_identity(p_flat), _oauth_identity(r_flat)
if p_ident and r_ident and p_ident != r_ident:
return None
adopted = _oauth_freshness(p_flat) > _oauth_freshness(r_flat)
if adopted:
r_providers[provider_id] = dict(p_block)
del p_providers[provider_id]
return adopted
def _heal_forked_single_use_oauth_grants(provider_id: str) -> Optional[Dict[str, Any]]:
from hermes_cli.auth import _auth_file_path, _auth_store_lock, _global_auth_file_path, _load_auth_store, _oauth_heal_clean_marks, _oauth_heal_notices, _same_path, _save_auth_store
root_path = _global_auth_file_path()
if root_path is None:
return None # classic mode: nothing to consolidate into
if os.environ.get("PYTEST_CURRENT_TEST"):
# Same seat belt as the write-through paths: never touch the real
# user's ~/.hermes/auth.json from a test that forgot to isolate HOME.
real_home_env = os.environ.get("HOME", "")
if real_home_env and _same_path(root_path, Path(real_home_env) / ".hermes" / "auth.json"):
return None
profile_path = _auth_file_path()
profile_home = profile_path.parent
root_home = root_path.parent
profile_singleton = profile_home / ".anthropic_oauth.json" if provider_id == "anthropic" else None
# Hot-path short-circuit: load_pool() runs per model call. Once this
# profile's store was verified clean for *provider_id*, skip the locked
# read-modify-write until the profile's own files change (mtime key).
def _stamp(p: Optional[Path]) -> Optional[int]:
try:
return p.stat().st_mtime_ns if p is not None else None
except OSError:
return None
fingerprint = (str(profile_path), _stamp(profile_path), _stamp(profile_singleton))
if _oauth_heal_clean_marks.get(provider_id) == fingerprint:
return None
if fingerprint[1] is None and fingerprint[2] is None:
_oauth_heal_clean_marks[provider_id] = fingerprint
return None
summary: Dict[str, Any] = {"adopted": False, "stripped_ids": [], "files": [], "providers_block": False}
log_bits: List[str] = []
# Lock order: active (profile) store first, then the root source store —
# the same order ``_provider_state_transaction`` uses.
with _auth_store_lock():
profile_store = _load_auth_store(profile_path) if profile_path.exists() else {"providers": {}}
with _auth_store_lock(target_path=root_path):
root_store = _load_auth_store(root_path) if root_path.exists() else {"providers": {}}
profile_changed = False
root_changed = False
p_pool = profile_store.get("credential_pool")
p_rows = p_pool.get(provider_id) if isinstance(p_pool, dict) else None
p_rows = p_rows if isinstance(p_rows, list) else []
r_pool = root_store.get("credential_pool")
r_rows = r_pool.get(provider_id) if isinstance(r_pool, dict) else None
r_rows = r_rows if isinstance(r_rows, list) else []
r_oauth = [r for r in r_rows if _is_oauth_pool_payload(r)]
root_singleton = root_home / ".anthropic_oauth.json" if provider_id == "anthropic" else None
root_singleton_row = (
_singleton_as_row(root_singleton)
if root_singleton is not None and root_singleton.exists() else None
)
# ── credential_pool rows ────────────────────────────────────
kept_rows: List[Any] = []
for row in p_rows:
if not _is_oauth_pool_payload(row):
kept_rows.append(row) # API keys are safe to duplicate
continue
match_idx = _find_root_counterpart(row, r_rows)
if match_idx is not None:
root_row = r_rows[match_idx]
if _oauth_freshness(row) > _oauth_freshness(root_row):
r_rows[match_idx] = _adopt_oauth_material(root_row, row)
root_changed = True
summary["adopted"] = True
summary["stripped_ids"].append(row.get("id"))
profile_changed = True
continue
# No root pool counterpart. Root's grant may live only in its
# .anthropic_oauth.json (the ``hermes auth`` PKCE shape); a
# profile hermes_pkce-family row is that grant's copy.
is_pkce = str(row.get("source") or "").endswith("hermes_pkce")
if is_pkce and root_singleton_row is not None and not r_oauth:
if _oauth_freshness(row) > _oauth_freshness(root_singleton_row):
root_singleton_row = _adopt_oauth_material(root_singleton_row, row)
summary["adopted"] = True
summary["stripped_ids"].append(row.get("id"))
profile_changed = True
continue
# Root holds no copy of this lineage (independent account, or
# root never had the grant): the profile's row may be the
# only surviving copy — leave it alone.
kept_rows.append(row)
if profile_changed and isinstance(p_pool, dict):
if kept_rows:
p_pool[provider_id] = kept_rows
else:
p_pool.pop(provider_id, None)
# ── providers.<id> device-code blocks (Codex / xAI) ─────────
if provider_id in ("openai-codex", "xai-oauth"):
block_result = _heal_forked_provider_block(profile_store, root_store, provider_id)
if block_result is not None:
profile_changed = True
summary["providers_block"] = True
if block_result:
root_changed = True
summary["adopted"] = True
# ── profile-local .anthropic_oauth.json singleton ───────────
if profile_singleton is not None and profile_singleton.exists():
p_single = _singleton_as_row(profile_singleton)
root_has_grant = bool(r_oauth) or root_singleton_row is not None
if p_single is not None and root_has_grant:
if root_singleton_row is not None:
if _oauth_freshness(p_single) > _oauth_freshness(root_singleton_row):
root_singleton_row = _adopt_oauth_material(root_singleton_row, p_single)
summary["adopted"] = True
else:
# Root only has pool rows: fold the singleton's pair
# into the freshest-matching root pkce row, if any.
idx = next(
(i for i, r in enumerate(r_rows)
if _is_oauth_pool_payload(r)
and str(r.get("source") or "").endswith("hermes_pkce")),
None,
)
if idx is not None and _oauth_freshness(p_single) > _oauth_freshness(r_rows[idx]):
r_rows[idx] = _adopt_oauth_material(r_rows[idx], p_single)
root_changed = True
summary["adopted"] = True
try:
profile_singleton.unlink()
summary["files"].append(profile_singleton.name)
except OSError:
logger.debug("could not remove %s", profile_singleton, exc_info=True)
# Otherwise root has NO grant for this provider (or the file
# is not a grant): the profile's singleton may be the only
# surviving copy — never delete it.
if not (profile_changed or root_changed or summary["adopted"]):
_oauth_heal_clean_marks[provider_id] = fingerprint
return None
if summary["adopted"] and root_singleton is not None and root_singleton_row is not None:
# Keep root's singleton and its ``hermes_pkce``-seeded pool row
# in step: root's next load_pool() re-seeds that row FROM the
# singleton file, so a stale file would resurrect the spent
# pair (and a stale row would be overwritten by a fresh file).
pkce_idx = next(
(i for i, r in enumerate(r_rows)
if _is_oauth_pool_payload(r) and r.get("source") == "hermes_pkce"),
None,
)
if pkce_idx is not None:
pkce_row = r_rows[pkce_idx]
if _oauth_freshness(pkce_row) > _oauth_freshness(root_singleton_row):
root_singleton_row = _adopt_oauth_material(root_singleton_row, pkce_row)
elif _oauth_freshness(root_singleton_row) > _oauth_freshness(pkce_row):
r_rows[pkce_idx] = _adopt_oauth_material(pkce_row, root_singleton_row)
root_changed = True
if root_changed:
if isinstance(r_pool, dict):
r_pool[provider_id] = r_rows
else:
root_store["credential_pool"] = {provider_id: r_rows}
_save_auth_store(root_store, target_path=root_path)
if summary["adopted"] and root_singleton is not None and root_singleton_row is not None:
from agent.anthropic_credentials import _write_hermes_oauth_credentials
_write_hermes_oauth_credentials(
root_singleton_row.get("access_token") or "",
root_singleton_row.get("refresh_token"),
root_singleton_row.get("expires_at_ms"),
target=root_singleton,
)
if profile_changed and profile_path.exists():
_save_auth_store(profile_store, target_path=profile_path)
if summary["stripped_ids"]:
log_bits.append(f"pool rows {summary['stripped_ids']}")
if summary["providers_block"]:
log_bits.append(f"providers.{provider_id} block")
if summary["files"]:
log_bits.append(", ".join(summary["files"]))
verdict = (
"profile copy was the live pair; root updated"
if summary["adopted"] else "root copy already newest; profile copy dropped"
)
message = (
f"profile {profile_home.name}: consolidated forked {provider_id} OAuth grant "
f"({'; '.join(log_bits) or 'no-op'}) into the root grant — {verdict}; "
f"this profile now borrows the root grant (#100339)"
)
logger.info(message)
_oauth_heal_notices.append(message)
return summary

222
hermes_cli/auth_zai_kimi.py Normal file
View File

@@ -0,0 +1,222 @@
"""Kimi Code and Z.AI endpoint auto-detection, LM Studio base-URL normalization.
Split out of ``hermes_cli/auth.py``; every moved name is re-imported there, so
``hermes_cli.auth.<name>`` keeps resolving (and monkeypatching) as before. Origin-internal
helpers are imported lazily inside each function (no import cycle; patches on
``hermes_cli.auth.<helper>`` still intercept).
"""
from __future__ import annotations
import logging
import hashlib
from typing import Dict, Optional
from hermes_cli.auth_constants import httpx
# Log-record parity with the origin module (caplog tests pin "hermes_cli.auth").
logger = logging.getLogger("hermes_cli.auth")
# Kimi Code (kimi.com/code) issues keys prefixed "sk-kimi-" that only work
# on api.kimi.com/coding. Legacy keys from platform.moonshot.ai work on
# api.moonshot.ai/v1 (the old default). Auto-detect when user hasn't set
# KIMI_BASE_URL explicitly.
#
# Note: the base URL intentionally has NO /v1 suffix. The /coding endpoint
# speaks the Anthropic Messages protocol, and the anthropic SDK appends
# "/v1/messages" internally — so "/coding" + SDK suffix → "/coding/v1/messages"
# (the correct target). Using "/coding/v1" here would produce
# "/coding/v1/v1/messages" (a 404).
KIMI_CODE_BASE_URL = "https://api.kimi.com/coding"
def _resolve_kimi_base_url(api_key: str, default_url: str, env_override: str) -> str:
"""Return the correct Kimi base URL based on the API key prefix.
If the user has explicitly set KIMI_BASE_URL, that always wins. Otherwise, sk-kimi- prefixed
keys route to api.kimi.com/coding/v1.
"""
if env_override:
return env_override
# No key → nothing to infer from. Return default without inspecting.
if not api_key:
return default_url
if api_key.startswith("sk-kimi-"):
return KIMI_CODE_BASE_URL
return default_url
ZAI_ENDPOINTS = [
# (id, base_url, probe_models, label)
("global", "https://api.z.ai/api/paas/v4", ["glm-5"], "Global"),
("cn", "https://open.bigmodel.cn/api/paas/v4", ["glm-5"], "China"),
("coding-global", "https://api.z.ai/api/coding/paas/v4", ["glm-5.3", "glm-5.3-flash", "glm-5.2", "glm-5.1", "glm-5v-turbo", "glm-4.7"], "Global (Coding Plan)"),
("coding-cn", "https://open.bigmodel.cn/api/coding/paas/v4", ["glm-5.3", "glm-5.3-flash", "glm-5.2", "glm-5.1", "glm-5v-turbo", "glm-4.7"], "China (Coding Plan)"),
]
def _probe_single_zai_endpoint(
api_key: str, endpoint: tuple, timeout: float,
) -> Optional[Dict[str, str]]:
"""Probe a single Z.AI endpoint. Returns endpoint info dict or None.
Preserves the per-endpoint candidate-model loop: endpoints carry a ``probe_models`` LIST and
each model is tried in order until one succeeds (some plans only accept newer/older GLM slugs).
"""
ep_id, base_url, probe_models, label = endpoint
for model in probe_models:
try:
resp = httpx.post(
f"{base_url}/chat/completions",
headers={
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
},
json={
"model": model,
"stream": False,
"max_tokens": 1,
"messages": [{"role": "user", "content": "ping"}],
},
timeout=timeout,
)
if resp.status_code == 200:
logger.debug("Z.AI endpoint probe: %s (%s) model=%s OK", ep_id, base_url, model)
return {
"id": ep_id,
"base_url": base_url,
"model": model,
"label": label,
}
logger.debug("Z.AI endpoint probe: %s model=%s returned %s", ep_id, model, resp.status_code)
except Exception as exc:
logger.debug("Z.AI endpoint probe: %s model=%s failed: %s", ep_id, model, exc)
return None
def detect_zai_endpoint(api_key: str, timeout: float = 8.0) -> Optional[Dict[str, str]]:
"""Probe z.ai endpoints in parallel to find one that accepts this API key.
Returns {"id": ..., "base_url": ..., "model": ..., "label": ...} for the first working endpoint
(in ZAI_ENDPOINTS priority order), or None if all fail. For endpoints with multiple candidate
models, each worker tries its endpoint's models in order and returns the first that succeeds.
"""
from concurrent.futures import ThreadPoolExecutor, as_completed
# No `with` block: a context manager would join ALL probe threads on
# exit, defeating the early return below. shutdown(wait=False) lets the
# surviving daemon-style probes drain in the background instead of
# blocking the caller on slow/unreachable endpoints.
pool = ThreadPoolExecutor(max_workers=len(ZAI_ENDPOINTS))
try:
futures = {
pool.submit(_probe_single_zai_endpoint, api_key, ep, timeout): ep[0]
for ep in ZAI_ENDPOINTS
}
by_id = {ep_id: f for f, ep_id in futures.items()}
results: Dict[str, Dict[str, str]] = {}
for future in as_completed(futures):
ep_id = futures[future]
try:
result = future.result()
if result is not None:
results[ep_id] = result
except Exception:
pass
# Early exit in PRIORITY order: walk endpoints highest-priority
# first; if one has succeeded and every higher-priority probe
# has already finished (without success), no later completion
# can win — return now instead of waiting out slow endpoints
# (main's sequential loop also stopped at first success).
for ep in ZAI_ENDPOINTS:
if not by_id[ep[0]].done():
break # a higher-priority probe is still in flight
if ep[0] in results:
return results[ep[0]]
# All probes finished: first match in priority order, if any.
for ep in ZAI_ENDPOINTS:
if ep[0] in results:
return results[ep[0]]
return None
finally:
pool.shutdown(wait=False)
def _resolve_zai_base_url(api_key: str, default_url: str, env_override: str) -> str:
"""Return the correct Z.AI base URL by probing endpoints.
If the user has explicitly set GLM_BASE_URL, that always wins. Otherwise, probe the candidate
endpoints to find one that accepts the key. The detected endpoint is cached in provider state
(auth.json) keyed on a hash of the API key so subsequent starts skip the probe.
"""
from hermes_cli.auth import _auth_store_lock, _load_auth_store, _load_provider_state, _save_auth_store, _store_provider_state, detect_zai_endpoint
if env_override:
return env_override
# No API key set → don't probe (would fire N×M HTTPS requests with an
# empty Bearer token, all returning 401). This path is hit during
# auxiliary-client auto-detection when the user has no Z.AI credentials
# at all — the caller discards the result immediately, so the probe is
# pure latency for every AIAgent construction.
if not api_key:
return default_url
# Check provider-state cache for a previously-detected endpoint.
auth_store = _load_auth_store()
state = _load_provider_state(auth_store, "zai") or {}
cached = state.get("detected_endpoint")
if isinstance(cached, dict) and cached.get("base_url"):
key_hash = cached.get("key_hash", "")
if key_hash == hashlib.sha256(api_key.encode()).hexdigest()[:16]:
logger.debug("Z.AI: using cached endpoint %s", cached["base_url"])
return cached["base_url"]
# Probe — may take up to ~8s per endpoint.
detected = detect_zai_endpoint(api_key)
if detected and detected.get("base_url"):
# Persist the detection result keyed on the API key hash.
key_hash = hashlib.sha256(api_key.encode()).hexdigest()[:16]
detected_endpoint = {
"base_url": detected["base_url"],
"endpoint_id": detected.get("id", ""),
"model": detected.get("model", ""),
"label": detected.get("label", ""),
"key_hash": key_hash,
}
# Persist failure (disk full, permissions, lock timeout) must not
# break resolution — detection already succeeded; worst case the
# next start re-probes.
try:
with _auth_store_lock():
# Reload auth_store under lock to avoid overwriting concurrent changes
auth_store = _load_auth_store()
state_under_lock = _load_provider_state(auth_store, "zai") or {}
state_under_lock["detected_endpoint"] = detected_endpoint
# set_active=False: this runs from credential-pool env seeding
# (agent/credential_pool.py) for ANY user with a Z.AI key in env,
# and caching a probe result must not flip their active provider.
_store_provider_state(auth_store, "zai", state_under_lock, set_active=False)
_save_auth_store(auth_store)
except Exception as exc:
logger.warning("Z.AI: could not persist detected endpoint (%s); will re-probe next start", exc)
logger.info("Z.AI: auto-detected endpoint %s (%s)", detected["label"], detected["base_url"])
return detected["base_url"]
logger.debug("Z.AI: probe failed, falling back to default %s", default_url)
return default_url
def _normalize_lmstudio_runtime_base_url(base_url: str) -> str:
"""Return the OpenAI-compatible LM Studio runtime base URL.
LM Studio's native management API lives under ``/api/v1`` while its OpenAI-compatible chat
endpoint lives under ``/v1``. Users often paste either form into ``LM_BASE_URL`` or
``model.base_url``; normalize before the OpenAI SDK appends ``/chat/completions``.
"""
root = str(base_url or "").strip().rstrip("/")
for suffix in ("/api/v1", "/api", "/v1"):
if root.endswith(suffix):
root = root[: -len(suffix)].rstrip("/")
break
return (root or "http://127.0.0.1:1234") + "/v1"