Files
neuron-tai/packages/node/meshnet_node/shard_engine.py

373 lines
13 KiB
Python

"""The project-owned ``ShardEngine`` contract (DGR-031).
A worker process (the gRPC surface in ``shard_runtime_server.py``, or any
future transport) never talks to llama.cpp directly. It talks to a
``ShardEngine``. This module is the *only* place that boundary is defined, and
every operation on it is built from project-owned dataclasses and plain
Python values (``str``, ``int``, ``bytes``, ``Mapping``) — never a
``ggml_tensor``, a llama context/scheduler handle, or a generated-protobuf
(ABI) message. A fake fixture engine (DGR-032) and a real llama.cpp-backed
engine (DGR-037) are both, structurally, nothing more than subclasses of
:class:`ShardEngine`; the worker code that calls them does not change when one
replaces the other.
This is deliberately a fourth, distinct layer from the three that already
exist:
- ``native_protocol`` — the generated gRPC/Protobuf wire ABI (DGR-021/024).
- ``protocol.ActivationEnvelope`` — the versioned wire envelope for activation
traffic between shard *hops* over the network (DGR-021).
- ``shard_lifecycle`` — the versioned RPC/session lifecycle contract a
generated gRPC binding consumes (DGR-022).
``ShardEngine`` sits *inside* one worker process, below all three: it is the
seam between "the code that speaks Meshnet's wire protocol" and "the code
that actually runs model layers." It reuses :class:`~meshnet_node.shard_lifecycle.StructuredStatus`,
:class:`~meshnet_node.shard_lifecycle.StatusCode`, :class:`~meshnet_node.shard_lifecycle.CacheExpectation`,
and :class:`~meshnet_node.shard_lifecycle.CacheResult` rather than inventing a
parallel status vocabulary, since those are already project-owned and
version-stable.
"""
from __future__ import annotations
import abc
from dataclasses import dataclass, field
from typing import Any, Mapping
from .shard_lifecycle import (
CacheExpectation,
CacheResult,
StatusCode,
StructuredStatus,
)
__all__ = [
"EngineError",
"EngineTensor",
"BoundaryBundle",
"TokenOutput",
"MtpHook",
"ArchitectureAuxStateHook",
"LoadRequest",
"LoadResult",
"EngineCapabilities",
"PrefillRequest",
"DecodeRequest",
"StepResult",
"HealthResult",
"MetricsResult",
"ShardEngine",
]
class EngineError(RuntimeError):
"""An engine-boundary failure represented by a structured status.
Mirrors :class:`~meshnet_node.shard_lifecycle.LifecycleContractError`:
callers pattern-match on ``error.status.code`` rather than on exception
subclasses, so a fake and a real engine can fail the exact same way for
the exact same reason.
"""
def __init__(self, status: StructuredStatus) -> None:
self.status = status
super().__init__(status.message)
@dataclass(frozen=True)
class EngineTensor:
"""One named tensor crossing the engine boundary.
Intentionally not a ``ggml_tensor`` or a framework tensor object: ``data``
is plain owned bytes, ``shape``/``dtype`` are plain metadata. An
implementation constructs this from whatever internal representation it
uses (a ``torch.Tensor``, a llama.cpp buffer, a synthetic fixture array)
without leaking that representation across the boundary.
"""
name: str
shape: tuple[int, ...]
dtype: str
data: bytes
def __post_init__(self) -> None:
if not self.name:
raise ValueError("engine tensor requires a name")
if not self.shape or any(dim <= 0 for dim in self.shape):
raise ValueError("engine tensor shape must be a non-empty tuple of positive ints")
if not self.dtype:
raise ValueError("engine tensor requires a dtype")
@dataclass(frozen=True)
class BoundaryBundle:
"""A named-tensor activation crossing a shard boundary (head/middle/tail-in).
``token_id_sideband`` carries token IDs alongside the activation only
where the architecture boundary requires them (V4's first three
hash-routed MoE layers); it is ``None`` everywhere else. Per-shard hot
KV/recurrent/CSA/HCA/SWA/indexer/compressor state never appears here — it
stays local to a shard via :class:`ArchitectureAuxStateHook` and is never
part of what crosses the wire.
"""
tensors: tuple[EngineTensor, ...]
architecture: str
boundary_point: str
token_id_sideband: tuple[int, ...] | None = None
def __post_init__(self) -> None:
if not self.tensors:
raise ValueError("boundary bundle requires at least one tensor")
if not self.architecture:
raise ValueError("boundary bundle requires an architecture name")
if not self.boundary_point:
raise ValueError("boundary bundle requires a boundary point name")
def tensor(self, name: str) -> EngineTensor:
for tensor in self.tensors:
if tensor.name == name:
return tensor
raise KeyError(name)
@dataclass(frozen=True)
class TokenOutput:
"""A tail shard's sampled decode result.
Never a raw logits tensor: the engine boundary only ever hands back the
already-sampled token (mirroring
:meth:`meshnet_node.architecture_boundary.TailOutput.sampled_token`, which
likewise refuses anything but a sampled token id).
"""
token_id: int
text: str | None = None
def __post_init__(self) -> None:
if self.token_id < 0:
raise ValueError("sampled token id must be non-negative")
@dataclass(frozen=True)
class MtpHook:
"""Reserved multi-token-prediction hook — typed, but refused when enabled.
RALPH-CONTEXT is explicit that "MTP is reserved and off for alpha; its
ownership contract, implementation, and benchmark are required before
beta" (DGR-065/DGR-066). Reserving the shape now means DGR-037's real
engine and DGR-051's V4 adapter do not have to change this dataclass's
field layout later; they only flip ``enabled`` once DGR-066 lands.
"""
enabled: bool = False
draft_token_count: int = 0
aux_state: Mapping[str, Any] | None = None
def __post_init__(self) -> None:
if self.enabled:
raise ValueError(
"MTP is reserved and must remain disabled before DGR-066; "
"this hook exists to fix its shape, not to enable it"
)
if self.draft_token_count < 0:
raise ValueError("draft_token_count must be non-negative")
@dataclass(frozen=True)
class ArchitectureAuxStateHook:
"""Reserved per-shard architecture auxiliary-state hook.
Covers V4's CSA/HCA/SWA/indexer/compressor state and any other
architecture-local state a future adapter needs. RALPH-CONTEXT locks this
as shard-local, keyed by route session/epoch, and explicitly never carried
over the WAN seam — so this hook has no wire encoding of its own and must
never be embedded inside a :class:`BoundaryBundle`.
"""
kind: str = ""
state: Mapping[str, Any] | None = None
@dataclass(frozen=True)
class LoadRequest:
"""One exact artifact/recipe/range identity for a worker to load."""
artifact_path: str
shard_start: int
shard_end: int
total_layers: int
recipe: Mapping[str, Any] = field(default_factory=dict)
def __post_init__(self) -> None:
if not self.artifact_path:
raise ValueError("load request requires an artifact path")
if self.shard_start < 0 or self.shard_end < self.shard_start:
raise ValueError("shard_start must be <= shard_end and non-negative")
if self.total_layers <= self.shard_end:
raise ValueError("total_layers must exceed shard_end (shard_end is inclusive)")
@dataclass(frozen=True)
class LoadResult:
status: StructuredStatus
effective_start: int = 0
architecture: str = ""
@dataclass(frozen=True)
class EngineCapabilities:
status: StructuredStatus
shard_start: int = 0
shard_end: int = 0
effective_start: int = 0
total_layers: int = 0
architecture: str = ""
max_concurrent_sessions: int = 0
max_context_tokens: int = 0
supports_mtp: bool = False
@property
def is_head(self) -> bool:
return self.shard_start == 0
@property
def is_tail(self) -> bool:
return self.shard_end >= self.total_layers - 1
@dataclass(frozen=True)
class PrefillRequest:
"""A prefill step. Exactly one of ``token_ids`` (head) or ``input`` (middle/tail) is set."""
session_id: str
route_epoch: int
position: int
idempotency_step: int
token_ids: tuple[int, ...] | None = None
input: BoundaryBundle | None = None
cache_expectation: CacheExpectation = CacheExpectation.NONE
mtp: MtpHook = field(default_factory=MtpHook)
architecture_aux_state: ArchitectureAuxStateHook | None = None
def __post_init__(self) -> None:
_require_exactly_one_input(self.token_ids, self.input)
if not self.session_id:
raise ValueError("prefill request requires a session id")
if self.route_epoch < 0 or self.position < 0 or self.idempotency_step < 0:
raise ValueError("route_epoch, position, and idempotency_step must be non-negative")
@dataclass(frozen=True)
class DecodeRequest:
"""A decode step. Exactly one of ``token_id`` (head) or ``input`` (middle/tail) is set."""
session_id: str
route_epoch: int
position: int
idempotency_step: int
token_id: int | None = None
input: BoundaryBundle | None = None
mtp: MtpHook = field(default_factory=MtpHook)
architecture_aux_state: ArchitectureAuxStateHook | None = None
def __post_init__(self) -> None:
_require_exactly_one_input(
None if self.token_id is None else (self.token_id,), self.input
)
if not self.session_id:
raise ValueError("decode request requires a session id")
if self.route_epoch < 0 or self.position < 0 or self.idempotency_step < 0:
raise ValueError("route_epoch, position, and idempotency_step must be non-negative")
def _require_exactly_one_input(
token_ids: tuple[int, ...] | None, bundle: BoundaryBundle | None
) -> None:
if (token_ids is None) == (bundle is None):
raise ValueError("exactly one of token ids or a boundary bundle must be set")
@dataclass(frozen=True)
class StepResult:
"""The result of a prefill or decode step.
``output`` is a :class:`BoundaryBundle` for a head/middle shard handing an
activation to the next hop, or a :class:`TokenOutput` for a tail shard
that sampled a token. It is ``None`` only when ``status.code`` is not
``OK``.
"""
status: StructuredStatus
cache_result: CacheResult = CacheResult.NOT_REQUESTED
output: BoundaryBundle | TokenOutput | None = None
def __post_init__(self) -> None:
if self.status.code is StatusCode.OK and self.output is None:
raise ValueError("a successful step result must carry an output")
@dataclass(frozen=True)
class HealthResult:
status: StructuredStatus
serving: bool = False
state: str = "UNKNOWN"
active_sessions: int = 0
@dataclass(frozen=True)
class MetricsResult:
status: StructuredStatus
active_sessions: int = 0
queued_frames: int = 0
inflight_bytes: int = 0
kv_entries: int = 0
generated_tokens: int = 0
cancelled_sessions: int = 0
class ShardEngine(abc.ABC):
"""The contract every shard execution engine (fake or real) must implement.
Every method returns a project-owned result carrying a
:class:`~meshnet_node.shard_lifecycle.StructuredStatus` rather than
raising for expected, protocol-visible outcomes (a cache miss, a stale
epoch, an unknown session); an :class:`EngineError` is reserved for
genuine programming errors at the call site (malformed request objects),
which the request dataclasses' own ``__post_init__`` validation already
catches before an implementation ever sees them.
"""
@abc.abstractmethod
def load(self, request: LoadRequest) -> LoadResult:
"""Load one exact artifact/recipe/range identity. Idempotent per engine instance."""
@abc.abstractmethod
def capabilities(self) -> EngineCapabilities:
"""Report this engine's authoritative range and limits after ``load``."""
@abc.abstractmethod
def prefill(self, request: PrefillRequest) -> StepResult:
"""Run one prefill step for a session."""
@abc.abstractmethod
def decode(self, request: DecodeRequest) -> StepResult:
"""Run one decode step for a session."""
@abc.abstractmethod
def cancel(self, session_id: str, *, work_id: str = "", reason: str = "") -> StructuredStatus:
"""Cancel a session (or one work item within it) in flight."""
@abc.abstractmethod
def release(self, session_id: str) -> StructuredStatus:
"""Release a session's held state. Idempotent."""
@abc.abstractmethod
def health(self) -> HealthResult:
"""Report liveness/serving state. Must never raise."""
@abc.abstractmethod
def metrics(self) -> MetricsResult:
"""Report point-in-time operational counters. Must never raise."""