This commit is contained in:
devin-ai-integration[bot] 2026-10-06 00:48:27 +00:00 • committed by GitHub
commit eecac2fd3f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
19 changed files with 2165 additions and 611 deletions

View file

@ -211,6 +211,7 @@ from litellm.types.mcp_server.mcp_server_manager import (
MCPServer,
parse_pinned_tools,
)
from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, ApprovedJwtAlgorithm
from litellm.types.utils import CallTypes
if TYPE_CHECKING:
@ -419,7 +420,7 @@ class MCPServerConfig(TypedDict, total=False):
id_jag_resource: str
client_private_key: str
client_private_key_id: str
client_assertion_signing_alg: str
client_assertion_signing_alg: ApprovedJwtAlgorithm
timeout: float
max_concurrent_requests: int
@ -1567,6 +1568,18 @@ def _warn_on_shared_identifier_prefixes(servers: Iterable[MCPServer]) -> None:
)
def _stored_client_assertion_signing_alg(value: object, server_name: str) -> ApprovedJwtAlgorithm:
if value in APPROVED_JWT_ALGORITHMS:
return cast(ApprovedJwtAlgorithm, value) # cast-ok: value was checked against APPROVED_JWT_ALGORITHMS
if value is not None:
verbose_logger.warning(
"MCP server %s: client_assertion_signing_alg %r is not an approved algorithm, using RS256",
server_name,
value,
)
return "RS256"
def _warn_legacy_delegate_auth_if_applicable(server: MCPServer, *, source: str) -> None:
"""Direct legacy delegated OAuth configurations to the admitted replacement."""
if server.auth_type != MCPAuth.oauth2:
@ -2707,7 +2720,10 @@ class MCPServerManager:
id_jag_resource=server_config.get("id_jag_resource", None),
client_private_key=server_config.get("client_private_key", None),
client_private_key_id=server_config.get("client_private_key_id", None),
client_assertion_signing_alg=server_config.get("client_assertion_signing_alg", "RS256"),
client_assertion_signing_alg=_stored_client_assertion_signing_alg(
server_config.get("client_assertion_signing_alg"),
str(server_config.get("alias") or server_config.get("server_name") or server_id),
),
token_exchange_profile=server_config.get("token_exchange_profile", "rfc8693"),
allow_sampling=bool(server_config.get("allow_sampling", False)),
allow_elicitation=bool(server_config.get("allow_elicitation", False)),
@ -3290,10 +3306,10 @@ class MCPServerManager:
credentials_are_encrypted,
),
client_private_key_id=(credentials_dict.get("client_private_key_id") if credentials_dict else None),
client_assertion_signing_alg=(
credentials_dict.get("client_assertion_signing_alg") if credentials_dict else None
)
or "RS256",
client_assertion_signing_alg=_stored_client_assertion_signing_alg(
credentials_dict.get("client_assertion_signing_alg") if credentials_dict else None,
name_for_prefix,
),
token_exchange_profile=mcp_server.token_exchange_profile
or (credentials_dict.get("token_exchange_profile") if credentials_dict else None)
or "rfc8693",

View file

