fix(identity): make JWTPrincipal hashable via tuple scopes and compare-excluded claims

This commit is contained in:
Yassin Kortam 2026-06-08 14:23:20 -07:00
parent bf33cc24e3
commit 30cb095859
5 changed files with 42 additions and 13 deletions

View file

@ -49,15 +49,16 @@ def _principal_from_uak(uak: "UserAPIKeyAuth") -> Principal:
claims = uak.jwt_claims
scope_claim = claims.get("scope") or claims.get("scp") or ""
if isinstance(scope_claim, list):
scopes = [str(s) for s in scope_claim if s]
scopes = tuple(str(s) for s in scope_claim if s)
elif isinstance(scope_claim, str):
scopes = [s for s in scope_claim.split(" ") if s]
scopes = tuple(s for s in scope_claim.split(" ") if s)
else:
scopes = []
scopes = ()
aud = claims.get("aud")
return JWTPrincipal(
sub=claims.get("sub"),
iss=claims.get("iss"),
aud=claims.get("aud"),
aud=tuple(aud) if isinstance(aud, list) else aud,
scopes=scopes,
claims=dict(claims),
mapped_user_id=uak.user_id,
@ -102,7 +103,9 @@ def identity_context_to_user_api_key_auth(
principal = ctx.principal
kwargs: dict = {
"end_user_id": ctx.end_user_id,
"access_group_ids": list(ctx.access_group_ids) if ctx.access_group_ids else None,
"access_group_ids": (
list(ctx.access_group_ids) if ctx.access_group_ids else None
),
}
if isinstance(principal, ApiKeyPrincipal):

View file

@ -10,7 +10,7 @@ for ``match``-style dispatch.
"""
from dataclasses import dataclass, field
from typing import Any, Dict, List, Literal, Optional, Union
from typing import Any, Dict, Literal, Optional, Tuple, Union
@dataclass(frozen=True)
@ -30,9 +30,9 @@ class JWTPrincipal:
kind: Literal["jwt"] = field(default="jwt", init=False)
sub: Optional[str] = None
iss: Optional[str] = None
aud: Optional[Union[str, List[str]]] = None
scopes: List[str] = field(default_factory=list)
claims: Dict[str, Any] = field(default_factory=dict)
aud: Optional[Union[str, Tuple[str, ...]]] = None
scopes: Tuple[str, ...] = field(default_factory=tuple)
claims: Dict[str, Any] = field(default_factory=dict, compare=False)
mapped_user_id: Optional[str] = None
mapped_team_id: Optional[str] = None
mapped_org_id: Optional[str] = None

View file

@ -75,7 +75,7 @@ def test_jwt_principal_roundtrip():
assert isinstance(ctx.principal, JWTPrincipal)
assert ctx.principal.sub == "jwt-user"
assert ctx.principal.iss == "idp"
assert ctx.principal.scopes == ["read", "write"]
assert ctx.principal.scopes == ("read", "write")
assert ctx.principal.mapped_user_id == "jwt-user"
assert ctx.principal.mapped_team_id == "jwt-team"
assert ctx.principal.mapped_org_id == "jwt-org"

View file

@ -127,7 +127,7 @@ def test_adapter_roundtrip_produces_jwt_principal():
assert isinstance(ctx.principal, JWTPrincipal)
assert ctx.principal.sub == "u-jwt"
assert ctx.principal.iss == "https://idp"
assert ctx.principal.scopes == ["read", "write"]
assert ctx.principal.scopes == ("read", "write")
assert ctx.principal.mapped_user_id == "u-jwt"
assert ctx.principal.mapped_team_id == "t-jwt"
assert ctx.principal.mapped_org_id == "org-jwt"

View file

@ -30,8 +30,34 @@ def test_principals_are_frozen():
def test_principals_are_hashable():
a = ApiKeyPrincipal(token_hash="x", user_id="u1")
b = ApiKeyPrincipal(token_hash="x", user_id="u1")
assert {a, b} == {a}
a_dup = ApiKeyPrincipal(token_hash="x", user_id="u1")
principals = [
a,
JWTPrincipal(
sub="s",
iss="idp",
aud=("aud-1", "aud-2"),
scopes=("read", "write"),
claims={"sub": "s", "scope": "read write"},
),
SSOPrincipal(sso_user_id="s"),
ServiceAccountPrincipal(name="n"),
AnonymousPrincipal(),
]
members = set(principals)
assert len(members) == len(principals)
for principal in principals:
assert principal in members
assert {a, a_dup} == {a}
def test_jwt_principal_hash_ignores_claims():
base = JWTPrincipal(sub="s", scopes=("read",))
with_claims = JWTPrincipal(sub="s", scopes=("read",), claims={"jti": "abc"})
assert base == with_claims
assert hash(base) == hash(with_claims)
def test_kind_is_not_constructor_arg():