79d0d3b600
- project_tree: split build_tree (207 LOC) into _auto_buckets/_home_project phase helpers; _field() replaces 9 '(x.get(k) or "").strip()' ladders; _project_node takes wire-shaped **flags; drop unused placement 'repo_path' and _strip_trailing_sep; _FolderIndex/ _project_for_session collapsed. Old-vs-new golden (67 synthetic build_tree runs) identical. - methods_projects: _project_ok/_path_param unify 4 handler tails; _project_tree_row via dict comprehensions; discovered-cache read via contextlib.nullcontext; policy loader helpers. Golden over rows/policy/junk predicates identical. - agent_callbacks: subagent mirror dispatch on a delta table; drop _render_personality_prompt (single caller); getattr shorthand in _background_agent_kwargs. - entry: _write_or_exit unifies 3 write-fail exits; suppress(); heartbeat/sweep start loop. - mcp_oauth_sessions: suppress(), gc/listener/redirect compaction, docstrings. WIRE-PARITY-OK; tests green.
289 lines
13 KiB
Python
289 lines
13 KiB
Python
"""Session-backed MCP OAuth flows for the gateway (mcp.servers.oauth.*).
|
|
|
|
Mirrors the dashboard's *provider* OAuth model: ``start`` kicks off a background
|
|
worker and returns ``{session_id, auth_url, flow}``; ``poll`` reports
|
|
``{status: pending|approved|error}`` until tokens land on disk. No OAuth logic is
|
|
reimplemented — the token machinery is ``hermes mcp login``'s
|
|
(``_probe_single_server`` under ``force_interactive_oauth``) and ``DashboardOAuthFlow``
|
|
is the thread-safe bridge; the only new piece is a loopback HTTP listener feeding
|
|
``deliver_callback``. Remote-backend variant: the client binds its OWN loopback
|
|
listener, passes ``client_redirect_uri`` to ``start`` and relays the redirect via
|
|
``deliver_callback_flow``; state verification stays server-side either way.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import http.server
|
|
import secrets
|
|
import threading
|
|
import time
|
|
from contextlib import suppress
|
|
from pathlib import Path
|
|
from typing import Any, Dict, Optional
|
|
from urllib.parse import parse_qs, urlparse
|
|
|
|
# session_id -> record wrapping the shared DashboardOAuthFlow bridge plus bookkeeping.
|
|
_sessions: Dict[str, Dict[str, Any]] = {}
|
|
_sessions_lock = threading.Lock()
|
|
|
|
# How long a completed/abandoned session lingers before GC (seconds).
|
|
_SESSION_TTL_SECONDS = 900
|
|
# Cap concurrent in-flight flows so a runaway client can't exhaust ports/threads.
|
|
_MAX_PENDING = 12
|
|
|
|
|
|
def _gc_sessions() -> None:
|
|
"""Drop expired sessions. Called opportunistically on start."""
|
|
cutoff = time.time() - _SESSION_TTL_SECONDS
|
|
with _sessions_lock:
|
|
for sid in [sid for sid, rec in _sessions.items() if rec["created_at"] < cutoff]:
|
|
_shutdown_listener(_sessions.pop(sid))
|
|
|
|
|
|
def _shutdown_listener(rec: Dict[str, Any]) -> None:
|
|
server = rec.get("httpd")
|
|
if server is None:
|
|
return
|
|
for stop in (server.shutdown, server.server_close):
|
|
with suppress(Exception):
|
|
stop()
|
|
rec["httpd"] = None
|
|
|
|
|
|
def _validate_client_redirect_uri(uri: str) -> str:
|
|
"""Accept only plain-http loopback URLs (RFC 8252 native-app rules) so the
|
|
gateway can't pin an attacker-controlled redirect into a DCR registration."""
|
|
parsed = urlparse(str(uri or "").strip())
|
|
host = (parsed.hostname or "").lower()
|
|
if (parsed.scheme != "http" or host not in ("127.0.0.1", "localhost", "::1") or not parsed.port
|
|
or parsed.username is not None or parsed.password is not None):
|
|
raise ValueError(
|
|
"client_redirect_uri must be a loopback http URL like "
|
|
"http://127.0.0.1:<port>/callback"
|
|
)
|
|
return f"http://{'[' + host + ']' if ':' in host else host}:{parsed.port}{parsed.path or '/callback'}"
|
|
|
|
|
|
def _start_loopback_listener(flow) -> "http.server.HTTPServer":
|
|
"""Bind a loopback callback listener feeding ``flow.deliver_callback``; returns the
|
|
HTTPServer already serving on a daemon thread. The caller pins ``flow.redirect_uri``
|
|
from ``server_address`` BEFORE the worker starts (fixed at authorization)."""
|
|
|
|
class _Handler(http.server.BaseHTTPRequestHandler):
|
|
def do_GET(self): # noqa: N802 — stdlib naming
|
|
parsed = urlparse(self.path)
|
|
if parsed.path.rstrip("/") not in ("/callback", ""):
|
|
self.send_response(404)
|
|
self.end_headers()
|
|
return
|
|
qs = parse_qs(parsed.query)
|
|
code, state, error = ((qs.get(k) or [None])[0] for k in ("code", "state", "error"))
|
|
body = b"<h1>Authorization received</h1><p>You can close this tab and return to Hermes.</p>"
|
|
status = 200
|
|
try:
|
|
flow.deliver_callback(code=code, state=state, error=error)
|
|
except Exception:
|
|
body = b"<h1>OAuth callback rejected</h1><p>The callback was invalid or already used.</p>"
|
|
status = 400
|
|
self.send_response(status)
|
|
self.send_header("Content-Type", "text/html; charset=utf-8")
|
|
self.end_headers()
|
|
with suppress(Exception):
|
|
self.wfile.write(body)
|
|
|
|
def log_message(self, *_a): # silence stdlib request logging
|
|
return
|
|
|
|
httpd = http.server.HTTPServer(("127.0.0.1", 0), _Handler)
|
|
threading.Thread(
|
|
target=httpd.serve_forever, kwargs={"poll_interval": 0.5}, daemon=True,
|
|
name=f"mcp-oauth-cb-{flow.server_name}").start()
|
|
return httpd
|
|
|
|
|
|
def _worker(session_id: str, hermes_home: str, server_name: str, cfg: dict, reconnect_live: bool) -> None:
|
|
"""Drive the interactive MCP OAuth probe under the shared dashboard bridge (same
|
|
wrapping as ``web_server._run_dashboard_mcp_oauth``). On success the token file
|
|
exists and the server config is (re)saved; on failure the prior token/manager
|
|
state is restored."""
|
|
from hermes_cli.mcp_config import _oauth_tokens_present, _probe_single_server, _save_mcp_server
|
|
from hermes_constants import reset_hermes_home_override, set_hermes_home_override
|
|
|
|
rec = _sessions.get(session_id)
|
|
flow = rec["flow"] if rec else None
|
|
try:
|
|
from agent.secret_scope import (
|
|
build_profile_secret_scope, reset_secret_scope, set_secret_scope)
|
|
from tools.mcp_dashboard_oauth import dashboard_oauth_flow
|
|
from tools.mcp_oauth import force_interactive_oauth
|
|
from tools.mcp_oauth_manager import get_manager
|
|
|
|
home_token = set_hermes_home_override(hermes_home)
|
|
secret_token = set_secret_scope(build_profile_secret_scope(Path(hermes_home)))
|
|
try:
|
|
with force_interactive_oauth(), dashboard_oauth_flow(flow):
|
|
from tools.mcp_oauth import HermesTokenStorage
|
|
|
|
manager = get_manager()
|
|
storage = HermesTokenStorage(server_name)
|
|
backup = storage.snapshot()
|
|
previous_entry = None
|
|
try:
|
|
previous_entry = manager.remove(server_name, hermes_home=hermes_home)
|
|
timeout = max(float(cfg.get("connect_timeout", 0) or 0), 315)
|
|
tools = _probe_single_server(server_name, cfg, connect_timeout=timeout)
|
|
if not _oauth_tokens_present(server_name):
|
|
raise RuntimeError(
|
|
"The server responded, but no OAuth token was obtained — "
|
|
"this provider may require a manually-registered OAuth client.")
|
|
_save_mcp_server(server_name, cfg)
|
|
if flow is not None:
|
|
flow.tools = [{"name": t, "description": d} for t, d in tools]
|
|
flow.mark_approved()
|
|
if reconnect_live:
|
|
from tools.mcp_tool import reconnect_mcp_server
|
|
|
|
reconnect_mcp_server(server_name)
|
|
except Exception:
|
|
storage.restore(backup, only_if_absent=True)
|
|
manager.restore_entry(server_name, previous_entry, hermes_home=hermes_home)
|
|
raise
|
|
finally:
|
|
reset_secret_scope(secret_token)
|
|
reset_hermes_home_override(home_token)
|
|
except Exception as exc:
|
|
msg = str(exc)
|
|
with suppress(Exception):
|
|
from tools.mcp_oauth import humanize_oauth_registration_error
|
|
|
|
msg = humanize_oauth_registration_error(
|
|
server_name, exc, server_url=cfg.get("url") if isinstance(cfg, dict) else None
|
|
) or msg
|
|
if flow is not None:
|
|
flow.mark_error(msg)
|
|
finally:
|
|
if flow is not None:
|
|
flow.mark_worker_done()
|
|
if rec is not None:
|
|
_shutdown_listener(rec)
|
|
|
|
|
|
def start_flow(
|
|
hermes_home: str,
|
|
server_name: str,
|
|
cfg: dict,
|
|
*,
|
|
reconnect_live: bool = False,
|
|
url_timeout: float = 30.0,
|
|
client_redirect_uri: Optional[str] = None) -> Dict[str, Any]:
|
|
"""Begin an MCP OAuth flow and return ``{session_id, auth_url, flow}``; blocks up
|
|
to ``url_timeout`` for the authorization URL. With ``client_redirect_uri`` (remote
|
|
backend; invalid values raise ``ValueError``) no gateway-side listener is bound and
|
|
the client relays ``code``/``state`` via ``deliver_callback_flow``."""
|
|
from tools.mcp_dashboard_oauth import DashboardOAuthFlow
|
|
|
|
if client_redirect_uri is not None:
|
|
client_redirect_uri = _validate_client_redirect_uri(client_redirect_uri)
|
|
|
|
_gc_sessions()
|
|
|
|
with _sessions_lock:
|
|
active = [r for r in _sessions.values() if not r["flow"].worker_done]
|
|
if len(active) >= _MAX_PENDING:
|
|
raise RuntimeError("Too many MCP OAuth flows are already in progress")
|
|
if any(r["server_name"] == server_name and r["hermes_home"] == hermes_home for r in active):
|
|
raise RuntimeError(f"MCP OAuth for '{server_name}' is already in progress")
|
|
|
|
session_id = secrets.token_urlsafe(24)
|
|
flow = DashboardOAuthFlow(
|
|
flow_id=session_id, server_name=server_name, profile=None, hermes_home=hermes_home,
|
|
redirect_uri="", # set below once the loopback port is known
|
|
reconnect_live=reconnect_live)
|
|
# Client-hosted listener: a 127.0.0.1 port here would be unreachable from the browser.
|
|
httpd = None if client_redirect_uri else _start_loopback_listener(flow)
|
|
flow.redirect_uri = (
|
|
client_redirect_uri or f"http://127.0.0.1:{httpd.server_address[1]}/callback")
|
|
|
|
rec = {
|
|
"session_id": session_id, "server_name": server_name, "hermes_home": hermes_home,
|
|
"flow": flow, "httpd": httpd, "created_at": time.time(),
|
|
}
|
|
with _sessions_lock:
|
|
_sessions[session_id] = rec
|
|
|
|
threading.Thread(
|
|
target=_worker, args=(session_id, hermes_home, server_name, dict(cfg), reconnect_live),
|
|
daemon=True, name=f"mcp-oauth-{server_name}").start()
|
|
|
|
try:
|
|
auth_url = None
|
|
# wait_for_authorization_url is async; run its wait synchronously.
|
|
deadline = time.time() + url_timeout
|
|
while time.time() < deadline:
|
|
snap = flow.snapshot()
|
|
if auth_url := snap.get("authorization_url"):
|
|
break
|
|
if snap.get("status") == "error":
|
|
raise RuntimeError(snap.get("error") or "MCP OAuth flow failed before authorization")
|
|
time.sleep(0.1)
|
|
if not auth_url:
|
|
raise TimeoutError("Timed out waiting for MCP authorization URL")
|
|
except Exception:
|
|
flow.mark_error("Timed out waiting for MCP authorization URL")
|
|
_shutdown_listener(rec)
|
|
raise
|
|
|
|
# ``flow`` mirrors the provider-OAuth discriminator: open a URL then poll
|
|
# (no user_code to type, unlike device_code).
|
|
return {"session_id": session_id, "auth_url": auth_url, "flow": "pkce"}
|
|
|
|
|
|
def _lookup(session_id: str, server_name: str) -> "tuple[Dict[str, Any] | None, str | None]":
|
|
"""Find a session record; returns ``(rec, None)`` or ``(None, error_message)``."""
|
|
with _sessions_lock:
|
|
rec = _sessions.get(session_id)
|
|
if rec is None:
|
|
return None, "OAuth session not found or expired"
|
|
if rec["server_name"] != server_name:
|
|
return None, "server name mismatch for session"
|
|
return rec, None
|
|
|
|
|
|
def poll_flow(session_id: str, server_name: str) -> Dict[str, Any]:
|
|
"""Poll a session → ``{status, error_message?, auth_url?, tools?}``; ``status``
|
|
is ``pending`` | ``approved`` | ``error`` (the bridge's ``authorization_required``
|
|
maps to ``pending`` — the client only needs to know whether to keep waiting)."""
|
|
rec, err = _lookup(session_id, server_name)
|
|
if rec is None:
|
|
return {"status": "error", "error_message": err}
|
|
|
|
flow = rec["flow"]
|
|
snap = flow.snapshot()
|
|
raw = snap.get("status")
|
|
status = raw if raw in ("approved", "error") else "pending"
|
|
out: Dict[str, Any] = {
|
|
"session_id": session_id, "status": status, "error_message": snap.get("error"),
|
|
"auth_url": snap.get("authorization_url"),
|
|
}
|
|
if status == "approved":
|
|
out["tools"] = list(getattr(flow, "tools", []) or [])
|
|
return out
|
|
|
|
|
|
def deliver_callback_flow(
|
|
session_id: str, server_name: str, *, code: Optional[str], state: Optional[str],
|
|
error: Optional[str] = None,
|
|
) -> Dict[str, Any]:
|
|
"""Relay a client-captured OAuth redirect into a session's flow (remote-backend
|
|
companion to ``start_flow(client_redirect_uri=...)``). Security is unchanged:
|
|
``DashboardOAuthFlow.deliver_callback`` verifies ``state`` (constant-time) and
|
|
rejects replays. Returns ``{ok: true}`` or ``{ok: false, error_message}``."""
|
|
rec, err = _lookup(session_id, server_name)
|
|
if rec is None:
|
|
return {"ok": False, "error_message": err}
|
|
try:
|
|
rec["flow"].deliver_callback(code=code, state=state, error=error)
|
|
except ValueError as exc:
|
|
return {"ok": False, "error_message": str(exc)}
|
|
return {"ok": True, "session_id": session_id}
|