feat: refresh oauth tokens for mcp sessions
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -79,6 +79,62 @@ async def test_mcp_binder_refreshes_oauth_for_http() -> None:
|
||||
assert len(refresher.calls) == 1
|
||||
|
||||
|
||||
async def test_httpx_oauth_refresher_posts_refresh_token_grant(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from pydantic import AnyUrl
|
||||
|
||||
from wf_sources_mcp import auth as mod
|
||||
from wf_sources_mcp.auth import HttpxOAuthTokenRefresher
|
||||
|
||||
captured_posts: list[tuple[str, dict[str, str]]] = []
|
||||
|
||||
class _Response:
|
||||
def raise_for_status(self) -> None:
|
||||
return None
|
||||
|
||||
def json(self) -> dict[str, object]:
|
||||
return {"access_token": "access-token", "expires_in": 3600}
|
||||
|
||||
class _Client:
|
||||
async def __aenter__(self) -> "_Client":
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
async def post(self, url: str, *, data: dict[str, str]) -> _Response:
|
||||
captured_posts.append((url, data))
|
||||
return _Response()
|
||||
|
||||
monkeypatch.setattr(mod.httpx, "AsyncClient", _Client)
|
||||
|
||||
token = await HttpxOAuthTokenRefresher().refresh(
|
||||
OAuthRefreshTokenAuth(
|
||||
client_id="client",
|
||||
client_secret="secret",
|
||||
refresh_token="refresh",
|
||||
token_url=AnyUrl("https://oauth2.googleapis.com/token"),
|
||||
scopes=("scope.one", "scope.two"),
|
||||
)
|
||||
)
|
||||
|
||||
assert token.access_token == "access-token"
|
||||
assert token.expires_in == 3600
|
||||
assert captured_posts == [
|
||||
(
|
||||
"https://oauth2.googleapis.com/token",
|
||||
{
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": "client",
|
||||
"client_secret": "secret",
|
||||
"refresh_token": "refresh",
|
||||
"scope": "scope.one scope.two",
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
async def test_mcp_binder_rejects_env_for_http() -> None:
|
||||
binder = McpAuthBinder()
|
||||
record = StoredAuthRecord(id="demo.auth", auth=EnvAuth(env={"TOKEN": "abc"}))
|
||||
|
||||
@@ -296,6 +296,73 @@ async def test_open_mcp_session_uses_binder_for_http_headers(
|
||||
assert captured_clients[0]["headers"] == {"Authorization": "Bearer token"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_open_mcp_session_refreshes_oauth_record_for_http(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from pydantic import AnyUrl
|
||||
|
||||
from wf_api.auth import OAuthRefreshTokenAuth, StoredAuthRecord
|
||||
from wf_sources_mcp.client.transport import open_mcp_session
|
||||
|
||||
captured_clients: list[dict[str, Any]] = []
|
||||
captured_posts: list[tuple[str, dict[str, str]]] = []
|
||||
|
||||
class _Response:
|
||||
def raise_for_status(self) -> None:
|
||||
return None
|
||||
|
||||
def json(self) -> dict[str, object]:
|
||||
return {"access_token": "fresh-access-token"}
|
||||
|
||||
class _CapturingClient:
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
captured_clients.append(kwargs)
|
||||
|
||||
async def __aenter__(self) -> "_CapturingClient":
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
async def post(self, url: str, *, data: dict[str, str]) -> _Response:
|
||||
captured_posts.append((url, data))
|
||||
return _Response()
|
||||
|
||||
import wf_sources_mcp.auth as auth_mod
|
||||
import wf_sources_mcp.client.transport as transport_mod
|
||||
|
||||
monkeypatch.setattr(auth_mod.httpx, "AsyncClient", _CapturingClient)
|
||||
monkeypatch.setattr(transport_mod.httpx, "AsyncClient", _CapturingClient)
|
||||
|
||||
connection = _http_connection()
|
||||
auth = StoredAuthRecord(
|
||||
id="google.drive.personal",
|
||||
auth=OAuthRefreshTokenAuth(
|
||||
client_id="client",
|
||||
client_secret="secret",
|
||||
refresh_token="refresh",
|
||||
token_url=AnyUrl("https://oauth2.googleapis.com/token"),
|
||||
),
|
||||
)
|
||||
|
||||
async with open_mcp_session(connection, auth):
|
||||
pass
|
||||
|
||||
assert captured_posts[0] == (
|
||||
"https://oauth2.googleapis.com/token",
|
||||
{
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": "client",
|
||||
"client_secret": "secret",
|
||||
"refresh_token": "refresh",
|
||||
},
|
||||
)
|
||||
assert captured_clients[-1]["headers"] == {
|
||||
"Authorization": "Bearer fresh-access-token"
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_open_mcp_session_uses_binder_for_stdio_env() -> None:
|
||||
import wf_sources_mcp.client.transport as mod # noqa: I001
|
||||
|
||||
Reference in New Issue
Block a user