test(mcp): update MCP suites for SDK2 handler signatures and ctx var

Call handlers with ServerRequestContext and params models, seed the
litellm contextvar instead of the removed SDK request_ctx, forward
headers/auth through the httpx2 MockTransport factory, and add
regressions for handler registration, context propagation, and modern
protocol-version rejection.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
joshua 2026-09-18 23:34:30 +00:00
parent 783038010b
commit 0d2963fe89
11 changed files with 390 additions and 151 deletions

View file

@ -1,29 +1,51 @@
import os
import pytest
import asyncio
import os
import subprocess
import sys
from pathlib import Path
from typing import Optional
from unittest.mock import AsyncMock, patch
import pytest
from mcp.types import CallToolResult, TextContent
from mcp.types import Tool as MCPTool
import litellm
from litellm.types.utils import StandardLoggingPayload
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
from litellm.proxy._experimental.mcp_server.server import (
mcp_server_tool_call,
set_auth_context,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
from litellm.types.mcp import MCPPostCallResponseObject
from litellm.types.utils import HiddenParams
from mcp.types import Tool as MCPTool, CallToolResult, TextContent
def _mcp_request_ctx(**overrides):
from types import SimpleNamespace
from mcp.server.context import ServerRequestContext
kwargs = {
"session": SimpleNamespace(),
"lifespan_context": {},
"protocol_version": "2025-06-18",
"method": "",
"params": None,
"request_id": 1,
"meta": None,
"request": None,
}
kwargs.update(overrides)
return ServerRequestContext(**kwargs)
def _call_tool_params(name, arguments=None):
from mcp.types import CallToolRequestParams
return CallToolRequestParams(name=name, arguments=arguments)
class TestMCPLogger(CustomLogger):
def __init__(self):
self.standard_logging_payload = None
@ -142,8 +164,8 @@ async def test_mcp_cost_tracking():
# Call mcp tool
response = await mcp_server_tool_call(
name="zapier_gmail_server-add_tools", # Use correct prefixed name with - separator
arguments={"test": "test"},
_mcp_request_ctx(),
_call_tool_params("zapier_gmail_server-add_tools", {"test": "test"}),
)
# wait 1-2 seconds for logging to be processed
@ -285,8 +307,8 @@ async def test_mcp_cost_tracking_per_tool():
# Test 1: Call expensive_tool - should cost 5.0
response1 = await mcp_server_tool_call(
name="test_server-expensive_tool", # Use correct prefixed name with - separator
arguments={"data": "test_expensive"},
_mcp_request_ctx(),
_call_tool_params("test_server-expensive_tool", {"data": "test_expensive"}),
)
# wait for logging to be processed
@ -313,8 +335,8 @@ async def test_mcp_cost_tracking_per_tool():
# Test 2: Call cheap_tool - should cost 0.1
response2 = await mcp_server_tool_call(
name="test_server-cheap_tool", # Use correct prefixed name with - separator
arguments={"data": "test_cheap"},
_mcp_request_ctx(),
_call_tool_params("test_server-cheap_tool", {"data": "test_cheap"}),
)
# wait for logging to be processed
@ -356,7 +378,7 @@ async def test_mcp_cost_tracking_per_tool():
class MCPLoggerHook(TestMCPLogger):
async def async_post_mcp_tool_call_hook(
self, kwargs, response_obj: MCPPostCallResponseObject, start_time, end_time
) -> Optional[MCPPostCallResponseObject]:
) -> MCPPostCallResponseObject | None:
print("post mcp tool call response_obj", response_obj)
# update the MCPPostCallResponseObject with the response_cost
response_obj.hidden_params.response_cost = 1.42
@ -443,8 +465,8 @@ async def test_mcp_tool_call_hook():
# Call mcp tool using the correct separator format (- not /)
response = await mcp_server_tool_call(
name="zapier_gmail_server-add_tools", # Use correct prefixed name with - separator
arguments={"test": "test"},
_mcp_request_ctx(),
_call_tool_params("zapier_gmail_server-add_tools", {"test": "test"}),
)
# wait 1-2 seconds for logging to be processed

View file

@ -19,7 +19,7 @@ import pytest
import uvicorn
import yaml
from mcp import ClientSession
from mcp.client.streamable_http import streamablehttp_client
from mcp.client.streamable_http import streamable_http_client
from mcp.types import CallToolResult
from starlette.requests import Request
@ -206,7 +206,7 @@ class TestProxyMcpSimpleConnections:
@pytest.mark.asyncio
async def test_proxy_mcp_stdio_roundtrip(self, proxy_server_url: str) -> None:
async with asyncio.timeout(20):
async with streamablehttp_client(
async with streamable_http_client(
url=f"{proxy_server_url}/mcp",
headers={
"Authorization": PROXY_AUTHORIZATION_HEADER,
@ -227,7 +227,7 @@ class TestProxyMcpSimpleConnections:
@pytest.mark.asyncio
async def test_proxy_mcp_streamable_http_roundtrip(self, proxy_server_url: str) -> None:
async with asyncio.timeout(20):
async with streamablehttp_client(
async with streamable_http_client(
url=f"{proxy_server_url}/mcp",
headers={
"Authorization": PROXY_AUTHORIZATION_HEADER,
@ -248,7 +248,7 @@ class TestProxyMcpSimpleConnections:
@pytest.mark.asyncio
async def test_proxy_mcp_lists_all_servers_without_header(self, proxy_server_url: str) -> None:
async with asyncio.timeout(20):
async with streamablehttp_client(
async with streamable_http_client(
url=f"{proxy_server_url}/mcp",
headers={"Authorization": PROXY_AUTHORIZATION_HEADER},
) as (read, write, _get_session_id):
@ -296,7 +296,7 @@ class TestProxyMcpStatelessBehavior:
"""Two independent clients connect and operate without sharing session state."""
async with asyncio.timeout(30):
# --- Client A: connect, initialize, call tool ---
async with streamablehttp_client(
async with streamable_http_client(
url=f"{proxy_server_url}/mcp",
headers={
"Authorization": PROXY_AUTHORIZATION_HEADER,
@ -316,7 +316,7 @@ class TestProxyMcpStatelessBehavior:
await asyncio.sleep(0.5)
# --- Client B: completely independent connection ---
async with streamablehttp_client(
async with streamable_http_client(
url=f"{proxy_server_url}/mcp",
headers={
"Authorization": PROXY_AUTHORIZATION_HEADER,
@ -342,7 +342,7 @@ def _payload(result: typing.Any) -> typing.Any:
def _proxy_session(proxy_server_url: str, **extra_headers: str):
return streamablehttp_client(
return streamable_http_client(
url=f"{proxy_server_url}/mcp/proxy",
headers={"Authorization": PROXY_AUTHORIZATION_HEADER, **extra_headers},
)

View file

@ -11,17 +11,12 @@ from unittest.mock import AsyncMock, MagicMock, patch
import anyio
import httpx2
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
from mcp import MCPError
from mcp.client.streamable_http import streamable_http_client
from pydantic import ValidationError
from mcp.shared.message import SessionMessage
from mcp_types.version import LATEST_HANDSHAKE_VERSION
from pydantic import TypeAdapter
from mcp.types import (
CONNECTION_CLOSED,
INTERNAL_ERROR,
LATEST_PROTOCOL_VERSION,
REQUEST_TIMEOUT,
CallToolResult,
ErrorData,
@ -33,9 +28,10 @@ from mcp.types import (
LoggingMessageNotificationParams,
ServerCapabilities,
)
from mcp_types.version import LATEST_HANDSHAKE_VERSION
from pydantic import TypeAdapter, ValidationError
# Add the parent directory to the path so we can import litellm
import litellm.experimental_mcp_client.client as mcp_client_module
from litellm.experimental_mcp_client.client import (
MCPClient,
@ -51,9 +47,9 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_format_byok_openapi_auth_header,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
from litellm.types.mcp import MCPAuth, MCPStdioConfig, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
_JSONRPC_MESSAGE_ADAPTER: Final = TypeAdapter(JSONRPCMessage)
@ -1188,7 +1184,7 @@ async def test_a_custom_credential_header_is_stripped_when_a_redirect_crosses_or
operator moved to its own slot would be replayed to whatever host the upstream redirects to.
Verified against real httpx redirect handling, not a hand-built request.
"""
seen: "list[tuple[str, str]]" = []
seen: list[tuple[str, str]] = []
def handler(request: httpx2.Request) -> httpx2.Response:
seen.append((request.url.host, request.headers.get("esb-oauth", "<stripped>")))
@ -1280,7 +1276,7 @@ async def test_the_guard_agrees_with_httpx_about_authorization(start: str, targe
outcomes. A future httpx that changes its redirect rule reds here instead of silently leaving
the custom slot forwarded where Authorization is not (or stripped where it is not needed).
"""
seen: "list[tuple[str, str, str]]" = []
seen: list[tuple[str, str, str]] = []
def handler(request: httpx2.Request) -> httpx2.Response:
seen.append(
@ -1325,11 +1321,11 @@ def test_a_differently_cased_injected_header_cannot_shadow_the_slot() -> None:
@pytest.mark.parametrize(
("content_type", "body", "expected_type"),
[
("text/html", b"<html>secret-page</html>", ValueError),
("application/json", b"secret-invalid-json", ValidationError),
("application/json", b"", ValidationError),
("application/json", b'{"secret":"invalid-rpc"}', ValidationError),
("application/json", b'{"jsonrpc":"2.0","id":0,"result":{"secret":"invalid-schema"}}', ValidationError),
("text/html", b"<html>secret-page</html>", MCPError),
("application/json", b"secret-invalid-json", MCPError),
("application/json", b"", MCPError),
("application/json", b'{"secret":"invalid-rpc"}', MCPError),
("application/json", b'{"jsonrpc":"2.0","id":0}', MCPError),
],
)
async def test_invalid_http_response_surfaces_without_waiting_for_timeout(
@ -1623,6 +1619,7 @@ async def test_sse_read_failure_is_preserved() -> None:
@pytest.mark.parametrize("mode", ["ok", "closed", "silent"])
async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str) -> None:
from mcp import ClientSession
from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message
logging_callback: Final = AsyncMock()
@ -1647,8 +1644,7 @@ async def test_transport_completion_and_normal_messages(transport: MCPTransport,
if mode == "closed":
assert "connection was closed" in _connection_error_message(caught.value, client.server_url, 0.2)
else:
assert caught.value.error.code == CONNECTION_CLOSED
assert "SSE stream ended" in caught.value.error.message
assert isinstance(as_mcp_read_timeout(caught.value), TimeoutError)
@pytest.mark.asyncio
@ -1843,6 +1839,7 @@ async def test_optional_discovery_capabilities_and_errors(
@pytest.mark.parametrize("supports_first", (True, False))
async def test_optional_discovery_uses_each_sessions_capabilities(supports_first: bool) -> None:
from unittest.mock import Mock
from mcp.types import JSONRPCRequest
capabilities: Final = iter(({"resources": {}}, {}) if supports_first else ({}, {"resources": {}}))

View file

@ -5,20 +5,17 @@ Tests for MCPDebug — MCP OAuth2 debug response headers.
import asyncio
from typing import Final
import httpx
import pytest
from starlette.types import Message
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import AuthResolution
import httpx
from litellm.proxy._experimental.mcp_server.mcp_debug import (
MCP_DEBUG_REQUEST_HEADER,
MCPAuthDiagnostics,
MCPDebug,
describe_upstream_http_failure,
MCPAuthDiagnostics,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import AuthResolution
class TestIsDebugEnabled:
@ -265,6 +262,24 @@ class TestDescribeUpstreamHttpFailure:
assert describe_upstream_http_failure(ConnectionError("refused")) is None
def _mcp_request_ctx(**overrides):
from types import SimpleNamespace
from mcp.server.context import ServerRequestContext
kwargs = {
"session": SimpleNamespace(),
"lifespan_context": {},
"protocol_version": "2025-06-18",
"method": "",
"params": None,
"request_id": 1,
"meta": None,
"request": None,
}
kwargs.update(overrides)
return ServerRequestContext(**kwargs)
@pytest.mark.parametrize("body", [
b'{"password":"first second","token":"demo-secret"}',
b'{"nested":[{"access_token":"first,second"}]}',
@ -467,10 +482,9 @@ def test_diagnostics_keep_requests_separate_and_do_not_collapse_multiple_servers
async def test_concurrent_mcp_messages_record_on_their_own_http_scope() -> None:
from unittest.mock import MagicMock
from mcp.server.lowlevel.server import request_ctx
from mcp.shared.context import RequestContext
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,
record_auth_resolution,
@ -481,16 +495,16 @@ async def test_concurrent_mcp_messages_record_on_their_own_http_scope() -> None:
second: Final = MCPAuthDiagnostics()
async def record(diagnostics: MCPAuthDiagnostics, source: AuthResolution) -> None:
context: Final = RequestContext(
request_id=1, meta=None, session=session, lifespan_context=None,
context: Final = _mcp_request_ctx(
session=session,
request=Request({"type": "http", MCP_AUTH_DIAGNOSTICS_SCOPE_KEY: diagnostics}),
)
token: Final = request_ctx.set(context)
token: Final = active_mcp_request_ctx_var.set(context)
try:
await asyncio.sleep(0)
record_auth_resolution("same-server", source)
finally:
request_ctx.reset(token)
active_mcp_request_ctx_var.reset(token)
await asyncio.gather(record(first, AuthResolution.stored_user_token), record(second, AuthResolution.per_request_header))
assert first.resolution() == "stored-user-token"
@ -543,6 +557,7 @@ def test_oversized_request_omits_potentially_reflected_response_credentials():
@pytest.mark.asyncio
async def test_streamed_error_redacts_reflected_credentials_before_capture():
import json
from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response
secret = "generic-credential-123"

View file

@ -44,16 +44,28 @@ async def test_proxy_rejects_non_tool_protocol_operations() -> None:
assert options.capabilities.resources is None
assert options.capabilities.tools is not None
from types import SimpleNamespace
from mcp.server.context import ServerRequestContext
from mcp.types import GetPromptRequestParams, PaginatedRequestParams, ReadResourceRequestParams
ctx = ServerRequestContext(
session=SimpleNamespace(),
lifespan_context={},
protocol_version="2025-06-18",
method="",
)
with pytest.raises(MCPError):
await server.list_prompts()
await server.list_prompts(ctx, PaginatedRequestParams())
with pytest.raises(MCPError):
await server.get_prompt("prompt", {})
await server.get_prompt(ctx, GetPromptRequestParams(name="prompt", arguments={}))
with pytest.raises(MCPError):
await server.list_resources()
await server.list_resources(ctx, PaginatedRequestParams())
with pytest.raises(MCPError):
await server.list_resource_templates()
await server.list_resource_templates(ctx, PaginatedRequestParams())
with pytest.raises(MCPError):
await server.read_resource(AnyUrl("https://example.com/resource"))
await server.read_resource(ctx, ReadResourceRequestParams(uri="https://example.com/resource"))
class FailureRecorder(CustomLogger):

View file

@ -28,14 +28,14 @@ def _params(**overrides):
role="user", content=SimpleNamespace(type="text", text="hi")
)
],
systemPrompt="be concise",
maxTokens=128,
system_prompt="be concise",
max_tokens=128,
temperature=None,
stopSequences=None,
stop_sequences=None,
tools=None,
toolChoice=None,
tool_choice=None,
metadata=None,
modelPreferences=None,
model_preferences=None,
)
base.update(overrides)
return SimpleNamespace(**base)
@ -52,13 +52,13 @@ class TestBuildCompletionKwargs:
async def test_should_include_sampling_options_and_tools(self):
params = _params(
temperature=0.3,
stopSequences=["STOP"],
stop_sequences=["STOP"],
tools=[
SimpleNamespace(
name="search", description="d", input_schema={"type": "object"}
)
],
toolChoice=SimpleNamespace(mode="required"),
tool_choice=SimpleNamespace(mode="required"),
metadata={"trace": "abc"},
)
with patch(

View file

@ -151,7 +151,7 @@ class TestConvertMcpToolChoiceToOpenAI:
class TestConvertImageAndAudioContent:
def test_should_convert_image_to_data_uri(self):
content = SimpleNamespace(type="image", data="aGVsbG8=", mimeType="image/jpeg")
content = SimpleNamespace(type="image", data="aGVsbG8=", mime_type="image/jpeg")
result = _convert_single_content(content)
assert result == {
"type": "image_url",
@ -159,20 +159,20 @@ class TestConvertImageAndAudioContent:
}
def test_should_map_audio_mime_to_format(self):
content = SimpleNamespace(type="audio", data="Zm9v", mimeType="audio/mp3")
content = SimpleNamespace(type="audio", data="Zm9v", mime_type="audio/mp3")
result = _convert_single_content(content)
assert result["type"] == "input_audio"
assert result["input_audio"] == {"data": "Zm9v", "format": "mp3"}
def test_should_default_unknown_audio_mime_to_wav(self):
content = SimpleNamespace(type="audio", data="Zm9v", mimeType="audio/weird")
content = SimpleNamespace(type="audio", data="Zm9v", mime_type="audio/weird")
result = _convert_single_content(content)
assert result["input_audio"]["format"] == "wav"
def test_should_flatten_list_content(self):
items = [
SimpleNamespace(type="text", text="a"),
SimpleNamespace(type="image", data="x", mimeType="image/png"),
SimpleNamespace(type="image", data="x", mime_type="image/png"),
]
result = _convert_mcp_content_to_openai(items)
assert isinstance(result, list)

View file

@ -1,5 +1,6 @@
import asyncio
import contextvars
import json
import os
from datetime import datetime, timedelta
from types import SimpleNamespace
@ -10,6 +11,7 @@ import pytest
from fastapi import HTTPException
from mcp import ReadResourceResult, Resource
from mcp.types import (
INVALID_REQUEST,
BlobResourceContents,
CallToolResult,
Prompt,
@ -17,7 +19,10 @@ from mcp.types import (
TextContent,
TextResourceContents,
)
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION
from starlette.types import Message, Scope
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
MCPTransport,
@ -75,6 +80,37 @@ def cleanup_mcp_global_state():
yield
def _mcp_request_ctx(**overrides):
from types import SimpleNamespace
from mcp.server.context import ServerRequestContext
kwargs = {
"session": SimpleNamespace(),
"lifespan_context": {},
"protocol_version": "2025-06-18",
"method": "",
"params": None,
"request_id": 1,
"meta": None,
"request": None,
}
kwargs.update(overrides)
return ServerRequestContext(**kwargs)
def _call_tool_params(name, arguments=None):
from mcp.types import CallToolRequestParams
return CallToolRequestParams(name=name, arguments=arguments)
def _paged_params():
from mcp.types import PaginatedRequestParams
return PaginatedRequestParams()
@pytest.mark.asyncio
async def test_mcp_server_tool_call_body_contains_request_data():
"""Test that proxy_server_request body contains name and arguments"""
@ -125,7 +161,7 @@ async def test_mcp_server_tool_call_body_contains_request_data():
MagicMock(),
):
# Call the function
await mcp_server_tool_call(tool_name, tool_arguments)
await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params(tool_name, tool_arguments))
# Verify the body contains the expected data
assert "proxy_server_request" in captured_data
@ -177,7 +213,7 @@ async def test_mcp_server_tool_call_forwards_client_headers_to_logging():
mock_call_mcp_tool,
):
with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()):
await mcp_server_tool_call("test_tool", {"param": "value"})
await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"}))
assert captured_headers.get("x-nuid") == "nuid-1"
assert captured_headers.get("x-app-id") == "app-1"
@ -229,7 +265,7 @@ async def test_mcp_server_tool_call_strips_custom_litellm_key_header():
{"litellm_key_header_name": "x-company-key"},
clear=False,
):
await mcp_server_tool_call("test_tool", {"param": "value"})
await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"}))
metadata_headers = captured_data["metadata"]["headers"]
assert metadata_headers.get("x-nuid") == "nuid-1"
@ -271,7 +307,7 @@ async def test_mcp_server_tool_call_relays_upstream_auth_error_as_iserror():
):
with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()):
with patch("litellm.proxy._experimental.mcp_server.server.verbose_logger", mock_logger):
result = await mcp_server_tool_call("test_tool", {"param": "value"})
result = await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"}))
assert result.is_error is True
# The dedicated MCPUpstreamAuthError branch (not the generic Exception fallthrough) produces this
@ -1725,7 +1761,7 @@ async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error(
),
):
with pytest.raises(MCPError) as exc_info:
await handle_list_tools()
await handle_list_tools(_mcp_request_ctx(), _paged_params())
assert exc_info.value.error.code == INVALID_REQUEST
assert exc_info.value.error.message == denial_message
@ -1751,7 +1787,7 @@ async def test_mcp_server_tool_call_renders_denial_message_not_detail_dict():
new=AsyncMock(side_effect=denial),
),
):
result = await mcp_server_tool_call("github-search_issues", {})
result = await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params("github-search_issues", {}))
assert result.is_error is True
assert result.content[0].text == f"Error: {denial_message}"
@ -1806,7 +1842,7 @@ async def test_mcp_server_tool_call_body_with_none_arguments():
MagicMock(),
):
# Call the function
await mcp_server_tool_call(tool_name, tool_arguments)
await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params(tool_name, tool_arguments))
# Verify the body contains the expected data
assert "proxy_server_request" in captured_data
@ -1978,8 +2014,6 @@ async def test_streamable_http_session_manager_is_stateless():
async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(
debug: bool, method: str, request_body: bytes, stateful: bool
) -> None:
from mcp.server.lowlevel.server import request_ctx
from mcp.shared.context import RequestContext
from starlette.requests import Request
from starlette.types import Message, Receive, Scope, Send
@ -1996,14 +2030,12 @@ async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(
async def handle_request(request_scope: Scope, receive: Receive, outgoing: Send) -> None:
await outgoing({"type": "http.response.start", "status": 200, "headers": []})
await observe_start(send.await_count)
context: Final = RequestContext(
request_id=1, meta=None, session=MagicMock(), lifespan_context=None, request=Request(request_scope)
)
token: Final = request_ctx.set(context)
context: Final = _mcp_request_ctx(request=Request(request_scope))
token: Final = active_mcp_request_ctx_var.set(context)
try:
record_auth_resolution("s1", AuthResolution.stored_user_token)
finally:
request_ctx.reset(token)
active_mcp_request_ctx_var.reset(token)
await outgoing(body)
stateless_handle: Final = AsyncMock(side_effect=handle_request)
@ -4922,11 +4954,12 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
Ensure list-tools logging path calls `async_success_handler` when enabled.
"""
try:
from mcp.types import Tool as MCPTool
from litellm.proxy._experimental.mcp_server.server import (
_get_tools_from_mcp_servers,
)
from litellm.proxy._types import UserAPIKeyAuth
from mcp.types import Tool as MCPTool
except ImportError:
pytest.skip("MCP server not available")
@ -7638,20 +7671,24 @@ class TestMCPMetaTraceCarrier:
(e.g. ``litellm.team.id``). Dropping it at the source is the regression guard."""
from types import SimpleNamespace
from mcp.types import RequestParams
from mcp.types import CallToolRequestParams
from litellm.proxy._experimental.mcp_server.server import (
_mcp_meta_trace_carrier,
)
meta = RequestParams.Meta.model_validate(
meta = CallToolRequestParams.model_validate(
{
"traceparent": "00-11111111111111111111111111111111-2222222222222222-01",
"tracestate": "rojo=1",
"baggage": "litellm.team.id=spoofed-team,litellm.metadata.user_api_key_user_id=attacker",
"progressToken": "p1",
}
)
"name": "t",
"_meta": {
"traceparent": "00-11111111111111111111111111111111-2222222222222222-01",
"tracestate": "rojo=1",
"baggage": "litellm.team.id=spoofed-team,litellm.metadata.user_api_key_user_id=attacker",
"progressToken": "p1",
},
},
by_name=False,
).meta
carrier = _mcp_meta_trace_carrier(SimpleNamespace(meta=meta))
assert carrier == {
"traceparent": "00-11111111111111111111111111111111-2222222222222222-01",
@ -7662,7 +7699,7 @@ class TestMCPMetaTraceCarrier:
def test_none_when_no_trace_context(self):
from types import SimpleNamespace
from mcp.types import RequestParams
from mcp.types import CallToolRequestParams
from litellm.proxy._experimental.mcp_server.server import (
_mcp_meta_trace_carrier,
@ -7670,7 +7707,7 @@ class TestMCPMetaTraceCarrier:
assert _mcp_meta_trace_carrier(None) is None
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=None)) is None
only_progress = RequestParams.Meta.model_validate({"progressToken": "p1"})
only_progress = CallToolRequestParams.model_validate({"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False).meta
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=only_progress)) is None
@ -7678,9 +7715,6 @@ class TestMCPMetaTraceCarrier:
async def test_stateful_mcp_tool_call_uses_current_requests_otel_destinations() -> None:
from types import SimpleNamespace
from mcp.server.lowlevel.server import request_ctx
from mcp.shared.context import RequestContext
from litellm.integrations.otel.model.destination import OtelDestination
from litellm.integrations.otel.plumbing.context import (
request_destinations,
@ -7723,20 +7757,14 @@ async def test_stateful_mcp_tool_call_uses_current_requests_otel_destinations()
set_auth_context(None, raw_headers={})
destinations_token = set_request_destinations((initialized_destination,))
scope = {_MCP_DESTINATIONS_SCOPE_KEY: (current_destination,)}
current_request_context = RequestContext(
request_id=1,
meta=None,
session=SimpleNamespace(),
lifespan_context=None,
request=SimpleNamespace(scope=scope),
)
request_token = request_ctx.set(current_request_context)
current_request_context = _mcp_request_ctx(request=SimpleNamespace(scope=scope))
request_token = active_mcp_request_ctx_var.set(current_request_context)
try:
result = await mcp_server_tool_call("otelcontext-observe", {})
result = await mcp_server_tool_call(current_request_context, _call_tool_params("otelcontext-observe", {}))
assert result.is_error is False
assert request_destinations() == (initialized_destination,)
finally:
request_ctx.reset(request_token)
active_mcp_request_ctx_var.reset(request_token)
reset_request_destinations(destinations_token)
global_mcp_tool_registry.tools.pop("otelcontext-observe", None)
global_mcp_server_manager.registry.pop(server.server_id, None)
@ -7876,10 +7904,10 @@ async def test_fire_mcp_tool_call_logging_iserror_logs_failure():
"""Regression test: a CallToolResult with is_error=True must go
down the failure logging path (async_failure_handler + post_call_failure_hook),
never async_success_handler."""
from litellm.proxy._experimental.mcp_server.exceptions import MCPToolResultError
from litellm.proxy._experimental.mcp_server.server import (
_fire_mcp_tool_call_logging,
)
from litellm.proxy._experimental.mcp_server.exceptions import MCPToolResultError
logging_obj = _mock_mcp_logging_obj()
proxy_logging_mock = _mock_mcp_proxy_logging()
@ -8229,11 +8257,11 @@ async def test_call_mcp_tool_skips_failure_hook_for_upstream_auth_error():
caller-must-reauth signal, not a failed call, so call_mcp_tool must re-raise it WITHOUT firing
post_call_failure_hook (which records a failure and can trip LLM exception alerts). The
streamable handler downgrades it to an informational isError result afterward."""
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._experimental.mcp_server.server import (
call_mcp_tool,
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._types import MCPTransport, UserAPIKeyAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -8421,7 +8449,7 @@ async def test_handle_list_tools_attaches_outcome_meta():
new=AsyncMock(return_value=listing),
),
):
result = await handle_list_tools()
result = await handle_list_tools(_mcp_request_ctx(), _paged_params())
assert isinstance(result, ListToolsResult)
wire = result.model_dump(by_alias=True)
@ -9210,3 +9238,123 @@ async def test_list_tools_injects_byok_credential_for_non_oauth2_auth_types(auth
assert seen_auth_headers == ["personal-api-key"]
assert [tool.name for tool in listing.tools] == ["byok-toolA"]
@pytest.mark.parametrize(
"method,handler_name",
[
("tools/list", "handle_list_tools"),
("tools/call", "mcp_server_tool_call"),
("prompts/list", "list_prompts"),
("prompts/get", "get_prompt"),
("resources/list", "list_resources"),
("resources/templates/list", "list_resource_templates"),
("resources/read", "read_resource"),
],
)
def test_mcp_server_registers_all_spec_handlers(method: str, handler_name: str) -> None:
from litellm.proxy._experimental.mcp_server import server as mcp_module
entry = mcp_module.server.get_request_handler(method)
assert entry is not None
assert getattr(mcp_module, handler_name) is entry.handler
@pytest.mark.asyncio
async def test_active_request_ctx_var_feeds_get_current_session() -> None:
from litellm.proxy._experimental.mcp_server.server import _get_current_session
session = SimpleNamespace()
ctx = _mcp_request_ctx(session=session)
token = active_mcp_request_ctx_var.set(ctx)
try:
assert _get_current_session() is session
finally:
active_mcp_request_ctx_var.reset(token)
assert _get_current_session() is None
@pytest.mark.asyncio
async def test_active_request_ctx_var_feeds_auth_resolution_recording() -> None:
from starlette.requests import Request
from litellm.proxy._experimental.mcp_server.mcp_debug import (
MCP_AUTH_DIAGNOSTICS_SCOPE_KEY,
MCPAuthDiagnostics,
record_auth_resolution,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import AuthResolution
diagnostics = MCPAuthDiagnostics()
ctx = _mcp_request_ctx(request=Request({"type": "http", MCP_AUTH_DIAGNOSTICS_SCOPE_KEY: diagnostics}))
token = active_mcp_request_ctx_var.set(ctx)
try:
record_auth_resolution("s1", AuthResolution.static_token)
finally:
active_mcp_request_ctx_var.reset(token)
assert diagnostics.resolution() == "static-token"
@pytest.mark.asyncio
@pytest.mark.parametrize(
("header_value", "expected_rejected"),
[
("2025-06-18", False),
("2025-11-25", False),
("2026-07-28", True),
("1999-01-01", True),
],
)
async def test_streamable_http_rejects_modern_protocol_version(header_value: str, expected_rejected: bool) -> None:
from litellm.proxy._experimental.mcp_server import server as mcp_module
from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version
scope: Scope = {
"type": "http",
"method": "POST",
"path": "/mcp",
"headers": [(b"mcp-protocol-version", header_value.encode("latin-1"))],
}
assert (unsupported_protocol_version(scope) == header_value) is expected_rejected
if not expected_rejected:
return
sent: list[Message] = []
async def receive() -> Message:
return {"type": "http.request", "body": b"", "more_body": False}
async def send(message: Message) -> None:
sent.append(message)
await mcp_module.handle_streamable_http_mcp(scope, receive, send)
start = next(m for m in sent if m["type"] == "http.response.start")
assert start["status"] == 400
body = json.loads(b"".join(m.get("body", b"") for m in sent if m["type"] == "http.response.body"))
assert body["error"]["code"] == INVALID_REQUEST
assert header_value in body["error"]["message"]
for version in body["error"]["message"].split("supported: ")[1].split(", "):
assert version in HANDSHAKE_PROTOCOL_VERSIONS
@pytest.mark.asyncio
async def test_initialize_never_negotiates_outside_handshake_versions() -> None:
from mcp.server.runner import ServerRunner
from litellm.proxy._experimental.mcp_server import server as mcp_module
negotiate = ServerRunner._negotiate_initialize
for requested in ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25", "9999-01-01"):
_, negotiated = negotiate({"protocolVersion": requested, "capabilities": {}, "clientInfo": {"name": "t", "version": "0"}})
assert negotiated in HANDSHAKE_PROTOCOL_VERSIONS
from mcp.server.connection import Connection
runner = ServerRunner(mcp_module.server, Connection.from_envelope(LATEST_HANDSHAKE_VERSION, None, None), None)
result = runner._handle_initialize(
{"protocolVersion": "9999-01-01", "capabilities": {}, "clientInfo": {"name": "t", "version": "0"}}
)
assert result.protocol_version in HANDSHAKE_PROTOCOL_VERSIONS

View file

@ -1,5 +1,6 @@
import importlib
import asyncio
import functools
import json
import logging
import os
@ -84,6 +85,23 @@ def _reload_mcp_manager_module():
return reloaded
def _mcp_request_ctx(**overrides):
from mcp.server.context import ServerRequestContext
from types import SimpleNamespace
kwargs = {
"session": SimpleNamespace(),
"lifespan_context": {},
"protocol_version": "2025-06-18",
"method": "",
"params": None,
"request_id": 1,
"meta": None,
"request": None,
}
kwargs.update(overrides)
return ServerRequestContext(**kwargs)
@pytest.fixture(autouse=True)
def enable_eager_mcp_oauth_discovery(monkeypatch):
monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "1")
@ -12719,8 +12737,7 @@ async def test_debug_resolution_matches_final_header_conflict_winner(
expected_source: str,
expected_authorization: str | None,
) -> None:
from mcp.server.lowlevel.server import request_ctx
from mcp.shared.context import RequestContext
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
from starlette.requests import Request
from pydantic import SecretStr
@ -12743,8 +12760,7 @@ async def test_debug_resolution_matches_final_header_conflict_winner(
store = Store()
context = MCPAuthenticatedUser(UserAPIKeyAuth(user_id="alice"))
diagnostics = MCPAuthDiagnostics()
token = request_ctx.set(RequestContext(
request_id=1, meta=None, session=MagicMock(), lifespan_context=None,
token = active_mcp_request_ctx_var.set(_mcp_request_ctx(
request=Request({"type": "http", MCP_AUTH_DIAGNOSTICS_SCOPE_KEY: diagnostics}),
))
selected = {
@ -12771,22 +12787,20 @@ async def test_debug_resolution_matches_final_header_conflict_winner(
assert request.headers.get("Authorization") == expected_authorization
assert store.calls == (1 if config == "stored" else 0)
finally:
request_ctx.reset(token)
active_mcp_request_ctx_var.reset(token)
@pytest.mark.asyncio
@pytest.mark.parametrize("transport", ["http", "stdio"])
async def test_debug_reports_legacy_signing_and_non_http_transport(transport: Literal["http", "stdio"]) -> None:
from mcp.server.lowlevel.server import request_ctx
from mcp.shared.context import RequestContext
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_debug import MCP_AUTH_DIAGNOSTICS_SCOPE_KEY, MCPAuthDiagnostics
from litellm.types.mcp_server.mcp_server_manager import MCPServer
diagnostics = MCPAuthDiagnostics()
token = request_ctx.set(RequestContext(
request_id=1, meta=None, session=MagicMock(), lifespan_context=None,
token = active_mcp_request_ctx_var.set(_mcp_request_ctx(
request=Request({"type": "http", MCP_AUTH_DIAGNOSTICS_SCOPE_KEY: diagnostics}),
))
try:
@ -12807,7 +12821,7 @@ async def test_debug_reports_legacy_signing_and_non_http_transport(transport: Li
assert request.headers["Authorization"].startswith("AWS4-HMAC-SHA256 ")
assert "Credential=AKIDEXAMPLE/" in request.headers["Authorization"]
finally:
request_ctx.reset(token)
active_mcp_request_ctx_var.reset(token)
@pytest.mark.asyncio
@ -13063,10 +13077,14 @@ def _mcp_upstream(respond):
"""Drive the SDK's streamable-HTTP transport off an httpx2 MockTransport; respx only sees httpx."""
from litellm.experimental_mcp_client.client import MCPClient
def factory(*args, **kwargs):
return httpx2.AsyncClient(transport=httpx2.MockTransport(respond))
def make_client(self, *args, **kwargs):
return httpx2.AsyncClient(
transport=httpx2.MockTransport(respond),
headers=kwargs.get("headers"),
auth=kwargs.get("auth") or self._resolved_auth or self._aws_auth,
)
with patch.object(MCPClient, "_create_httpx_client_factory", lambda self: factory):
with patch.object(MCPClient, "_create_httpx_client_factory", lambda self: functools.partial(make_client, self)):
yield

View file

@ -85,6 +85,30 @@ FAKE_VECTORS: dict[str, Vector] = {
}
def _mcp_request_ctx(**overrides):
from types import SimpleNamespace
from mcp.server.context import ServerRequestContext
kwargs = {
"session": SimpleNamespace(),
"lifespan_context": {},
"protocol_version": "2025-06-18",
"method": "",
"params": None,
"request_id": 1,
"meta": None,
"request": None,
}
kwargs.update(overrides)
return ServerRequestContext(**kwargs)
def _paged_params():
from mcp.types import PaginatedRequestParams
return PaginatedRequestParams()
class RecordingEmbedder:
def __init__(self) -> None:
self.calls: list[tuple[str, ...]] = []
@ -1146,25 +1170,23 @@ class TestDispatchVirtualMcpTool:
class TestCaptureHostProgressCallback:
"""Covers the host progress-forwarding helper extracted from the tool call path."""
def test_returns_none_when_request_context_unavailable(self) -> None:
def test_returns_none_when_no_meta(self) -> None:
from types import SimpleNamespace
from litellm.proxy._experimental.mcp_server.server import (
_capture_host_progress_callback,
)
class _NoCtx:
@property
def request_context(self): # type: ignore[no-untyped-def]
raise RuntimeError("no context")
assert _capture_host_progress_callback(_NoCtx()) is None
assert _capture_host_progress_callback(SimpleNamespace(meta=None, session=object())) is None
def test_returns_none_when_no_progress_token(self) -> None:
from litellm.proxy._experimental.mcp_server.server import (
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progress_token = None
from types import SimpleNamespace
host = SimpleNamespace(meta=SimpleNamespace(progress_token=None), session=MagicMock())
assert _capture_host_progress_callback(host) is None
def test_returns_callable_when_token_present(self) -> None:
@ -1172,9 +1194,9 @@ class TestCaptureHostProgressCallback:
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progress_token = "tok12345"
host.request_context.session = MagicMock()
from types import SimpleNamespace
host = SimpleNamespace(meta=SimpleNamespace(progress_token="tok12345"), session=MagicMock())
assert callable(_capture_host_progress_callback(host))
def test_returns_callable_when_token_is_integer(self) -> None:
@ -1182,9 +1204,9 @@ class TestCaptureHostProgressCallback:
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progress_token = 12345
host.request_context.session = MagicMock()
from types import SimpleNamespace
host = SimpleNamespace(meta=SimpleNamespace(progress_token=12345), session=MagicMock())
assert callable(_capture_host_progress_callback(host))
def test_returns_callable_when_token_is_zero(self) -> None:
@ -1192,9 +1214,9 @@ class TestCaptureHostProgressCallback:
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progress_token = 0
host.request_context.session = MagicMock()
from types import SimpleNamespace
host = SimpleNamespace(meta=SimpleNamespace(progress_token=0), session=MagicMock())
assert callable(_capture_host_progress_callback(host))
@pytest.mark.asyncio
@ -1203,10 +1225,10 @@ class TestCaptureHostProgressCallback:
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progress_token = 12345
from types import SimpleNamespace
session = AsyncMock()
host.request_context.session = session
host = SimpleNamespace(meta=SimpleNamespace(progress_token=12345), session=session)
callback = _capture_host_progress_callback(host)
assert callback is not None
@ -1232,9 +1254,9 @@ class TestHandleListToolsVirtual:
new_callable=AsyncMock,
return_value=(uak, None, None, None, None, None, None),
):
tools = await srv.handle_list_tools()
result = await srv.handle_list_tools(_mcp_request_ctx(), _paged_params())
assert {t.name for t in tools} == {
assert {t.name for t in result.tools} == {
MCP_TOOL_SEARCH_TOOL_NAME,
MCP_TOOL_CALL_TOOL_NAME,
AGENT_SEARCH_TOOL_NAME,
@ -1265,9 +1287,14 @@ class TestMcpServerToolCallErrorHandling:
side_effect=HTTPException(status_code=403, detail="User not allowed to call this tool"),
),
):
from mcp.types import CallToolRequestParams
result = await srv.mcp_server_tool_call(
name=MCP_TOOL_CALL_TOOL_NAME,
arguments={"tool_name": "other-server-tool", "arguments": {}},
_mcp_request_ctx(),
CallToolRequestParams(
name=MCP_TOOL_CALL_TOOL_NAME,
arguments={"tool_name": "other-server-tool", "arguments": {}},
),
)
assert result.is_error is True

View file

@ -3921,7 +3921,7 @@ class TestConnectionErrorMessage:
@pytest.mark.parametrize("read_timeout", [0, 1])
async def test_timeout_message_uses_the_deadline_that_expired(self, sdk_timeout: bool, read_timeout: int) -> None:
from mcp import MCPError
from mcp.types import ErrorData
from mcp.types import REQUEST_TIMEOUT, ErrorData
async def operation(client: rest_endpoints.MCPClient) -> dict[str, object]:
try:
@ -3930,7 +3930,7 @@ class TestConnectionErrorMessage:
if not sdk_timeout:
raise
try:
raise MCPError(code=408, message="secret-sdk-timeout") from elapsed
raise MCPError(code=REQUEST_TIMEOUT, message="secret-sdk-timeout") from elapsed
except MCPError as sdk_error:
raise TimeoutError() from sdk_error