diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 111fde86ea0..8b41735dd7b 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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", } diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 68a61b85175..c043d66103d 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/tool_search.py b/litellm/proxy/_experimental/mcp_server/tool_search.py index fa57a2b3eb2..23f5a19456b 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_search.py +++ b/litellm/proxy/_experimental/mcp_server/tool_search.py @@ -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"], diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py index b8f0b205831..bfefe764f69 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -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: