From b257feb0b9fe9f8a40ec88e3120b7e0e91e913d0 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Mon, 8 Jun 2026 15:16:03 -0700 Subject: [PATCH] refactor(identity): use PyJWT directly in extract_jwt_principal --- litellm/identity/extractors/jwt.py | 14 ++++++------- .../identity/extractors/test_jwt.py | 21 +++++++++++++++++++ 2 files changed, 28 insertions(+), 7 deletions(-) diff --git a/litellm/identity/extractors/jwt.py b/litellm/identity/extractors/jwt.py index 0c882b070e4..441a170864c 100644 --- a/litellm/identity/extractors/jwt.py +++ b/litellm/identity/extractors/jwt.py @@ -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( diff --git a/tests/test_litellm/identity/extractors/test_jwt.py b/tests/test_litellm/identity/extractors/test_jwt.py index 43cabc607a7..a8970ef69cc 100644 --- a/tests/test_litellm/identity/extractors/test_jwt.py +++ b/tests/test_litellm/identity/extractors/test_jwt.py @@ -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")