""" Smoke test for the EPANET MCP Space. Runs the full MCP handshake (initialize -> tools/list -> tools/call) against a running server over BOTH transports, and verifies the API-key gate. This is the transport-level proof that every target client can connect: Streamable HTTP -> HuggingChat, Codex, Perplexity, Gemini, ChatGPT (dev) SSE -> Claude Desktop (via mcp-remote) and other legacy clients Usage: python smoke_test.py --base-url http://127.0.0.1:7860 [--api-key KEY] """ from __future__ import annotations import argparse import asyncio import sys import httpx from mcp import ClientSession from mcp.client.sse import sse_client from mcp.client.streamable_http import streamablehttp_client PASS, FAIL = "✅", "❌" results: list[tuple[str, bool, str]] = [] def record(name: str, ok: bool, detail: str = "") -> None: results.append((name, ok, detail)) print(f" {PASS if ok else FAIL} {name}" + (f" — {detail}" if detail else "")) def _headers(api_key: str | None) -> dict: return {"Authorization": f"Bearer {api_key}"} if api_key else {} async def exercise(session: ClientSession, transport: str) -> None: await session.initialize() record(f"[{transport}] initialize", True) tools = (await session.list_tools()).tools names = {t.name for t in tools} record(f"[{transport}] tools/list", len(tools) > 0, f"{len(tools)} tools") expected = {"load_network", "get_network_summary", "run_hydraulic_simulation", "list_bundled_networks", "create_leakage_event"} missing = expected - names record(f"[{transport}] core tools present", not missing, "all present" if not missing else f"missing {missing}") # Functional round-trip: list bundled -> load Net1 -> summarise r = await session.call_tool("list_bundled_networks", {}) record(f"[{transport}] call list_bundled_networks", not r.isError) r = await session.call_tool("load_network", {"path": "Net1.inp", "network_id": f"smoke_{transport}"}) record(f"[{transport}] call load_network(Net1)", not r.isError) r = await session.call_tool("get_network_summary", {"network_id": f"smoke_{transport}"}) text = (r.content[0].text if r.content else "") ok = (not r.isError) and ("junctions" in text or "nodes" in text) record(f"[{transport}] call get_network_summary", ok, "summary returned" if ok else "unexpected payload") async def test_streamable(base_url: str, api_key: str | None) -> None: print("\n▶ Streamable HTTP (/mcp)") url = f"{base_url}/mcp" try: async with streamablehttp_client(url, headers=_headers(api_key)) as (r, w, _): async with ClientSession(r, w) as session: await exercise(session, "http") except Exception as e: # noqa: BLE001 record("[http] connection", False, repr(e)) async def test_sse(base_url: str, api_key: str | None) -> None: print("\n▶ SSE (/sse)") url = f"{base_url}/sse" try: async with sse_client(url, headers=_headers(api_key)) as (r, w): async with ClientSession(r, w) as session: await exercise(session, "sse") except Exception as e: # noqa: BLE001 record("[sse] connection", False, repr(e)) async def test_health_and_auth(base_url: str, api_key: str | None) -> None: print("\n▶ Health + auth gate") async with httpx.AsyncClient(timeout=15) as c: try: resp = await c.get(f"{base_url}/health") record("health returns 200/ok", resp.status_code == 200 and resp.text.strip() == "ok") except Exception as e: # noqa: BLE001 record("health probe", False, repr(e)) if api_key: # No key -> must be 401 on a protected endpoint try: resp = await c.post(f"{base_url}/mcp", json={"jsonrpc": "2.0", "id": 1, "method": "ping"}, headers={"Accept": "application/json, text/event-stream"}) record("unauthenticated /mcp rejected (401)", resp.status_code == 401, f"status {resp.status_code}") except Exception as e: # noqa: BLE001 record("unauthenticated /mcp rejected", False, repr(e)) else: record("auth gate", True, "open mode (no key configured) — skipped") async def main() -> int: ap = argparse.ArgumentParser() ap.add_argument("--base-url", default="http://127.0.0.1:7860") ap.add_argument("--api-key", default=None) args = ap.parse_args() base = args.base_url.rstrip("/") print(f"Target: {base} auth={'on' if args.api_key else 'off'}") await test_health_and_auth(base, args.api_key) await test_streamable(base, args.api_key) await test_sse(base, args.api_key) passed = sum(1 for _, ok, _ in results if ok) total = len(results) print(f"\n{'='*48}\nRESULT: {passed}/{total} checks passed") failed = [n for n, ok, _ in results if not ok] if failed: print("Failed:", ", ".join(failed)) return 0 if passed == total else 1 if __name__ == "__main__": sys.exit(asyncio.run(main()))