refactor(hermes_cli): auth_oauth_grants/device_flow — fold guard pairs, single browser-open try path, list-comp log bits

This commit is contained in:
Teknium
2026-09-02 23:18:23 -07:00
parent 45bce4faef
commit 0a142eff58
2 changed files with 31 additions and 47 deletions

View File

@@ -101,8 +101,7 @@ def _print_loopback_ssh_hint(redirect_uri: str, *, docs_url: str | None = None)
parsed = urlparse(redirect_uri)
except Exception:
return
host = parsed.hostname or ""
port = parsed.port
host, port = parsed.hostname or "", parsed.port
if host not in {"127.0.0.1", "::1", "localhost"} or not port:
return
divider = "-" * 60
@@ -168,8 +167,8 @@ def _request_device_code(
response.raise_for_status()
data = response.json()
required_fields = [
"device_code", "user_code", "verification_uri",
"verification_uri_complete", "expires_in", "interval"]
"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)}")
@@ -202,13 +201,12 @@ def _print_device_code_instructions(
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:
try:
opened = webbrowser.open(verification_url)
except Exception:
if not swallow_open_errors:
raise
opened = False
if opened:
print(" (Opened browser for verification)")
else:
@@ -273,10 +271,9 @@ def _poll_for_token(
"device_code": device_code}),
expires_in=expires_in,
poll_interval=max(1, min(poll_interval, DEVICE_AUTH_POLL_INTERVAL_CAP_SECONDS)),
validate_success=_validate,
validate_success=_validate, on_error=_error,
on_non_json_error=lambda _r: RuntimeError(
"Token endpoint returned a non-JSON error response"),
on_error=_error,
# Enriched at the SOURCE so the CLI login and the dashboard/desktop poller
# (web_server._nous_poller surfaces str(e) to the UI) both inherit the guidance.
on_timeout=lambda: TimeoutError(_nous_device_auth_timeout_message(portal_base_url)))

View File

@@ -60,8 +60,8 @@ def strip_cloned_single_use_oauth_grants(profile_dir: Path) -> Dict[str, Any]:
stripped: Dict[str, Any] = {"pool": [], "providers": [], "files": []}
profile_dir = Path(profile_dir)
for name in SINGLE_USE_OAUTH_SINGLETON_FILES:
target = profile_dir / name
try:
target = profile_dir / name
if target.is_file() or target.is_symlink():
target.unlink()
stripped["files"].append(name)
@@ -73,17 +73,16 @@ def strip_cloned_single_use_oauth_grants(profile_dir: Path) -> Dict[str, Any]:
try:
store = json.loads(auth_path.read_text(encoding="utf-8-sig"))
except (OSError, json.JSONDecodeError):
return stripped
store = None
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):
if (provider_id not in SINGLE_USE_REFRESH_POOL_PROVIDERS
or not isinstance(entries, list)):
continue
kept = [e for e in entries if not _is_oauth_pool_payload(e)]
if len(kept) != len(entries):
@@ -163,16 +162,11 @@ def _oauth_freshness(entry: Dict[str, Any]) -> float:
the live copy; ``last_refresh`` and the JWT ``exp`` claim are fallbacks.
"""
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
stamps = [entry.get(k) for k in ("expires_at_ms", "expires_at", "last_refresh")]
best = max((ts for ts in map(_parse_absolute_timestamp, stamps) if ts), default=0.0)
if best == 0.0:
exp = _decode_jwt_claims(entry.get("access_token")).get("exp")
ts = _parse_absolute_timestamp(exp)
if ts:
best = ts
best = _parse_absolute_timestamp(exp) or 0.0
return best
@@ -223,8 +217,7 @@ def _adopt_oauth_material(target: Dict[str, Any], winner: Dict[str, Any]) -> Dic
merged[key] = winner[key]
else:
merged.pop(key, None)
for status_field in _POOL_STATUS_FIELDS:
merged[status_field] = None
merged.update(dict.fromkeys(_POOL_STATUS_FIELDS))
return merged
@@ -272,12 +265,10 @@ def _heal_forked_provider_block(
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")
p_providers, r_providers = profile_store.get("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)
p_block, r_block = p_providers.get(provider_id), r_providers.get(provider_id)
if not (isinstance(p_block, dict) and p_block and isinstance(r_block, dict) and r_block):
return None
@@ -313,9 +304,8 @@ def _pool_rows(store: Dict[str, Any], provider_id: str) -> Tuple[Any, List[Any]]
def _adopt_if_fresher(
target: Dict[str, Any], candidate: Dict[str, Any]) -> Optional[Dict[str, Any]]:
"""*target* carrying *candidate*'s pair when the candidate rotated later, else None."""
if _oauth_freshness(candidate) > _oauth_freshness(target):
return _adopt_oauth_material(target, candidate)
return None
fresher = _oauth_freshness(candidate) > _oauth_freshness(target)
return _adopt_oauth_material(target, candidate) if fresher else None
class _HealPass:
@@ -399,8 +389,7 @@ class _HealPass:
# root pkce row, if any.
idx = next(
(i for i, r in enumerate(self.r_rows)
if _is_oauth_pool_payload(r) and _is_pkce_row(r)),
None)
if _is_oauth_pool_payload(r) and _is_pkce_row(r)), None)
if idx is not None:
self._adopt_root_row(idx, p_single)
try:
@@ -421,14 +410,15 @@ class _HealPass:
return
pkce_idx = next(
(i for i, r in enumerate(self.r_rows)
if _is_oauth_pool_payload(r) and r.get("source") == "hermes_pkce"),
None)
if _is_oauth_pool_payload(r) and r.get("source") == "hermes_pkce"), None)
if pkce_idx is None:
return
pkce_row = self.r_rows[pkce_idx]
if _oauth_freshness(pkce_row) > _oauth_freshness(self.root_singleton_row):
row_fresh = _oauth_freshness(pkce_row)
single_fresh = _oauth_freshness(self.root_singleton_row)
if row_fresh > single_fresh:
self.root_singleton_row = _adopt_oauth_material(self.root_singleton_row, pkce_row)
elif _oauth_freshness(self.root_singleton_row) > _oauth_freshness(pkce_row):
elif single_fresh > row_fresh:
self.r_rows[pkce_idx] = _adopt_oauth_material(pkce_row, self.root_singleton_row)
self.root_changed = True
@@ -438,13 +428,10 @@ class _HealPass:
def notice(self, profile_name: str) -> str:
summary = self.summary
log_bits: List[str] = []
if summary["stripped_ids"]:
log_bits.append(f"pool rows {summary['stripped_ids']}")
if summary["providers_block"]:
log_bits.append(f"providers.{self.provider_id} block")
if summary["files"]:
log_bits.append(", ".join(summary["files"]))
log_bits = [bit for bit, present in (
(f"pool rows {summary['stripped_ids']}", summary["stripped_ids"]),
(f"providers.{self.provider_id} block", summary["providers_block"]),
(", ".join(summary["files"]), summary["files"])) if present]
verdict = (
"profile copy was the live pair; root updated"
if summary["adopted"] else "root copy already newest; profile copy dropped")