mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: keep mcp catalog uncompressed
This commit is contained in:
parent
7a4bc18271
commit
9221010a1f
4 changed files with 89 additions and 1 deletions
|
|
@ -36,7 +36,7 @@ async def get_platform_mcp_settings() -> tuple[bool, int]:
|
|||
if isinstance(param_value, dict):
|
||||
settings.update(param_value)
|
||||
|
||||
enabled = settings.get("platform_mcp_enabled") is True
|
||||
enabled = _coerce_enabled(settings.get("platform_mcp_enabled"))
|
||||
threshold = _coerce_positive_threshold(settings.get("platform_mcp_tool_threshold"))
|
||||
return enabled, threshold
|
||||
|
||||
|
|
@ -201,11 +201,27 @@ def serialize_tool(tool: Any) -> dict[str, Any]:
|
|||
|
||||
|
||||
def _coerce_positive_threshold(value: Any) -> int:
|
||||
if isinstance(value, bool):
|
||||
return DEFAULT_PLATFORM_MCP_TOOL_THRESHOLD
|
||||
if isinstance(value, int) and value > 0:
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
value = value.strip()
|
||||
if value.isdigit():
|
||||
threshold = int(value)
|
||||
if threshold > 0:
|
||||
return threshold
|
||||
return DEFAULT_PLATFORM_MCP_TOOL_THRESHOLD
|
||||
|
||||
|
||||
def _coerce_enabled(value: Any) -> bool:
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
return value.strip().lower() in {"1", "true", "yes", "on"}
|
||||
return False
|
||||
|
||||
|
||||
def _server_match_name(server: MCPServer) -> str:
|
||||
return server.alias or server.server_name or server.name
|
||||
|
||||
|
|
|
|||
|
|
@ -2182,6 +2182,7 @@ if MCP_AVAILABLE:
|
|||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
log_list_tools_to_spendlogs: bool = False,
|
||||
list_tools_log_source: Optional[str] = None,
|
||||
enable_platform_mcp_compression: bool = True,
|
||||
) -> List[MCPTool]:
|
||||
"""
|
||||
List all available MCP tools.
|
||||
|
|
@ -2212,6 +2213,9 @@ if MCP_AVAILABLE:
|
|||
platform_mcp_enabled, platform_mcp_threshold = (
|
||||
await get_platform_mcp_settings()
|
||||
)
|
||||
platform_mcp_enabled = (
|
||||
platform_mcp_enabled and enable_platform_mcp_compression
|
||||
)
|
||||
enabled_server_names = get_enabled_server_names_for_session(
|
||||
get_active_mcp_session()
|
||||
)
|
||||
|
|
|
|||
|
|
@ -734,6 +734,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header=None,
|
||||
mcp_servers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
enable_platform_mcp_compression=False,
|
||||
)
|
||||
dumped_tools = [dict(tool) for tool in tools]
|
||||
|
||||
|
|
|
|||
|
|
@ -99,6 +99,43 @@ def test_platform_mcp_advertises_tool_list_changed_capability():
|
|||
assert options.capabilities.tools.listChanged is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_platform_mcp_settings_accept_config_field_string_values(monkeypatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setattr(
|
||||
proxy_server,
|
||||
"general_settings",
|
||||
{
|
||||
"platform_mcp_enabled": "true",
|
||||
"platform_mcp_tool_threshold": "12",
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
|
||||
assert await platform_mcp.get_platform_mcp_settings() == (True, 12)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_platform_mcp_threshold_ignores_bool_values(monkeypatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setattr(
|
||||
proxy_server,
|
||||
"general_settings",
|
||||
{
|
||||
"platform_mcp_enabled": True,
|
||||
"platform_mcp_tool_threshold": True,
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
|
||||
assert await platform_mcp.get_platform_mcp_settings() == (
|
||||
True,
|
||||
platform_mcp.DEFAULT_PLATFORM_MCP_TOOL_THRESHOLD,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_platform_mcp_compresses_aggregate_tools_over_threshold(monkeypatch):
|
||||
normal_tools = [_tool(f"tool_{idx}") for idx in range(11)]
|
||||
|
|
@ -127,6 +164,36 @@ async def test_platform_mcp_compresses_aggregate_tools_over_threshold(monkeypatc
|
|||
assert [tool.name for tool in tools] == ["list_servers", "enable_server"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_platform_mcp_catalog_call_can_bypass_compression(monkeypatch):
|
||||
normal_tools = [_tool(f"tool_{idx}") for idx in range(11)]
|
||||
|
||||
async def fake_get_tools(**kwargs):
|
||||
return normal_tools
|
||||
|
||||
monkeypatch.setattr(
|
||||
mcp_server_module,
|
||||
"_merge_toolset_permissions",
|
||||
_merge_toolset_permissions,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
mcp_server_module,
|
||||
"_get_tools_from_mcp_servers",
|
||||
fake_get_tools,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
platform_mcp,
|
||||
"get_platform_mcp_settings",
|
||||
_enabled_platform_settings,
|
||||
)
|
||||
|
||||
tools = await mcp_server_module._list_mcp_tools(
|
||||
enable_platform_mcp_compression=False
|
||||
)
|
||||
|
||||
assert tools == normal_tools
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_platform_mcp_does_not_compress_scoped_server_tools(monkeypatch):
|
||||
normal_tools = [_tool(f"tool_{idx}") for idx in range(11)]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue