mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(mcp): configure tool search default top k
This commit is contained in:
parent
24a438adfd
commit
12a579c9fa
4 changed files with 173 additions and 18 deletions
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue