mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
refactor(identity): extract parse_jwt_scopes helper
This commit is contained in:
parent
287468ba0c
commit
d02d161477
5 changed files with 35 additions and 21 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue