From 75e698feefe16dee5f152036a84be7f41da82e39 Mon Sep 17 00:00:00 2001 From: iabhi4 Date: Sun, 31 Aug 2025 15:58:26 -0700 Subject: [PATCH] bug(auth): support for ES256/ES384/ES512 and EdDSA JWT verification --- litellm/proxy/auth/handle_jwt.py | 23 +++-- tests/proxy_unit_tests/test_jwt.py | 139 +++++++++++++++++++++++++++++ 2 files changed, 153 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index f20c3debc2..6a7aa794e5 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -484,7 +484,7 @@ class JWTHandler: # Supported algos: https://pyjwt.readthedocs.io/en/stable/algorithms.html # "Warning: Make sure not to mix symmetric and asymmetric algorithms that interpret # the key in different ways (e.g. HS* and RS*)." - algorithms = ["RS256", "RS384", "RS512", "PS256", "PS384", "PS512"] + algorithms = ["RS256", "RS384", "RS512", "PS256", "PS384", "PS512", "ES256", "ES384", "ES512", "EdDSA"] audience = os.getenv("JWT_AUDIENCE") decode_options = None @@ -492,7 +492,7 @@ class JWTHandler: decode_options = {"verify_aud": False} import jwt - from jwt.algorithms import RSAAlgorithm + from jwt.api_jwk import PyJWK header = jwt.get_unverified_header(token) @@ -512,14 +512,21 @@ class JWTHandler: jwk["n"] = public_key["n"] if "e" in public_key: jwk["e"] = public_key["e"] + if "x" in public_key: + jwk["x"] = public_key["x"] + if "y" in public_key: + jwk["y"] = public_key["y"] + if "crv" in public_key: + jwk["crv"] = public_key["crv"] - public_key_rsa = RSAAlgorithm.from_jwk(json.dumps(jwk)) + # parse RSA/EC/OKP keys + public_key_obj = PyJWK.from_dict(jwk).key try: # decode the token using the public key payload = jwt.decode( token, - public_key_rsa, # type: ignore + public_key_obj, # type: ignore algorithms=algorithms, options=decode_options, audience=audience, @@ -534,9 +541,7 @@ class JWTHandler: raise Exception(f"Validation fails: {str(e)}") elif public_key is not None and isinstance(public_key, str): try: - cert = x509.load_pem_x509_certificate( - public_key.encode(), default_backend() - ) + cert = x509.load_pem_x509_certificate(public_key.encode(), default_backend()) # Extract public key key = cert.public_key().public_bytes( @@ -561,7 +566,7 @@ class JWTHandler: raise Exception(f"Validation fails: {str(e)}") raise Exception("Invalid JWT Submitted") - + async def close(self): await self.http_handler.close() @@ -1210,4 +1215,4 @@ class JWTAuthManager: end_user_object=end_user_object, token=api_key, team_membership=team_membership_object, - ) + ) \ No newline at end of file diff --git a/tests/proxy_unit_tests/test_jwt.py b/tests/proxy_unit_tests/test_jwt.py index d040376942..57514817e8 100644 --- a/tests/proxy_unit_tests/test_jwt.py +++ b/tests/proxy_unit_tests/test_jwt.py @@ -1375,3 +1375,142 @@ async def test_custom_validate_called(): pass # Assert custom_validate was called with the jwt token mock_custom_validate.assert_called_once_with({"sub": "test_user"}) + + +@pytest.mark.asyncio +async def test_auth_jwt_es256_jwk_path(monkeypatch): + import time, base64, jwt + from cryptography.hazmat.primitives.asymmetric import ec + from cryptography.hazmat.primitives import serialization + + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + + def b64url_uint(n: int, size: int) -> str: + return base64.urlsafe_b64encode(n.to_bytes(size, "big")).rstrip(b"=").decode() + + ec_key = ec.generate_private_key(ec.SECP256R1()) + ec_priv_pem = ec_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + + pub = ec_key.public_key().public_numbers() + ec_jwk = { + "kty": "EC", + "crv": "P-256", + "x": b64url_uint(pub.x, 32), + "y": b64url_uint(pub.y, 32), + "kid": "ec1", + "alg": "ES256", + "use": "sig", + } + + now = int(time.time()) + token = jwt.encode( + {"sub": "alice", "aud": "litellm-proxy", "iss": "http://example", "iat": now, "exp": now + 300}, + ec_priv_pem, + algorithm="ES256", + headers={"kid": "ec1"}, + ) + + h = JWTHandler() + with patch.object(h, "get_public_key", new=AsyncMock(return_value=ec_jwk)): + claims = await h.auth_jwt(token) + assert claims["sub"] == "alice" + + +@pytest.mark.asyncio +async def test_auth_jwt_rs256_regression(monkeypatch): + """ + Regression: RSA path must still work (kty RSA, n/e) after EC support. + """ + import time, base64, jwt + from cryptography.hazmat.primitives.asymmetric import rsa + from cryptography.hazmat.primitives import serialization + + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + + rsa_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + rsa_priv_pem = rsa_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + pub = rsa_key.public_key().public_numbers() + + def b64url(b: bytes) -> str: + return base64.urlsafe_b64encode(b).rstrip(b"=").decode() + + n = pub.n.to_bytes((pub.n.bit_length() + 7) // 8, "big") + e = pub.e.to_bytes((pub.e.bit_length() + 7) // 8, "big") + rsa_jwk = { + "kty": "RSA", + "n": b64url(n), + "e": b64url(e), + "kid": "rsa1", + "alg": "RS256", + "use": "sig", + } + + now = int(time.time()) + token = jwt.encode( + {"sub": "bob", "aud": "litellm-proxy", "iss": "http://example", "iat": now, "exp": now + 300}, + rsa_priv_pem, + algorithm="RS256", + headers={"kid": "rsa1"}, + ) + + h = JWTHandler() + with patch.object(h, "get_public_key", new=AsyncMock(return_value=rsa_jwk)): + claims = await h.auth_jwt(token) + assert claims["sub"] == "bob" + + +@pytest.mark.asyncio +async def test_auth_jwt_mismatched_key_fails(monkeypatch): + """ + Negative: ES256 token must fail if JWKS returns an RSA key (mismatch). + """ + import time, base64, jwt + from cryptography.hazmat.primitives.asymmetric import ec, rsa + from cryptography.hazmat.primitives import serialization + + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + + # ES256 token + ec_key = ec.generate_private_key(ec.SECP256R1()) + ec_priv_pem = ec_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + now = int(time.time()) + token = jwt.encode( + {"sub": "mallory", "aud": "litellm-proxy", "iss": "http://example", "iat": now, "exp": now + 300}, + ec_priv_pem, + algorithm="ES256", + headers={"kid": "ec1"}, + ) + + # RSA JWK (wrong key) + rsa_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + pub = rsa_key.public_key().public_numbers() + + def b64url(b: bytes) -> str: + return base64.urlsafe_b64encode(b).rstrip(b"=").decode() + + rsa_jwk = { + "kty": "RSA", + "n": b64url(pub.n.to_bytes((pub.n.bit_length() + 7) // 8, "big")), + "e": b64url(pub.e.to_bytes((pub.e.bit_length() + 7) // 8, "big")), + "kid": "rsa1", + "alg": "RS256", + "use": "sig", + } + + h = JWTHandler() + with patch.object(h, "get_public_key", new=AsyncMock(return_value=rsa_jwk)): + with pytest.raises(Exception) as exc: + await h.auth_jwt(token) + assert "Validation fails" in str(exc.value) \ No newline at end of file