fix: keep mcp catalog uncompressed

This commit is contained in:
Krrish Dholakia 2026-06-22 18:12:21 -07:00
parent 7a4bc18271
commit 9221010a1f
4 changed files with 89 additions and 1 deletions

View file

@ -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

View file

@ -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()
)

View file

@ -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]

View file

@ -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)]