refactor(mcp): compact mcp_oauth_provider token-request stamping and dashboard flow dataclass
- _stamp_token_user_agent -> _prepare_token_request (no external refs) - DashboardOAuthFlow event fields via one _event_field() factory
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user