This commit is contained in:
Yug 2026-05-01 07:41:59 +05:30
parent f7a5909019
commit 12f1b36089
8 changed files with 311 additions and 120 deletions

View file

@ -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,

View file

@ -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."""

View file

@ -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"

View file

@ -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"

View file

@ -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"

View file

@ -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"

View file

@ -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

View file

@ -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