@ -23,20 +23,12 @@ from typing_extensions import assert_never
from litellm._logging import verbose_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy.auth.jwt_algorithms import jwks_keys_for
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer
from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS
_ALLOWED_ID_TOKEN_ALGORITHMS: Final = (
"RS256",
"RS384",
"RS512",
"ES256",
"ES384",
"ES512",
"PS256",
"PS384",
"PS512",
)
_ALLOWED_ID_TOKEN_ALGORITHMS: Final = APPROVED_JWT_ALGORITHMS
_JWKS_CACHE_TTL_SECONDS: Final = 3600
_jwks_cache: Final = InMemoryCache(default_ttl=_JWKS_CACHE_TTL_SECONDS)
@ -126,7 +118,7 @@ async def _discover_jwks_url(issuer: str) -> str:
def _select_signing_key(id_token: str, keys: Sequence[Mapping[str, object]]) -> "jwt.PyJWK | _BindingRejection":
header: Final = jwt.get_unverified_header(id_token)
kid: Final = header.get("kid")
for key in keys:
for key in jwks_keys_for(keys, _ALLOWED_ID_TOKEN_ALGORITHMS):
if kid is None or key.get("kid") == kid:
return jwt.PyJWK(dict(key))
return _BindingRejection(

View file

@ -45,6 +45,7 @@ from litellm.types.mcp import (
DEFAULT_SUBJECT_TOKEN_TYPE,
normalize_upstream_header_name,
)
from litellm.types.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm
class AuthResolution(str, Enum):
@ -326,7 +327,7 @@ class PrivateKeyJwtAuth(BaseModel):
source: Literal["private_key_jwt"] = "private_key_jwt"
private_key: SecretStr
key_id: str | None = None
signing_alg: str = "RS256"
signing_alg: ApprovedJwtAlgorithm = "RS256"
class ClientSecretAuth(BaseModel):

View file

@ -17,6 +17,7 @@ from pydantic import (
PositiveInt,
PrivateAttr,
TypeAdapter,
ValidationError,
field_validator,
model_validator,
)
@ -52,6 +53,10 @@ from litellm.types.mcp import (
)
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
from litellm.types.proxy.agent_identity import ManagedAgentContext
from litellm.types.proxy.auth.jwt_algorithms import (
APPROVED_JWT_ALGORITHMS,
ApprovedJwtAlgorithm,
)
from litellm.types.proxy.carried_budget_state import (
OrgBudgetSnapshot,
TeamBudgetSnapshot,
@ -1593,6 +1598,19 @@ def _reject_unsupported_per_server_oauth_discovery(values: object, require_auth_
raise _per_server_oauth_discovery_error()
_APPROVED_JWT_ALGORITHM_ADAPTER: Final = TypeAdapter(ApprovedJwtAlgorithm)
def _validated_client_assertion_signing_alg(credentials: MCPCredentials | None) -> MCPCredentials | None:
if credentials is None or credentials.get("client_assertion_signing_alg") is None:
return credentials
try:
_APPROVED_JWT_ALGORITHM_ADAPTER.validate_python(credentials["client_assertion_signing_alg"])
except ValidationError as exc:
raise ValueError(f"client_assertion_signing_alg must be one of {', '.join(APPROVED_JWT_ALGORITHMS)}") from exc
return credentials
def _validate_mcp_transport_fields(values: object) -> None:
if not isinstance(values, dict):
return
@ -1680,6 +1698,11 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
description="Server-managed: set by the endpoint; caller values are overridden.",
)
@field_validator("credentials")
@classmethod
def check_client_assertion_signing_alg(cls, credentials: MCPCredentials | None) -> MCPCredentials | None:
return _validated_client_assertion_signing_alg(credentials)
@model_validator(mode="after")
def validate_protocol_transport(self) -> "NewMCPServerRequest":
validate_mcp_protocol_transport(
@ -1773,6 +1796,11 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
timeout: float | None = None
max_concurrent_requests: int | None = None
@field_validator("credentials")
@classmethod
def check_client_assertion_signing_alg(cls, credentials: MCPCredentials | None) -> MCPCredentials | None:
return _validated_client_assertion_signing_alg(credentials)
@model_validator(mode="after")
def validate_protocol_transport(self) -> "UpdateMCPServerRequest":
if not {"transport", "mcp_info"}.issubset(self.model_fields_set):

View file

@ -16,6 +16,7 @@ import re
import time
from collections.abc import Awaitable, Callable, Collection, Mapping, Sequence
from dataclasses import dataclass
from functools import lru_cache
from typing import Any, Final, Literal, NoReturn, Protocol, TypeVar, cast
import httpx
@ -58,6 +59,7 @@ from litellm.proxy.agent_endpoints.identity import has_legacy_identity
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore, resolve_managed_agent
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
from litellm.proxy.auth.auth_checks import can_team_access_model
from litellm.proxy.auth.jwt_algorithms import allowed_jwt_algorithms, jwks_keys_for
from litellm.proxy.auth.model_access_denied import (
ModelAccessDeniedHTTPException,
model_access_denied_client_message,
@ -65,6 +67,7 @@ from litellm.proxy.auth.model_access_denied import (
from litellm.proxy.auth.resolvers.grants import GrantResolver, UserLookup, canonical_user_id
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.team_grants import team_grants, team_model_aliases
from litellm.proxy.common_utils.fips import is_fips_mode
from litellm.proxy.common_utils.user_api_key_cache import (
AUTH_OBJECTS_TARGET,
UserApiKeyCache,
@ -75,6 +78,7 @@ from litellm.repositories.user_repository import UserRepository
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import AgentIdentityFailure
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, LEGACY_JWT_ALGORITHMS
from .auth_checks import (
TeamNotFoundError,
@ -111,6 +115,14 @@ UNREACHABLE_CACHE_KEY_PREFIX: Final = "litellm_jwks_unreachable_"
_CachedValueT = TypeVar("_CachedValueT", bound=JWKKeyValue | str)
@lru_cache(maxsize=1)
def _log_eddsa_deprecation() -> None:
verbose_proxy_logger.warning(
"JWT Auth: accepted a token signed with EdDSA, which is deprecated and not FIPS 140-3 approved; "
"it is rejected when LITELLM_FIPS_MODE=true. Move the IdP signing key to RS256, PS256 or ES256"
)
class _JWTAuthSettings(Protocol):
"""The JWT auth settings block this handler reads back through ``getattr``, when one is configured."""
@ -222,16 +234,8 @@ class JWTHandler:
# "Warning: Make sure not to mix symmetric and asymmetric algorithms that interpret
# the key in different ways (e.g. HS* and RS*)."
SUPPORTED_JWT_ALGORITHMS = [
"RS256",
"RS384",
"RS512",
"PS256",
"PS384",
"PS512",
"ES256",
"ES384",
"ES512",
"EdDSA",
*APPROVED_JWT_ALGORITHMS,
*LEGACY_JWT_ALGORITHMS,
]
LITELLM_JWT_ISSUER_CLAIM = "_litellm_jwt_issuer"
LITELLM_USER_ID_CLAIM = "_litellm_user_id"
@ -252,13 +256,24 @@ class JWTHandler:
def __init__(
self,
fips_mode: Callable[[], bool] = is_fips_mode,
) -> None:
self._fips_mode: Final = fips_mode
self.http_handler = HTTPHandler()
self.leeway = 0
# Per-cache-key locks so a TTL lapse triggers one refresh instead of one per in-flight request.
self._refresh_locks: dict[str, asyncio.Lock] = {} # mutable-ok: lock registry, keyed by JWKS url
self.agent_lookup: AgentLookup = _NoRegisteredAgents()
def allowed_algorithms(self) -> tuple[str, ...]:
return allowed_jwt_algorithms(self._fips_mode())
def _warn_deprecated_signing_algorithm(self, token: str) -> None:
if self._fips_mode():
return
if jwt.get_unverified_header(token).get("alg") == "EdDSA":
_log_eddsa_deprecation()
def bind_agent_lookup(self, agent_lookup: AgentLookup) -> None:
self.agent_lookup = agent_lookup
@ -958,7 +973,13 @@ class JWTHandler:
log_context=f"kid={kid}",
)
public_key: Final = self.parse_keys(keys=keys, kid=kid)
allowed: Final = self.allowed_algorithms()
usable_keys: Final[JWKKeyValue] = (
list(jwks_keys_for(keys, allowed))
if isinstance(keys, list)
else next(iter(jwks_keys_for((keys,), allowed)), {})
)
public_key: Final = self.parse_keys(keys=usable_keys, kid=kid)
if public_key is not None:
return cast(dict, public_key)
@ -1239,32 +1260,27 @@ class JWTHandler:
)
)
if isinstance(public_key, dict):
public_key_obj: Final = PyJWK.from_dict(self._get_jwk_from_public_key(public_key=public_key)).key
return jwt.decode(
token,
public_key_obj,
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
options=decode_options,
audience=audience,
issuer=issuer,
leeway=self.leeway,
key_obj: Final = (
PyJWK.from_dict(self._get_jwk_from_public_key(public_key=public_key)).key
if isinstance(public_key, dict)
else x509.load_pem_x509_certificate(public_key.encode(), default_backend())
.public_key()
.public_bytes(
serialization.Encoding.PEM,
serialization.PublicFormat.SubjectPublicKeyInfo,
)
cert: Final = x509.load_pem_x509_certificate(public_key.encode(), default_backend())
key: Final = cert.public_key().public_bytes(
serialization.Encoding.PEM,
serialization.PublicFormat.SubjectPublicKeyInfo,
)
return jwt.decode(
payload: Final = jwt.decode(
token,
key,
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
key_obj,
algorithms=self.allowed_algorithms(),
audience=audience,
issuer=issuer,
options=decode_options,
leeway=self.leeway,
)
self._warn_deprecated_signing_algorithm(token)
return payload
async def _auth_jwt_with_issuer(self, token: str, issuer_config: JWTIssuerConfig, kid: str | None) -> dict:
try:

View file

@ -0,0 +1,35 @@
from collections.abc import Collection, Mapping, Sequence
from types import MappingProxyType
from typing import Final
from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, LEGACY_JWT_ALGORITHMS
_KEY_TYPE_ALGORITHMS: Final = MappingProxyType(
{
"RSA": frozenset(APPROVED_JWT_ALGORITHMS[:6]),
"EC": frozenset(APPROVED_JWT_ALGORITHMS[6:]),
"OKP": frozenset(LEGACY_JWT_ALGORITHMS),
}
)
def allowed_jwt_algorithms(fips_mode: bool) -> tuple[str, ...]:
return APPROVED_JWT_ALGORITHMS if fips_mode else (*APPROVED_JWT_ALGORITHMS, *LEGACY_JWT_ALGORITHMS)
def jwks_keys_for(
keys: Sequence[Mapping[str, object]], algorithms: Collection[str]
) -> tuple[Mapping[str, object], ...]:
"""Keep keys whose declared alg is allowed; keys without alg are kept only when their kty can sign with an
allowed algorithm."""
return tuple(key for key in keys if _key_allowed(key, frozenset(algorithms)))
def _key_allowed(key: Mapping[str, object], algorithms: frozenset[str]) -> bool:
alg: Final = key.get("alg")
if isinstance(alg, str):
return alg in algorithms
key_type: Final = key.get("kty")
if not isinstance(key_type, str):
return False
return bool(_KEY_TYPE_ALGORITHMS.get(key_type, frozenset()) & algorithms)

View file

@ -90,8 +90,10 @@ from litellm.integrations.custom_guardrail import (
log_guardrail_information,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.jwt_algorithms import jwks_keys_for
from litellm.types.guardrail_base_init import GuardrailBaseInitKwargs
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS
from litellm.types.utils import CallTypesLiteral
if TYPE_CHECKING:
@ -452,12 +454,16 @@ class MCPJWTSigner(CustomGuardrail):
unverified_header: Final = jwt.get_unverified_header(raw_token)
kid: Final = unverified_header.get("kid")
approved_keys: Final = jwks_keys_for(jwks_keys, APPROVED_JWT_ALGORITHMS)
if not approved_keys:
raise jwt.exceptions.PyJWKSetError(f"No JWKS key at {jwks_uri!r} uses an approved signing algorithm")
# Build a JWKS object and pick the matching key.
# PyJWT's PyJWKSet handles key-type parsing and kid matching correctly.
from jwt import PyJWKSet
try:
jwks_set: Final = PyJWKSet.from_dict({"keys": jwks_keys})
jwks_set: Final = PyJWKSet.from_dict({"keys": list(approved_keys)})
except Exception as exc:
raise jwt.exceptions.PyJWKSetError(f"Failed to parse JWKS from {jwks_uri!r}: {exc}") from exc

View file

@ -15,6 +15,7 @@ from litellm.types.mcp import (
normalize_upstream_header_name,
validate_mcp_protocol_transport,
)
from litellm.types.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm
# MCPInfo now allows arbitrary additional fields for custom metadata
@ -178,7 +179,7 @@ class MCPServer(BaseModel):
id_jag_resource: str | None = None
client_private_key: str | None = None
client_private_key_id: str | None = None
client_assertion_signing_alg: str = "RS256"
client_assertion_signing_alg: ApprovedJwtAlgorithm = "RS256"
# Wire dialect: "rfc8693" (standard token-exchange grant) or "entra_obo" (Microsoft Entra
# On-Behalf-Of, the RFC 7523 jwt-bearer grant + requested_token_use extension)
token_exchange_profile: str = "rfc8693"

View file

@ -0,0 +1,17 @@
from typing import Final, Literal
ApprovedJwtAlgorithm = Literal["RS256", "RS384", "RS512", "PS256", "PS384", "PS512", "ES256", "ES384", "ES512"]
APPROVED_JWT_ALGORITHMS: Final[tuple[ApprovedJwtAlgorithm, ...]] = (
"RS256",
"RS384",
"RS512",
"PS256",
"PS384",
"PS512",
"ES256",
"ES384",
"ES512",
)
LEGACY_JWT_ALGORITHMS: Final = ("EdDSA",)

View file

@ -0,0 +1,721 @@
import asyncio
import base64
import json
import os
import socket
import time
import uuid
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from typing import Final
import httpx
import jwt
import psutil
import yaml
from anthropic import Anthropic
from cryptography.hazmat.primitives.asymmetric import ec, ed25519, rsa
from jwt.algorithms import ECAlgorithm, OKPAlgorithm, RSAAlgorithm
from openai import AsyncOpenAI, OpenAI
from tests.integration._support.client import Gateway, eventually
from tests.integration._support.database import read_rows
from tests.integration._support.process import owned_proxy_process
from tests.integration._support.wire import Reply, Request, wire_server
KEY_ID: Final = "integration-jwt-allowlist-key"
APPROVED_WARNING_PARTS: Final = ("EdDSA", "deprecated", "LITELLM_FIPS_MODE")
def _jwks_reply(public_jwk: str) -> Reply:
return _jwks_reply_keys([{**json.loads(public_jwk), "kid": KEY_ID}])
def _jwks_reply_keys(keys: list[dict[str, object]]) -> Reply:
return Reply(body=json.dumps({"keys": keys}).encode())
def _eddsa_keypair() -> tuple[ed25519.Ed25519PrivateKey, dict[str, object]]:
private_key: Final = ed25519.Ed25519PrivateKey.generate()
return private_key, {**json.loads(OKPAlgorithm.to_jwk(private_key.public_key())), "kid": "ed"}
def _rsa_keypair() -> tuple[rsa.RSAPrivateKey, dict[str, object]]:
private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
return private_key, {**json.loads(RSAAlgorithm.to_jwk(private_key.public_key())), "kid": "rsa"}
def _ec_keypair() -> tuple[ec.EllipticCurvePrivateKey, dict[str, object]]:
private_key: Final = ec.generate_private_key(ec.SECP256R1())
return private_key, {**json.loads(ECAlgorithm.to_jwk(private_key.public_key())), "kid": "ec"}
def _free_port() -> int:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
return reserve.getsockname()[1]
def _observed_bodies() -> tuple[str, ...]:
response: Final = httpx.get(f"{os.environ['INTEGRATION_UPSTREAM_URL']}/__observations", timeout=15)
assert response.status_code == 200, response.text
return tuple(json.dumps(entry["body"]) for entry in response.json()["requests"])
def _jwt_auth_config(tmp_path: Path) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["general_settings"] = {
**config["general_settings"],
"enable_jwt_auth": True,
"litellm_jwtauth": {"user_id_jwt_field": "sub", "user_id_upsert": True},
}
path: Final = tmp_path / "jwt_allowlist.yaml"
path.write_text(yaml.safe_dump(config))
return path
def _token(key: object, algorithm: str, subject: str, kid: str = KEY_ID) -> str:
return jwt.encode(
{"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300},
key, # pyright: ignore[reportArgumentType] # jwt.encode takes Any key material
algorithm=algorithm,
headers={"kid": kid},
)
def test_eddsa_signed_token_is_accepted_and_logged_as_deprecated_outside_fips_mode(
gateway: Gateway, tmp_path: Path
) -> None:
private_key: Final = ed25519.Ed25519PrivateKey.generate()
public_jwk: Final = OKPAlgorithm.to_jwk(private_key.public_key())
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply(public_jwk)
with wire_server(respond) as jwks, gateway.scenario() as scenario:
subject: Final = f"integration-jwt-{uuid.uuid4().hex}"
token: Final = _token(private_key, "EdDSA", subject)
with owned_proxy_process(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path)
) as owned:
model: Final = scenario.model()
scenario.cleanups.callback(scenario.delete_user, subject)
response: Final = owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist control"}]},
key=token,
)
assert response.status_code == 200, response.text
deprecation: Final = eventually(
lambda: owned.log.read_text(),
lambda text: "EdDSA" in text and "deprecated" in text and "LITELLM_FIPS_MODE" in text,
seconds=30,
)
assert "EdDSA" in deprecation and "deprecated" in deprecation and "LITELLM_FIPS_MODE" in deprecation
def test_rs256_signed_token_is_accepted_without_deprecation_log(gateway: Gateway, tmp_path: Path) -> None:
private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
public_jwk: Final = RSAAlgorithm.to_jwk(private_key.public_key())
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply(public_jwk)
with wire_server(respond) as jwks, gateway.scenario() as scenario:
subject: Final = f"integration-jwt-{uuid.uuid4().hex}"
token: Final = _token(private_key, "RS256", subject)
with owned_proxy_process(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path)
) as owned:
model: Final = scenario.model()
scenario.cleanups.callback(scenario.delete_user, subject)
response: Final = owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist control"}]},
key=token,
)
assert response.status_code == 200, response.text
assert "EdDSA" not in owned.log.read_text(), owned.log.read_text()
def _deprecation_lines(log_text: str) -> list[str]:
return [line for line in log_text.splitlines() if all(part in line for part in APPROVED_WARNING_PARTS)]
def _alg_lied_token(private_key: ed25519.Ed25519PrivateKey, kid: str, subject: str) -> str:
def segment(value: bytes) -> str:
return base64.urlsafe_b64encode(value).rstrip(b"=").decode()
signing_input: Final = (
f"{segment(json.dumps({'alg': 'RS256', 'typ': 'JWT', 'kid': kid}).encode())}."
f"{segment(json.dumps({'sub': subject, 'iat': int(time.time()), 'exp': int(time.time()) + 300}).encode())}"
)
signature: Final = segment(private_key.sign(signing_input.encode()))
return f"{signing_input}.{signature}"
def test_eddsa_non_stream_request_reaches_upstream_and_logs_deprecation(gateway: Gateway, tmp_path: Path) -> None:
private_key, jwks_key = _eddsa_keypair()
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([jwks_key])
with wire_server(respond) as jwks, gateway.scenario() as scenario:
subject: Final = f"integration-jwt-{uuid.uuid4().hex}"
token: Final = _token(private_key, "EdDSA", subject, kid="ed")
with owned_proxy_process(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2
) as owned:
model: Final = scenario.model()
scenario.cleanups.callback(scenario.delete_user, subject)
marker: Final = f"jwt-allowlist-a1-{uuid.uuid4().hex}"
_observed_bodies()
response: Final = owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": marker}]},
key=token,
)
assert response.status_code == 200, response.text
deprecation: Final = eventually(
lambda: owned.log.read_text(),
lambda text: _deprecation_lines(text) != [],
seconds=30,
)
assert _deprecation_lines(deprecation) != []
observed: Final = _observed_bodies()
assert any(marker in body for body in observed), observed
def test_eddsa_streamed_request_through_openai_sdk_reaches_upstream_and_logs_deprecation(
gateway: Gateway, tmp_path: Path
) -> None:
private_key, jwks_key = _eddsa_keypair()
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([jwks_key])
with wire_server(respond) as jwks, gateway.scenario() as scenario:
subject: Final = f"integration-jwt-{uuid.uuid4().hex}"
token: Final = _token(private_key, "EdDSA", subject, kid="ed")
with owned_proxy_process(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2
) as owned:
model: Final = scenario.model()
scenario.cleanups.callback(scenario.delete_user, subject)
marker: Final = f"jwt-allowlist-a3-{uuid.uuid4().hex}"
_observed_bodies()
sdk: Final = OpenAI(base_url=f"{owned.gateway.client.base_url}/v1", api_key=token, timeout=30)
chunks: Final = [
chunk
for chunk in sdk.chat.completions.create(
model=model, messages=[{"role": "user", "content": marker}], stream=True
)
]
assert chunks, "stream produced no chunks"
deprecation: Final = eventually(
lambda: owned.log.read_text(),
lambda text: _deprecation_lines(text) != [],
seconds=30,
)
assert _deprecation_lines(deprecation) != []
observed: Final = _observed_bodies()
assert any(marker in body for body in observed), observed
def test_rs256_messages_request_through_anthropic_sdk_reaches_upstream(gateway: Gateway, tmp_path: Path) -> None:
private_key, jwks_key = _rsa_keypair()
completion: Final = json.dumps(
{
"id": "msg_allowlist_a4",
"type": "message",
"role": "assistant",
"model": "claude-3-5-haiku-20241022",
"content": [{"type": "text", "text": "a4-reply"}],
"stop_reason": "end_turn",
"usage": {"input_tokens": 5, "output_tokens": 2},
}
).encode()
def jwks(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([jwks_key])
def provider(request: Request) -> Reply:
return Reply(body=completion)
with wire_server(jwks) as keys_server, wire_server(provider) as provider_server, gateway.scenario() as scenario:
subject: Final = f"integration-jwt-{uuid.uuid4().hex}"
token: Final = _token(private_key, "RS256", subject, kid="rsa")
with owned_proxy_process(
gateway,
tmp_path,
{"JWT_PUBLIC_KEY_URL": keys_server.url},
config=_jwt_auth_config(tmp_path),
workers=2,
) as owned:
model: Final = scenario.model(model="anthropic/claude-3-5-haiku-latest", api_base=provider_server.url)
scenario.cleanups.callback(scenario.delete_user, subject)
sdk: Final = Anthropic(
base_url=str(owned.gateway.client.base_url).rstrip("/"),
auth_token=token,
timeout=30,
max_retries=0,
)
message: Final = sdk.messages.create(
model=model, max_tokens=8, messages=[{"role": "user", "content": "a4"}]
)
text: Final = message.content[0].text
assert "a4-reply" in text, message
received: Final = provider_server.drain()
assert received, "upstream never received the request"
def test_rs256_responses_request_through_openai_async_sdk_reaches_upstream(gateway: Gateway, tmp_path: Path) -> None:
private_key, jwks_key = _rsa_keypair()
response_object: Final = json.dumps(
{
"id": "resp_allowlist_a5",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o-mini",
"output": [
{
"type": "message",
"id": "msg_a5",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "a5-reply", "annotations": []}],
}
],
"usage": {"input_tokens": 5, "output_tokens": 2, "total_tokens": 7},
}
).encode()
def jwks(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([jwks_key])
def provider(request: Request) -> Reply:
return Reply(body=response_object)
with wire_server(jwks) as keys_server, wire_server(provider) as provider_server, gateway.scenario() as scenario:
subject: Final = f"integration-jwt-{uuid.uuid4().hex}"
token: Final = _token(private_key, "RS256", subject, kid="rsa")
with owned_proxy_process(
gateway,
tmp_path,
{"JWT_PUBLIC_KEY_URL": keys_server.url},
config=_jwt_auth_config(tmp_path),
workers=2,
) as owned:
model: Final = scenario.model(api_base=provider_server.url)
scenario.cleanups.callback(scenario.delete_user, subject)
async def call() -> str:
sdk: Final = AsyncOpenAI(base_url=f"{owned.gateway.client.base_url}/v1", api_key=token, timeout=30)
response: Final = await sdk.responses.create(model=model, input="a5")
return response.id
identity: Final = asyncio.run(call())
assert identity.startswith("resp_"), identity
received: Final = provider_server.drain()
assert received, "upstream never received the request"
def test_es256_signed_token_is_accepted(gateway: Gateway, tmp_path: Path) -> None:
private_key, jwks_key = _ec_keypair()
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([jwks_key])
with wire_server(respond) as jwks, gateway.scenario() as scenario:
subject: Final = f"integration-jwt-{uuid.uuid4().hex}"
token: Final = jwt.encode(
{"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300},
private_key,
algorithm="ES256",
headers={"kid": "ec"},
)
with owned_proxy_process(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2
) as owned:
model: Final = scenario.model()
scenario.cleanups.callback(scenario.delete_user, subject)
response: Final = owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist es256"}]},
key=token,
)
assert response.status_code == 200, response.text
def test_hs256_token_against_oct_jwks_key_is_rejected(gateway: Gateway, tmp_path: Path) -> None:
secret: Final = b"integration-hs256-jwt-secret-0123456789abcdef"
oct_key: Final = {
"kty": "oct",
"kid": "sym",
"alg": "HS256",
"k": base64.urlsafe_b64encode(secret).rstrip(b"=").decode(),
}
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([oct_key])
with wire_server(respond) as jwks, gateway.scenario() as scenario:
subject: Final = f"integration-jwt-{uuid.uuid4().hex}"
token: Final = jwt.encode(
{"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300},
secret,
algorithm="HS256",
headers={"kid": "sym"},
)
with owned_proxy_process(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2
) as owned:
model: Final = scenario.model()
response: Final = owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist hs256"}]},
key=token,
)
assert response.status_code == 401, response.text
assert "error" in response.text, response.text
def test_token_signed_eddsa_with_rs256_header_is_rejected(gateway: Gateway, tmp_path: Path) -> None:
private_key, jwks_key = _eddsa_keypair()
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([jwks_key])
with wire_server(respond) as jwks, gateway.scenario() as scenario:
subject: Final = f"integration-jwt-{uuid.uuid4().hex}"
token: Final = _alg_lied_token(private_key, "ed", subject)
with owned_proxy_process(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2
) as owned:
model: Final = scenario.model()
response: Final = owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist lied"}]},
key=token,
)
assert response.status_code == 401, response.text
def test_request_without_authorization_header_is_rejected(gateway: Gateway, tmp_path: Path) -> None:
_, jwks_key = _rsa_keypair()
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([jwks_key])
with wire_server(respond) as jwks, gateway.scenario() as scenario:
with owned_proxy_process(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2
) as owned:
model: Final = scenario.model()
response: Final = owned.gateway.client.post(
"/v1/chat/completions",
json={"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist noauth"}]},
)
assert response.status_code == 401, response.text
def test_empty_jwks_document_rejects_tokens(gateway: Gateway, tmp_path: Path) -> None:
private_key, _jwks_key = _rsa_keypair()
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([])
with wire_server(respond) as jwks, gateway.scenario() as scenario:
subject: Final = f"integration-jwt-{uuid.uuid4().hex}"
token: Final = jwt.encode(
{"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300},
private_key,
algorithm="RS256",
headers={"kid": "rsa"},
)
with owned_proxy_process(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2
) as owned:
model: Final = scenario.model()
response: Final = owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist empty"}]},
key=token,
)
assert response.status_code == 401, response.text
def test_five_eddsa_requests_warn_at_most_once_per_worker(gateway: Gateway, tmp_path: Path) -> None:
private_key, jwks_key = _eddsa_keypair()
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([jwks_key])
with wire_server(respond) as jwks, gateway.scenario() as scenario:
subject: Final = f"integration-jwt-{uuid.uuid4().hex}"
token: Final = _token(private_key, "EdDSA", subject, kid="ed")
with owned_proxy_process(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2
) as owned:
model: Final = scenario.model()
scenario.cleanups.callback(scenario.delete_user, subject)
for _ in range(5):
response: Final = owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist burst"}]},
key=token,
)
assert response.status_code == 200, response.text
log_text: Final = eventually(
lambda: owned.log.read_text(),
lambda text: _deprecation_lines(text) != [],
seconds=30,
)
warnings: Final = _deprecation_lines(log_text)
assert 1 <= len(warnings) <= 2, warnings
def test_master_key_request_still_works_on_jwt_enabled_proxy(gateway: Gateway, tmp_path: Path) -> None:
_, jwks_key = _rsa_keypair()
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([jwks_key])
with wire_server(respond) as jwks, gateway.scenario() as scenario:
with owned_proxy_process(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2
) as owned:
model: Final = scenario.model()
response: Final = owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist master"}]},
)
assert response.status_code == 200, response.text
def _response_ids(status: int, text: str) -> tuple[str, ...]:
if status != 200:
return ()
stripped: Final = text.strip()
if stripped.startswith("{"):
body: Final = json.loads(stripped)
return (str(body["id"]),) if "id" in body else ()
ids: Final = {
str(chunk["id"])
for line in stripped.splitlines()
if line.startswith("data: ") and line[len("data: ") :].strip() != "[DONE]"
for chunk in (json.loads(line[len("data: ") :]),)
if "id" in chunk
}
return tuple(ids)
def _spend_row_counts(identifiers: set[str]) -> dict[str, int]:
if not identifiers:
return {}
rows: Final = read_rows(
'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s::text[])',
("{" + ",".join(sorted(identifiers)) + "}",),
)
counts: Final = {identifier: 0 for identifier in identifiers}
for row in rows:
identifier = str(row["request_id"])
if identifier in counts:
counts[identifier] += 1
return counts
def test_mixed_jwt_burst_survives_jwks_restart_and_writes_spend_rows(gateway: Gateway, tmp_path: Path) -> None:
rsa_private, rsa_jwks_key = _rsa_keypair()
ed_private, ed_jwks_key = _eddsa_keypair()
port: Final = _free_port()
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([rsa_jwks_key, ed_jwks_key])
with gateway.scenario() as scenario:
rs_token: Final = jwt.encode(
{"sub": f"integration-jwt-{uuid.uuid4().hex}", "iat": int(time.time()), "exp": int(time.time()) + 900},
rsa_private,
algorithm="RS256",
headers={"kid": "rsa"},
)
ed_token: Final = jwt.encode(
{"sub": f"integration-jwt-{uuid.uuid4().hex}", "iat": int(time.time()), "exp": int(time.time()) + 900},
ed_private,
algorithm="EdDSA",
headers={"kid": "ed"},
)
with owned_proxy_process(
gateway,
tmp_path,
{"JWT_PUBLIC_KEY_URL": f"http://127.0.0.1:{port}/jwks"},
config=_jwt_auth_config(tmp_path),
workers=2,
) as owned:
model: Final = scenario.model()
scenario.cleanups.callback(
scenario.delete_user, str(jwt.decode(rs_token, options={"verify_signature": False})["sub"])
)
scenario.cleanups.callback(
scenario.delete_user, str(jwt.decode(ed_token, options={"verify_signature": False})["sub"])
)
def chat(token: str, index: int, stream: bool = False) -> httpx.Response:
with owned.gateway.client.stream(
"POST",
"/v1/chat/completions",
json={
"model": model,
"messages": [{"role": "user", "content": f"d1-{index}"}],
**({"stream": True} if stream else {}),
},
headers={"Authorization": f"Bearer {token}"},
) as response:
return httpx.Response(
status_code=response.status_code,
headers=dict(response.headers),
content=response.read(),
)
with wire_server(respond, port=port):
warmup_rs: Final = chat(rs_token, 0)
warmup_ed: Final = chat(ed_token, 1)
assert warmup_rs.status_code == 200, warmup_rs.text
assert warmup_ed.status_code == 200, warmup_ed.text
statuses: Final[list[int]] = []
bodies: Final[list[str]] = []
with ThreadPoolExecutor(max_workers=15) as pool:
futures = [
pool.submit(chat, rs_token if index % 2 == 0 else ed_token, index, index % 5 == 0)
for index in range(30)
]
for future in futures:
result = future.result(timeout=60)
statuses.append(result.status_code)
bodies.append(result.text)
assert all(status in (200, 401) for status in statuses), (statuses, bodies[:3])
assert any(status == 200 for status in statuses), statuses
with wire_server(respond, port=port):
fresh_rs: Final = chat(rs_token, 100)
fresh_ed: Final = chat(ed_token, 101)
assert fresh_rs.status_code == 200, fresh_rs.text
assert fresh_ed.status_code == 200, fresh_ed.text
succeeded: Final = {
identifier for status, text in zip(statuses, bodies) for identifier in _response_ids(status, text)
}
succeeded.update(_response_ids(200, fresh_rs.text) + _response_ids(200, fresh_ed.text))
assert len(succeeded) == sum(1 for status in statuses if status == 200) + 2, (
statuses,
bodies[:3],
)
counts: Final = eventually(
lambda: _spend_row_counts(succeeded),
lambda value: value == {identifier: 1 for identifier in succeeded},
seconds=70,
)
assert all(count == 1 for count in counts.values()), counts
def test_burst_survives_killed_worker_and_stays_alive(gateway: Gateway, tmp_path: Path) -> None:
rsa_private, rsa_jwks_key = _rsa_keypair()
ed_private, ed_jwks_key = _eddsa_keypair()
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return _jwks_reply_keys([rsa_jwks_key, ed_jwks_key])
with wire_server(respond) as jwks, gateway.scenario() as scenario:
rs_token: Final = jwt.encode(
{"sub": f"integration-jwt-{uuid.uuid4().hex}", "iat": int(time.time()), "exp": int(time.time()) + 900},
rsa_private,
algorithm="RS256",
headers={"kid": "rsa"},
)
ed_token: Final = jwt.encode(
{"sub": f"integration-jwt-{uuid.uuid4().hex}", "iat": int(time.time()), "exp": int(time.time()) + 900},
ed_private,
algorithm="EdDSA",
headers={"kid": "ed"},
)
with owned_proxy_process(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2
) as owned:
model: Final = scenario.model()
scenario.cleanups.callback(
scenario.delete_user, str(jwt.decode(rs_token, options={"verify_signature": False})["sub"])
)
scenario.cleanups.callback(
scenario.delete_user, str(jwt.decode(ed_token, options={"verify_signature": False})["sub"])
)
def chat(token: str, index: int) -> httpx.Response:
return owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": f"d2-{index}"}]},
key=token,
)
with ThreadPoolExecutor(max_workers=12) as pool:
futures = [pool.submit(chat, rs_token if index % 2 == 0 else ed_token, index) for index in range(24)]
owned_port: Final = owned.gateway.client.base_url.port
listeners: Final = [
child
for child in psutil.Process(owned.process.pid).children(recursive=True)
if any(
connection.status == psutil.CONN_LISTEN and connection.laddr.port == owned_port
for connection in child.net_connections(kind="inet")
)
]
assert len(listeners) >= 2, listeners
killed_pid: Final = listeners[0].pid
listeners[0].kill()
results: Final[list[int]] = []
for future in futures:
try:
result = future.result(timeout=60)
results.append(result.status_code)
except httpx.TransportError:
results.append(0)
assert all(status in (200, 401, 0) for status in results), results
assert any(status == 200 for status in results), results
dead: Final = eventually(
lambda: psutil.pid_exists(killed_pid),
lambda exists: not exists,
seconds=30,
)
assert not dead, f"killed worker pid {killed_pid} still exists"
liveliness: Final = eventually(
lambda: owned.gateway.client.get("/health/liveliness"),
lambda response: response.status_code == 200,
seconds=45,
)
assert liveliness.status_code == 200, liveliness.text

View file

@ -0,0 +1,277 @@
import json
import os
import uuid
from pathlib import Path
from typing import Final
import httpx
import psycopg
import pytest
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.mcp import forget_mcp, mcp_peer, register_mcp
from integration._support.process import owned_proxy_process
def _pem_private_key() -> str:
return (
rsa.generate_private_key(public_exponent=65537, key_size=2048)
.private_bytes(
serialization.Encoding.PEM,
serialization.PrivateFormat.PKCS8,
serialization.NoEncryption(),
)
.decode()
)
def _stored_client_assertion_signing_alg(identity: str) -> str:
rows: Final = read_rows('SELECT credentials FROM "LiteLLM_MCPServerTable" WHERE server_id = %s', (identity,))
assert len(rows) == 1, rows
credentials: Final = rows[0]["credentials"]
if credentials is None:
return "RS256"
blob: Final = credentials if isinstance(credentials, dict) else json.loads(credentials)
return str(blob.get("client_assertion_signing_alg") or "RS256")
def test_non_approved_client_assertion_signing_alg_is_rejected_on_create_and_update(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
alias: Final = "alg" + uuid.uuid4().hex[:8]
created: Final = gateway.request(
"POST",
"/v1/mcp/server",
{
"server_name": alias,
"alias": alias,
**peer.registration(),
"token_exchange_endpoint": "https://idp.integration.invalid/oauth2/token",
"credentials": {
"client_private_key": _pem_private_key(),
"client_assertion_signing_alg": "HS256",
},
},
)
if created.status_code == 201:
scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"]))
assert created.status_code in (400, 422), created.text
assert "client_assertion_signing_alg" in created.text, created.text
identity: Final = register_mcp(scenario, peer, alias + "v")
edited: Final = gateway.request(
"PUT",
"/v1/mcp/server",
{"server_id": identity, "credentials": {"client_assertion_signing_alg": "EdDSA"}},
)
assert edited.status_code in (400, 422), edited.text
assert "client_assertion_signing_alg" in edited.text, edited.text
assert _stored_client_assertion_signing_alg(identity) == "RS256"
def test_approved_client_assertion_signing_alg_round_trips(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
alias: Final = "algok" + uuid.uuid4().hex[:8]
created: Final = gateway.request(
"POST",
"/v1/mcp/server",
{
"server_name": alias,
"alias": alias,
**peer.registration(),
"token_exchange_endpoint": "https://idp.integration.invalid/oauth2/token",
"credentials": {
"client_private_key": _pem_private_key(),
"client_assertion_signing_alg": "ES256",
},
},
)
assert created.status_code == 201, created.text
identity: Final = str(created.json()["server_id"])
scenario.cleanups.callback(forget_mcp, gateway, identity)
assert _stored_client_assertion_signing_alg(identity) == "ES256"
def test_server_row_with_stale_client_assertion_signing_alg_still_loads(gateway: Gateway) -> None:
"""A row persisted before the allowlist (credentials.client_assertion_signing_alg = "HS256")
must keep loading with the RS256 fallback instead of disappearing on upgrade."""
identity: Final = "stalealg" + uuid.uuid4().hex[:8]
with psycopg.connect(os.environ["DATABASE_URL"]) as connection:
connection.execute(
'INSERT INTO "LiteLLM_MCPServerTable" (server_id, server_name, url, transport, credentials,'
" created_at, updated_at) VALUES (%s, %s, %s, %s, %s::jsonb, NOW(), NOW())",
(
identity,
"stalealg" + uuid.uuid4().hex[:8],
"https://mcp.integration.invalid",
"http",
json.dumps({"client_assertion_signing_alg": "HS256"}),
),
)
try:
with mcp_peer() as peer, gateway.scenario() as scenario:
register_mcp(scenario, peer, "staletrigger" + uuid.uuid4().hex[:8])
listed: Final = eventually(
lambda: gateway.request("GET", "/v1/mcp/server"),
lambda response: (
response.status_code == 200
and any(server.get("server_id") == identity for server in response.json())
),
seconds=30,
)
assert listed.status_code == 200, listed.text
finally:
with psycopg.connect(os.environ["DATABASE_URL"]) as connection:
connection.execute('DELETE FROM "LiteLLM_MCPServerTable" WHERE server_id = %s', (identity,))
APPROVED_ALGS: Final = ("RS256", "RS384", "RS512", "PS256", "PS384", "PS512", "ES256", "ES384", "ES512")
def _post_server(gateway: Gateway, name: str, credentials: object, omit_alg: bool = False) -> httpx.Response:
blob: Final = {} if omit_alg else {"client_assertion_signing_alg": credentials}
return gateway.request(
"POST",
"/v1/mcp/server",
{
"server_name": name,
"alias": name,
"url": "https://mcp.integration.invalid",
"transport": "http",
"credentials": blob,
},
)
def test_post_with_lowercase_hs256_is_rejected_naming_the_field(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
created: Final = _post_server(gateway, "b1hs" + uuid.uuid4().hex[:8], "hs256")
if created.status_code == 201:
scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"]))
assert created.status_code == 422, created.text
assert "client_assertion_signing_alg" in created.text, created.text
for algorithm in APPROVED_ALGS:
assert algorithm in created.text, created.text
def test_put_with_eddsa_on_existing_server_is_rejected(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
identity: Final = register_mcp(scenario, peer, "b2ed" + uuid.uuid4().hex[:8])
edited: Final = gateway.request(
"PUT",
"/v1/mcp/server",
{"server_id": identity, "credentials": {"client_assertion_signing_alg": "EdDSA"}},
)
assert edited.status_code == 422, edited.text
assert "client_assertion_signing_alg" in edited.text, edited.text
@pytest.mark.parametrize("algorithm", APPROVED_ALGS)
def test_each_approved_client_assertion_signing_alg_is_accepted_and_stored(gateway: Gateway, algorithm: str) -> None:
with gateway.scenario() as scenario:
name: Final = "b3" + algorithm.lower() + uuid.uuid4().hex[:6]
created: Final = _post_server(gateway, name, algorithm)
assert created.status_code == 201, created.text
identity: Final = str(created.json()["server_id"])
scenario.cleanups.callback(forget_mcp, gateway, identity)
listed: Final = gateway.request("GET", "/v1/mcp/server")
assert listed.status_code == 200, listed.text
assert any(server.get("server_id") == identity for server in listed.json()), listed.text
assert _stored_client_assertion_signing_alg(identity) == algorithm
@pytest.mark.parametrize(
("label", "value", "expected"),
(
("empty", "", 422),
("five-kb", "x" * 5120, 422),
("integer", 7, 422),
("list", ["RS256"], 422),
("null", None, 201),
),
)
def test_client_assertion_signing_alg_payload_variants(
gateway: Gateway, label: str, value: object, expected: int
) -> None:
with gateway.scenario() as scenario:
created: Final = _post_server(gateway, f"b{label}" + uuid.uuid4().hex[:8], value)
if created.status_code == 201:
scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"]))
assert created.status_code == expected, created.text
if expected == 422 and label in ("empty", "five-kb"):
assert "client_assertion_signing_alg" in created.text, created.text
def test_server_post_without_credentials_alg_key_is_accepted(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
created: Final = _post_server(gateway, "b9absent" + uuid.uuid4().hex[:8], None, omit_alg=True)
assert created.status_code == 201, created.text
scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"]))
def test_unauthenticated_server_post_with_hs256_is_rejected(gateway: Gateway) -> None:
response: Final = gateway.client.post(
"/v1/mcp/server",
json={
"server_name": "b10unauth" + uuid.uuid4().hex[:8],
"credentials": {"client_assertion_signing_alg": "hs256"},
},
)
assert response.status_code == 401, response.text
def _seeded_row_loads_with_fallback_warning(gateway: Gateway, tmp_path: Path, algorithm: str) -> None:
identity: Final = "dualread" + uuid.uuid4().hex[:8]
name: Final = "dualread" + uuid.uuid4().hex[:8]
with psycopg.connect(os.environ["DATABASE_URL"]) as connection:
connection.execute(
'INSERT INTO "LiteLLM_MCPServerTable" (server_id, server_name, url, transport, credentials,'
" created_at, updated_at) VALUES (%s, %s, %s, %s, %s::jsonb, NOW(), NOW())",
(
identity,
name,
"https://mcp.integration.invalid",
"http",
json.dumps({"client_assertion_signing_alg": algorithm}),
),
)
try:
with owned_proxy_process(gateway, tmp_path, {}) as owned:
listed: Final = eventually(
lambda: owned.gateway.request("GET", "/v1/mcp/server"),
lambda response: (
response.status_code == 200
and any(server.get("server_id") == identity for server in response.json())
),
seconds=30,
)
assert listed.status_code == 200, listed.text
log_text: Final = eventually(
lambda: owned.log.read_text(),
lambda text: name in text and "not an approved algorithm" in text and "using RS256" in text,
seconds=30,
)
assert name in log_text and "not an approved algorithm" in log_text and "using RS256" in log_text
finally:
with psycopg.connect(os.environ["DATABASE_URL"]) as connection:
connection.execute('DELETE FROM "LiteLLM_MCPServerTable" WHERE server_id = %s', (identity,))
def test_seeded_hs256_row_loads_with_rs256_fallback_warning(gateway: Gateway, tmp_path: Path) -> None:
_seeded_row_loads_with_fallback_warning(gateway, tmp_path, "HS256")
def test_seeded_eddsa_row_loads_with_rs256_fallback_warning(gateway: Gateway, tmp_path: Path) -> None:
_seeded_row_loads_with_fallback_warning(gateway, tmp_path, "EdDSA")
def test_chat_still_works_after_alg_rejection(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
created: Final = _post_server(gateway, "b13chat" + uuid.uuid4().hex[:8], "hs256")
assert created.status_code in (201, 422), created.text
if created.status_code == 201:
scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"]))
model: Final = scenario.model()
body: Final = gateway.chat(model)
assert body["id"], body

View file

@ -0,0 +1,291 @@
import base64
import json
import os
import signal
import socket
import subprocess
import sys
import time
import uuid
from collections.abc import Callable, Generator, Mapping
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from pathlib import Path
from typing import Final
import httpx
import jwt
import pytest
from cryptography.hazmat.primitives.asymmetric import ed25519, rsa
from integration._support.client import Gateway, eventually
from integration._support.mcp import mcp_peer, register_mcp, tool_calls, tool_names
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
from jwt.algorithms import OKPAlgorithm, RSAAlgorithm
@contextmanager
def _jwt_signer_guardrail(gateway: Gateway, discovery_uri: str) -> Generator[None]:
created: Final = gateway.client.post(
"/guardrails",
headers={"x-litellm-api-key": gateway.key},
json={
"guardrail": {
"guardrail_name": "signer" + uuid.uuid4().hex[:8],
"litellm_params": {
"guardrail": "mcp_jwt_signer",
"mode": "pre_mcp_call",
"default_on": True,
"access_token_discovery_uri": discovery_uri,
},
}
},
)
assert created.status_code == 200, created.text
identity: Final = str(created.json()["guardrail_id"])
try:
yield
finally:
deleted: Final = gateway.client.delete(f"/guardrails/{identity}", headers={"x-litellm-api-key": gateway.key})
assert deleted.status_code == 200, deleted.text
@contextmanager
def _idp_server(jwks_keys: list[Mapping[str, object]]) -> Generator[str]:
holder: Final = {"url": ""}
def respond(request: Request) -> Reply:
if request.target == "/.well-known/openid-configuration":
return Reply(body=json.dumps({"jwks_uri": holder["url"] + "/.well-known/jwks.json"}).encode())
if request.target == "/.well-known/jwks.json":
return Reply(body=json.dumps({"keys": list(jwks_keys)}).encode())
return Reply(status=404)
with wire_server(respond) as idp:
holder["url"] = idp.url
yield idp.url + "/.well-known/openid-configuration"
def _claims() -> dict[str, object]:
now: Final = int(time.time())
return {"sub": "integration-mcp-user", "iat": now, "exp": now + 300}
def _hs256_key_and_token() -> tuple[dict[str, object], str]:
secret: Final = b"integration-hs256-client-secret-0123456789abcdef"
key: Final = {
"kty": "oct",
"alg": "HS256",
"kid": "sym",
"k": base64.urlsafe_b64encode(secret).rstrip(b"=").decode(),
}
token: Final = jwt.encode(_claims(), secret, algorithm="HS256", headers={"kid": "sym"})
return key, token
def _eddsa_key_and_token() -> tuple[dict[str, object], str]:
private_key: Final = ed25519.Ed25519PrivateKey.generate()
public_jwk: Final = json.loads(OKPAlgorithm.to_jwk(private_key.public_key()))
key: Final = {**public_jwk, "alg": "EdDSA", "kid": "ed"}
token: Final = jwt.encode(_claims(), private_key, algorithm="EdDSA", headers={"kid": "ed"})
return key, token
def _rs256_key_and_token() -> tuple[dict[str, object], str]:
private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
public_jwk: Final = json.loads(RSAAlgorithm.to_jwk(private_key.public_key()))
key: Final = {**public_jwk, "alg": "RS256", "kid": "rsa"}
token: Final = jwt.encode(_claims(), private_key, algorithm="RS256", headers={"kid": "rsa"})
return key, token
def _call_with_bearer(
gateway: Gateway, key: str, identity: str, name: str, arguments: dict[str, object], bearer: str
) -> httpx.Response:
return gateway.client.post(
"/mcp-rest/tools/call",
headers={"x-litellm-api-key": key, "Authorization": f"Bearer {bearer}"},
json={"server_id": identity, "name": name, "arguments": arguments},
)
@pytest.mark.parametrize("key_and_token", (_hs256_key_and_token, _eddsa_key_and_token), ids=("oct-HS256", "OKP-EdDSA"))
def test_jwks_key_with_non_approved_alg_cannot_verify_the_incoming_token(
gateway: Gateway, key_and_token: Callable[[], tuple[dict[str, object], str]]
) -> None:
jwks_key, token = key_and_token()
with _idp_server([jwks_key]) as discovery_uri, _jwt_signer_guardrail(gateway, discovery_uri):
with mcp_peer() as peer, gateway.scenario() as scenario:
alias: Final = "jwksalg" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
name: Final = tool_names(gateway, key, identity)["add"]
peer.drain()
response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, token)
assert response.status_code == 401, response.text
assert "incoming token verification failed" in response.text, response.text
assert tool_calls(peer.drain()) == (), "rejected call reached the peer"
def test_jwks_rs256_key_still_verifies_the_incoming_token(gateway: Gateway) -> None:
jwks_key, token = _rs256_key_and_token()
with _idp_server([jwks_key]) as discovery_uri, _jwt_signer_guardrail(gateway, discovery_uri):
with mcp_peer() as peer, gateway.scenario() as scenario:
alias: Final = "jwksok" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
name: Final = tool_names(gateway, key, identity)["add"]
peer.drain()
response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, token)
assert response.status_code == 200, response.text
assert response.json()["content"][0]["text"] == "9", response.text
assert len(tool_calls(peer.drain())) == 1, "accepted call never reached the peer"
def _rs256_keypair() -> tuple[rsa.RSAPrivateKey, dict[str, object]]:
private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
return private_key, json.loads(RSAAlgorithm.to_jwk(private_key.public_key()))
def _rs256_key_and_token_no_alg() -> tuple[dict[str, object], str]:
private_key, public_jwk = _rs256_keypair()
key: Final = {**public_jwk, "kid": "rsa"}
token: Final = jwt.encode(_claims(), private_key, algorithm="RS256", headers={"kid": "rsa"})
return key, token
def _signer_call(gateway: Gateway, key: str, identity: str, name: str, bearer: str) -> tuple[int, str]:
try:
response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, bearer)
return response.status_code, response.text
except httpx.HTTPError as error:
return 0, repr(error)
def test_jwks_rs256_key_without_alg_field_still_verifies_the_token(gateway: Gateway) -> None:
jwks_key, token = _rs256_key_and_token_no_alg()
with _idp_server([jwks_key]) as discovery_uri, _jwt_signer_guardrail(gateway, discovery_uri):
with mcp_peer() as peer, gateway.scenario() as scenario:
identity: Final = register_mcp(scenario, peer, "jwksc4" + uuid.uuid4().hex[:8])
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
name: Final = tool_names(gateway, key, identity)["add"]
peer.drain()
response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, token)
assert response.status_code == 200, response.text
assert response.json()["content"][0]["text"] == "9", response.text
assert len(tool_calls(peer.drain())) == 1, "accepted call never reached the peer"
def test_rs256_token_verifies_when_jwks_also_carries_a_non_approved_key(gateway: Gateway) -> None:
okp_key, _eddsa_token = _eddsa_key_and_token()
jwks_key, token = _rs256_key_and_token()
with _idp_server([okp_key, jwks_key]) as discovery_uri, _jwt_signer_guardrail(gateway, discovery_uri):
with mcp_peer() as peer, gateway.scenario() as scenario:
identity: Final = register_mcp(scenario, peer, "jwksc5" + uuid.uuid4().hex[:8])
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
name: Final = tool_names(gateway, key, identity)["add"]
peer.drain()
response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, token)
assert response.status_code == 200, response.text
assert response.json()["content"][0]["text"] == "9", response.text
assert len(tool_calls(peer.drain())) == 1, "accepted call never reached the peer"
@pytest.mark.parametrize(
("label", "jwks_reply"),
(
("jwks-404", Reply(status=404, body=b'{"error": "missing"}')),
("jwks-not-json", Reply(body=b"this is not json")),
),
ids=("jwks-404", "jwks-not-json"),
)
def test_unusable_jwks_document_rejects_the_incoming_token(gateway: Gateway, label: str, jwks_reply: Reply) -> None:
_, token = _rs256_key_and_token()
def respond(request: Request) -> Reply:
if request.target == "/.well-known/openid-configuration":
return Reply(body=json.dumps({"jwks_uri": holder["url"] + "/.well-known/jwks.json"}).encode())
return jwks_reply
holder: Final = {"url": ""}
with wire_server(respond) as idp:
holder["url"] = idp.url
with _jwt_signer_guardrail(gateway, idp.url + "/.well-known/openid-configuration"):
with mcp_peer() as peer, gateway.scenario() as scenario:
identity: Final = register_mcp(scenario, peer, "jwks" + uuid.uuid4().hex[:8])
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
name: Final = tool_names(gateway, key, identity)["add"]
peer.drain()
status, text = _signer_call(gateway, key, identity, name, token)
assert status == 401, (status, text)
assert "incoming token verification failed" in text, text
assert tool_calls(peer.drain()) == (), "rejected call reached the peer"
alive: Final = gateway.client.get("/health/liveliness")
assert alive.status_code == 200, alive.text
def test_signer_jwks_pause_rejects_then_recovers(gateway: Gateway, tmp_path: Path) -> None:
private_key, public_jwk = _rs256_keypair()
jwks_dir: Final = tmp_path / "jwks" / ".well-known"
jwks_dir.mkdir(parents=True)
port: Final = _reserve_port()
(jwks_dir / "jwks.json").write_text(json.dumps({"keys": [{**public_jwk, "kid": "rsa", "alg": "RS256"}]}))
(jwks_dir / "openid-configuration").write_text(
json.dumps({"jwks_uri": f"http://127.0.0.1:{port}/.well-known/jwks.json"})
)
token: Final = jwt.encode(_claims(), private_key, algorithm="RS256", headers={"kid": "rsa"})
server: Final = subprocess.Popen(
[
sys.executable,
"-I",
"-m",
"http.server",
str(port),
"--bind",
"127.0.0.1",
"--directory",
str(tmp_path / "jwks"),
],
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)
try:
with owned_proxy_process(gateway, tmp_path, {}) as owned:
discovery: Final = f"http://127.0.0.1:{port}/.well-known/openid-configuration"
with _jwt_signer_guardrail(owned.gateway, discovery):
with mcp_peer() as peer, owned.gateway.scenario() as scenario:
identity: Final = register_mcp(scenario, peer, "jwksd3" + uuid.uuid4().hex[:8])
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
name: Final = tool_names(owned.gateway, key, identity)["add"]
peer.drain()
os.kill(server.pid, signal.SIGSTOP)
outcomes: Final = []
with ThreadPoolExecutor(max_workers=8) as pool:
futures = [
pool.submit(_signer_call, owned.gateway, key, identity, name, token) for _ in range(8)
]
for future in futures:
outcomes.append(future.result(timeout=80))
assert all(status != 200 for status, _text in outcomes), outcomes
assert any(status in (0, 401, 500) for status, _text in outcomes), outcomes
os.kill(server.pid, signal.SIGCONT)
recovered: Final = eventually(
lambda: _signer_call(owned.gateway, key, identity, name, token),
lambda outcome: outcome[0] == 200,
seconds=70,
)
assert recovered[0] == 200, recovered
assert tool_calls(peer.drain()) != (), "resumed call never reached the peer"
finally:
try:
os.kill(server.pid, signal.SIGCONT)
except ProcessLookupError:
pass
server.terminate()
server.wait(timeout=10)
def _reserve_port() -> int:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
return reserve.getsockname()[1]

View file

@ -16,7 +16,7 @@ import jwt
import litellm
import pytest
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from cryptography.hazmat.primitives.asymmetric import ec, rsa
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
Error,
@ -50,6 +50,15 @@ _PRIVATE_PEM = _RSA_KEY.private_bytes(
serialization.PrivateFormat.PKCS8,
serialization.NoEncryption(),
).decode()
_EC_PRIVATE_PEM = (
ec.generate_private_key(ec.SECP256R1())
.private_bytes(
serialization.Encoding.PEM,
serialization.PrivateFormat.PKCS8,
serialization.NoEncryption(),
)
.decode()
)
_PUBLIC_PEM = (
_RSA_KEY.public_key()
.public_bytes(
@ -185,9 +194,9 @@ async def test_fetch_network_error_maps_to_upstream_unavailable(raised):
"auth",
[
PrivateKeyJwtAuth(private_key=SecretStr("not-a-pem-key"), signing_alg="RS256"),
PrivateKeyJwtAuth(private_key=SecretStr(_PRIVATE_PEM), signing_alg="XX999"),
PrivateKeyJwtAuth(private_key=SecretStr(_EC_PRIVATE_PEM), signing_alg="RS256"),
],
ids=["garbage-key", "unknown-alg"],
ids=["garbage-key", "key-alg-mismatch"],
)
async def test_fetch_unsignable_client_assertion_is_misconfigured_not_a_crash(auth):
client = AsyncMock()

View file

@ -1794,3 +1794,20 @@ async def test_unverified_legacy_cache_cannot_bypass_enforcement(monkeypatch):
await mcp_per_user_token_cache.set("alice", "srv", "bob", 60)
assert await module.resolve_user_oauth_access_token("alice", server) is None
assert await mcp_per_user_token_cache.get("alice", "srv") is None
def test_server_table_row_with_stale_client_assertion_signing_alg_still_validates() -> None:
"""Rows written before the approved-algorithm allowlist carry values like HS256 in the
credentials blob; the stored shape stays lenient so reload_servers_from_database still
parses the row and _stored_client_assertion_signing_alg falls back to RS256 instead of
the server silently disappearing on upgrade."""
row: Final = {
"server_id": "srv-stale-alg",
"server_name": "stale_alg_server",
"url": "https://mcp.example.com",
"transport": "http",
"credentials": {"client_assertion_signing_alg": "HS256"},
}
parsed: Final = LiteLLM_MCPServerTable.model_validate(row)
assert parsed.credentials is not None
assert parsed.credentials["client_assertion_signing_alg"] == "HS256"

View file

@ -1,7 +1,9 @@
from litellm.proxy._experimental.mcp_server.upstream import resolve_upstream_auth
import importlib
import asyncio
# Add the parent directory to the path so we can import litellm
import contextlib
import functools
import importlib
import json
import logging
import os
@ -10,26 +12,13 @@ import time
from collections.abc import AsyncIterator
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, Final, Literal, Optional
from typing import Any, Final, Literal
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from respx import MockRouter
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPServerListError,
MCPUpstreamAuthError,
)
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ServerListFault
# Add the parent directory to the path so we can import litellm
import contextlib
import httpx
import httpx2
import pytest
from fastapi import HTTPException
from mcp import ReadResourceResult, Resource
from mcp.types import (
CallToolResult,
@ -40,24 +29,37 @@ from mcp.types import (
)
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl, TypeAdapter
from respx import MockRouter
import litellm
import litellm.llms as litellm_llms
from litellm.caching.caching import DualCache
from litellm.caching.llm_caching_handler import LLMClientCache
from litellm.constants import MCP_METADATA_TIMEOUT
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._experimental.mcp_server.tool_outcome import TextResult
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPServerListError,
MCPUpstreamAuthError,
)
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ServerListFault
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
ListedToolsCaller,
MCPServerManager,
_deserialize_json_dict,
_deserialize_json_list,
_flow_endpoints_missing,
_mcp_oauth_discovery_on_startup_enabled,
_oauth_endpoints_unresolved,
_deserialize_json_list,
_normalize_mcp_server_cost_info,
_oauth_endpoints_unresolved,
_obo_retry_applies,
_resolve_openapi_tool_auth,
_should_strip_caller_authorization,
listed_tools_caller_for,
)
from litellm.proxy._experimental.mcp_server.tool_outcome import TextResult
from litellm.proxy._experimental.mcp_server.upstream import resolve_upstream_auth
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
LiteLLM_ObjectPermissionTable,
@ -68,18 +70,12 @@ from litellm.proxy._types import (
MCPTransport,
UserAPIKeyAuth,
)
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import MCPAuth, MCPAuthType, MCPUpstreamProtocol
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer, PinnedMCPTool
from litellm.caching.caching import DualCache
from litellm.caching.llm_caching_handler import LLMClientCache
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
import litellm.llms as litellm_llms
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.integrations.slack_alerting import AlertType
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import MCPAuth, MCPAuthType, MCPUpstreamProtocol
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer, PinnedMCPTool
@pytest.mark.asyncio
@ -2134,7 +2130,7 @@ class TestMCPServerManager:
assert exc_info.value.fault == ServerListFault(tag="internal", status_code=412)
assert exc_info.value.server_name == "te-412-server"
def _upstream_status_error(self, status_code: int, www_authenticate: Optional[str] = None) -> httpx.HTTPStatusError:
def _upstream_status_error(self, status_code: int, www_authenticate: str | None = None) -> httpx.HTTPStatusError:
"""Build an httpx.HTTPStatusError shaped like the one the MCP SDK surfaces for an upstream
HTTP failure, so _extract_upstream_auth_failure can read status_code and WWW-Authenticate."""
request = httpx.Request("POST", "https://up.example.com/mcp")
@ -2804,11 +2800,11 @@ class TestMCPServerManager:
assert built.token_url == "https://idp.example.com/manual-token"
assert built.scopes == ["calendar.read"]
async def _capture_subject_token(self, call) -> Optional[str]:
async def _capture_subject_token(self, call) -> str | None:
"""Run a manager method (via ``call(manager)``) and return the subject_token it threaded
into ``_create_mcp_client``."""
manager = MCPServerManager()
captured: Dict[str, Any] = {}
captured: dict[str, Any] = {}
async def capture_create_mcp_client(
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None, **kwargs
@ -9103,7 +9099,7 @@ class TestMCPServerTimestamps:
_carry_forward_resolved_oauth_endpoints,
)
def make_server(url: str, auth_type: MCPAuth, authorization_url: Optional[str]) -> MCPServer:
def make_server(url: str, auth_type: MCPAuth, authorization_url: str | None) -> MCPServer:
return MCPServer(
server_id="s1",
name="s1",
@ -10476,7 +10472,7 @@ class TestOAuthDiscoverySSRFGuard:
}
mock_response.raise_for_status = MagicMock()
captured_kwargs: Dict[str, Any] = {}
captured_kwargs: dict[str, Any] = {}
async def fake_get(url, **kwargs):
captured_kwargs.update(kwargs)
@ -10507,7 +10503,7 @@ class TestOAuthDiscoverySSRFGuard:
}
mock_response.raise_for_status = MagicMock()
captured_kwargs: Dict[str, Any] = {}
captured_kwargs: dict[str, Any] = {}
async def fake_get(url, **kwargs):
captured_kwargs.update(kwargs)
@ -10748,8 +10744,8 @@ class TestHealthCheckInterpolatesGlobalEnvVars:
)
@staticmethod
def _capture_headers(manager: MCPServerManager) -> Dict[str, Any]:
captured: Dict[str, Any] = {}
def _capture_headers(manager: MCPServerManager) -> dict[str, Any]:
captured: dict[str, Any] = {}
mock_client = AsyncMock()
mock_client.run_with_session = AsyncMock(return_value="ok")
@ -10796,7 +10792,7 @@ class TestUserEnvVarsCacheEviction:
def _patch_cache(monkeypatch, max_size):
from litellm.proxy._experimental.mcp_server import mcp_server_manager as m
cache: Dict[Any, Any] = {}
cache: dict[Any, Any] = {}
monkeypatch.setattr(m, "_user_env_vars_cache", cache)
monkeypatch.setattr(m, "_USER_ENV_VARS_CACHE_MAX_SIZE", max_size)
return m, cache
@ -11022,7 +11018,7 @@ class TestCreateMcpClientV2Graft:
"""
def _http_server(self, **overrides: Any) -> MCPServer:
base: Dict[str, Any] = dict(
base: dict[str, Any] = dict(
server_id="http-graft",
name="graft_server",
url="https://upstream.example.com/mcp",
@ -13943,7 +13939,7 @@ class TestConfigServerIdPinning:
@pytest.mark.asyncio
async def test_two_servers_pinning_the_same_id_are_rejected(self, config_only_mcp_manager_factory):
manager = config_only_mcp_manager_factory()
config: Dict[str, Any] = {
config: dict[str, Any] = {
"docs_server": {"url": "https://a.example.com/mcp", "server_id": "shared-id"},
"wiki_server": {"url": "https://b.example.com/mcp", "server_id": "shared-id"},
}
@ -13962,7 +13958,7 @@ class TestConfigServerIdPinning:
auth_type=None,
alias=None,
)
config: Dict[str, Any] = {
config: dict[str, Any] = {
"docs_server": {"url": "https://a.example.com/mcp", "transport": MCPTransport.http},
"wiki_server": {"url": "https://b.example.com/mcp", "server_id": derived},
}
@ -14793,11 +14789,11 @@ async def test_debug_resolution_matches_final_header_conflict_winner(
expected_source: str,
expected_authorization: str | None,
) -> None:
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
from starlette.requests import Request
from pydantic import SecretStr
from starlette.requests import Request
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import MCPAuthenticatedUser
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
from litellm.proxy._experimental.mcp_server.mcp_debug import MCP_AUTH_DIAGNOSTICS_SCOPE_KEY, MCPAuthDiagnostics
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
ApiKeyConfig,
@ -14863,9 +14859,9 @@ async def test_debug_reports_legacy_signing_and_non_http_transport(
_mcp_request_ctx, monkeypatch, transport: Literal["http", "stdio"]
) -> None:
monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true")
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
from starlette.requests import Request
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
from litellm.proxy._experimental.mcp_server.mcp_debug import MCP_AUTH_DIAGNOSTICS_SCOPE_KEY, MCPAuthDiagnostics
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -15172,7 +15168,6 @@ class _DiscoveryClock:
return self.now
from pydantic import TypeAdapter
from mcp.types import JSONRPCMessage
_JSONRPC_ADAPTER = TypeAdapter(JSONRPCMessage)
@ -15269,7 +15264,6 @@ def _discovery_server() -> MCPServer:
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ("prompts", "resources", "templates"))
async def test_discovery_cache_reuses_raw_results_and_expires(kind: str) -> None:
import respx
clock: Final = _DiscoveryClock()
manager: Final = MCPServerManager(discovery_clock=clock)
@ -15301,7 +15295,6 @@ async def test_discovery_cache_reuses_raw_results_and_expires(kind: str) -> None
@pytest.mark.parametrize("kind", ("prompts", "resources", "templates"))
@pytest.mark.parametrize("outcome", ("unsupported", "rejected", "failure"))
async def test_discovery_cache_empty_results_and_failures(kind: str, outcome: str) -> None:
import respx
manager: Final = MCPServerManager()
upstream: Final = _DiscoveryUpstream()
@ -15346,7 +15339,6 @@ async def test_discovery_cache_retries_failed_pagination_before_caching_complete
@pytest.mark.asyncio
async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_auth() -> None:
import respx
manager: Final = MCPServerManager()
upstream: Final = _DiscoveryUpstream()
@ -15376,7 +15368,6 @@ async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_
@pytest.mark.asyncio
async def test_discovery_cache_coalesces_and_survives_waiter_cancellation() -> None:
import respx
manager: Final = MCPServerManager()
upstream: Final = _DiscoveryUpstream()
@ -15400,7 +15391,6 @@ async def test_discovery_cache_coalesces_and_survives_waiter_cancellation() -> N
@pytest.mark.asyncio
async def test_discovery_cache_invalidation_during_fetch_does_not_repopulate_old_results() -> None:
import respx
manager: Final = MCPServerManager()
upstream: Final = _DiscoveryUpstream()
@ -15420,7 +15410,6 @@ async def test_discovery_cache_invalidation_during_fetch_does_not_repopulate_old
@pytest.mark.asyncio
async def test_discovery_cache_can_be_disabled(monkeypatch: pytest.MonkeyPatch) -> None:
import respx
monkeypatch.setenv("LITELLM_MCP_DISCOVERY_CACHE_TTL", "0")
manager: Final = MCPServerManager()
@ -15586,7 +15575,6 @@ async def test_discovery_cache_bounds_detached_fetches_without_dropping_results(
@pytest.mark.asyncio
async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> None:
import respx
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import UpstreamCredentialProvider
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result
@ -15649,7 +15637,6 @@ async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> N
@pytest.mark.asyncio
async def test_discovery_resolves_stored_oauth_for_the_requesting_user() -> None:
import respx
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken
class TokenStore:
@ -16090,6 +16077,7 @@ class TestProtectedCredentialPreparation:
@pytest.mark.asyncio
async def test_static_resolution_cancellation_closes_flow(self) -> None:
from collections.abc import AsyncGenerator
from litellm.experimental_mcp_client.client import MCPClient
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import prepare_mcp_client
@ -16355,8 +16343,8 @@ class TestProtectedCredentialPreparation:
@pytest.mark.asyncio
@pytest.mark.parametrize("selected", [False, True])
async def test_request_selected_during_guardrail_runs_concurrently_with_tool(monkeypatch, selected):
from litellm.responses.mcp.request_context import MCPRequestContext
from litellm.proxy._experimental.mcp_server import tool_registry
from litellm.responses.mcp.request_context import MCPRequestContext
tool_started = asyncio.Event()
guardrail_started = asyncio.Event()
@ -16434,6 +16422,7 @@ async def test_server_response_identifies_read_only_config(in_config, in_db, exp
@pytest.mark.parametrize("with_caller,legacy_factory", [(True, False), (False, False), (True, True)])
async def test_client_sampling_does_not_fill_explicit_context_from_another_ambient_caller(with_caller, legacy_factory):
from mcp.server.auth.middleware.auth_context import auth_context_var
from litellm.proxy._experimental.mcp_server import server as legacy_server
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_sampling_callback
@ -16546,6 +16535,40 @@ class TestSharedIdentifierPrefixWarning:
assert "'shared'" in shared_warnings[0]
def test_stored_client_assertion_signing_alg_falls_back_to_rs256_for_non_approved(caplog):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_stored_client_assertion_signing_alg,
)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
assert _stored_client_assertion_signing_alg("HS256", "srv") == "RS256"
assert "client_assertion_signing_alg" in caplog.text and "HS256" in caplog.text and "srv" in caplog.text
def test_stored_client_assertion_signing_alg_passes_through_approved(caplog):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_stored_client_assertion_signing_alg,
)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
assert _stored_client_assertion_signing_alg("PS384", "srv") == "PS384"
assert _stored_client_assertion_signing_alg(None, "srv") == "RS256"
assert "approved algorithm" not in caplog.text
def test_mcp_server_model_rejects_non_approved_client_assertion_signing_alg():
from pydantic import ValidationError
with pytest.raises(ValidationError) as exc:
MCPServer(
server_id="srv",
name="srv",
transport="http",
client_assertion_signing_alg="EdDSA",
)
assert "client_assertion_signing_alg" in str(exc.value)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"flag,transports,expected_warnings",
@ -17244,8 +17267,11 @@ async def test_upstream_preparation_honors_case_sensitive_extra_command(monkeypa
monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true")
monkeypatch.setattr(upstream, "MCP_STDIO_ALLOWED_COMMANDS", frozenset({"CustomRunner"}))
server: Final = MCPServer(
server_id="custom-stdio", name="custom-stdio", transport=MCPTransport.stdio,
command="/opt/tools/CustomRunner", args=[],
server_id="custom-stdio",
name="custom-stdio",
transport=MCPTransport.stdio,
command="/opt/tools/CustomRunner",
args=[],
)
client: Final = await MCPServerManager()._create_mcp_client(server)
assert client.stdio_config is not None

View file

@ -665,3 +665,35 @@ async def test_audit_matching_login_without_nonce_does_not_report_failure(caplog
)
assert result is None
assert "oauth_identity_binding audit" not in caplog.text
def test_select_signing_key_rejects_kid_matching_key_with_non_approved_alg() -> None:
import json
from cryptography.hazmat.primitives.asymmetric import ed25519
from litellm.proxy._experimental.mcp_server.oauth_identity_binding import _BindingRejection
okp_jwk: Final = {
**json.loads(jwt.algorithms.OKPAlgorithm.to_jwk(ed25519.Ed25519PrivateKey.generate().public_key())),
"kid": KID,
"alg": "EdDSA",
}
result: Final = _select_signing_key(_sign_id_token({}), [okp_jwk])
assert isinstance(result, _BindingRejection)
def test_select_signing_key_picks_approved_key_when_kid_is_shared() -> None:
import json
from cryptography.hazmat.primitives.asymmetric import ed25519
okp_jwk: Final = {
**json.loads(jwt.algorithms.OKPAlgorithm.to_jwk(ed25519.Ed25519PrivateKey.generate().public_key())),
"kid": KID,
"alg": "EdDSA",
}
hs_jwk: Final = {"kty": "oct", "kid": KID, "alg": "HS256", "k": "c2VjcmV0"}
result: Final = _select_signing_key(_sign_id_token({}), [okp_jwk, hs_jwk, _PUBLIC_JWK])
assert isinstance(result, jwt.PyJWK)
assert result.key_id == KID

File diff suppressed because it is too large Load diff

View file

@ -1292,3 +1292,101 @@ async def test_inject_mcp_jwt_signs_for_tool_call_path():
scopes = set(decoded["scope"].split())
assert "mcp:tools/call" in scopes
assert "mcp:tools/search_web:call" in scopes
# ---------------------------------------------------------------------------
# B5: incoming-JWT JWKS allowlist (LIT-8429)
# ---------------------------------------------------------------------------
def _okp_idp_key_and_token(now: int):
import json
from cryptography.hazmat.primitives.asymmetric import ed25519
private_key = ed25519.Ed25519PrivateKey.generate()
jwk = {
**json.loads(jwt.algorithms.OKPAlgorithm.to_jwk(private_key.public_key())),
"kid": "ed",
"alg": "EdDSA",
}
token = jwt.encode(
{"sub": "idp-user", "iat": now, "exp": now + 300},
private_key,
algorithm="EdDSA",
headers={"kid": "ed"},
)
return jwk, token
def _oct_idp_key_and_token(now: int):
secret = b"integration-hs256-client-secret-0123456789abcdef"
jwk = {
"kty": "oct",
"kid": "sym",
"alg": "HS256",
"k": base64.urlsafe_b64encode(secret).rstrip(b"=").decode(),
}
token = jwt.encode(
{"sub": "idp-user", "iat": now, "exp": now + 300},
secret,
algorithm="HS256",
headers={"kid": "sym"},
)
return jwk, token
@pytest.mark.asyncio
@pytest.mark.parametrize(
"key_and_token", (_oct_idp_key_and_token, _okp_idp_key_and_token), ids=("oct-HS256", "OKP-EdDSA")
)
async def test_verify_incoming_jwt_rejects_jwks_without_approved_algorithms(key_and_token):
"""A JWKS whose only keys use non-approved algorithms can never verify the incoming token."""
jwks_key, incoming_token = key_and_token(int(time.time()))
signer = _make_signer(
access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration",
)
with patch.object(
signer,
"_get_oidc_discovery",
new_callable=AsyncMock,
return_value={"jwks_uri": "https://idp.example.com/jwks"},
):
with patch(
"litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer._fetch_jwks",
new_callable=AsyncMock,
return_value=[jwks_key],
):
with pytest.raises(jwt.exceptions.PyJWKSetError, match="approved signing algorithm"):
await signer._verify_incoming_jwt(incoming_token)
@pytest.mark.asyncio
async def test_verify_incoming_jwt_ignores_filtered_keys_sharing_kid():
"""An approved RS256 key still verifies when the JWKS also carries a non-approved key with the same kid."""
signer = _make_signer(
access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration",
)
now = int(time.time())
oct_jwk, _ = _oct_idp_key_and_token(now)
oct_jwk["kid"] = signer._kid
incoming_token = jwt.encode(
{"sub": "idp-user", "iat": now, "exp": now + 300},
signer._private_key,
algorithm="RS256",
headers={"kid": signer._kid},
)
jwks = [oct_jwk, *signer.get_jwks()["keys"]]
with patch.object(
signer,
"_get_oidc_discovery",
new_callable=AsyncMock,
return_value={"jwks_uri": "https://idp.example.com/jwks"},
):
with patch(
"litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer._fetch_jwks",
new_callable=AsyncMock,
return_value=jwks,
):
payload = await signer._verify_incoming_jwt(incoming_token)
assert payload["sub"] == "idp-user"

View file

@ -387,6 +387,38 @@ def test_change_password_request_passwords_hidden_from_repr():
for rendered in (repr(request), str(request)):
assert "hunter2hunter2" not in rendered
assert "NewP@ssw0rd-2026" not in rendered
def test_mcp_server_requests_reject_non_approved_client_assertion_signing_alg() -> None:
"""The REST boundary is strict even though the stored blob stays lenient:
a write carrying HS256/hs256/"" must 422 naming the field."""
from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest
for cls, base in (
(NewMCPServerRequest, {"transport": "http", "url": "https://mcp.example.com"}),
(UpdateMCPServerRequest, {"server_id": "srv-1", "transport": "http", "url": "https://mcp.example.com"}),
):
for alg in ("HS256", "hs256", "", "EdDSA"):
with pytest.raises(ValidationError) as exc:
cls(**base, credentials={"client_assertion_signing_alg": alg})
assert "client_assertion_signing_alg" in str(exc.value), f"{cls.__name__} accepted {alg!r}"
def test_mcp_server_requests_accept_approved_client_assertion_signing_alg() -> None:
from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest
for alg in ("ES256", "PS384", "RS256", None):
request = NewMCPServerRequest(
transport="http",
url="https://mcp.example.com",
credentials={"client_assertion_signing_alg": alg},
)
assert request.credentials is not None
assert request.credentials["client_assertion_signing_alg"] == alg
update = UpdateMCPServerRequest(server_id="srv-1")
assert update.credentials is None
@pytest.mark.parametrize("versions", [[], ["2099-01-01"], ["2026-07-28"]])
def test_mcp_advertised_versions_reject_unavailable_revisions(versions):
from pydantic import ValidationError
@ -401,7 +433,12 @@ def test_mcp_advertised_versions_reject_unavailable_revisions(versions):
def test_mcp_metadata_rejects_unavailable_upstream_protocol(revision):
from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest
payload = {"server_id": "test", "transport": "http", "url": "https://example.com/mcp", "mcp_info": {"protocol_version": revision}}
payload = {
"server_id": "test",
"transport": "http",
"url": "https://example.com/mcp",
"mcp_info": {"protocol_version": revision},
}
for model in (NewMCPServerRequest, UpdateMCPServerRequest):
with pytest.raises(ValidationError):
model.model_validate(payload)
@ -457,7 +494,10 @@ def test_an_enabled_stdio_mcp_server_still_needs_a_command_and_args(monkeypatch,
def test_an_http_mcp_server_is_unaffected_by_the_stdio_flag(monkeypatch, request_model):
monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False)
assert request_model(server_id="http-1", transport="http", url="https://mcp.example.com").url == "https://mcp.example.com"
assert (
request_model(server_id="http-1", transport="http", url="https://mcp.example.com").url
== "https://mcp.example.com"
)
with pytest.raises(ValidationError, match="url or spec_path is required"):
request_model(server_id="http-1", transport="http")
@ -470,9 +510,13 @@ def test_a_non_mapping_mcp_server_payload_gets_a_validation_error(request_model)
@pytest.mark.parametrize("request_model", MCP_SERVER_REQUESTS)
def test_modern_http_upstream_protocol_is_available(request_model):
parsed = request_model.model_validate({
"server_id": "modern", "transport": "http", "url": "https://example.com/mcp",
"mcp_info": {"protocol_version": "2026-07-28"},
})
parsed = request_model.model_validate(
{
"server_id": "modern",
"transport": "http",
"url": "https://example.com/mcp",
"mcp_info": {"protocol_version": "2026-07-28"},
}
)
assert parsed.mcp_info["protocol_version"] == "2026-07-28"
assert parsed.transport == "http"