diff --git a/litellm/proxy/auth/v2/jwt_verifier.py b/litellm/proxy/auth/v2/jwt_verifier.py index 32f3f0fc840..9b1e7e12f0a 100644 --- a/litellm/proxy/auth/v2/jwt_verifier.py +++ b/litellm/proxy/auth/v2/jwt_verifier.py @@ -1,7 +1,7 @@ import time -from typing import Any, Dict, Optional +from typing import Any, Dict, List, Optional -from authlib.jose import JsonWebKey, jwt +from authlib.jose import JsonWebKey, JsonWebToken from authlib.jose.errors import JoseError @@ -9,6 +9,25 @@ class JWTVerificationError(Exception): """Raised when a token fails signature or standard-claim validation.""" +# Pin verification to asymmetric algorithms. An OIDC JWKS holds public keys, so +# permitting HMAC (HS*) here is what enables the RS256->HS256 confusion attack: +# an attacker signs HS256 using the public key as the secret. Refusing HS* (and +# the unsigned ``none`` alg, which authlib already rejects) closes that class +# regardless of the library's internal key-type checks. +_DEFAULT_ALGORITHMS: List[str] = [ + "RS256", + "RS384", + "RS512", + "ES256", + "ES384", + "ES512", + "PS256", + "PS384", + "PS512", + "EdDSA", +] + + def build_claims_options( issuer: Optional[str], audience: Optional[str] ) -> Dict[str, Any]: @@ -25,16 +44,20 @@ def verify( key_set: Any, issuer: Optional[str] = None, audience: Optional[str] = None, + algorithms: Optional[List[str]] = None, ) -> Dict[str, Any]: """Verify ``token`` against ``key_set`` and validate exp/iss/aud. authlib owns the crypto: signature verification, key selection by ``kid``, - and standard-claim checks. ``key_set`` is an imported JWKS (injected so this - is testable without network). Raises :class:`JWTVerificationError` on any - failure so callers never branch on authlib's internal exception types. + and standard-claim checks. The accepted ``algorithms`` are pinned (asymmetric + only by default) so the token header cannot downgrade verification to HMAC. + ``key_set`` is an imported JWKS (injected so this is testable without + network). Raises :class:`JWTVerificationError` on any failure so callers + never branch on authlib's internal exception types. """ + decoder = JsonWebToken(algorithms or _DEFAULT_ALGORITHMS) try: - claims = jwt.decode( + claims = decoder.decode( token, key_set, claims_options=build_claims_options(issuer, audience) ) claims.validate(now=int(time.time())) diff --git a/tests/test_litellm/proxy/auth/v2/test_jwt_verifier.py b/tests/test_litellm/proxy/auth/v2/test_jwt_verifier.py index a083b9875d4..3a02ff8bb72 100644 --- a/tests/test_litellm/proxy/auth/v2/test_jwt_verifier.py +++ b/tests/test_litellm/proxy/auth/v2/test_jwt_verifier.py @@ -82,3 +82,30 @@ def test_garbage_is_rejected(signing): _, _, key_set = signing with pytest.raises(JWTVerificationError): verify("not.a.jwt", key_set, ISSUER, AUDIENCE) + + +def test_rs256_to_hs256_algorithm_confusion_is_rejected(signing): + # The JWKS holds an RSA public key. An attacker forges an HS256 token using + # that public key as the HMAC secret. Pinning to asymmetric algorithms must + # refuse it; accepting it would be a full authentication bypass. + import base64 + import hashlib + import hmac + + key, kid, key_set = signing + public_pem = key.as_pem(is_private=False) + + def b64(raw: bytes) -> bytes: + return base64.urlsafe_b64encode(raw).rstrip(b"=") + + header = b64(b'{"alg":"HS256","kid":"%s"}' % kid.encode()) + payload = b64( + b'{"sub":"attacker","iss":"%s","aud":"%s","exp":9999999999}' + % (ISSUER.encode(), AUDIENCE.encode()) + ) + signing_input = header + b"." + payload + signature = b64(hmac.new(public_pem, signing_input, hashlib.sha256).digest()) + forged = (signing_input + b"." + signature).decode("utf-8") + + with pytest.raises(JWTVerificationError): + verify(forged, key_set, ISSUER, AUDIENCE)