mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
refactor(identity): use PyJWT directly in extract_jwt_principal
This commit is contained in:
parent
33bf08a6e1
commit
b257feb0b9
2 changed files with 28 additions and 7 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue