diff --git a/src/fastmcp/server/auth/providers/discord.py b/src/fastmcp/server/auth/providers/discord.py index 7ef187202..c07dfd86f 100644 --- a/src/fastmcp/server/auth/providers/discord.py +++ b/src/fastmcp/server/auth/providers/discord.py @@ -48,6 +48,7 @@ class DiscordTokenVerifier(TokenVerifier): def __init__( self, *, + expected_client_id: str, required_scopes: list[str] | None = None, timeout_seconds: int = 10, http_client: httpx.AsyncClient | None = None, @@ -55,6 +56,7 @@ class DiscordTokenVerifier(TokenVerifier): """Initialize the Discord token verifier. Args: + expected_client_id: Expected Discord OAuth client ID for audience binding required_scopes: Required OAuth scopes (e.g., ['email']) timeout_seconds: HTTP request timeout http_client: Optional httpx.AsyncClient for connection pooling. When provided, @@ -62,6 +64,7 @@ class DiscordTokenVerifier(TokenVerifier): lifecycle. When None (default), a fresh client is created per call. """ super().__init__(required_scopes=required_scopes) + self.expected_client_id = expected_client_id self.timeout_seconds = timeout_seconds self._http_client = http_client @@ -121,6 +124,13 @@ class DiscordTokenVerifier(TokenVerifier): user_data = token_info.get("user", {}) application = token_info.get("application") or {} client_id = str(application.get("id", "unknown")) + if client_id != self.expected_client_id: + logger.debug( + "Discord token app ID mismatch: expected %s, got %s", + self.expected_client_id, + client_id, + ) + return None # Create AccessToken with Discord user info access_token = AccessToken( @@ -235,6 +245,7 @@ class DiscordProvider(OAuthProxy): # Create Discord token verifier token_verifier = DiscordTokenVerifier( + expected_client_id=client_id, required_scopes=required_scopes_final, timeout_seconds=timeout_seconds, http_client=http_client, diff --git a/tests/server/auth/providers/test_discord.py b/tests/server/auth/providers/test_discord.py index 509eb0826..d40e180b1 100644 --- a/tests/server/auth/providers/test_discord.py +++ b/tests/server/auth/providers/test_discord.py @@ -1,9 +1,11 @@ """Tests for Discord OAuth provider.""" +from unittest.mock import AsyncMock, MagicMock, patch + import pytest from key_value.aio.stores.memory import MemoryStore -from fastmcp.server.auth.providers.discord import DiscordProvider +from fastmcp.server.auth.providers.discord import DiscordProvider, DiscordTokenVerifier @pytest.fixture @@ -81,3 +83,45 @@ class TestDiscordProvider: # Provider should initialize successfully with these scopes assert provider is not None + + def test_token_verifier_is_bound_to_provider_client_id( + self, memory_storage: MemoryStore + ): + """Test DiscordProvider binds token verifier to the configured client ID.""" + provider = DiscordProvider( + client_id="expected-client-id", + client_secret="GOCSPX-test123", + base_url="https://myserver.com", + jwt_signing_key="test-secret", + client_storage=memory_storage, + ) + + verifier = provider._token_validator + assert isinstance(verifier, DiscordTokenVerifier) + assert verifier.expected_client_id == "expected-client-id" + + +class TestDiscordTokenVerifier: + """Test DiscordTokenVerifier behavior.""" + + async def test_rejects_token_from_different_discord_application(self): + """Token must be bound to configured Discord client_id.""" + verifier = DiscordTokenVerifier(expected_client_id="expected-app-id") + + mock_client = AsyncMock() + token_info_response = MagicMock() + token_info_response.status_code = 200 + token_info_response.json.return_value = { + "application": {"id": "different-app-id"}, + "user": {"id": "123"}, + "scopes": ["identify"], + } + mock_client.get.return_value = token_info_response + + with patch( + "fastmcp.server.auth.providers.discord.httpx.AsyncClient" + ) as mock_client_class: + mock_client_class.return_value.__aenter__.return_value = mock_client + result = await verifier.verify_token("token") + + assert result is None diff --git a/tests/server/auth/providers/test_http_client.py b/tests/server/auth/providers/test_http_client.py index 34c118f38..fa6126183 100644 --- a/tests/server/auth/providers/test_http_client.py +++ b/tests/server/auth/providers/test_http_client.py @@ -245,7 +245,10 @@ class TestDiscordHttpClient: from fastmcp.server.auth.providers.discord import DiscordTokenVerifier client = httpx.AsyncClient() - verifier = DiscordTokenVerifier(http_client=client) + verifier = DiscordTokenVerifier( + expected_client_id="test-client-id", + http_client=client, + ) assert verifier._http_client is client