diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 52117fbcbac..fe2fb19e10b 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -288,7 +288,6 @@ class MCPClient: transport = await transport_ctx.__aenter__() try: read_stream, write_stream = transport[0], transport[1] - session_ctx = ClientSession(read_stream, write_stream) # Build session kwargs with optional callbacks session_kwargs: Dict[str, Any] = {} if self._sampling_callback is not None: diff --git a/litellm/proxy/_experimental/mcp_server/sampling_handler.py b/litellm/proxy/_experimental/mcp_server/sampling_handler.py index 8efd1d5a422..e600c300b59 100644 --- a/litellm/proxy/_experimental/mcp_server/sampling_handler.py +++ b/litellm/proxy/_experimental/mcp_server/sampling_handler.py @@ -88,8 +88,8 @@ def _resolve_model_from_preferences( # Fall back to first available model if available_model_names: return available_model_names[0] - # Last resort - return "gpt-4o-mini" + # Last resort - use LiteLLM default or return None + return getattr(litellm, "default_mcp_sampling_model", None) or "gpt-4o-mini" def _convert_mcp_content_to_openai( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 36886b98a9b..bd6b615259a 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -221,9 +221,6 @@ if MCP_AVAILABLE: ######################################################## ############ Initialize the MCP Server ################# ######################################################## - from fastapi import APIRouter - - router = APIRouter() server: Server = Server( name=LITELLM_MCP_SERVER_NAME, version=LITELLM_MCP_SERVER_VERSION, @@ -2825,7 +2822,6 @@ if MCP_AVAILABLE: return {"enabled": MCP_AVAILABLE} # Include the MCP router - app.include_router(router) # Mount SSE handlers using the SDK's documented pattern. # We use app.mount for raw ASGI callables to avoid Starlette's request/response wrapper. app.mount("/sse", handle_sse_mcp_endpoint) diff --git a/tests/mcp_tests/test_coverage_boost.py b/tests/mcp_tests/test_coverage_boost.py new file mode 100644 index 00000000000..9955012f40b --- /dev/null +++ b/tests/mcp_tests/test_coverage_boost.py @@ -0,0 +1,113 @@ +import pytest +from unittest.mock import MagicMock, AsyncMock, patch +from litellm.proxy._experimental.mcp_server.sampling_handler import ( + _convert_single_content, + _convert_openai_response_to_mcp_result, + handle_sampling_create_message, +) +from litellm.proxy._experimental.mcp_server.server import get_or_extract_auth_context +from litellm.proxy._types import UserAPIKeyAuth + +# Mock MCP types +try: + from mcp.types import ( + TextContent, ImageContent, SamplingMessage, + CreateMessageRequestParams, ToolUseContent, ToolResultContent + ) +except ImportError: + class TextContent: + 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 + class SamplingMessage: + 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 + 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 + class ToolResultContent: + def __init__(self, type="tool_result", toolUseId=None, content=None): + self.type = type; self.toolUseId = toolUseId; self.content = content + +class MockAudioContent: + def __init__(self, data="audio_data", mimeType="audio/wav"): + self.type = "audio" + self.data = data + self.mimeType = mimeType + +def test_convert_audio_content(): + audio = MockAudioContent() + result = _convert_single_content(audio) + assert result["type"] == "input_audio" + 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" + mock_tool_call = MagicMock() + 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 + ) + 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_read_stream._litellm_auth_context = MCPAuthenticatedUser( + user_api_key_auth=mock_user_auth, + mcp_auth_header=None, + mcp_servers=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers=None, + client_ip=None + ) + mock_session._read_stream = mock_read_stream + + 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("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")): + result = await get_or_extract_auth_context() + assert result[0] is None diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 19ea5715e57..c3f78bb7fe2 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -13,7 +13,7 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServer, MCPTransport, ) -from litellm.proxy._types import LiteLLM_ObjectPermissionTable +from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth from mcp.types import Tool as MCPTool, CallToolResult from mcp.types import TextContent @@ -791,18 +791,31 @@ async def test_list_tools_rest_api_success(): side_effect=lambda server_ids, client_ip: (server_ids, 0) ) - # Mock the _get_tools_for_single_server function + # Mock the get_auth_context function to return our mock user auth with patch( - "litellm.proxy._experimental.mcp_server.rest_endpoints._get_tools_for_single_server" - ) as mock_get_tools: - mock_get_tools.return_value = mock_tools + "litellm.proxy._experimental.mcp_server.server.get_auth_context", + return_value=( + mock_user_auth, + None, + ["test-server-123"], + None, + None, + {}, + None, + ), + ): + # Mock the _get_tools_for_single_server function + with patch( + "litellm.proxy._experimental.mcp_server.rest_endpoints._get_tools_for_single_server" + ) as mock_get_tools: + mock_get_tools.return_value = mock_tools - # Test successful case - response = await list_tool_rest_api( - request=mock_request, - server_id="test-server-123", - user_api_key_dict=mock_user_auth, - ) + # Test successful case + response = await list_tool_rest_api( + request=mock_request, + server_id="test-server-123", + user_api_key_dict=mock_user_auth, + ) assert isinstance(response, dict) assert len(response["tools"]) == 1 @@ -2764,6 +2777,18 @@ async def test_call_mcp_tool_uses_manager_permission_lookup(): "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", return_value=True, ), + patch( + "litellm.proxy._experimental.mcp_server.server.get_auth_context", + return_value=( + UserAPIKeyAuth(api_key="test", user_id="test"), + None, + ["test_server"], + None, + None, + {}, + None, + ), + ), ): mock_get_allowed.return_value = [mock_server.server_id] mock_tool_registry.get_tool.return_value = None @@ -2840,6 +2865,18 @@ async def test_call_mcp_tool_resolves_unprefixed_tool_name_and_checks_permission "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", return_value=True, ) as mock_is_allowed, + patch( + "litellm.proxy._experimental.mcp_server.server.get_auth_context", + return_value=( + UserAPIKeyAuth(api_key="test", user_id="test"), + None, + ["test_server"], + None, + None, + {}, + None, + ), + ), ): mock_get_allowed.return_value = [mock_server.server_id] mock_tool_registry.get_tool.return_value = None diff --git a/tests/mcp_tests/test_sampling_handler.py b/tests/mcp_tests/test_sampling_handler.py new file mode 100644 index 00000000000..846a1861c7f --- /dev/null +++ b/tests/mcp_tests/test_sampling_handler.py @@ -0,0 +1,170 @@ +import pytest +from unittest.mock import MagicMock, AsyncMock, patch +from litellm.proxy._experimental.mcp_server.sampling_handler import ( + _resolve_model_from_preferences, + _convert_mcp_content_to_openai, + _convert_mcp_messages_to_openai, + _convert_mcp_tools_to_openai, + _convert_mcp_tool_choice_to_openai, + _convert_openai_response_to_mcp_result, + handle_sampling_create_message, +) + +# Mock MCP types if not available +try: + from mcp.types import ( + ModelPreferences, ModelHint, SamplingMessage, TextContent, + ImageContent, Tool, ToolChoice, CreateMessageRequestParams, + ToolUseContent, ToolResultContent + ) +except ImportError: + # Minimal mocks for testing when mcp package is not installed + class ModelHint: + def __init__(self, name=None): self.name = name + class ModelPreferences: + def __init__(self, hints=None): self.hints = hints + class SamplingMessage: + def __init__(self, role, content): self.role = role; self.content = content + class TextContent: + 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 + class Tool: + def __init__(self, name, description=None, inputSchema=None): + self.name = name; self.description = description; self.inputSchema = inputSchema + class ToolChoice: + def __init__(self, mode="auto"): self.mode = mode + class CreateMessageRequestParams: + def __init__(self, messages, modelPreferences=None, systemPrompt=None, + maxTokens=None, temperature=None, stopSequences=None, + tools=None, toolChoice=None, metadata=None): + self.messages = messages; self.modelPreferences = modelPreferences + self.systemPrompt = systemPrompt; self.maxTokens = maxTokens + self.temperature = temperature; self.stopSequences = stopSequences + self.tools = tools; self.toolChoice = toolChoice; self.metadata = metadata + 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 + class ToolResultContent: + def __init__(self, type="tool_result", toolUseId=None, content=None): + self.type = type; self.toolUseId = toolUseId; self.content = content + +def test_resolve_model_from_preferences(): + # Test 1: Direct match + prefs = ModelPreferences(hints=[ModelHint(name="gpt-4")]) + with patch("litellm.proxy.proxy_server.llm_router") as mock_router: + mock_router.get_model_names.return_value = ["gpt-4", "gpt-3.5-turbo"] + assert _resolve_model_from_preferences(prefs) == "gpt-4" + + # Test 2: Substring match + prefs = ModelPreferences(hints=[ModelHint(name="claude")]) + with patch("litellm.proxy.proxy_server.llm_router") as mock_router: + mock_router.get_model_names.return_value = ["anthropic/claude-3"] + assert _resolve_model_from_preferences(prefs) == "anthropic/claude-3" + + # Test 3: Default fallback + assert _resolve_model_from_preferences(None, default_model="fallback") == "fallback" + +def test_convert_mcp_content_to_openai(): + # Text content + text = TextContent(type="text", text="hello") + assert _convert_mcp_content_to_openai(text) == {"type": "text", "text": "hello"} + + # Image content + img = ImageContent(type="image", data="base64data", mimeType="image/jpeg") + assert _convert_mcp_content_to_openai(img) == { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,base64data"} + } + + # List of content + content_list = [text, img] + result = _convert_mcp_content_to_openai(content_list) + assert len(result) == 2 + assert result[0]["type"] == "text" + assert result[1]["type"] == "image_url" + +def test_convert_mcp_messages_to_openai(): + msg1 = SamplingMessage(role="user", content=TextContent(type="text", text="hi")) + msg2 = SamplingMessage(role="assistant", content=TextContent(type="text", text="hello")) + + # Standard messages + openai_msgs = _convert_mcp_messages_to_openai([msg1, msg2], system_prompt="system") + assert len(openai_msgs) == 3 + assert openai_msgs[0] == {"role": "system", "content": "system"} + assert openai_msgs[1]["role"] == "user" + assert openai_msgs[2]["role"] == "assistant" + + # Tool use/result conversion + tool_use = ToolUseContent(type="tool_use", id="call_1", name="search", input={"q": "test"}) + msg_tool_use = SamplingMessage(role="assistant", content=[TextContent(type="text", text="searching..."), tool_use]) + + openai_msgs = _convert_mcp_messages_to_openai([msg_tool_use]) + assert len(openai_msgs) == 1 + assert openai_msgs[0]["role"] == "assistant" + assert "tool_calls" in openai_msgs[0] + assert openai_msgs[0]["tool_calls"][0]["function"]["name"] == "search" + assert openai_msgs[0]["content"] == "searching..." + + tool_result = ToolResultContent(type="tool_result", toolUseId="call_1", content=[TextContent(type="text", text="found it")]) + msg_tool_result = SamplingMessage(role="user", content=[tool_result]) + openai_msgs = _convert_mcp_messages_to_openai([msg_tool_result]) + assert len(openai_msgs) == 1 + assert openai_msgs[0]["role"] == "tool" + assert openai_msgs[0]["tool_call_id"] == "call_1" + assert openai_msgs[0]["content"] == "found it" + +def test_convert_mcp_tools_to_openai(): + mcp_tool = Tool(name="my_tool", description="desc", inputSchema={"type": "object"}) + openai_tools = _convert_mcp_tools_to_openai([mcp_tool]) + assert len(openai_tools) == 1 + assert openai_tools[0]["type"] == "function" + assert openai_tools[0]["function"]["name"] == "my_tool" + +def test_convert_mcp_tool_choice_to_openai(): + assert _convert_mcp_tool_choice_to_openai(ToolChoice(mode="auto")) == "auto" + assert _convert_mcp_tool_choice_to_openai(ToolChoice(mode="required")) == "required" + assert _convert_mcp_tool_choice_to_openai(ToolChoice(mode="none")) == "none" + +@pytest.mark.asyncio +async def test_handle_sampling_create_message_success(): + params = CreateMessageRequestParams( + messages=[SamplingMessage(role="user", content=TextContent(type="text", text="hi"))], + maxTokens=100 + ) + + mock_response = MagicMock() + mock_response.choices = [MagicMock(message=MagicMock(content="hello response", tool_calls=None), finish_reason="stop")] + mock_response.model = "gpt-4o-mini" + + with patch("litellm.acompletion", new_callable=AsyncMock) as mock_completion: + mock_completion.return_value = mock_response + result = await handle_sampling_create_message(context=None, params=params) + + assert result.role == "assistant" + assert result.content.text == "hello response" + assert result.model == "gpt-4o-mini" + +@pytest.mark.asyncio +async def test_handle_sampling_with_auth_cost_tracking(): + from litellm.proxy._types import UserAPIKeyAuth + params = CreateMessageRequestParams( + messages=[SamplingMessage(role="user", content=TextContent(type="text", text="hi"))], + maxTokens=100 + ) + user_auth = UserAPIKeyAuth(api_key="sk-123", user_id="user-456", team_id="team-789") + + mock_response = MagicMock() + mock_response.choices = [MagicMock(message=MagicMock(content="ok", tool_calls=None), finish_reason="stop")] + mock_response.model = "gpt-4o-mini" + + with patch("litellm.acompletion", new_callable=AsyncMock) as mock_completion: + mock_completion.return_value = mock_response + await handle_sampling_create_message(context=None, params=params, user_api_key_auth=user_auth) + + # Verify auth was injected into metadata + kwargs = mock_completion.call_args.kwargs + assert kwargs["user"] == "user-456" + assert kwargs["metadata"]["user_api_key"] is not None + assert kwargs["metadata"]["user_api_key_team_id"] == "team-789"