mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): pin auth_v2 JWT verification to asymmetric algorithms
The JWT verifier used authlib's default decoder, which permits HMAC alongside RSA/EC. With a JWKS of public keys that is the RS256->HS256 confusion setup: an attacker signs HS256 using the public key as the HMAC secret. authlib happens to reject this today via key-type checks, but relying on that is fragile and breaks the moment a symmetric key enters the key set. Pinning the accepted algorithms to an asymmetric allowlist closes the class outright. Adds a regression test that forges an HS256 token from the JWKS public key and asserts it is rejected.
This commit is contained in:
parent
d2ab6197eb
commit
25823cba7f
2 changed files with 56 additions and 6 deletions
|
|
@ -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()))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue