refactor(auth): share the loopback PKCE listener between OAuth flows

Move the S256 verifier/challenge pair, the loopback callback handler, the bind-first
listener and the serve-until-redirect loop out of auth_spotify into auth_device_flow so a
second loopback PKCE provider does not copy 80 lines of HTTP-server plumbing. Spotify's
behaviour and error codes are unchanged; only its private copies are deleted.
This commit is contained in:
Teknium
2026-09-06 04:13:45 -07:00
parent 41380ccef9
commit 3310298a37
2 changed files with 100 additions and 74 deletions
+89 -2
View File
@@ -1,4 +1,4 @@
"""Shared device-code / browser / TLS helpers for interactive OAuth logins.
"""Shared device-code / loopback-PKCE / browser / TLS helpers for interactive OAuth logins.
Split out of ``hermes_cli/auth.py`` and re-exported there; origin helpers are imported lazily
inside each function so ``hermes_cli.auth.<name>`` patches still intercept (and no import cycle).
@@ -6,15 +6,19 @@ inside each function so ``hermes_cli.auth.<name>`` patches still intercept (and
from __future__ import annotations
import base64
import hashlib
import logging
import os
import ssl
import sys
import threading
import time
import webbrowser
from http.server import BaseHTTPRequestHandler, HTTPServer
from pathlib import Path
from typing import Any, Callable, Dict, FrozenSet, Optional
from urllib.parse import urlparse
from urllib.parse import parse_qs, 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)
@@ -95,6 +99,89 @@ def _ssh_user_at_host() -> str:
return f"{user}@{hostname}"
def _pkce_code_verifier(length: int = 64) -> str:
return base64.urlsafe_b64encode(os.urandom(length)).decode("ascii").rstrip("=")[:128]
def _pkce_code_challenge(code_verifier: str) -> str:
digest = hashlib.sha256(code_verifier.encode("utf-8")).digest()
return base64.urlsafe_b64encode(digest).decode("ascii").rstrip("=")
def _make_loopback_callback_handler(
expected_path: str, *, display_name: str,
) -> tuple[type[BaseHTTPRequestHandler], dict[str, Any]]:
"""Handler class for an RFC 8252 loopback redirect plus the dict it fills in.
Only a GET on *expected_path* is accepted (anything else is a 404 and leaves the result
untouched), so a nonce embedded in the path acts as the CSRF ``state`` for authorization
servers that do not echo an explicit ``state`` parameter.
"""
result: dict[str, Any] = {"code": None, "state": None, "error": None, "error_description": None}
class _LoopbackCallbackHandler(BaseHTTPRequestHandler):
def do_GET(self) -> None: # noqa: N802
parsed = urlparse(self.path)
if parsed.path != expected_path:
self.send_response(404)
self.end_headers()
self.wfile.write(b"Not found.")
return
params = parse_qs(parsed.query)
for key in result:
result[key] = params.get(key, [None])[0]
self.send_response(200)
self.send_header("Content-Type", "text/html; charset=utf-8")
self.end_headers()
outcome = "failed" if result["error"] else "received"
self.wfile.write(
f"<html><body><h1>{display_name} authorization {outcome}.</h1>"
"You can close this tab.</body></html>".encode("utf-8"))
def log_message(self, format: str, *args: Any) -> None: # noqa: A003
return
return _LoopbackCallbackHandler, result
def _bind_loopback_callback_server(
host: str, port: int, handler_cls: type[BaseHTTPRequestHandler], *, err: Callable[..., AuthError],
bind_failed_code: str,
) -> HTTPServer:
"""Bind the loopback listener up front (``port=0`` = OS-assigned) so the redirect URI sent to
the authorization server names a port we already own — no probe-close-rebind race."""
class _ReuseHTTPServer(HTTPServer):
allow_reuse_address = True
try:
return _ReuseHTTPServer((host, port), handler_cls)
except OSError as exc:
raise err(f"Could not bind callback server on {host}:{port}: {exc}", bind_failed_code) from exc
def _serve_loopback_callback(
server: HTTPServer, result: dict[str, Any], *, timeout_seconds: float, err: Callable[..., AuthError],
timeout_code: str,
) -> dict[str, Any]:
"""Serve *server* until the redirect lands in *result* or the deadline passes; always closes."""
thread = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.1}, daemon=True)
thread.start()
deadline = time.monotonic() + max(5.0, timeout_seconds)
try:
while time.monotonic() < deadline:
if result["code"] or result["error"]:
return result
time.sleep(0.1)
finally:
server.shutdown()
server.server_close()
thread.join(timeout=1.0)
raise err("Authorization timed out waiting for the local callback.", timeout_code)
def _print_loopback_ssh_hint(redirect_uri: str, *, docs_url: str | None = None) -> None:
"""Print an SSH tunnel hint when a loopback-redirect OAuth flow runs on a remote host.
+11 -72
View File
@@ -7,27 +7,22 @@ lazily per function so ``hermes_cli.auth.<helper>`` patches still intercept and
from __future__ import annotations
import logging
import base64
import hashlib
import os
import threading
import time
import uuid
import webbrowser
from datetime import datetime, timezone
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import Any, Dict, Optional, Tuple
from urllib.parse import parse_qs, urlencode, urlparse
from urllib.parse import urlencode, urlparse
from hermes_cli.auth_constants import (
AuthError, DEFAULT_SPOTIFY_ACCOUNTS_BASE_URL, DEFAULT_SPOTIFY_API_BASE_URL, DEFAULT_SPOTIFY_REDIRECT_URI,
DEFAULT_SPOTIFY_SCOPE, SPOTIFY_ACCESS_TOKEN_REFRESH_SKEW_SECONDS, SPOTIFY_DASHBOARD_URL, SPOTIFY_DOCS_URL,
_spotify_err, httpx,
)
from hermes_cli.auth_device_flow import (
_bind_loopback_callback_server, _make_loopback_callback_handler, _pkce_code_challenge,
_pkce_code_verifier, _serve_loopback_callback)
logger = logging.getLogger("hermes_cli.auth")
_CALLBACK_HTML = "<html><body><h1>Spotify authorization {}.</h1>You can close this tab.</body></html>"
def _clean(value: Any) -> str:
return str(value or "").strip()
@@ -89,15 +84,6 @@ def _spotify_accounts_base_url(state: Optional[Dict[str, Any]] = None) -> str:
)
def _spotify_code_verifier(length: int = 64) -> str:
return base64.urlsafe_b64encode(os.urandom(length)).decode("ascii").rstrip("=")[:128]
def _spotify_code_challenge(code_verifier: str) -> str:
digest = hashlib.sha256(code_verifier.encode("utf-8")).digest()
return base64.urlsafe_b64encode(digest).decode("ascii").rstrip("=")
def _spotify_build_authorize_url(
*, client_id: str, redirect_uri: str, scope: str, state: str, code_challenge: str,
accounts_base_url: str,
@@ -124,60 +110,13 @@ def _spotify_validate_redirect_uri(redirect_uri: str) -> tuple[str, int, str]:
return host, parsed.port, parsed.path or "/"
def _make_spotify_callback_handler(expected_path: str) -> tuple[type[BaseHTTPRequestHandler], dict[str, Any]]:
result: dict[str, Any] = {"code": None, "state": None, "error": None, "error_description": None}
class _SpotifyCallbackHandler(BaseHTTPRequestHandler):
def do_GET(self) -> None: # noqa: N802
parsed = urlparse(self.path)
if parsed.path != expected_path:
self.send_response(404)
self.end_headers()
self.wfile.write(b"Not found.")
return
params = parse_qs(parsed.query)
for key in result:
result[key] = params.get(key, [None])[0]
self.send_response(200)
self.send_header("Content-Type", "text/html; charset=utf-8")
self.end_headers()
self.wfile.write(_CALLBACK_HTML.format("failed" if result["error"] else "received").encode("utf-8"))
def log_message(self, format: str, *args: Any) -> None: # noqa: A003
return
return _SpotifyCallbackHandler, result
def _spotify_wait_for_callback(redirect_uri: str, *, timeout_seconds: float = 180.0) -> dict[str, Any]:
host, port, path = _spotify_validate_redirect_uri(redirect_uri)
handler_cls, result = _make_spotify_callback_handler(path)
class _ReuseHTTPServer(HTTPServer):
allow_reuse_address = True
try:
server = _ReuseHTTPServer((host, port), handler_cls)
except OSError as exc:
raise _spotify_err(
f"Could not bind Spotify callback server on {host}:{port}: {exc}", "spotify_callback_bind_failed",
) from exc
thread = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.1}, daemon=True)
thread.start()
deadline = time.monotonic() + max(5.0, timeout_seconds)
try:
while time.monotonic() < deadline:
if result["code"] or result["error"]:
return result
time.sleep(0.1)
finally:
server.shutdown()
server.server_close()
thread.join(timeout=1.0)
raise _spotify_err("Spotify authorization timed out waiting for the local callback.", "spotify_callback_timeout")
handler_cls, result = _make_loopback_callback_handler(path, display_name="Spotify")
server = _bind_loopback_callback_server(
host, port, handler_cls, err=_spotify_err, bind_failed_code="spotify_callback_bind_failed")
return _serve_loopback_callback(
server, result, timeout_seconds=timeout_seconds, err=_spotify_err, timeout_code="spotify_callback_timeout")
def _spotify_token_payload_to_state(
@@ -392,11 +331,11 @@ def login_spotify_command(args) -> None:
api_base_url = _spotify_api_base_url(existing_state)
open_browser = not getattr(args, "no_browser", False)
code_verifier = _spotify_code_verifier()
code_verifier = _pkce_code_verifier()
state_nonce = uuid.uuid4().hex
authorize_url = _spotify_build_authorize_url(
client_id=client_id, redirect_uri=redirect_uri, scope=scope, state=state_nonce,
code_challenge=_spotify_code_challenge(code_verifier), accounts_base_url=accounts_base_url,
code_challenge=_pkce_code_challenge(code_verifier), accounts_base_url=accounts_base_url,
)
print(