fix(mcp): apply semantic filter to expanded litellm_proxy tools and show filtered-out count (#32285)

* fix(mcp): apply semantic filter to expanded litellm_proxy tools and show filtered-out count

* fix(mcp): guard expansion-path filtering on enabled flag and isolate metadata emission
This commit is contained in:
tin-berri 2026-07-07 09:23:48 -07:00 • committed by GitHub
parent 65be4c16cd
commit 6041d37414
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 334 additions and 24 deletions

View file

@ -129,6 +129,27 @@ class SemanticToolFilterHook(CustomLogger):
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.
@ -184,6 +205,32 @@ class SemanticToolFilterHook(CustomLogger):
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",
@ -206,9 +253,6 @@ class SemanticToolFilterHook(CustomLogger):
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")
@ -227,11 +271,26 @@ class SemanticToolFilterHook(CustomLogger):
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
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
@ -297,19 +356,13 @@ class SemanticToolFilterHook(CustomLogger):
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

View file

@ -773,6 +773,251 @@ async def test_semantic_filter_hook_responses_api_name_collision():
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():
"""

View file

@ -90,7 +90,7 @@ describe("MCPSemanticFilterTestPanel", () => {
expect(screen.queryByText("Semantic filtering is disabled")).not.toBeInTheDocument();
});
it("should display test results when testResult is provided", () => {
it("should display selected and filtered-out counts when testResult is provided", () => {
const testResult: TestResult = {
totalTools: 10,
selectedTools: 3,
@ -98,8 +98,8 @@ describe("MCPSemanticFilterTestPanel", () => {
};
render(<MCPSemanticFilterTestPanel {...buildProps({ testResult })} />);
expect(screen.getByText("3 tools selected")).toBeInTheDocument();
expect(screen.getByText("Filtered from 10 available tools")).toBeInTheDocument();
expect(screen.getByText("3 of 10 tools selected")).toBeInTheDocument();
expect(screen.getByText("7 tools filtered out")).toBeInTheDocument();
expect(screen.getByText("wiki-fetch")).toBeInTheDocument();
expect(screen.getByText("github-search")).toBeInTheDocument();
expect(screen.getByText("slack-post")).toBeInTheDocument();
@ -117,6 +117,18 @@ describe("MCPSemanticFilterTestPanel", () => {
expect(screen.getByText("+5 more selected tools not shown")).toBeInTheDocument();
});
it("should surface a zero filtered-out count when the filter selected every tool", () => {
const testResult: TestResult = {
totalTools: 207,
selectedTools: 207,
tools: ["tool-a", "tool-b"],
};
render(<MCPSemanticFilterTestPanel {...buildProps({ testResult })} />);
expect(screen.getByText("207 of 207 tools selected")).toBeInTheDocument();
expect(screen.getByText("0 tools filtered out")).toBeInTheDocument();
});
it("should not render the results section when testResult is null", () => {
render(<MCPSemanticFilterTestPanel {...buildProps({ testResult: null })} />);
expect(screen.queryByText("Results")).not.toBeInTheDocument();

View file

@ -86,9 +86,9 @@ export default function MCPSemanticFilterTestPanel({
<div>
<Typography.Title level={5}>Results</Typography.Title>
<Alert
type="success"
message={`${testResult.selectedTools} tools selected`}
description={`Filtered from ${testResult.totalTools} available tools`}
type={testResult.totalTools - testResult.selectedTools > 0 ? "success" : "warning"}
message={`${testResult.selectedTools} of ${testResult.totalTools} tools selected`}
description={`${testResult.totalTools - testResult.selectedTools} tools filtered out`}
showIcon
style={{ marginBottom: 16 }}
/>