mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(identity): make JWTPrincipal hashable via tuple scopes and compare-excluded claims
This commit is contained in:
parent
bf33cc24e3
commit
30cb095859
5 changed files with 42 additions and 13 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue