diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260716000000_add_mcp_tool_search_top_k/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260716000000_add_mcp_tool_search_top_k/migration.sql new file mode 100644 index 00000000000..94240b245ea --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260716000000_add_mcp_tool_search_top_k/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "mcp_tool_search_top_k" INTEGER; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index d9959677116..24bdd8c09e3 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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[] diff --git a/litellm/models/object_permission.py b/litellm/models/object_permission.py index a09d50ddc33..8236c64111d 100644 --- a/litellm/models/object_permission.py +++ b/litellm/models/object_permission.py @@ -23,3 +23,4 @@ class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase): blocked_tools: list[str] | None = [] search_tools: list[str] | None = [] mcp_tool_search_enabled: bool | None = None + mcp_tool_search_top_k: int | None = None diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 3a8fd6de5e5..94b85f193c9 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -169,9 +169,9 @@ if MCP_AVAILABLE: ) from litellm.proxy._experimental.mcp_server.tool_search import ( MCP_TOOL_SEARCH_TOOL_NAME, - coerce_top_k, 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 @@ -190,9 +190,15 @@ if MCP_AVAILABLE: ) = _extract_mcp_headers_from_request(request, MCPRequestHandler) virtual_oauth2_headers: Final = MCPRequestHandler._get_oauth2_headers_from_headers(request.headers) if tool_name == MCP_TOOL_SEARCH_TOOL_NAME: + raw_settings = proxy_config.get_config_state().get("litellm_settings") + settings = raw_settings if isinstance(raw_settings, Mapping) else None 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, + settings, + ), user_api_key_dict=user_api_key_dict, client_ip=rest_client_ip, mcp_auth_header=virtual_mcp_auth_header, @@ -759,11 +765,20 @@ 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 + raw_settings = proxy_config.get_config_state().get("litellm_settings") + settings = raw_settings if isinstance(raw_settings, Mapping) else None 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, + settings, + ) + ), "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 c6b2ac489bb..17b2614eccf 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -790,10 +790,22 @@ 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.model_validate(d) for d in get_virtual_tool_definitions()] + raw_settings = proxy_config.get_config_state().get("litellm_settings") + settings = raw_settings if isinstance(raw_settings, Mapping) else None + return [ + Tool.model_validate(d) + for d in get_virtual_tool_definitions( + default_top_k=get_mcp_tool_search_default_top_k( + user_api_key_auth, + settings, + ) + ) + ] # Get mcp_servers from context variable verbose_logger.debug("MCP list_tools - Calling _list_mcp_tools") @@ -914,9 +926,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): @@ -939,9 +951,17 @@ if MCP_AVAILABLE: args: Final = arguments or {} if name == MCP_TOOL_SEARCH_TOOL_NAME: + from litellm.proxy.proxy_server import proxy_config + + raw_settings = proxy_config.get_config_state().get("litellm_settings") + settings = raw_settings if isinstance(raw_settings, Mapping) else None 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, + settings, + ), user_api_key_dict=user_api_key_auth, client_ip=client_ip, mcp_servers=mcp_servers, diff --git a/litellm/proxy/_experimental/mcp_server/tool_search.py b/litellm/proxy/_experimental/mcp_server/tool_search.py index 3b0dd2071ae..226a54bbf5c 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, Final @@ -12,16 +13,55 @@ if TYPE_CHECKING: MCP_TOOL_SEARCH_TOOL_NAME: Final[str] = "mcp_tool_search" MCP_TOOL_CALL_TOOL_NAME: Final[str] = "mcp_tool_call" +DEFAULT_MCP_TOOL_SEARCH_TOP_K: int = 5 -def coerce_top_k(value: Any, default: int = 5) -> int: +def get_mcp_tool_search_default_top_k( + user_api_key_dict: UserAPIKeyAuth | None = None, + litellm_settings: Mapping[str, object] | None = 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) + + if litellm_settings is None: + global_top_k = None + else: + global_top_k = 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: object, + user_api_key_dict: UserAPIKeyAuth | None = None, + litellm_settings: Mapping[str, object] | None = None, +) -> int: + default_top_k = get_mcp_tool_search_default_top_k(user_api_key_dict, litellm_settings) + 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: object, default: int = DEFAULT_MCP_TOOL_SEARCH_TOP_K) -> int: + if not isinstance(value, (int, float, str, bytes, bytearray)): + return default try: - return int(value) + result = int(value) except (TypeError, ValueError): return default + return result if result > 0 else 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: Final = query.lower().split() @@ -34,7 +74,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 +91,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/litellm/proxy/_types.py b/litellm/proxy/_types.py index ed49ca2caa9..f69dfb87f0b 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1098,6 +1098,7 @@ class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase): models: list[str] | None = None search_tools: list[str] | None = None mcp_tool_search_enabled: bool | None = None + mcp_tool_search_top_k: int | None = None from litellm.models.team import BudgetLimitEntry as BudgetLimitEntry # noqa: E402 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index d9959677116..24bdd8c09e3 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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[] diff --git a/litellm/types/object_permission.py b/litellm/types/object_permission.py index 1b391a3a1ef..52cf9f06422 100644 --- a/litellm/types/object_permission.py +++ b/litellm/types/object_permission.py @@ -23,3 +23,4 @@ class ObjectPermissionDict(TypedDict, total=False): models: list[str] | None search_tools: list[str] | None mcp_tool_search_enabled: bool | None + mcp_tool_search_top_k: int | None # writable-ok: mutated on in-memory dict payloads before persistence diff --git a/schema.prisma b/schema.prisma index d9959677116..24bdd8c09e3 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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[] 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 0e442102e53..6ef9ad4d475 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 @@ -20,8 +20,11 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import Aggregat 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 @@ -69,10 +72,50 @@ class TestCoerceTopK: def test_none_returns_default(self) -> None: assert coerce_top_k(None) == 5 + @pytest.mark.parametrize("value", [0, -1, "-3"]) + def test_non_positive_value_returns_default(self, value: Any) -> None: + assert coerce_top_k(value) == 5 + def test_custom_default(self) -> None: 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) -> None: + assert get_mcp_tool_search_default_top_k(litellm_settings={"mcp_tool_search_default_top_k": 12}) == 12 + + def test_per_key_beats_global(self) -> None: + uak = UserAPIKeyAuth( + api_key="k", + object_permission=_make_perm(mcp_tool_search_top_k=8), + ) + assert get_mcp_tool_search_default_top_k(uak, {"mcp_tool_search_default_top_k": 12}) == 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) @@ -140,6 +183,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" diff --git a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py index 5c163c44cb3..2710c36da00 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py @@ -719,6 +719,7 @@ _EXPECTED_CUSTOMER = { "blocked_tools": [], "search_tools": [], "mcp_tool_search_enabled": None, + "mcp_tool_search_top_k": None, }, } diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 25fbd53018a..961c12edb17 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -28162,6 +28162,8 @@ export interface components { } | null; /** Mcp Tool Search Enabled */ mcp_tool_search_enabled?: boolean | null; + /** Mcp Tool Search Top K */ + mcp_tool_search_top_k?: number | null; /** Mcp Toolsets */ mcp_toolsets?: string[] | null; /** Models */ @@ -28207,6 +28209,8 @@ export interface components { } | null; /** Mcp Tool Search Enabled */ mcp_tool_search_enabled?: boolean | null; + /** Mcp Tool Search Top K */ + mcp_tool_search_top_k?: number | null; /** Mcp Toolsets */ mcp_toolsets?: string[] | null; /**