This repository has no description
Something went wrong. Try again.
10 kB · 257 lines
Python
at main
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258"""Golden-bytes agreement between the Python endpoint and the sharedenvelope fixture, plus the wire-level target/namespace binding checks.The same fixture bytes are asserted in Rust (klbr-runtime/tests/protocol.rs)and TypeScript (klbr-web/src/lib/protocol-v2.test.ts)."""from __future__ import annotations
import jsonimport osimport unittestfrom pathlib import Path
from klbr_runtime.wire import ( BindingToken, ProtocolError, RELAY_PROTOCOL_VERSION, VERSION, Wire, measured_runtime_identity, validate_target_stamp, validate_worker_identity,)
FIXTURE = Path(__file__).resolve().parents[3] / "tests" / "runtime-fixtures" / "envelopes.json"
def load_fixture() -> list[dict]: return json.loads(FIXTURE.read_text())
def remote_stamp() -> dict: return { "target_id": "linux-build", "namespace_id": "main", "namespace_generation": "ns-gen-remote-1", }
class SharedFixtureAgreement(unittest.TestCase): def test_fixture_parses_and_has_the_remote_handshake_entries(self): entries = load_fixture() self.assertEqual(len(entries), 11, "the shared fixture drives all three consumers") hello_remote = [e for e in entries if e["message"]["kind"] == "hello" and e["message"].get("identity")] execute_remote = [e for e in entries if "target" in e] self.assertEqual(len(hello_remote), 1) self.assertEqual(len(execute_remote), 1) identity = hello_remote[0]["message"]["identity"] validate_worker_identity(identity) validate_target_stamp(identity["target"]["stamp"]) validate_target_stamp(execute_remote[0]["target"])
def test_local_golden_hello_stays_valid_without_identity(self): local_hello = load_fixture()[0] self.assertEqual(local_hello["message"]["kind"], "hello") self.assertNotIn("identity", local_hello["message"]) self.assertNotIn("target", local_hello)
class WorkerIdentityChecks(unittest.TestCase): def base_identity(self) -> dict: return { "protocol_version": VERSION, "relay_protocol_version": RELAY_PROTOCOL_VERSION, "runtime": { "platform": "Linux", "arch": "aarch64", "interpreter_path": "/usr/bin/python3", "interpreter_version": "3.11.9", "capabilities": [ {"capability": "shell", "shell": "/bin/zsh"}, {"capability": "closures", "dill_available": False}, ], }, }
def test_incompatible_worker_version_is_a_clean_refusal(self): identity = self.base_identity() identity["protocol_version"] = 2 with self.assertRaises(ProtocolError) as ctx: validate_worker_identity(identity) self.assertIn("unsupported worker protocol version 2", str(ctx.exception))
def test_incompatible_relay_version_is_a_clean_refusal(self): identity = self.base_identity() identity["relay_protocol_version"] = RELAY_PROTOCOL_VERSION + 1 with self.assertRaises(ProtocolError) as ctx: validate_worker_identity(identity) self.assertIn("unsupported relay protocol version", str(ctx.exception))
def test_incomplete_identity_is_refused_field_by_field(self): identity = self.base_identity() identity["runtime"]["platform"] = "" with self.assertRaises(ProtocolError) as ctx: validate_worker_identity(identity) self.assertIn("runtime identity platform", str(ctx.exception))
def test_advertised_but_unusable_capability_is_expressible(self): identity = self.base_identity() validate_worker_identity(identity) caps = {c["capability"]: c for c in identity["runtime"]["capabilities"]} self.assertFalse(caps["closures"]["dill_available"], "advertised but unusable")
broken = self.base_identity() broken["runtime"]["capabilities"] = [ {"capability": "relay_framing", "version": RELAY_PROTOCOL_VERSION + 3} ] with self.assertRaises(ProtocolError) as ctx: validate_worker_identity(broken) self.assertIn("newer than this endpoint supports", str(ctx.exception))
def test_malformed_capability_entry_is_refused(self): identity = self.base_identity() identity["runtime"]["capabilities"] = [{"capability": 7}] with self.assertRaises(ProtocolError): validate_worker_identity(identity)
class EnvelopeTargetBindingChecks(unittest.TestCase): def _stubs(self): class FakeWriter: def __init__(self): self.sent = b""
def write(self, data): self.sent += data
async def drain(self): return None
class FakeReader: def __init__(self, data: bytes): self.data = data self.pos = 0
async def readexactly(self, n: int) -> bytes: chunk = self.data[self.pos:self.pos + n] self.pos += n if len(chunk) < n: raise asyncio.IncompleteReadError(chunk, n) return chunk
return FakeReader, FakeWriter
def test_local_envelope_claiming_a_stamp_is_forged(self): import asyncio
import struct as _struct
body = json.dumps({ "version": VERSION, "session_id": "test-session", "generation": "test-generation", "message": {"kind": "shutdown"}, "target": remote_stamp(), }).encode("utf-8") FakeReader, _ = self._stubs() reader = FakeReader(_struct.pack(">I", len(body)) + body) wire = Wire(reader, None, "test-session", "test-generation") # no target with self.assertRaises(ProtocolError) as ctx: asyncio.run(wire.receive()) self.assertIn("must not claim a remote target stamp", str(ctx.exception))
def test_remote_envelope_missing_or_foreign_stamp_is_fatal(self): import asyncio
import struct as _struct
def frame(packet: dict) -> bytes: body = json.dumps(packet).encode("utf-8") return _struct.pack(">I", len(body)) + body
FakeReader, _ = self._stubs() missing = frame({ "version": VERSION, "session_id": "test-session", "generation": "test-generation", "message": {"kind": "shutdown"}, }) foreign = frame({ "version": VERSION, "session_id": "test-session", "generation": "test-generation", "message": {"kind": "shutdown"}, "target": {**remote_stamp(), "namespace_generation": "ns-gen-other"}, }) matching = frame({ "version": VERSION, "session_id": "test-session", "generation": "test-generation", "message": {"kind": "shutdown"}, "target": remote_stamp(), }) wire = Wire(FakeReader(missing), None, "test-session", "test-generation", target=remote_stamp()) with self.assertRaises(ProtocolError) as ctx: asyncio.run(wire.receive()) self.assertIn("missing its target stamp", str(ctx.exception))
wire = Wire(FakeReader(foreign), None, "test-session", "test-generation", target=remote_stamp()) with self.assertRaises(ProtocolError) as ctx: asyncio.run(wire.receive()) self.assertIn("foreign target stamp", str(ctx.exception))
wire = Wire(FakeReader(matching), None, "test-session", "test-generation", target=remote_stamp()) message = asyncio.run(wire.receive()) self.assertEqual(message["kind"], "shutdown")
def test_remote_wire_stamps_its_own_sends(self): import asyncio
FakeReader, FakeWriter = self._stubs() writer = FakeWriter() wire = Wire(FakeReader(b""), writer, "test-session", "test-generation", target=remote_stamp()) asyncio.run(wire.send({"kind": "shutdown"})) packet = json.loads(writer.sent[4:].decode("utf-8")) self.assertEqual(packet["target"], remote_stamp()) local_writer = FakeWriter() local = Wire(FakeReader(b""), local_writer, "test-session", "test-generation") asyncio.run(local.send({"kind": "shutdown"})) local_packet = json.loads(local_writer.sent[4:].decode("utf-8")) self.assertNotIn("target", local_packet, "local workers omit the stamp")
if __name__ == "__main__": unittest.main()
class LaunchBindingVsMeasuredIdentity(unittest.TestCase): """The echoed --target binding is NOT measured identity (parent's trap)."""
def test_binding_token_carries_no_identity_fields(self): token = BindingToken.parse(remote_stamp()) stamp = token.as_stamp() self.assertEqual(stamp, remote_stamp()) # Structurally: no account/machine/platform/cwd fields exist on the token. self.assertEqual(sorted(token.__dict__), ["namespace_generation", "namespace_id", "target_id"]) for forbidden in ("account", "machine", "platform", "cwd", "interpreter_path"): self.assertNotIn(forbidden, stamp)
def test_hello_identity_runtime_is_measured_not_echoed(self): runtime = measured_runtime_identity() this_platform = __import__("platform").system() this_arch = __import__("platform").machine() self.assertEqual(runtime["platform"], this_platform, "the endpoint measures its own platform") self.assertEqual(runtime["arch"], this_arch) import sys self.assertEqual(runtime["interpreter_path"], sys.executable) # The binding token is absent from the runtime identity. self.assertNotIn("namespace_generation", runtime)
def test_parse_rejects_a_malformed_binding(self): with self.assertRaises(ProtocolError): BindingToken.parse({**remote_stamp(), "namespace_generation": "bad id"}) with self.assertRaises(ProtocolError): BindingToken.parse({"target_id": "t", "namespace_id": "n"})