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:
Teknium
2026-09-02 16:24:53 -07:00
parent 088eda22ae
commit b2d1d4cc1f
2 changed files with 16 additions and 25 deletions
+10 -14
View File
@@ -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:
+6 -11
View File
@@ -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: