mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
resolve
This commit is contained in:
parent
4218c08fd8
commit
81bb48dacd
5 changed files with 177 additions and 71 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue