refactor(identity): use PyJWT directly in extract_jwt_principal

This commit is contained in:
Yassin Kortam 2026-06-08 15:16:03 -07:00
parent 33bf08a6e1
commit b257feb0b9
2 changed files with 28 additions and 7 deletions

View file

@ -1,13 +1,14 @@
"""JWT principal extraction.
Uses ``JWTHandler.is_jwt`` for shape detection and
``JWTHandler.get_unverified_claims`` for claim peek. Signature
Decodes claims without verifying the signature (PyJWT direct). Signature
verification and DB-backed claim mapping live in ``JWTHandler.auth_builder``;
that path runs from the resolver when DB access is available.
"""
from typing import Optional
import jwt
from litellm.identity.jwt import parse_jwt_scopes
from litellm.identity.principal import JWTPrincipal
@ -22,13 +23,12 @@ def extract_jwt_principal(token: Optional[str]) -> Optional[JWTPrincipal]:
if not token:
return None
from litellm.proxy.auth.handle_jwt import JWTHandler
if not JWTHandler.is_jwt(token=token):
try:
jwt.get_unverified_header(token)
claims = jwt.decode(token, options={"verify_signature": False})
except jwt.PyJWTError:
return None
claims = JWTHandler.get_unverified_claims(token=token) or {}
aud = claims.get("aud")
return JWTPrincipal(

View file

@ -2,6 +2,7 @@ import base64
import json
import os
import sys
from unittest.mock import patch
sys.path.insert(0, os.path.abspath("../../.."))
@ -63,3 +64,23 @@ def test_raw_claims_preserved():
p = extract_jwt_principal(token)
assert p is not None
assert p.claims["custom"] == {"groups": ["g1"]}
def test_extract_jwt_principal_uses_pyjwt_not_jwt_handler():
token = _build_unverified_jwt({"sub": "u-pyjwt", "scope": "read write"})
with (
patch(
"litellm.proxy.auth.handle_jwt.JWTHandler.is_jwt",
side_effect=AssertionError("extractor must not call JWTHandler"),
),
patch(
"litellm.proxy.auth.handle_jwt.JWTHandler.get_unverified_claims",
side_effect=AssertionError("extractor must not call JWTHandler"),
),
):
p = extract_jwt_principal(token)
assert p is not None
assert p.sub == "u-pyjwt"
assert p.scopes == ("read", "write")