diff --git a/docs/servers/auth/oauth-proxy.mdx b/docs/servers/auth/oauth-proxy.mdx index d60fd2c98..9e01d92c6 100644 --- a/docs/servers/auth/oauth-proxy.mdx +++ b/docs/servers/auth/oauth-proxy.mdx @@ -125,6 +125,16 @@ mcp = FastMCP(name="My Server", auth=auth) - `False`: Disable only if upstream provider doesn't support PKCE + + Token endpoint authentication method for the upstream OAuth server. Controls how the proxy authenticates when exchanging authorization codes and refresh tokens with the upstream provider. + - `"client_secret_basic"`: Send credentials in Authorization header (most common) + - `"client_secret_post"`: Send credentials in request body (required by some providers) + - `"none"`: No authentication (for public clients) + - `None` (default): Uses authlib's default (typically `"client_secret_basic"`) + + Set this if your provider requires a specific authentication method and the default doesn't work. + + List of allowed redirect URI patterns for MCP clients. Patterns support wildcards (e.g., `"http://localhost:*"`, `"https://*.example.com/*"`). - `None` (default): All redirect URIs allowed (for MCP/DCR compatibility) diff --git a/src/fastmcp/server/auth/oauth_proxy.py b/src/fastmcp/server/auth/oauth_proxy.py index dcb7d56d5..25feb4a64 100644 --- a/src/fastmcp/server/auth/oauth_proxy.py +++ b/src/fastmcp/server/auth/oauth_proxy.py @@ -248,6 +248,8 @@ class OAuthProxy(OAuthProvider): valid_scopes: list[str] | None = None, # PKCE configuration forward_pkce: bool = True, + # Token endpoint authentication + token_endpoint_auth_method: str | None = None, ): """Initialize the OAuth proxy provider. @@ -273,6 +275,9 @@ class OAuthProxy(OAuthProvider): forward_pkce: Whether to forward PKCE to upstream server (default True). Enable for providers that support/require PKCE (Google, Azure, etc.). Disable only if upstream provider doesn't support PKCE. + token_endpoint_auth_method: Token endpoint authentication method for upstream server. + Common values: "client_secret_basic", "client_secret_post", "none". + If None, authlib will use its default (typically "client_secret_basic"). """ # Always enable DCR since we implement it locally for MCP clients client_registration_options = ClientRegistrationOptions( @@ -311,6 +316,9 @@ class OAuthProxy(OAuthProvider): # PKCE configuration self._forward_pkce = forward_pkce + # Token endpoint authentication + self._token_endpoint_auth_method = token_endpoint_auth_method + # Local state for DCR and token bookkeeping self._clients: dict[str, OAuthClientInformationFull] = {} self._access_tokens: dict[str, AccessToken] = {} @@ -622,6 +630,7 @@ class OAuthProxy(OAuthProvider): oauth_client = AsyncOAuth2Client( client_id=self._upstream_client_id, client_secret=self._upstream_client_secret.get_secret_value(), + token_endpoint_auth_method=self._token_endpoint_auth_method, timeout=HTTP_TIMEOUT_SECONDS, ) @@ -849,6 +858,7 @@ class OAuthProxy(OAuthProvider): oauth_client = AsyncOAuth2Client( client_id=self._upstream_client_id, client_secret=self._upstream_client_secret.get_secret_value(), + token_endpoint_auth_method=self._token_endpoint_auth_method, timeout=HTTP_TIMEOUT_SECONDS, ) diff --git a/tests/server/auth/test_oauth_proxy.py b/tests/server/auth/test_oauth_proxy.py index f171f196a..de151092e 100644 --- a/tests/server/auth/test_oauth_proxy.py +++ b/tests/server/auth/test_oauth_proxy.py @@ -1,7 +1,7 @@ """Comprehensive tests for OAuth Proxy Provider functionality.""" import time -from unittest.mock import Mock +from unittest.mock import AsyncMock, Mock, patch from urllib.parse import parse_qs, urlparse import pytest @@ -9,7 +9,7 @@ from mcp.server.auth.provider import AuthorizationParams from mcp.shared.auth import OAuthClientInformationFull from pydantic import AnyUrl -from fastmcp.server.auth.auth import AccessToken +from fastmcp.server.auth.auth import AccessToken, RefreshToken from fastmcp.server.auth.oauth_proxy import OAuthProxy from fastmcp.server.auth.providers.jwt import JWTVerifier @@ -719,3 +719,178 @@ class TestOAuthProxyPKCE: txn_id = query_params["state"][0] transaction = oauth_proxy_with_pkce._oauth_transactions[txn_id] assert "proxy_code_verifier" not in transaction + + +class TestOAuthProxyTokenAuthMethod: + """Test suite for OAuth Proxy token_endpoint_auth_method parameter.""" + + @pytest.fixture + def jwt_verifier(self): + """Create a mock JWT verifier for testing.""" + verifier = Mock() + verifier.required_scopes = ["read", "write"] + return verifier + + def test_initialization_with_token_auth_method(self, jwt_verifier): + """Test that token_endpoint_auth_method is stored correctly when provided.""" + proxy = OAuthProxy( + upstream_authorization_endpoint="https://oauth.example.com/authorize", + upstream_token_endpoint="https://oauth.example.com/token", + upstream_client_id="client-id", + upstream_client_secret="client-secret", + token_verifier=jwt_verifier, + base_url="https://proxy.example.com", + token_endpoint_auth_method="client_secret_post", + ) + + assert proxy._token_endpoint_auth_method == "client_secret_post" + + def test_initialization_without_token_auth_method(self, jwt_verifier): + """Test that token_endpoint_auth_method defaults to None when not provided.""" + proxy = OAuthProxy( + upstream_authorization_endpoint="https://oauth.example.com/authorize", + upstream_token_endpoint="https://oauth.example.com/token", + upstream_client_id="client-id", + upstream_client_secret="client-secret", + token_verifier=jwt_verifier, + base_url="https://proxy.example.com", + ) + + assert proxy._token_endpoint_auth_method is None + + def test_different_auth_methods(self, jwt_verifier): + """Test initialization with different authentication methods.""" + # Test client_secret_basic + proxy_basic = OAuthProxy( + upstream_authorization_endpoint="https://oauth.example.com/authorize", + upstream_token_endpoint="https://oauth.example.com/token", + upstream_client_id="client-id", + upstream_client_secret="client-secret", + token_verifier=jwt_verifier, + base_url="https://proxy.example.com", + token_endpoint_auth_method="client_secret_basic", + ) + assert proxy_basic._token_endpoint_auth_method == "client_secret_basic" + + # Test none (for public clients) + proxy_none = OAuthProxy( + upstream_authorization_endpoint="https://oauth.example.com/authorize", + upstream_token_endpoint="https://oauth.example.com/token", + upstream_client_id="client-id", + upstream_client_secret="client-secret", + token_verifier=jwt_verifier, + base_url="https://proxy.example.com", + token_endpoint_auth_method="none", + ) + assert proxy_none._token_endpoint_auth_method == "none" + + @pytest.mark.asyncio + async def test_token_auth_method_passed_to_oauth_client(self, jwt_verifier): + """Test that token_endpoint_auth_method is passed to AsyncOAuth2Client during token exchange.""" + proxy = OAuthProxy( + upstream_authorization_endpoint="https://oauth.example.com/authorize", + upstream_token_endpoint="https://oauth.example.com/token", + upstream_client_id="client-id", + upstream_client_secret="client-secret", + token_verifier=jwt_verifier, + base_url="https://proxy.example.com", + token_endpoint_auth_method="client_secret_post", + ) + + # Mock AsyncOAuth2Client + with patch( + "fastmcp.server.auth.oauth_proxy.AsyncOAuth2Client" + ) as MockOAuth2Client: + mock_client = AsyncMock() + mock_client.refresh_token = AsyncMock( + return_value={ + "access_token": "new-access-token", + "refresh_token": "new-refresh-token", + "expires_in": 3600, + } + ) + MockOAuth2Client.return_value = mock_client + + # Test exchange_refresh_token + client = OAuthClientInformationFull( + client_id="test-client", + client_secret="test-secret", + redirect_uris=[AnyUrl("http://localhost:12345/callback")], + ) + + refresh_token = RefreshToken( + token="old-refresh-token", + client_id="test-client", + scopes=["read"], + expires_at=None, + ) + + await proxy.exchange_refresh_token(client, refresh_token, ["read"]) + + # Verify AsyncOAuth2Client was called with token_endpoint_auth_method + MockOAuth2Client.assert_called_with( + client_id="client-id", + client_secret="client-secret", + token_endpoint_auth_method="client_secret_post", + timeout=30.0, + ) + + @pytest.mark.asyncio + async def test_token_auth_method_in_idp_callback(self, jwt_verifier): + """Test that token_endpoint_auth_method is passed during IdP callback handling.""" + proxy = OAuthProxy( + upstream_authorization_endpoint="https://oauth.example.com/authorize", + upstream_token_endpoint="https://oauth.example.com/token", + upstream_client_id="client-id", + upstream_client_secret="client-secret", + token_verifier=jwt_verifier, + base_url="https://proxy.example.com", + token_endpoint_auth_method="client_secret_basic", + ) + + # Set up OAuth transaction + txn_id = "test-transaction" + proxy._oauth_transactions[txn_id] = { + "client_id": "test-client", + "redirect_uri": "http://localhost:12345/callback", + "state": "client-state", + "scopes": ["read"], + "proxy_code_verifier": "verifier123", + } + + # Register client + client = OAuthClientInformationFull( + client_id="test-client", + client_secret="test-secret", + redirect_uris=[AnyUrl("http://localhost:12345/callback")], + ) + await proxy.register_client(client) + + # Mock AsyncOAuth2Client and Request + with patch( + "fastmcp.server.auth.oauth_proxy.AsyncOAuth2Client" + ) as MockOAuth2Client: + mock_client = AsyncMock() + mock_client.fetch_token = AsyncMock( + return_value={ + "access_token": "idp-access-token", + "refresh_token": "idp-refresh-token", + "expires_in": 3600, + } + ) + MockOAuth2Client.return_value = mock_client + + # Create mock request with query params + mock_request = Mock() + mock_request.query_params = {"code": "idp-auth-code", "state": txn_id} + + # Simulate IdP callback + await proxy._handle_idp_callback(mock_request) + + # Verify AsyncOAuth2Client was called with token_endpoint_auth_method + MockOAuth2Client.assert_called_with( + client_id="client-id", + client_secret="client-secret", + token_endpoint_auth_method="client_secret_basic", + timeout=30.0, + )