Files
dmf-studio-rnd/tests/test_ws_api.py
T
JohannesItten 59b2ee437c test: WebSocket API test suite (19 cases)
Covers connection, protocol robustness, get_status, load_graph,
stop/start_node, crash detection, and multi-client push.
Runs against a live studio-manager; --skip-nodes skips tests
that need node binaries and MXL.

Run: ./tests/.venv/bin/python3 tests/test_ws_api.py

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-07-07 22:26:28 +03:00

466 lines
18 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
"""
WebSocket API test suite for dmf-studio-manager.
Usage:
./tests/.venv/bin/python3 tests/test_ws_api.py [options]
Options:
--url URL WebSocket server URL (default: ws://localhost:7070/ws)
--skip-nodes Skip tests that require node binaries and MXL
Prerequisites:
- studio-manager running (no graph loaded, or will be reset)
- Node binaries in the same dir as studio-manager (for node tests)
"""
import argparse
import asyncio
import json
import os
import signal
import sys
import websockets
# ── Minimal graph: testpattern → fakesink (no hardware needed) ────────────────
SIMPLE_GRAPH = {
"nodes": [
{"id": "src", "type": "testpattern", "params": {}},
{"id": "sink", "type": "fakesink", "params": {}},
],
"edges": [
{
"from": "src", "from_port": "video_flow_id",
"to": "sink", "to_port": "video_flow_id",
"format": {"kind": "video", "width": 1920, "height": 1080,
"fps_num": 25, "fps_den": 1},
}
],
}
EMPTY_GRAPH = {"nodes": [], "edges": []}
# ── Suite / reporter ──────────────────────────────────────────────────────────
class Suite:
def __init__(self):
self.passed = self.failed = self.skipped = 0
def ok(self, name):
print(f" \033[32m✓\033[0m {name}")
self.passed += 1
def fail(self, name, reason=""):
msg = f": {reason}" if reason else ""
print(f" \033[31m✗\033[0m {name}{msg}")
self.failed += 1
def skip(self, name, reason=""):
msg = f" ({reason})" if reason else ""
print(f" \033[33m\033[0m {name}{msg}")
self.skipped += 1
def section(self, title):
print(f"\n{title}")
print("" * len(title))
def summary(self):
total = self.passed + self.failed + self.skipped
parts = [f"{self.passed}/{total} passed"]
if self.skipped:
parts.append(f"{self.skipped} skipped")
if self.failed:
parts.append(f"\033[31m{self.failed} failed\033[0m")
print("\n" + " ".join(parts))
return self.failed == 0
# ── Low-level helpers ─────────────────────────────────────────────────────────
async def recv(ws, timeout=3.0):
raw = await asyncio.wait_for(ws.recv(), timeout=timeout)
return json.loads(raw)
async def send_recv(ws, payload, timeout=3.0):
await ws.send(json.dumps(payload))
return await recv(ws, timeout)
async def open_ws(url):
return await websockets.connect(url)
async def reset(ws):
"""Load empty graph to clear all nodes between tests."""
r = await send_recv(ws, {"type": "load_graph", "graph": EMPTY_GRAPH})
return r
def is_status(msg):
return isinstance(msg, dict) and msg.get("type") == "status"
def is_error(msg):
return isinstance(msg, dict) and msg.get("type") == "error"
def nodes_by_id(status):
return {n["id"]: n for n in status.get("nodes", [])}
# ── Test groups ───────────────────────────────────────────────────────────────
async def run_connection_tests(url, s):
s.section("Connection")
# 1. Connect → immediate status push
ws = await open_ws(url)
try:
msg = await recv(ws, timeout=3.0)
if is_status(msg) and "nodes" in msg:
s.ok("connect → immediate status received")
else:
s.fail("connect → immediate status received", f"got {msg}")
except asyncio.TimeoutError:
s.fail("connect → immediate status received", "timeout")
finally:
await ws.close()
# 2. Two clients connect → both receive status independently
ws1 = await open_ws(url)
ws2 = await open_ws(url)
try:
m1 = await recv(ws1, timeout=3.0)
m2 = await recv(ws2, timeout=3.0)
if is_status(m1) and is_status(m2):
s.ok("two clients connect → both receive status")
else:
s.fail("two clients connect → both receive status",
f"ws1={m1.get('type')} ws2={m2.get('type')}")
except asyncio.TimeoutError:
s.fail("two clients connect → both receive status", "timeout")
finally:
await ws1.close()
await ws2.close()
async def run_protocol_tests(url, s):
s.section("Protocol robustness")
ws = await open_ws(url)
await recv(ws) # discard initial status
# 3. Invalid JSON → error, connection stays open
await ws.send("not json {{")
try:
msg = await recv(ws, timeout=3.0)
if is_error(msg):
s.ok("invalid JSON → error response")
else:
s.fail("invalid JSON → error response", f"got type={msg.get('type')}")
except asyncio.TimeoutError:
s.fail("invalid JSON → error response", "timeout")
# 4. Unknown command type → error, connection stays open
try:
msg = await send_recv(ws, {"type": "explode"})
if is_error(msg):
s.ok("unknown command → error response")
else:
s.fail("unknown command → error response", f"got type={msg.get('type')}")
except asyncio.TimeoutError:
s.fail("unknown command → error response", "timeout")
# 5. Connection still alive after errors
try:
msg = await send_recv(ws, {"type": "get_status"})
if is_status(msg):
s.ok("connection alive after error responses")
else:
s.fail("connection alive after error responses", f"got {msg}")
except Exception as e:
s.fail("connection alive after error responses", str(e))
# 6. load_graph with missing 'graph' key → error
try:
msg = await send_recv(ws, {"type": "load_graph"})
if is_error(msg):
s.ok("load_graph missing 'graph' key → error")
else:
s.fail("load_graph missing 'graph' key → error", f"got {msg}")
except asyncio.TimeoutError:
s.fail("load_graph missing 'graph' key → error", "timeout")
# 7. stop_node with missing 'id' key → error
try:
msg = await send_recv(ws, {"type": "stop_node"})
if is_error(msg):
s.ok("stop_node missing 'id' key → error")
else:
s.fail("stop_node missing 'id' key → error", f"got {msg}")
except asyncio.TimeoutError:
s.fail("stop_node missing 'id' key → error", "timeout")
await ws.close()
async def run_status_tests(url, s):
s.section("get_status")
ws = await open_ws(url)
push = await recv(ws)
# 8. get_status response matches last push
try:
reply = await send_recv(ws, {"type": "get_status"})
if reply.get("nodes") == push.get("nodes"):
s.ok("get_status matches initial push")
else:
s.fail("get_status matches initial push",
f"push={push['nodes']} reply={reply['nodes']}")
except asyncio.TimeoutError:
s.fail("get_status matches initial push", "timeout")
await ws.close()
async def run_node_tests(url, s):
s.section("Graph and node lifecycle (requires node binaries + MXL)")
ws = await open_ws(url)
await recv(ws) # discard initial status
# 9. load_graph with valid simple graph → nodes Running
try:
r = await send_recv(ws, {"type": "load_graph", "graph": SIMPLE_GRAPH}, timeout=10.0)
nodes = nodes_by_id(r)
if (is_status(r)
and "src" in nodes and nodes["src"]["state"] == "running"
and "sink" in nodes and nodes["sink"]["state"] == "running"):
s.ok("load valid graph → nodes running")
else:
s.fail("load valid graph → nodes running",
f"states={[(n['id'], n['state']) for n in r.get('nodes', [])]}")
except asyncio.TimeoutError:
s.fail("load valid graph → nodes running", "timeout")
# 10. get_status consistent with load_graph response
try:
r = await send_recv(ws, {"type": "get_status"})
nodes = nodes_by_id(r)
if "src" in nodes and "sink" in nodes:
s.ok("get_status consistent after load")
else:
s.fail("get_status consistent after load", f"nodes={list(nodes.keys())}")
except asyncio.TimeoutError:
s.fail("get_status consistent after load", "timeout")
# 11. load_graph again → previous nodes stopped, new ones started
try:
r1 = await send_recv(ws, {"type": "get_status"})
old_pids = {n["id"]: n["pid"] for n in r1.get("nodes", [])}
r2 = await send_recv(ws, {"type": "load_graph", "graph": SIMPLE_GRAPH}, timeout=10.0)
new_nodes = nodes_by_id(r2)
new_pids = {nid: n["pid"] for nid, n in new_nodes.items()}
if (all(n["state"] == "running" for n in new_nodes.values())
and any(new_pids.get(k) != old_pids.get(k) for k in old_pids)):
s.ok("reload graph → new PIDs, all running")
else:
s.fail("reload graph → new PIDs, all running",
f"old={old_pids} new={new_pids}")
except asyncio.TimeoutError:
s.fail("reload graph → new PIDs, all running", "timeout")
# 12. stop_node → state becomes stopped
try:
r = await send_recv(ws, {"type": "stop_node", "id": "src"})
nodes = nodes_by_id(r)
if nodes.get("src", {}).get("state") == "stopped":
s.ok("stop_node → state=stopped")
else:
s.fail("stop_node → state=stopped",
f"state={nodes.get('src', {}).get('state')}")
except asyncio.TimeoutError:
s.fail("stop_node → state=stopped", "timeout")
# 13. stop already-stopped node → no crash, state still stopped
try:
r = await send_recv(ws, {"type": "stop_node", "id": "src"})
nodes = nodes_by_id(r)
if is_status(r) and nodes.get("src", {}).get("state") == "stopped":
s.ok("stop already-stopped node → graceful no-op")
else:
s.fail("stop already-stopped node → graceful no-op", f"got {r}")
except asyncio.TimeoutError:
s.fail("stop already-stopped node → graceful no-op", "timeout")
# 14. stop nonexistent node → error or status (not a crash)
try:
r = await send_recv(ws, {"type": "stop_node", "id": "does-not-exist"})
if is_status(r) or is_error(r):
s.ok("stop nonexistent node → no crash")
else:
s.fail("stop nonexistent node → no crash", f"got {r}")
except asyncio.TimeoutError:
s.fail("stop nonexistent node → no crash", "timeout")
# 15. start_node → state becomes running
try:
r = await send_recv(ws, {"type": "start_node", "id": "src"}, timeout=5.0)
nodes = nodes_by_id(r)
if nodes.get("src", {}).get("state") == "running":
s.ok("start_node → state=running")
else:
s.fail("start_node → state=running",
f"state={nodes.get('src', {}).get('state')}")
except asyncio.TimeoutError:
s.fail("start_node → state=running", "timeout")
# 16. start already-running node → graceful no-op
try:
r = await send_recv(ws, {"type": "start_node", "id": "src"})
nodes = nodes_by_id(r)
if is_status(r) and nodes.get("src", {}).get("state") == "running":
s.ok("start already-running node → graceful no-op")
else:
s.fail("start already-running node → graceful no-op", f"got {r}")
except asyncio.TimeoutError:
s.fail("start already-running node → graceful no-op", "timeout")
# 17. Node crash detection: kill a node externally → pushed as crashed
try:
status = await send_recv(ws, {"type": "get_status"})
src_pid = nodes_by_id(status).get("src", {}).get("pid", -1)
if src_pid > 0:
os.kill(src_pid, signal.SIGKILL)
# Monitor polls every 500ms; wait up to 2s for crash push
deadline = asyncio.get_event_loop().time() + 2.0
crashed = False
while asyncio.get_event_loop().time() < deadline:
try:
push = await recv(ws, timeout=1.5)
if is_status(push):
node = nodes_by_id(push).get("src", {})
if node.get("state") == "crashed":
crashed = True
break
except asyncio.TimeoutError:
break
if crashed:
s.ok("crashed node detected and pushed within 2s")
else:
s.fail("crashed node detected and pushed within 2s", "no push received")
else:
s.fail("crashed node detected and pushed within 2s", "couldn't get src pid")
except Exception as e:
s.fail("crashed node detected and pushed within 2s", str(e))
# 18. Push sent to all connected clients simultaneously
ws2 = await open_ws(url)
await recv(ws2) # discard initial status
try:
r = await send_recv(ws, {"type": "load_graph", "graph": SIMPLE_GRAPH}, timeout=10.0)
# ws gets the response; ws2 should also get a push
if is_status(r):
try:
push2 = await recv(ws2, timeout=3.0)
if is_status(push2) and len(push2.get("nodes", [])) > 0:
s.ok("load_graph push reaches all connected clients")
else:
s.fail("load_graph push reaches all connected clients",
f"ws2 got {push2}")
except asyncio.TimeoutError:
s.fail("load_graph push reaches all connected clients",
"ws2 got no push")
else:
s.fail("load_graph push reaches all connected clients",
f"load response was {r}")
except asyncio.TimeoutError:
s.fail("load_graph push reaches all connected clients", "timeout on load")
finally:
await ws2.close()
# 19. Unknown node type → node quickly transitions to crashed
try:
r = await send_recv(ws,
{"type": "load_graph", "graph": {
"nodes": [{"id": "bad", "type": "nonexistent-node-type", "params": {}}],
"edges": []
}}, timeout=5.0)
# The load itself may succeed (fork happens), but the node should crash fast
if is_status(r):
nodes = nodes_by_id(r)
initial_state = nodes.get("bad", {}).get("state")
if initial_state == "crashed":
s.ok("unknown node type → immediately crashed")
else:
# Wait for crash push (exec fails → child exits → monitor detects)
try:
push = await recv(ws, timeout=2.0)
if is_status(push) and nodes_by_id(push).get("bad", {}).get("state") == "crashed":
s.ok("unknown node type → crashes quickly")
else:
s.fail("unknown node type → crashes quickly",
f"initial={initial_state}, push={push}")
except asyncio.TimeoutError:
s.fail("unknown node type → crashes quickly",
f"initial state={initial_state}, no crash push")
else:
s.fail("unknown node type → crashes quickly", f"load returned {r}")
except asyncio.TimeoutError:
s.fail("unknown node type → crashes quickly", "timeout")
# Cleanup
await reset(ws)
await ws.close()
# ── Entry point ───────────────────────────────────────────────────────────────
async def main(url, skip_nodes):
print(f"studio-manager WebSocket API tests")
print(f"server: {url}")
s = Suite()
try:
await run_connection_tests(url, s)
await run_protocol_tests(url, s)
await run_status_tests(url, s)
if skip_nodes:
s.section("Graph and node lifecycle (requires node binaries + MXL)")
for name in [
"load valid graph → nodes running",
"get_status consistent after load",
"reload graph → new PIDs, all running",
"stop_node → state=stopped",
"stop already-stopped node → graceful no-op",
"stop nonexistent node → no crash",
"start_node → state=running",
"start already-running node → graceful no-op",
"crashed node detected and pushed within 2s",
"load_graph push reaches all connected clients",
"unknown node type → crashes quickly",
]:
s.skip(name, "--skip-nodes")
else:
await run_node_tests(url, s)
except OSError as e:
print(f"\n\033[31mCannot connect to {url}: {e}\033[0m")
print("Is studio-manager running?")
sys.exit(2)
s.summary()
sys.exit(0 if s.failed == 0 else 1)
if __name__ == "__main__":
parser = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
parser.add_argument("--url", default="ws://localhost:7070/ws",
help="WebSocket server URL")
parser.add_argument("--skip-nodes", action="store_true",
help="Skip tests requiring node binaries and MXL")
args = parser.parse_args()
asyncio.run(main(args.url, args.skip_nodes))