mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(mcp): coerce mcp_server_cost_info values to float at ingest (#30109)
* fix(mcp): coerce mcp_server_cost_info values to float at ingest YAML 1.1 parses scientific notation without a decimal point (e.g. 7e-05) as a string, and MCPServerCostInfo is a TypedDict with no runtime validation, so a string-typed default_cost_per_query from config.yaml flowed through the proxy untouched and crashed the MCP server settings page with '.toFixed is not a function'. Normalize mcp_server_cost_info on both the config and DB load paths, dropping non-numeric values with a warning instead of failing the server load. Fixes #27097. * fix(mcp): drop non-numeric default_cost_per_query instead of nulling it Keeping the key with a None value still exposes a null to the UI, which can crash .toFixed formatting when the consumer checks key existence rather than truthiness. Delete the key on coercion failure, matching how non-numeric per-tool cost entries are already omitted.
This commit is contained in:
parent
ea29b28f40
commit
d47268939d
2 changed files with 112 additions and 0 deletions
|
|
@ -354,6 +354,52 @@ def _deserialize_json_list(data: Any) -> Optional[List[Dict[str, Any]]]:
|
|||
]
|
||||
|
||||
|
||||
def _normalize_mcp_server_cost_info(mcp_info: MCPInfo) -> None:
|
||||
"""Coerce ``mcp_server_cost_info`` numeric fields to ``float`` at ingest.
|
||||
|
||||
YAML 1.1 parses scientific notation without a decimal point (e.g.
|
||||
``7e-05``) as a string, and ``MCPServerCostInfo`` is a TypedDict with no
|
||||
runtime validation, so string-typed costs flow through to the UI and
|
||||
crash its ``.toFixed`` formatting. Values that cannot be coerced are
|
||||
dropped with a warning instead of failing the server load.
|
||||
"""
|
||||
cost_info = mcp_info.get("mcp_server_cost_info")
|
||||
if not isinstance(cost_info, dict):
|
||||
return
|
||||
|
||||
server_name = mcp_info.get("server_name")
|
||||
normalized = dict(cost_info)
|
||||
|
||||
default_cost = normalized.get("default_cost_per_query")
|
||||
if default_cost is not None:
|
||||
try:
|
||||
normalized["default_cost_per_query"] = float(default_cost)
|
||||
except (TypeError, ValueError):
|
||||
verbose_logger.warning(
|
||||
"MCP server '%s' has non-numeric default_cost_per_query %r; ignoring it",
|
||||
server_name,
|
||||
default_cost,
|
||||
)
|
||||
del normalized["default_cost_per_query"]
|
||||
|
||||
tool_costs = normalized.get("tool_name_to_cost_per_query")
|
||||
if isinstance(tool_costs, dict):
|
||||
normalized_tool_costs = {}
|
||||
for tool_name, cost in tool_costs.items():
|
||||
try:
|
||||
normalized_tool_costs[tool_name] = float(cost)
|
||||
except (TypeError, ValueError):
|
||||
verbose_logger.warning(
|
||||
"MCP server '%s' has non-numeric cost %r for tool '%s'; ignoring it",
|
||||
server_name,
|
||||
cost,
|
||||
tool_name,
|
||||
)
|
||||
normalized["tool_name_to_cost_per_query"] = normalized_tool_costs
|
||||
|
||||
mcp_info["mcp_server_cost_info"] = normalized
|
||||
|
||||
|
||||
def _create_sampling_callback(user_api_key_auth: Optional[Any] = None):
|
||||
"""
|
||||
Create a sampling callback for MCP ClientSession.
|
||||
|
|
@ -621,6 +667,7 @@ class MCPServerManager:
|
|||
mcp_info["server_name"] = server_name
|
||||
if "description" not in mcp_info and server_config.get("description"):
|
||||
mcp_info["description"] = server_config.get("description")
|
||||
_normalize_mcp_server_cost_info(mcp_info)
|
||||
|
||||
# Use alias for name if present, else server_name
|
||||
alias = server_config.get("alias", None)
|
||||
|
|
@ -1091,6 +1138,7 @@ class MCPServerManager:
|
|||
mcp_info["server_name"] = mcp_server.server_name or mcp_server.server_id
|
||||
if "description" not in mcp_info and mcp_server.description:
|
||||
mcp_info["description"] = mcp_server.description
|
||||
_normalize_mcp_server_cost_info(mcp_info)
|
||||
|
||||
auth_type = cast(MCPAuthType, mcp_server.auth_type)
|
||||
server_url = mcp_server.url
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
|||
MCPServerManager,
|
||||
_deserialize_json_dict,
|
||||
_deserialize_json_list,
|
||||
_normalize_mcp_server_cost_info,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
|
|
@ -257,6 +258,69 @@ class TestMCPServerManager:
|
|||
assert server.alias == "friendly_alias"
|
||||
assert server.server_name == "validserver"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_servers_from_config_coerces_cost_string_to_float(self):
|
||||
"""YAML 1.1 parses `7e-05` as a string; ingest must coerce it to float."""
|
||||
manager = MCPServerManager()
|
||||
config = {
|
||||
"google_maps": {
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
"mcp_info": {
|
||||
"mcp_server_cost_info": {
|
||||
"default_cost_per_query": "7e-05",
|
||||
"tool_name_to_cost_per_query": {"geocode": "1e-3"},
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
await manager.load_servers_from_config(config)
|
||||
|
||||
server = next(iter(manager.config_mcp_servers.values()))
|
||||
cost_info = server.mcp_info["mcp_server_cost_info"]
|
||||
assert cost_info["default_cost_per_query"] == 7e-05
|
||||
assert isinstance(cost_info["default_cost_per_query"], float)
|
||||
assert cost_info["tool_name_to_cost_per_query"]["geocode"] == 1e-3
|
||||
assert isinstance(cost_info["tool_name_to_cost_per_query"]["geocode"], float)
|
||||
|
||||
def test_normalize_mcp_server_cost_info_preserves_float_values(self):
|
||||
mcp_info = {
|
||||
"server_name": "maps",
|
||||
"mcp_server_cost_info": {
|
||||
"default_cost_per_query": 0.01,
|
||||
"tool_name_to_cost_per_query": {"search": 0.05},
|
||||
},
|
||||
}
|
||||
|
||||
_normalize_mcp_server_cost_info(mcp_info)
|
||||
|
||||
cost_info = mcp_info["mcp_server_cost_info"]
|
||||
assert cost_info["default_cost_per_query"] == 0.01
|
||||
assert cost_info["tool_name_to_cost_per_query"] == {"search": 0.05}
|
||||
|
||||
def test_normalize_mcp_server_cost_info_drops_non_numeric_values(self):
|
||||
mcp_info = {
|
||||
"server_name": "maps",
|
||||
"mcp_server_cost_info": {
|
||||
"default_cost_per_query": "not-a-number",
|
||||
"tool_name_to_cost_per_query": {"search": "free", "geocode": "2e-4"},
|
||||
},
|
||||
}
|
||||
|
||||
_normalize_mcp_server_cost_info(mcp_info)
|
||||
|
||||
cost_info = mcp_info["mcp_server_cost_info"]
|
||||
assert "default_cost_per_query" not in cost_info
|
||||
assert cost_info["tool_name_to_cost_per_query"] == {"geocode": 2e-4}
|
||||
|
||||
def test_normalize_mcp_server_cost_info_leaves_missing_cost_info_alone(self):
|
||||
mcp_info = {"server_name": "maps"}
|
||||
|
||||
_normalize_mcp_server_cost_info(mcp_info)
|
||||
|
||||
assert "mcp_server_cost_info" not in mcp_info
|
||||
|
||||
def test_warns_when_custom_separator_invalid(self, monkeypatch, caplog):
|
||||
"""Invalid MCP_TOOL_PREFIX_SEPARATOR values should log a warning."""
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue