mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
f1f33f560f
commit
5d19cb6ee2
11 changed files with 129 additions and 10 deletions
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "mcp_tool_search_top_k" INTEGER;
|
||||
|
|
@ -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[]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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[]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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[]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue