diff --git a/src/fastmcp/server/auth/providers/aws.py b/src/fastmcp/server/auth/providers/aws.py index 1e4f913b6..58f1ab5a1 100644 --- a/src/fastmcp/server/auth/providers/aws.py +++ b/src/fastmcp/server/auth/providers/aws.py @@ -26,7 +26,6 @@ from __future__ import annotations from key_value.aio.protocols import AsyncKeyValue from pydantic import AnyHttpUrl -from fastmcp.server.auth import TokenVerifier from fastmcp.server.auth.auth import AccessToken from fastmcp.server.auth.oidc_proxy import OIDCProxy from fastmcp.server.auth.providers.jwt import JWTVerifier @@ -147,6 +146,7 @@ class AWSCognitoProvider(OIDCProxy): # Store Cognito-specific info for claim filtering self.user_pool_id = user_pool_id self.aws_region = aws_region + self.client_id = client_id # Initialize OIDC proxy with Cognito discovery super().__init__( @@ -178,7 +178,7 @@ class AWSCognitoProvider(OIDCProxy): audience: str | None = None, required_scopes: list[str] | None = None, timeout_seconds: int | None = None, - ) -> TokenVerifier: + ) -> AWSCognitoTokenVerifier: """Creates a Cognito-specific token verifier with claim filtering. Args: @@ -189,7 +189,7 @@ class AWSCognitoProvider(OIDCProxy): """ return AWSCognitoTokenVerifier( issuer=str(self.oidc_config.issuer), - audience=audience, + audience=audience or self.client_id, algorithm=algorithm, jwks_uri=str(self.oidc_config.jwks_uri), required_scopes=required_scopes, diff --git a/tests/server/auth/providers/test_aws.py b/tests/server/auth/providers/test_aws.py index 9831dfaf8..59d51e4c8 100644 --- a/tests/server/auth/providers/test_aws.py +++ b/tests/server/auth/providers/test_aws.py @@ -102,6 +102,36 @@ class TestAWSCognitoProvider: assert provider._upstream_token_endpoint is not None assert "amazoncognito.com" in provider._upstream_authorization_endpoint + def test_token_verifier_defaults_audience_to_client_id(self): + """Test Cognito token verifier enforces the configured client ID by default.""" + with mock_cognito_oidc_discovery(): + provider = AWSCognitoProvider( + user_pool_id="us-east-1_XXXXXXXXX", + client_id="test_client", + client_secret="test_secret", + base_url="https://example.com", + jwt_signing_key="test-secret", + ) + + verifier = provider.get_token_verifier() + + assert verifier.audience == "test_client" + + def test_token_verifier_supports_audience_override(self): + """Test Cognito token verifier still allows explicit audience overrides.""" + with mock_cognito_oidc_discovery(): + provider = AWSCognitoProvider( + user_pool_id="us-east-1_XXXXXXXXX", + client_id="test_client", + client_secret="test_secret", + base_url="https://example.com", + jwt_signing_key="test-secret", + ) + + verifier = provider.get_token_verifier(audience="custom-audience") + + assert verifier.audience == "custom-audience" + # Token verification functionality is now tested as part of the OIDC provider integration # The CognitoTokenVerifier class is an internal implementation detail