mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #41364 from BerriAI/litellm_fix_mcp_auth_fail_closed_4501
fix(mcp): fail closed on missing upstream credentials
This commit is contained in:
commit
41410e9556
8 changed files with 723 additions and 58 deletions
|
|
@ -346,13 +346,17 @@ class MCPClient:
|
|||
self.update_auth_value(auth_value)
|
||||
|
||||
async def discovery_auth_fingerprint(self) -> str:
|
||||
return self._hash_discovery_auth(await self.prepare_request_auth())
|
||||
|
||||
async def prepare_request_auth(self) -> httpx.Request:
|
||||
"""Preview the authenticated request without sending it, closing the auth flow afterwards."""
|
||||
request: Final = httpx.Request("POST", self.server_url or "http://localhost/", headers=self._get_auth_headers())
|
||||
if self._resolved_auth is None:
|
||||
return self._hash_discovery_auth(request)
|
||||
return request
|
||||
flow: Final = self._resolved_auth.async_auth_flow(request)
|
||||
try:
|
||||
authenticated: Final = await flow.__anext__()
|
||||
return self._hash_discovery_auth(authenticated)
|
||||
return authenticated
|
||||
finally:
|
||||
await flow.aclose()
|
||||
|
||||
|
|
|
|||
|
|
@ -102,6 +102,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials import (
|
|||
UpstreamCredentialProvider,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
prepare_mcp_client,
|
||||
raise_public,
|
||||
raise_token_exchange_challenge,
|
||||
raise_user_oauth_challenge,
|
||||
|
|
@ -2804,6 +2805,8 @@ class MCPServerManager:
|
|||
headers=headers,
|
||||
server_label=server.name or server.server_name or server.alias or server.server_id,
|
||||
relays_upstream_auth=server.is_client_forwarded_token,
|
||||
auth_type=server.auth_type,
|
||||
upstream_token_header=server.upstream_token_header,
|
||||
)
|
||||
tool_func.__name__ = prefixed_tool_name
|
||||
tool_func.__doc__ = description
|
||||
|
|
@ -4259,15 +4262,20 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
extra_headers=extra_headers,
|
||||
)
|
||||
return MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
auth_type=resolved_server.auth_type,
|
||||
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
|
||||
extra_headers=extra_headers,
|
||||
resolved_auth=resolved_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
return await prepare_mcp_client(
|
||||
resolved_server,
|
||||
MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
auth_type=resolved_server.auth_type,
|
||||
timeout=(
|
||||
resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT
|
||||
),
|
||||
extra_headers=extra_headers,
|
||||
resolved_auth=resolved_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
),
|
||||
)
|
||||
|
||||
# Create SigV4 auth if configured
|
||||
|
|
@ -4297,17 +4305,20 @@ class MCPServerManager:
|
|||
else AuthResolution.no_auth
|
||||
)
|
||||
record_auth_resolution(server.server_id, legacy_source)
|
||||
return MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
auth_type=resolved_server.auth_type,
|
||||
auth_value=auth_value,
|
||||
auth_header_name=auth_header_name,
|
||||
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
|
||||
extra_headers=extra_headers,
|
||||
aws_auth=aws_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
return await prepare_mcp_client(
|
||||
resolved_server,
|
||||
MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
auth_type=resolved_server.auth_type,
|
||||
auth_value=auth_value,
|
||||
auth_header_name=auth_header_name,
|
||||
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
|
||||
extra_headers=extra_headers,
|
||||
aws_auth=aws_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
),
|
||||
)
|
||||
|
||||
async def _get_tools_from_server(
|
||||
|
|
|
|||
|
|
@ -54,7 +54,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
from litellm.types.mcp import credential_redirect_hook, custom_credential_slot
|
||||
from litellm.types.mcp import MCPAuthType, credential_redirect_hook, custom_credential_slot
|
||||
|
||||
|
||||
class _OpenAPIJSONSchema(TypedDict, total=False):
|
||||
|
|
@ -471,6 +471,8 @@ def create_tool_function(
|
|||
headers: dict[str, str] | None = None,
|
||||
server_label: str | None = None,
|
||||
relays_upstream_auth: bool = False,
|
||||
auth_type: MCPAuthType = None,
|
||||
upstream_token_header: str | None = None,
|
||||
):
|
||||
"""Create a tool function for an OpenAPI operation.
|
||||
|
||||
|
|
@ -503,6 +505,18 @@ def create_tool_function(
|
|||
by using **kwargs instead of named parameters.
|
||||
"""
|
||||
effective_headers: Final = _merge_openapi_tool_request_headers(headers)
|
||||
if auth_type is not None:
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_public,
|
||||
validate_static_credential,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok
|
||||
|
||||
match validate_static_credential(auth_type, effective_headers, upstream_token_header):
|
||||
case Error(error):
|
||||
raise_public(error)
|
||||
case Ok():
|
||||
pass
|
||||
|
||||
# Build URL from base_url and path
|
||||
url = base_url + path
|
||||
|
|
|
|||
|
|
@ -13,15 +13,17 @@ from __future__ import annotations
|
|||
|
||||
import base64
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Final, Literal, NoReturn
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import SecretStr
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.experimental_mcp_client.client import strip_auth_scheme, to_basic_credentials
|
||||
from litellm.experimental_mcp_client.client import MCPClient, strip_auth_scheme, to_basic_credentials
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import MCPServerURLCredentialsError
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import resolve_upstream_resource
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
ApiKeyConfig,
|
||||
|
|
@ -39,7 +41,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
|||
Subject,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth
|
||||
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth, MCPAuthType, MCPTransport
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
@ -79,7 +81,7 @@ def to_server_spec(server: MCPServer) -> ServerSpec | None:
|
|||
|
||||
BYOK is the per-user source of the ``api_key`` mode; its scheme rides on ``auth_type`` just
|
||||
like a shared key, but the value is per-user and not migrated yet, so a BYOK server defers
|
||||
to v1 regardless of ``auth_type`` (this guard is the seam the BYOK arm replaces later).
|
||||
to v1 for its static schemes. Declared OBO always stays with the exchange arm.
|
||||
|
||||
Dispatches on the declared ``auth_type``. The match is exhaustive over ``MCPAuthType`` with
|
||||
an ``assert_never`` tail, so a newly added auth mode fails the type gate here until it is
|
||||
|
|
@ -90,8 +92,8 @@ def to_server_spec(server: MCPServer) -> ServerSpec | None:
|
|||
modes ``true_passthrough`` / ``oauth_delegate`` (``PassthroughConfig``); delegated/passthrough
|
||||
oauth2 and SigV4 return None and stay on v1.
|
||||
"""
|
||||
if server.is_byok:
|
||||
return None # per-user BYOK source not migrated yet -> defer to v1 (any auth_type)
|
||||
if server.is_byok and server.auth_type != MCPAuth.oauth2_token_exchange:
|
||||
return None # per-user BYOK source not migrated yet -> defer to v1
|
||||
resource: Final = server.url or server.server_id
|
||||
auth_type: Final = server.auth_type
|
||||
match auth_type:
|
||||
|
|
@ -165,21 +167,9 @@ def _client_credentials_spec(server: MCPServer, resource: str) -> ServerSpec:
|
|||
)
|
||||
|
||||
|
||||
def _token_exchange_spec(server: MCPServer, resource: str) -> ServerSpec | None:
|
||||
"""Build a token_exchange (OBO) spec, or defer (None) when it is not OBO-configured.
|
||||
|
||||
An OBO server with ``client_id``/``client_secret`` is owned by the v2 arm even if the
|
||||
``token_exchange_endpoint``/``token_url`` is absent: a missing endpoint then fails closed (412) at
|
||||
the exchanger rather than silently deferring to v1 and connecting unauthenticated, since the
|
||||
gateway must not guess the IdP or fall back to a weaker source. Without client credentials there is
|
||||
nothing to own, so the server stays on v1 (parity-safe). ``profile`` selects the wire dialect
|
||||
(``rfc8693`` default, ``entra_obo`` for Microsoft Entra On-Behalf-Of); an unrecognized value
|
||||
normalizes to ``rfc8693`` so a bad config value cannot crash spec-building. ``audience`` is
|
||||
forwarded only when the operator set it; a missing one is omitted, not derived.
|
||||
"""
|
||||
def _token_exchange_spec(server: MCPServer, resource: str) -> ServerSpec:
|
||||
"""Keep declared OBO owned by the resolver, including incomplete client configuration."""
|
||||
endpoint: Final = server.token_exchange_endpoint or server.effective_token_url
|
||||
if not server.client_id or not server.client_secret:
|
||||
return None
|
||||
profile: Final[Literal["rfc8693", "entra_obo"]] = (
|
||||
"entra_obo" if server.token_exchange_profile == "entra_obo" else "rfc8693"
|
||||
)
|
||||
|
|
@ -193,7 +183,7 @@ def _token_exchange_spec(server: MCPServer, resource: str) -> ServerSpec | None:
|
|||
token_exchange_endpoint=endpoint,
|
||||
audience=server.audience,
|
||||
client_id=server.client_id,
|
||||
client_secret=SecretStr(server.client_secret),
|
||||
client_secret=SecretStr(server.client_secret) if server.client_secret else None,
|
||||
token_endpoint_auth_method=server.token_endpoint_auth_method,
|
||||
scopes=tuple(server.scopes or ()),
|
||||
),
|
||||
|
|
@ -397,3 +387,69 @@ def raise_token_exchange_challenge(
|
|||
detail="Unauthorized",
|
||||
headers={"WWW-Authenticate": www_authenticate},
|
||||
)
|
||||
|
||||
|
||||
_STATIC_MODES: Final = frozenset(
|
||||
(MCPAuth.api_key, MCPAuth.bearer_token, MCPAuth.basic, MCPAuth.token, MCPAuth.authorization)
|
||||
)
|
||||
|
||||
|
||||
def _usable_credential_value(auth_type: MCPAuthType, name: str, value: str) -> bool:
|
||||
if not value:
|
||||
return False
|
||||
if auth_type == MCPAuth.api_key and name != "authorization":
|
||||
return True
|
||||
if value.lower() in ("bearer", "basic", "token", "apikey"):
|
||||
return False
|
||||
if auth_type == MCPAuth.api_key:
|
||||
api_scheme: Final = value.split(None, 1)[0]
|
||||
if api_scheme.lower() in ("bearer", "token", "apikey"):
|
||||
api_credential: Final = strip_auth_scheme(value, api_scheme).strip()
|
||||
return api_credential.lower() != api_scheme.lower()
|
||||
if auth_type in (MCPAuth.bearer_token, MCPAuth.token):
|
||||
scheme: Final = "Bearer" if auth_type == MCPAuth.bearer_token else "token"
|
||||
credential: Final = strip_auth_scheme(value, scheme).strip()
|
||||
return bool(credential) and credential.lower() != scheme.lower()
|
||||
if auth_type == MCPAuth.basic:
|
||||
parts: Final = value.split(None, 1)
|
||||
if len(parts) != 2 or parts[0].lower() != "basic":
|
||||
return False
|
||||
try:
|
||||
decoded: Final = base64.b64decode(parts[1], validate=True).strip()
|
||||
return b":" in decoded
|
||||
except ValueError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def validate_static_credential(
|
||||
auth_type: MCPAuthType,
|
||||
headers: Mapping[str, str],
|
||||
upstream_token_header: str | None = None,
|
||||
) -> Result[None, CredError]:
|
||||
if auth_type not in _STATIC_MODES:
|
||||
return Ok(None)
|
||||
default_slot: Final = "X-API-Key" if auth_type == MCPAuth.api_key else "Authorization"
|
||||
slots: Final = frozenset(
|
||||
name.lower()
|
||||
for name in (
|
||||
upstream_token_header or default_slot,
|
||||
default_slot,
|
||||
"Authorization",
|
||||
)
|
||||
)
|
||||
values: Final = tuple((name.lower(), value.strip()) for name, value in headers.items() if name.lower() in slots)
|
||||
if any(_usable_credential_value(auth_type, name, value) for name, value in values):
|
||||
return Ok(None)
|
||||
return Error(CredError.of_misconfigured(f"{auth_type} requires a usable upstream credential"))
|
||||
|
||||
|
||||
async def prepare_mcp_client(server: MCPServer, client: MCPClient) -> MCPClient:
|
||||
if server.auth_type not in _STATIC_MODES or client.transport_type == MCPTransport.stdio:
|
||||
return client
|
||||
request: Final = await client.prepare_request_auth()
|
||||
match validate_static_credential(server.auth_type, request.headers, server.upstream_token_header):
|
||||
case Error(error):
|
||||
raise_public(error)
|
||||
case Ok():
|
||||
return client
|
||||
|
|
|
|||
|
|
@ -1934,3 +1934,18 @@ async def test_discovery_auth_fingerprint_tracks_effective_credentials(resolved:
|
|||
assert original != replaced
|
||||
assert len(original) == 64
|
||||
assert "private-original-credential" not in original
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_auth_preview_uses_the_same_effective_headers_as_egress() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
|
||||
|
||||
client: Final = MCPClient(
|
||||
server_url="https://upstream.example/mcp", auth_type=MCPAuth.bearer_token,
|
||||
resolved_auth=StaticHeaderAuth("Bearer resolved"), extra_headers={"X-Trace": "trace"},
|
||||
)
|
||||
request: Final = await client.prepare_request_auth()
|
||||
assert request.method == "POST"
|
||||
assert str(request.url) == "https://upstream.example/mcp"
|
||||
assert request.headers["Authorization"] == "Bearer resolved"
|
||||
assert request.headers["X-Trace"] == "trace"
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ maps each CredError onto its HTTP status. These pin the parity-critical mapping
|
|||
|
||||
import base64
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -20,7 +21,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import
|
|||
raise_user_oauth_challenge,
|
||||
to_server_spec,
|
||||
to_subject,
|
||||
validate_static_credential,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
ApiKeyConfig,
|
||||
AuthorizationCodeConfig,
|
||||
|
|
@ -34,10 +37,28 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
|||
SharedKey,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp import MCPAuth, MCPAuthType, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
@pytest.mark.parametrize("auth_type,header,value", [
|
||||
(MCPAuth.api_key, "Authorization", "Bearer fixture-key"),
|
||||
(MCPAuth.api_key, "Authorization", "ApiKey fixture-key"),
|
||||
(MCPAuth.api_key, "Authorization", "token fixture-key"),
|
||||
(MCPAuth.api_key, "Authorization", "Bearer token"),
|
||||
(MCPAuth.api_key, "Authorization", "opaque-key"),
|
||||
(MCPAuth.api_key, "Authorization", "Custom Custom"),
|
||||
(MCPAuth.api_key, "X-API-Key", "Bearer Bearer"),
|
||||
(MCPAuth.api_key, "X-Custom", "ApiKey ApiKey"),
|
||||
(MCPAuth.authorization, "Authorization", "opaque-secret-value"),
|
||||
])
|
||||
def test_static_credential_preserves_supported_api_key_and_raw_headers(
|
||||
auth_type: MCPAuthType, header: str, value: str,
|
||||
) -> None:
|
||||
result: Final = validate_static_credential(auth_type, {header: value}, upstream_token_header=header)
|
||||
assert isinstance(result, Ok)
|
||||
|
||||
|
||||
def _server(**kwargs) -> MCPServer:
|
||||
return MCPServer(server_id="s", name="n", transport=MCPTransport.http, **kwargs)
|
||||
|
||||
|
|
@ -155,12 +176,6 @@ def test_oauth2_user_token_maps_to_authorization_code(oauth2_flow):
|
|||
_server(auth_type=MCPAuth.api_key), # no token configured
|
||||
_server(auth_type=MCPAuth.bearer_token), # no token configured
|
||||
_server(auth_type=MCPAuth.oauth2, delegate_auth_to_upstream=True), # delegated upstream OAuth -> v1
|
||||
_server(auth_type=MCPAuth.oauth2_token_exchange), # no endpoint/client creds -> incomplete -> v1
|
||||
_server(
|
||||
auth_type=MCPAuth.oauth2_token_exchange,
|
||||
token_exchange_endpoint="https://idp/token",
|
||||
client_id="cid",
|
||||
), # missing client_secret -> incomplete -> v1
|
||||
_server(auth_type=MCPAuth.aws_sigv4),
|
||||
_server(auth_type=None, oauth_passthrough=True, extra_headers=["Authorization"]),
|
||||
],
|
||||
|
|
@ -802,3 +817,14 @@ def test_a_blank_header_name_means_unset_rather_than_an_error(blank):
|
|||
spec = to_server_spec(server)
|
||||
assert spec is not None
|
||||
assert spec.config.header_name == "Authorization"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client_secret", [None, ""])
|
||||
@pytest.mark.parametrize("is_byok", [False, True])
|
||||
def test_incomplete_obo_keeps_exchange_ownership(client_secret: str | None, is_byok: bool) -> None:
|
||||
spec = to_server_spec(_server(auth_type=MCPAuth.oauth2_token_exchange, client_id="client",
|
||||
client_secret=client_secret, is_byok=is_byok))
|
||||
assert spec is not None
|
||||
assert isinstance(spec.config, TokenExchangeConfig)
|
||||
assert spec.config.client_id == "client"
|
||||
assert spec.config.client_secret is None
|
||||
|
|
|
|||
|
|
@ -5,11 +5,13 @@ import logging
|
|||
import os
|
||||
import sys
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Final, Literal, Optional
|
||||
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,
|
||||
|
|
@ -5127,7 +5129,8 @@ class TestMCPServerManager:
|
|||
captured: dict = {}
|
||||
|
||||
def fake_create_tool_function(
|
||||
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False
|
||||
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False,
|
||||
auth_type=None, upstream_token_header=None,
|
||||
):
|
||||
captured["headers"] = headers
|
||||
captured["server_label"] = server_label
|
||||
|
|
@ -5212,7 +5215,8 @@ class TestMCPServerManager:
|
|||
captured: dict = {}
|
||||
|
||||
def fake_create_tool_function(
|
||||
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False
|
||||
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False,
|
||||
auth_type=None, upstream_token_header=None,
|
||||
):
|
||||
captured["headers"] = headers
|
||||
|
||||
|
|
@ -9401,12 +9405,13 @@ class TestCreateMcpClientV2Graft:
|
|||
assert "misconfigured" in str(exc_info.value.detail)
|
||||
assert "token_url" in str(exc_info.value.detail)
|
||||
|
||||
async def test_static_token_missing_defers_to_v1(self):
|
||||
client = await MCPServerManager()._create_mcp_client(
|
||||
self._http_server(auth_type=MCPAuth.api_key, authentication_token=None)
|
||||
)
|
||||
|
||||
assert client._resolved_auth is None
|
||||
async def test_static_token_missing_rejects_before_connecting(self):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(
|
||||
self._http_server(auth_type=MCPAuth.api_key, authentication_token=None)
|
||||
)
|
||||
assert exc.value.status_code == 500
|
||||
assert "credential" in str(exc.value.detail)
|
||||
|
||||
async def test_stdio_migrated_auth_type_still_defers_to_v1(self):
|
||||
client = await MCPServerManager()._create_mcp_client(
|
||||
|
|
@ -13467,3 +13472,400 @@ async def test_discovery_cache_returns_oversized_results_without_retaining_them(
|
|||
result: Final = await cache.get(("server", None), fetch)
|
||||
assert result[0].description == description
|
||||
assert fetch.await_count == 2
|
||||
|
||||
|
||||
class TestProtectedCredentialPreparation:
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,credential", [
|
||||
(MCPAuth.bearer_token, None),
|
||||
(MCPAuth.bearer_token, "Bearer"),
|
||||
(MCPAuth.api_key, None),
|
||||
(MCPAuth.basic, "Basic"),
|
||||
])
|
||||
@pytest.mark.parametrize("dispatch", ["managed", "local"])
|
||||
async def test_openapi_dispatch_rejects_unusable_effective_credentials(
|
||||
self, tmp_path: Path, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
|
||||
auth_type: MCPAuthType, credential: str | None, dispatch: str,
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.server import _handle_local_mcp_tool
|
||||
from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name, get_server_prefix
|
||||
|
||||
spec_path: Final = tmp_path / "openapi.json"
|
||||
spec_path.write_text(json.dumps({"openapi": "3.0.0", "info": {"title": "Auth", "version": "1"},
|
||||
"paths": {"/echo": {"get": {"operationId": "echo"}}}}))
|
||||
server: Final = MCPServer(
|
||||
server_id="dispatch-auth", name="dispatch-auth", url="https://upstream.example",
|
||||
transport=MCPTransport.http, auth_type=auth_type, authentication_token=credential,
|
||||
)
|
||||
manager: Final = MCPServerManager()
|
||||
await manager._register_openapi_tools(str(spec_path), server, server.url)
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="unexpected success")
|
||||
result: Final = (
|
||||
await manager._call_openapi_tool_handler(server, "echo", {})
|
||||
if dispatch == "managed"
|
||||
else await _handle_local_mcp_tool(add_server_prefix_to_name("echo", get_server_prefix(server)), {})
|
||||
)
|
||||
assert result.isError is True
|
||||
assert "requires a usable upstream credential" in result.content[0].text
|
||||
assert destination.call_count == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("transport", [MCPTransport.http, MCPTransport.sse])
|
||||
@pytest.mark.parametrize("client_secret", [None, ""])
|
||||
@pytest.mark.parametrize("subject", [None, "caller-subject"])
|
||||
async def test_incomplete_obo_rejects_caller_and_static_fallback(
|
||||
self, transport: MCPTransport, client_secret: str | None, subject: str | None
|
||||
) -> None:
|
||||
server = MCPServer(
|
||||
server_id="incomplete-obo", name="incomplete-obo", url="https://upstream.example/mcp",
|
||||
transport=transport, auth_type=MCPAuth.oauth2_token_exchange,
|
||||
client_id="gateway", client_secret=client_secret,
|
||||
token_exchange_endpoint="https://idp.example/token", authentication_token="static-fallback",
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(
|
||||
server, mcp_auth_header="Bearer override", subject_token=subject,
|
||||
)
|
||||
assert exc.value.status_code == (401 if subject is None else 500)
|
||||
assert "static-fallback" not in str(exc.value.detail)
|
||||
assert "override" not in str(exc.value.detail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type", [MCPAuth.api_key, MCPAuth.bearer_token])
|
||||
@pytest.mark.parametrize("credential", [None, "", " ", {"X-Trace": "trace"}])
|
||||
async def test_static_auth_without_usable_credential_rejects(
|
||||
self, auth_type: MCPAuthType, credential: str | dict[str, str] | None
|
||||
) -> None:
|
||||
server = MCPServer(
|
||||
server_id="empty-static", name="empty-static", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type,
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header=credential)
|
||||
assert exc.value.status_code == 500
|
||||
assert "credential" in str(exc.value.detail).lower()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,headers", [
|
||||
(MCPAuth.api_key, {"X-API-Key": "key"}),
|
||||
(MCPAuth.bearer_token, {"Authorization": "Bearer token"}),
|
||||
])
|
||||
async def test_static_auth_accepts_actual_forwarded_credential(
|
||||
self, auth_type: MCPAuthType, headers: dict[str, str]
|
||||
) -> None:
|
||||
server = MCPServer(
|
||||
server_id="header-static", name="header-static", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type,
|
||||
)
|
||||
client = await MCPServerManager()._create_mcp_client(server, extra_headers=headers)
|
||||
assert client._get_auth_headers() == headers
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type", [MCPAuth.oauth2_token_exchange])
|
||||
async def test_openapi_protected_auth_rejects_missing_credentials(self, auth_type: MCPAuthType) -> None:
|
||||
server = MCPServer(
|
||||
server_id="openapi-empty", name="openapi-empty", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type,
|
||||
token_exchange_endpoint="https://idp.example/token",
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager().resolve_openapi_upstream_auth(
|
||||
mcp_server=server, oauth2_headers=None, raw_headers=None, mcp_auth_header=None,
|
||||
user_api_key_auth=None, forwarded_headers=None,
|
||||
)
|
||||
assert exc.value.status_code in (401, 500)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,slot,value", [
|
||||
(MCPAuth.api_key, "X-API-Key", "token"),
|
||||
(MCPAuth.authorization, "Authorization", "opaque-secret-value"),
|
||||
(MCPAuth.authorization, "Authorization", "Bearer abc"),
|
||||
(MCPAuth.authorization, "Authorization", "Custom abc"),
|
||||
])
|
||||
async def test_raw_static_credentials_are_forwarded_unchanged(
|
||||
self, auth_type: MCPAuthType, slot: str, value: str,
|
||||
) -> None:
|
||||
server = MCPServer(server_id="raw-key", name="raw-key", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type, authentication_token=value)
|
||||
client = await MCPServerManager()._create_mcp_client(server)
|
||||
assert client._resolved_auth is not None
|
||||
request = httpx.Request("GET", server.url)
|
||||
flow = client._resolved_auth.auth_flow(request)
|
||||
try:
|
||||
assert next(flow).headers[slot] == value
|
||||
finally:
|
||||
flow.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", ["Bearer", "basic", "token", "ApiKey", " bEaReR ", "\tTOKEN\t"])
|
||||
@pytest.mark.parametrize("source", ["configured", "caller", "forwarded"])
|
||||
async def test_raw_authorization_rejects_bare_schemes_before_dispatch(
|
||||
self, respx_mock: MockRouter, value: str, source: str,
|
||||
) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="raw-empty", name="raw-empty", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.authorization,
|
||||
authentication_token=value if source == "configured" else None,
|
||||
)
|
||||
destination: Final = respx_mock.route().respond(200)
|
||||
with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc:
|
||||
await MCPServerManager()._create_mcp_client(
|
||||
server, mcp_auth_header=value if source == "caller" else None,
|
||||
extra_headers={"Authorization": value} if source == "forwarded" else None,
|
||||
)
|
||||
assert exc.value.status_code == 500
|
||||
assert destination.call_count == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_byok_flag_cannot_bypass_incomplete_obo(self) -> None:
|
||||
server = MCPServer(server_id="obo-byok", name="obo-byok", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.oauth2_token_exchange, is_byok=True,
|
||||
token_exchange_endpoint="https://idp.example/token")
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header="Bearer override")
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("configured,override", [(None, "Bearer usable"), ("shared", "Bearer usable")])
|
||||
async def test_bearer_override_remains_usable(self, configured: str | None, override: str) -> None:
|
||||
server = MCPServer(server_id="override", name="override", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=configured)
|
||||
client = await MCPServerManager()._create_mcp_client(server, mcp_auth_header=override)
|
||||
assert client._get_auth_headers()["Authorization"] == override
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("token", [None, "shared"])
|
||||
async def test_empty_injected_header_cannot_satisfy_protected_auth(self, token: str | None) -> None:
|
||||
server = MCPServer(server_id="empty-header", name="empty-header", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=token)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, extra_headers={"authorization": " "})
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_slot_uses_its_actual_credential(self) -> None:
|
||||
server = MCPServer(server_id="custom", name="custom", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.api_key,
|
||||
upstream_token_header="X-Custom", authentication_token="key")
|
||||
client = await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Trace": "trace"})
|
||||
assert client._credential_slot == "X-Custom"
|
||||
assert await client.discovery_auth_fingerprint()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("static,forwarded,caller", [
|
||||
({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None),
|
||||
({}, {"X-API-Key": "forwarded"}, None),
|
||||
({}, None, "ApiKey caller"),
|
||||
({"X-API-Key": "static"}, {"Authorization": ""}, None),
|
||||
])
|
||||
async def test_openapi_static_credentials_remain_supported(
|
||||
self, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
|
||||
static: dict[str, str], forwarded: dict[str, str] | None, caller: str | None
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_auth_header, _request_extra_headers, create_tool_function,
|
||||
)
|
||||
tool: Final = create_tool_function(
|
||||
"/echo", "get", {}, "https://upstream.example", headers=static, auth_type=MCPAuth.api_key,
|
||||
)
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
|
||||
caller_token: Final = _request_auth_header.set(caller)
|
||||
extra_token: Final = _request_extra_headers.set(forwarded)
|
||||
try:
|
||||
assert await tool() == "authenticated"
|
||||
sent: Final = destination.calls.last.request.headers
|
||||
assert sent.get("x-api-key") == static.get("X-API-Key", (forwarded or {}).get("X-API-Key"))
|
||||
if caller:
|
||||
assert sent["authorization"] == caller
|
||||
assert destination.call_count == 1
|
||||
finally:
|
||||
_request_auth_header.reset(caller_token)
|
||||
_request_extra_headers.reset(extra_token)
|
||||
|
||||
@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
|
||||
|
||||
class CancelledAuth(httpx.Auth):
|
||||
closed = False
|
||||
|
||||
async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx.Request, httpx.Response]:
|
||||
try:
|
||||
raise asyncio.CancelledError()
|
||||
yield request
|
||||
finally:
|
||||
self.closed = True
|
||||
|
||||
auth = CancelledAuth()
|
||||
server = MCPServer(server_id="cancel", name="cancel", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.api_key)
|
||||
client = MCPClient(server_url=server.url, auth_type=MCPAuth.api_key, resolved_auth=auth)
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await prepare_mcp_client(server, client)
|
||||
assert auth.closed
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type", [MCPAuth.basic, MCPAuth.token, MCPAuth.authorization])
|
||||
async def test_other_static_schemes_reject_whitespace_credentials(self, auth_type: MCPAuthType) -> None:
|
||||
server = MCPServer(server_id="blank-static", name="blank-static", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type, authentication_token=" ")
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server)
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("header", ["Basic", "Basic @@@", "Other abc", "Basic QmFzaWM=", "Basic bm8tY29sb24="])
|
||||
async def test_basic_headers_without_usable_credentials_reject(self, header: str) -> None:
|
||||
server = MCPServer(server_id="bad-basic", name="bad-basic", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.basic)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, extra_headers={"Authorization": header})
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", ["Basic", "Basic ", "basic"])
|
||||
@pytest.mark.parametrize("source", ["configured", "caller"])
|
||||
async def test_basic_scheme_alone_is_not_a_credential(self, value: str, source: str) -> None:
|
||||
server = MCPServer(server_id="basic-scheme", name="basic-scheme", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.basic,
|
||||
authentication_token=value if source == "configured" else None)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None)
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,value,default_slot", [
|
||||
(MCPAuth.api_key, "fixture-key", "X-API-Key"),
|
||||
(MCPAuth.bearer_token, "fixture-key", "Authorization"),
|
||||
(MCPAuth.basic, "user:pass", "Authorization"),
|
||||
(MCPAuth.token, "fixture-key", "Authorization"),
|
||||
(MCPAuth.authorization, "fixture-key", "Authorization"),
|
||||
])
|
||||
@pytest.mark.parametrize("source", ["configured", "caller"])
|
||||
async def test_usable_credential_survives_an_empty_alternate_header(
|
||||
self, auth_type: MCPAuthType, value: str, default_slot: str, source: str
|
||||
) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="alternate", name="alternate", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type, upstream_token_header="X-Custom",
|
||||
authentication_token=value if source == "configured" else None,
|
||||
)
|
||||
empty_slot: Final = default_slot if source == "configured" else "X-Custom"
|
||||
selected_slot: Final = "X-Custom" if source == "configured" else default_slot
|
||||
client: Final = await MCPServerManager()._create_mcp_client(
|
||||
server, mcp_auth_header=value if source == "caller" else None, extra_headers={empty_slot: ""},
|
||||
)
|
||||
request: Final = await client.prepare_request_auth()
|
||||
assert request.headers[selected_slot]
|
||||
assert request.headers[empty_slot] == ""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_custom_and_default_headers_do_not_satisfy_auth(self) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="both-empty", name="both-empty", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header="X-Custom",
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Custom": "", "X-API-Key": ""})
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("custom_slot", [None, "X-Custom"])
|
||||
@pytest.mark.parametrize("source", ["caller", "forwarded"])
|
||||
async def test_api_key_preserves_explicit_authorization_credential(
|
||||
self, custom_slot: str | None, source: str
|
||||
) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="caller-auth", name="caller-auth", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header=custom_slot,
|
||||
)
|
||||
headers: Final = {"Authorization": "Bearer caller-credential", "X-API-Key": ""}
|
||||
client: Final = await MCPServerManager()._create_mcp_client(
|
||||
server, mcp_auth_header=headers if source == "caller" else None,
|
||||
extra_headers=headers if source == "forwarded" else None,
|
||||
)
|
||||
request: Final = await client.prepare_request_auth()
|
||||
assert request.headers["Authorization"] == "Bearer caller-credential"
|
||||
assert request.headers["X-API-Key"] == ""
|
||||
assert custom_slot is None or custom_slot not in request.headers
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", [
|
||||
"", " ", "Bearer", "Basic", "token", "ApiKey",
|
||||
"Bearer Bearer", "ApiKey ApiKey", "token token", "bEaReR BEARER", "aPiKeY\tAPIKEY",
|
||||
])
|
||||
async def test_api_key_rejects_authorization_without_a_credential(self, value: str) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="caller-empty", name="caller-empty", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.api_key,
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header={"Authorization": value})
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", ["no-colon", "Basic bm8tY29sb24="])
|
||||
@pytest.mark.parametrize("source", ["configured", "caller"])
|
||||
async def test_basic_requires_a_username_password_separator(self, value: str, source: str) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="basic-pair", name="basic-pair", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.basic,
|
||||
authentication_token=value if source == "configured" else None,
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None)
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", ["user:pass", "user:", ":pass", ":"])
|
||||
async def test_basic_preserves_username_password_pairs(self, value: str) -> None:
|
||||
import base64
|
||||
|
||||
server: Final = MCPServer(
|
||||
server_id="basic-valid", name="basic-valid", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.basic, authentication_token=value,
|
||||
)
|
||||
client: Final = await MCPServerManager()._create_mcp_client(server)
|
||||
request: Final = await client.prepare_request_auth()
|
||||
scheme, encoded = request.headers["Authorization"].split(" ", 1)
|
||||
assert scheme == "Basic"
|
||||
assert base64.b64decode(encoded) == value.encode()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,value", [
|
||||
(MCPAuth.bearer_token, "Bearer"), (MCPAuth.bearer_token, "Bearer "), (MCPAuth.bearer_token, "bearer"),
|
||||
(MCPAuth.token, "token"), (MCPAuth.token, "token "), (MCPAuth.token, "TOKEN"),
|
||||
])
|
||||
@pytest.mark.parametrize("source", ["configured", "caller"])
|
||||
async def test_static_scheme_only_input_cannot_hide_behind_rendered_prefix(
|
||||
self, auth_type: MCPAuthType, value: str, source: str
|
||||
) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="empty-scheme", name="empty-scheme", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type,
|
||||
authentication_token=value if source == "configured" else None,
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None)
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,value,expected", [
|
||||
(MCPAuth.bearer_token, "token", "Bearer token"),
|
||||
(MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"),
|
||||
(MCPAuth.token, "tokenish", "token tokenish"),
|
||||
])
|
||||
async def test_static_credentials_that_resemble_schemes_remain_usable(
|
||||
self, auth_type: MCPAuthType, value: str, expected: str
|
||||
) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="real-token", name="real-token", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type, authentication_token=value,
|
||||
)
|
||||
client: Final = await MCPServerManager()._create_mcp_client(server)
|
||||
request: Final = await client.prepare_request_auth()
|
||||
assert request.headers["Authorization"] == expected
|
||||
|
|
|
|||
|
|
@ -10,9 +10,14 @@ This test suite ensures that:
|
|||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from respx import MockRouter
|
||||
|
||||
from litellm.types.mcp import MCPAuth, MCPAuthType
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_auth_header,
|
||||
|
|
@ -35,6 +40,120 @@ from litellm.proxy._experimental.mcp_server.exceptions import (
|
|||
GET_ASYNC_CLIENT_TARGET = "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.get_async_httpx_client"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,value,accepted", [
|
||||
(MCPAuth.api_key, "Bearer Bearer", False), (MCPAuth.api_key, "ApiKey ApiKey", False),
|
||||
(MCPAuth.api_key, "token token", False), (MCPAuth.api_key, "bEaReR BEARER", False),
|
||||
(MCPAuth.api_key, "aPiKeY\tAPIKEY", False), (MCPAuth.api_key, "Bearer fixture-key", True),
|
||||
(MCPAuth.api_key, "ApiKey fixture-key", True), (MCPAuth.api_key, "token fixture-key", True),
|
||||
(MCPAuth.authorization, "Bearer", False), (MCPAuth.authorization, "basic", False),
|
||||
(MCPAuth.authorization, "token", False), (MCPAuth.authorization, "ApiKey", False),
|
||||
(MCPAuth.authorization, " bEaReR ", False), (MCPAuth.authorization, "\tTOKEN\t", False),
|
||||
(MCPAuth.authorization, "opaque-secret-value", True), (MCPAuth.authorization, "Bearer abc", True),
|
||||
(MCPAuth.authorization, "Custom abc", True),
|
||||
])
|
||||
async def test_authorization_validates_credentials_before_http(
|
||||
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, auth_type: MCPAuthType, value: str, accepted: bool,
|
||||
) -> None:
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
tool: Final = create_tool_function(
|
||||
"/echo", "get", {}, "https://upstream.example", auth_type=auth_type,
|
||||
)
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
|
||||
caller_token: Final = _request_auth_header.set(value)
|
||||
try:
|
||||
if accepted:
|
||||
assert await tool() == "authenticated"
|
||||
assert destination.call_count == 1
|
||||
assert destination.calls.last.request.headers["authorization"] == value
|
||||
else:
|
||||
with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc:
|
||||
await tool()
|
||||
assert exc.value.status_code == 500
|
||||
assert destination.call_count == 0
|
||||
finally:
|
||||
_request_auth_header.reset(caller_token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("static,forwarded,caller,resolved,expected", [
|
||||
({"Authorization": "Bearer configured"}, {"authorization": "Bearer forwarded"}, None, None, "Bearer configured"),
|
||||
({"Authorization": "Bearer configured"}, None, "Bearer caller", None, "Bearer caller"),
|
||||
({"Authorization": "Bearer configured"}, None, "Bearer", None, None),
|
||||
({"Authorization": "Bearer configured"}, None, "Bearer caller", {"authorization": " "}, None),
|
||||
({"Authorization": "Bearer configured"}, None, "Bearer", {"authorization": "Bearer resolved"}, "Bearer resolved"),
|
||||
])
|
||||
async def test_static_auth_validates_headers_after_existing_precedence(
|
||||
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
|
||||
static: dict[str, str], forwarded: dict[str, str] | None, caller: str | None,
|
||||
resolved: dict[str, str] | None, expected: str | None,
|
||||
) -> None:
|
||||
tool: Final = create_tool_function(
|
||||
"/echo", "get", {}, "https://upstream.example", headers=static, auth_type=MCPAuth.bearer_token,
|
||||
)
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
|
||||
caller_token: Final = _request_auth_header.set(caller)
|
||||
extra_token: Final = _request_extra_headers.set(forwarded)
|
||||
resolved_token: Final = _request_resolved_auth_headers.set(resolved)
|
||||
try:
|
||||
if expected is None:
|
||||
with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc:
|
||||
await tool()
|
||||
assert exc.value.status_code == 500
|
||||
assert destination.call_count == 0
|
||||
else:
|
||||
assert await tool() == "authenticated"
|
||||
assert destination.call_count == 1
|
||||
assert destination.calls.last.request.headers["authorization"] == expected
|
||||
finally:
|
||||
_request_auth_header.reset(caller_token)
|
||||
_request_extra_headers.reset(extra_token)
|
||||
_request_resolved_auth_headers.reset(resolved_token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("credential", ["custom-key", ""])
|
||||
async def test_static_auth_uses_configured_custom_header(
|
||||
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, credential: str,
|
||||
) -> None:
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
tool: Final = create_tool_function(
|
||||
"/echo", "get", {}, "https://upstream.example", headers={"x-custom": credential},
|
||||
auth_type=MCPAuth.api_key, upstream_token_header="X-Custom",
|
||||
)
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
|
||||
if credential:
|
||||
assert await tool() == "authenticated"
|
||||
assert destination.call_count == 1
|
||||
assert destination.calls.last.request.headers["x-custom"] == credential
|
||||
else:
|
||||
with pytest.raises(HTTPException, match="requires a usable upstream credential"):
|
||||
await tool()
|
||||
assert destination.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,resolved", [
|
||||
(MCPAuth.none, None),
|
||||
(MCPAuth.oauth2, {"Authorization": "Bearer user-oauth"}),
|
||||
])
|
||||
async def test_static_validation_preserves_no_auth_and_resolved_oauth(
|
||||
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
|
||||
auth_type: MCPAuthType, resolved: dict[str, str] | None,
|
||||
) -> None:
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
tool: Final = create_tool_function("/echo", "get", {}, "https://upstream.example", auth_type=auth_type)
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="echo")
|
||||
token: Final = _request_resolved_auth_headers.set(resolved)
|
||||
try:
|
||||
assert await tool() == "echo"
|
||||
assert destination.call_count == 1
|
||||
assert destination.calls.last.request.headers.get("authorization") == (resolved or {}).get("Authorization")
|
||||
finally:
|
||||
_request_resolved_auth_headers.reset(token)
|
||||
|
||||
|
||||
def _create_mock_client(method: str, response_text: str, status_code: int = 200) -> AsyncMock:
|
||||
"""Utility to create a mocked async httpx client for the given method.
|
||||
|
||||
|
|
@ -1458,3 +1577,21 @@ class TestBoundedOpenAPISpecLoading:
|
|||
else:
|
||||
assert await load_openapi_spec_async("https://93.184.216.34/spec.json", max_bytes=100) == {"paths": {}}
|
||||
assert destination.call_count == 1
|
||||
|
||||
|
||||
def test_openapi_generator_import_does_not_require_mcp_sdk() -> None:
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
script = """
|
||||
import builtins
|
||||
original_import = builtins.__import__
|
||||
def without_mcp(name, *args, **kwargs):
|
||||
if name == 'mcp' or name.startswith('mcp.'):
|
||||
raise ModuleNotFoundError('MCP SDK unavailable')
|
||||
return original_import(name, *args, **kwargs)
|
||||
builtins.__import__ = without_mcp
|
||||
import litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator
|
||||
"""
|
||||
result = subprocess.run([sys.executable, "-c", script], capture_output=True, text=True)
|
||||
assert result.returncode == 0, result.stderr
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue