fix(mcp): make open_envelope total over hostile jwt claim types and cap candidate size

This commit is contained in:
Tin Chi Lo 2026-07-10 01:48:55 -07:00
parent dd38e9f1a0
commit d6503d1d87
2 changed files with 149 additions and 12 deletions

View file

@ -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)

View file

@ -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"})