Files
app/backend/tests/flow/test_remote.py
T
stroblmeandClaude Opus 5 4355c917f8 Node settings arrive as keyword arguments, not a params dict
A python node's settings are constants of its own function, so they are passed
the way its ports are: by name. The controller binds them to the compiled
function, the `params` field is gone from the worker and remote protocols, and
a setting sharing a port's name is reported as a node error rather than
shadowing it. The panel's scaffold follows suit and keeps the header in step
with both ports and settings.

The demo's `pace` moves from a flow input to a setting of the training node,
which is what it always was.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01NUb8YpL2s3gmN9WTACTt4q
2026-08-20 17:47:45 +02:00

223 lines
6.7 KiB
Python

"""A worker on another host, and the thread-to-loop bridge that reaches it.
The socket lives on an event loop; a node executes on a worker thread. These
run a real loop in a thread of its own, because that split is the whole
difficulty: what is checked is that a call handed across it comes back — with
its answer, with the reports it made on the way, or with a failure that says
what happened rather than a wait that never ends.
"""
import asyncio
import threading
from collections.abc import Iterator
import pytest
from app.flow.remote import NoWorker, RemoteWorker, RemoteWorkerHub
from app.flow.workers import NodeTimeout, RemoteError
@pytest.fixture
def loop() -> Iterator[asyncio.AbstractEventLoop]:
"""An event loop running in a thread, as the server's does."""
running = asyncio.new_event_loop()
thread = threading.Thread(target=running.run_forever, daemon=True)
thread.start()
yield running
running.call_soon_threadsafe(running.stop)
thread.join(timeout=5)
running.close()
class FakeSocket:
"""Stands in for the websocket: records frames, and says when one lands."""
def __init__(self) -> None:
self.sent: list[dict] = []
self.arrived = threading.Event()
async def send_json(self, payload: dict) -> None:
self.sent.append(payload)
self.arrived.set()
def attach(hub: RemoteWorkerHub, loop: asyncio.AbstractEventLoop, name: str = "gpu1"):
socket = FakeSocket()
worker = RemoteWorker(
name=name, labels=["gpu"], send=socket.send_json, loop=loop, max_parallel=2
)
hub.attach(worker)
return worker, socket
def call_in_thread(target) -> threading.Thread:
thread = threading.Thread(target=target, daemon=True)
thread.start()
return thread
def test_a_call_crosses_to_the_thread_and_the_answer_comes_back(loop):
hub = RemoteWorkerHub()
worker, socket = attach(hub, loop)
result: dict = {}
thread = call_in_thread(
lambda: result.update(
value=hub.run(
"gpu", "flow", "node", "src", {"x": 1}, "flow.node", timeout=5
)
)
)
assert socket.arrived.wait(5)
assert socket.sent[0]["source"] == "src"
assert socket.sent[0]["kwargs"] == {"x": 1}
worker.deliver(
{"call_id": socket.sent[0]["call_id"], "ok": True, "result": {"out": 2}}
)
thread.join(timeout=5)
assert result["value"] == {"out": 2}
def test_reports_arrive_before_the_answer_and_a_heartbeat_is_not_one(loop):
hub = RemoteWorkerHub()
worker, socket = attach(hub, loop)
seen: list[dict] = []
result: dict = {}
thread = call_in_thread(
lambda: result.update(
value=hub.run(
"gpu",
"flow",
"node",
"src",
{},
"flow.node",
timeout=5,
run_id="r1",
on_event=seen.append,
)
)
)
assert socket.arrived.wait(5)
call_id = socket.sent[0]["call_id"]
# The call names its run, which is how a metric finds the run that made it.
assert call_id == "r1:flow.node"
worker.deliver(
{"call_id": call_id, "event": "metric", "name": "loss", "value": 1.0}
)
worker.deliver({"call_id": call_id, "event": "heartbeat"})
worker.deliver({"call_id": call_id, "ok": True, "result": {"done": True}})
thread.join(timeout=5)
assert result["value"] == {"done": True}
# Liveness is not a measurement; only the metric is kept.
assert [event["event"] for event in seen] == ["metric"]
def test_a_failure_keeps_its_class_across_the_socket(loop):
hub = RemoteWorkerHub()
worker, socket = attach(hub, loop)
caught: list[Exception] = []
def call() -> None:
try:
hub.run("gpu", "flow", "node", "src", {}, "flow.node", timeout=5)
except Exception as exc:
caught.append(exc)
thread = call_in_thread(call)
assert socket.arrived.wait(5)
worker.deliver(
{
"call_id": socket.sent[0]["call_id"],
"ok": False,
"error": {"type": "ValueError", "message": "bad input", "traceback": "tb"},
}
)
thread.join(timeout=5)
assert type(caught[0]).__name__ == "ValueError"
assert str(caught[0]) == "bad input"
def test_a_worker_that_goes_away_fails_the_call_rather_than_hanging(loop):
hub = RemoteWorkerHub()
worker, socket = attach(hub, loop)
caught: list[Exception] = []
def call() -> None:
try:
hub.run("gpu", "flow", "node", "src", {}, "flow.node", timeout=30)
except Exception as exc:
caught.append(exc)
thread = call_in_thread(call)
assert socket.arrived.wait(5)
# Pulling the cable mid-training: the node fails, and does not wait out its
# thirty-second deadline to do it.
worker.detach()
thread.join(timeout=5)
assert not thread.is_alive()
assert isinstance(caught[0], RemoteError)
assert "went away" in str(caught[0])
def test_silence_past_the_deadline_is_a_timeout(loop):
hub = RemoteWorkerHub()
attach(hub, loop)
caught: list[Exception] = []
def call() -> None:
try:
hub.run("gpu", "flow", "node", "src", {}, "flow.node", timeout=0.3)
except Exception as exc:
caught.append(exc)
call_in_thread(call).join(timeout=5)
assert isinstance(caught[0], NodeTimeout)
def test_a_label_nothing_carries_is_named_rather_than_waited_on(loop):
hub = RemoteWorkerHub()
attach(hub, loop)
with pytest.raises(NoWorker, match="tpu"):
hub.run("tpu", "flow", "node", "src", {}, "flow.node", timeout=5)
# Compiling against a machine that is not attached is not a broken node —
# a node importing torch is correct there and missing here.
assert hub.compile("tpu", "flow", "node", "src") is None
def test_reattaching_replaces_the_old_socket(loop):
hub = RemoteWorkerHub()
first, _ = attach(hub, loop)
second, _ = attach(hub, loop)
assert first.gone
assert hub.pick("gpu") is second
assert [worker.name for worker in hub.workers()] == ["gpu1"]
def test_cancelling_a_run_reaches_only_that_run(loop):
hub = RemoteWorkerHub()
worker, socket = attach(hub, loop)
def call(run_id: str) -> None:
try:
hub.run("gpu", "flow", "node", "src", {}, "flow.node", 30, run_id=run_id)
except Exception:
pass
call_in_thread(lambda: call("run-a"))
assert socket.arrived.wait(5)
socket.arrived.clear()
call_in_thread(lambda: call("run-b"))
assert socket.arrived.wait(5)
assert hub.cancel_run("run-a") == 1
cancels = [frame for frame in socket.sent if frame.get("op") == "cancel"]
assert [frame["call_id"] for frame in cancels] == ["run-a:flow.node"]