fix(mcp): configure tool search default top k

This commit is contained in:
Devin AI 2026-07-15 21:14:24 +00:00
parent 24a438adfd
commit 12a579c9fa
4 changed files with 173 additions and 18 deletions

View file

@ -140,6 +140,7 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.tool_search import (
MCP_TOOL_SEARCH_TOOL_NAME,
coerce_top_k,
get_mcp_tool_search_default_top_k,
handle_mcp_tool_call,
handle_mcp_tool_search,
)
@ -160,9 +161,13 @@ if MCP_AVAILABLE:
) = _extract_mcp_headers_from_request(request, MCPRequestHandler)
virtual_oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(request.headers)
if tool_name == MCP_TOOL_SEARCH_TOOL_NAME:
default_top_k = get_mcp_tool_search_default_top_k(proxy_config.get_config_state().get("litellm_settings"))
return await handle_mcp_tool_search(
query=tool_arguments.get("query", ""),
top_k=coerce_top_k(tool_arguments.get("top_k", 5)),
top_k=coerce_top_k(
tool_arguments.get("top_k"),
default=default_top_k,
),
user_api_key_dict=user_api_key_dict,
client_ip=rest_client_ip,
mcp_auth_header=virtual_mcp_auth_header,
@ -733,11 +738,16 @@ if MCP_AVAILABLE:
)
):
from litellm.proxy._experimental.mcp_server.tool_search import (
get_mcp_tool_search_default_top_k,
get_virtual_tool_definitions,
)
from litellm.proxy.proxy_server import proxy_config
default_top_k = get_mcp_tool_search_default_top_k(
proxy_config.get_config_state().get("litellm_settings")
)
return {
"tools": get_virtual_tool_definitions(),
"tools": get_virtual_tool_definitions(default_top_k=default_top_k),
"error": None,
"message": "Successfully retrieved tools",
}

View file

@ -702,10 +702,15 @@ if MCP_AVAILABLE:
from mcp.types import Tool
from litellm.proxy._experimental.mcp_server.tool_search import (
get_mcp_tool_search_default_top_k,
get_virtual_tool_definitions,
)
from litellm.proxy.proxy_server import proxy_config
return [Tool(**d) for d in get_virtual_tool_definitions()]
default_top_k = get_mcp_tool_search_default_top_k(
proxy_config.get_config_state().get("litellm_settings")
)
return [Tool(**definition) for definition in get_virtual_tool_definitions(default_top_k=default_top_k)]
# Get mcp_servers from context variable
verbose_logger.debug("MCP list_tools - Calling _list_mcp_tools")
@ -812,6 +817,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers: Optional[dict[str, dict[str, str]]] = None,
oauth2_headers: Optional[dict[str, str]] = None,
raw_headers: Optional[dict[str, str]] = None,
default_top_k: Optional[int] = None,
) -> Optional[CallToolResult]:
"""Handle the mcp_tool_search / mcp_tool_call virtual tools.
@ -819,6 +825,7 @@ if MCP_AVAILABLE:
the caller falls through to normal tool routing.
"""
from litellm.proxy._experimental.mcp_server.tool_search import (
DEFAULT_MCP_TOOL_SEARCH_TOP_K,
MCP_TOOL_CALL_TOOL_NAME,
MCP_TOOL_SEARCH_TOOL_NAME,
coerce_top_k,
@ -848,7 +855,10 @@ if MCP_AVAILABLE:
if name == MCP_TOOL_SEARCH_TOOL_NAME:
return await handle_mcp_tool_search(
query=args.get("query", ""),
top_k=coerce_top_k(args.get("top_k", 5)),
top_k=coerce_top_k(
args.get("top_k"),
default=(default_top_k if default_top_k is not None else DEFAULT_MCP_TOOL_SEARCH_TOP_K),
),
user_api_key_dict=user_api_key_auth,
client_ip=client_ip,
mcp_servers=mcp_servers,
@ -892,6 +902,9 @@ if MCP_AVAILABLE:
from mcp.types import CallToolResult
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
from litellm.proxy._experimental.mcp_server.tool_search import (
get_mcp_tool_search_default_top_k,
)
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.proxy_server import proxy_config
@ -932,6 +945,9 @@ if MCP_AVAILABLE:
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
default_top_k=get_mcp_tool_search_default_top_k(
proxy_config.get_config_state().get("litellm_settings")
),
)
if virtual_tool_result is not None:
return virtual_tool_result

View file

@ -1,6 +1,7 @@
from __future__ import annotations
import json
from collections.abc import Mapping
from datetime import datetime
from typing import TYPE_CHECKING, Any, Optional
@ -12,16 +13,32 @@ if TYPE_CHECKING:
MCP_TOOL_SEARCH_TOOL_NAME: str = "mcp_tool_search"
MCP_TOOL_CALL_TOOL_NAME: str = "mcp_tool_call"
DEFAULT_MCP_TOOL_SEARCH_TOP_K: int = 5
def coerce_top_k(value: Any, default: int = 5) -> int:
def coerce_top_k(value: Any, default: int = DEFAULT_MCP_TOOL_SEARCH_TOP_K) -> int:
try:
return int(value)
except (TypeError, ValueError):
return default
def search_tools(query: str, tools: list[dict[str, Any]], top_k: int = 5) -> list[dict[str, Any]]:
def get_mcp_tool_search_default_top_k(
litellm_settings: Optional[Mapping[str, object]],
) -> int:
if litellm_settings is None:
return DEFAULT_MCP_TOOL_SEARCH_TOP_K
return coerce_top_k(
litellm_settings.get("mcp_tool_search_default_top_k"),
default=DEFAULT_MCP_TOOL_SEARCH_TOP_K,
)
def search_tools(
query: str,
tools: list[dict[str, Any]],
top_k: int = DEFAULT_MCP_TOOL_SEARCH_TOP_K,
) -> list[dict[str, Any]]:
if not query:
return []
tokens = query.lower().split()
@ -34,7 +51,9 @@ def search_tools(query: str, tools: list[dict[str, Any]], top_k: int = 5) -> lis
return [tool for _, tool in sorted(scored, key=lambda x: x[0], reverse=True)[:top_k]]
def get_virtual_tool_definitions() -> list[dict[str, Any]]:
def get_virtual_tool_definitions(
default_top_k: int = DEFAULT_MCP_TOOL_SEARCH_TOP_K,
) -> list[dict[str, Any]]:
return [
{
"name": MCP_TOOL_SEARCH_TOOL_NAME,
@ -49,7 +68,7 @@ def get_virtual_tool_definitions() -> list[dict[str, Any]]:
"top_k": {
"type": "integer",
"description": "Maximum number of results to return.",
"default": 5,
"default": default_top_k,
},
},
"required": ["query"],

View file

@ -17,9 +17,11 @@ import pytest
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
from litellm.proxy._experimental.mcp_server.tool_search import (
DEFAULT_MCP_TOOL_SEARCH_TOP_K,
MCP_TOOL_CALL_TOOL_NAME,
MCP_TOOL_SEARCH_TOOL_NAME,
coerce_top_k,
get_mcp_tool_search_default_top_k,
get_virtual_tool_definitions,
search_tools,
)
@ -72,6 +74,20 @@ class TestCoerceTopK:
assert coerce_top_k("nope", default=10) == 10
class TestMcpToolSearchDefaultTopK:
def test_uses_builtin_default_when_unconfigured(self) -> None:
assert get_mcp_tool_search_default_top_k(None) == DEFAULT_MCP_TOOL_SEARCH_TOP_K
def test_uses_litellm_settings_default(self) -> None:
assert get_mcp_tool_search_default_top_k({"mcp_tool_search_default_top_k": 10}) == 10
def test_invalid_litellm_settings_default_uses_builtin_default(self) -> None:
assert (
get_mcp_tool_search_default_top_k({"mcp_tool_search_default_top_k": "invalid"})
== DEFAULT_MCP_TOOL_SEARCH_TOP_K
)
class TestSearchTools:
def test_returns_matching_tools(self) -> None:
results = search_tools("github issue", SAMPLE_TOOLS)
@ -131,6 +147,11 @@ class TestGetVirtualToolDefinitions:
assert "query" in props
assert search_tool["inputSchema"]["required"] == ["query"]
def test_mcp_tool_search_schema_uses_configured_default(self) -> None:
tools = get_virtual_tool_definitions(default_top_k=10)
search_tool = next(tool for tool in tools if tool["name"] == MCP_TOOL_SEARCH_TOOL_NAME)
assert search_tool["inputSchema"]["properties"]["top_k"]["default"] == 10
def test_mcp_tool_call_schema_has_tool_name_and_arguments(self) -> None:
tools = get_virtual_tool_definitions()
call_tool = next(t for t in tools if t["name"] == MCP_TOOL_CALL_TOOL_NAME)
@ -177,16 +198,22 @@ class TestListToolRestApiWithToolSearch:
if hasattr(r, "path") and r.path.endswith("/tools/list") and hasattr(r, "methods") and "GET" in r.methods
)
result = await list_fn(
request=mock_request,
server_id=None,
include_disabled_tools=False,
user_api_key_dict=user_api_key_dict,
)
with patch(
"litellm.proxy.proxy_server.proxy_config.get_config_state",
return_value={"litellm_settings": {"mcp_tool_search_default_top_k": 10}},
):
result = await list_fn(
request=mock_request,
server_id=None,
include_disabled_tools=False,
user_api_key_dict=user_api_key_dict,
)
assert result["error"] is None
tool_names = [t["name"] for t in result["tools"]]
assert set(tool_names) == {MCP_TOOL_SEARCH_TOOL_NAME, MCP_TOOL_CALL_TOOL_NAME}
search_tool = next(tool for tool in result["tools"] if tool["name"] == MCP_TOOL_SEARCH_TOOL_NAME)
assert search_tool["inputSchema"]["properties"]["top_k"]["default"] == 10
@pytest.mark.asyncio
async def test_returns_full_catalog_when_flag_disabled(self) -> None:
@ -394,6 +421,35 @@ class TestCallToolRestApiVirtualTools:
assert isinstance(returned_tools, list)
assert any(t["name"] == "github-create_issue" for t in returned_tools)
@pytest.mark.asyncio
async def test_mcp_tool_search_call_uses_configured_default_top_k(
self,
) -> None:
user_api_key_dict = UserAPIKeyAuth(
api_key="test_key",
object_permission=_make_perm(mcp_tool_search_enabled=True),
)
request = self._make_request({"name": MCP_TOOL_SEARCH_TOOL_NAME, "arguments": {"query": "issue"}})
with (
patch(
"litellm.proxy.proxy_server.proxy_config.get_config_state",
return_value={"litellm_settings": {"mcp_tool_search_default_top_k": 10}},
),
patch(
"litellm.proxy._experimental.mcp_server.tool_search.handle_mcp_tool_search",
new_callable=AsyncMock,
return_value="SEARCH_RESULT",
) as mock_search,
):
result = await self._get_call_fn()(
request=request,
user_api_key_dict=user_api_key_dict,
)
assert result == "SEARCH_RESULT"
assert mock_search.await_args.kwargs["top_k"] == 10
@pytest.mark.asyncio
async def test_mcp_tool_call_executes_discovered_tool(self) -> None:
from mcp.types import CallToolResult, TextContent
@ -591,6 +647,52 @@ class TestDispatchVirtualMcpTool:
assert mock_search.await_args.kwargs["query"] == "q"
assert mock_search.await_args.kwargs["top_k"] == 3
@pytest.mark.asyncio
async def test_routes_search_with_configured_default_top_k(self) -> None:
from litellm.proxy._experimental.mcp_server import server as srv
uak = UserAPIKeyAuth(
api_key="k",
object_permission=_make_perm(mcp_tool_search_enabled=True),
)
with patch(
"litellm.proxy._experimental.mcp_server.tool_search.handle_mcp_tool_search",
new_callable=AsyncMock,
return_value="SEARCH_RESULT",
) as mock_search:
await srv._dispatch_virtual_mcp_tool(
name=MCP_TOOL_SEARCH_TOOL_NAME,
arguments={"query": "q"},
user_api_key_auth=uak,
client_ip=None,
default_top_k=10,
)
assert mock_search.await_args.kwargs["top_k"] == 10
@pytest.mark.asyncio
async def test_explicit_top_k_overrides_configured_default(self) -> None:
from litellm.proxy._experimental.mcp_server import server as srv
uak = UserAPIKeyAuth(
api_key="k",
object_permission=_make_perm(mcp_tool_search_enabled=True),
)
with patch(
"litellm.proxy._experimental.mcp_server.tool_search.handle_mcp_tool_search",
new_callable=AsyncMock,
return_value="SEARCH_RESULT",
) as mock_search:
await srv._dispatch_virtual_mcp_tool(
name=MCP_TOOL_SEARCH_TOOL_NAME,
arguments={"query": "q", "top_k": 3},
user_api_key_auth=uak,
client_ip=None,
default_top_k=10,
)
assert mock_search.await_args.kwargs["top_k"] == 3
@pytest.mark.asyncio
async def test_routes_call_with_client_ip(self) -> None:
from litellm.proxy._experimental.mcp_server import server as srv
@ -839,10 +941,16 @@ class TestHandleListToolsVirtual:
from litellm.proxy._experimental.mcp_server import server as srv
uak = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=True))
with patch(
"litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context",
new_callable=AsyncMock,
return_value=(uak, None, None, None, None, None, None),
with (
patch(
"litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context",
new_callable=AsyncMock,
return_value=(uak, None, None, None, None, None, None),
),
patch(
"litellm.proxy.proxy_server.proxy_config.get_config_state",
return_value={"litellm_settings": {"mcp_tool_search_default_top_k": 10}},
),
):
tools = await srv.handle_list_tools()
@ -850,6 +958,8 @@ class TestHandleListToolsVirtual:
MCP_TOOL_SEARCH_TOOL_NAME,
MCP_TOOL_CALL_TOOL_NAME,
}
search_tool = next(tool for tool in tools if tool.name == MCP_TOOL_SEARCH_TOOL_NAME)
assert search_tool.inputSchema["properties"]["top_k"]["default"] == 10
class TestMcpServerToolCallErrorHandling: