diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 820824f3331..6ede5d3cf13 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -240,50 +240,57 @@ class MCPServerManager: ) def add_update_server(self, mcp_server: LiteLLM_MCPServerTable): - if mcp_server.server_id not in self.get_registry(): - _mcp_info: MCPInfo = mcp_server.mcp_info or {} - # Use helper to deserialize environment dictionary - # Safely access env field which may not exist on Prisma model objects - env_data = getattr(mcp_server, "env", None) - env_dict = _deserialize_env_dict(env_data) - # Use alias for name if present, else server_name - name_for_prefix = ( - mcp_server.alias or mcp_server.server_name or mcp_server.server_id - ) - # Preserve all custom fields from database while setting defaults for core fields - mcp_info: MCPInfo = _mcp_info.copy() - # Set default values for core fields if not present - if "server_name" not in mcp_info: - mcp_info["server_name"] = mcp_server.server_name or mcp_server.server_id - if "description" not in mcp_info and mcp_server.description: - mcp_info["description"] = mcp_server.description + try: + if mcp_server.server_id not in self.get_registry(): + _mcp_info: MCPInfo = mcp_server.mcp_info or {} + # Use helper to deserialize environment dictionary + # Safely access env field which may not exist on Prisma model objects + env_data = getattr(mcp_server, "env", None) + env_dict = _deserialize_env_dict(env_data) + # Use alias for name if present, else server_name + name_for_prefix = ( + mcp_server.alias or mcp_server.server_name or mcp_server.server_id + ) + # Preserve all custom fields from database while setting defaults for core fields + mcp_info: MCPInfo = _mcp_info.copy() + # Set default values for core fields if not present + if "server_name" not in mcp_info: + mcp_info["server_name"] = ( + mcp_server.server_name or mcp_server.server_id + ) + if "description" not in mcp_info and mcp_server.description: + mcp_info["description"] = mcp_server.description - new_server = MCPServer( - server_id=mcp_server.server_id, - name=name_for_prefix, - alias=getattr(mcp_server, "alias", None), - server_name=getattr(mcp_server, "server_name", None), - url=mcp_server.url, - transport=cast(MCPTransportType, mcp_server.transport), - auth_type=cast(MCPAuthType, mcp_server.auth_type), - mcp_info=mcp_info, - extra_headers=getattr(mcp_server, "extra_headers", None), - # oauth specific fields - client_id=getattr(mcp_server, "client_id", None), - client_secret=getattr(mcp_server, "client_secret", None), - scopes=getattr(mcp_server, "scopes", None), - authorization_url=getattr(mcp_server, "authorization_url", None), - token_url=getattr(mcp_server, "token_url", None), - # Stdio-specific fields - command=getattr(mcp_server, "command", None), - args=getattr(mcp_server, "args", None) or [], - env=env_dict, - access_groups=getattr(mcp_server, "mcp_access_groups", None), - allowed_tools=getattr(mcp_server, "allowed_tools", None), - disallowed_tools=getattr(mcp_server, "disallowed_tools", None), - ) - self.registry[mcp_server.server_id] = new_server - verbose_logger.debug(f"Added MCP Server: {name_for_prefix}") + new_server = MCPServer( + server_id=mcp_server.server_id, + name=name_for_prefix, + alias=getattr(mcp_server, "alias", None), + server_name=getattr(mcp_server, "server_name", None), + url=mcp_server.url, + transport=cast(MCPTransportType, mcp_server.transport), + auth_type=cast(MCPAuthType, mcp_server.auth_type), + mcp_info=mcp_info, + extra_headers=getattr(mcp_server, "extra_headers", None), + # oauth specific fields + client_id=getattr(mcp_server, "client_id", None), + client_secret=getattr(mcp_server, "client_secret", None), + scopes=getattr(mcp_server, "scopes", None), + authorization_url=getattr(mcp_server, "authorization_url", None), + token_url=getattr(mcp_server, "token_url", None), + # Stdio-specific fields + command=getattr(mcp_server, "command", None), + args=getattr(mcp_server, "args", None) or [], + env=env_dict, + access_groups=getattr(mcp_server, "mcp_access_groups", None), + allowed_tools=getattr(mcp_server, "allowed_tools", None), + disallowed_tools=getattr(mcp_server, "disallowed_tools", None), + ) + self.registry[mcp_server.server_id] = new_server + verbose_logger.debug(f"Added MCP Server: {name_for_prefix}") + + except Exception as e: + verbose_logger.debug(f"Failed to add MCP server: {str(e)}") + raise e def get_all_mcp_server_ids(self) -> Set[str]: """ @@ -1183,49 +1190,19 @@ class MCPServerManager: } ) - ## filter out invalid servers + ## mark invalid servers w/ reason for being invalid valid_server_ids = self.get_all_mcp_server_ids() - filtered_list_mcp_servers = [] for server in list_mcp_servers: if server.server_id in valid_server_ids: - filtered_list_mcp_servers.append(server) + server.status = "unhealthy" + ## try adding server to registry to get error + try: + self.add_update_server(server) + except Exception as e: + server.health_check_error = str(e) + server.health_check_error = "Server is not in in memory registry yet. This could be a temporary sync issue." - # Map servers to their teams and return with health data - from typing import cast - - return [ - LiteLLM_MCPServerTable( - server_id=server.server_id, - server_name=server.server_name, - alias=server.alias, - description=server.description, - url=server.url, - transport=server.transport, - auth_type=server.auth_type, - created_at=server.created_at, - created_by=server.created_by, - updated_at=server.updated_at, - updated_by=server.updated_by, - mcp_access_groups=( - server.mcp_access_groups - if server.mcp_access_groups is not None - else [] - ), - allowed_tools=( - server.allowed_tools if server.allowed_tools is not None else [] - ), - mcp_info=server.mcp_info, - teams=cast( - List[Dict[str, str | None]], - server_to_teams_map.get(server.server_id, []), - ), - # Stdio-specific fields - command=getattr(server, "command", None), - args=getattr(server, "args", None) or [], - env=getattr(server, "env", None) or {}, - ) - for server in filtered_list_mcp_servers - ] + return list_mcp_servers async def reload_servers_from_database(self): """ diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index efe7ff90973..12f2b7b297e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -989,7 +989,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase): allowed_tools: List[str] = Field(default_factory=list) mcp_info: Optional[MCPInfo] = None # Health check status - status: Optional[str] = Field( + status: Optional[Literal["healthy", "unhealthy", "unknown"]] = Field( default="unknown", description="Health status: 'healthy', 'unhealthy', 'unknown'", )