From 6ccb20d69d7074c52093cae07d8cb94bbbf6d79e Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Fri, 3 Oct 2025 17:25:08 -0700 Subject: [PATCH] fix(mcp_server_manager.py): mark mcp server as unhealthy and show reason for not being able to serve it on list servers on ui make it easy for user to debug why server can't be added --- .../mcp_server/mcp_server_manager.py | 141 ++++++++---------- litellm/proxy/_types.py | 2 +- 2 files changed, 60 insertions(+), 83 deletions(-) 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'", )