mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
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:
parent
783038010b
commit
0d2963fe89
11 changed files with 390 additions and 151 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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": {}}))
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue