refactor(identity): extract parse_jwt_scopes helper

This commit is contained in:
Yassin Kortam 2026-06-08 14:39:15 -07:00
parent 287468ba0c
commit d02d161477
5 changed files with 35 additions and 21 deletions

View file

@ -15,6 +15,7 @@ from typing import TYPE_CHECKING
from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME
from litellm.identity.context import AuditInfo, ClientInfo, IdentityContext, RequestIds
from litellm.identity.jwt import parse_jwt_scopes
from litellm.identity.principal import (
AnonymousPrincipal,
ApiKeyPrincipal,
@ -38,19 +39,12 @@ def _principal_from_uak(uak: "UserAPIKeyAuth") -> Principal:
if kind == "jwt":
claims = uak.jwt_claims
scope_claim = claims.get("scope") or claims.get("scp") or ""
if isinstance(scope_claim, list):
scopes = tuple(str(s) for s in scope_claim if s)
elif isinstance(scope_claim, str):
scopes = tuple(s for s in scope_claim.split(" ") if s)
else:
scopes = ()
aud = claims.get("aud")
return JWTPrincipal(
sub=claims.get("sub"),
iss=claims.get("iss"),
aud=tuple(aud) if isinstance(aud, list) else aud,
scopes=scopes,
scopes=parse_jwt_scopes(claims),
claims=dict(claims),
mapped_user_id=uak.user_id,
mapped_team_id=uak.team_id,

View file

@ -8,6 +8,7 @@ that path runs from the resolver when DB access is available.
from typing import Optional
from litellm.identity.jwt import parse_jwt_scopes
from litellm.identity.principal import JWTPrincipal
@ -29,18 +30,11 @@ def extract_jwt_principal(token: Optional[str]) -> Optional[JWTPrincipal]:
claims = JWTHandler.get_unverified_claims(token=token) or {}
aud = claims.get("aud")
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]
elif isinstance(scope_claim, str):
scopes = [s for s in scope_claim.split(" ") if s]
else:
scopes = []
return JWTPrincipal(
sub=claims.get("sub"),
iss=claims.get("iss"),
aud=aud,
scopes=scopes,
aud=tuple(aud) if isinstance(aud, list) else aud,
scopes=parse_jwt_scopes(claims),
claims=claims,
)

View file

@ -19,12 +19,26 @@ upstream. Here we just build the carrier.
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Optional
from typing import TYPE_CHECKING, Any, Tuple
if TYPE_CHECKING:
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
def parse_jwt_scopes(claims: dict) -> Tuple[str, ...]:
"""Normalize the ``scope``/``scp`` claim into a tuple of scope strings.
Accepts the space-delimited string form (``"read write"``) and the JSON
array form (``["read", "write"]``); anything else yields ``()``.
"""
scope_claim = claims.get("scope") or claims.get("scp") or ""
if isinstance(scope_claim, list):
return tuple(str(s) for s in scope_claim if s)
if isinstance(scope_claim, str):
return tuple(s for s in scope_claim.split(" ") if s)
return ()
def build_user_api_key_auth_from_jwt_result(
*,
result: dict,

View file

@ -40,21 +40,21 @@ def test_string_scope_is_split():
token = _build_unverified_jwt({"sub": "u", "scope": "read write admin"})
p = extract_jwt_principal(token)
assert p is not None
assert p.scopes == ["read", "write", "admin"]
assert p.scopes == ("read", "write", "admin")
def test_list_scope_is_preserved():
token = _build_unverified_jwt({"sub": "u", "scope": ["a", "b"]})
p = extract_jwt_principal(token)
assert p is not None
assert p.scopes == ["a", "b"]
assert p.scopes == ("a", "b")
def test_scp_claim_supported():
token = _build_unverified_jwt({"sub": "u", "scp": "read"})
p = extract_jwt_principal(token)
assert p is not None
assert p.scopes == ["read"]
assert p.scopes == ("read",)
def test_raw_claims_preserved():

View file

@ -5,6 +5,7 @@ from types import SimpleNamespace
sys.path.insert(0, os.path.abspath("../.."))
from litellm.identity import build_user_api_key_auth_from_jwt_result
from litellm.identity.jwt import parse_jwt_scopes
from litellm.identity.principal import JWTPrincipal
from litellm.proxy._types import (
LiteLLM_TeamTableCachedObj,
@ -64,6 +65,17 @@ def _auth_builder_result(
}
def test_parse_jwt_scopes_handles_str_list_garbage_empty():
assert parse_jwt_scopes({"scope": "read write"}) == ("read", "write")
assert parse_jwt_scopes({"scope": ["read", "", "write"]}) == ("read", "write")
assert parse_jwt_scopes({"scp": "admin"}) == ("admin",)
assert parse_jwt_scopes({"scope": 123}) == ()
assert parse_jwt_scopes({"scope": {"nested": "obj"}}) == ()
assert parse_jwt_scopes({}) == ()
assert parse_jwt_scopes({"scope": ""}) == ()
assert parse_jwt_scopes({"scope": " "}) == ()
def test_proxy_admin_path_marks_role_and_drops_membership_fields():
result = _auth_builder_result(is_proxy_admin=True, team=_team(), user=_user())
uak = build_user_api_key_auth_from_jwt_result(