feat: option to add upstream claims to the FastMCP proxy JWT (#2997)

This commit is contained in:
Jonas Krüger Svensson 2026-01-28 21:57:24 +01:00 committed by GitHub
commit fa5b136205
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 658 additions and 11 deletions

View file

@ -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"}

View file

@ -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"