mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge dd4d2f44e3 into 44d84360fb
This commit is contained in:
commit
0fa7f930c5
13 changed files with 149 additions and 11 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[]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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[]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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[]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -719,6 +719,7 @@ _EXPECTED_CUSTOMER = {
|
|||
"blocked_tools": [],
|
||||
"search_tools": [],
|
||||
"mcp_tool_search_enabled": None,
|
||||
"mcp_tool_search_top_k": None,
|
||||
},
|
||||
}
|
||||
|
||||
|
|
|
|||
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
/**
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue