Files
neuron-tai/tests/test_shard_runtime_harness.py
Dobromir Popov 966aa10854 distributed-gguf-runtime: add CMake skeleton, gRPC harness, split-GGUF provisioning, performance contracts
DGR-019  Lock alpha/beta performance contracts (evidence + contract framework)
DGR-020  Run controlled whole-model GGUF baseline (benchmark results & contracts)
DGR-024  Real generated-gRPC protocol harness (shard_runtime_server.py + tests)
DGR-026  split-GGUF provisioning outside /home (provision script + manifest + tests)
DGR-028  Numbered patch-stack apply & verify (llama_cpp_dependency.py + UPSTREAM_LOCK.json)
DGR-029  Native CMake skeleton + deterministic CPU lane (UPSTREAM_LOCK.json + cmake gating)

New modules:
  packages/node/meshnet_node/dgr_performance/  — performance contract framework
  packages/node/meshnet_node/split_gguf/        — split-GGUF manifest & provisioning
  scripts/provision_split_gguf.py               — artifact provisioning CLI
  tests/test_dgr_performance_contract.py        — contract validation tests
  tests/test_split_gguf_manifest.py             — manifest tests
  tests/test_split_gguf_provision.py            — provisioning tests
  tests/test_shard_runtime_harness.py           — gRPC harness tests
2026-07-23 09:55:00 +03:00

570 lines
22 KiB
Python

"""REAL DGR-024 generated-gRPC protocol harness.
This test drives the *committed* generated stubs over a *real* localhost TCP
socket to a *separately spawned* gRPC server subprocess. It proves:
(a) the committed stubs have not drifted from ``shard_runtime.proto``;
(b) a generated ``ShardRuntimeStub`` client reaches a real server over a
socket (GetCapability, Health, and a bidirectional Session stream);
(c) a DIRECT localhost hop and an OPAQUE RELAY carry of the exact captured
request bytes produce BYTE-IDENTICAL server responses — the relay forwards
raw captured frames without reinterpreting or inventing anything;
(d) the server echoed the *real* payload bytes, not a synthesized response.
No in-memory pipe, no fake channel, no synthetic model output. Any failure
(server unreachable, stubs out of date, handshake mismatch, byte inequality)
fails the test.
"""
from __future__ import annotations
import contextlib
import hashlib
import json
import os
import socket
import subprocess
import sys
import textwrap
import time
import grpc
import pytest
import zlib
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
# Package root is packages/node (import "meshnet_node...") and packages/tracker.
_PYTHONPATH = os.pathsep.join(
[os.path.join(REPO_ROOT, "packages", "node"), os.path.join(REPO_ROOT, "packages", "tracker")]
)
from meshnet_node.native_protocol.generated import ( # noqa: E402
shard_runtime_pb2 as pb,
shard_runtime_pb2_grpc as pb_grpc,
)
def _free_port() -> int:
s = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
s.bind(("127.0.0.1", 0))
port = s.getsockname()[1]
s.close()
return port
def _start_server(listen_addr: str, capture_path: str) -> subprocess.Popen:
env = dict(os.environ)
env["PYTHONPATH"] = _PYTHONPATH
env["MESHNET_SHARD_LISTEN_ADDR"] = listen_addr
env["MESHNET_WIRE_CAPTURE_PATH"] = capture_path
proc = subprocess.Popen(
[sys.executable, "-m", "meshnet_node.shard_runtime_server"],
cwd=REPO_ROOT,
env=env,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
)
# Wait until the server reports it is listening (real readiness signal).
deadline = time.time() + 30.0
while time.time() < deadline:
line = proc.stdout.readline()
if not line:
if proc.poll() is not None:
out, _ = proc.communicate()
raise RuntimeError(f"server exited early:\n{out}")
continue
if "listening on" in line:
return proc
raise RuntimeError("server did not start listening in time")
def _build_activation_chunk(
work_id: str,
payload: bytes,
step: int,
*,
route_session_id: str = "route-session-1",
route_epoch: int = 7,
deadline_unix_nanos: int = 0,
bad_checksum: bool = False,
bad_offset: bool = False,
) -> pb.SessionRequest:
checksum_value = b"\x00\x00\x00\x00" if bad_checksum else zlib_crc32c(payload)
fragment_offset = 5 if bad_offset else 0
tensor = pb.NamedTensor(
name="hidden_states",
shape=[1, 1, 4096],
dtype=pb.DTYPE_BFLOAT16,
byte_order=pb.BYTE_ORDER_LITTLE_ENDIAN,
total_bytes=len(payload),
compression=pb.COMPRESSION_NONE,
checksum=pb.Checksum(
algorithm=pb.CHECKSUM_ALGORITHM_CRC32C,
value=checksum_value,
),
fragments=[
pb.TensorFragment(
fragment_index=0, fragment_count=1, byte_offset=fragment_offset, payload=payload
)
],
)
bundle = pb.TensorBundle(
bundle_version=1,
tensors=[tensor],
architecture=pb.ARCHITECTURE_TYPE_DENSE,
boundary_point="pre_tail_residual",
)
envelope = pb.Envelope(
schema_version=pb.SCHEMA_VERSION_1,
work_id=work_id,
route_session_id=route_session_id,
route_epoch=route_epoch,
idempotency_step=step,
phase=pb.PHASE_PREFILL,
position=pb.PositionSpan(first_position=0, token_count=1),
deadline_unix_nanos=deadline_unix_nanos,
)
return pb.SessionRequest(chunk=pb.ActivationChunk(envelope=envelope, bundle=bundle))
def _build_decode_step(work_id: str, payload: bytes, step: int, position: int) -> pb.SessionRequest:
tensor = pb.NamedTensor(
name="hidden_states",
shape=[1, 1, 4096],
dtype=pb.DTYPE_BFLOAT16,
byte_order=pb.BYTE_ORDER_LITTLE_ENDIAN,
total_bytes=len(payload),
compression=pb.COMPRESSION_NONE,
checksum=pb.Checksum(
algorithm=pb.CHECKSUM_ALGORITHM_CRC32C,
value=zlib_crc32c(payload),
),
fragments=[pb.TensorFragment(fragment_index=0, fragment_count=1, byte_offset=0, payload=payload)],
)
bundle = pb.TensorBundle(
bundle_version=1,
tensors=[tensor],
architecture=pb.ARCHITECTURE_TYPE_DENSE,
boundary_point="pre_tail_residual",
)
step_msg = pb.DecodeStep(
idempotency_step=step,
position=position,
expected_past_len=position,
work_id=work_id,
deadline_unix_nanos=0,
bundle=bundle,
)
return pb.SessionRequest(decode=step_msg)
def _build_open(
*, route_session_id: str = "route-session-1", route_epoch: int = 7, credits_granted: int = 16
) -> pb.SessionRequest:
return pb.SessionRequest(
open=pb.SessionOpen(
schema_version=pb.SCHEMA_VERSION_1,
route_session_id=route_session_id,
route_epoch=route_epoch,
fingerprint=pb.Fingerprint(
model_artifact_digest="sha256:native-test-artifact",
runtime_recipe_digest="sha256:native-test-recipe",
recipe_id="native-test",
recipe_version="1",
catalogue_version="1",
),
shard_range=pb.ShardRange(start_layer=0, end_layer=32, effective_start_layer=0),
proposed_flow_control=pb.FlowControl(
credits_granted=credits_granted,
max_inflight_chunks=16,
max_chunk_bytes=4 * 1024 * 1024,
max_prefill_chunk_tokens=512,
),
accepted_compression=[pb.COMPRESSION_NONE],
)
)
def _build_release() -> pb.SessionRequest:
return pb.SessionRequest(
release=pb.ReleaseSignal(route_session_id="route-session-1", route_epoch=7, work_id="work-final")
)
def _build_cancel(*, route_session_id: str = "route-session-1", work_id: str = "", reason: str = "test cancel") -> pb.SessionRequest:
return pb.SessionRequest(
cancel=pb.CancelSignal(
route_session_id=route_session_id, route_epoch=7, work_id=work_id, reason=reason
)
)
@contextlib.contextmanager
def _running_server():
"""Spawn a fresh server subprocess + real socket channel for one test."""
port = _free_port()
listen_addr = f"127.0.0.1:{port}"
capture_path = os.path.join(REPO_ROOT, "tests", f".dgr024_wire_capture_{port}.jsonl")
proc = _start_server(listen_addr, capture_path)
channel = None
try:
channel = grpc.insecure_channel(listen_addr)
grpc.channel_ready_future(channel).result(timeout=15.0)
yield channel
finally:
if channel is not None:
channel.close()
proc.terminate()
try:
proc.wait(timeout=10)
except subprocess.TimeoutExpired:
proc.kill()
if os.path.exists(capture_path):
os.remove(capture_path)
def _session_call(channel, requests):
"""Drive a real Session stream, returning parsed SessionResponse messages."""
call = channel.stream_stream(
"/meshnet.shard.v1.ShardRuntime/Session",
request_serializer=lambda m: m.SerializeToString(),
response_deserializer=pb.SessionResponse.FromString,
)
return list(call(iter(requests)))
def zlib_crc32c(payload: bytes) -> bytes:
return zlib.crc32(payload).to_bytes(4, "big")
def _open_session(channel, requests, *, identity_send: bool) -> tuple[list[bytes], list[bytes]]:
"""Open a real Session over the socket.
Returns (client_request_bytes, raw_response_bytes). When ``identity_send``
is True the request objects are already serialized bytes (the opaque relay
path); otherwise real ``SessionRequest`` objects are sent with the generated
serializer (the direct path). Responses are always captured as raw wire bytes.
"""
if identity_send:
request_serializer = lambda b: b # raw captured bytes, no reinterpretation
req_iter = iter(requests)
else:
request_serializer = lambda m: m.SerializeToString()
req_iter = iter(requests)
sent = [m.SerializeToString() for m in requests]
call = channel.stream_stream(
"/meshnet.shard.v1.ShardRuntime/Session",
request_serializer=request_serializer,
response_deserializer=lambda b: b, # capture exact wire bytes
)
responses = list(call(req_iter))
if identity_send:
return list(requests), responses
return sent, responses
def test_native_protocol_not_drifted():
"""(a) The committed generated stubs match shard_runtime.proto exactly."""
env = dict(os.environ)
env["PYTHONPATH"] = _PYTHONPATH
result = subprocess.run(
[sys.executable, "scripts/generate_native_protocol.py", "--check"],
cwd=REPO_ROOT,
env=env,
capture_output=True,
text=True,
)
assert result.returncode == 0, (
f"generate_native_protocol.py --check failed:\n{result.stdout}\n{result.stderr}"
)
assert "up to date" in result.stdout
def test_shard_runtime_real_subprocess_harness():
"""(b-d) Real subprocess + socket + generated client; direct vs opaque relay byte equality."""
port = _free_port()
listen_addr = f"127.0.0.1:{port}"
capture_path = os.path.join(REPO_ROOT, "tests", ".dgr024_wire_capture.jsonl")
proc = _start_server(listen_addr, capture_path)
try:
# Real channel to a real listening socket.
channel = grpc.insecure_channel(listen_addr)
grpc.channel_ready_future(channel).result(timeout=15.0)
stub = pb_grpc.ShardRuntimeStub(channel)
# Unary RPCs over the socket.
cap = stub.GetCapability(pb.CapabilityRequest(schema_version=pb.SCHEMA_VERSION_1))
assert cap.schema_version == pb.SCHEMA_VERSION_1
assert cap.validated is True
health = stub.Health(pb.HealthRequest(schema_version=pb.SCHEMA_VERSION_1))
assert health.state == pb.SERVING_STATE_SERVING
# Build real requests with real byte payloads.
payload_chunk = b"REAL_ACTIVATION_BYTES_prefill_chunk_A9F2"
payload_decode = b"REAL_ACTIVATION_BYTES_decode_step_B7C1"
open_req = _build_open()
chunk_req = _build_activation_chunk("work-1", payload_chunk, step=1)
decode_req = _build_decode_step("work-2", payload_decode, step=2, position=1)
release_req = _build_release()
direct_requests = [open_req, chunk_req, decode_req, release_req]
# ---- DIRECT localhost hop ----
direct_req_bytes, direct_resp_bytes = _open_session(
channel, direct_requests, identity_send=False
)
# ---- OPAQUE RELAY: re-carry the EXACT captured request bytes ----
# The relay forwards raw captured frames; it must not reinterpret them.
relay_req_bytes, relay_resp_bytes = _open_session(
channel, list(direct_req_bytes), identity_send=True
)
# (c) The relay carried the exact same request bytes and the server's
# responses are byte-identical for both paths.
assert relay_req_bytes == direct_req_bytes, (
"opaque relay must forward the exact captured request bytes"
)
assert len(relay_resp_bytes) == len(direct_resp_bytes) == 4, (
f"expected 4 responses (accepted,chunk,chunk,status), "
f"got direct={len(direct_resp_bytes)} relay={len(relay_resp_bytes)}"
)
for i, (d, r) in enumerate(zip(direct_resp_bytes, relay_resp_bytes)):
assert d == r, (
f"server response #{i} differs between direct and opaque relay:\n"
f" direct ={d.hex()}\n relay ={r.hex()}"
)
# (e) The server echoed the REAL payload bytes, not a synthesized response.
# Parse the direct responses; index 1 = ActivationChunk echo, 2 = DecodeStep echo.
echoed_chunk = pb.SessionResponse.FromString(direct_resp_bytes[1]).chunk
echoed_decode = pb.SessionResponse.FromString(direct_resp_bytes[2]).chunk
assert echoed_chunk.bundle.tensors[0].fragments[0].payload == payload_chunk, (
"ActivationChunk payload was not echoed faithfully"
)
assert echoed_decode.bundle.tensors[0].fragments[0].payload == payload_decode, (
"DecodeStep payload was not echoed faithfully"
)
# Checksums over the echoed bytes must match the originals.
assert echoed_chunk.bundle.tensors[0].checksum.value == zlib_crc32c(payload_chunk)
assert echoed_decode.bundle.tensors[0].checksum.value == zlib_crc32c(payload_decode)
# Handshake: first response is a SessionAccepted.
accepted = pb.SessionResponse.FromString(direct_resp_bytes[0]).accepted
assert accepted.route_session_id == "route-session-1"
assert accepted.schema_version == pb.SCHEMA_VERSION_1
# Final response is a terminal status.
status = pb.SessionResponse.FromString(direct_resp_bytes[3]).status
assert status.terminal is True
# Out-of-process capture cross-check: the server's WireCapture file must
# record the same request bytes the client sent (proof the frames
# traversed the wire and were captured by the real server, not the test).
with open(capture_path, "r", encoding="utf-8") as fh:
lines = [ln for ln in fh.read().splitlines() if ln.strip()]
assert len(lines) >= 2, f"expected >=2 capture lines (direct+relay), got {len(lines)}"
direct_capture = json.loads(lines[-2])
relay_capture = json.loads(lines[-1])
assert [bytes.fromhex(h) for h in direct_capture["requests"]] == direct_req_bytes
assert [bytes.fromhex(h) for h in relay_capture["requests"]] == relay_req_bytes
assert direct_capture["requests"] == relay_capture["requests"]
# Wire-frame hashes: the server-persisted SHA-256 over the exact
# captured request/response frame bytes must match independently
# computed hashes over what the client sent/received, and must be
# identical between the direct hop and the opaque relay carry.
expected_req_sha256 = hashlib.sha256(b"".join(direct_req_bytes)).hexdigest()
expected_resp_sha256 = hashlib.sha256(b"".join(direct_resp_bytes)).hexdigest()
assert direct_capture["requests_sha256"] == expected_req_sha256
assert direct_capture["responses_sha256"] == expected_resp_sha256
assert relay_capture["requests_sha256"] == expected_req_sha256
assert relay_capture["responses_sha256"] == expected_resp_sha256
print(
f"wire-frame sha256: requests={expected_req_sha256} responses={expected_resp_sha256}"
)
channel.close()
finally:
proc.terminate()
try:
proc.wait(timeout=10)
except subprocess.TimeoutExpired:
proc.kill()
if os.path.exists(capture_path):
os.remove(capture_path)
# ---------------------------------------------------------------------------
# Negative paths (DGR-024 acceptance criterion 2 & 4): flow-control, deadlines,
# malformed input, checksum failure, duplicates, stale epochs, cancel — each
# exercised over the same real subprocess/socket/generated-stub harness above,
# never an in-memory fake.
# ---------------------------------------------------------------------------
def test_stale_route_epoch_is_rejected():
with _running_server() as channel:
responses = _session_call(
channel,
[
_build_open(route_epoch=7),
_build_activation_chunk("work-stale", b"payload", step=1, route_epoch=5),
],
)
assert responses[0].WhichOneof("kind") == "accepted"
status = responses[1].status
assert status.error.code == pb.ERROR_CODE_EPOCH_STALE
assert status.terminal is False
def test_expired_deadline_is_rejected():
with _running_server() as channel:
responses = _session_call(
channel,
[
_build_open(),
_build_activation_chunk(
"work-late", b"payload", step=1, deadline_unix_nanos=1
),
],
)
status = responses[1].status
assert status.error.code == pb.ERROR_CODE_DEADLINE_EXCEEDED
assert status.terminal is False
def test_malformed_fragment_tiling_is_rejected():
with _running_server() as channel:
responses = _session_call(
channel,
[
_build_open(),
_build_activation_chunk("work-gap", b"payload", step=1, bad_offset=True),
],
)
status = responses[1].status
assert status.error.code == pb.ERROR_CODE_PAYLOAD_CORRUPT
assert "tile" in status.error.detail
def test_checksum_failure_is_rejected():
with _running_server() as channel:
responses = _session_call(
channel,
[
_build_open(),
_build_activation_chunk("work-corrupt", b"payload", step=1, bad_checksum=True),
],
)
status = responses[1].status
assert status.error.code == pb.ERROR_CODE_PAYLOAD_CORRUPT
assert "checksum" in status.error.detail
def test_duplicate_idempotency_step_is_acked_not_reapplied():
with _running_server() as channel:
chunk = _build_activation_chunk("work-dup", b"payload", step=1)
responses = _session_call(channel, [_build_open(), chunk, chunk])
first = responses[1]
second = responses[2]
assert first.WhichOneof("kind") == "chunk"
assert second.WhichOneof("kind") == "ack"
assert second.ack.duplicate is True
assert second.ack.work_id == "work-dup"
assert second.ack.idempotency_step == 1
def test_flow_control_violation_and_topup():
with _running_server() as channel:
responses = _session_call(
channel,
[
_build_open(credits_granted=1),
_build_activation_chunk("work-a", b"payload-a", step=1),
_build_activation_chunk("work-b", b"payload-b", step=2),
pb.SessionRequest(flow_control=pb.FlowControl(credits_granted=5)),
_build_activation_chunk("work-c", b"payload-c", step=3),
],
)
assert responses[1].WhichOneof("kind") == "chunk", "first chunk should consume the one granted credit"
violation = responses[2].status
assert violation.error.code == pb.ERROR_CODE_FLOW_CONTROL_VIOLATION
assert violation.error.retryable is True
assert responses[3].WhichOneof("kind") == "flow_control"
assert responses[3].flow_control.credits_granted >= 5
assert responses[4].WhichOneof("kind") == "chunk", "chunk after top-up should succeed"
def test_in_band_cancel_of_single_work_item_does_not_end_stream():
with _running_server() as channel:
responses = _session_call(
channel,
[
_build_open(),
_build_cancel(work_id="work-x"),
_build_activation_chunk("work-x", b"payload", step=1),
_build_activation_chunk("work-y", b"payload", step=2),
_build_release(),
],
)
cancel_ack = responses[1].status
assert cancel_ack.error.code == pb.ERROR_CODE_CANCELLED
assert cancel_ack.terminal is False
cancelled_work_status = responses[2].status
assert cancelled_work_status.error.code == pb.ERROR_CODE_CANCELLED
still_alive = responses[3]
assert still_alive.WhichOneof("kind") == "chunk", "an unrelated work_id must still be served"
assert responses[4].status.terminal is True
def test_in_band_cancel_of_whole_session_is_terminal():
with _running_server() as channel:
responses = _session_call(
channel,
[
_build_open(),
_build_cancel(work_id=""),
],
)
status = responses[1].status
assert status.error.code == pb.ERROR_CODE_CANCELLED
assert status.terminal is True
def test_out_of_band_cancel_rpc_fails_closed_even_before_open():
"""The unary Cancel RPC can race ahead of SessionOpen; the eventual Session
for that route_session_id/work_id must still fail closed (ADR-0020)."""
with _running_server() as channel:
stub = pb_grpc.ShardRuntimeStub(channel)
cancel_response = stub.Cancel(
pb.CancelRequest(
schema_version=pb.SCHEMA_VERSION_1,
route_session_id="route-session-precancel",
route_epoch=1,
work_id="work-precancelled",
reason="operator abort",
)
)
assert cancel_response.cancelled_work_items == 1
responses = _session_call(
channel,
[
_build_open(route_session_id="route-session-precancel"),
_build_activation_chunk(
"work-precancelled",
b"payload",
step=1,
route_session_id="route-session-precancel",
),
],
)
status = responses[1].status
assert status.error.code == pb.ERROR_CODE_CANCELLED