mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
resolve
This commit is contained in:
parent
f7a5909019
commit
12f1b36089
8 changed files with 311 additions and 120 deletions
|
|
@ -239,7 +239,7 @@ class MCPClient:
|
|||
server_params = StdioServerParameters(
|
||||
command=self.stdio_config.get("command", ""),
|
||||
args=self.stdio_config.get("args", []),
|
||||
env=self.stdio_config.get("env", None),
|
||||
env=self._get_safe_stdio_env(self.stdio_config.get("env")),
|
||||
)
|
||||
return stdio_client(server_params), None
|
||||
if self.transport_type == MCPTransport.sse:
|
||||
|
|
@ -273,6 +273,53 @@ class MCPClient:
|
|||
)
|
||||
return transport_ctx, http_client
|
||||
|
||||
def _get_safe_stdio_env(
|
||||
self, provided_env: Optional[Dict[str, str]]
|
||||
) -> Optional[Dict[str, str]]:
|
||||
"""
|
||||
Return a safe environment for the stdio subprocess.
|
||||
|
||||
If provided_env is set, we use it as-is.
|
||||
If provided_env is None, we return a minimal allowlist from the parent environment
|
||||
to avoid leaking sensitive LiteLLM keys (OPENAI_API_KEY, etc.) to sub-processes.
|
||||
"""
|
||||
if provided_env is not None:
|
||||
return provided_env
|
||||
|
||||
import os
|
||||
|
||||
# Minimal allowlist of safe/standard environment variables
|
||||
safe_keys = {
|
||||
"PATH",
|
||||
"HOME",
|
||||
"USER",
|
||||
"LOGNAME",
|
||||
"TMPDIR",
|
||||
"TMP",
|
||||
"TEMP",
|
||||
"SHELL",
|
||||
"LANG",
|
||||
"LC_ALL",
|
||||
# Node/Package manager caches
|
||||
"NPM_CONFIG_CACHE",
|
||||
"PNPM_HOME",
|
||||
"XDG_CACHE_HOME",
|
||||
"XDG_CONFIG_HOME",
|
||||
"XDG_DATA_HOME",
|
||||
# System info
|
||||
"SYSTEMROOT",
|
||||
"COMSPEC",
|
||||
"PATHEXT",
|
||||
"WINDIR",
|
||||
}
|
||||
|
||||
safe_env = {}
|
||||
for key in safe_keys:
|
||||
if key in os.environ:
|
||||
safe_env[key] = os.environ[key]
|
||||
|
||||
return safe_env
|
||||
|
||||
async def _execute_session_operation(
|
||||
self,
|
||||
transport_ctx: Any,
|
||||
|
|
|
|||
|
|
@ -2559,17 +2559,24 @@ if MCP_AVAILABLE:
|
|||
return user_api_key_auth.model_copy(update={"object_permission": updated_op})
|
||||
|
||||
async def _render_mcp_error(
|
||||
e: Exception, scope: Scope, receive: Receive, send: Send
|
||||
e: Exception,
|
||||
scope: Scope,
|
||||
receive: Receive,
|
||||
send: Send,
|
||||
error_prefix: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Render an exception as a JSON response for ASGI handlers."""
|
||||
from starlette.responses import JSONResponse
|
||||
|
||||
status_code = 500
|
||||
detail = str(e)
|
||||
headers = {}
|
||||
|
||||
if isinstance(e, HTTPException):
|
||||
status_code = e.status_code
|
||||
detail = e.detail
|
||||
if e.headers:
|
||||
headers.update(e.headers)
|
||||
elif hasattr(e, "status_code"):
|
||||
status_code = getattr(e, "status_code")
|
||||
elif hasattr(e, "code"):
|
||||
|
|
@ -2579,9 +2586,14 @@ if MCP_AVAILABLE:
|
|||
except (ValueError, TypeError):
|
||||
status_code = 500
|
||||
|
||||
error_msg = error_prefix or "MCP request failed"
|
||||
if error_prefix is None and status_code in (401, 403):
|
||||
error_msg = "Authentication processing failed"
|
||||
|
||||
error_response = JSONResponse(
|
||||
status_code=status_code,
|
||||
content={"error": "MCP request failed", "details": detail},
|
||||
content={"error": error_msg, "details": detail},
|
||||
headers=headers,
|
||||
)
|
||||
await error_response(scope, receive, send)
|
||||
|
||||
|
|
@ -2773,7 +2785,9 @@ if MCP_AVAILABLE:
|
|||
_captured_session_id_container_var.reset(_capture_token)
|
||||
|
||||
except Exception as e:
|
||||
await _render_mcp_error(e, scope, receive, send)
|
||||
await _render_mcp_error(
|
||||
e, scope, receive, send, error_prefix="Authentication processing failed"
|
||||
)
|
||||
# No need to return Response for raw ASGI app.
|
||||
|
||||
async def handle_sse_post_messages(
|
||||
|
|
@ -2835,7 +2849,9 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
await sse.handle_post_message(scope, receive, send)
|
||||
except Exception as e:
|
||||
await _render_mcp_error(e, scope, receive, send)
|
||||
await _render_mcp_error(
|
||||
e, scope, receive, send, error_prefix="Authentication processing failed"
|
||||
)
|
||||
|
||||
def get_active_mcp_session() -> Optional[_McpServerSession]:
|
||||
"""Get the active downstream MCP session from the current context."""
|
||||
|
|
|
|||
|
|
@ -13,42 +13,89 @@ from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
|||
# Mock MCP types if not available
|
||||
try:
|
||||
from mcp.types import (
|
||||
ModelPreferences, ModelHint, SamplingMessage, TextContent,
|
||||
ImageContent, Tool, ToolChoice, CreateMessageRequestParams,
|
||||
ToolUseContent, ToolResultContent
|
||||
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
|
||||
def __init__(self, name=None):
|
||||
self.name = name
|
||||
|
||||
class ModelPreferences:
|
||||
def __init__(self, hints=None): self.hints = hints
|
||||
def __init__(self, hints=None):
|
||||
self.hints = hints
|
||||
|
||||
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 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 Tool:
|
||||
def __init__(self, name, description=None, inputSchema=None):
|
||||
self.name = name; self.description = description; self.inputSchema = inputSchema
|
||||
self.name = name
|
||||
self.description = description
|
||||
self.inputSchema = inputSchema
|
||||
|
||||
class ToolChoice:
|
||||
def __init__(self, mode="auto"): self.mode = mode
|
||||
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
|
||||
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
|
||||
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
|
||||
|
||||
|
||||
def test_resolve_model_from_preferences():
|
||||
# Test 1: Direct match
|
||||
|
|
@ -66,6 +113,7 @@ def test_resolve_model_from_preferences():
|
|||
# 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")
|
||||
|
|
@ -75,7 +123,7 @@ def test_convert_mcp_content_to_openai():
|
|||
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"}
|
||||
"image_url": {"url": "data:image/jpeg;base64,base64data"},
|
||||
}
|
||||
|
||||
# List of content
|
||||
|
|
@ -85,10 +133,13 @@ def test_convert_mcp_content_to_openai():
|
|||
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"))
|
||||
|
||||
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
|
||||
|
|
@ -97,9 +148,14 @@ def test_convert_mcp_messages_to_openai():
|
|||
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])
|
||||
|
||||
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"
|
||||
|
|
@ -107,7 +163,11 @@ def test_convert_mcp_messages_to_openai():
|
|||
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")])
|
||||
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
|
||||
|
|
@ -115,6 +175,7 @@ def test_convert_mcp_messages_to_openai():
|
|||
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])
|
||||
|
|
@ -122,47 +183,66 @@ def test_convert_mcp_tools_to_openai():
|
|||
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
|
||||
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.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
|
||||
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.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)
|
||||
|
||||
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"
|
||||
|
|
|
|||
|
|
@ -3,28 +3,37 @@ Unit tests for the MCP Elicitation Handler.
|
|||
Tests the elicitation/create handler that relays elicitation requests
|
||||
from upstream MCP servers to downstream clients or declines them.
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
import pytest
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Helper factories
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
def _make_form_params(message="Please provide info", schema=None):
|
||||
"""Create a mock ElicitRequestFormParams."""
|
||||
from mcp.types import ElicitRequestFormParams
|
||||
|
||||
params = MagicMock(spec=ElicitRequestFormParams)
|
||||
params.mode = "form"
|
||||
params.message = message
|
||||
params.requestedSchema = schema
|
||||
return params
|
||||
|
||||
|
||||
def _make_url_params(message="Click the link", url="https://auth.example.com"):
|
||||
"""Create a mock ElicitRequestURLParams."""
|
||||
from mcp.types import ElicitRequestURLParams
|
||||
|
||||
params = MagicMock(spec=ElicitRequestURLParams)
|
||||
params.mode = "url"
|
||||
params.message = message
|
||||
params.url = url
|
||||
params.elicitationId = "elicit-123"
|
||||
return params
|
||||
|
||||
|
||||
def _make_capabilities(form=True, url=True):
|
||||
"""Create mock client capabilities with elicitation support."""
|
||||
caps = MagicMock()
|
||||
|
|
@ -33,21 +42,27 @@ def _make_capabilities(form=True, url=True):
|
|||
elicit.url = MagicMock() if url else None
|
||||
caps.elicitation = elicit
|
||||
return caps
|
||||
|
||||
|
||||
def _make_capabilities_no_elicitation():
|
||||
"""Create mock client capabilities without elicitation."""
|
||||
caps = MagicMock()
|
||||
caps.elicitation = None
|
||||
return caps
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Tests: No downstream session (Tool Bridge mode)
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
class TestElicitationNoDownstream:
|
||||
"""Tests when no downstream client is available."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_decline_when_no_downstream_session(self):
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
handle_elicitation_request,
|
||||
)
|
||||
|
||||
params = _make_form_params()
|
||||
result = await handle_elicitation_request(
|
||||
context=MagicMock(),
|
||||
|
|
@ -55,16 +70,20 @@ class TestElicitationNoDownstream:
|
|||
downstream_session=None,
|
||||
)
|
||||
assert result.action == "decline"
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Tests: Downstream session relay
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
class TestElicitationRelay:
|
||||
"""Tests for relaying elicitation to downstream clients."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_relay_form_mode(self):
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
handle_elicitation_request,
|
||||
)
|
||||
|
||||
params = _make_form_params(message="Enter your name")
|
||||
mock_session = AsyncMock()
|
||||
mock_result = MagicMock()
|
||||
|
|
@ -80,11 +99,13 @@ class TestElicitationRelay:
|
|||
)
|
||||
mock_session.elicit_form.assert_called_once()
|
||||
assert result.action == "submit"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_relay_url_mode(self):
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
handle_elicitation_request,
|
||||
)
|
||||
|
||||
params = _make_url_params(
|
||||
message="Authenticate", url="https://oauth.example.com"
|
||||
)
|
||||
|
|
@ -101,11 +122,13 @@ class TestElicitationRelay:
|
|||
)
|
||||
mock_session.elicit_url.assert_called_once()
|
||||
assert result.action == "submit"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_decline_when_client_lacks_elicitation(self):
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
handle_elicitation_request,
|
||||
)
|
||||
|
||||
params = _make_form_params()
|
||||
mock_session = AsyncMock()
|
||||
caps = _make_capabilities_no_elicitation()
|
||||
|
|
@ -116,11 +139,13 @@ class TestElicitationRelay:
|
|||
downstream_capabilities=caps,
|
||||
)
|
||||
assert result.action == "decline"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_decline_when_client_lacks_url_mode(self):
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
handle_elicitation_request,
|
||||
)
|
||||
|
||||
params = _make_url_params()
|
||||
mock_session = AsyncMock()
|
||||
caps = _make_capabilities(form=True, url=False)
|
||||
|
|
@ -131,11 +156,13 @@ class TestElicitationRelay:
|
|||
downstream_capabilities=caps,
|
||||
)
|
||||
assert result.action == "decline"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_decline_when_client_lacks_form_mode(self):
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
handle_elicitation_request,
|
||||
)
|
||||
|
||||
params = _make_form_params()
|
||||
mock_session = AsyncMock()
|
||||
caps = _make_capabilities(form=False, url=True)
|
||||
|
|
@ -146,11 +173,13 @@ class TestElicitationRelay:
|
|||
downstream_capabilities=caps,
|
||||
)
|
||||
assert result.action == "decline"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_decline_gracefully_on_relay_failure(self):
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
handle_elicitation_request,
|
||||
)
|
||||
|
||||
params = _make_form_params()
|
||||
mock_session = AsyncMock()
|
||||
mock_session.elicit_form.side_effect = Exception("Connection lost")
|
||||
|
|
@ -162,16 +191,20 @@ class TestElicitationRelay:
|
|||
downstream_capabilities=caps,
|
||||
)
|
||||
assert result.action == "decline"
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Tests: Error handling
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
class TestElicitationErrorHandling:
|
||||
"""Tests for error handling in the elicitation handler."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_relay_without_capability_check_when_caps_none(self):
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
handle_elicitation_request,
|
||||
)
|
||||
|
||||
params = _make_form_params()
|
||||
mock_session = AsyncMock()
|
||||
mock_result = MagicMock()
|
||||
|
|
@ -185,4 +218,4 @@ class TestElicitationErrorHandling:
|
|||
downstream_capabilities=None,
|
||||
)
|
||||
mock_session.elicit_form.assert_called_once()
|
||||
assert result.action == "submit"
|
||||
assert result.action == "submit"
|
||||
|
|
|
|||
|
|
@ -201,7 +201,9 @@ class TestResolveModel:
|
|||
|
||||
with patch("litellm.model_list", []):
|
||||
# Simulate configured default model
|
||||
with patch.object(litellm, "default_mcp_sampling_model", "claude-3-haiku", create=True):
|
||||
with patch.object(
|
||||
litellm, "default_mcp_sampling_model", "claude-3-haiku", create=True
|
||||
):
|
||||
result = _resolve_model_from_preferences(None)
|
||||
assert result == "claude-3-haiku"
|
||||
|
||||
|
|
|
|||
|
|
@ -317,6 +317,7 @@ class TestMCPServerManager:
|
|||
mcp_auth_header=None,
|
||||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
**kwargs,
|
||||
):
|
||||
if server.name == "github":
|
||||
tool1 = MagicMock()
|
||||
|
|
@ -371,6 +372,7 @@ class TestMCPServerManager:
|
|||
mcp_auth_header=None,
|
||||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
**kwargs,
|
||||
):
|
||||
assert mcp_auth_header == "legacy-token" # Should use legacy header
|
||||
tool = MagicMock()
|
||||
|
|
@ -409,6 +411,7 @@ class TestMCPServerManager:
|
|||
mcp_auth_header=None,
|
||||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
**kwargs,
|
||||
):
|
||||
assert (
|
||||
mcp_auth_header == "server-specific-token"
|
||||
|
|
@ -990,6 +993,7 @@ class TestMCPServerManager:
|
|||
mcp_auth_header=None,
|
||||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
**kwargs,
|
||||
):
|
||||
assert (
|
||||
mcp_auth_header == "server-specific-token"
|
||||
|
|
|
|||
|
|
@ -497,37 +497,52 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401():
|
|||
oauth_server.auth_type = MCPAuth.oauth2
|
||||
oauth_server.needs_user_oauth_token = True
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(user_auth, None, ["repro_oauth_server"], None, None, None),
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
|
||||
True,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session",
|
||||
new_callable=AsyncMock,
|
||||
return_value=False,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db",
|
||||
new_callable=AsyncMock,
|
||||
return_value=None,
|
||||
) as mock_get_stored_token, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=oauth_server,
|
||||
), patch.object(
|
||||
session_manager,
|
||||
"handle_request",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handle_request:
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(user_auth, None, ["repro_oauth_server"], None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
|
||||
True,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session",
|
||||
new_callable=AsyncMock,
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db",
|
||||
new_callable=AsyncMock,
|
||||
return_value=None,
|
||||
) as mock_get_stored_token,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=oauth_server,
|
||||
),
|
||||
patch.object(
|
||||
session_manager,
|
||||
"handle_request",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handle_request,
|
||||
):
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 401
|
||||
assert "www-authenticate" in exc.headers
|
||||
# Verify 401 response was sent
|
||||
assert send.called
|
||||
# Extract status code from mock_send
|
||||
response_start = next(
|
||||
call.args[0]
|
||||
for call in send.mock_calls
|
||||
if call.args[0].get("type") == "http.response.start"
|
||||
)
|
||||
assert response_start["status"] == 401
|
||||
headers = dict(response_start["headers"])
|
||||
assert b"www-authenticate" in headers
|
||||
assert mock_get_stored_token.await_count == 1
|
||||
assert mock_handle_request.await_count == 0
|
||||
|
||||
|
|
@ -562,31 +577,39 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
|
|||
oauth_server.auth_type = MCPAuth.oauth2
|
||||
oauth_server.needs_user_oauth_token = True
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(user_auth, None, ["repro_oauth_server"], None, None, None),
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
|
||||
True,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session",
|
||||
new_callable=AsyncMock,
|
||||
return_value=False,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"Authorization": "Bearer cached-token"},
|
||||
) as mock_get_stored_token, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=oauth_server,
|
||||
), patch.object(
|
||||
session_manager,
|
||||
"handle_request",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handle_request:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(user_auth, None, ["repro_oauth_server"], None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
|
||||
True,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session",
|
||||
new_callable=AsyncMock,
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"Authorization": "Bearer cached-token"},
|
||||
) as mock_get_stored_token,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=oauth_server,
|
||||
),
|
||||
patch.object(
|
||||
session_manager,
|
||||
"handle_request",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handle_request,
|
||||
):
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
|
||||
assert mock_get_stored_token.await_count == 1
|
||||
|
|
|
|||
|
|
@ -489,9 +489,7 @@ class TestGetToolsByNames:
|
|||
{"name": "send_email", "description": "send mail"},
|
||||
]
|
||||
|
||||
matched = filter_instance._get_tools_by_names(
|
||||
["send_email"], available_tools
|
||||
)
|
||||
matched = filter_instance._get_tools_by_names(["send_email"], available_tools)
|
||||
|
||||
assert len(matched) == 1
|
||||
assert matched[0]["name"] == "send_email"
|
||||
|
|
@ -503,9 +501,7 @@ class TestGetToolsByNames:
|
|||
client_name = "litellm_" + canonical
|
||||
available_tools = [{"name": client_name, "description": "scrape"}]
|
||||
|
||||
matched = filter_instance._get_tools_by_names(
|
||||
[canonical], available_tools
|
||||
)
|
||||
matched = filter_instance._get_tools_by_names([canonical], available_tools)
|
||||
|
||||
assert len(matched) == 1
|
||||
# Must return the incoming tool unchanged so the client-facing
|
||||
|
|
@ -516,13 +512,9 @@ class TestGetToolsByNames:
|
|||
"""Some clients use dash as alias separator; accept that too."""
|
||||
filter_instance = self._make_filter()
|
||||
canonical = "weather_svc-get_weather"
|
||||
available_tools = [
|
||||
{"name": "mcp-" + canonical, "description": "weather"}
|
||||
]
|
||||
available_tools = [{"name": "mcp-" + canonical, "description": "weather"}]
|
||||
|
||||
matched = filter_instance._get_tools_by_names(
|
||||
[canonical], available_tools
|
||||
)
|
||||
matched = filter_instance._get_tools_by_names([canonical], available_tools)
|
||||
|
||||
assert len(matched) == 1
|
||||
assert matched[0]["name"] == "mcp-" + canonical
|
||||
|
|
@ -552,9 +544,7 @@ class TestGetToolsByNames:
|
|||
{"name": "litellm_" + canonical, "description": "wrapped"},
|
||||
]
|
||||
|
||||
matched = filter_instance._get_tools_by_names(
|
||||
[canonical], available_tools
|
||||
)
|
||||
matched = filter_instance._get_tools_by_names([canonical], available_tools)
|
||||
|
||||
assert len(matched) == 1
|
||||
assert matched[0]["name"] == canonical
|
||||
|
|
@ -567,9 +557,7 @@ class TestGetToolsByNames:
|
|||
separator-anchored suffixes of ``litellm_api-fs-read_file``.
|
||||
"""
|
||||
filter_instance = self._make_filter()
|
||||
available_tools = [
|
||||
{"name": "litellm_api-fs-read_file", "description": "read"}
|
||||
]
|
||||
available_tools = [{"name": "litellm_api-fs-read_file", "description": "read"}]
|
||||
|
||||
matched = filter_instance._get_tools_by_names(
|
||||
["fs-read_file", "api-fs-read_file"], available_tools
|
||||
|
|
@ -590,9 +578,7 @@ class TestGetToolsByNames:
|
|||
{"name": "my_" + canonical, "description": "plain search"},
|
||||
]
|
||||
|
||||
matched = filter_instance._get_tools_by_names(
|
||||
[canonical], available_tools
|
||||
)
|
||||
matched = filter_instance._get_tools_by_names([canonical], available_tools)
|
||||
|
||||
assert len(matched) == 1
|
||||
assert matched[0]["name"] == "my_" + canonical
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue