fix: stop including spec_version in MCP server registration inserts

This commit is contained in:
Yuta Saito 2025-09-19 06:44:43 +09:00
parent 6c291093e9
commit 654f1d3290
14 changed files with 307 additions and 302 deletions

View file

@ -114,7 +114,6 @@ mcp_servers:
description: "My custom MCP server"
auth_type: "api_key"
auth_value: "abc123"
spec_version: "2025-03-26"
```
**Configuration Options:**
@ -716,7 +715,6 @@ mcp_servers:
url: https://mcp.deepwiki.com/mcp
transport: "http"
auth_type: "none"
spec_version: "2025-03-26"
access_groups: ["dev_group"]
```

View file

@ -0,0 +1,8 @@
/*
Warnings:
- You are about to drop the column `spec_version` on the `LiteLLM_MCPServerTable` table. All the data in the column will be lost.
*/
-- AlterTable
ALTER TABLE "public"."LiteLLM_MCPServerTable" DROP COLUMN "spec_version";

View file

@ -171,7 +171,6 @@ model LiteLLM_MCPServerTable {
description String?
url String?
transport String @default("sse")
spec_version String @default("2025-03-26")
auth_type String?
created_at DateTime? @default(now()) @map("created_at")
created_by String?

View file

@ -34,8 +34,6 @@ from litellm.proxy._experimental.mcp_server.utils import (
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
MCPAuthType,
MCPSpecVersion,
MCPSpecVersionType,
MCPTransport,
MCPTransportType,
UserAPIKeyAuth,
@ -70,38 +68,6 @@ def _deserialize_env_dict(env_data: Any) -> Optional[Dict[str, str]]:
return env_data
def _convert_protocol_version_to_enum(
protocol_version: Optional[str | MCPSpecVersionType],
) -> MCPSpecVersionType:
"""
Convert string protocol version to MCPSpecVersion enum.
Args:
protocol_version: String protocol version, enum, or None
Returns:
MCPSpecVersionType: The enum value
"""
if not protocol_version:
return cast(MCPSpecVersionType, MCPSpecVersion.jun_2025)
# If it's already an MCPSpecVersion enum, return it
if isinstance(protocol_version, MCPSpecVersion):
return cast(MCPSpecVersionType, protocol_version)
# If it's a string, try to match it to enum values
if isinstance(protocol_version, str):
for version in MCPSpecVersion:
if version.value == protocol_version:
return cast(MCPSpecVersionType, version)
# If no match found, return default
verbose_logger.warning(
f"Unknown protocol version '{protocol_version}', using default"
)
return cast(MCPSpecVersionType, MCPSpecVersion.jun_2025)
class MCPServerManager:
def __init__(self):
self.registry: Dict[str, MCPServer] = {}
@ -113,8 +79,7 @@ class MCPServerManager:
"name": "zapier_mcp_server",
"url": "https://actions.zapier.com/mcp/sk-ak-2ew3bofIeQIkNoeKIdXrF1Hhhp/sse"
"transport": "sse",
"auth_type": "api_key",
"spec_version": "2025-03-26"
"auth_type": "api_key"
},
"uuid-2": {
"name": "google_drive_mcp_server",
@ -223,7 +188,6 @@ class MCPServerManager:
server_name=server_name,
url=server_config.get("url", None) or "",
transport=server_config.get("transport", MCPTransport.http),
spec_version=server_config.get("spec_version", MCPSpecVersion.jun_2025),
auth_type=server_config.get("auth_type", None),
alias=alias,
)
@ -873,7 +837,6 @@ class MCPServerManager:
server_name: str,
url: str,
transport: str,
spec_version: str,
auth_type: Optional[str] = None,
alias: Optional[str] = None,
) -> str:
@ -889,7 +852,6 @@ class MCPServerManager:
server_name: Name of the server
url: Server URL
transport: Transport type (sse, http, etc.)
spec_version: MCP spec version
auth_type: Authentication type (optional)
alias: Server alias (optional)
@ -897,7 +859,9 @@ class MCPServerManager:
A deterministic server ID string
"""
# Create a string from all the identifying parameters
params_string = f"{server_name}|{url}|{transport}|{spec_version}|{auth_type or ''}|{alias or ''}"
params_string = (
f"{server_name}|{url}|{transport}|{auth_type or ''}|{alias or ''}"
)
# Generate SHA-256 hash
hash_object = hashlib.sha256(params_string.encode("utf-8"))
@ -1054,7 +1018,6 @@ class MCPServerManager:
alias=_server_config.alias,
url=_server_config.url,
transport=_server_config.transport,
spec_version=_server_config.spec_version,
auth_type=_server_config.auth_type,
created_at=datetime.datetime.now(),
updated_at=datetime.datetime.now(),
@ -1117,7 +1080,6 @@ class MCPServerManager:
description=server.description,
url=server.url,
transport=server.transport,
spec_version=server.spec_version,
auth_type=server.auth_type,
created_at=server.created_at,
created_by=server.created_by,

View file

@ -23,7 +23,6 @@ router = APIRouter(
if MCP_AVAILABLE:
from litellm.experimental_mcp_client.client import MCPTool
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_convert_protocol_version_to_enum,
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.server import (

View file

@ -28,8 +28,6 @@ from litellm.types.integrations.slack_alerting import AlertType
from litellm.types.llms.openai import AllMessageValues, OpenAIFileObject
from litellm.types.mcp import (
MCPAuthType,
MCPSpecVersion,
MCPSpecVersionType,
MCPTransport,
MCPTransportType,
)
@ -748,9 +746,9 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
allowed_cache_controls: Optional[list] = []
config: Optional[dict] = {}
permissions: Optional[dict] = {}
model_max_budget: Optional[dict] = (
{}
) # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {}
model_max_budget: Optional[
dict
] = {} # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {}
model_config = ConfigDict(protected_namespaces=())
model_rpm_limit: Optional[dict] = None
@ -789,6 +787,7 @@ class GenerateKeyRequest(KeyRequestBase):
description="Type of key that determines default allowed routes.",
)
class GenerateKeyResponse(KeyRequestBase):
key: str # type: ignore
key_name: Optional[str] = None
@ -916,7 +915,6 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
alias: Optional[str] = None
description: Optional[str] = None
transport: MCPTransportType = MCPTransport.sse
spec_version: MCPSpecVersionType = MCPSpecVersion.jun_2025
auth_type: Optional[MCPAuthType] = None
url: Optional[str] = None
mcp_info: Optional[MCPInfo] = None
@ -948,7 +946,6 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
alias: Optional[str] = None
description: Optional[str] = None
transport: MCPTransportType = MCPTransport.sse
spec_version: MCPSpecVersionType = MCPSpecVersion.jun_2025
auth_type: Optional[MCPAuthType] = None
url: Optional[str] = None
mcp_info: Optional[MCPInfo] = None
@ -983,7 +980,6 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
description: Optional[str] = None
url: Optional[str] = None
transport: MCPTransportType
spec_version: MCPSpecVersionType
auth_type: Optional[MCPAuthType] = None
created_at: Optional[datetime] = None
created_by: Optional[str] = None
@ -1150,12 +1146,12 @@ class NewCustomerRequest(BudgetNewRequest):
blocked: bool = False # allow/disallow requests for this end-user
budget_id: Optional[str] = None # give either a budget_id or max_budget
spend: Optional[float] = None
allowed_model_region: Optional[AllowedModelRegion] = (
None # require all user requests to use models in this specific region
)
default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
default_model: Optional[
str
] = None # if no equivalent model in allowed region - default all requests to this model
@model_validator(mode="before")
@classmethod
@ -1177,12 +1173,12 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase):
blocked: bool = False # allow/disallow requests for this end-user
max_budget: Optional[float] = None
budget_id: Optional[str] = None # give either a budget_id or max_budget
allowed_model_region: Optional[AllowedModelRegion] = (
None # require all user requests to use models in this specific region
)
default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
default_model: Optional[
str
] = None # if no equivalent model in allowed region - default all requests to this model
class DeleteCustomerRequest(LiteLLMPydanticObjectBase):
@ -1256,15 +1252,15 @@ class NewTeamRequest(TeamBase):
guardrails: Optional[List[str]] = None
prompts: Optional[List[str]] = None
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
team_member_budget: Optional[float] = (
None # allow user to set a budget for all team members
)
team_member_rpm_limit: Optional[int] = (
None # allow user to set RPM limit for all team members
)
team_member_tpm_limit: Optional[int] = (
None # allow user to set TPM limit for all team members
)
team_member_budget: Optional[
float
] = None # allow user to set a budget for all team members
team_member_rpm_limit: Optional[
int
] = None # allow user to set RPM limit for all team members
team_member_tpm_limit: Optional[
int
] = None # allow user to set TPM limit for all team members
team_member_key_duration: Optional[str] = None # e.g. "1d", "1w", "1m"
model_config = ConfigDict(protected_namespaces=())
@ -1343,9 +1339,9 @@ class BlockKeyRequest(LiteLLMPydanticObjectBase):
class AddTeamCallback(LiteLLMPydanticObjectBase):
callback_name: str
callback_type: Optional[Literal["success", "failure", "success_and_failure"]] = (
"success_and_failure"
)
callback_type: Optional[
Literal["success", "failure", "success_and_failure"]
] = "success_and_failure"
callback_vars: Dict[str, str]
@model_validator(mode="before")
@ -1614,15 +1610,16 @@ class ConfigList(LiteLLMPydanticObjectBase):
stored_in_db: Optional[bool]
field_default_value: Any
premium_field: bool = False
nested_fields: Optional[List[FieldDetail]] = (
None # For nested dictionary or Pydantic fields
)
nested_fields: Optional[
List[FieldDetail]
] = None # For nested dictionary or Pydantic fields
class UserHeaderMapping(LiteLLMPydanticObjectBase):
"""
Map an incoming HTTP header to a LiteLLM user role.
"""
header_name: str
litellm_user_role: Literal[
LitellmUserRoles.INTERNAL_USER,
@ -1633,6 +1630,7 @@ class UserHeaderMapping(LiteLLMPydanticObjectBase):
"extra": "forbid",
}
class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
"""
Documents all the fields supported by `general_settings` in config.yaml
@ -1943,9 +1941,9 @@ class LiteLLM_OrganizationMembershipTable(LiteLLMPydanticObjectBase):
budget_id: Optional[str] = None
created_at: datetime
updated_at: datetime
user: Optional[Any] = (
None # You might want to replace 'Any' with a more specific type if available
)
user: Optional[
Any
] = None # You might want to replace 'Any' with a more specific type if available
litellm_budget_table: Optional[LiteLLM_BudgetTable] = None
model_config = ConfigDict(protected_namespaces=())
@ -2840,9 +2838,9 @@ class TeamModelDeleteRequest(BaseModel):
# Organization Member Requests
class OrganizationMemberAddRequest(OrgMemberAddRequest):
organization_id: str
max_budget_in_organization: Optional[float] = (
None # Users max budget within the organization
)
max_budget_in_organization: Optional[
float
] = None # Users max budget within the organization
class OrganizationMemberDeleteRequest(MemberDeleteRequest):
@ -2941,10 +2939,12 @@ class LitellmDataForBackendLLMCall(TypedDict, total=False):
user: Optional[str]
num_retries: Optional[int]
class LitellmMetadataFromRequestHeaders(TypedDict, total=False):
"""
Headers a user can pass that will get added to litellm metadata for the request
"""
spend_logs_metadata: Optional[dict]
@ -3050,9 +3050,9 @@ class ProviderBudgetResponse(LiteLLMPydanticObjectBase):
Maps provider names to their budget configs.
"""
providers: Dict[str, ProviderBudgetResponseObject] = (
{}
) # Dictionary mapping provider names to their budget configurations
providers: Dict[
str, ProviderBudgetResponseObject
] = {} # Dictionary mapping provider names to their budget configurations
class ProxyStateVariables(TypedDict):
@ -3186,9 +3186,9 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
enforce_rbac: bool = False
roles_jwt_field: Optional[str] = None # v2 on role mappings
role_mappings: Optional[List[RoleMapping]] = None
object_id_jwt_field: Optional[str] = (
None # can be either user / team, inferred from the role mapping
)
object_id_jwt_field: Optional[
str
] = None # can be either user / team, inferred from the role mapping
scope_mappings: Optional[List[ScopeMapping]] = None
enforce_scope_based_access: bool = False
enforce_team_based_model_access: bool = False

View file

@ -171,7 +171,6 @@ model LiteLLM_MCPServerTable {
description String?
url String?
transport String @default("sse")
spec_version String @default("2025-03-26")
auth_type String?
created_at DateTime? @default(now()) @map("created_at")
created_by String?

View file

@ -1,9 +1,9 @@
from typing import TYPE_CHECKING, Dict, List, Optional
from typing import Dict, List, Optional
from pydantic import BaseModel, ConfigDict
from typing_extensions import TypedDict
from litellm.proxy._types import MCPAuthType, MCPSpecVersionType, MCPTransportType
from litellm.proxy._types import MCPAuthType, MCPTransportType
from litellm.types.mcp import MCPServerCostInfo

View file

@ -171,7 +171,6 @@ model LiteLLM_MCPServerTable {
description String?
url String?
transport String @default("sse")
spec_version String @default("2025-03-26")
auth_type String?
created_at DateTime? @default(now()) @map("created_at")
created_by String?

View file

@ -25,7 +25,6 @@ def test_mcp_server_works_without_config_auth_value():
alias="test_no_config",
url="https://api.example.com/mcp",
transport=MCPTransport.http,
spec_version=MCPSpecVersion.jun_2025,
auth_type=MCPAuth.authorization,
authentication_token=None, # No config auth
)

View file

@ -479,7 +479,6 @@ def test_generate_stable_server_id():
"server_name": "zapier_mcp_server",
"url": "https://actions.zapier.com/mcp/sse",
"transport": "sse",
"spec_version": "2025-03-26",
"auth_type": "api_key",
},
"expected_hash": "8d5c9f8a12e3b7c4f6a2d8e1b5c9f2a4",
@ -489,7 +488,6 @@ def test_generate_stable_server_id():
"server_name": "google_drive_mcp_server",
"url": "https://drive.google.com/mcp/http",
"transport": "http",
"spec_version": "2024-11-20",
"auth_type": None,
},
"expected_hash": "7a4b2e8f3c1d9e6b5a7c8f2d4e1b9c6a",
@ -499,7 +497,6 @@ def test_generate_stable_server_id():
"server_name": "local_test_server",
"url": "http://localhost:8080/mcp",
"transport": "http",
"spec_version": "2025-03-26",
"auth_type": "basic",
},
"expected_hash": "2f1e8d7c6b5a4e3f2d1c9b8a7e6f5d4c",
@ -527,7 +524,6 @@ def test_generate_stable_server_id():
"server_name": "test_server",
"url": "https://test.com/mcp",
"transport": "sse",
"spec_version": "2025-03-26",
"auth_type": "api_key",
}
@ -538,7 +534,6 @@ def test_generate_stable_server_id():
{"server_name": "different_server"},
{"url": "https://different.com/mcp"},
{"transport": "http"},
{"spec_version": "2024-11-20"},
{"auth_type": "basic"},
{"auth_type": None},
]
@ -558,7 +553,6 @@ def test_generate_stable_server_id():
"server_name": "test_server",
"url": "https://test.com/mcp",
"transport": "sse",
"spec_version": "2025-03-26",
"auth_type": None,
}
@ -566,7 +560,6 @@ def test_generate_stable_server_id():
"server_name": "test_server",
"url": "https://test.com/mcp",
"transport": "sse",
"spec_version": "2025-03-26",
"auth_type": "",
}
@ -584,7 +577,6 @@ def test_generate_stable_server_id():
server_name="zapier_mcp_server",
url="https://actions.zapier.com/mcp/sk-ak-example/sse",
transport="sse",
spec_version="2025-03-26",
auth_type="api_key",
)
@ -592,7 +584,6 @@ def test_generate_stable_server_id():
server_name="github_mcp_server",
url="https://api.github.com/mcp/http",
transport="http",
spec_version="2025-03-26",
auth_type=None,
)
@ -601,7 +592,6 @@ def test_generate_stable_server_id():
server_name="zapier_mcp_server",
url="https://actions.zapier.com/mcp/sk-ak-example/sse",
transport="sse",
spec_version="2025-03-26",
auth_type="api_key",
)
@ -609,7 +599,6 @@ def test_generate_stable_server_id():
server_name="github_mcp_server",
url="https://api.github.com/mcp/http",
transport="http",
spec_version="2025-03-26",
auth_type=None,
)
@ -1038,7 +1027,6 @@ def test_mcp_server_manager_config_integration_with_database():
server_name="database-server",
url="https://db-server.com/mcp",
transport="http",
spec_version="2025-03-26",
auth_type="none",
description="Database server description",
created_at=datetime.datetime.now(),
@ -1356,7 +1344,6 @@ def test_add_update_server_with_alias():
mock_mcp_server.server_name = "Test Server"
mock_mcp_server.url = "https://test-server.com/mcp"
mock_mcp_server.transport = MCPTransport.http
mock_mcp_server.spec_version = "2025-03-26"
mock_mcp_server.auth_type = None
mock_mcp_server.description = "Test server description"
mock_mcp_server.mcp_info = {}
@ -1388,7 +1375,6 @@ def test_add_update_server_without_alias():
mock_mcp_server.server_name = "Test Server"
mock_mcp_server.url = "https://test-server.com/mcp"
mock_mcp_server.transport = MCPTransport.http
mock_mcp_server.spec_version = "2025-03-26"
mock_mcp_server.auth_type = None
mock_mcp_server.description = "Test server description"
mock_mcp_server.mcp_info = {}
@ -1420,7 +1406,6 @@ def test_add_update_server_fallback_to_server_id():
mock_mcp_server.server_name = None
mock_mcp_server.url = "https://test-server.com/mcp"
mock_mcp_server.transport = MCPTransport.http
mock_mcp_server.spec_version = "2025-03-26"
mock_mcp_server.auth_type = None
mock_mcp_server.description = "Test server description"
mock_mcp_server.mcp_info = {}

View file

@ -11,15 +11,27 @@ from fastapi import FastAPI
from starlette import status
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
from litellm.proxy._types import MCPSpecVersion, MCPSpecVersionType, MCPTransportType, MCPTransport, NewMCPServerRequest, LiteLLM_MCPServerTable, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import (
MCPSpecVersion,
MCPSpecVersionType,
MCPTransportType,
MCPTransport,
NewMCPServerRequest,
LiteLLM_MCPServerTable,
LitellmUserRoles,
UserAPIKeyAuth,
)
from litellm.types.mcp import MCPAuth
from litellm.proxy.management_endpoints.mcp_management_endpoints import does_mcp_server_exist
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
does_mcp_server_exist,
)
TEST_MASTER_KEY = os.getenv("LITELLM_MASTER_KEY", "sk-1234")
def generate_mcpserver_record(url: Optional[str] = None,
transport: Optional[MCPTransportType] = None,
spec_version: Optional[MCPSpecVersionType] = None) -> LiteLLM_MCPServerTable:
def generate_mcpserver_record(
url: Optional[str] = None, transport: Optional[MCPTransportType] = None
) -> LiteLLM_MCPServerTable:
"""
Generate a mock record for testing.
"""
@ -30,11 +42,11 @@ def generate_mcpserver_record(url: Optional[str] = None,
alias="Test Server",
url=url or "http://localhost.com:8080/mcp",
transport=transport or MCPTransport.sse,
spec_version=spec_version or MCPSpecVersion.mar_2025,
created_at=now,
updated_at=now,
)
# Cheers SO
def is_valid_uuid(val):
try:
@ -43,11 +55,12 @@ def is_valid_uuid(val):
except ValueError:
return False
def generate_mcpserver_create_request(
server_id: Optional[str] = None,
url: Optional[str] = None,
transport: Optional[MCPTransportType] = None,
spec_version: Optional[MCPSpecVersionType] = None) -> NewMCPServerRequest:
server_id: Optional[str] = None,
url: Optional[str] = None,
transport: Optional[MCPTransportType] = None,
) -> NewMCPServerRequest:
"""
Generate a mock create request for testing.
"""
@ -56,10 +69,12 @@ def generate_mcpserver_create_request(
alias="Test Server",
url=url or "http://localhost.com:8080/mcp",
transport=transport or MCPTransport.sse,
spec_version=spec_version or MCPSpecVersion.mar_2025,
)
def assert_mcp_server_record_same(mcp_server: NewMCPServerRequest, resp: LiteLLM_MCPServerTable):
def assert_mcp_server_record_same(
mcp_server: NewMCPServerRequest, resp: LiteLLM_MCPServerTable
):
"""
Assert that the mcp server record is created correctly.
"""
@ -71,7 +86,6 @@ def assert_mcp_server_record_same(mcp_server: NewMCPServerRequest, resp: LiteLLM
assert resp.url == mcp_server.url
assert resp.description == mcp_server.description
assert resp.transport == mcp_server.transport
assert resp.spec_version == mcp_server.spec_version
assert resp.auth_type == mcp_server.auth_type
assert resp.created_at is not None
assert resp.updated_at is not None
@ -83,224 +97,263 @@ def test_does_mcp_server_exist():
"""
Unit Test if the MCP server exists in the list.
"""
mcp_server_records: List[LiteLLM_MCPServerTable] = [generate_mcpserver_record(), generate_mcpserver_record()]
mcp_server_records: List[LiteLLM_MCPServerTable] = [
generate_mcpserver_record(),
generate_mcpserver_record(),
]
# test all records are found
for record in mcp_server_records:
assert does_mcp_server_exist(mcp_server_records, record.server_id)
# test record not found
not_found_record = str(uuid.uuid4())
assert False == does_mcp_server_exist(mcp_server_records, not_found_record)
@pytest.mark.asyncio
async def test_create_mcp_server_direct():
"""
Direct test of the MCP server creation logic without HTTP calls.
"""
# Mock the database functions directly
with mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE", True), \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw") as mock_get_prisma, \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server", new_callable=mock.AsyncMock) as mock_create, \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server", new_callable=mock.AsyncMock) as mock_get_server, \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager") as mock_manager:
with mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE",
True,
), mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw"
) as mock_get_prisma, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server",
new_callable=mock.AsyncMock,
) as mock_create, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
new_callable=mock.AsyncMock,
) as mock_get_server, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager"
) as mock_manager:
# Import after mocking
from litellm.proxy.management_endpoints.mcp_management_endpoints import add_mcp_server
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
add_mcp_server,
)
# Mock database client
mock_prisma = mock.Mock()
mock_get_prisma.return_value = mock_prisma
# Mock server manager
mock_manager.add_update_server = mock.Mock()
mock_manager.reload_servers_from_database = mock.AsyncMock()
# Set up test data
server_id = str(uuid.uuid4())
mcp_server_request = generate_mcpserver_create_request(server_id=server_id)
# The function will normalize the alias by replacing spaces with underscores
expected_alias = mcp_server_request.alias.replace(' ', '_') if mcp_server_request.alias else None
expected_alias = (
mcp_server_request.alias.replace(" ", "_")
if mcp_server_request.alias
else None
)
expected_response = LiteLLM_MCPServerTable(
server_id=server_id,
alias=expected_alias, # Use the normalized alias
description=mcp_server_request.description,
url=mcp_server_request.url,
transport=mcp_server_request.transport,
spec_version=mcp_server_request.spec_version,
auth_type=mcp_server_request.auth_type,
created_at=datetime.now(),
updated_at=datetime.now(),
created_by=LITELLM_PROXY_ADMIN_NAME,
updated_by=LITELLM_PROXY_ADMIN_NAME,
teams=[]
teams=[],
)
# Mock the database calls
mock_get_server.return_value = None # Server doesn't exist yet
# Set up async mock for create_mcp_server using AsyncMock
mock_create.return_value = expected_response
# Create mock user auth
user_auth = UserAPIKeyAuth(
api_key=TEST_MASTER_KEY,
user_id="test-user",
user_role=LitellmUserRoles.PROXY_ADMIN
user_role=LitellmUserRoles.PROXY_ADMIN,
)
# Call the function directly
result = await add_mcp_server(
payload=mcp_server_request,
user_api_key_dict=user_auth
payload=mcp_server_request, user_api_key_dict=user_auth
)
# Verify the result
assert result.server_id == server_id
assert result.alias == expected_alias # Check against normalized alias
assert result.url == mcp_server_request.url
assert result.transport == mcp_server_request.transport
assert result.spec_version == mcp_server_request.spec_version
# Verify mocks were called
mock_get_server.assert_called_once_with(mock_prisma, server_id)
mock_create.assert_called_once()
mock_manager.add_update_server.assert_called_once_with(expected_response)
@pytest.mark.asyncio
@pytest.mark.asyncio
async def test_create_duplicate_mcp_server():
"""
Test that creating a duplicate MCP server fails appropriately.
"""
# Mock the database functions directly
with mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE", True), \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw") as mock_get_prisma, \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server", new_callable=mock.AsyncMock) as mock_get_server:
with mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE",
True,
), mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw"
) as mock_get_prisma, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
new_callable=mock.AsyncMock,
) as mock_get_server:
# Import after mocking
from litellm.proxy.management_endpoints.mcp_management_endpoints import add_mcp_server
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
add_mcp_server,
)
from fastapi import HTTPException
# Mock database client
mock_prisma = mock.Mock()
mock_get_prisma.return_value = mock_prisma
# Set up test data
server_id = str(uuid.uuid4())
mcp_server_request = generate_mcpserver_create_request(server_id=server_id)
existing_server = LiteLLM_MCPServerTable(
server_id=server_id,
alias="Existing Server",
url="http://existing.com",
transport=MCPTransport.sse,
spec_version=MCPSpecVersion.mar_2025,
created_at=datetime.now(),
updated_at=datetime.now(),
teams=[]
teams=[],
)
# Mock that server already exists
mock_get_server.return_value = existing_server
# Create mock user auth
user_auth = UserAPIKeyAuth(
api_key=TEST_MASTER_KEY,
user_id="test-user",
user_role=LitellmUserRoles.PROXY_ADMIN
user_role=LitellmUserRoles.PROXY_ADMIN,
)
# Expect HTTPException to be raised
with pytest.raises(HTTPException) as exc_info:
await add_mcp_server(
payload=mcp_server_request,
user_api_key_dict=user_auth
payload=mcp_server_request, user_api_key_dict=user_auth
)
# Verify the exception details
assert exc_info.value.status_code == 400
assert "already exists" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_create_mcp_server_auth_failure():
"""
Test that non-admin users cannot create MCP servers.
"""
# Mock the database functions directly
with mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE", True), \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw") as mock_get_prisma:
with mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE",
True,
), mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw"
) as mock_get_prisma:
# Import after mocking
from litellm.proxy.management_endpoints.mcp_management_endpoints import add_mcp_server
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
add_mcp_server,
)
from fastapi import HTTPException
# Mock database client
# Mock database client
mock_prisma = mock.Mock()
mock_get_prisma.return_value = mock_prisma
# Set up test data
server_id = str(uuid.uuid4())
mcp_server_request = generate_mcpserver_create_request(server_id=server_id)
# Create mock user auth without admin role
user_auth = UserAPIKeyAuth(
api_key=TEST_MASTER_KEY,
user_id="test-user",
user_role=LitellmUserRoles.INTERNAL_USER # Not an admin
user_role=LitellmUserRoles.INTERNAL_USER, # Not an admin
)
# Expect HTTPException to be raised
with pytest.raises(HTTPException) as exc_info:
await add_mcp_server(
payload=mcp_server_request,
user_api_key_dict=user_auth
payload=mcp_server_request, user_api_key_dict=user_auth
)
# Verify the exception details
assert exc_info.value.status_code == 403
assert "permission" in str(exc_info.value.detail)
@pytest.mark.asyncio
@pytest.mark.asyncio
async def test_create_mcp_server_invalid_alias():
"""
Test that creating an MCP server with a '-' in the alias fails with the correct error.
"""
with mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE", True), \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw") as mock_get_prisma, \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server") as mock_get_server, \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server") as mock_create:
from litellm.proxy.management_endpoints.mcp_management_endpoints import add_mcp_server
with mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE",
True,
), mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw"
) as mock_get_prisma, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server"
) as mock_get_server, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server"
) as mock_create:
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
add_mcp_server,
)
from fastapi import HTTPException
mock_prisma = mock.Mock()
mock_get_prisma.return_value = mock_prisma
# Set up test data with invalid alias
server_id = str(uuid.uuid4())
mcp_server_request = generate_mcpserver_create_request(server_id=server_id)
mcp_server_request.alias = "invalid-alias" # This should trigger the validation error
mcp_server_request.alias = (
"invalid-alias" # This should trigger the validation error
)
# Mock that server does not exist
mock_get_server.return_value = None
# Mock create_mcp_server to prevent 500 error (this should not be called due to validation)
mock_create.return_value = None
user_auth = UserAPIKeyAuth(
api_key=TEST_MASTER_KEY,
user_id="test-user",
user_role=LitellmUserRoles.PROXY_ADMIN
user_role=LitellmUserRoles.PROXY_ADMIN,
)
with pytest.raises(HTTPException) as exc_info:
await add_mcp_server(
payload=mcp_server_request,
user_api_key_dict=user_auth
payload=mcp_server_request, user_api_key_dict=user_auth
)
assert exc_info.value.status_code == 400
assert "Server name cannot contain '-'. Use an alternative character instead Found: invalid-alias" in str(exc_info.value.detail)
assert (
"Server name cannot contain '-'. Use an alternative character instead Found: invalid-alias"
in str(exc_info.value.detail)
)
def test_validate_mcp_server_name_direct():
"""
@ -308,16 +361,16 @@ def test_validate_mcp_server_name_direct():
"""
from litellm.proxy._experimental.mcp_server.utils import validate_mcp_server_name
from fastapi import HTTPException
# Test that valid names pass
validate_mcp_server_name("valid_name")
validate_mcp_server_name("valid name")
# Test that invalid names with hyphens raise exceptions
with pytest.raises(Exception) as exc_info:
validate_mcp_server_name("invalid-name")
assert "cannot contain" in str(exc_info.value)
# Test that invalid names with hyphens raise HTTPException when requested
with pytest.raises(HTTPException) as exc_info:
validate_mcp_server_name("invalid-name", raise_http_exception=True)

View file

@ -5,7 +5,7 @@ from unittest.mock import MagicMock, AsyncMock
import pytest
# Add the parent directory to the path so we can import litellm
sys.path.insert(0, '../../../../../')
sys.path.insert(0, "../../../../../")
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
@ -24,12 +24,12 @@ class TestMCPServerManager:
env_json = '{"PATH": "/usr/bin", "DEBUG": "1"}'
result = _deserialize_env_dict(env_json)
assert result == {"PATH": "/usr/bin", "DEBUG": "1"}
# Test already dict
env_dict = {"PATH": "/usr/bin", "DEBUG": "1"}
result = _deserialize_env_dict(env_dict)
assert result == {"PATH": "/usr/bin", "DEBUG": "1"}
# Test invalid JSON
invalid_json = '{"PATH": "/usr/bin", "DEBUG": 1'
result = _deserialize_env_dict(invalid_json)
@ -38,27 +38,26 @@ class TestMCPServerManager:
def test_add_update_server_stdio(self):
"""Test adding stdio MCP server"""
manager = MCPServerManager()
stdio_server = LiteLLM_MCPServerTable(
server_id="stdio-server-1",
alias="test_stdio_server",
description="Test stdio server",
url=None,
transport=MCPTransport.stdio,
spec_version=MCPSpecVersion.mar_2025,
command="python",
args=["-m", "server"],
env={"DEBUG": "1", "TEST": "1"},
created_at=datetime.now(),
updated_at=datetime.now()
updated_at=datetime.now(),
)
manager.add_update_server(stdio_server)
# Verify server was added
assert "stdio-server-1" in manager.registry
added_server = manager.registry["stdio-server-1"]
assert added_server.server_id == "stdio-server-1"
assert added_server.name == "test_stdio_server"
assert added_server.transport == MCPTransport.stdio
@ -69,20 +68,19 @@ class TestMCPServerManager:
def test_create_mcp_client_stdio(self):
"""Test creating MCP client for stdio transport"""
manager = MCPServerManager()
stdio_server = MCPServer(
server_id="stdio-server-2",
name="test_stdio_server",
url=None,
transport=MCPTransport.stdio,
spec_version=MCPSpecVersion.mar_2025,
command="node",
args=["server.js"],
env={"NODE_ENV": "test"}
env={"NODE_ENV": "test"},
)
client = manager._create_mcp_client(stdio_server)
assert client.transport_type == MCPTransport.stdio
assert client.stdio_config is not None
assert client.stdio_config["command"] == "node"
@ -93,24 +91,28 @@ class TestMCPServerManager:
async def test_list_tools_with_server_specific_auth_headers(self):
"""Test list_tools method with server-specific auth headers"""
manager = MCPServerManager()
# Mock servers
server1 = MagicMock()
server1.name = "github"
server1.alias = "github"
server1.server_name = "github"
server2 = MagicMock()
server2.name = "zapier"
server2.alias = "zapier"
server2.server_name = "zapier"
# Mock get_allowed_mcp_servers to return our test servers
manager.get_allowed_mcp_servers = AsyncMock(return_value=["github", "zapier"])
manager.get_mcp_server_by_id = MagicMock(side_effect=lambda x: server1 if x == "github" else server2)
manager.get_mcp_server_by_id = MagicMock(
side_effect=lambda x: server1 if x == "github" else server2
)
# Mock _get_tools_from_server to return different results
async def mock_get_tools_from_server(server, mcp_auth_header=None, mcp_protocol_version=None):
async def mock_get_tools_from_server(
server, mcp_auth_header=None, mcp_protocol_version=None
):
if server.name == "github":
tool1 = MagicMock()
tool1.name = "github_tool_1"
@ -121,20 +123,22 @@ class TestMCPServerManager:
tool1 = MagicMock()
tool1.name = "zapier_tool_1"
return [tool1]
manager._get_tools_from_server = mock_get_tools_from_server
# Test with server-specific auth headers
mcp_server_auth_headers = {
"github": "Bearer github-token",
"zapier": "zapier-api-key"
"zapier": "zapier-api-key",
}
result = await manager.list_tools(mcp_server_auth_headers=mcp_server_auth_headers)
result = await manager.list_tools(
mcp_server_auth_headers=mcp_server_auth_headers
)
# Verify that both servers were called with their specific auth headers
assert len(result) == 3 # 2 from github + 1 from zapier
# Verify the tools have the expected names
tool_names = [tool.name for tool in result]
assert "github_tool_1" in tool_names
@ -145,32 +149,34 @@ class TestMCPServerManager:
async def test_list_tools_fallback_to_legacy_auth_header(self):
"""Test that list_tools falls back to legacy auth header when server-specific not available"""
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.name = "github"
server.alias = "github"
server.server_name = "github"
# Mock get_allowed_mcp_servers
manager.get_allowed_mcp_servers = AsyncMock(return_value=["github"])
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock _get_tools_from_server
async def mock_get_tools_from_server(server, mcp_auth_header=None, mcp_protocol_version=None):
async def mock_get_tools_from_server(
server, mcp_auth_header=None, mcp_protocol_version=None
):
assert mcp_auth_header == "legacy-token" # Should use legacy header
tool = MagicMock()
tool.name = "github_tool_1"
return [tool]
manager._get_tools_from_server = mock_get_tools_from_server
# Test with only legacy auth header (no server-specific headers)
result = await manager.list_tools(
mcp_auth_header="legacy-token",
mcp_server_auth_headers={} # Empty server-specific headers
mcp_server_auth_headers={}, # Empty server-specific headers
)
assert len(result) == 1
assert result[0].name == "github_tool_1"
@ -178,32 +184,36 @@ class TestMCPServerManager:
async def test_list_tools_prioritizes_server_specific_over_legacy(self):
"""Test that server-specific auth headers take priority over legacy header"""
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.name = "github"
server.alias = "github"
server.server_name = "github"
# Mock get_allowed_mcp_servers
manager.get_allowed_mcp_servers = AsyncMock(return_value=["github"])
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock _get_tools_from_server
async def mock_get_tools_from_server(server, mcp_auth_header=None, mcp_protocol_version=None):
assert mcp_auth_header == "server-specific-token" # Should use server-specific header
async def mock_get_tools_from_server(
server, mcp_auth_header=None, mcp_protocol_version=None
):
assert (
mcp_auth_header == "server-specific-token"
) # Should use server-specific header
tool = MagicMock()
tool.name = "github_tool_1"
return [tool]
manager._get_tools_from_server = mock_get_tools_from_server
# Test with both legacy and server-specific headers
result = await manager.list_tools(
mcp_auth_header="legacy-token",
mcp_server_auth_headers={"github": "server-specific-token"}
mcp_server_auth_headers={"github": "server-specific-token"},
)
assert len(result) == 1
assert result[0].name == "github_tool_1"
@ -211,32 +221,36 @@ class TestMCPServerManager:
async def test_list_tools_handles_missing_server_alias(self):
"""Test that list_tools handles servers without alias gracefully"""
manager = MCPServerManager()
# Mock server without alias
server = MagicMock()
server.name = "github"
server.alias = None # No alias
server.server_name = "github"
# Mock get_allowed_mcp_servers
manager.get_allowed_mcp_servers = AsyncMock(return_value=["github"])
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock _get_tools_from_server
async def mock_get_tools_from_server(server, mcp_auth_header=None, mcp_protocol_version=None):
assert mcp_auth_header == "server-specific-token" # Should use server-specific header via server_name
async def mock_get_tools_from_server(
server, mcp_auth_header=None, mcp_protocol_version=None
):
assert (
mcp_auth_header == "server-specific-token"
) # Should use server-specific header via server_name
tool = MagicMock()
tool.name = "github_tool_1"
return [tool]
manager._get_tools_from_server = mock_get_tools_from_server
# Test with server-specific headers that match server_name (even without alias)
result = await manager.list_tools(
mcp_auth_header="legacy-token",
mcp_server_auth_headers={"github": "server-specific-token"}
mcp_server_auth_headers={"github": "server-specific-token"},
)
assert len(result) == 1
assert result[0].name == "github_tool_1"
@ -244,14 +258,14 @@ class TestMCPServerManager:
async def test_health_check_server_healthy(self):
"""Test health check for a healthy server"""
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.server_id = "test-server"
server.name = "test-server"
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock successful _get_tools_from_server
async def mock_get_tools_from_server(server, mcp_auth_header=None):
tool1 = MagicMock()
@ -259,12 +273,12 @@ class TestMCPServerManager:
tool2 = MagicMock()
tool2.name = "tool2"
return [tool1, tool2]
manager._get_tools_from_server = mock_get_tools_from_server
# Perform health check
result = await manager.health_check_server("test-server")
# Verify results
assert result["server_id"] == "test-server"
assert result["status"] == "healthy"
@ -278,23 +292,23 @@ class TestMCPServerManager:
async def test_health_check_server_unhealthy(self):
"""Test health check for an unhealthy server"""
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.server_id = "test-server"
server.name = "test-server"
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock failed _get_tools_from_server
async def mock_get_tools_from_server(server, mcp_auth_header=None):
raise Exception("Connection timeout")
manager._get_tools_from_server = mock_get_tools_from_server
# Perform health check
result = await manager.health_check_server("test-server")
# Verify results
assert result["server_id"] == "test-server"
assert result["status"] == "unhealthy"
@ -307,13 +321,13 @@ class TestMCPServerManager:
async def test_health_check_server_not_found(self):
"""Test health check for a server that doesn't exist"""
manager = MCPServerManager()
# Mock server not found
manager.get_mcp_server_by_id = MagicMock(return_value=None)
# Perform health check
result = await manager.health_check_server("non-existent-server")
# Verify results
assert result["server_id"] == "non-existent-server"
assert result["status"] == "unknown"
@ -325,22 +339,19 @@ class TestMCPServerManager:
async def test_health_check_all_servers(self):
"""Test health check for all servers"""
manager = MCPServerManager()
# Mock servers
server1 = MagicMock()
server1.server_id = "server1"
server1.name = "server1"
server2 = MagicMock()
server2.server_id = "server2"
server2.name = "server2"
# Mock registry
manager.registry = {
"server1": server1,
"server2": server2
}
manager.registry = {"server1": server1, "server2": server2}
# Mock get_mcp_server_by_id
def mock_get_server_by_id(server_id):
if server_id == "server1":
@ -348,9 +359,9 @@ class TestMCPServerManager:
elif server_id == "server2":
return server2
return None
manager.get_mcp_server_by_id = mock_get_server_by_id
# Mock _get_tools_from_server with different results
async def mock_get_tools_from_server(server, mcp_auth_header=None):
if server.server_id == "server1":
@ -360,22 +371,22 @@ class TestMCPServerManager:
elif server.server_id == "server2":
raise Exception("Connection failed")
return []
manager._get_tools_from_server = mock_get_tools_from_server
# Perform health check for all servers
result = await manager.health_check_all_servers()
# Verify results
assert len(result) == 2
assert "server1" in result
assert "server2" in result
# Check server1 (healthy)
assert result["server1"]["status"] == "healthy"
assert result["server1"]["tools_count"] == 1
assert result["server1"]["error"] is None
# Check server2 (unhealthy)
assert result["server2"]["status"] == "unhealthy"
assert result["server2"]["error"] == "Connection failed"
@ -384,26 +395,26 @@ class TestMCPServerManager:
async def test_health_check_server_with_auth_header(self):
"""Test health check with authentication header"""
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.server_id = "test-server"
server.name = "test-server"
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock _get_tools_from_server to verify auth header is passed
async def mock_get_tools_from_server(server, mcp_auth_header=None):
assert mcp_auth_header == "test-token"
tool = MagicMock()
tool.name = "tool1"
return [tool]
manager._get_tools_from_server = mock_get_tools_from_server
# Perform health check with auth header
result = await manager.health_check_server("test-server", "test-token")
# Verify results
assert result["server_id"] == "test-server"
assert result["status"] == "healthy"
@ -411,4 +422,4 @@ class TestMCPServerManager:
if __name__ == "__main__":
pytest.main([__file__])
pytest.main([__file__])

View file

@ -31,7 +31,6 @@ def generate_mock_mcp_server_db_record(
alias: str = "Test DB Server",
url: str = "https://db-server.example.com/mcp",
transport: str = "sse",
spec_version: str = "2025-03-26",
auth_type: Optional[str] = None,
) -> LiteLLM_MCPServerTable:
"""Generate a mock MCP server record from database"""
@ -41,11 +40,6 @@ def generate_mock_mcp_server_db_record(
alias=alias,
url=url,
transport=MCPTransport.sse if transport == "sse" else MCPTransport.http,
spec_version=(
MCPSpecVersion.mar_2025
if spec_version == "2025-03-26"
else MCPSpecVersion.nov_2024
),
auth_type=MCPAuth.api_key if auth_type == "api_key" else None,
created_at=now,
updated_at=now,
@ -59,7 +53,6 @@ def generate_mock_mcp_server_config_record(
name: str = "Test Config Server",
url: str = "https://config-server.example.com/mcp",
transport: str = "http",
spec_version: str = "2025-03-26",
auth_type: Optional[str] = None,
) -> MCPServer:
"""Generate a mock MCP server record from config.yaml"""