Merge pull request #894 from jlowin/claude-wt-20250620-120811

This commit is contained in:
Jeremiah Lowin 2025-06-20 12:25:01 -04:00 committed by GitHub
commit f97ace175f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 61 additions and 8 deletions

View file

@ -89,7 +89,7 @@ class RSAKeyPair:
self,
subject: str = "fastmcp-user",
issuer: str = "https://fastmcp.example.com",
audience: str | None = None,
audience: str | list[str] | None = None,
scopes: list[str] | None = None,
expires_in_seconds: int = 3600,
additional_claims: dict[str, Any] | None = None,
@ -102,7 +102,7 @@ class RSAKeyPair:
private_key_pem: RSA private key in PEM format
subject: Subject claim (usually user ID)
issuer: Issuer claim
audience: Audience claim (optional)
audience: Audience claim - can be a string or list of strings (optional)
scopes: List of scopes to include
expires_in_seconds: Token expiration time in seconds
additional_claims: Any additional claims to include
@ -161,7 +161,7 @@ class BearerAuthProvider(OAuthProvider):
public_key: str | None = None,
jwks_uri: str | None = None,
issuer: str | None = None,
audience: str | None = None,
audience: str | list[str] | None = None,
required_scopes: list[str] | None = None,
):
"""
@ -171,7 +171,7 @@ class BearerAuthProvider(OAuthProvider):
public_key: RSA public key in PEM format (for static key)
jwks_uri: URI to fetch keys from (for key rotation)
issuer: Expected issuer claim (optional)
audience: Expected audience claim (optional)
audience: Expected audience claim - can be a string or list of strings (optional)
required_scopes: List of required scopes for access (optional)
"""
if not (public_key or jwks_uri):
@ -312,11 +312,25 @@ class BearerAuthProvider(OAuthProvider):
# Validate audience if configured
if self.audience:
aud = claims.get("aud")
if isinstance(aud, list):
if self.audience not in aud:
# Handle different combinations of audience types
if isinstance(self.audience, list):
# self.audience is a list - check if any expected audience is present
if isinstance(aud, list):
# Both are lists - check for intersection
if not any(expected in aud for expected in self.audience):
return None
else:
# aud is a string - check if it's in our expected list
if aud not in self.audience:
return None
else:
# self.audience is a string - use original logic
if isinstance(aud, list):
if self.audience not in aud:
return None
elif aud != self.audience:
return None
elif aud != self.audience:
return None
# Extract claims - prefer client_id over sub for OAuth application identification
client_id = claims.get("client_id") or claims.get("sub") or "unknown"

View file

@ -446,6 +446,45 @@ class TestBearerToken:
access_token = await provider.load_access_token(token)
assert access_token is not None
async def test_provider_with_multiple_expected_audiences(
self, rsa_key_pair: RSAKeyPair
):
"""Test provider configured with multiple expected audiences."""
provider = BearerAuthProvider(
public_key=rsa_key_pair.public_key,
issuer="https://test.example.com",
audience=["https://api.example.com", "https://other-api.example.com"],
)
# Token with single audience that matches one of the expected
token1 = rsa_key_pair.create_token(
subject="test-user",
issuer="https://test.example.com",
audience="https://api.example.com",
)
access_token1 = await provider.load_access_token(token1)
assert access_token1 is not None
# Token with multiple audiences, one of which matches
token2 = rsa_key_pair.create_token(
subject="test-user",
issuer="https://test.example.com",
additional_claims={
"aud": ["https://api.example.com", "https://third-party.example.com"]
},
)
access_token2 = await provider.load_access_token(token2)
assert access_token2 is not None
# Token with audience that doesn't match any expected
token3 = rsa_key_pair.create_token(
subject="test-user",
issuer="https://test.example.com",
audience="https://wrong-api.example.com",
)
access_token3 = await provider.load_access_token(token3)
assert access_token3 is None
async def test_scope_extraction_string(
self, rsa_key_pair: RSAKeyPair, bearer_provider: BearerAuthProvider
):