fix(mcp): make mcp_tool_search default top_k configurable

Allow operators to set a global litellm_settings default or per-key
object_permission override instead of always defaulting to 5.

Fixes #33440
This commit is contained in:
Hashim1999164 2026-07-16 01:28:59 +05:00
parent f1f33f560f
commit 5d19cb6ee2
11 changed files with 129 additions and 10 deletions

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "mcp_tool_search_top_k" INTEGER;

View file

@ -280,6 +280,7 @@ model LiteLLM_ObjectPermissionTable {
mcp_toolsets String[] @default([]) // Toolset IDs granted to this key/team/user
search_tools String[] @default([]) // search_tool_name values this key/team/user may call
mcp_tool_search_enabled Boolean?
mcp_tool_search_top_k Int?
teams LiteLLM_TeamTable[]
projects LiteLLM_ProjectTable[]
verification_tokens LiteLLM_VerificationToken[]

View file

@ -25,3 +25,4 @@ class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase):
blocked_tools: Optional[List[str]] = []
search_tools: Optional[List[str]] = []
mcp_tool_search_enabled: Optional[bool] = None
mcp_tool_search_top_k: Optional[int] = None

View file

@ -139,9 +139,11 @@ 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,
get_virtual_tool_definitions,
handle_mcp_tool_call,
handle_mcp_tool_search,
resolve_mcp_tool_search_top_k,
)
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.proxy_server import general_settings, proxy_config, proxy_logging_obj
@ -162,7 +164,7 @@ if MCP_AVAILABLE:
if tool_name == MCP_TOOL_SEARCH_TOOL_NAME:
return await handle_mcp_tool_search(
query=tool_arguments.get("query", ""),
top_k=coerce_top_k(tool_arguments.get("top_k", 5)),
top_k=resolve_mcp_tool_search_top_k(tool_arguments.get("top_k"), user_api_key_dict),
user_api_key_dict=user_api_key_dict,
client_ip=rest_client_ip,
mcp_auth_header=virtual_mcp_auth_header,
@ -733,11 +735,14 @@ if MCP_AVAILABLE:
)
):
from litellm.proxy._experimental.mcp_server.tool_search import (
get_mcp_tool_search_default_top_k,
get_virtual_tool_definitions,
)
return {
"tools": get_virtual_tool_definitions(),
"tools": get_virtual_tool_definitions(
default_top_k=get_mcp_tool_search_default_top_k(user_api_key_dict)
),
"error": None,
"message": "Successfully retrieved tools",
}

View file

@ -702,10 +702,16 @@ 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,
)
return [Tool(**d) for d in get_virtual_tool_definitions()]
return [
Tool(**d)
for d in get_virtual_tool_definitions(
default_top_k=get_mcp_tool_search_default_top_k(user_api_key_auth)
)
]
# Get mcp_servers from context variable
verbose_logger.debug("MCP list_tools - Calling _list_mcp_tools")
@ -821,9 +827,9 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.tool_search import (
MCP_TOOL_CALL_TOOL_NAME,
MCP_TOOL_SEARCH_TOOL_NAME,
coerce_top_k,
handle_mcp_tool_call,
handle_mcp_tool_search,
resolve_mcp_tool_search_top_k,
)
if name not in (MCP_TOOL_SEARCH_TOOL_NAME, MCP_TOOL_CALL_TOOL_NAME):
@ -848,7 +854,7 @@ 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=resolve_mcp_tool_search_top_k(args.get("top_k"), user_api_key_auth),
user_api_key_dict=user_api_key_auth,
client_ip=client_ip,
mcp_servers=mcp_servers,

View file

@ -12,16 +12,54 @@ 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 _get_litellm_settings() -> dict[str, Any]:
try:
from litellm.proxy.proxy_server import proxy_config
return proxy_config.get_config_state().get("litellm_settings") or {}
except Exception:
return {}
def get_mcp_tool_search_default_top_k(
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
) -> int:
"""Resolve the default top_k for mcp_tool_search (per-key, then global, then 5)."""
if user_api_key_dict is not None:
object_permission = getattr(user_api_key_dict, "object_permission", None)
if object_permission is not None:
key_top_k = getattr(object_permission, "mcp_tool_search_top_k", None)
if key_top_k is not None:
return coerce_top_k(key_top_k, default=DEFAULT_MCP_TOOL_SEARCH_TOP_K)
global_top_k = _get_litellm_settings().get("mcp_tool_search_default_top_k")
if global_top_k is not None:
return coerce_top_k(global_top_k, default=DEFAULT_MCP_TOOL_SEARCH_TOP_K)
return DEFAULT_MCP_TOOL_SEARCH_TOP_K
def resolve_mcp_tool_search_top_k(
explicit_top_k: Any,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
) -> int:
default_top_k = get_mcp_tool_search_default_top_k(user_api_key_dict)
if explicit_top_k is None:
return default_top_k
return coerce_top_k(explicit_top_k, default=default_top_k)
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 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 +72,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 +89,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

@ -1009,6 +1009,7 @@ class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase):
models: Optional[List[str]] = None
search_tools: Optional[List[str]] = None
mcp_tool_search_enabled: Optional[bool] = None
mcp_tool_search_top_k: Optional[int] = None
from litellm.types.object_permission import ( # noqa: E402

View file

@ -280,6 +280,7 @@ model LiteLLM_ObjectPermissionTable {
mcp_toolsets String[] @default([]) // Toolset IDs granted to this key/team/user
search_tools String[] @default([]) // search_tool_name values this key/team/user may call
mcp_tool_search_enabled Boolean?
mcp_tool_search_top_k Int?
teams LiteLLM_TeamTable[]
projects LiteLLM_ProjectTable[]
verification_tokens LiteLLM_VerificationToken[]

View file

@ -25,3 +25,4 @@ class ObjectPermissionDict(TypedDict, total=False):
models: Optional[list[str]]
search_tools: Optional[list[str]]
mcp_tool_search_enabled: Optional[bool]
mcp_tool_search_top_k: Optional[int]

View file

@ -280,6 +280,7 @@ model LiteLLM_ObjectPermissionTable {
mcp_toolsets String[] @default([]) // Toolset IDs granted to this key/team/user
search_tools String[] @default([]) // search_tool_name values this key/team/user may call
mcp_tool_search_enabled Boolean?
mcp_tool_search_top_k Int?
teams LiteLLM_TeamTable[]
projects LiteLLM_ProjectTable[]
verification_tokens LiteLLM_VerificationToken[]

View file

@ -19,8 +19,11 @@ from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
from litellm.proxy._experimental.mcp_server.tool_search import (
MCP_TOOL_CALL_TOOL_NAME,
MCP_TOOL_SEARCH_TOOL_NAME,
DEFAULT_MCP_TOOL_SEARCH_TOP_K,
coerce_top_k,
get_mcp_tool_search_default_top_k,
get_virtual_tool_definitions,
resolve_mcp_tool_search_top_k,
search_tools,
)
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
@ -72,6 +75,58 @@ class TestCoerceTopK:
assert coerce_top_k("nope", default=10) == 10
class TestMcpToolSearchDefaultTopK:
def test_builtin_default(self) -> None:
assert get_mcp_tool_search_default_top_k() == DEFAULT_MCP_TOOL_SEARCH_TOP_K
def test_per_key_override(self) -> None:
uak = UserAPIKeyAuth(
api_key="k",
object_permission=_make_perm(mcp_tool_search_top_k=10),
)
assert get_mcp_tool_search_default_top_k(uak) == 10
def test_global_litellm_settings_override(self, monkeypatch: pytest.MonkeyPatch) -> None:
mock_config = MagicMock()
mock_config.get_config_state.return_value = {
"litellm_settings": {"mcp_tool_search_default_top_k": 12}
}
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_config",
mock_config,
)
assert get_mcp_tool_search_default_top_k() == 12
def test_per_key_beats_global(self, monkeypatch: pytest.MonkeyPatch) -> None:
mock_config = MagicMock()
mock_config.get_config_state.return_value = {
"litellm_settings": {"mcp_tool_search_default_top_k": 12}
}
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_config",
mock_config,
)
uak = UserAPIKeyAuth(
api_key="k",
object_permission=_make_perm(mcp_tool_search_top_k=8),
)
assert get_mcp_tool_search_default_top_k(uak) == 8
def test_resolve_uses_explicit_top_k(self) -> None:
uak = UserAPIKeyAuth(
api_key="k",
object_permission=_make_perm(mcp_tool_search_top_k=10),
)
assert resolve_mcp_tool_search_top_k(3, uak) == 3
def test_resolve_uses_default_when_omitted(self) -> None:
uak = UserAPIKeyAuth(
api_key="k",
object_permission=_make_perm(mcp_tool_search_top_k=10),
)
assert resolve_mcp_tool_search_top_k(None, uak) == 10
class TestSearchTools:
def test_returns_matching_tools(self) -> None:
results = search_tools("github issue", SAMPLE_TOOLS)
@ -139,6 +194,11 @@ class TestGetVirtualToolDefinitions:
assert "arguments" in props
assert "tool_name" in call_tool["inputSchema"]["required"]
def test_mcp_tool_search_schema_top_k_default(self) -> None:
tools = get_virtual_tool_definitions(default_top_k=10)
search_tool = next(t for t in tools if t["name"] == MCP_TOOL_SEARCH_TOOL_NAME)
assert search_tool["inputSchema"]["properties"]["top_k"]["default"] == 10
def test_all_tools_have_description(self) -> None:
for tool in get_virtual_tool_definitions():
assert tool.get("description"), f"{tool['name']} missing description"