This commit is contained in:
Yug 2026-04-29 17:21:23 +05:30
parent 4218c08fd8
commit 81bb48dacd
5 changed files with 177 additions and 71 deletions

View file

@ -242,7 +242,7 @@ if MCP_AVAILABLE:
# client-visible POST path becomes /mcp/messages.
from mcp.server.sse import SseServerTransport as _McpSseServerTransport
sse = _McpSseServerTransport("/messages")
sse = _McpSseServerTransport("/mcp/messages")
# Create session managers (StreamableHTTP — stateless by default)
session_manager = StreamableHTTPSessionManager(
app=server,
@ -2743,7 +2743,7 @@ if MCP_AVAILABLE:
raw_headers,
) = await extract_mcp_auth_context(scope, path)
_sse_client_ip = IPAddressUtils.get_mcp_client_ip(request)
# set_auth_context here is a no-op for actual tool execution since the SDK
# set_auth_context here is a no-op for actual tool execution since the SDK
# processes messages in background tasks that don't inherit this ContextVar.
# Authentication must be recovered from the session-auth-storage during execution.
except HTTPException:
@ -2788,7 +2788,7 @@ if MCP_AVAILABLE:
session = request_ctx.get().session
read_stream = getattr(session, "_read_stream", None)
return getattr(read_stream, "_litellm_auth_context", None)
return _session_auth_storage.get(read_stream) if read_stream else None
except Exception:
return None

View file

@ -1,5 +1,5 @@
import pytest
from unittest.mock import MagicMock, AsyncMock, patch
from unittest.mock import MagicMock, patch
from litellm.proxy._experimental.mcp_server.sampling_handler import (
_convert_single_content,
_convert_openai_response_to_mcp_result,
@ -11,25 +11,49 @@ from litellm.proxy._types import UserAPIKeyAuth
# Mock MCP types
try:
from mcp.types import (
TextContent, ImageContent, SamplingMessage,
CreateMessageRequestParams, ToolUseContent, ToolResultContent
TextContent,
ImageContent,
SamplingMessage,
CreateMessageRequestParams,
ToolUseContent,
ToolResultContent,
)
except ImportError:
class TextContent:
def __init__(self, type="text", text=""): self.type = type; self.text = text
def __init__(self, type="text", text=""):
self.type = type
self.text = text
class ImageContent:
def __init__(self, type="image", data="", mimeType="image/png"):
self.type = type; self.data = data; self.mimeType = mimeType
self.type = type
self.data = data
self.mimeType = mimeType
class SamplingMessage:
def __init__(self, role, content): self.role = role; self.content = content
def __init__(self, role, content):
self.role = role
self.content = content
class CreateMessageRequestParams:
def __init__(self, messages, maxTokens=100): self.messages = messages; self.maxTokens = maxTokens
def __init__(self, messages, maxTokens=100):
self.messages = messages
self.maxTokens = maxTokens
class ToolUseContent:
def __init__(self, type="tool_use", id=None, name=None, input=None):
self.type = type; self.id = id; self.name = name; self.input = input
self.type = type
self.id = id
self.name = name
self.input = input
class ToolResultContent:
def __init__(self, type="tool_result", toolUseId=None, content=None):
self.type = type; self.toolUseId = toolUseId; self.content = content
self.type = type
self.toolUseId = toolUseId
self.content = content
class MockAudioContent:
def __init__(self, data="audio_data", mimeType="audio/wav"):
@ -37,6 +61,7 @@ class MockAudioContent:
self.data = data
self.mimeType = mimeType
def test_convert_audio_content():
audio = MockAudioContent()
result = _convert_single_content(audio)
@ -44,6 +69,7 @@ def test_convert_audio_content():
assert result["input_audio"]["data"] == "audio_data"
assert result["input_audio"]["format"] == "wav"
def test_convert_openai_response_to_mcp_result_with_tool_calls():
mock_choice = MagicMock()
mock_choice.message.content = "I will search now"
@ -51,42 +77,51 @@ def test_convert_openai_response_to_mcp_result_with_tool_calls():
mock_tool_call.id = "call_1"
mock_tool_call.function.name = "search"
mock_tool_call.function.arguments = '{"q": "test"}'
mock_choice.message.tool_calls = [mock_tool_call]
mock_choice.finish_reason = "tool_calls"
mock_response = MagicMock()
mock_response.choices = [mock_choice]
mock_response.model = "gpt-4"
result = _convert_openai_response_to_mcp_result(mock_response, model_name="gpt-4")
assert result.role == "assistant"
# It should have both text and tool use content
# Depending on implementation it might return CreateMessageResultWithTools
assert hasattr(result, "content")
@pytest.mark.asyncio
async def test_handle_sampling_no_package_error():
params = CreateMessageRequestParams(
messages=[SamplingMessage(role="user", content=TextContent(type="text", text="hi"))],
maxTokens=100
messages=[
SamplingMessage(role="user", content=TextContent(type="text", text="hi"))
],
maxTokens=100,
)
with patch("litellm.proxy._experimental.mcp_server.sampling_handler.MCP_SAMPLING_AVAILABLE", False):
with patch(
"litellm.proxy._experimental.mcp_server.sampling_handler.MCP_SAMPLING_AVAILABLE",
False,
):
result = await handle_sampling_create_message(context=None, params=params)
assert hasattr(result, "message")
assert "MCP sampling is not available" in result.message
@pytest.mark.asyncio
async def test_get_or_extract_auth_context_fallback():
# Test fallback to session read_stream when ContextVar is empty
mock_session = MagicMock()
mock_read_stream = MagicMock()
mock_user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="user-1")
from litellm.proxy._experimental.mcp_server.server import MCPAuthenticatedUser
mock_session._read_stream = mock_read_stream
from litellm.proxy._experimental.mcp_server.server import _session_auth_storage
_session_auth_storage[mock_read_stream] = MCPAuthenticatedUser(
user_api_key_auth=mock_user_auth,
mcp_auth_header=None,
@ -94,22 +129,32 @@ async def test_get_or_extract_auth_context_fallback():
mcp_server_auth_headers=None,
oauth2_headers=None,
raw_headers=None,
client_ip=None
client_ip=None,
)
mock_request_ctx = MagicMock()
mock_request_ctx.get.return_value.session = mock_session
with patch("litellm.proxy._experimental.mcp_server.server.get_auth_context", return_value=(None, None, None, None, None, {}, None)):
with patch(
"litellm.proxy._experimental.mcp_server.server.get_auth_context",
return_value=(None, None, None, None, None, {}, None),
):
with patch("mcp.server.lowlevel.server.request_ctx", mock_request_ctx):
result = await get_or_extract_auth_context()
assert result[0] == mock_user_auth
assert result[0].api_key is not None
@pytest.mark.asyncio
async def test_get_or_extract_auth_context_exception_handling():
# Test that it handles exceptions in fallback gracefully
with patch("litellm.proxy._experimental.mcp_server.server.get_auth_context", return_value=(None, None, None, None, None, {}, None)):
with patch("mcp.server.lowlevel.server.request_ctx", side_effect=Exception("Context error")):
with patch(
"litellm.proxy._experimental.mcp_server.server.get_auth_context",
return_value=(None, None, None, None, None, {}, None),
):
with patch(
"mcp.server.lowlevel.server.request_ctx",
side_effect=Exception("Context error"),
):
result = await get_or_extract_auth_context()
assert result[0] is None

View file

@ -1,6 +1,5 @@
import pytest
from unittest.mock import MagicMock, AsyncMock, patch
from litellm.proxy._types import MCPTransportType
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.proxy._experimental.mcp_server.server import (
_get_prompts_from_mcp_servers,
@ -9,6 +8,7 @@ from litellm.proxy._experimental.mcp_server.server import (
_get_tools_from_mcp_servers,
)
@pytest.mark.asyncio
async def test_get_prompts_from_mcp_servers_coverage():
server1 = MCPServer(
@ -19,18 +19,25 @@ async def test_get_prompts_from_mcp_servers_coverage():
)
mock_prompt = MagicMock()
mock_prompt.name = "test_prompt"
with patch("litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", return_value=[server1, server2]):
with patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_prompts_from_server", new_callable=AsyncMock) as mock_get:
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
return_value=[server1, server2],
):
with patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_prompts_from_server",
new_callable=AsyncMock,
) as mock_get:
mock_get.side_effect = [[mock_prompt], Exception("Server error")]
result = await _get_prompts_from_mcp_servers(
user_api_key_auth=None,
mcp_auth_header=None,
mcp_servers=["test1", "test2"]
mcp_servers=["test1", "test2"],
)
assert len(result) == 1
assert result[0] == mock_prompt
@pytest.mark.asyncio
async def test_get_resources_from_mcp_servers_coverage():
server1 = MCPServer(
@ -38,18 +45,23 @@ async def test_get_resources_from_mcp_servers_coverage():
)
mock_resource = MagicMock()
mock_resource.name = "test_resource"
with patch("litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", return_value=[server1]):
with patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_resources_from_server", new_callable=AsyncMock) as mock_get:
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
return_value=[server1],
):
with patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_resources_from_server",
new_callable=AsyncMock,
) as mock_get:
mock_get.return_value = [mock_resource]
result = await _get_resources_from_mcp_servers(
user_api_key_auth=None,
mcp_auth_header=None,
mcp_servers=["test1"]
user_api_key_auth=None, mcp_auth_header=None, mcp_servers=["test1"]
)
assert len(result) == 1
assert result[0] == mock_resource
@pytest.mark.asyncio
async def test_get_resource_templates_from_mcp_servers_coverage():
server1 = MCPServer(
@ -57,18 +69,23 @@ async def test_get_resource_templates_from_mcp_servers_coverage():
)
mock_template = MagicMock()
mock_template.name = "test_template"
with patch("litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", return_value=[server1]):
with patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_resource_templates_from_server", new_callable=AsyncMock) as mock_get:
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
return_value=[server1],
):
with patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_resource_templates_from_server",
new_callable=AsyncMock,
) as mock_get:
mock_get.return_value = [mock_template]
result = await _get_resource_templates_from_mcp_servers(
user_api_key_auth=None,
mcp_auth_header=None,
mcp_servers=["test1"]
user_api_key_auth=None, mcp_auth_header=None, mcp_servers=["test1"]
)
assert len(result) == 1
assert result[0] == mock_template
@pytest.mark.asyncio
async def test_get_tools_from_mcp_servers_coverage():
server1 = MCPServer(
@ -76,9 +93,15 @@ async def test_get_tools_from_mcp_servers_coverage():
)
mock_tool = MagicMock()
mock_tool.name = "test_tool"
with patch("litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", return_value=[server1]):
with patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server", new_callable=AsyncMock) as mock_get:
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
return_value=[server1],
):
with patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server",
new_callable=AsyncMock,
) as mock_get:
mock_get.return_value = [mock_tool]
# test with some tracking headers
result = await _get_tools_from_mcp_servers(
@ -86,7 +109,61 @@ async def test_get_tools_from_mcp_servers_coverage():
mcp_auth_header=None,
mcp_servers=["test1"],
log_list_tools_to_spendlogs=True,
litellm_trace_id="test-trace"
litellm_trace_id="test-trace",
)
assert len(result) == 1
assert result[0] == mock_tool
@pytest.mark.asyncio
async def test_handle_stale_mcp_session():
"""Test the handle_stale_mcp_session logic for handling missing session IDs on multiple workers."""
from litellm.proxy._experimental.mcp_server.server import _handle_stale_mcp_session
# 1. DELETE request for non-existent session
scope_delete = {
"headers": [(b"mcp-session-id", b"stale-session-123")],
"method": "DELETE",
"type": "http",
}
mock_mgr = MagicMock()
mock_mgr._server_instances = {}
mock_receive = AsyncMock()
mock_send = AsyncMock()
result_delete = await _handle_stale_mcp_session(
scope_delete, mock_receive, mock_send, mock_mgr
)
assert result_delete is True
# The JSONResponse success should have been sent
assert mock_send.call_count >= 1
# 2. POST request for non-existent session -> header should be stripped
scope_post = {
"headers": [
(b"mcp-session-id", b"stale-session-123"),
(b"content-type", b"application/json"),
],
"method": "POST",
}
result_post = await _handle_stale_mcp_session(
scope_post, mock_receive, mock_send, mock_mgr
)
assert result_post is False
headers = dict(scope_post["headers"])
assert b"mcp-session-id" not in headers
assert b"content-type" in headers
# 3. Request with valid session -> should return False immediately
mock_mgr._server_instances = {"valid-session-123": MagicMock()}
scope_valid = {
"headers": [(b"mcp-session-id", b"valid-session-123")],
"method": "POST",
}
result_valid = await _handle_stale_mcp_session(
scope_valid, mock_receive, mock_send, mock_mgr
)
assert result_valid is False
# Header should not be stripped
headers_valid = dict(scope_valid["headers"])
assert b"mcp-session-id" in headers_valid

View file

@ -46,6 +46,7 @@ async def test_get_or_extract_auth_context_fallback():
mock_session._read_stream = mock_read_stream
from litellm.proxy._experimental.mcp_server.server import _session_auth_storage
_session_auth_storage[mock_read_stream] = auth_user
mock_request_ctx = MagicMock()

View file

@ -853,21 +853,16 @@ async def test_concurrent_initialize_session_managers():
# Reset state before test
original_initialized = mcp_server._SESSION_MANAGERS_INITIALIZED
original_session_cm = mcp_server._session_manager_cm
original_sse_session_cm = mcp_server._sse_session_manager_cm
try:
mcp_server._SESSION_MANAGERS_INITIALIZED = False
mcp_server._session_manager_cm = None
mcp_server._sse_session_manager_cm = None
# Mock the session managers to avoid actual MCP initialization
with (
patch(
"litellm.proxy._experimental.mcp_server.server.session_manager"
) as mock_session_manager,
patch(
"litellm.proxy._experimental.mcp_server.server.sse_session_manager"
) as mock_sse_session_manager,
patch("litellm.proxy._experimental.mcp_server.server.verbose_logger"),
):
# Mock the run() method to return a mock context manager
@ -876,7 +871,6 @@ async def test_concurrent_initialize_session_managers():
mock_cm.__aexit__ = AsyncMock()
mock_session_manager.run.return_value = mock_cm
mock_sse_session_manager.run.return_value = mock_cm
# Create multiple concurrent tasks that call initialize_session_managers
async def init_task():
@ -896,14 +890,11 @@ async def test_concurrent_initialize_session_managers():
assert (
mock_session_manager.run.call_count == 1
), f"Expected 1 call to session_manager.run(), got {mock_session_manager.run.call_count}"
assert (
mock_sse_session_manager.run.call_count == 1
), f"Expected 1 call to sse_session_manager.run(), got {mock_sse_session_manager.run.call_count}"
# The context managers should only be entered once each
assert (
mock_cm.__aenter__.call_count == 2
), f"Expected 2 calls to __aenter__ (one for each session manager), got {mock_cm.__aenter__.call_count}"
mock_cm.__aenter__.call_count == 1
), f"Expected 1 call to __aenter__ (one for each session manager), got {mock_cm.__aenter__.call_count}"
# State should be properly set
assert mcp_server._SESSION_MANAGERS_INITIALIZED is True
@ -912,7 +903,6 @@ async def test_concurrent_initialize_session_managers():
# Restore original state
mcp_server._SESSION_MANAGERS_INITIALIZED = original_initialized
mcp_server._session_manager_cm = original_session_cm
mcp_server._sse_session_manager_cm = original_sse_session_cm
@pytest.mark.asyncio
@ -1077,20 +1067,13 @@ async def test_oauth2_headers_passed_to_mcp_client():
captured_client_args = {}
async def mock_create_mcp_client(
server,
mcp_auth_header=None,
extra_headers=None,
stdio_env=None,
*args,
**kwargs,
):
# Capture the arguments for verification
captured_client_args.update(
{
"server": server,
"mcp_auth_header": mcp_auth_header,
"extra_headers": extra_headers,
"stdio_env": stdio_env,
}
)
captured_client_args.update(kwargs)
if args and len(args) > 0:
captured_client_args["server"] = args[0]
# Return a mock client that doesn't actually connect
mock_client = MagicMock()
return mock_client