feat: refresh oauth tokens for mcp sessions

This commit is contained in:
lda
2026-06-13 05:44:00 +07:00 Verified
parent 87a7cf9e86
commit dc415b6410
4 changed files with 153 additions and 2 deletions
+28
View File
@@ -167,6 +167,33 @@ class OAuthTokenRefresher(Protocol):
async def refresh(self, auth: OAuthRefreshTokenAuth) -> OAuthAccessToken: ...
class HttpxOAuthTokenRefresher:
"""Refresh OAuth2 access tokens from stored refresh-token credentials."""
async def refresh(self, auth: OAuthRefreshTokenAuth) -> OAuthAccessToken:
data: dict[str, str] = {
"grant_type": "refresh_token",
"client_id": auth.client_id,
"refresh_token": auth.refresh_token,
}
if auth.client_secret:
data["client_secret"] = auth.client_secret
if auth.scopes:
data["scope"] = " ".join(auth.scopes)
async with httpx.AsyncClient() as client:
response = await client.post(str(auth.token_url), data=data)
response.raise_for_status()
payload = response.json()
access_token = payload.get("access_token")
if not isinstance(access_token, str) or not access_token:
raise ValueError("OAuth token refresh response did not include access_token")
expires_in = payload.get("expires_in")
return OAuthAccessToken(
access_token=access_token,
expires_in=expires_in if isinstance(expires_in, int) else None,
)
@dataclass(frozen=True, slots=True)
class BoundMcpHttpAuth:
headers: dict[str, str] = field(default_factory=dict)
@@ -226,6 +253,7 @@ __all__ = [
"AuthRecord",
"BoundMcpHttpAuth",
"BoundMcpStdioAuth",
"HttpxOAuthTokenRefresher",
"McpAuthBinder",
"OAuthAccessToken",
"OAuthTokenRefresher",
+2 -2
View File
@@ -11,7 +11,7 @@ from mcp.client.stdio import StdioServerParameters, stdio_client
from mcp.client.streamable_http import streamable_http_client
from wf_api.auth import StoredAuthRecord, auth_record_from_compat
from wf_sources_mcp.auth import AuthRecord, McpAuthBinder
from wf_sources_mcp.auth import AuthRecord, HttpxOAuthTokenRefresher, McpAuthBinder
from wf_sources_mcp.connections import McpSourceConnection
from wf_sources_mcp.transports import HttpSourceTransport, StdioSourceTransport
@@ -49,7 +49,7 @@ async def open_mcp_session(
if transport is None:
raise ValueError(f"connection {connection.id!r} requires metadata.transport")
binder = auth_binder or McpAuthBinder()
binder = auth_binder or McpAuthBinder(oauth_refresher=HttpxOAuthTokenRefresher())
stored_auth = _as_stored_auth(auth)
if isinstance(transport, StdioSourceTransport):