diff --git a/litellm/identity/adapter.py b/litellm/identity/adapter.py index efd0f3c867e..478c7f9a77a 100644 --- a/litellm/identity/adapter.py +++ b/litellm/identity/adapter.py @@ -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): diff --git a/litellm/identity/principal.py b/litellm/identity/principal.py index 58c264ebe5b..28377f5f281 100644 --- a/litellm/identity/principal.py +++ b/litellm/identity/principal.py @@ -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 diff --git a/tests/test_litellm/identity/test_adapter.py b/tests/test_litellm/identity/test_adapter.py index fd34e1b2af9..9f8d2d8ee6a 100644 --- a/tests/test_litellm/identity/test_adapter.py +++ b/tests/test_litellm/identity/test_adapter.py @@ -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" diff --git a/tests/test_litellm/identity/test_jwt_builder.py b/tests/test_litellm/identity/test_jwt_builder.py index 9da31b04f47..f3b5194c6b4 100644 --- a/tests/test_litellm/identity/test_jwt_builder.py +++ b/tests/test_litellm/identity/test_jwt_builder.py @@ -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" diff --git a/tests/test_litellm/identity/test_principal.py b/tests/test_litellm/identity/test_principal.py index d3ef95888ec..f2e663ec895 100644 --- a/tests/test_litellm/identity/test_principal.py +++ b/tests/test_litellm/identity/test_principal.py @@ -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():