"""Dashboard-mediated callback bridge for MCP OAuth. The MCP SDK remains responsible for discovery, DCR, PKCE, state validation and token exchange. This module only moves the two human/browser callbacks from a loopback listener into the already-authenticated dashboard session. """ from __future__ import annotations import asyncio import contextvars import secrets import threading import time from contextlib import contextmanager from dataclasses import dataclass, field from typing import Iterator from urllib.parse import parse_qs, urlparse @contextmanager def contextvar_set(var: contextvars.ContextVar, value) -> Iterator[None]: """Set *var* to *value* for the block, restoring the previous value after.""" token = var.set(value) try: yield finally: var.reset(token) def _event_field(): return field(default_factory=threading.Event, init=False, repr=False) @dataclass class DashboardOAuthFlow: flow_id: str server_name: str profile: str | None hermes_home: str redirect_uri: str reconnect_live: bool = False created_at: float = field(default_factory=time.time) status: str = "starting" authorization_url: str | None = None error: str | None = None tools: list[dict] = field(default_factory=list) expected_state: str | None = field(default=None, init=False) _callback: tuple[str, str | None] | None = field(default=None, init=False, repr=False) _callback_error: str | None = field(default=None, init=False, repr=False) _authorization_ready: threading.Event = _event_field() _callback_ready: threading.Event = _event_field() _worker_done: threading.Event = _event_field() _lock: threading.Lock = field(default_factory=threading.Lock, init=False, repr=False) async def publish_authorization_url(self, url: str) -> None: """Record the SDK's authorization URL (with its ``state``) for the dashboard to show.""" state = parse_qs(urlparse(url).query).get("state", [None])[0] if not state: raise ValueError("OAuth authorization URL did not include state") with self._lock: if self.status in {"approved", "error"}: raise RuntimeError("OAuth flow already ended") self.expected_state = state self.authorization_url = url self.status = "authorization_required" self._authorization_ready.set() @staticmethod async def _await_event(event: threading.Event, timeout: float, message: str) -> None: if not await asyncio.to_thread(event.wait, timeout): raise TimeoutError(message) async def wait_for_authorization_url(self, timeout: float = 30.0) -> str: await self._await_event(self._authorization_ready, timeout, "Timed out waiting for MCP authorization URL") if not self.authorization_url: raise RuntimeError(self.error or "MCP OAuth flow ended before authorization") return self.authorization_url def deliver_callback(self, *, code: str | None, state: str | None, error: str | None) -> None: """Hand the browser redirect to the waiting flow; ``state`` must match exactly.""" with self._lock: if self._callback_ready.is_set(): raise ValueError("OAuth callback already received") if self.expected_state is None or state is None or not secrets.compare_digest(self.expected_state, state): raise ValueError("OAuth callback state mismatch") if error: self._callback_error = error elif code: self._callback = (code, state) else: self._callback_error = "OAuth callback did not include code or error" self._callback_ready.set() async def wait_for_callback(self, timeout: float = 300.0) -> tuple[str, str | None]: await self._await_event(self._callback_ready, timeout, "Timed out waiting for MCP OAuth callback") if self._callback_error: raise RuntimeError(f"OAuth authorization failed: {self._callback_error}") if self._callback is None: raise RuntimeError("OAuth callback did not include an authorization code") return self._callback def mark_approved(self) -> None: with self._lock: if self.status == "error": raise RuntimeError("OAuth flow already ended") self.status = "approved" self.error = None def mark_error(self, error: str) -> None: with self._lock: if self.status == "approved": return self.status = "error" self.error = error self._authorization_ready.set() self._callback_ready.set() def snapshot(self) -> dict: with self._lock: return { "flow_id": self.flow_id, "server_name": self.server_name, "status": self.status, "authorization_url": self.authorization_url, "error": self.error, } def mark_worker_done(self) -> None: self._worker_done.set() @property def worker_done(self) -> bool: return self._worker_done.is_set() _current_dashboard_flow: contextvars.ContextVar[DashboardOAuthFlow | None] = ( contextvars.ContextVar("mcp_dashboard_oauth_flow", default=None) ) def dashboard_oauth_flow(flow: DashboardOAuthFlow): """Make *flow* the active dashboard OAuth flow for the block (ContextVar-scoped).""" return contextvar_set(_current_dashboard_flow, flow) def get_dashboard_oauth_flow() -> DashboardOAuthFlow | None: return _current_dashboard_flow.get()