Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 71 additions & 18 deletions litellm/proxy/hooks/mcp_semantic_filter/hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,27 @@ async def _expand_mcp_tools(

return openai_tools_as_dicts

async def _filter_expanded_tools(
self,
data: dict,
expanded_tools: list[dict[str, Any]],
) -> list[dict[str, Any]]:
"""
Apply the semantic filter to expanded MCP tool definitions.

Expanded tools are flat OpenAI function dicts with a top-level
"name" (see transform_mcp_tool_to_openai_responses_api_tool), so
filter_tools can name-match them against the semantic router.
"""
raw_messages = data.get("messages") or data.get("input") or []
messages = [{"role": "user", "content": raw_messages}] if isinstance(raw_messages, str) else raw_messages
user_query = self.filter.extract_user_query(messages)
if not user_query:
verbose_proxy_logger.debug("No user query found, skipping semantic filter on expanded MCP tools")
return expanded_tools

return await self.filter.filter_tools(query=user_query, available_tools=expanded_tools)

def _is_mcp_tool(self, tool: object) -> bool:
"""
Check whether *tool* is registered in the MCP semantic router.
Expand Down Expand Up @@ -172,6 +193,32 @@ def _emit_filter_metadata(
f"Semantic tool filter: all {len(native_tools)} tools are native, no MCP filtering applied"
)

def _emit_filter_metadata_safe(
self,
data: dict,
mcp_tools: list[object],
filtered_mcp_tools: list[object],
native_tools: list[object],
filtered_tools: list[object],
) -> None:
"""
Emit filter metadata without letting an emission failure abort the
already-filtered request.
"""
try:
self._emit_filter_metadata(
data=data,
mcp_tools=mcp_tools,
filtered_mcp_tools=filtered_mcp_tools,
native_tools=native_tools,
filtered_tools=filtered_tools,
)
except Exception as e:
verbose_proxy_logger.warning(
f"Failed to emit semantic filter metadata: {e}",
exc_info=True,
)

async def async_pre_call_hook(
self,
user_api_key_dict: "UserAPIKeyAuth",
Expand All @@ -194,9 +241,6 @@ async def async_pre_call_hook(
verbose_proxy_logger.debug("No tools in request, skipping semantic filter")
return None

# Expanded MCP tools are in OpenAI nested format which
# filter_tools/_extract_tool_info cannot name-match, so we skip
# semantic filtering and return early.
if self._should_expand_mcp_tools(tools):
verbose_proxy_logger.debug("Detected litellm_proxy MCP references, expanding before semantic filtering")

Expand All @@ -215,11 +259,26 @@ async def async_pre_call_hook(
verbose_proxy_logger.warning("No tools expanded from MCP references")
return None

data["tools"] = native_tools_before_expand + expanded_tools
if not self.filter.enabled:
data["tools"] = native_tools_before_expand + expanded_tools
verbose_proxy_logger.debug("Semantic filter disabled, forwarding expanded MCP tools unfiltered")
return data

filtered_expanded_tools = await self._filter_expanded_tools(data=data, expanded_tools=expanded_tools)

combined_tools = native_tools_before_expand + filtered_expanded_tools
Comment thread
greptile-apps[bot] marked this conversation as resolved.
data["tools"] = combined_tools
self._emit_filter_metadata_safe(
data=data,
mcp_tools=expanded_tools,
filtered_mcp_tools=filtered_expanded_tools,
native_tools=native_tools_before_expand,
filtered_tools=combined_tools,
)
verbose_proxy_logger.info(
f"Expanded MCP references to {len(expanded_tools)} tools "
f"({len(native_tools_before_expand)} native preserved), "
f"skipping semantic filter (OpenAI nested format)"
f"semantic filter selected {len(filtered_expanded_tools)}"
)
return data

Expand Down Expand Up @@ -285,19 +344,13 @@ async def async_pre_call_hook(

data["tools"] = filtered_tools

try:
self._emit_filter_metadata(
data=data,
mcp_tools=mcp_tools,
filtered_mcp_tools=filtered_mcp_tools,
native_tools=native_tools,
filtered_tools=filtered_tools,
)
except Exception as e:
verbose_proxy_logger.warning(
f"Failed to emit semantic filter metadata: {e}",
exc_info=True,
)
self._emit_filter_metadata_safe(
data=data,
mcp_tools=mcp_tools,
filtered_mcp_tools=filtered_mcp_tools,
native_tools=native_tools,
filtered_tools=filtered_tools,
)

return data

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -773,6 +773,251 @@ async def mock_embedding_async(*args, **kwargs):
print("✅ Responses API tool with MCP-matching name correctly classified as native")


@pytest.mark.asyncio
async def test_semantic_filter_hook_filters_expanded_litellm_proxy_tools():
"""
Regression test (LIT-4214): litellm_proxy MCP references must be
semantically filtered after expansion, with real filter stats.

Given: A /v1/responses-style request whose tools are a single
{"type": "mcp", "server_url": "litellm_proxy"} reference that
expands to 5 flat OpenAI function dicts
When: The hook processes the request
Then: The expanded tools go through the semantic filter (top_k=2)
and litellm_semantic_filter_stats reports pre/post counts, so
the x-litellm-semantic-filter header shows how many tools
were filtered out instead of silently forwarding all tools
with no stats.
"""
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
SemanticMCPToolFilter,
)
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
from litellm.types.utils import Embedding, EmbeddingResponse

mock_router = Mock()

def mock_embedding_sync(*args, **kwargs):
return EmbeddingResponse(
data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")],
model="text-embedding-3-small",
object="list",
usage={"prompt_tokens": 10, "total_tokens": 10},
)

async def mock_embedding_async(*args, **kwargs):
return mock_embedding_sync()

mock_router.embedding = mock_embedding_sync
mock_router.aembedding = mock_embedding_async

filter_instance = SemanticMCPToolFilter(
embedding_model="text-embedding-3-small",
litellm_router_instance=mock_router,
top_k=2,
similarity_threshold=0.3,
enabled=True,
)

registry_tools = [
MCPTool(
name=f"srv-tool_{i}",
description=f"Registry tool {i}",
inputSchema={"type": "object"},
)
for i in range(5)
]
filter_instance._build_router(registry_tools)

expanded_tools = [
{
"type": "function",
"name": f"srv-tool_{i}",
"description": f"Registry tool {i}",
"parameters": {"type": "object", "properties": {}},
}
for i in range(5)
]

hook = SemanticToolFilterHook(filter_instance)
hook._expand_mcp_tools = AsyncMock( # type: ignore[method-assign]
return_value=expanded_tools
)

data = {
"model": "gpt-4",
"input": [{"role": "user", "content": "Send an email", "type": "message"}],
"tools": [
{
"type": "mcp",
"server_url": "litellm_proxy",
"require_approval": "never",
}
],
"metadata": {},
}

result = await hook.async_pre_call_hook(
user_api_key_dict=Mock(),
cache=Mock(),
data=data,
call_type="aresponses",
)

assert result is not None, "Hook should return modified data"
filtered = result["tools"]

assert len(filtered) <= 2, f"Expanded tools should be filtered to top_k=2, got {len(filtered)}"
assert len(filtered) < len(expanded_tools), (
f"Hook must not forward all {len(expanded_tools)} expanded tools unfiltered, got {len(filtered)}"
)
for tool in filtered:
assert tool in expanded_tools, "Filtered tools must be the original expanded tool dicts"

assert (
"litellm_semantic_filter_stats" in result["metadata"]
), "Filter stats must be emitted for the litellm_proxy expansion path"
stats = result["metadata"]["litellm_semantic_filter_stats"]
total, selected = stats.split("->")
assert int(total) == 5, f"Stats 'from' should be pre-filter expanded count (5), got {total}"
assert int(selected) == len(filtered), f"Stats 'to' should match post-filter count, got {selected}"

print(f"✅ Expanded litellm_proxy tools filtered: {len(expanded_tools)} -> {len(filtered)}, stats={stats}")


@pytest.mark.asyncio
async def test_semantic_filter_hook_filters_expanded_tools_with_string_input():
"""
Responses API requests may pass ``input`` as a plain string; the
expanded-tool filtering must treat it as the user query instead of
crashing (which would silently disable MCP expansion).
"""
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
SemanticMCPToolFilter,
)
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
from litellm.types.utils import Embedding, EmbeddingResponse

mock_router = Mock()

def mock_embedding_sync(*args, **kwargs):
return EmbeddingResponse(
data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")],
model="text-embedding-3-small",
object="list",
usage={"prompt_tokens": 10, "total_tokens": 10},
)

async def mock_embedding_async(*args, **kwargs):
return mock_embedding_sync()

mock_router.embedding = mock_embedding_sync
mock_router.aembedding = mock_embedding_async

filter_instance = SemanticMCPToolFilter(
embedding_model="text-embedding-3-small",
litellm_router_instance=mock_router,
top_k=2,
similarity_threshold=0.3,
enabled=True,
)

registry_tools = [
MCPTool(
name=f"srv-tool_{i}",
description=f"Registry tool {i}",
inputSchema={"type": "object"},
)
for i in range(5)
]
filter_instance._build_router(registry_tools)

expanded_tools = [
{
"type": "function",
"name": f"srv-tool_{i}",
"description": f"Registry tool {i}",
"parameters": {"type": "object", "properties": {}},
}
for i in range(5)
]

hook = SemanticToolFilterHook(filter_instance)

filtered = await hook._filter_expanded_tools(
data={"input": "Send an email"},
expanded_tools=expanded_tools,
)

assert len(filtered) <= 2, f"String input must still drive semantic filtering, got {len(filtered)} tools"

print(f"✅ String input filtered expanded tools: {len(expanded_tools)} -> {len(filtered)}")


@pytest.mark.asyncio
async def test_semantic_filter_hook_expansion_skips_filter_when_disabled():
"""
When the filter is disabled at runtime (e.g. via the UI toggle), the
expansion path must forward all expanded tools and emit NO filter
stats, mirroring the generic path's enabled guard.
"""
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
SemanticMCPToolFilter,
)
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook

filter_instance = SemanticMCPToolFilter(
embedding_model="text-embedding-3-small",
litellm_router_instance=Mock(),
top_k=2,
similarity_threshold=0.3,
enabled=False,
)

expanded_tools = [
{
"type": "function",
"name": f"srv-tool_{i}",
"description": f"Registry tool {i}",
"parameters": {"type": "object", "properties": {}},
}
for i in range(5)
]

hook = SemanticToolFilterHook(filter_instance)
hook._expand_mcp_tools = AsyncMock( # type: ignore[method-assign]
return_value=expanded_tools
)

data = {
"model": "gpt-4",
"input": [{"role": "user", "content": "Send an email", "type": "message"}],
"tools": [
{
"type": "mcp",
"server_url": "litellm_proxy",
"require_approval": "never",
}
],
"metadata": {},
}

result = await hook.async_pre_call_hook(
user_api_key_dict=Mock(),
cache=Mock(),
data=data,
call_type="aresponses",
)

assert result is not None, "Hook should still expand MCP references when the filter is disabled"
assert len(result["tools"]) == 5, f"All expanded tools must be forwarded when disabled, got {len(result['tools'])}"
assert (
"litellm_semantic_filter_stats" not in result["metadata"]
), "No filter stats may be emitted when the filter is disabled"

print("✅ Disabled filter: expansion preserved, no spurious stats")


@pytest.mark.asyncio
async def test_semantic_filter_hook_preserves_tool_order():
"""
Expand Down
Loading
Loading