mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
fix: stop including spec_version in MCP server registration inserts
This commit is contained in:
parent
6c291093e9
commit
654f1d3290
14 changed files with 307 additions and 302 deletions
|
|
@ -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"]
|
||||
```
|
||||
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue