From dca1d06065f7feba84a846662120502c7d8890ee Mon Sep 17 00:00:00 2001 From: moe-berri Date: Thu, 24 Sep 2026 15:47:10 -0700 Subject: [PATCH] feat(mcp): check reachability without per-user credentials --- litellm/models/mcp_server.py | 3 + .../mcp_server/mcp_server_manager.py | 50 +++- litellm/proxy/_lazy_openapi_snapshot.json | 110 +++++++- .../mcp_management_endpoints.py | 31 ++- .../mcp_server/test_mcp_env_vars.py | 29 +-- .../mcp_server/test_mcp_server_manager.py | 246 +++++++++++++++--- .../test_mcp_management_endpoints.py | 16 +- .../mcpServers/useMCPServerHealth.test.ts | 19 +- .../hooks/mcpServers/useMCPServerHealth.ts | 6 +- .../_components/MCPServerCard.test.tsx | 14 + .../mcp-servers/_components/MCPServerCard.tsx | 14 +- ...t.tsx => mcp_servers.integration.test.tsx} | 13 +- .../mcp-servers/_components/mcp_servers.tsx | 5 +- .../src/components/mcp_tools/types.tsx | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 17 +- 15 files changed, 486 insertions(+), 88 deletions(-) rename ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/{mcp_servers.test.tsx => mcp_servers.integration.test.tsx} (97%) diff --git a/litellm/models/mcp_server.py b/litellm/models/mcp_server.py index 9125d708e79..6cc6936449b 100644 --- a/litellm/models/mcp_server.py +++ b/litellm/models/mcp_server.py @@ -79,6 +79,9 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase): ) last_health_check: datetime | None = None health_check_error: str | None = None + health_check_type: Literal["liveness", "protocol"] | None = Field( + default=None, json_schema_extra={"readOnly": True} + ) command: str | None = None args: list[str] = Field(default_factory=list) env: dict[str, str] = Field(default_factory=dict) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 312dcb27d89..9e2dbc05a65 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -67,7 +67,7 @@ from litellm.integrations.custom_guardrail import ( _sync_guardrail_info_to_logging_obj, # pyright: ignore[reportPrivateUsage] - the same bridge @log_guardrail_information uses; reimplementing it here would fork the metadata-key logic ) from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get -from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client, get_ssl_configuration from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, MCPServerAccess, @@ -6913,6 +6913,29 @@ class MCPServerManager: # Take first 32 characters and format as UUID-like string return hash_hex[:32] + @staticmethod + async def _check_mcp_liveness(url: str) -> tuple[Literal["healthy", "unhealthy", "unknown"], str | None]: + async def probe() -> None: + endpoint: Final = httpx.URL(url).copy_with(username="", password="") + async with httpx.AsyncClient( + verify=get_ssl_configuration(), + trust_env=False, + follow_redirects=False, + timeout=MCP_HEALTH_CHECK_TIMEOUT, + ) as client: + async with client.stream("GET", endpoint, auth=None): + pass + + try: + await asyncio.wait_for(probe(), timeout=MCP_HEALTH_CHECK_TIMEOUT) + return "healthy", None + except asyncio.TimeoutError: + return "unhealthy", f"Health check timed out after {MCP_HEALTH_CHECK_TIMEOUT} seconds" + except asyncio.CancelledError: + return "unknown", "Health check was cancelled" + except Exception as exc: + return "unhealthy", f"Liveness check failed ({type(exc).__name__})" + async def health_check_server(self, server_id: str, mcp_auth_header: str | None = None) -> LiteLLM_MCPServerTable: """ Perform a health check on a specific MCP server. @@ -6953,11 +6976,7 @@ class MCPServerManager: status: Literal["healthy", "unhealthy", "unknown"] = "unknown" health_check_error = None - # Check if we should skip health check based on auth configuration - should_skip_health_check = False - - # Skip if server requires per-user authentication (OAuth2 or passthrough auth) - if ( + should_skip_health_check: Final = ( server.requires_per_user_auth or ( server.auth_type @@ -6966,8 +6985,22 @@ class MCPServerManager: and not server.authentication_token ) or self._references_per_user_env_var(server) + ) + if ( + should_skip_health_check + and server.transport in (MCPTransport.http, MCPTransport.sse) + and server.url is not None + and server.url.lower().startswith(("http://", "https://")) ): - should_skip_health_check = True + liveness_status, liveness_error = await self._check_mcp_liveness(server.url) + return self._build_mcp_server_table(server).model_copy( + update={ + "status": liveness_status, + "health_check_error": liveness_error, + "health_check_type": "liveness", + "last_health_check": datetime.now(), + } + ) if not should_skip_health_check: try: @@ -6998,7 +7031,7 @@ class MCPServerManager: health_check_error = "Health check was cancelled" status = "unknown" except Exception as e: - health_check_error = str(e) + health_check_error = f"Health check failed ({type(e).__name__})" status = "unhealthy" return LiteLLM_MCPServerTable( @@ -7021,6 +7054,7 @@ class MCPServerManager: status=status, last_health_check=datetime.now(), health_check_error=health_check_error, + health_check_type="protocol" if not should_skip_health_check else None, command=getattr(server, "command", None), args=getattr(server, "args", None) or [], env=getattr(server, "env", None) or {}, diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 0b43c3864ab..5064b713425 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -31685,6 +31685,22 @@ ], "title": "Health Check Error" }, + "health_check_type": { + "anyOf": [ + { + "enum": [ + "liveness", + "protocol" + ], + "type": "string" + }, + { + "type": "null" + } + ], + "readOnly": true, + "title": "Health Check Type" + }, "instructions": { "anyOf": [ { @@ -34741,6 +34757,22 @@ ], "title": "Health Check Error" }, + "health_check_type": { + "anyOf": [ + { + "enum": [ + "liveness", + "protocol" + ], + "type": "string" + }, + { + "type": "null" + } + ], + "readOnly": true, + "title": "Health Check Type" + }, "instructions": { "anyOf": [ { @@ -35955,6 +35987,76 @@ "title": "MCPOAuthUserCredentialStatus", "type": "object" }, + "MCPServerHealthResponse": { + "properties": { + "health_check_error": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Health Check Error" + }, + "health_check_type": { + "anyOf": [ + { + "enum": [ + "liveness", + "protocol" + ], + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Health Check Type" + }, + "last_health_check": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Last Health Check" + }, + "server_id": { + "title": "Server Id", + "type": "string" + }, + "status": { + "anyOf": [ + { + "enum": [ + "healthy", + "unhealthy", + "unknown" + ], + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Status" + } + }, + "required": [ + "server_id", + "status", + "last_health_check", + "health_check_error" + ], + "title": "MCPServerHealthResponse", + "type": "object" + }, "MCPServerUserCredentialListItem": { "description": "One user's stored credential for an MCP server, as an admin sees it. Never carries the secret.", "properties": { @@ -37808,7 +37910,13 @@ "200": { "content": { "application/json": { - "schema": {} + "schema": { + "items": { + "$ref": "#/components/schemas/MCPServerHealthResponse" + }, + "title": "Response Health Check Servers V1 Mcp Server Health Get", + "type": "array" + } } }, "description": "Successful Response" diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index aa218f42023..03fab44566e 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -44,7 +44,7 @@ from fastapi import ( status, ) from fastapi.responses import JSONResponse -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import NotRequired, ReadOnly, TypedDict try: from prisma.errors import RecordNotFoundError, UniqueViolationError @@ -96,6 +96,14 @@ class _HasServerId(Protocol): server_id: str +class MCPServerHealthResponse(TypedDict): + server_id: ReadOnly[str] + status: ReadOnly[Literal["healthy", "unhealthy", "unknown"] | None] + health_check_type: ReadOnly[NotRequired[Literal["liveness", "protocol"] | None]] + last_health_check: ReadOnly[datetime | None] + health_check_error: ReadOnly[str | None] + + def does_mcp_server_exist(mcp_server_records: Iterable[_HasServerId], mcp_server_id: str) -> bool: """ Check if the mcp server with the given id exists in the iterable of mcp servers. @@ -237,6 +245,15 @@ if MCP_AVAILABLE: ) from litellm.types.mcp_server.mcp_server_manager import MCPServer + def _mcp_server_health_response(server: LiteLLM_MCPServerTable) -> MCPServerHealthResponse: + return { + "server_id": server.server_id, + "status": server.status, + "health_check_type": server.health_check_type, + "last_health_check": server.last_health_check, + "health_check_error": server.health_check_error, + } + @dataclass class _TemporaryMCPServerEntry: server: MCPServer @@ -1297,6 +1314,7 @@ if MCP_AVAILABLE: @router.get( "/server/health", description="Health check for MCP servers", + response_model=list[MCPServerHealthResponse], dependencies=[Depends(user_api_key_auth)], ) async def health_check_servers( @@ -1305,7 +1323,7 @@ if MCP_AVAILABLE: description="Server IDs to check. If not provided, checks all accessible servers.", ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - ): + ) -> list[MCPServerHealthResponse]: """ Perform health checks on one or more MCP servers. @@ -1329,11 +1347,11 @@ if MCP_AVAILABLE: if user_mcp_management_mode == "view_all" and not _is_restricted_virtual_key_request(user_api_key_dict): servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_unfiltered(server_ids=server_ids) - return [{"server_id": server.server_id, "status": server.status} for server in servers] + return [_mcp_server_health_response(server) for server in servers] auth_contexts: Final = await build_effective_auth_contexts(user_api_key_dict) - server_status_map: Final[dict[str, Literal["healthy", "unhealthy", "unknown"] | None]] = {} + server_status_map: Final[dict[str, MCPServerHealthResponse]] = {} for auth_context in auth_contexts: servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_and_teams( user_api_key_auth=auth_context, @@ -1341,9 +1359,9 @@ if MCP_AVAILABLE: ) for server in servers: if server.server_id not in server_status_map: - server_status_map[server.server_id] = server.status + server_status_map[server.server_id] = _mcp_server_health_response(server) - return [{"server_id": server_id, "status": status} for server_id, status in server_status_map.items()] + return list(server_status_map.values()) @router.post( "/server/register", @@ -1688,6 +1706,7 @@ if MCP_AVAILABLE: mcp_server.status = health_result.status if health_result.status else "unknown" mcp_server.last_health_check = health_result.last_health_check mcp_server.health_check_error = health_result.health_check_error + mcp_server.health_check_type = health_result.health_check_type except Exception as e: verbose_proxy_logger.debug("Error performing health check on server %s: %s", server_id, e) mcp_server.status = "unknown" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py index 93b894f7645..70609308490 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py @@ -615,33 +615,18 @@ def test_references_per_user_env_var(static_headers, env_vars, expected): @pytest.mark.asyncio -async def test_health_check_skips_servers_referencing_per_user_env_var( - mock_server, monkeypatch -): - """A userless health probe cannot fill per-user ${NAME} placeholders, so a - server whose static_headers reference one must report 'unknown' without - connecting. Otherwise it forwards the literal placeholder upstream, gets a - 401, and flips to 'unhealthy' even though real user calls succeed.""" - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - MCPServerManager, - ) +async def test_health_check_probes_without_per_user_env_var(mock_server, respx_mock): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager manager = MCPServerManager() manager.registry[mock_server.server_id] = mock_server - - created = [] - - async def fake_create_client(*args, **kwargs): - created.append((args, kwargs)) - raise RuntimeError("upstream rejected literal ${NAME}") - - monkeypatch.setattr(manager, "_create_mcp_client", fake_create_client) - + upstream = respx_mock.get(mock_server.url).respond(401) result = await manager.health_check_server(mock_server.server_id) - - assert created == [] - assert result.status == "unknown" + assert result.status == "healthy" + assert result.health_check_type == "liveness" assert result.health_check_error is None + assert "authorization" not in upstream.calls.last.request.headers + assert not any("${" in value for value in upstream.calls.last.request.headers.values()) # ── _load_user_env_vars guard paths ──────────────────────────────────────── 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 cd5dae1269a..6bf3b22ddb5 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 @@ -1196,7 +1196,7 @@ class TestMCPServerManager: with ( patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)), - pytest.raises(ValueError) as exc_info, + pytest.raises(ValueError, match="oauth2_flow: client_credentials") as exc_info, ): await manager.load_servers_from_config(self._oauth2_config()) @@ -1209,7 +1209,7 @@ class TestMCPServerManager: with ( patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)), - pytest.raises(ValueError) as exc_info, + pytest.raises(ValueError, match="got 'm2m'") as exc_info, ): await manager.load_servers_from_config(self._oauth2_config(oauth2_flow="m2m")) @@ -1476,7 +1476,7 @@ class TestMCPServerManager: with ( patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)), - pytest.raises(ValueError, match="per_server_oauth_discovery.*must be a boolean"), + pytest.raises(ValueError, match=r"per_server_oauth_discovery.*must be a boolean"), ): await manager.load_servers_from_config( self._oauth2_config(oauth2_flow="authorization_code", per_server_oauth_discovery="yes") @@ -1488,7 +1488,7 @@ class TestMCPServerManager: with ( patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)), - pytest.raises(ValueError) as exc_info, + pytest.raises(ValueError, match="dcr_bridge is only supported") as exc_info, ): await manager.load_servers_from_config( self._oauth2_config(oauth2_flow="authorization_code", dcr_bridge=True) @@ -1502,7 +1502,7 @@ class TestMCPServerManager: with ( patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)), - pytest.raises(ValueError) as exc_info, + pytest.raises(ValueError, match="must be a boolean") as exc_info, ): await manager.load_servers_from_config( self._client_forwarded_config(MCPAuth.true_passthrough, dcr_bridge="yes") @@ -4815,6 +4815,7 @@ class TestMCPServerManager: assert isinstance(result, LiteLLM_MCPServerTable) assert result.server_id == "test-server" assert result.status == "healthy" + assert result.health_check_type == "protocol" assert result.health_check_error is None assert result.last_health_check is not None @@ -4847,7 +4848,7 @@ class TestMCPServerManager: assert isinstance(result, LiteLLM_MCPServerTable) assert result.server_id == "test-server" assert result.status == "unhealthy" - assert result.health_check_error == "Connection timeout" + assert result.health_check_error == "Health check failed (Exception)" assert result.last_health_check is not None @pytest.mark.asyncio @@ -4871,7 +4872,7 @@ class TestMCPServerManager: result = await manager.health_check_server(server.server_id) assert result.status == "unhealthy" - assert "OAuth discovery unavailable" in (result.health_check_error or "") + assert result.health_check_error == "Health check failed (HTTPException)" @pytest.mark.asyncio async def test_health_check_server_not_found(self): @@ -4893,8 +4894,9 @@ class TestMCPServerManager: assert result.last_health_check is not None @pytest.mark.asyncio - async def test_health_check_server_oauth2_skips_check(self): - """Test that health check is skipped for OAuth2 servers and returns unknown status""" + async def test_health_check_server_oauth2_checks_liveness(self, respx_mock): + """OAuth servers can report reachability without a protocol handshake""" + respx_mock.get("http://oauth2-server.com").respond(401) manager = MCPServerManager() # Mock OAuth2 server @@ -4920,13 +4922,15 @@ class TestMCPServerManager: # Verify results assert isinstance(result, LiteLLM_MCPServerTable) assert result.server_id == "oauth2-server" - assert result.status == "unknown" + assert result.status == "healthy" + assert result.health_check_type == "liveness" assert result.health_check_error is None assert result.last_health_check is not None @pytest.mark.asyncio - async def test_health_check_server_no_token_skips_check(self): - """Test that health check is skipped when auth_type is set but authentication_token is missing""" + async def test_health_check_server_no_token_checks_liveness(self, respx_mock): + """Missing static credentials still permit a credential-free liveness check""" + respx_mock.get("http://no-token-server.com").respond(401) manager = MCPServerManager() # Mock server with auth_type but no authentication_token @@ -4953,7 +4957,8 @@ class TestMCPServerManager: # Verify results assert isinstance(result, LiteLLM_MCPServerTable) assert result.server_id == "no-token-server" - assert result.status == "unknown" + assert result.status == "healthy" + assert result.health_check_type == "liveness" assert result.health_check_error is None assert result.last_health_check is not None @@ -4999,11 +5004,13 @@ class TestMCPServerManager: assert isinstance(result, LiteLLM_MCPServerTable) assert result.server_id == "test-server" assert result.status == "healthy" + assert result.health_check_type == "protocol" assert result.health_check_error is None @pytest.mark.asyncio - async def test_health_check_skips_passthrough_auth_with_authorization_header(self): - """Test that health check is skipped for servers with passthrough Authorization header""" + async def test_health_check_checks_liveness_with_forwarded_authorization(self, respx_mock): + """Forwarded Authorization does not reach the liveness probe""" + respx_mock.get("http://github-server.com").respond(401) manager = MCPServerManager() # Mock server with auth_type=none and Authorization in extra_headers (passthrough auth) @@ -5031,13 +5038,15 @@ class TestMCPServerManager: # Verify results assert isinstance(result, LiteLLM_MCPServerTable) assert result.server_id == "github-server" - assert result.status == "unknown" + assert result.status == "healthy" + assert result.health_check_type == "liveness" assert result.health_check_error is None assert result.last_health_check is not None @pytest.mark.asyncio - async def test_health_check_skips_passthrough_auth_with_api_key_header(self): - """Test that health check is skipped for servers with passthrough x-api-key header""" + async def test_health_check_checks_liveness_with_forwarded_api_key(self, respx_mock): + """Forwarded API keys do not reach the liveness probe""" + respx_mock.get("http://sourcegraph-server.com").respond(401) manager = MCPServerManager() # Mock server with auth_type=none and x-api-key in extra_headers @@ -5065,7 +5074,8 @@ class TestMCPServerManager: # Verify results assert isinstance(result, LiteLLM_MCPServerTable) assert result.server_id == "sourcegraph-server" - assert result.status == "unknown" + assert result.status == "healthy" + assert result.health_check_type == "liveness" assert result.health_check_error is None assert result.last_health_check is not None @@ -5811,7 +5821,7 @@ class TestMCPServerManager: def test_resolve_mcp_server_for_tool_call_raises_when_not_found(self): """ValueError is raised when no resolution path finds the tool.""" manager = MCPServerManager() - with pytest.raises(ValueError, match="Tool .* not found"): + with pytest.raises(ValueError, match=r"Tool .* not found"): manager._resolve_mcp_server_for_tool_call("nonexistent", "ghost_tool") def test_resolve_mcp_server_for_tool_call_unscoped_cached_tool_still_fails(self): @@ -6046,25 +6056,24 @@ class TestMCPServerManager: assert result is None @pytest.mark.asyncio - async def test_has_user_oauth_token_delegates_to_provider(self): + @pytest.mark.parametrize("verdict", (True, False)) + async def test_has_user_oauth_token_delegates_to_provider(self, verdict: bool): """has_user_oauth_token maps the server and delegates the verdict to the v2 resolver.""" from litellm.proxy._types import UserAPIKeyAuth - for verdict in (True, False): + class _Provider: + async def has_user_token(self, subject, spec): + return verdict - class _Provider: - async def has_user_token(self, subject, spec): - return verdict - - manager = MCPServerManager(cred_provider=_Provider()) - server = MCPServer( - server_id="s", - name="n", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, - ) - user_auth = UserAPIKeyAuth(api_key="sk", user_id="alice") - assert await manager.has_user_oauth_token(server, user_auth) is verdict + manager = MCPServerManager(cred_provider=_Provider()) + server = MCPServer( + server_id="s", + name="n", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + ) + user_auth = UserAPIKeyAuth(api_key="sk", user_id="alice") + assert await manager.has_user_oauth_token(server, user_auth) is verdict @pytest.mark.asyncio async def test_has_user_oauth_token_short_circuits_for_unmigrated_server(self): @@ -13485,7 +13494,6 @@ class _DiscoveryClock: return self.now -from pydantic import TypeAdapter from mcp.types import JSONRPCMessage _JSONRPC_ADAPTER = TypeAdapter(JSONRPCMessage) @@ -14575,3 +14583,169 @@ class TestSharedIdentifierPrefixWarning: assert "srv-b" in shared_warnings[0] assert "srv-c" not in shared_warnings[0] assert "'shared'" in shared_warnings[0] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("transport", (MCPTransport.http, MCPTransport.sse)) +@pytest.mark.parametrize("status_code", (200, 302, 401, 403, 404, 405, 500)) +async def test_liveness_accepts_headers_without_body_auth_or_redirects( + respx_mock: MockRouter, transport: MCPTransport, status_code: int +) -> None: + from collections.abc import AsyncIterator + + class UnreadBody(httpx.AsyncByteStream): + closed = False + + async def __aiter__(self) -> AsyncIterator[bytes]: + raise AssertionError("A liveness probe must not read the response body") + yield b"" + + async def aclose(self) -> None: + self.closed = True + + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="liveness", name="liveness", transport=transport, + url="http://url-user:url-secret@127.0.0.1/mcp", auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", authentication_token="stored-secret", + static_headers={"Authorization": "Bearer static-secret", "Cookie": "session=secret", "X-Api-Key": "secret"}, + extra_headers=["Authorization", "Cookie"], + ) + manager.registry[server.server_id] = server + body: Final = UnreadBody() + upstream: Final = respx_mock.get("http://127.0.0.1/mcp").mock( + return_value=httpx.Response(status_code, headers={"Location": "https://other.invalid/", "Content-Type": "text/event-stream"}, stream=body) + ) + + result: Final = await manager.health_check_server(server.server_id, mcp_auth_header="Bearer forwarded-secret") + + assert result.status == "healthy" + assert result.health_check_type == "liveness" + assert result.health_check_error is None + assert result.last_health_check is not None + assert body.closed + assert len(respx_mock.calls) == 1 + request: Final = upstream.calls.last.request + assert not {"authorization", "cookie", "x-api-key"}.intersection(request.headers) + assert request.url.username == request.url.password == "" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", (httpx.ConnectError, httpx.ReadTimeout)) +async def test_liveness_network_failures_do_not_expose_credentials(respx_mock: MockRouter, failure: type[httpx.RequestError]) -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="liveness", name="liveness", transport="http", url="https://example.invalid/mcp", auth_type="oauth2") + manager.registry[server.server_id] = server + respx_mock.get(server.url).mock(side_effect=failure("secret in https://user:password@example.invalid/mcp")) + result: Final = await manager.health_check_server(server.server_id) + assert result.status == "unhealthy" + assert result.health_check_type == "liveness" + assert result.health_check_error == f"Liveness check failed ({failure.__name__})" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("cancel", (False, True)) +async def test_liveness_whole_operation_timeout_and_cancellation( + respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, cancel: bool +) -> None: + started: Final = asyncio.Event() + finished: Final = asyncio.Event() + async def pending(request: httpx.Request) -> httpx.Response: + started.set() + try: + await asyncio.Event().wait() + finally: + finished.set() + return httpx.Response(200) + + monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCP_HEALTH_CHECK_TIMEOUT", 0.2) + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="liveness", name="liveness", transport="http", url="http://127.0.0.1/mcp", auth_type="oauth2") + manager.registry[server.server_id] = server + respx_mock.get(server.url).mock(side_effect=pending) + task: Final = asyncio.create_task(manager.health_check_server(server.server_id)) + await asyncio.wait_for(started.wait(), timeout=2) + if cancel: + task.cancel() + result: Final = await asyncio.wait_for(task, timeout=2) + assert finished.is_set() + assert result.status == ("unknown" if cancel else "unhealthy") + assert result.health_check_error == ("Health check was cancelled" if cancel else "Health check timed out after 0.2 seconds") + assert result.health_check_type == "liveness" + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("transport", "url"), (("stdio", "https://example.invalid"), ("http", None), ("sse", "file:///tmp/mcp"))) +async def test_liveness_unsupported_servers_stay_unknown(respx_mock: MockRouter, transport: str, url: str | None) -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="unsupported", name="unsupported", transport=transport, url=url, auth_type="oauth2", command="unused") + manager.registry[server.server_id] = server + result: Final = await manager.health_check_server(server.server_id) + assert result.status == "unknown" + assert result.health_check_type is None + assert not respx_mock.calls + + +@pytest.mark.asyncio +@pytest.mark.parametrize("ssl_setting", ("untrusted", "ca_file", "verify_path", "disabled")) +async def test_liveness_tls_and_infinite_sse_close_after_headers( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ssl_setting: str +) -> None: + import ipaddress + import ssl + from datetime import timedelta, timezone + from cryptography import x509 + from cryptography.hazmat.primitives import hashes, serialization + from cryptography.hazmat.primitives.asymmetric import rsa + from cryptography.x509.oid import NameOID + + key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + subject: Final = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "localhost")]) + now: Final = datetime.now(timezone.utc) + cert: Final = ( + x509.CertificateBuilder().subject_name(subject).issuer_name(subject).public_key(key.public_key()) + .serial_number(x509.random_serial_number()).not_valid_before(now - timedelta(days=1)) + .not_valid_after(now + timedelta(days=1)) + .add_extension(x509.SubjectAlternativeName([x509.IPAddress(ipaddress.ip_address("127.0.0.1"))]), critical=False) + .sign(key, hashes.SHA256()) + ) + cert_path: Final = tmp_path / "ca.pem" + key_path: Final = tmp_path / "key.pem" + cert_path.write_bytes(cert.public_bytes(serialization.Encoding.PEM)) + key_path.write_bytes(key.private_bytes(serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption())) + context: Final = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + context.load_cert_chain(cert_path, key_path) + monkeypatch.delenv("SSL_CERT_FILE", raising=False) + monkeypatch.setenv("HTTPS_PROXY", "http://proxy-user:proxy-secret@127.0.0.1:1") + monkeypatch.setenv("NO_PROXY", "") + monkeypatch.setenv("SSL_VERIFY", "False" if ssl_setting == "disabled" else str(cert_path) if ssl_setting == "verify_path" else "True") + if ssl_setting == "ca_file": + monkeypatch.setenv("SSL_CERT_FILE", str(cert_path)) + closed: Final = asyncio.Event() + + async def serve(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + try: + await reader.readuntil(b"\r\n\r\n") + writer.write(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n") + await writer.drain() + await reader.read() + finally: + writer.close() + await writer.wait_closed() + closed.set() + + upstream: Final = await asyncio.start_server(serve, "127.0.0.1", 0, ssl=context) + async with upstream: + port: Final = upstream.sockets[0].getsockname()[1] + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="tls", name="tls", transport="sse", url=f"https://127.0.0.1:{port}/sse", auth_type="oauth2") + manager.registry[server.server_id] = server + result: Final = await asyncio.wait_for(manager.health_check_server(server.server_id), timeout=3) + if ssl_setting == "untrusted": + assert result.status == "unhealthy" + assert result.health_check_error == "Liveness check failed (ConnectError)" + else: + assert result.status == "healthy" + assert result.health_check_error is None + await asyncio.wait_for(closed.wait(), timeout=1) + assert result.health_check_type == "liveness" 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 557e753a76f..d2e622b0a8e 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 @@ -4225,7 +4225,9 @@ class TestHealthCheckServers: ("view_all", True, ("server-x",), None, ("server-x",), 503), ], ) +@pytest.mark.parametrize("probe_kind", ("openapi", "liveness")) async def test_health_discovery_respects_route_restricted_key_grants( + probe_kind: str, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, mode: str, @@ -4247,13 +4249,14 @@ async def test_health_discovery_respects_route_restricted_key_grants( server_id=server_id, name=server_id, transport=MCPTransport.http, - spec_path=f"https://93.184.216.34/{server_id}.json", - auth_type=MCPAuth.none, + spec_path=f"https://93.184.216.34/{server_id}.json" if probe_kind == "openapi" else None, + url=f"http://127.0.0.1/{server_id}" if probe_kind == "liveness" else None, + auth_type=MCPAuth.oauth2 if probe_kind == "liveness" else MCPAuth.none, ) for server_id in ("server-x", "server-y") } routes: Final = { - server_id: respx_mock.get(server.spec_path).respond(upstream_status, json={"paths": {}}) + server_id: respx_mock.get(server.spec_path or server.url).respond(upstream_status, json={"paths": {}}) for server_id, server in manager.registry.items() } caller: Final = UserAPIKeyAuth( @@ -4288,9 +4291,14 @@ async def test_health_discovery_respects_route_restricted_key_grants( assert {row["server_id"] for row in result} == set(expected) assert {server_id for server_id, route in routes.items() if route.called} == set(expected) - expected_status: Final = {200: "healthy", 503: "unhealthy"}[upstream_status] + expected_status: Final = "healthy" if probe_kind == "liveness" else {200: "healthy", 503: "unhealthy"}[upstream_status] assert all(row["status"] == expected_status for row in result) + assert all(row["last_health_check"] is not None for row in result) + if probe_kind == "liveness": + assert all(row["health_check_type"] == "liveness" for row in result) + assert all(row["health_check_error"] is None for row in result) + class TestMCPRegistryEndpoint: def test_registry_returns_404_when_flag_missing(self): diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServerHealth.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServerHealth.test.ts index 567e1d23013..11ef743cf90 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServerHealth.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServerHealth.test.ts @@ -1,6 +1,6 @@ /* @vitest-environment jsdom */ import React from "react"; -import { renderHook, waitFor } from "@testing-library/react"; +import { act, renderHook, waitFor } from "@testing-library/react"; import { describe, it, expect, vi, beforeEach } from "vitest"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { useMCPServerHealth } from "./useMCPServerHealth"; @@ -73,6 +73,23 @@ describe("useMCPServerHealth", () => { expect(result.current.error).toEqual(mockError); }); + it("replaces probe metadata when a server is rechecked", async () => { + const initial = [ + { server_id: "server-1", status: "healthy", health_check_type: "liveness", health_check_error: null }, + ]; + const updated = [ + { server_id: "server-1", status: "unhealthy", health_check_type: "protocol", health_check_error: "Check failed" }, + ]; + vi.mocked(networking.fetchMCPServerHealth).mockResolvedValueOnce(initial).mockResolvedValueOnce(updated); + const { result } = renderHook(() => useMCPServerHealth(), { wrapper }); + await waitFor(() => expect(result.current.data).toEqual(initial)); + await act(async () => { + await result.current.recheckServerHealth("server-1"); + }); + expect(result.current.data).toEqual(updated); + expect(result.current.recheckingServerIds.size).toBe(0); + }); + it("should not fetch when accessToken is not available", async () => { // Mock useAuthorized to return no token const useAuthorizedModule = await import("../useAuthorized"); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServerHealth.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServerHealth.ts index 9ad8a6f43fa..4fdd47eab9d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServerHealth.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServerHealth.ts @@ -1,3 +1,4 @@ +import type { components } from "@/lib/http/schema"; import { useCallback, useState } from "react"; import { useQuery, useQueryClient } from "@tanstack/react-query"; import { createQueryKeys } from "../common/queryKeysFactory"; @@ -6,10 +7,7 @@ import useAuthorized from "../useAuthorized"; const mcpServerHealthKeys = createQueryKeys("mcpServerHealth"); -interface MCPServerHealth { - server_id: string; - status: string; -} +type MCPServerHealth = components["schemas"]["MCPServerHealthResponse"]; export const useMCPServerHealth = () => { const { accessToken } = useAuthorized(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx index 71c2e107774..aa3087a2163 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx @@ -112,3 +112,17 @@ describe("MCPServerCard per-user credentials", () => { expect(screen.queryByRole("button", { name: "Set" })).not.toBeInTheDocument(); }); }); + +describe("MCPServerCard health qualification", () => { + it.each(["protocol", null, undefined] as const)("keeps Healthy for %s checks", (health_check_type) => { + renderCard({ status: "healthy", health_check_type }); + expect(screen.getByText("Healthy")).toBeInTheDocument(); + expect(screen.queryByText("Reachable")).not.toBeInTheDocument(); + }); + + it.each(["unhealthy", "unknown"] as const)("does not call %s liveness reachable", (status) => { + renderCard({ status, health_check_type: "liveness" }); + expect(screen.getByText(status === "unhealthy" ? "Unhealthy" : "Unknown")).toBeInTheDocument(); + expect(screen.queryByText("Reachable")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx index 42fb95d5951..b347044c2ed 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx @@ -213,6 +213,7 @@ const MCPServerCard: FC = ({ isLoadingHealth={isLoadingHealth} isRechecking={isRechecking} onRecheck={onRecheckHealth} + healthCheckType={server.health_check_type} lastCheck={server.last_health_check} error={server.health_check_error} dotClass={healthTone.dot} @@ -307,6 +308,7 @@ const MCPServerCard: FC = ({ interface HealthChipProps { status: string; + healthCheckType?: MCPServer["health_check_type"]; isLoadingHealth?: boolean; isRechecking?: boolean; onRecheck?: () => void; @@ -317,6 +319,7 @@ interface HealthChipProps { const HealthChip: FC = ({ status, + healthCheckType, isLoadingHealth, isRechecking, onRecheck, @@ -332,6 +335,10 @@ const HealthChip: FC = ({ ); } + const label = + status === "healthy" && healthCheckType === "liveness" + ? "Reachable" + : status.charAt(0).toUpperCase() + status.slice(1); return ( = ({ } > - {status.charAt(0).toUpperCase() + status.slice(1)} + {label} } /> -
Health: {status}
+
Health: {label}
+ {healthCheckType === "liveness" && ( +
Authentication and tools were not checked
+ )} {lastCheck &&
Last check: {new Date(lastCheck).toLocaleString()}
} {error && (
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.integration.test.tsx similarity index 97% rename from ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.integration.test.tsx index 1217d878489..ebb12c42e38 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.integration.test.tsx @@ -397,7 +397,12 @@ describe("MCPServers", () => { // Mock health status data const mockHealthStatuses = [ - { server_id: "server-1", status: "healthy" }, + { + server_id: "server-1", + status: "healthy", + health_check_type: "liveness", + last_health_check: "2026-01-02T00:00:00Z", + }, { server_id: "server-2", status: "unhealthy" }, ]; @@ -416,6 +421,12 @@ describe("MCPServers", () => { expect(screen.getByText("MCP Servers")).toBeInTheDocument(); }); + expect(await screen.findByText("Reachable")).toBeInTheDocument(); + expect(await screen.findByText("Unhealthy")).toBeInTheDocument(); + await userEvent.hover(screen.getByText("Reachable")); + expect(await screen.findByText("Authentication and tools were not checked")).toBeInTheDocument(); + expect(await screen.findByText(/Last check:/)).toBeInTheDocument(); + // Verify the health check API was called (without a server ID filter — the hook always // fetches health for all servers so the query key stays stable) await waitFor(() => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx index b4b7ab6b3c8..6d7b8156b7a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx @@ -185,13 +185,14 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i if (!mcpServers) return []; if (!healthStatuses) return mcpServers; - const healthMap = new Map(healthStatuses.map((h) => [h.server_id, h.status])); + const healthMap = new Map(healthStatuses.map((h) => [h.server_id, h])); return mcpServers.map((server) => { const healthStatus = healthMap.get(server.server_id); return { ...server, - status: healthStatus ? (healthStatus as "healthy" | "unhealthy" | "unknown") : server.status, + ...healthStatus, + status: healthStatus?.status ?? server.status, }; }); }, [mcpServers, healthStatuses]); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 47df369fb8a..e5dc7280148 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -438,6 +438,7 @@ export interface MCPServer { status?: "healthy" | "unhealthy" | "unknown"; last_health_check?: string | null; health_check_error?: string | null; + health_check_type?: "liveness" | "protocol" | null; teams?: Team[]; mcp_access_groups?: string[]; allowed_tools?: string[]; diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 5d0bb56936a..73143535732 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -32250,6 +32250,8 @@ export interface components { has_user_credential?: boolean | null; /** Health Check Error */ health_check_error?: string | null; + /** Health Check Type */ + readonly health_check_type?: ("liveness" | "protocol") | null; /** Instructions */ instructions?: string | null; /** @@ -35153,6 +35155,19 @@ export interface components { [key: string]: unknown; }; }; + /** MCPServerHealthResponse */ + MCPServerHealthResponse: { + /** Health Check Error */ + health_check_error: string | null; + /** Health Check Type */ + health_check_type?: ("liveness" | "protocol") | null; + /** Last Health Check */ + last_health_check: string | null; + /** Server Id */ + server_id: string; + /** Status */ + status: ("healthy" | "unhealthy" | "unknown") | null; + }; /** * MCPServerUserCredentialListItem * @description One user's stored credential for an MCP server, as an admin sees it. Never carries the secret. @@ -71895,7 +71910,7 @@ export interface operations { [name: string]: unknown; }; content: { - "application/json": unknown; + "application/json": components["schemas"]["MCPServerHealthResponse"][]; }; }; /** @description Validation Error */