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 6bf3b22ddb5..aa6d6dd5de6 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 @@ -14605,16 +14605,24 @@ async def test_liveness_accepts_headers_without_body_auth_or_redirects( 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", + 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) + 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") @@ -14632,9 +14640,13 @@ async def test_liveness_accepts_headers_without_body_auth_or_redirects( @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: +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") + 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) @@ -14650,6 +14662,7 @@ async def test_liveness_whole_operation_timeout_and_cancellation( ) -> None: started: Final = asyncio.Event() finished: Final = asyncio.Event() + async def pending(request: httpx.Request) -> httpx.Response: started.set() try: @@ -14660,7 +14673,9 @@ async def test_liveness_whole_operation_timeout_and_cancellation( 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") + 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)) @@ -14670,15 +14685,23 @@ async def test_liveness_whole_operation_timeout_and_cancellation( 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_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: +@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") + 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" @@ -14703,8 +14726,12 @@ async def test_liveness_tls_and_infinite_sse_close_after_headers( 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)) + 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()) @@ -14712,13 +14739,18 @@ async def test_liveness_tls_and_infinite_sse_close_after_headers( 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())) + 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") + 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() @@ -14738,7 +14770,9 @@ async def test_liveness_tls_and_infinite_sse_close_after_headers( 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") + 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": 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 eb5ea57c6b3..3f708ede95d 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 @@ -4293,7 +4293,9 @@ 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 = "healthy" if probe_kind == "liveness" else {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)