diff --git a/tools/mcp_dashboard_oauth.py b/tools/mcp_dashboard_oauth.py index 4f0841971b..fba7b29977 100644 --- a/tools/mcp_dashboard_oauth.py +++ b/tools/mcp_dashboard_oauth.py @@ -28,6 +28,10 @@ def contextvar_set(var: contextvars.ContextVar, value) -> Iterator[None]: var.reset(token) +def _event_field(): + return field(default_factory=threading.Event, init=False, repr=False) + + @dataclass class DashboardOAuthFlow: flow_id: str @@ -44,9 +48,9 @@ class DashboardOAuthFlow: 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 = field(default_factory=threading.Event, init=False, repr=False) - _callback_ready: threading.Event = field(default_factory=threading.Event, init=False, repr=False) - _worker_done: threading.Event = field(default_factory=threading.Event, 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: @@ -68,9 +72,7 @@ class DashboardOAuthFlow: 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" - ) + 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 @@ -80,11 +82,7 @@ class DashboardOAuthFlow: 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) - ): + 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 @@ -95,9 +93,7 @@ class DashboardOAuthFlow: 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" - ) + 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: diff --git a/tools/mcp_oauth_provider.py b/tools/mcp_oauth_provider.py index 144dfc0fb6..80224338b6 100644 --- a/tools/mcp_oauth_provider.py +++ b/tools/mcp_oauth_provider.py @@ -41,7 +41,8 @@ class HermesProviderMixin: super().__init__(*args, **kwargs) self._hermes_token_user_agent = token_user_agent - def _stamp_token_user_agent(self, request): + def _prepare_token_request(self, request): + """Stamp the configured User-Agent onto a token/refresh request.""" ua = getattr(self, "_hermes_token_user_agent", None) # tests build via __new__ if ua: request.headers["User-Agent"] = ua @@ -61,13 +62,11 @@ class HermesProviderMixin: async def _exchange_token_authorization_code(self, *args: Any, **kwargs: Any): self._coerce_client_secret_post() - request = await super()._exchange_token_authorization_code(*args, **kwargs) - return self._stamp_token_user_agent(request) + return self._prepare_token_request(await super()._exchange_token_authorization_code(*args, **kwargs)) async def _refresh_token(self): self._coerce_client_secret_post() - request = await super()._refresh_token() - return self._stamp_token_user_agent(request) + return self._prepare_token_request(await super()._refresh_token()) async def _store_tokens(self, token_response) -> None: self.context.current_tokens = token_response @@ -109,9 +108,7 @@ class HermesProviderMixin: return True -def prepare_oauth_config( - server_name: str, server_url: str, oauth_config: dict | None -) -> tuple[dict, "HermesTokenStorage"]: +def prepare_oauth_config(server_name: str, server_url: str, oauth_config: dict | None) -> tuple[dict, "HermesTokenStorage"]: """Copy the ``oauth:`` block, apply provider defaults, open its token storage. The copy matters: later steps record ``_resolved_port`` / ``_cimd_url`` in @@ -124,9 +121,7 @@ def prepare_oauth_config( return cfg, mo.HermesTokenStorage(server_name) -def build_provider_kwargs( - cfg: dict, storage: "HermesTokenStorage", *, ssh_proxy_hint: bool -) -> dict[str, Any]: +def build_provider_kwargs(cfg: dict, storage: "HermesTokenStorage", *, ssh_proxy_hint: bool) -> dict[str, Any]: """Resolve the callback port and return the shared provider constructor kwargs. Runs the port → client-metadata → pre-registration sequence (order matters: