From 3e9540f58951aee86ea9d5f0920f9008544e87c1 Mon Sep 17 00:00:00 2001 From: Uzair Ali <72073401+uzaxirr@users.noreply.github.com> Date: Tue, 23 Sep 2025 09:45:35 +0530 Subject: [PATCH] feat: enable custom fields in mcp_info configuration (#14794) Allow proxy admins to add arbitrary metadata fields to MCP servers in config.yaml under mcp_servers..mcp_info, similar to how model_info already works. Changes: - Changed MCPInfo from TypedDict to Dict[str, Any] for flexibility - Updated load_servers_from_config to preserve all custom fields - Updated add_update_server to handle arbitrary fields from database - Added comprehensive unit tests covering all scenarios --- .../mcp_server/mcp_server_manager.py | 30 +-- .../types/mcp_server/mcp_server_manager.py | 9 +- .../mcp_server/test_mcp_custom_fields.py | 212 ++++++++++++++++++ 3 files changed, 231 insertions(+), 20 deletions(-) create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_custom_fields.py diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index d0eadb36ba3..ab5d1b10bf3 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -121,15 +121,13 @@ class MCPServerManager: for server_name, server_config in mcp_servers_config.items(): validate_mcp_server_name(server_name) _mcp_info: Dict[str, Any] = server_config.get("mcp_info", None) or {} - # Convert Dict[str, Any] to MCPInfo properly - mcp_info: MCPInfo = { - "server_name": _mcp_info.get("server_name", server_name), - "description": _mcp_info.get( - "description", server_config.get("description", None) - ), - "logo_url": _mcp_info.get("logo_url", None), - "mcp_server_cost_info": _mcp_info.get("mcp_server_cost_info", None), - } + # Preserve all custom fields from config while setting defaults for core fields + mcp_info: MCPInfo = _mcp_info.copy() + # Set default values for core fields if not present + if "server_name" not in mcp_info: + mcp_info["server_name"] = server_name + if "description" not in mcp_info and server_config.get("description"): + mcp_info["description"] = server_config.get("description") # Use alias for name if present, else server_name alias = server_config.get("alias", None) @@ -243,6 +241,14 @@ class MCPServerManager: name_for_prefix = ( mcp_server.alias or mcp_server.server_name or mcp_server.server_id ) + # Preserve all custom fields from database while setting defaults for core fields + mcp_info: MCPInfo = _mcp_info.copy() + # Set default values for core fields if not present + if "server_name" not in mcp_info: + 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 + new_server = MCPServer( server_id=mcp_server.server_id, name=name_for_prefix, @@ -251,11 +257,7 @@ class MCPServerManager: url=mcp_server.url, transport=cast(MCPTransportType, mcp_server.transport), auth_type=cast(MCPAuthType, mcp_server.auth_type), - mcp_info=MCPInfo( - server_name=mcp_server.server_name or mcp_server.server_id, - description=mcp_server.description, - mcp_server_cost_info=_mcp_info.get("mcp_server_cost_info", None), - ), + mcp_info=mcp_info, # Stdio-specific fields command=getattr(mcp_server, "command", None), args=getattr(mcp_server, "args", None) or [], diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index eb1eb3250ba..bd09ef9c199 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -1,4 +1,4 @@ -from typing import Dict, List, Optional +from typing import Any, Dict, List, Optional from pydantic import BaseModel, ConfigDict from typing_extensions import TypedDict @@ -7,11 +7,8 @@ from litellm.proxy._types import MCPAuthType, MCPTransportType from litellm.types.mcp import MCPServerCostInfo -class MCPInfo(TypedDict, total=False): - server_name: str - description: Optional[str] - logo_url: Optional[str] - mcp_server_cost_info: Optional[MCPServerCostInfo] +# MCPInfo now allows arbitrary additional fields for custom metadata +MCPInfo = Dict[str, Any] class MCPServer(BaseModel): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_custom_fields.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_custom_fields.py new file mode 100644 index 00000000000..d2fa7853e78 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_custom_fields.py @@ -0,0 +1,212 @@ +""" +Test suite for MCP server custom fields functionality. + +Tests that mcp_info can accept arbitrary custom fields in addition to predefined ones. +""" +import pytest +import sys +import os +from unittest.mock import Mock, patch +from typing import Dict, Any + +# Add the path to find the modules +sys.path.insert( + 0, os.path.abspath("../../../..") +) # Adjust the path as needed + +from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager +from litellm.types.mcp import MCPAuth +from litellm.proxy._types import LiteLLM_MCPServerTable + + +class TestMCPCustomFields: + """Test custom fields functionality in MCP server configuration.""" + + def test_custom_fields_preserved_from_config(self): + """Test that custom fields in mcp_info are preserved when loading from config.""" + manager = MCPServerManager() + + # Mock config with custom fields + mock_config = { + "test_server": { + "url": "http://localhost:3000", + "transport": "http", + "auth_type": "bearer_token", + "authentication_token": "test-token", + "mcp_info": { + "server_name": "Test Server", + "description": "A test server", + "custom_field_1": "custom_value_1", + "custom_field_2": {"nested": "value"}, + "custom_field_3": ["list", "values"], + "priority": 10, + "tags": ["production", "api"] + } + } + } + + # Load servers from config + manager.load_servers_from_config(mock_config) + + # Get the loaded server + servers = list(manager.config_mcp_servers.values()) + assert len(servers) == 1 + + server = servers[0] + mcp_info = server.mcp_info + + # Verify standard fields are preserved + assert mcp_info["server_name"] == "Test Server" + assert mcp_info["description"] == "A test server" + + # Verify custom fields are preserved + assert mcp_info["custom_field_1"] == "custom_value_1" + assert mcp_info["custom_field_2"] == {"nested": "value"} + assert mcp_info["custom_field_3"] == ["list", "values"] + assert mcp_info["priority"] == 10 + assert mcp_info["tags"] == ["production", "api"] + + def test_custom_fields_preserved_from_database(self): + """Test that custom fields in mcp_info are preserved when adding from database.""" + manager = MCPServerManager() + + # Mock database record with custom fields + mock_server = Mock(spec=LiteLLM_MCPServerTable) + mock_server.server_id = "test-server-id" + mock_server.server_name = "Test Server" + mock_server.description = "A test server" + mock_server.url = "http://localhost:3000" + mock_server.transport = "http" + mock_server.auth_type = MCPAuth.bearer_token + mock_server.alias = None + mock_server.mcp_info = { + "server_name": "Test Server", + "description": "A test server", + "custom_db_field": "database_value", + "metadata": {"source": "database"}, + "version": "1.0.0" + } + mock_server.command = None + mock_server.args = None + mock_server.env = None + mock_server.mcp_access_groups = None + + # Add server to manager + manager.add_update_server(mock_server) + + # Get the added server + server = manager.get_mcp_server_by_id("test-server-id") + assert server is not None + + mcp_info = server.mcp_info + + # Verify standard fields are preserved + assert mcp_info["server_name"] == "Test Server" + assert mcp_info["description"] == "A test server" + + # Verify custom fields are preserved + assert mcp_info["custom_db_field"] == "database_value" + assert mcp_info["metadata"] == {"source": "database"} + assert mcp_info["version"] == "1.0.0" + + def test_empty_mcp_info_handled_gracefully(self): + """Test that empty or missing mcp_info is handled gracefully.""" + manager = MCPServerManager() + + # Config with empty mcp_info + mock_config = { + "test_server": { + "url": "http://localhost:3000", + "transport": "http", + "mcp_info": {} + } + } + + manager.load_servers_from_config(mock_config) + + servers = list(manager.config_mcp_servers.values()) + assert len(servers) == 1 + + server = servers[0] + mcp_info = server.mcp_info + + # Should have default server_name + assert mcp_info["server_name"] == "test_server" + + def test_missing_mcp_info_creates_defaults(self): + """Test that missing mcp_info creates appropriate defaults.""" + manager = MCPServerManager() + + # Config without mcp_info + mock_config = { + "test_server": { + "url": "http://localhost:3000", + "transport": "http", + "description": "Server description" + } + } + + manager.load_servers_from_config(mock_config) + + servers = list(manager.config_mcp_servers.values()) + assert len(servers) == 1 + + server = servers[0] + mcp_info = server.mcp_info + + # Should have default server_name and description from config + assert mcp_info["server_name"] == "test_server" + assert mcp_info["description"] == "Server description" + + def test_config_description_fallback(self): + """Test that description from config level is used as fallback.""" + manager = MCPServerManager() + + # Config with description at server level but not in mcp_info + mock_config = { + "test_server": { + "url": "http://localhost:3000", + "transport": "http", + "description": "Config level description", + "mcp_info": { + "custom_field": "custom_value" + } + } + } + + manager.load_servers_from_config(mock_config) + + servers = list(manager.config_mcp_servers.values()) + server = servers[0] + mcp_info = server.mcp_info + + # Should use config level description as fallback + assert mcp_info["description"] == "Config level description" + assert mcp_info["custom_field"] == "custom_value" + + def test_mcp_info_description_takes_precedence(self): + """Test that description in mcp_info takes precedence over config level.""" + manager = MCPServerManager() + + # Config with description at both levels + mock_config = { + "test_server": { + "url": "http://localhost:3000", + "transport": "http", + "description": "Config level description", + "mcp_info": { + "description": "MCP info description", + "custom_field": "custom_value" + } + } + } + + manager.load_servers_from_config(mock_config) + + servers = list(manager.config_mcp_servers.values()) + server = servers[0] + mcp_info = server.mcp_info + + # Should use mcp_info description, not config level + assert mcp_info["description"] == "MCP info description" + assert mcp_info["custom_field"] == "custom_value" \ No newline at end of file