mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 4713920478 into 6e75f28289
This commit is contained in:
commit
eecac2fd3f
19 changed files with 2165 additions and 611 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
35
litellm/proxy/auth/jwt_algorithms.py
Normal file
35
litellm/proxy/auth/jwt_algorithms.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
17
litellm/types/proxy/auth/jwt_algorithms.py
Normal file
17
litellm/types/proxy/auth/jwt_algorithms.py
Normal 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",)
|
||||
721
tests/integration/authorization/test_jwt_algorithm_allowlist.py
Normal file
721
tests/integration/authorization/test_jwt_algorithm_allowlist.py
Normal 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
|
||||
277
tests/integration/mcp/test_mcp_client_assertion_signing_alg.py
Normal file
277
tests/integration/mcp/test_mcp_client_assertion_signing_alg.py
Normal 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
|
||||
291
tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py
Normal file
291
tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py
Normal 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]
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue