refactor(auth): split OAuth-grant hygiene, device-flow helpers, model picker and Kimi/Z.AI detection into modules
This commit is contained in:
1486
hermes_cli/auth.py
1486
hermes_cli/auth.py
File diff suppressed because it is too large
Load Diff
415
hermes_cli/auth_device_flow.py
Normal file
415
hermes_cli/auth_device_flow.py
Normal 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
|
||||
387
hermes_cli/auth_model_picker.py
Normal file
387
hermes_cli/auth_model_picker.py
Normal 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)
|
||||
513
hermes_cli/auth_oauth_grants.py
Normal file
513
hermes_cli/auth_oauth_grants.py
Normal 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
222
hermes_cli/auth_zai_kimi.py
Normal 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"
|
||||
Reference in New Issue
Block a user