diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 06fdc4d..ada4421 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -28,8 +28,10 @@ jobs: - uses: astral-sh/setup-uv@v5 with: enable-cache: true + # With the checker's extra: tests/test_conformance_probe.py runs it, and + # skips without websockets. - name: Sync - run: uv sync --frozen + run: uv sync --frozen --extra conformance - name: Test run: uv run pytest -q @@ -70,6 +72,15 @@ jobs: --port 8765 --steps 10 --horizon 4 --action-dim 7 --timeout 120 \ --client "/tmp/plugrl_client 127.0.0.1 8765 10" + # SPEC section 8.1 for the C++ client too: what it does with each chunk, + # and how it handles a resync, a stop and a text frame. + - name: ... and in probe mode, in every scenario + run: | + uv run --extra conformance plugrl-conformance \ + --probe --scenario all --port 8769 --steps 12 --horizon 4 \ + --action-dim 3 --action-dtype float64 --timeout 120 \ + --client "/tmp/plugrl_client --probe 127.0.0.1 8769 3" + # The action dtype is the environment's and is never renegotiated, so a # client has to read the typestr. A client that hard-codes float32 still # passes every clause the server can see - it sends valid messages - and @@ -94,3 +105,13 @@ jobs: uv run --extra conformance plugrl-conformance \ --port 8767 --steps 10 --horizon 3 --action-dim 5 --timeout 120 \ --client "python examples/raw_client.py --host 127.0.0.1 --port 8767 --steps 10" + + # Section 8.1: the probe env lets the checker see what the client did + # with each chunk, and the scenarios drive a resync, a stop and a text + # frame. The client is started once per scenario. + - name: ... and in probe mode, in every scenario + run: | + uv run --extra conformance plugrl-conformance \ + --probe --scenario all --port 8768 --steps 12 --horizon 4 \ + --action-dim 3 --action-dtype float64 --timeout 120 \ + --client "python examples/raw_client.py --probe --batch 3 --host 127.0.0.1 --port 8768" diff --git a/README.md b/README.md index 580e022..9d6415d 100644 --- a/README.md +++ b/README.md @@ -40,6 +40,37 @@ uv run --extra conformance plugrl-conformance \ It exits non-zero on a violation, so it can sit in a CI job. What it does not require, because SPEC.md does not, is listed at the top of the module. +That form watches one connection, so it sees the messages and not what the +client did with them. To check that too, have the client run the probe +environment of [SPEC.md section 8.1](SPEC.md#81-the-probe-environment), a +few lines in any language, and add `--probe`: + +```bash +uv run --extra conformance plugrl-conformance --probe --scenario all \ + --client "./my_client --probe 127.0.0.1 8000" +``` + +The checker then works out from each feedback whether the client summed the +chunk's reward, sent the terminal observation, and applied the actions +time-major and in order. It also drives the connection: it speaks late, sends +a metadata frame larger than 1 MiB, closes for a resync, stops the run, and +answers with a text frame, starting the client once per scenario. + +## Checking a server + +The other direction: a client that drives a training server the way SPEC.md +lets a client behave, and checks what comes back. That covers the action +layout and `env_ids`, ragged batches, large frames, a resync close on each +kind of malformed message with the server staying up afterwards, and the +stop at the end of the run ([SPEC.md section 8.2](SPEC.md#82-checking-a-server)): + +```bash +uv run --extra conformance plugrl-conformance-server --port 8000 --state-dim 3 --until-stop +``` + +[`examples/reference_server.py`](examples/reference_server.py) is a server +written against the specification that trains nothing, and passes. + ## The protocol in one screen Four message types, four lowercase strings. The server speaks first. diff --git a/SPEC.md b/SPEC.md index 489e241..28d8787 100644 --- a/SPEC.md +++ b/SPEC.md @@ -657,38 +657,127 @@ A client conforms to version 1 if it: `feedback` for an `action` that arrived on an earlier connection; - [ ] treats a text frame as a fatal error. -`examples/conformance_server.py` checks the clauses above that one passive -connection can observe - framing, alternation, env indices, observation -shape, and the feedback payload's keys, dtypes and lengths - and reports -what it accepts but cannot require as a note rather than a failure. Both -reference clients pass it with one note: they send `text` as a msgpack -string array rather than a ` #include @@ -35,10 +41,29 @@ #include #include #include +#include +#include #include namespace { +// SPEC section 7: why the server closed. The reason is how a client tells a +// finished run (plugrl-server-stop) from a request to resync. +struct ServerClosed : std::runtime_error { + int code; + std::string reason; + ServerClosed(int c, std::string r) + : std::runtime_error("server closed the connection (" + std::to_string(c) + + (r.empty() ? std::string() : ", " + r) + ")"), + code(c), + reason(std::move(r)) {} +}; + +// SPEC section 7.4: a text frame stands for a server-side error. +struct TextFrame : std::runtime_error { + using std::runtime_error::runtime_error; +}; + // ---------------------------------------------------------------- SHA-1 // RFC 3174. Needed only to verify the server's Sec-WebSocket-Accept. @@ -535,7 +560,13 @@ class WebSocket { if (masked) for (uint64_t i = 0; i < len; ++i) chunk[i] ^= key[i % 4]; - if (opcode == 0x8) throw std::runtime_error("server closed the connection"); + if (opcode == 0x8) { + // RFC 6455 section 5.5.1: a two-byte status code, then the reason. + int code = chunk.size() >= 2 + ? (static_cast(chunk[0]) << 8) | static_cast(chunk[1]) + : 1005; + throw ServerClosed(code, chunk.size() > 2 ? chunk.substr(2) : std::string()); + } if (opcode == 0x9) { send_pong(chunk); continue; } if (opcode == 0xa) continue; @@ -546,8 +577,7 @@ class WebSocket { // msgpack type 0x47". if (first_data_frame) { if (opcode == 0x1) - throw std::runtime_error("server sent a text frame: " + - chunk.substr(0, 2048)); + throw TextFrame("server sent a text frame: " + chunk.substr(0, 2048)); if (opcode != 0x2) throw std::runtime_error("unexpected websocket opcode " + std::to_string(static_cast(opcode))); @@ -824,9 +854,205 @@ int run(const std::string& host, int port, int steps, int batch) { return 0; } +// ---------------------------------------------------------------- probe mode +// +// SPEC section 8.1. Env i counts the steps of its episode in t, remembers the +// action it applied last in a, pays 1 per step, and ends its episode, +// terminated, at t = 3 + 2 * (i % 3). plugrl-conformance --probe works out +// from each feedback what the client did with the chunk it was sent. + +struct ProbeEnv { + int length = 3; + int t = 0; + std::vector a; +}; + +void pack_probe_observation(Packer& p, const std::vector& envs, int64_t width) { + int64_t n = static_cast(envs.size()); + std::vector t, a; + for (const auto& env : envs) { + t.push_back(env.t); + if (env.a.empty()) a.insert(a.end(), static_cast(width), 0.0); + else a.insert(a.end(), env.a.begin(), env.a.end()); + } + p.map(3); + p.str("images"); p.map(0); + p.str("states"); p.map(2); + p.str("t"); pack_ndarray(p, pack_f8(t), "kind == Value::Kind::Int ? v->i : static_cast(v->u); +} + +// One connection, until the server closes it - which throws ServerClosed. +void probe_session(WebSocket& ws, std::vector& envs, int64_t& width, + bool& stepped) { + const int64_t n = static_cast(envs.size()); + + // The server speaks first. Unknown keys - and the probe server sends a + // 2 MiB one - are ignored, as section 5.1 says. + auto meta_raw = ws.recv_message(); + Unpacker mu(reinterpret_cast(meta_raw.data()), meta_raw.size()); + auto meta = mu.parse(); + auto mt = meta->find("message_type"); + if (!mt || mt->s != "metadata") throw std::runtime_error("expected metadata"); + auto md = meta->find("data"); + auto ad = md ? md->find("action_dim") : nullptr; + if (ad && !stepped && + (ad->kind == Value::Kind::Int || ad->kind == Value::Kind::UInt) && as_int(ad) > 0) + width = as_int(ad); + + std::vector env_idx(static_cast(n)); + for (int64_t i = 0; i < n; ++i) env_idx[static_cast(i)] = i; + std::vector step_ids(static_cast(n), 0); + + for (;;) { + Packer p; + p.map(4); + p.str("message_type"); p.str("infer"); + p.str("data"); pack_probe_observation(p, envs, width); + p.str("env_indices"); pack_ndarray(p, pack_i8(env_idx), "(reply.data()), reply.size()); + auto msg = up.parse(); + auto type = msg->find("message_type"); + if (!type || type->s != "action") throw std::runtime_error("expected action"); + auto data = msg->find("data"); + auto action = data ? data->find("action") : nullptr; + auto dtype = action ? action->find("dtype") : nullptr; + auto shape = action ? action->find("shape") : nullptr; + auto blob = action ? action->find("data") : nullptr; + if (!dtype || !shape || !blob || shape->arr.size() < 2) + throw std::runtime_error("malformed action"); + + // Time-major, [H, n, *da], in the dtype the typestr names. + TypeStr at = parse_typestr(dtype->s); + const int64_t horizon = as_int(shape->arr[0]); + const int64_t rows = as_int(shape->arr[1]); + int64_t w = 1; + for (size_t i = 2; i < shape->arr.size(); ++i) w *= as_int(shape->arr[i]); + width = w; + stepped = true; + + // Section 5.3: read env_ids when present, else the request's own order. + std::vector ids = env_idx; + auto eid = data->find("env_ids"); + if (eid && eid->find("data") && eid->find("dtype")) { + TypeStr it = parse_typestr(eid->find("dtype")->s); + const std::string& raw = eid->find("data")->s; + ids.clear(); + for (size_t i = 0; i < raw.size() / static_cast(it.size); ++i) + ids.push_back(static_cast(read_element(raw, i, it))); + } + + std::vector rewards(static_cast(n), 0.0f); + std::string terminated(static_cast(n), '\0'); + for (int64_t row = 0; row < rows && row < static_cast(ids.size()); ++row) { + const int64_t index = ids[static_cast(row)]; + if (index < 0 || index >= n) + throw std::runtime_error("action for an env this client does not run"); + ProbeEnv& env = envs[static_cast(index)]; + for (int64_t k = 0; k < horizon; ++k) { + const size_t base = static_cast((k * rows + row) * w); + env.a.assign(static_cast(w), 0.0); + for (int64_t d = 0; d < w; ++d) + env.a[static_cast(d)] = read_element(blob->s, base + static_cast(d), at); + env.t += 1; + rewards[static_cast(index)] += 1.0f; // the chunk's sum + if (env.t == env.length) { + terminated[static_cast(index)] = '\x01'; + break; // a chunk is flushed at the terminal step + } + } + } + + Packer fb; + fb.map(4); + fb.str("message_type"); fb.str("feedback"); + fb.str("env_indices"); pack_ndarray(fb, pack_i8(env_idx), "(n), '\0'), "|b1", {n}); + fb.str("info"); fb.map(0); + + // Reset before sending: an episode that ended has ended in the + // environment whether or not its feedback gets through, so a close + // arriving now must not carry it onto the next connection. The packed + // feedback above already holds the terminal observation. + for (int64_t i = 0; i < n; ++i) { + ProbeEnv& env = envs[static_cast(i)]; + if (terminated[static_cast(i)]) { + env.t = 0; + env.a.clear(); + step_ids[static_cast(i)] = 0; + } else { + step_ids[static_cast(i)] += 1; + } + } + ws.send_binary(fb.data()); + } +} + +int run_probe(const std::string& host, int port, int batch) { + std::vector envs(static_cast(batch)); + for (int i = 0; i < batch; ++i) envs[static_cast(i)].length = 3 + 2 * (i % 3); + int64_t width = 1; + bool stepped = false; + for (;;) { + try { + WebSocket ws; + std::cout << "connecting to ws://" << host << ":" << port << " (probe mode)\n"; + ws.connect(host, port); + probe_session(ws, envs, width, stepped); + } catch (const ServerClosed& closed) { + if (closed.reason == "plugrl-server-stop") { + std::cout << "the server stopped the run\n"; + return 0; + } + if (closed.reason == "plugrl-server-resync") { + // Section 7.6: whatever was in flight went with the old connection; + // the new one starts with an infer, and nothing is resent. + std::cout << "the server asked for a resync; reconnecting\n"; + std::this_thread::sleep_for(std::chrono::milliseconds(200)); + continue; + } + std::cerr << "error: " << closed.what() << "\n"; + return 1; + } catch (const TextFrame& frame) { + std::cerr << "error: " << frame.what() << "\n"; + return 2; + } + } +} + } // namespace int main(int argc, char** argv) { + if (argc > 1 && std::string(argv[1]) == "--probe") { + // --probe host port batch + std::string host = argc > 2 ? argv[2] : "127.0.0.1"; + int port = argc > 3 ? std::stoi(argv[3]) : 8000; + int batch = argc > 4 ? std::stoi(argv[4]) : 1; + try { + return run_probe(host, port, batch); + } catch (const std::exception& e) { + std::cerr << "error: " << e.what() << "\n"; + return 1; + } + } + // host port steps batch img_size cameras std::string host = argc > 1 ? argv[1] : "127.0.0.1"; int port = argc > 2 ? std::stoi(argv[2]) : 8000; // the server default diff --git a/examples/raw_client.py b/examples/raw_client.py index a396954..707157a 100644 --- a/examples/raw_client.py +++ b/examples/raw_client.py @@ -18,15 +18,28 @@ Usage: python raw_client.py [--host 127.0.0.1] [--port 8000] [--steps 20] + +With `--probe` it runs the probe environment of SPEC.md section 8.1 instead +of a fake sensor, executes each action chunk the way section 5.3 says, and +handles the server's close reasons: it reconnects after a resync, dropping the +feedback it was holding, exits 0 when the server stops the run, and treats a +text frame as fatal. That is what `plugrl-conformance --probe` checks: + + python raw_client.py --probe --batch 3 [--host 127.0.0.1] [--port 8000] + +`--bug NAME` makes it break one rule on purpose, so the checker's tests can +show that each rule is actually checked. """ from __future__ import annotations import argparse import sys +import time from array import array import msgpack +import websockets.exceptions import websockets.sync.client # ---------------------------------------------------------------- wire format @@ -261,13 +274,262 @@ def run(host: str, port: int, steps: int, batch: int) -> int: return 0 +# ---------------------------------------------------------------- probe mode +# +# SPEC.md section 8.1. Env i counts the steps of its episode in t, remembers +# the action it applied last in a, pays a reward of 1 per step, and ends its +# episode, terminated, at t = 3 + 2 * (i % 3). + +STOP_REASON = "plugrl-server-stop" +RESYNC_REASON = "plugrl-server-resync" + +BUGS = ( + "early-infer", # sends an infer before reading the metadata + "compress", # offers permessage-deflate in the handshake + "frame-cap", # keeps the websockets library's 1 MiB frame cap + "float32", # assumes the action is float32 whatever its typestr says + "env-major", # reads the action chunk as [n, H, da] + "last-step-reward", # reports the last step's reward, not the chunk's sum + "reset-in-step", # resets inside step, so a done step sends the next episode's obs + "resend-feedback", # resends the feedback it was holding after a resync + "ignore-stop", # reconnects after plugrl-server-stop + "ignore-text", # carries on after a text frame instead of stopping + "stale-chunk", # after a resync, feeds back chunks from the old connection +) + + +class TextFrame(Exception): + """The server sent text where binary was expected: section 7.4, fatal.""" + + +class ProbeEnv: + def __init__(self, index: int) -> None: + self.index = index + self.length = 3 + 2 * (index % 3) + self.t = 0 + self.a: list[float] = [] + + def step(self, action: list[float]) -> tuple[float, bool]: + self.t += 1 + self.a = list(action) + return 1.0, self.t == self.length + + def reset(self) -> None: + self.t = 0 + self.a = [] + + +def probe_observation(envs: list[ProbeEnv], width: int) -> dict: + n = len(envs) + a: list[float] = [] + for env in envs: + a.extend(env.a if env.a else [0.0] * width) + return { + "images": {}, + "states": { + "t": encode_array([float(env.t) for env in envs], " None: + """One connection, until the server closes it.""" + packer = msgpack.Packer() + n = len(envs) + env_indices = encode_array(range(n), " bytes: + asking = list(range(n)) if asking is None else asking + return packer.pack( + { + "message_type": INFER, + "data": probe_observation([envs[i] for i in asking], state["width"]), + "env_indices": encode_array(asking, " 0 and not state["stepped"]: + state["width"] = width + + # A correct client dropped this when the old connection closed. + if state["held"] is not None: + ws.send(packer.pack(state["held"])) + state["held"] = None + + first_on_connection = True + while True: + # The stale-chunk bug: back on a new connection, ask only for env 0, + # as if the other envs were still executing their old chunks. + asking = list(range(n)) + if bug == "stale-chunk" and state["reconnects"] and first_on_connection: + asking = [0] + first_on_connection = False + ws.send(infer(asking)) + try: + reply = _recv(ws) + except TextFrame: + if bug != "ignore-text": + raise + continue # and ask again, as if nothing had happened + if reply.get("message_type") != ACTION: + raise RuntimeError(f"expected {ACTION}, got {reply.get('message_type')!r}") + field = reply["data"]["action"] + if bug == "float32": + field = dict(field) + field[b"dtype"] = " int: + uri = f"ws://{host}:{port}" + envs = [ProbeEnv(i) for i in range(batch)] + state = { + "width": 1, + "stepped": False, + "held": None, + "last_feedback": None, + "reconnects": 0, + } + stops_ignored = 0 + while True: + try: + with websockets.sync.client.connect( + uri, + compression="deflate" if bug == "compress" else None, + max_size=2**20 if bug == "frame-cap" else None, + open_timeout=30, + ) as ws: + _probe_session(ws, envs, bug, state) + except TextFrame as frame: + print(f"server sent a text frame, which means a server error: {frame}") + return 2 + except websockets.exceptions.ConnectionClosed as closed: + reason = closed.rcvd.reason if closed.rcvd is not None else "" + code = closed.rcvd.code if closed.rcvd is not None else None + if reason == STOP_REASON: + if bug == "ignore-stop" and stops_ignored == 0: + stops_ignored += 1 + time.sleep(0.2) + continue + print("the server stopped the run") + return 0 + if reason == RESYNC_REASON: + # Section 7.6: the server's state for these envs went with the + # connection, so the feedback in flight cannot be completed. + state["held"] = ( + state["last_feedback"] if bug == "resend-feedback" else None + ) + state["reconnects"] += 1 + print("the server asked for a resync; reconnecting") + time.sleep(0.2) + continue + print(f"connection closed: code {code}, reason {reason!r}") + return 1 + + def main() -> int: p = argparse.ArgumentParser(description=__doc__) p.add_argument("--host", default="127.0.0.1") p.add_argument("--port", type=int, default=8000) p.add_argument("--steps", type=int, default=20) p.add_argument("--batch", type=int, default=1) + p.add_argument( + "--probe", + action="store_true", + help="run the SPEC.md section 8.1 probe env until the server stops the run", + ) + p.add_argument( + "--bug", choices=BUGS, default=None, help="break one rule on purpose" + ) a = p.parse_args() + if a.probe: + return run_probe(a.host, a.port, a.batch, a.bug) return run(a.host, a.port, a.steps, a.batch) diff --git a/examples/reference_server.py b/examples/reference_server.py new file mode 100644 index 0000000..d0b6ac0 --- /dev/null +++ b/examples/reference_server.py @@ -0,0 +1,256 @@ +"""A PlugRL training server that trains nothing, written against SPEC.md alone. + +The other side of `raw_client.py`. It answers every infer with a random +action chunk and counts the feedback it gets, and that is all - no policy, +no algorithm, nothing from plugrl-server. What it does do is every server +clause of SPEC.md: it speaks first; it answers time-major, with `env_ids`; it +validates what it receives and closes for a resync on anything malformed, +including an `info` that cannot be split per environment; it keeps one +client's mistake from ending the run for the others; and when it has had +`--steps` frames it ends the run with `plugrl-server-stop`. + + python reference_server.py --port 8000 --horizon 4 --action-dim 3 --steps 300 + +`plugrl-conformance-server` grades it, and `--bug NAME` makes it break one +server clause on purpose, so the grader's tests can show that each one is +checked. +""" + +from __future__ import annotations + +import argparse +import asyncio +import pathlib +import sys + +import numpy as np +import websockets.asyncio.server as ws_server +import websockets.exceptions as ws_exceptions + +# So a fresh checkout works without installing anything first. +sys.path.insert(0, str(pathlib.Path(__file__).resolve().parents[1] / "src")) + +from plugrl_protocol import msgpack_numpy # noqa: E402 +from plugrl_protocol.websocket_protocol import ( # noqa: E402 + SERVER_RESYNC_REASON, + SERVER_STOP_REASON, + MessageType, +) + +BUGS = ( + "no-metadata", # waits for the client instead of speaking first + "env-major", # sends the action as [n, H, da] + "no-env-ids", # leaves env_ids out of the action + "wrong-horizon", # declares H in the metadata and sends H + 1 + "lenient", # accepts a malformed infer instead of closing for a resync + "crash-on-info", # dies on an info it cannot split, as plugrl-server did before #108 + "plain-stop", # ends the run with a plain close, no plugrl-server-stop +) + + +class ProtocolError(ValueError): + """The client's message breaks SPEC.md; its connection closes for a resync.""" + + +def _int_vector(value, what: str) -> np.ndarray: + if not ( + isinstance(value, np.ndarray) and value.ndim == 1 and value.dtype.kind == "i" + ): + raise ProtocolError(f"{what} is not a one-dimensional integer array") + return value + + +def _leading(observation, what: str) -> int: + if not isinstance(observation, dict) or set(observation) != { + "images", + "states", + "text", + }: + raise ProtocolError(f"{what} is not a map of images, states and text") + sizes = { + int(array.shape[0]) + for group in ("images", "states") + for array in (observation[group] or {}).values() + if isinstance(array, np.ndarray) and array.ndim >= 1 + } + if len(sizes) > 1: + raise ProtocolError( + f"{what} arrays disagree on their leading dimension: {sorted(sizes)}" + ) + return sizes.pop() if sizes else -1 + + +def _info_entries(info: dict) -> int | None: + """How many per-env entries an info splits into: SPEC section 5.4's rule. + + The leading dimension of the first ndarray at the top level or one level + into a nested map; None if there is none, which only works for m = 1. + """ + for value in info.values(): + if isinstance(value, np.ndarray): + return int(value.shape[0]) + if isinstance(value, dict): + for nested in value.values(): + if isinstance(nested, np.ndarray): + return int(nested.shape[0]) + return None + + +class ReferenceServer: + def __init__(self, args: argparse.Namespace): + self.args = args + self.frames = 0 + self.stopping = asyncio.Event() + self.connections: set = set() + self.rng = np.random.default_rng(0) + + def check_infer(self, message: dict) -> np.ndarray: + if message.get("message_type") != str(MessageType.INFER): + raise ProtocolError(f"expected infer, got {message.get('message_type')!r}") + if self.args.bug == "lenient": + return np.asarray(message.get("env_indices", [0]), dtype=np.int64) + for key in ("data", "env_indices", "step_ids"): + if key not in message: + raise ProtocolError(f"infer has no {key!r}") + envs = _int_vector(message["env_indices"], "env_indices") + n = _leading(message["data"], "the observation") + if n not in (-1, len(envs)): + raise ProtocolError(f"observation batch {n}, env_indices {len(envs)}") + return envs + + def check_feedback(self, message: dict, holding: set) -> list[int]: + if message.get("message_type") != str(MessageType.FEEDBACK): + raise ProtocolError( + f"expected feedback, got {message.get('message_type')!r}" + ) + for key in ("data", "env_indices", "step_ids"): + if key not in message: + raise ProtocolError(f"feedback has no {key!r}") + envs = _int_vector(message["env_indices"], "env_indices").tolist() + m = len(envs) + data = message["data"] + if not isinstance(data, dict) or set(data) != { + "obs", "rewards", "terminated", "truncated", "info" + }: # fmt: skip + raise ProtocolError("feedback data does not have exactly its five keys") + for key in ("rewards", "terminated", "truncated"): + if not (isinstance(data[key], np.ndarray) and data[key].shape == (m,)): + raise ProtocolError(f"{key} is not an array of length {m}") + if _leading(data["obs"], "the feedback observation") not in (-1, m): + raise ProtocolError("feedback observation batch does not match env_indices") + info = data["info"] + if not isinstance(info, dict): + raise ProtocolError("info is not a map") + if info: + entries = _info_entries(info) + if (entries is None and m > 1) or (entries is not None and entries != m): + if self.args.bug == "crash-on-info": + raise RuntimeError( + "could not split info" + ) # and takes the server down + raise ProtocolError(f"info cannot be split into {m} entries") + # Feedback for an env this connection never answered is a transition + # with no start: dropped, with a word (SPEC section 7.6 SHOULD). + stale = [e for e in envs if e not in holding] + if stale: + print( + f"feedback for envs {stale} with no action on this connection; dropped" + ) + return [e for e in envs if e in holding] + + def action(self, envs: np.ndarray) -> dict: + horizon = self.args.horizon + (1 if self.args.bug == "wrong-horizon" else 0) + chunk = self.rng.uniform(-1, 1, (horizon, len(envs), self.args.action_dim)) + if self.args.bug == "env-major": + chunk = chunk.swapaxes(0, 1) + data = {"action": chunk.astype(np.float32)} + if self.args.bug != "no-env-ids": + data["env_ids"] = np.asarray(envs, dtype=np.int64) + return {"message_type": str(MessageType.ACTION), "data": data} + + async def handle(self, websocket) -> None: + packer = msgpack_numpy.Packer() + self.connections.add(websocket) + holding: set[int] = set() + try: + if self.args.bug != "no-metadata": + await websocket.send( + packer.pack( + { + "message_type": str(MessageType.METADATA), + "data": { + "protocol_version": 1, + "server": "reference_server.py", + "action_horizon": self.args.horizon, + "action_dim": self.args.action_dim, + }, + } + ) + ) + while True: + raw = await websocket.recv() + if not isinstance(raw, bytes): + raise ProtocolError("a text frame") + envs = self.check_infer(msgpack_numpy.unpackb(raw)) + await websocket.send(packer.pack(self.action(envs))) + holding.update(envs.tolist()) + raw = await websocket.recv() + if not isinstance(raw, bytes): + raise ProtocolError("a text frame") + for env in self.check_feedback(msgpack_numpy.unpackb(raw), holding): + holding.discard(env) + self.frames += 1 + if self.frames >= self.args.steps: + # serve() closes every connection, this one included, + # with the stop reason. + self.stopping.set() + await websocket.wait_closed() + return + except ProtocolError as exc: + print(f"protocol error: {exc}; closing for a resync") + await websocket.close(1001, SERVER_RESYNC_REASON) + except ws_exceptions.ConnectionClosed: + pass + except RuntimeError: + if self.args.bug == "crash-on-info": + await websocket.close(1011, "Internal server error.") + self.stopping.set() # the whole server goes, as it did + return + raise + finally: + self.connections.discard(websocket) + + async def serve(self) -> None: + async with ws_server.serve( + self.handle, self.args.host, self.args.port, compression=None, max_size=None + ) as server: + print( + f"reference server on ws://{self.args.host}:{self.args.port}", + flush=True, + ) + await self.stopping.wait() + for websocket in list(self.connections): + if self.args.bug == "plain-stop": + await websocket.close() + else: + await websocket.close(1001, SERVER_STOP_REASON) + server.close() + print(f"stopped after {self.frames} frames") + + +def main() -> int: + p = argparse.ArgumentParser(description=__doc__) + p.add_argument("--host", default="127.0.0.1") + p.add_argument("--port", type=int, default=8000) + p.add_argument("--horizon", type=int, default=4) + p.add_argument("--action-dim", type=int, default=3) + p.add_argument("--steps", type=int, default=300, help="frames before the run ends") + p.add_argument( + "--bug", choices=BUGS, default=None, help="break one rule on purpose" + ) + asyncio.run(ReferenceServer(p.parse_args()).serve()) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/pyproject.toml b/pyproject.toml index eb3b7df..ba7c3f4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -25,6 +25,7 @@ conformance = [ [project.scripts] plugrl-conformance = "plugrl_protocol.conformance:main" +plugrl-conformance-server = "plugrl_protocol.server_conformance:main" [dependency-groups] dev = [ diff --git a/src/plugrl_protocol/conformance.py b/src/plugrl_protocol/conformance.py index 4a7cf15..e534bed 100644 --- a/src/plugrl_protocol/conformance.py +++ b/src/plugrl_protocol/conformance.py @@ -6,9 +6,9 @@ someone writing a client in a new language, who wants to be told which clause they violated rather than to watch the loss fail to move. -This server speaks the same protocol and checks every clause of SPEC.md that -can be checked from one side of the wire. It prints a report and exits -non-zero if anything failed, so it can sit in a CI job. +This server speaks the same protocol and checks the clauses of SPEC.md that +can be checked from the server's side of the wire. It prints a report and +exits non-zero if anything failed, so it can sit in a CI job. One command, which starts your client once the server is listening and waits for both: @@ -25,11 +25,27 @@ before the server is listening fails for a reason that has nothing to do with the specification. +Two modes. By default the server watches one well-behaved connection, which +is all it can do with a client running an environment it knows nothing about. +With `--probe` the client runs the probe environment of SPEC.md section 8.1 +instead. Its observation says how many steps an episode has taken and which +action was applied last, so the server can tell what the client did with each +action chunk: whether it summed the reward, sent the terminal observation, +applied the actions time-major and in order, and read the dtype. Probe mode +also drives the connection rather than watching it - it speaks late, sends an +oversized metadata frame, and with `--scenario` closes the connection for a +resync, stops the run, or sends a text frame - so the clauses about closing +and reconnecting are checked too: + + plugrl-conformance --probe --scenario all \ + --client "python examples/raw_client.py --probe --batch 3" + What it deliberately does NOT require, because SPEC.md does not: * that the feedback's env set matches the infer's (section 4.3); * that `n` stays the same between requests (section 5.2); - * any particular key in the observation, or any camera at all; + * any particular key in the observation, or any camera at all, outside + probe mode; * the ` str: + return f"[{self.context}] {detail}" if self.context else detail def ok(self, clause: str) -> None: self.checked[clause] = self.checked.get(clause, 0) + 1 def fail(self, clause: str, detail: str) -> None: self.checked.setdefault(clause, 0) - self.failures.append((clause, detail)) + self.failures.append((clause, self._where(detail))) def require(self, condition: bool, clause: str, detail: str) -> bool: if condition: @@ -86,7 +122,7 @@ def advise(self, condition: bool, clause: str, detail: str) -> None: """Accepted, but worth saying out loud. Never fails the run.""" self.ok(clause) if not condition: - self.advisories.setdefault(clause, detail) + self.advisories.setdefault(clause, self._where(detail)) def render(self) -> str: lines = ["", "=" * 68, "SPEC.md conformance report", "=" * 68] @@ -191,17 +227,294 @@ def check_observation(report: Report, observation, where: str, expect_n: int) -> ) +# ------------------------------------------------------------- probe mode +# +# SPEC.md section 8.1. Each probe env counts the steps of its episode in +# `states["t"]` and echoes the action it applied last in `states["a"]`; every +# step pays a reward of 1, and env `i`'s episode ends, terminated, when `t` +# reaches `probe_episode_length(i)`. The server makes every action value say +# where in the chunk it was, so the echo shows exactly which one was applied. + + +def probe_episode_length(env: int) -> int: + """3, 5, 7, 3, ...: so chunks end at different steps in different envs.""" + return 3 + 2 * (env % 3) + + +def probe_action_value(step: int, env: int, dim: int) -> float: + """10000 per step of the chunk, 10 per env index, 1 per action dimension.""" + return 10000.0 * step + 10.0 * env + dim + + +def _describe_probe_value(value: float) -> str: + v = int(round(value)) + return f"step {v // 10000}, env {(v % 10000) // 10}, dim {v % 10}" + + +class ProbeTracker: + """What one connection's probe envs must look like, given what was sent.""" + + def __init__(self, report: Report, horizon: int, action_dim: int, fresh: bool): + self.report = report + self.horizon = horizon + self.action_dim = action_dim + # Env indices are connection-scoped (section 4.4), but the client's + # envs are not: after a reconnect they carry on mid-episode. So only + # the first connection can expect every env to start at t = 0. + self.fresh = fresh + self.state: dict[int, tuple[float, bool]] = {} + self.pending: dict[int, tuple[float, np.ndarray | None]] = {} + + def _states(self, observation, where: str, n: int): + report = self.report + states = observation.get("states") if isinstance(observation, dict) else None + t = states.get("t") if isinstance(states, dict) else None + a = states.get("a") if isinstance(states, dict) else None + if not report.require( + isinstance(t, np.ndarray) and isinstance(a, np.ndarray), + "8.1 probe observation carries states t and a", + f"{where}: states were " + f"{sorted(states) if isinstance(states, dict) else type(states).__name__}", + ): + return None + if not report.require( + t.shape == (n, 1) and a.ndim == 2 and a.shape[0] == n, + "8.1 probe t is [n, 1] and a is [n, d]", + f"{where}: t {t.shape}, a {a.shape}, n={n}", + ): + return None + return t[:, 0].astype(np.float64), a.astype(np.float64) + + def on_infer(self, env_indices: np.ndarray, observation) -> None: + report = self.report + parsed = self._states(observation, "infer", len(env_indices)) + if parsed is None: + return + t, _ = parsed + for row, env in enumerate(env_indices.tolist()): + report.require( + env not in self.pending, + "8 exactly one feedback for every action received", + f"env {env} asked for a chunk while still holding the last one", + ) + known = self.state.get(env) + if known is None: + if self.fresh: + report.require( + t[row] == 0, + "8.1 a probe env starts its first episode at t = 0", + f"env {env} began at t={t[row]:g}", + ) + else: + last_t, terminal = known + if terminal: + report.require( + t[row] == 0, + "8.1 after a terminal step, the env's next infer starts a new episode", + f"env {env} ended at t={last_t:g}, then asked from t={t[row]:g}", + ) + else: + report.require( + t[row] == last_t, + "8.1 an infer carries the observation the last feedback reported", + f"env {env}: feedback said t={last_t:g}, infer said t={t[row]:g}", + ) + self.pending[env] = (float(t[row]), None) + + def action(self, env_indices: np.ndarray, dtype: np.dtype) -> np.ndarray: + envs = env_indices.tolist() + action = np.zeros((self.horizon, len(envs), self.action_dim), dtype) + for row, env in enumerate(envs): + for k in range(self.horizon): + for d in range(self.action_dim): + action[k, row, d] = probe_action_value(k, env, d) + start, _ = self.pending.get(env, (0.0, None)) + self.pending[env] = (start, action[:, row, :].astype(np.float64)) + return action + + def on_feedback(self, env_indices: np.ndarray, data) -> None: + report = self.report + n = len(env_indices) + parsed = self._states( + data.get("obs") if isinstance(data, dict) else None, "feedback", n + ) + if parsed is None: + return + t, a = parsed + rewards = data.get("rewards") + terminated = data.get("terminated") + truncated = data.get("truncated") + shapes_ok = all( + isinstance(x, np.ndarray) and x.shape == (n,) + for x in (rewards, terminated, truncated) + ) + if not shapes_ok: + return # check_feedback has already said why + for row, env in enumerate(env_indices.tolist()): + holding = env in self.pending and self.pending[env][1] is not None + if not holding and not self.fresh and env not in self.state: + # Never given an action on this connection, so the chunk it + # is feeding back came from an earlier one. + report.fail( + "7.6 no feedback for an action that arrived on an earlier connection", + f"env {env} sent feedback on a reconnected connection that " + "never gave it an action", + ) + continue + if not report.require( + holding, + "4.3 feedback names only envs that are holding an action", + f"env {env} had no action outstanding", + ): + continue + start, chunk = self.pending.pop(env) + end = float(t[row]) + length = probe_episode_length(env) + done = bool(terminated[row]) + + report.require( + not bool(truncated[row]), + "8.1 the probe env never truncates", + f"env {env} reported truncated", + ) + if done and not report.require( + end == length, + "5.4 a done step reports the observation that step returned", + f"env {env} terminated, but its observation says t={end:g}; its " + f"episode ends at t={length} - this is the next episode's " + "observation", + ): + self.state[env] = (end, True) + continue + + steps = end - start + if not report.require( + steps.is_integer() and 1 <= steps <= self.horizon, + "5.3 a client executes between 1 and H steps of a chunk", + f"env {env}: t went from {start:g} to {end:g} with H={self.horizon}", + ): + self.state[env] = (end, done) + continue + steps = int(steps) + + reward = float(rewards[row]) + hint = ( + " - the last step's reward, not the sum" + if reward == 1 and steps > 1 + else "" + ) + report.require( + reward == steps, + "5.4 rewards is the sum over the steps of the chunk", + f"env {env}: {steps} steps of reward 1, reported {reward:g}{hint}", + ) + + expected = chunk[steps - 1] + applied = a[row] + if applied.shape == expected.shape and np.array_equal(applied, expected): + report.ok( + "5.3 the action chunk is applied time-major, in order, as sent" + ) + else: + first = applied[0] if applied.size else float("nan") + report.fail( + "5.3 the action chunk is applied time-major, in order, as sent", + f"env {env}: after {steps} step(s) it applied {applied[:3].tolist()}, " + f"expected step {steps - 1}'s action {expected[:3].tolist()}" + + ( + f" (what it applied decodes as {_describe_probe_value(first)})" + if np.isfinite(first) and 0 <= first < 1e7 + else " (not a value the server sent - check the dtype)" + ), + ) + + report.require( + end <= length, + "5.4 a chunk stops at the step that ends the episode", + f"env {env}: its episode ends at t={length}, and it reached t={end:g}", + ) + report.require( + done == (end == length), + "5.4 terminated is set on the terminal step, and only there", + f"env {env}: t={end:g} of {length}, terminated={done}", + ) + if done: + episode = (data.get("info") or {}).get("episode") + if isinstance(episode, dict) and "r" in episode: + r = np.asarray(episode.get("r")).reshape(-1) + report.advise( + r.size > row and float(r[row]) == length, + "5.4 info.episode, when sent, reports the episode's return", + f"env {env}: an episode of {length} steps returned {length}, " + f"info.episode.r said {r.tolist()}", + ) + self.state[env] = (end, done) + + +# ------------------------------------------------------------- the server + + class ConformanceServer: - def __init__(self, horizon: int, action_dim: int, action_dtype: str, steps: int): + def __init__( + self, + horizon: int, + action_dim: int, + action_dtype: str, + steps: int, + *, + probe: bool = False, + scenario: str = "basic", + report: Report | None = None, + ): self.horizon = horizon self.action_dim = action_dim self.action_dtype = np.dtype(action_dtype) self.steps = steps - self.report = Report() + self.probe = probe + self.scenario = scenario + self.report = report if report is not None else Report() self.exchanges = 0 + self.connections = 0 + self.resynced = False + self.stopped = False self.finished = asyncio.Event() + def metadata(self) -> dict: + data = { + "protocol_version": 1, + "server": "conformance_server.py", + "action_horizon": self.horizon, + "action_dim": self.action_dim, + } + if self.probe: + # A key no client knows, and big enough that a client which kept + # its library's frame cap cannot receive this message. Section + # 5.1 says to ignore unknown keys, section 1.1 forbids the cap. + data["conformance_padding"] = "x" * METADATA_PADDING + return {"message_type": str(MessageType.METADATA), "data": data} + async def handle(self, websocket) -> None: + if self.probe: + await self._handle_probe(websocket) + else: + await self._handle_passive(websocket) + + def _envelope(self, payload) -> str | None: + """The message_type of a well-formed envelope, or None.""" + report = self.report + if not isinstance(payload, dict): + report.fail("2 message is a map", f"got {type(payload).__name__}") + return None + report.ok("2 message is a map") + kind = payload.get("message_type") + report.require( + isinstance(kind, str), + "2 message_type is a string, not bytes", + f"got {type(kind).__name__}", + ) + return kind + + async def _handle_passive(self, websocket) -> None: report = self.report packer = msgpack_numpy.Packer() @@ -210,36 +523,15 @@ async def handle(self, websocket) -> None: # being told out of band. They are sent here so that a client which # reads them is exercised, and one which ignores them is proved not # to break - section 5.1 requires no key to be present. - await websocket.send( - packer.pack( - { - "message_type": str(MessageType.METADATA), - "data": { - "protocol_version": 1, - "server": "conformance_server.py", - "action_horizon": self.horizon, - "action_dim": self.action_dim, - }, - } - ) - ) + await websocket.send(packer.pack(self.metadata())) awaiting = "infer" try: while self.exchanges < self.steps: payload = msgpack_numpy.unpackb(await websocket.recv()) - - if not isinstance(payload, dict): - report.fail("2 message is a map", f"got {type(payload).__name__}") + kind = self._envelope(payload) + if kind is None: return - report.ok("2 message is a map") - - kind = payload.get("message_type") - report.require( - isinstance(kind, str), - "2 message_type is a string, not bytes", - f"got {type(kind).__name__}", - ) if not report.require( kind == awaiting, "4.2 infer and feedback strictly alternate", @@ -277,6 +569,168 @@ async def handle(self, websocket) -> None: finally: self.finished.set() + async def _handle_probe(self, websocket) -> None: + report = self.report + packer = msgpack_numpy.Packer() + self.connections += 1 + index = self.connections + + if self.stopped: + report.fail( + "7.1 a client does not reconnect after plugrl-server-stop", + f"connection {index} arrived after the server stopped the run", + ) + await websocket.close(1001, SERVER_STOP_REASON) + return + if index > 1 and self.scenario == "resync": + report.ok("7.2 a client reconnects after plugrl-server-resync") + + request = getattr(websocket, "request", None) + offered = request.headers.get("Sec-WebSocket-Extensions", "") if request else "" + report.require( + "permessage-deflate" not in offered.lower(), + "1.1 a client does not offer compression", + f"the handshake offered {offered!r}", + ) + + # Speak late, and see whether the client waits for us. + first = asyncio.ensure_future(websocket.recv()) + done, _ = await asyncio.wait({first}, timeout=METADATA_DELAY) + if done: + if first.exception() is None: + report.fail( + "5.1 a client reads metadata before sending anything", + f"a message arrived within {METADATA_DELAY}s of the handshake, " + "before the server had spoken", + ) + else: + report.fail( + "connection", f"closed before metadata: {first.exception()}" + ) + self.finished.set() + return + report.ok("5.1 a client reads metadata before sending anything") + await websocket.send(packer.pack(self.metadata())) + + tracker = ProbeTracker(report, self.horizon, self.action_dim, fresh=index == 1) + awaiting = str(MessageType.INFER) + opening = True + try: + while True: + raw = await first if first is not None else await websocket.recv() + first = None + if not report.require( + isinstance(raw, bytes), + "2 every frame is binary", + "the client sent a text frame", + ): + return + payload = msgpack_numpy.unpackb(raw) + kind = self._envelope(payload) + if kind is None: + return + if opening and index > 1: + report.require( + kind == str(MessageType.INFER), + "7.6 a reconnecting client drops held feedback and starts with infer", + f"its first message on connection {index} was {kind!r}", + ) + opening = False + if not report.require( + kind == awaiting, + "4.2 infer and feedback strictly alternate", + f"expected {awaiting!r}, received {kind!r}", + ): + return + + if kind == str(MessageType.INFER): + if self.check_infer(payload) is None: + return + tracker.on_infer(payload["env_indices"], payload["data"]) + if self.scenario == "text": + await self._send_text_frame(websocket) + return + action = tracker.action(payload["env_indices"], self.action_dtype) + await websocket.send( + packer.pack( + { + "message_type": str(MessageType.ACTION), + "data": { + "env_ids": np.asarray( + payload["env_indices"], dtype=np.int64 + ), + "action": action, + }, + } + ) + ) + awaiting = str(MessageType.FEEDBACK) + if ( + self.scenario == "resync" + and not self.resynced + and self.exchanges >= self.steps // 2 + ): + # The client now holds a feedback it can never send. + self.resynced = True + await websocket.close(1001, SERVER_RESYNC_REASON) + return + else: + self.check_feedback(payload) + tracker.on_feedback(payload["env_indices"], payload["data"]) + awaiting = str(MessageType.INFER) + self.exchanges += 1 + if self.exchanges >= self.steps: + self.stopped = True + await websocket.close(1001, SERVER_STOP_REASON) + self.finished.set() + return + except ws_exceptions.ConnectionClosed as closed: + code = closed.rcvd.code if closed.rcvd is not None else None + if code == 1009: + report.fail( + "1.1 a client accepts frames larger than 1 MiB", + "it closed with 1009 (message too big) on the metadata, " + f"which is {METADATA_PADDING // (1024 * 1024)} MiB", + ) + else: + report.fail( + "connection", + f"the client closed the connection (code {code}) after " + f"{self.exchanges} exchanges, before the server ended the run", + ) + self.finished.set() + except Exception as exc: # noqa: BLE001 - the report is the output + report.fail("connection", f"{type(exc).__name__}: {exc}") + self.finished.set() + + async def _send_text_frame(self, websocket) -> None: + """Section 7.4: a text frame where binary was expected is fatal.""" + report = self.report + await websocket.send( + "Traceback (most recent call last): this text frame stands for a " + "server-side error, sent by the conformance checker" + ) + try: + extra = await asyncio.wait_for(websocket.recv(), timeout=10) + except ws_exceptions.ConnectionClosed: + report.ok("7.4 a client treats a text frame as fatal") + except asyncio.TimeoutError: + report.fail( + "7.4 a client treats a text frame as fatal", + "the connection was still open 10 s after the text frame", + ) + else: + kind = "a message" + try: + kind = repr(msgpack_numpy.unpackb(extra).get("message_type")) + except Exception: # noqa: BLE001 - only for the detail + pass + report.fail( + "7.4 a client treats a text frame as fatal", + f"after the text frame it went on and sent {kind}", + ) + self.finished.set() + def check_infer(self, payload: dict) -> int | None: report = self.report for key in ("data", "env_indices", "step_ids"): @@ -392,18 +846,24 @@ def check_feedback(self, payload: dict) -> None: check_observation(report, data.get("obs"), "feedback", m) -async def serve(args: argparse.Namespace) -> int: +async def run_scenario(args: argparse.Namespace, scenario: str, report: Report) -> int: + """One server, one client launch. Returns the number of exchanges.""" + report.context = scenario if args.probe and args.scenario == "all" else "" server = ConformanceServer( horizon=args.horizon, action_dim=args.action_dim, action_dtype=args.action_dtype, steps=args.steps, + probe=args.probe, + scenario=scenario, + report=report, ) async with ws_server.serve( server.handle, args.host, args.port, compression=None, max_size=None ): + mode = f"probe mode, scenario {scenario}" if args.probe else "passive mode" print( - f"conformance server on ws://{args.host}:{args.port} - " + f"conformance server on ws://{args.host}:{args.port} ({mode}) - " f"expecting {args.steps} exchanges, horizon {args.horizon}, " f"action dtype {args.action_dtype}", flush=True, @@ -420,26 +880,55 @@ async def serve(args: argparse.Namespace) -> int: try: await asyncio.wait_for(server.finished.wait(), timeout=args.timeout) except asyncio.TimeoutError: - server.report.fail("connection", f"no client finished in {args.timeout}s") + if scenario == "resync" and server.resynced and server.connections < 2: + report.fail( + "7.2 a client reconnects after plugrl-server-resync", + f"no new connection within {args.timeout}s of the resync", + ) + else: + report.fail("connection", f"no client finished in {args.timeout}s") if client is not None: + # Still listening, so a client that reconnects after the stop is + # seen doing it. try: - code = await asyncio.wait_for(client.wait(), timeout=30) + code = await asyncio.wait_for( + client.wait(), timeout=CLIENT_EXIT_TIMEOUT + ) except asyncio.TimeoutError: client.kill() await client.wait() - server.report.fail( - "client process", "still running 30s after the exchanges ended" + report.fail( + "client process", + f"still running {CLIENT_EXIT_TIMEOUT:.0f}s after the exchanges ended", ) else: # A client that crashes on the way out has not passed, however - # clean the exchanges looked. - if code != 0: - server.report.fail("client process", f"exited with status {code}") + # clean the exchanges looked. After a text frame (section 7.4) + # it is expected to stop with an error. + if not args.probe: + if code != 0: + report.fail("client process", f"exited with status {code}") + elif scenario != "text": + report.require( + code == 0, + "7.1 a client exits cleanly on plugrl-server-stop", + f"exited with status {code}", + ) + elif args.probe and server.stopped: + await asyncio.sleep(2) # long enough to see an immediate reconnect + return server.exchanges + - print(f"\ncompleted {server.exchanges} infer/action/feedback exchanges") - print(server.report.render()) - return 1 if server.report.failures else 0 +async def serve(args: argparse.Namespace) -> int: + report = Report() + scenarios = list(SCENARIOS) if args.scenario == "all" else [args.scenario] + exchanges = 0 + for scenario in scenarios: + exchanges += await run_scenario(args, scenario, report) + print(f"\ncompleted {exchanges} infer/action/feedback exchanges") + print(report.render()) + return 1 if report.failures else 0 def main() -> int: @@ -455,6 +944,25 @@ def main() -> int: help="the environment's action dtype; a client must read the typestr", ) parser.add_argument("--timeout", type=float, default=120.0) + parser.add_argument( + "--probe", + action="store_true", + help=( + "the client runs the probe environment of SPEC.md section 8.1, so " + "chunk sums, terminal observations and action order can be checked" + ), + ) + parser.add_argument( + "--scenario", + choices=(*SCENARIOS, "all"), + default="basic", + help=( + "probe mode only. basic: run, then stop with plugrl-server-stop. " + "resync: close for a resync halfway, then stop. text: answer an " + "infer with a text frame. all: each in turn, starting the client " + "once per scenario" + ), + ) parser.add_argument( "--client", default=None, @@ -466,7 +974,18 @@ def main() -> int: "started elsewhere, as it always has." ), ) - return asyncio.run(serve(parser.parse_args())) + args = parser.parse_args() + if args.scenario != "basic" and not args.probe: + parser.error("--scenario needs --probe") + if args.scenario == "all" and not args.client: + parser.error( + "--scenario all starts the client once per scenario: pass --client" + ) + if args.probe and args.action_dim > 9: + parser.error( + "probe mode encodes the action dimension in one digit: --action-dim <= 9" + ) + return asyncio.run(serve(args)) if __name__ == "__main__": diff --git a/src/plugrl_protocol/server_conformance.py b/src/plugrl_protocol/server_conformance.py new file mode 100644 index 0000000..7394ded --- /dev/null +++ b/src/plugrl_protocol/server_conformance.py @@ -0,0 +1,485 @@ +"""A client that checks a training server against SPEC.md and says what it broke. + +`plugrl-conformance` grades an env client. This is the other half: it +connects to a training server - `plugrl-server`, or any other implementation +of the protocol - drives it the way a client may, and checks every server +clause it can see from the wire: + + plugrl-conformance-server --port 8000 --state-dim 3 + +It runs in phases, each on connections of its own: + + * exchange: the server speaks first, with a metadata message; every action + is time-major, `[H, n, *da]`, with `env_ids` equal to the request's + `env_indices` in order, and agrees with the `action_horizon` and + `action_dim` the metadata declares. The client then does what SPEC.md + lets it do: ragged batches, a feedback env set that differs from the + infer set, terminated episodes, and an observation frame over 1 MiB. The + server must keep the connection open through all of it. + * errors: a malformed infer, two infers in a row, and a feedback whose + `info` cannot be split per environment. Each must close that connection + with 1001 and `plugrl-server-resync` (sections 4.2, 5.4, 7.2), and the + server must still accept a new connection afterwards. + * scoping: two connections at once, both using env index 0 (section 4.4). + * stop, with `--until-stop`: exchange until the server ends the run, which + it must do with 1001 and `plugrl-server-stop` (section 7.1). + +The observation it sends is states only, `states[--state-key]` of +`--state-dim` float32 values per env, which is what a state-based policy +reads; add `--image-key` for one that needs a camera. A server whose policy +needs anything else cannot be driven by this checker. + +A server batches inference across its connections and may wait for every +connected client before it infers, so the phases never leave a connection +idle while another one waits. +""" + +from __future__ import annotations + +import argparse +import asyncio + +import numpy as np +import websockets.asyncio.client as ws_client +import websockets.exceptions as ws_exceptions + +from plugrl_protocol import msgpack_numpy +from plugrl_protocol.conformance import Report +from plugrl_protocol.websocket_protocol import ( + SERVER_RESYNC_REASON, + SERVER_STOP_REASON, + MessageType, +) + +BIG_FRAME_PIXELS = 512 # three envs of 512x512x3 uint8 is 2.25 MiB +# What a failed connection attempt raises: refused, timed out, or a server +# that answered the handshake with an HTTP error, as one shutting down does. +CONNECT_ERRORS = (OSError, asyncio.TimeoutError, ws_exceptions.WebSocketException) + + +class Probe: + """One connection to the server under test, with the checks it makes.""" + + def __init__(self, args: argparse.Namespace, report: Report): + self.args = args + self.report = report + self.packer = msgpack_numpy.Packer() + self.ws = None + self.metadata: dict = {} + self.rng = np.random.default_rng(0) + + # ---------------------------------------------------------- wire + + async def connect(self) -> bool: + report = self.report + self.ws = await ws_client.connect( + f"ws://{self.args.host}:{self.args.port}", + compression=None, + max_size=None, + open_timeout=self.args.timeout, + ) + try: + raw = await asyncio.wait_for(self.ws.recv(), timeout=self.args.timeout) + except asyncio.TimeoutError: + report.fail( + "5.1 the server speaks first, with metadata", + f"nothing arrived within {self.args.timeout}s of the handshake", + ) + return False + if not report.require( + isinstance(raw, bytes), + "2 every frame is binary", + "the metadata was a text frame", + ): + return False + message = msgpack_numpy.unpackb(raw) + if not report.require( + isinstance(message, dict) + and message.get("message_type") == str(MessageType.METADATA) + and isinstance(message.get("data"), dict), + "5.1 the server speaks first, with metadata", + f"the first message was {str(message)[:120]}", + ): + return False + self.metadata = message["data"] + version = self.metadata.get("protocol_version") + if version is not None: + report.require( + version == 1, + "10 protocol_version, when sent, is 1", + f"protocol_version was {version!r}", + ) + return True + + def observation(self, n: int, *, big: bool = False) -> dict: + args = self.args + states = { + args.state_key: self.rng.normal(size=(n, args.state_dim)).astype(np.float32) + } + images = {} + if args.image_key: + side = args.image_size + images[args.image_key] = np.zeros((n, side, side, 3), dtype=np.uint8) + if big: + images["conformance_padding"] = np.zeros( + (n, BIG_FRAME_PIXELS, BIG_FRAME_PIXELS, 3), dtype=np.uint8 + ) + return {"images": images, "states": states, "text": np.asarray(["probe"] * n)} + + async def infer(self, envs: list[int], *, big: bool = False) -> np.ndarray | None: + """Send an infer and check the action that comes back.""" + report = self.report + await self.ws.send( + self.packer.pack( + { + "message_type": str(MessageType.INFER), + "data": self.observation(len(envs), big=big), + "env_indices": np.asarray(envs, dtype=np.int64), + "step_ids": np.zeros(len(envs), dtype=np.int64), + } + ) + ) + raw = await self.recv() + if raw is None: + return None + message = msgpack_numpy.unpackb(raw) + if not report.require( + isinstance(message, dict) + and message.get("message_type") == str(MessageType.ACTION), + "5.3 an infer is answered with an action", + f"got {str(message)[:120]}", + ): + return None + data = message.get("data") or {} + action = data.get("action") + if not report.require( + isinstance(action, np.ndarray) and action.ndim >= 2, + "5.3 action is an array of at least [H, n]", + f"action was {type(action).__name__} {getattr(action, 'shape', None)}", + ): + return None + report.require( + action.dtype.kind in "fiu", + "5.3 action has a numeric dtype", + f"dtype {action.dtype}", + ) + horizon, rows = action.shape[0], action.shape[1] + report.require( + rows == len(envs), + "5.3 action is time-major, [H, n, *da]", + f"{len(envs)} envs asked, action shape {action.shape}", + ) + report.require(horizon >= 1, "5.3 the horizon is at least 1", f"H = {horizon}") + env_ids = data.get("env_ids") + report.require( + isinstance(env_ids, np.ndarray) and env_ids.tolist() == list(envs), + "4.3 action env_ids equal the infer's env_indices, in order", + f"asked {list(envs)}, env_ids " + f"{env_ids.tolist() if isinstance(env_ids, np.ndarray) else env_ids!r}", + ) + declared_h = self.metadata.get("action_horizon") + if isinstance(declared_h, int): + report.require( + horizon == declared_h, + "5.1 the action matches the metadata's action_horizon", + f"metadata said {declared_h}, action has H = {horizon}", + ) + declared_d = self.metadata.get("action_dim") + if isinstance(declared_d, int): + width = int(np.prod(action.shape[2:])) if action.ndim > 2 else 1 + report.require( + width == declared_d, + "5.1 the action matches the metadata's action_dim", + f"metadata said {declared_d}, action has {width} per env", + ) + return action + + async def feedback( + self, envs: list[int], *, done: list[bool] | None = None, info=None + ): + done = done or [False] * len(envs) + if info is None: + info = {} + if any(done): + info = { + "episode": { + "r": np.asarray([5.0 if d else 0.0 for d in done], np.float32), + "l": np.asarray([5 if d else 0 for d in done], np.int64), + "s": np.asarray([False] * len(envs)), + "mask": np.asarray(done), + } + } + await self.ws.send( + self.packer.pack( + { + "message_type": str(MessageType.FEEDBACK), + "env_indices": np.asarray(envs, dtype=np.int64), + "step_ids": np.zeros(len(envs), dtype=np.int64), + "data": { + "obs": self.observation(len(envs)), + "rewards": np.ones(len(envs), dtype=np.float32), + "terminated": np.asarray(done, dtype=np.bool_), + "truncated": np.zeros(len(envs), dtype=np.bool_), + "info": info, + }, + } + ) + ) + + async def recv(self) -> bytes | None: + """The next binary message, or None with the reason recorded.""" + report = self.report + try: + raw = await asyncio.wait_for(self.ws.recv(), timeout=self.args.timeout) + except asyncio.TimeoutError: + report.fail("connection", f"no reply within {self.args.timeout}s") + return None + except ws_exceptions.ConnectionClosed as closed: + code = closed.rcvd.code if closed.rcvd else None + reason = closed.rcvd.reason if closed.rcvd else "" + report.fail( + "connection", + f"the server closed the connection (code {code}, reason {reason!r}) " + "where SPEC.md lets a client go on", + ) + return None + report.require( + isinstance(raw, bytes), "2 every frame is binary", "got a text frame" + ) + return raw if isinstance(raw, bytes) else None + + async def expect_resync(self, clause: str, what: str) -> None: + """The server must close this connection for a resync, and nothing else.""" + report = self.report + try: + message = await asyncio.wait_for(self.ws.recv(), timeout=self.args.timeout) + except ws_exceptions.ConnectionClosed as closed: + code = closed.rcvd.code if closed.rcvd else None + reason = closed.rcvd.reason if closed.rcvd else "" + report.require( + code == 1001 and reason == SERVER_RESYNC_REASON, + clause, + f"after {what}, the server closed with code {code}, reason {reason!r}", + ) + except asyncio.TimeoutError: + report.fail(clause, f"after {what}, the connection stayed open") + else: + report.fail( + clause, f"after {what}, the server answered: {str(message)[:80]}" + ) + + async def close(self) -> None: + if self.ws is not None: + await self.ws.close() + self.ws = None + + +# ---------------------------------------------------------------- phases + + +async def phase_exchange(args: argparse.Namespace, report: Report) -> None: + probe = Probe(args, report) + if not await probe.connect(): + return + try: + # All three ask; env 0 finishes first and asks again on its own; the + # others report afterwards. So the feedback set differs from the + # infer set, and n varies between infers (sections 4.3 and 5.2). + if await probe.infer([0, 1, 2]) is None: + return + await probe.feedback([0], done=[True]) + if await probe.infer([0]) is None: + return + await probe.feedback([1, 2]) + report.ok("4.3 the server accepts a feedback env set unlike the infer's") + if await probe.infer([2, 1]) is None: # and in another order + return + await probe.feedback([0, 2, 1], done=[False, True, False]) + if await probe.infer([0, 1, 2], big=True) is None: + return + report.ok("1.1 the server accepts a frame larger than 1 MiB") + await probe.feedback([0, 1, 2]) + for _ in range(args.exchanges): + if await probe.infer([0, 1, 2]) is None: + return + await probe.feedback([0, 1, 2]) + report.ok("4.2 the server keeps a well-behaved connection open") + finally: + await probe.close() + + +async def phase_errors(args: argparse.Namespace, report: Report) -> None: + cases = [ + ( + "7.2 a malformed infer closes the connection for a resync", + "an infer with no env_indices", + "malformed", + ), + ( + "4.2 two infers in a row close the connection for a resync", + "a second infer where a feedback was due", + "double", + ), + ( + "5.4 an info that cannot be split per env closes the connection for a resync", + 'a feedback for 2 envs with info {"task": "pick"}', + "info", + ), + ] + for clause, what, kind in cases: + probe = Probe(args, report) + if not await probe.connect(): + return + try: + if kind == "malformed": + await probe.ws.send( + probe.packer.pack( + { + "message_type": str(MessageType.INFER), + "data": probe.observation(1), + "step_ids": np.zeros(1, dtype=np.int64), + } + ) + ) + else: + if await probe.infer([0, 1]) is None: + return + if kind == "double": + await probe.ws.send( + probe.packer.pack( + { + "message_type": str(MessageType.INFER), + "data": probe.observation(2), + "env_indices": np.asarray([0, 1], dtype=np.int64), + "step_ids": np.zeros(2, dtype=np.int64), + } + ) + ) + else: + await probe.feedback([0, 1], info={"task": "pick"}) + await probe.expect_resync(clause, what) + finally: + await probe.close() + + # One client's mistake must not end the run for everyone else. + alive = Probe(args, report) + try: + ok = await alive.connect() + except CONNECT_ERRORS as exc: + report.fail( + "7.2 the server survives a client's protocol error", + f"after {what}, a new connection was refused: {exc}", + ) + return # the server is gone; the later cases would only repeat it + else: + report.require( + ok, + "7.2 the server survives a client's protocol error", + f"after {what}, a new connection got no metadata", + ) + finally: + await alive.close() + + +async def phase_scoping(args: argparse.Namespace, report: Report) -> None: + a, b = Probe(args, report), Probe(args, report) + if not (await a.connect() and await b.connect()): + await a.close() + await b.close() + return + try: + for _ in range(3): + actions = await asyncio.gather(a.infer([0]), b.infer([0])) + if any(x is None for x in actions): + return + await asyncio.gather(a.feedback([0]), b.feedback([0])) + report.ok("4.4 env indices are connection-scoped: two clients can both use 0") + finally: + await a.close() + await b.close() + + +async def phase_stop(args: argparse.Namespace, report: Report) -> None: + probe = Probe(args, report) + if not await probe.connect(): + return + try: + for count in range(args.max_exchanges): + try: + await probe.ws.send( + probe.packer.pack( + { + "message_type": str(MessageType.INFER), + "data": probe.observation(3), + "env_indices": np.asarray([0, 1, 2], dtype=np.int64), + "step_ids": np.zeros(3, dtype=np.int64), + } + ) + ) + await asyncio.wait_for(probe.ws.recv(), timeout=args.timeout) + await probe.feedback([0, 1, 2], done=[count % 5 == 4] * 3) + except ws_exceptions.ConnectionClosed as closed: + code = closed.rcvd.code if closed.rcvd else None + reason = closed.rcvd.reason if closed.rcvd else "" + report.require( + code == 1001 and reason == SERVER_STOP_REASON, + "7.1 the server ends a run with plugrl-server-stop", + f"it closed with code {code}, reason {reason!r}", + ) + return + except asyncio.TimeoutError: + report.fail("connection", f"no reply within {args.timeout}s") + return + report.fail( + "7.1 the server ends a run with plugrl-server-stop", + f"still running after {args.max_exchanges} exchanges; raise --max-exchanges", + ) + finally: + await probe.close() + + +async def run(args: argparse.Namespace) -> int: + report = Report() + phases = [ + ("exchange", phase_exchange), + ("errors", phase_errors), + ("scoping", phase_scoping), + ] + if args.until_stop: + phases.append(("stop", phase_stop)) + for name, phase in phases: + report.context = name + print(f"phase: {name}", flush=True) + try: + await phase(args, report) + except CONNECT_ERRORS as exc: + report.fail("connection", f"cannot connect to the server: {exc}") + break + # Let the server drop the phase's connections before the next one, + # so it does not wait for a client that has gone. + await asyncio.sleep(0.5) + print(report.render()) + return 1 if report.failures else 0 + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", type=int, default=8000) + parser.add_argument("--state-key", default="obs") + parser.add_argument("--state-dim", type=int, default=3) + parser.add_argument("--image-key", default=None) + parser.add_argument("--image-size", type=int, default=224) + parser.add_argument("--exchanges", type=int, default=5) + parser.add_argument("--timeout", type=float, default=60.0) + parser.add_argument( + "--until-stop", + action="store_true", + help="last, exchange until the server ends the run, and check how it does", + ) + parser.add_argument("--max-exchanges", type=int, default=2000) + return asyncio.run(run(parser.parse_args())) + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/test_conformance_probe.py b/tests/test_conformance_probe.py new file mode 100644 index 0000000..eb3b0ea --- /dev/null +++ b/tests/test_conformance_probe.py @@ -0,0 +1,87 @@ +"""The probe checker catches each rule a client can break, and passes one that keeps them. + +`examples/raw_client.py --probe` runs the probe environment of SPEC.md +section 8.1 and follows every client clause of section 8. Its `--bug` option +breaks one clause on purpose. Each case below runs the checker against one +bug and expects it to name that clause, so a checker that stopped checking a +clause fails here instead of quietly printing "no violations". +""" + +from __future__ import annotations + +import pathlib +import socket +import subprocess +import sys + +import pytest + +pytest.importorskip("websockets") + +ROOT = pathlib.Path(__file__).resolve().parents[1] +CLIENT = ROOT / "examples" / "raw_client.py" + + +def _free_port() -> int: + with socket.socket() as s: + s.bind(("127.0.0.1", 0)) + return s.getsockname()[1] + + +def _check(scenario: str, bug: str | None = None) -> subprocess.CompletedProcess: + port = _free_port() + client = f'"{sys.executable}" "{CLIENT}" --probe --batch 3 --port {port}' + if bug: + client += f" --bug {bug}" + return subprocess.run( + [ + sys.executable, "-m", "plugrl_protocol.conformance", + "--probe", "--scenario", scenario, "--port", str(port), + "--steps", "8", "--horizon", "4", "--action-dim", "3", + "--action-dtype", "float64", "--timeout", "40", + "--client", client, + ], + capture_output=True, + text=True, + timeout=180, + ) # fmt: skip + + +def test_a_conforming_client_passes_every_scenario(): + result = _check("all") + + assert result.returncode == 0, result.stdout + result.stderr + assert "no violations" in result.stdout + for clause in ( + "5.4 rewards is the sum over the steps of the chunk", + "5.4 a done step reports the observation that step returned", + "5.3 the action chunk is applied time-major, in order, as sent", + "7.2 a client reconnects after plugrl-server-resync", + "7.6 a reconnecting client drops held feedback and starts with infer", + "7.4 a client treats a text frame as fatal", + "7.1 a client exits cleanly on plugrl-server-stop", + ): + assert f"ok {clause}" in result.stdout, clause + + +@pytest.mark.parametrize( + ("bug", "scenario", "clause"), + [ + ("early-infer", "basic", "5.1 a client reads metadata before sending anything"), + ("compress", "basic", "1.1 a client does not offer compression"), + ("frame-cap", "basic", "1.1 a client accepts frames larger than 1 MiB"), + ("float32", "basic", "5.3 the action chunk is applied time-major, in order, as sent"), + ("env-major", "basic", "5.3 the action chunk is applied time-major, in order, as sent"), + ("last-step-reward", "basic", "5.4 rewards is the sum over the steps of the chunk"), + ("reset-in-step", "basic", "5.4 a done step reports the observation that step returned"), + ("resend-feedback", "resync", "7.6 a reconnecting client drops held feedback and starts with infer"), + ("ignore-stop", "basic", "7.1 a client does not reconnect after plugrl-server-stop"), + ("ignore-text", "text", "7.4 a client treats a text frame as fatal"), + ("stale-chunk", "resync", "7.6 no feedback for an action that arrived on an earlier connection"), + ], +) # fmt: skip +def test_each_broken_rule_is_named(bug, scenario, clause): + result = _check(scenario, bug) + + assert result.returncode == 1, result.stdout + result.stderr + assert f"FAIL {clause}" in result.stdout, result.stdout diff --git a/tests/test_server_conformance.py b/tests/test_server_conformance.py new file mode 100644 index 0000000..a16831e --- /dev/null +++ b/tests/test_server_conformance.py @@ -0,0 +1,96 @@ +"""The server checker passes a conforming server and names each rule a server breaks. + +`examples/reference_server.py` implements every server clause of SPEC.md and +trains nothing. Its `--bug` option breaks one clause on purpose; each case +below runs `plugrl-conformance-server` against one bug and expects it to +name that clause. +""" + +from __future__ import annotations + +import pathlib +import socket +import subprocess +import sys +import time + +import pytest + +pytest.importorskip("websockets") + +ROOT = pathlib.Path(__file__).resolve().parents[1] +SERVER = ROOT / "examples" / "reference_server.py" + + +def _free_port() -> int: + with socket.socket() as s: + s.bind(("127.0.0.1", 0)) + return s.getsockname()[1] + + +def _check(bug: str | None = None, *extra: str) -> subprocess.CompletedProcess: + port = _free_port() + command = [sys.executable, str(SERVER), "--port", str(port), "--steps", "60"] + if bug: + command += ["--bug", bug] + server = subprocess.Popen( + command, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL + ) + try: + for _ in range(100): # until it listens + with socket.socket() as s: + if s.connect_ex(("127.0.0.1", port)) == 0: + break + time.sleep(0.1) + return subprocess.run( + [ + sys.executable, "-m", "plugrl_protocol.server_conformance", + "--port", str(port), "--timeout", "5", *extra, + ], + capture_output=True, + text=True, + timeout=180, + ) # fmt: skip + finally: + server.kill() + server.wait() + + +def test_a_conforming_server_passes_every_phase(): + result = _check(None, "--until-stop") + + assert result.returncode == 0, result.stdout + result.stderr + assert "no violations" in result.stdout + for clause in ( + "5.1 the server speaks first, with metadata", + "5.3 action is time-major, [H, n, *da]", + "4.3 action env_ids equal the infer's env_indices, in order", + "4.3 the server accepts a feedback env set unlike the infer's", + "1.1 the server accepts a frame larger than 1 MiB", + "7.2 a malformed infer closes the connection for a resync", + "4.2 two infers in a row close the connection for a resync", + "5.4 an info that cannot be split per env closes the connection for a resync", + "7.2 the server survives a client's protocol error", + "4.4 env indices are connection-scoped: two clients can both use 0", + "7.1 the server ends a run with plugrl-server-stop", + ): + assert f"ok {clause}" in result.stdout, clause + + +@pytest.mark.parametrize( + ("bug", "extra", "clause"), + [ + ("no-metadata", (), "5.1 the server speaks first, with metadata"), + ("env-major", (), "5.3 action is time-major, [H, n, *da]"), + ("no-env-ids", (), "4.3 action env_ids equal the infer's env_indices, in order"), + ("wrong-horizon", (), "5.1 the action matches the metadata's action_horizon"), + ("lenient", (), "7.2 a malformed infer closes the connection for a resync"), + ("crash-on-info", (), "7.2 the server survives a client's protocol error"), + ("plain-stop", ("--until-stop",), "7.1 the server ends a run with plugrl-server-stop"), + ], +) # fmt: skip +def test_each_broken_rule_is_named(bug, extra, clause): + result = _check(bug, *extra) + + assert result.returncode == 1, result.stdout + result.stderr + assert f"FAIL {clause}" in result.stdout, result.stdout