mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(mcp): make open_envelope total over hostile jwt claim types and cap candidate size
This commit is contained in:
parent
dd38e9f1a0
commit
d6503d1d87
2 changed files with 149 additions and 12 deletions
|
|
@ -21,8 +21,15 @@ Failures are values: :func:`open_envelope` returns one of the frozen
|
|||
``EnvelopeOpenError`` variants (discriminated on ``tag``) for invalid, expired,
|
||||
tampered, or undecryptable input, and :func:`mint_envelope` returns
|
||||
``EnvelopeTooLarge`` for oversized grants. Error values carry tags and sizes only,
|
||||
never token material. Raising is reserved for programmer errors, which the pydantic
|
||||
models reject at construction (e.g. a non-positive ``expires_in``).
|
||||
never token material.
|
||||
|
||||
The pydantic input models reject programmer errors at construction (e.g. a
|
||||
non-positive ``expires_in`` or an empty required field). :func:`open_envelope` is
|
||||
additionally total over hostile, attacker-controlled input: it never raises, only
|
||||
returns an ``EnvelopeOpenError``. :func:`mint_envelope` operates on a
|
||||
gateway-supplied grant (an upstream IdP's UTF-8 JSON token response), so it does not
|
||||
defend against non-UTF-8 field content that cannot survive JSON parsing; its only
|
||||
value-typed failure is ``EnvelopeTooLarge``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -85,10 +92,15 @@ class UpstreamTokenGrant(BaseModel):
|
|||
|
||||
|
||||
class EnvelopeKeys(BaseModel):
|
||||
"""Injected key material: the HS256 signing key and the symmetric encryption key."""
|
||||
"""Injected key material: the HS256 signing key and the symmetric encryption key.
|
||||
|
||||
``signing_key`` must be at least 32 bytes: HS256's HMAC-SHA256 has a 256-bit
|
||||
security level, RFC 7518 requires a key of at least that size, and a shorter key
|
||||
makes PyJWT emit ``InsecureKeyLengthWarning``.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
signing_key: SecretStr = Field(min_length=1)
|
||||
signing_key: SecretStr = Field(min_length=32)
|
||||
encryption_key: SecretStr = Field(min_length=1)
|
||||
|
||||
|
||||
|
|
@ -160,12 +172,22 @@ EnvelopeOpenError: TypeAlias = NotAnEnvelope | BadSignature | Expired | Malforme
|
|||
|
||||
|
||||
class _EnvelopeClaims(BaseModel):
|
||||
"""Decoded-claims boundary. ``user_id``/``server_id`` mirror the ``min_length``
|
||||
constraints of :class:`EnvelopeIdentity` so any claim set that validates here also
|
||||
constructs an identity, keeping :func:`open_envelope` raise-free: a correctly signed
|
||||
JWT with an empty identity claim fails here and maps to ``MalformedPayload``."""
|
||||
"""Decoded-claims boundary that pins the exact shape :func:`mint_envelope` emits.
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
``user_id``/``server_id`` mirror the ``min_length`` constraints of
|
||||
:class:`EnvelopeIdentity` so any claim set that validates here also constructs an
|
||||
identity, keeping :func:`open_envelope` raise-free: a correctly signed JWT with an
|
||||
empty identity claim fails here and maps to ``MalformedPayload``.
|
||||
|
||||
``strict`` rejects coerced types (``exp: "123"``, ``exp: 123.0``) rather than opening
|
||||
on them, and ``extra="forbid"`` rejects any claim the gateway never mints (a hostile
|
||||
``nbf``/``aud``/... rides along on a re-signed token). Since PyJWT's own ``iat``/
|
||||
``nbf``/``exp`` validators are disabled at decode (they raise on hostile claim types
|
||||
and, for ``iat``/``nbf``, compare against the wall clock rather than the injected
|
||||
``now``), this model is the sole, total type gate for every registered claim.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(frozen=True, strict=True, extra="forbid")
|
||||
iss: str
|
||||
iat: int
|
||||
exp: int
|
||||
|
|
@ -228,10 +250,15 @@ def open_envelope(
|
|||
"""Validate ``candidate`` and recover the identity and inner grant.
|
||||
|
||||
Never raises for bad input: every invalid, expired, tampered, or undecryptable
|
||||
candidate maps to a distinct ``EnvelopeOpenError`` variant.
|
||||
candidate maps to a distinct ``EnvelopeOpenError`` variant. The recovered
|
||||
``grant.expires_in`` is the value the upstream reported at mint time and is not
|
||||
re-derived, so it is stale by up to the envelope's lifetime; callers that need a
|
||||
live remaining lifetime should use ``now`` against the upstream, not this field.
|
||||
"""
|
||||
if not is_envelope(candidate):
|
||||
return NotAnEnvelope()
|
||||
if len(candidate) > MAX_ENVELOPE_BYTES:
|
||||
return MalformedPayload()
|
||||
claims = _decode_claims(candidate.removeprefix(ENVELOPE_PREFIX), keys.signing_key)
|
||||
if not isinstance(claims, _EnvelopeClaims):
|
||||
return claims
|
||||
|
|
@ -267,17 +294,34 @@ def _decode_claims(
|
|||
compact: str,
|
||||
signing_key: SecretStr,
|
||||
) -> _EnvelopeClaims | BadSignature | MalformedPayload:
|
||||
"""Verify the HS256 signature and shape of an attacker-controlled compact JWT.
|
||||
|
||||
``compact`` is fully hostile and bounded to ``MAX_ENVELOPE_BYTES`` by the caller.
|
||||
PyJWT's ``iat``/``nbf``/``exp`` validators are disabled: they raise on hostile claim
|
||||
types and, for ``iat``/``nbf``, compare against the wall clock rather than the
|
||||
injected ``now`` (``exp`` is checked by the caller against ``now``). Apart from a
|
||||
signature mismatch (``BadSignature``), every decode failure is ``MalformedPayload``:
|
||||
a non-UTF-8 candidate surfaces as ``UnicodeEncodeError`` (a ``ValueError``), a
|
||||
non-string registered claim such as ``iss`` as a ``TypeError`` from PyJWT's claim
|
||||
validators, and a wrong issuer or structurally invalid token as an
|
||||
``InvalidTokenError``. ``_EnvelopeClaims`` is the total type gate for the payload.
|
||||
"""
|
||||
try:
|
||||
payload = jwt.decode(
|
||||
compact,
|
||||
signing_key.get_secret_value(),
|
||||
algorithms=[_ENVELOPE_JWT_ALGORITHM],
|
||||
issuer=ENVELOPE_ISSUER,
|
||||
options={"verify_exp": False, "require": ["iss", "iat", "exp"]},
|
||||
options={
|
||||
"verify_exp": False,
|
||||
"verify_iat": False,
|
||||
"verify_nbf": False,
|
||||
"require": ["iss", "iat", "exp"],
|
||||
},
|
||||
)
|
||||
except jwt.InvalidSignatureError:
|
||||
return BadSignature()
|
||||
except jwt.InvalidTokenError:
|
||||
except (jwt.InvalidTokenError, ValueError, TypeError):
|
||||
return MalformedPayload()
|
||||
try:
|
||||
return _EnvelopeClaims.model_validate(payload)
|
||||
|
|
|
|||
|
|
@ -9,11 +9,14 @@ error value, model repr, or raised exception ever contains the inner access toke
|
|||
"""
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import jwt
|
||||
import pytest
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from pydantic import SecretStr, ValidationError
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
|
||||
|
|
@ -79,6 +82,22 @@ def _forge(claims: dict[str, object], signing_key: str = _SIGNING_KEY) -> str:
|
|||
return ENVELOPE_PREFIX + jwt.encode(claims, signing_key, algorithm="HS256")
|
||||
|
||||
|
||||
def _b64url(raw: bytes) -> str:
|
||||
return base64.urlsafe_b64encode(raw).rstrip(b"=").decode("ascii")
|
||||
|
||||
|
||||
def _hand_crafted_hs256(payload: dict[str, object], signing_key: str = _SIGNING_KEY) -> str:
|
||||
"""Assemble an HS256 envelope from raw bytes, bypassing PyJWT's encode-side claim
|
||||
guards (it refuses to build a token with a non-string ``iss``). This is the real
|
||||
attacker path: a client crafts the compact JWT directly, so any registered claim can
|
||||
carry a hostile JSON type."""
|
||||
header = _b64url(json.dumps({"alg": "HS256", "typ": "JWT"}).encode("utf-8"))
|
||||
body = _b64url(json.dumps(payload).encode("utf-8"))
|
||||
signing_input = f"{header}.{body}".encode("ascii")
|
||||
signature = _b64url(hmac.new(signing_key.encode("utf-8"), signing_input, hashlib.sha256).digest())
|
||||
return ENVELOPE_PREFIX + f"{header}.{body}.{signature}"
|
||||
|
||||
|
||||
def _tampered(sealed_token: str, segment: int, index: int) -> str:
|
||||
parts = sealed_token.removeprefix(ENVELOPE_PREFIX).split(".")
|
||||
original = parts[segment][index]
|
||||
|
|
@ -221,6 +240,80 @@ def test_signed_empty_identity_claim_is_malformed_payload_not_a_raise(identity_c
|
|||
assert intact.identity == _IDENTITY
|
||||
|
||||
|
||||
def test_lone_surrogate_candidate_is_malformed_payload_not_a_raise():
|
||||
surrogate_candidate = ENVELOPE_PREFIX + "\ud800abc.def.ghi"
|
||||
result = open_envelope(surrogate_candidate, _KEYS, _NOW)
|
||||
assert isinstance(result, MalformedPayload)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"override",
|
||||
[
|
||||
{"iat": [1]},
|
||||
{"iat": {}},
|
||||
{"iat": float("inf")},
|
||||
{"nbf": None},
|
||||
{"nbf": [1]},
|
||||
],
|
||||
)
|
||||
def test_hostile_iat_nbf_types_are_malformed_payload_not_a_raise(override):
|
||||
claims = _unverified_claims(_sealed_token(_full_grant()))
|
||||
forged = _forge({**claims, **override})
|
||||
result = open_envelope(forged, _KEYS, _NOW)
|
||||
assert isinstance(result, MalformedPayload)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("hostile_iss", [["litellm-mcp-bridge"], 5, {"iss": "x"}])
|
||||
def test_non_string_issuer_claim_is_malformed_payload_not_a_raise(hostile_iss):
|
||||
claims = _unverified_claims(_sealed_token(_full_grant()))
|
||||
forged = _hand_crafted_hs256({**claims, "iss": hostile_iss})
|
||||
result = open_envelope(forged, _KEYS, _NOW)
|
||||
assert isinstance(result, MalformedPayload)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("hostile_exp", ["600", 600.5, [600]])
|
||||
def test_non_int_exp_claim_is_malformed_payload_not_a_raise(hostile_exp):
|
||||
claims = _unverified_claims(_sealed_token(_full_grant()))
|
||||
forged = _hand_crafted_hs256({**claims, "exp": hostile_exp})
|
||||
result = open_envelope(forged, _KEYS, _NOW)
|
||||
assert isinstance(result, MalformedPayload)
|
||||
|
||||
|
||||
def test_unexpected_extra_claim_is_malformed_payload():
|
||||
claims = _unverified_claims(_sealed_token(_full_grant()))
|
||||
forged = _forge({**claims, "role": "admin"})
|
||||
assert isinstance(open_envelope(forged, _KEYS, _NOW), MalformedPayload)
|
||||
|
||||
|
||||
def test_future_iat_opens_against_injected_now_not_wall_clock():
|
||||
future = _NOW + timedelta(seconds=100_000)
|
||||
sealed = mint_envelope(_IDENTITY, _full_grant(), _KEYS, future)
|
||||
assert isinstance(sealed, SealedEnvelope)
|
||||
opened = open_envelope(sealed.token.get_secret_value(), _KEYS, future)
|
||||
assert isinstance(opened, OpenedEnvelope)
|
||||
assert opened.identity == _IDENTITY
|
||||
assert opened.grant == _full_grant()
|
||||
|
||||
|
||||
def test_rs256_signed_token_is_rejected_against_the_hs256_pin():
|
||||
claims = _unverified_claims(_sealed_token(_full_grant()))
|
||||
private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
rs256_token = ENVELOPE_PREFIX + jwt.encode(claims, private_key, algorithm="RS256")
|
||||
result = open_envelope(rs256_token, _KEYS, _NOW)
|
||||
assert isinstance(result, MalformedPayload)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("short_key", ["", "too-short", "x" * 31])
|
||||
def test_signing_key_below_hs256_minimum_is_rejected_at_construction(short_key):
|
||||
with pytest.raises(ValidationError):
|
||||
EnvelopeKeys(signing_key=SecretStr(short_key), encryption_key=SecretStr(_ENCRYPTION_KEY))
|
||||
|
||||
|
||||
def test_signing_key_at_hs256_minimum_is_accepted():
|
||||
keys = EnvelopeKeys(signing_key=SecretStr("y" * 32), encryption_key=SecretStr(_ENCRYPTION_KEY))
|
||||
assert keys.signing_key.get_secret_value() == "y" * 32
|
||||
|
||||
|
||||
def test_correctly_signed_garbage_grant_blob_is_decrypt_failed():
|
||||
claims = _unverified_claims(_sealed_token(_full_grant()))
|
||||
forged = _forge({**claims, "grant": "not-a-ciphertext"})
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue