nevzheng commented on code in PR #12531: URL: https://github.com/apache/gravitino/pull/12531#discussion_r3877237541
########## mcp-server/mcp_server/core/oauth.py: ########## @@ -0,0 +1,280 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +"""httpx ``auth=`` hook for MCP → Gravitino OAuth2 client-credentials. + +Uses ``httpx-auth`` for fetch and cache on the existing ``httpx.AsyncClient``. +Credentials go in the form body (``client_secret_post``), matching the +Java/Python Gravitino clients. httpx-auth defaults to HTTP Basic. This class +retries once after Gravitino HTTP 401. +""" + +import asyncio +import base64 +import json +import logging +import threading +import time +from collections.abc import AsyncGenerator, Generator +from typing import Optional, Union + +import httpx +from httpx_auth import AuthenticationFailed, OAuth2, OAuth2ClientCredentials + +_LOG = logging.getLogger(__name__) + +# Refresh this many seconds before recorded expiry. httpx-auth default is 30. +DEFAULT_REFRESH_SKEW_SECONDS = 60 + +_TokenTuple = Union[tuple[str, str], tuple[str, str, Union[int, str]]] + + +class RefreshableBearerAuth(OAuth2ClientCredentials): + """httpx-auth client-credentials with form POST and one 401 retry.""" + + requires_request_body = True + requires_response_body = True + + def __init__( + self, + *, + token_endpoint: str, + client_id: str, + client_secret: str, + scope: str = "", + refresh_skew_seconds: int = DEFAULT_REFRESH_SKEW_SECONDS, + client: Optional[httpx.Client] = None, + ): + """Build an ``auth=`` hook for the service hop. + + Args: + token_endpoint: Identity-provider token URL. + client_id: OAuth2 client id. + client_secret: OAuth2 client secret. + scope: Optional OAuth2 scope. + refresh_skew_seconds: httpx-auth ``early_expiry``. + client: Optional sync httpx client used only for token POSTs + (tests inject ``MockTransport`` here). + """ + kwargs = {"early_expiry": float(refresh_skew_seconds)} + if scope: + kwargs["scope"] = scope + if client is not None: + kwargs["client"] = client + super().__init__(token_endpoint, client_id, client_secret, **kwargs) + self._token_lock = asyncio.Lock() + self._sync_lock = threading.Lock() + self._rejected_tokens: set[str] = set() + self._retried_tokens: set[str] = set() + + def invalidate(self) -> None: + """Drop the cached token so the next call fetches a new one.""" + cache = OAuth2.token_cache + # TokenMemoryCache.clear() wipes every client; only drop ours. + with cache._forbid_concurrent_cache_access: # pylint: disable=protected-access + cache.tokens.pop(self.state, None) + + def request_new_token(self) -> _TokenTuple: + """POST ``client_credentials`` with id/secret in the form body.""" + data = self._token_form_data() + client = self.client or httpx.Client() + self._configure_client(client) + try: + response = client.post(self.token_url, data=data) + self._log_token_http_error(response) + response.raise_for_status() + body = response.json() + finally: + if self.client is None: + client.close() + return self._token_tuple(body) + + async def request_new_token_async(self) -> _TokenTuple: + """POST ``client_credentials`` without blocking the event loop.""" + if self.client is not None: + return await asyncio.to_thread(self.request_new_token) + data = self._token_form_data() + async with httpx.AsyncClient() as client: + client.timeout = self.timeout + response = await client.post(self.token_url, data=data) + self._log_token_http_error(response) + response.raise_for_status() + body = response.json() + return self._token_tuple(body) + + def auth_flow( + self, request: httpx.Request + ) -> Generator[httpx.Request, httpx.Response, None]: + """Attach a cached or freshly fetched Bearer; retry once on HTTP 401.""" + if self.requires_request_body: + request.read() + token, fetched = self._apply_token(request) + response = yield request + if response.status_code != 401: + return + with self._sync_lock: + if not self._begin_401_retry(token, fetched): + return + self._invalidate_if_still_cached(token) + retry_token, _ = self._apply_token(request) + response = yield request + if response.status_code == 401: + self._rejected_tokens.add(retry_token) + + async def async_auth_flow( + self, request: httpx.Request + ) -> AsyncGenerator[httpx.Request, httpx.Response]: + """Attach a Bearer without a blocking IdP POST on the event loop.""" + if self.requires_request_body: + await request.aread() + token, fetched = await self._apply_token_async(request) + response = yield request + if response.status_code != 401: + return + async with self._token_lock: + if not self._begin_401_retry(token, fetched): + return + self._invalidate_if_still_cached(token) + retry_token, _ = await self._apply_token_async(request) Review Comment: @yuqi1129 Recommend keeping **request-driven refresh** with the 60s `early_expiry` skew. TTL 3600s → cache is stale at ~3540s. The **next** tool call POSTs the IdP, then calls Gravitino. Idle MCP (stdio / Cursor overnight) does not hit the IdP. ```mermaid sequenceDiagram participant Tool participant MCP participant IdP participant Gravitino Tool->>MCP: call (t=0) MCP->>IdP: POST /token IdP-->>MCP: tok-1 expires_in=3600 MCP->>Gravitino: Bearer tok-1 Note over MCP: t=1..3540 reuse, no IdP Tool->>MCP: call after idle (t>3540) MCP->>IdP: POST /token IdP-->>MCP: tok-2 MCP->>Gravitino: Bearer tok-2 ``` TTL expiry itself is not a user-facing error. That first post-idle call waits one IdP RTT. If Gravitino still 401s (clock), we already retry once with a new token. Do you have a CUJ that needs a background timer so the first post-idle call never waits on the IdP? If so, please reopen and we can follow up. Nevin Sent from my 🤖 (Cursor) -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected]
