"""DGR-031 ``ShardEngine`` contract tests. ``_ReferenceEngine`` below is a minimal, in-memory ``ShardEngine`` that exists only to prove :func:`assert_shard_engine_contract` is non-vacuous and to pin the abstract contract's own validation rules. It is deliberately not the DGR-032 deterministic fixture (delay/memory-pressure/malformed/crash injection, full session/epoch modeling for the fake worker) — that is a separate, larger story. DGR-032 and DGR-037 are expected to import ``assert_shard_engine_contract`` from ``tests/shard_engine_contract.py`` against their own engines. """ from __future__ import annotations import hashlib import pytest from meshnet_node.shard_engine import ( ArchitectureAuxStateHook, BoundaryBundle, DecodeRequest, EngineCapabilities, EngineTensor, HealthResult, LoadRequest, LoadResult, MetricsResult, MtpHook, PrefillRequest, ShardEngine, StepResult, TokenOutput, ) from meshnet_node.shard_lifecycle import CacheResult, StatusCode, StructuredStatus from shard_engine_contract import assert_shard_engine_contract class _ReferenceEngine(ShardEngine): """Minimal in-memory engine used only to exercise the shared contract.""" def __init__(self) -> None: self._loaded: LoadRequest | None = None self._sessions: dict[str, dict] = {} self._cancelled_total = 0 def load(self, request: LoadRequest) -> LoadResult: self._loaded = request return LoadResult( status=StructuredStatus(StatusCode.OK, "loaded"), effective_start=request.shard_start, architecture="dense", ) def capabilities(self) -> EngineCapabilities: if self._loaded is None: return EngineCapabilities(status=StructuredStatus(StatusCode.FAILED_PRECONDITION, "not loaded")) request = self._loaded return EngineCapabilities( status=StructuredStatus(StatusCode.OK, "ready"), shard_start=request.shard_start, shard_end=request.shard_end, effective_start=request.shard_start, total_layers=request.total_layers, architecture="dense", max_concurrent_sessions=8, max_context_tokens=131072, supports_mtp=False, ) def prefill(self, request: PrefillRequest) -> StepResult: if self._loaded is None: return StepResult(status=StructuredStatus(StatusCode.FAILED_PRECONDITION, "engine not loaded")) self._sessions[request.session_id] = {"epoch": request.route_epoch, "cancelled": False} output = self._transform(self._seed_bytes(request.token_ids, request.input), request.idempotency_step) return StepResult(status=StructuredStatus(StatusCode.OK, "prefilled"), cache_result=CacheResult.STORED, output=output) def decode(self, request: DecodeRequest) -> StepResult: session = self._sessions.get(request.session_id) if session is None: return StepResult( status=StructuredStatus(StatusCode.NOT_FOUND, "no cached session state"), cache_result=CacheResult.MISS, ) if session["cancelled"]: return StepResult(status=StructuredStatus(StatusCode.CANCELLED, "session cancelled")) if request.route_epoch < session["epoch"]: return StepResult(status=StructuredStatus(StatusCode.FAILED_PRECONDITION, "stale route epoch")) session["epoch"] = request.route_epoch token_ids = (request.token_id,) if request.token_id is not None else None output = self._transform(self._seed_bytes(token_ids, request.input), request.idempotency_step) return StepResult(status=StructuredStatus(StatusCode.OK, "decoded"), cache_result=CacheResult.HIT, output=output) def cancel(self, session_id: str, *, work_id: str = "", reason: str = "") -> StructuredStatus: session = self._sessions.setdefault(session_id, {"epoch": 0, "cancelled": False}) if not session["cancelled"]: self._cancelled_total += 1 session["cancelled"] = True return StructuredStatus(StatusCode.CANCELLED, reason or "cancelled") def release(self, session_id: str) -> StructuredStatus: self._sessions.pop(session_id, None) return StructuredStatus(StatusCode.OK, "released") def health(self) -> HealthResult: return HealthResult( status=StructuredStatus(StatusCode.OK, "ok"), serving=self._loaded is not None, state="SERVING" if self._loaded is not None else "NOT_LOADED", active_sessions=len(self._sessions), ) def metrics(self) -> MetricsResult: return MetricsResult( status=StructuredStatus(StatusCode.OK, "ok"), active_sessions=len(self._sessions), cancelled_sessions=self._cancelled_total, ) @staticmethod def _seed_bytes(token_ids, bundle: BoundaryBundle | None) -> bytes: if token_ids: return b"".join(int(t).to_bytes(4, "big") for t in token_ids) if bundle is not None: return b"".join(tensor.data for tensor in bundle.tensors) return b"" def _transform(self, seed: bytes, idempotency_step: int) -> BoundaryBundle | TokenOutput: digest = hashlib.sha256(seed + idempotency_step.to_bytes(4, "big")).digest() assert self._loaded is not None if self._loaded.shard_end >= self._loaded.total_layers - 1: token_id = int.from_bytes(digest[:4], "big") % 50_000 return TokenOutput(token_id=token_id) tensor = EngineTensor(name="hidden_states", shape=(1, max(len(seed) // 4, 1)), dtype="bfloat16", data=digest) return BoundaryBundle(tensors=(tensor,), architecture="dense", boundary_point="pre_tail_residual") def test_reference_engine_obeys_the_shared_shard_engine_contract(): assert_shard_engine_contract(_ReferenceEngine) def test_shard_engine_is_abstract_and_cannot_be_instantiated_directly(): with pytest.raises(TypeError): ShardEngine() # type: ignore[abstract] def test_engine_tensor_rejects_empty_name_shape_or_dtype(): with pytest.raises(ValueError): EngineTensor(name="", shape=(1,), dtype="bfloat16", data=b"x") with pytest.raises(ValueError): EngineTensor(name="t", shape=(), dtype="bfloat16", data=b"x") with pytest.raises(ValueError): EngineTensor(name="t", shape=(0,), dtype="bfloat16", data=b"x") with pytest.raises(ValueError): EngineTensor(name="t", shape=(1,), dtype="", data=b"x") def test_boundary_bundle_requires_at_least_one_tensor(): with pytest.raises(ValueError): BoundaryBundle(tensors=(), architecture="dense", boundary_point="pre_tail_residual") def test_boundary_bundle_tensor_lookup_by_name(): tensor = EngineTensor(name="hidden_states", shape=(1, 1), dtype="bfloat16", data=b"\x00\x00") bundle = BoundaryBundle(tensors=(tensor,), architecture="dense", boundary_point="pre_tail_residual") assert bundle.tensor("hidden_states") is tensor with pytest.raises(KeyError): bundle.tensor("router_logits") def test_token_output_rejects_negative_token_id(): with pytest.raises(ValueError): TokenOutput(token_id=-1) def test_mtp_hook_is_reserved_and_refuses_to_enable(): MtpHook() # disabled is fine with pytest.raises(ValueError): MtpHook(enabled=True) with pytest.raises(ValueError): MtpHook(draft_token_count=-1) def test_architecture_aux_state_hook_carries_opaque_shard_local_state(): hook = ArchitectureAuxStateHook(kind="csa", state={"window": 128}) assert hook.kind == "csa" assert hook.state == {"window": 128} def test_prefill_and_decode_requests_require_exactly_one_input_kind(): with pytest.raises(ValueError): PrefillRequest(session_id="s", route_epoch=0, position=0, idempotency_step=0) with pytest.raises(ValueError): PrefillRequest( session_id="s", route_epoch=0, position=0, idempotency_step=0, token_ids=(1,), input=BoundaryBundle( tensors=(EngineTensor(name="hidden_states", shape=(1,), dtype="bfloat16", data=b"x"),), architecture="dense", boundary_point="pre_tail_residual", ), ) with pytest.raises(ValueError): DecodeRequest(session_id="s", route_epoch=0, position=0, idempotency_step=0) def test_load_request_validates_shard_range_against_total_layers(): LoadRequest(artifact_path="a", shard_start=0, shard_end=3, total_layers=4) with pytest.raises(ValueError): LoadRequest(artifact_path="a", shard_start=0, shard_end=4, total_layers=4) with pytest.raises(ValueError): LoadRequest(artifact_path="", shard_start=0, shard_end=0, total_layers=1) with pytest.raises(ValueError): LoadRequest(artifact_path="a", shard_start=3, shard_end=1, total_layers=4) def test_step_result_requires_an_output_when_status_is_ok(): with pytest.raises(ValueError): StepResult(status=StructuredStatus(StatusCode.OK, "ok"), output=None) # A non-OK status is allowed to carry no output. StepResult(status=StructuredStatus(StatusCode.NOT_FOUND, "missing"), output=None) def test_shard_engine_module_imports_no_native_or_grpc_or_wire_abi_types(): import meshnet_node.shard_engine as shard_engine_module # The boundary module must not *import* anything that would let a # ggml_tensor, llama context/scheduler handle, ctypes native handle, or a # generated-protobuf (ABI) message leak into a project-owned dataclass # field. Checking bound globals (not docstring prose) proves this # structurally rather than by convention. forbidden_modules = {"ctypes", "grpc", "meshnet_node.native_protocol"} for name, value in vars(shard_engine_module).items(): module_name = getattr(value, "__name__", None) assert module_name not in forbidden_modules, ( f"shard_engine.{name} binds forbidden module {module_name!r}" )