mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-24 06:24:18 +02:00
feat: option to add upstream claims to the FastMCP proxy JWT (#2997)
This commit is contained in:
parent
6edd5e699e
commit
fa5b136205
8 changed files with 658 additions and 11 deletions
|
|
@ -807,3 +807,244 @@ class TestOIDCScopeHandling:
|
|||
assert "profile" in result
|
||||
assert "api://my-api/openid" not in result
|
||||
assert "api://my-api/profile" not in result
|
||||
|
||||
|
||||
class TestAzureExtractUpstreamClaims:
|
||||
"""Tests for Azure provider's _extract_upstream_claims method."""
|
||||
|
||||
@staticmethod
|
||||
def create_test_jwt(claims: dict) -> str:
|
||||
"""Create a test JWT token with the given claims."""
|
||||
import base64
|
||||
import json
|
||||
|
||||
header = base64.urlsafe_b64encode(
|
||||
json.dumps({"alg": "RS256", "typ": "JWT"}).encode()
|
||||
).rstrip(b"=")
|
||||
payload = base64.urlsafe_b64encode(json.dumps(claims).encode()).rstrip(b"=")
|
||||
signature = base64.urlsafe_b64encode(b"fake-signature").rstrip(b"=")
|
||||
return f"{header.decode()}.{payload.decode()}.{signature.decode()}"
|
||||
|
||||
async def test_extract_claims_from_azure_jwt(self):
|
||||
"""Test that Azure identity claims are extracted from access token."""
|
||||
provider = AzureProvider(
|
||||
client_id="test_client",
|
||||
client_secret="test_secret",
|
||||
tenant_id="test-tenant",
|
||||
base_url="https://myserver.com",
|
||||
required_scopes=["read"],
|
||||
jwt_signing_key="test-secret",
|
||||
)
|
||||
|
||||
azure_jwt = self.create_test_jwt(
|
||||
{
|
||||
"sub": "user-subject-id",
|
||||
"oid": "user-object-id",
|
||||
"tid": "tenant-id-123",
|
||||
"azp": "client-app-id",
|
||||
"name": "Test User",
|
||||
"given_name": "Test",
|
||||
"family_name": "User",
|
||||
"preferred_username": "testuser@example.com",
|
||||
"upn": "testuser@example.com",
|
||||
"email": "test@example.com",
|
||||
"roles": ["Admin", "Reader"],
|
||||
"groups": ["group-1", "group-2"],
|
||||
"exp": 9999999999,
|
||||
"iat": 1234567890,
|
||||
"iss": "https://login.microsoftonline.com/test-tenant/v2.0",
|
||||
}
|
||||
)
|
||||
|
||||
idp_tokens = {
|
||||
"access_token": azure_jwt,
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
|
||||
claims = await provider._extract_upstream_claims(idp_tokens)
|
||||
|
||||
assert claims is not None
|
||||
assert claims["sub"] == "user-subject-id"
|
||||
assert claims["oid"] == "user-object-id"
|
||||
assert claims["tid"] == "tenant-id-123"
|
||||
assert claims["azp"] == "client-app-id"
|
||||
assert claims["name"] == "Test User"
|
||||
assert claims["given_name"] == "Test"
|
||||
assert claims["family_name"] == "User"
|
||||
assert claims["preferred_username"] == "testuser@example.com"
|
||||
assert claims["upn"] == "testuser@example.com"
|
||||
assert claims["email"] == "test@example.com"
|
||||
assert claims["roles"] == ["Admin", "Reader"]
|
||||
assert claims["groups"] == ["group-1", "group-2"]
|
||||
|
||||
async def test_extract_claims_only_includes_identity_claims(self):
|
||||
"""Test that only identity claims are extracted, not all JWT claims."""
|
||||
provider = AzureProvider(
|
||||
client_id="test_client",
|
||||
client_secret="test_secret",
|
||||
tenant_id="test-tenant",
|
||||
base_url="https://myserver.com",
|
||||
required_scopes=["read"],
|
||||
jwt_signing_key="test-secret",
|
||||
)
|
||||
|
||||
azure_jwt = self.create_test_jwt(
|
||||
{
|
||||
"sub": "user-id",
|
||||
"oid": "object-id",
|
||||
"name": "Test User",
|
||||
"exp": 9999999999,
|
||||
"iat": 1234567890,
|
||||
"iss": "https://issuer.example.com",
|
||||
"aud": "test-audience",
|
||||
"nbf": 1234567890,
|
||||
"scp": "read write",
|
||||
"azp": "some-client",
|
||||
}
|
||||
)
|
||||
|
||||
idp_tokens = {"access_token": azure_jwt}
|
||||
|
||||
claims = await provider._extract_upstream_claims(idp_tokens)
|
||||
|
||||
# Only identity claims should be present
|
||||
assert claims is not None
|
||||
assert "sub" in claims
|
||||
assert "oid" in claims
|
||||
assert "name" in claims
|
||||
assert "azp" in claims # azp is an identity claim we extract
|
||||
# Standard JWT claims should NOT be extracted
|
||||
assert "exp" not in claims
|
||||
assert "iat" not in claims
|
||||
assert "iss" not in claims
|
||||
assert "aud" not in claims
|
||||
assert "nbf" not in claims
|
||||
assert "scp" not in claims
|
||||
|
||||
async def test_extract_claims_returns_none_for_missing_access_token(self):
|
||||
"""Test that None is returned when access_token is missing."""
|
||||
provider = AzureProvider(
|
||||
client_id="test_client",
|
||||
client_secret="test_secret",
|
||||
tenant_id="test-tenant",
|
||||
base_url="https://myserver.com",
|
||||
required_scopes=["read"],
|
||||
jwt_signing_key="test-secret",
|
||||
)
|
||||
|
||||
idp_tokens = {"token_type": "Bearer", "expires_in": 3600}
|
||||
|
||||
claims = await provider._extract_upstream_claims(idp_tokens)
|
||||
|
||||
assert claims is None
|
||||
|
||||
async def test_extract_claims_returns_none_for_opaque_token(self):
|
||||
"""Test that None is returned for opaque (non-JWT) tokens."""
|
||||
provider = AzureProvider(
|
||||
client_id="test_client",
|
||||
client_secret="test_secret",
|
||||
tenant_id="test-tenant",
|
||||
base_url="https://myserver.com",
|
||||
required_scopes=["read"],
|
||||
jwt_signing_key="test-secret",
|
||||
)
|
||||
|
||||
idp_tokens = {
|
||||
"access_token": "gho_opaque_token_not_a_jwt", # Not a JWT
|
||||
"token_type": "Bearer",
|
||||
}
|
||||
|
||||
claims = await provider._extract_upstream_claims(idp_tokens)
|
||||
|
||||
assert claims is None
|
||||
|
||||
async def test_extract_claims_returns_none_for_malformed_jwt(self):
|
||||
"""Test that None is returned for malformed JWT tokens."""
|
||||
provider = AzureProvider(
|
||||
client_id="test_client",
|
||||
client_secret="test_secret",
|
||||
tenant_id="test-tenant",
|
||||
base_url="https://myserver.com",
|
||||
required_scopes=["read"],
|
||||
jwt_signing_key="test-secret",
|
||||
)
|
||||
|
||||
# Only two parts (missing signature)
|
||||
idp_tokens = {"access_token": "header.payload"}
|
||||
|
||||
claims = await provider._extract_upstream_claims(idp_tokens)
|
||||
|
||||
assert claims is None
|
||||
|
||||
async def test_extract_claims_returns_none_for_invalid_base64(self):
|
||||
"""Test that None is returned for JWT with invalid base64."""
|
||||
provider = AzureProvider(
|
||||
client_id="test_client",
|
||||
client_secret="test_secret",
|
||||
tenant_id="test-tenant",
|
||||
base_url="https://myserver.com",
|
||||
required_scopes=["read"],
|
||||
jwt_signing_key="test-secret",
|
||||
)
|
||||
|
||||
# Invalid base64 in payload
|
||||
idp_tokens = {"access_token": "header.not-valid-base64!!!.signature"}
|
||||
|
||||
claims = await provider._extract_upstream_claims(idp_tokens)
|
||||
|
||||
assert claims is None
|
||||
|
||||
async def test_extract_claims_returns_none_for_empty_identity_claims(self):
|
||||
"""Test that None is returned when no identity claims are present."""
|
||||
provider = AzureProvider(
|
||||
client_id="test_client",
|
||||
client_secret="test_secret",
|
||||
tenant_id="test-tenant",
|
||||
base_url="https://myserver.com",
|
||||
required_scopes=["read"],
|
||||
jwt_signing_key="test-secret",
|
||||
)
|
||||
|
||||
# JWT with only standard claims, no identity claims
|
||||
azure_jwt = self.create_test_jwt(
|
||||
{
|
||||
"exp": 9999999999,
|
||||
"iat": 1234567890,
|
||||
"iss": "https://issuer.example.com",
|
||||
"aud": "test-audience",
|
||||
}
|
||||
)
|
||||
|
||||
idp_tokens = {"access_token": azure_jwt}
|
||||
|
||||
claims = await provider._extract_upstream_claims(idp_tokens)
|
||||
|
||||
assert claims is None
|
||||
|
||||
async def test_extract_claims_partial_identity_claims(self):
|
||||
"""Test extraction when only some identity claims are present."""
|
||||
provider = AzureProvider(
|
||||
client_id="test_client",
|
||||
client_secret="test_secret",
|
||||
tenant_id="test-tenant",
|
||||
base_url="https://myserver.com",
|
||||
required_scopes=["read"],
|
||||
jwt_signing_key="test-secret",
|
||||
)
|
||||
|
||||
# JWT with only sub and name
|
||||
azure_jwt = self.create_test_jwt(
|
||||
{
|
||||
"sub": "user-id",
|
||||
"name": "Test User",
|
||||
"exp": 9999999999,
|
||||
}
|
||||
)
|
||||
|
||||
idp_tokens = {"access_token": azure_jwt}
|
||||
|
||||
claims = await provider._extract_upstream_claims(idp_tokens)
|
||||
|
||||
assert claims is not None
|
||||
assert claims == {"sub": "user-id", "name": "Test User"}
|
||||
|
|
|
|||
|
|
@ -228,3 +228,60 @@ class TestJWTIssuer:
|
|||
|
||||
with pytest.raises(JoseError):
|
||||
issuer.verify_token("header.payload") # Missing signature
|
||||
|
||||
def test_issue_access_token_with_upstream_claims(self, issuer):
|
||||
"""Test that upstream claims are included when provided."""
|
||||
upstream_claims = {
|
||||
"sub": "user-123",
|
||||
"oid": "object-id-456",
|
||||
"name": "Test User",
|
||||
"email": "test@example.com",
|
||||
"roles": ["Admin", "Reader"],
|
||||
}
|
||||
token = issuer.issue_access_token(
|
||||
client_id="client-abc",
|
||||
scopes=["read", "write"],
|
||||
jti="token-id-123",
|
||||
expires_in=3600,
|
||||
upstream_claims=upstream_claims,
|
||||
)
|
||||
|
||||
payload = issuer.verify_token(token)
|
||||
assert "upstream_claims" in payload
|
||||
assert payload["upstream_claims"]["sub"] == "user-123"
|
||||
assert payload["upstream_claims"]["oid"] == "object-id-456"
|
||||
assert payload["upstream_claims"]["name"] == "Test User"
|
||||
assert payload["upstream_claims"]["email"] == "test@example.com"
|
||||
assert payload["upstream_claims"]["roles"] == ["Admin", "Reader"]
|
||||
|
||||
def test_issue_access_token_without_upstream_claims(self, issuer):
|
||||
"""Test that upstream_claims is not present when not provided."""
|
||||
token = issuer.issue_access_token(
|
||||
client_id="client-abc",
|
||||
scopes=["read"],
|
||||
jti="token-id-123",
|
||||
expires_in=3600,
|
||||
)
|
||||
|
||||
payload = issuer.verify_token(token)
|
||||
assert "upstream_claims" not in payload
|
||||
|
||||
def test_issue_refresh_token_with_upstream_claims(self, issuer):
|
||||
"""Test that refresh tokens also include upstream claims when provided."""
|
||||
upstream_claims = {
|
||||
"sub": "user-123",
|
||||
"name": "Test User",
|
||||
}
|
||||
token = issuer.issue_refresh_token(
|
||||
client_id="client-abc",
|
||||
scopes=["read"],
|
||||
jti="refresh-token-id",
|
||||
expires_in=60 * 60 * 24 * 30,
|
||||
upstream_claims=upstream_claims,
|
||||
)
|
||||
|
||||
payload = issuer.verify_token(token)
|
||||
assert "upstream_claims" in payload
|
||||
assert payload["upstream_claims"]["sub"] == "user-123"
|
||||
assert payload["upstream_claims"]["name"] == "Test User"
|
||||
assert payload["token_use"] == "refresh"
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue