mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
errors
This commit is contained in:
parent
2e837326c3
commit
869df308ef
6 changed files with 333 additions and 18 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
113
tests/mcp_tests/test_coverage_boost.py
Normal file
113
tests/mcp_tests/test_coverage_boost.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
170
tests/mcp_tests/test_sampling_handler.py
Normal file
170
tests/mcp_tests/test_sampling_handler.py
Normal file
|
|
@ -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"
|
||||
Loading…
Add table
Reference in a new issue