mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
65be4c16cd
commit
6041d37414
4 changed files with 334 additions and 24 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
/>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue