From 654f1d3290e1632abeceae84568ff6027ae71451 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Fri, 19 Sep 2025 06:44:43 +0900 Subject: [PATCH] fix: stop including spec_version in MCP server registration inserts --- docs/my-website/docs/mcp.md | 2 - .../migration.sql | 8 + .../litellm_proxy_extras/schema.prisma | 1 - .../mcp_server/mcp_server_manager.py | 46 +--- .../mcp_server/rest_endpoints.py | 1 - litellm/proxy/_types.py | 94 +++---- litellm/proxy/schema.prisma | 1 - .../types/mcp_server/mcp_server_manager.py | 4 +- schema.prisma | 1 - tests/mcp_tests/test_mcp_auth_priority.py | 1 - tests/mcp_tests/test_mcp_server.py | 15 -- .../test_mcp_servers.py | 237 +++++++++++------- .../mcp_server/test_mcp_server_manager.py | 191 +++++++------- .../test_mcp_management_endpoints.py | 7 - 14 files changed, 307 insertions(+), 302 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20250918083359_drop_spec_version_column_from_mcp_table/migration.sql diff --git a/docs/my-website/docs/mcp.md b/docs/my-website/docs/mcp.md index 18c99051709..80b4c32d0ab 100644 --- a/docs/my-website/docs/mcp.md +++ b/docs/my-website/docs/mcp.md @@ -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"] ``` diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20250918083359_drop_spec_version_column_from_mcp_table/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20250918083359_drop_spec_version_column_from_mcp_table/migration.sql new file mode 100644 index 00000000000..5686876b37c --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20250918083359_drop_spec_version_column_from_mcp_table/migration.sql @@ -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"; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index b8f2201d6b5..2b1e20820f9 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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? diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 79b309d6a30..d0eadb36ba3 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 9d88b979b62..2a9174717d1 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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 ( diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 2ef67c507b2..b0a4e71e23a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index b8f2201d6b5..2b1e20820f9 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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? diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index e7bd67b23e6..eb1eb3250ba 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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 diff --git a/schema.prisma b/schema.prisma index b8f2201d6b5..2b1e20820f9 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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? diff --git a/tests/mcp_tests/test_mcp_auth_priority.py b/tests/mcp_tests/test_mcp_auth_priority.py index bd159ef1a00..1f120b46697 100644 --- a/tests/mcp_tests/test_mcp_auth_priority.py +++ b/tests/mcp_tests/test_mcp_auth_priority.py @@ -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 ) diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 6b9e6e75b57..90ee80ad4bd 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -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 = {} diff --git a/tests/store_model_in_db_tests/test_mcp_servers.py b/tests/store_model_in_db_tests/test_mcp_servers.py index 40273f051e7..8bbcc1f97d9 100644 --- a/tests/store_model_in_db_tests/test_mcp_servers.py +++ b/tests/store_model_in_db_tests/test_mcp_servers.py @@ -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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index b874a834ab8..5f5755077b4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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__]) \ No newline at end of file + pytest.main([__file__]) diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 24bd8be6353..7694fae8405 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -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"""