feat: add activation stream envelope
This commit is contained in:
376
packages/node/meshnet_node/protocol.py
Normal file
376
packages/node/meshnet_node/protocol.py
Normal file
@@ -0,0 +1,376 @@
|
||||
"""Versioned activation-stream envelope for shard hops.
|
||||
|
||||
The transport still moves raw bytes over HTTP/WebSocket, but the payload now has
|
||||
a stable, extensible envelope that names each tensor, preserves unknown fields,
|
||||
and can round-trip deterministically across direct and relayed hops.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
SCHEMA_NAME = "meshnet.activation-stream"
|
||||
SCHEMA_VERSION = 1
|
||||
DEFAULT_FRAGMENT_BYTES = 64 * 1024
|
||||
|
||||
|
||||
def _canonical_json(data: Any) -> bytes:
|
||||
return json.dumps(data, sort_keys=True, separators=(",", ":"), ensure_ascii=False).encode("utf-8")
|
||||
|
||||
|
||||
def _sha256_hex(data: bytes) -> str:
|
||||
return hashlib.sha256(data).hexdigest()
|
||||
|
||||
|
||||
def _normalize_shape(shape: list[int] | tuple[int, ...]) -> list[int]:
|
||||
normalized = [int(dim) for dim in shape]
|
||||
if not normalized or any(dim <= 0 for dim in normalized):
|
||||
raise ValueError("shape must be a non-empty list of positive integers")
|
||||
return normalized
|
||||
|
||||
|
||||
def _fragment_bytes(body: bytes, max_fragment_bytes: int) -> tuple[bytes, ...]:
|
||||
if max_fragment_bytes <= 0:
|
||||
raise ValueError("max_fragment_bytes must be positive")
|
||||
if not body:
|
||||
return (b"",)
|
||||
return tuple(body[offset : offset + max_fragment_bytes] for offset in range(0, len(body), max_fragment_bytes))
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TensorFragment:
|
||||
"""One bounded chunk of a named tensor."""
|
||||
|
||||
offset: int
|
||||
body: bytes
|
||||
checksum: str
|
||||
compression: str = "identity"
|
||||
extensions: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@classmethod
|
||||
def from_bytes(
|
||||
cls,
|
||||
body: bytes,
|
||||
*,
|
||||
offset: int,
|
||||
compression: str = "identity",
|
||||
extensions: dict[str, Any] | None = None,
|
||||
) -> "TensorFragment":
|
||||
return cls(
|
||||
offset=int(offset),
|
||||
body=bytes(body),
|
||||
checksum=_sha256_hex(body),
|
||||
compression=compression,
|
||||
extensions=dict(extensions or {}),
|
||||
)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
data = {
|
||||
"offset": self.offset,
|
||||
"compression": self.compression,
|
||||
"checksum": self.checksum,
|
||||
"body_base64": base64.b64encode(self.body).decode("ascii"),
|
||||
}
|
||||
data.update(self.extensions)
|
||||
return data
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any]) -> "TensorFragment":
|
||||
known = {"offset", "compression", "checksum", "body_base64"}
|
||||
body = base64.b64decode(data.get("body_base64", ""))
|
||||
fragment = cls(
|
||||
offset=int(data["offset"]),
|
||||
body=body,
|
||||
checksum=str(data.get("checksum") or _sha256_hex(body)),
|
||||
compression=str(data.get("compression") or "identity"),
|
||||
extensions={k: v for k, v in data.items() if k not in known},
|
||||
)
|
||||
fragment.validate()
|
||||
return fragment
|
||||
|
||||
def validate(self) -> None:
|
||||
if self.checksum != _sha256_hex(self.body):
|
||||
raise ValueError("fragment checksum mismatch")
|
||||
if self.offset < 0:
|
||||
raise ValueError("fragment offset must be non-negative")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class NamedTensor:
|
||||
"""A tensor named within a versioned activation envelope."""
|
||||
|
||||
name: str
|
||||
shape: list[int]
|
||||
dtype: str
|
||||
byte_order: str
|
||||
checksum: str
|
||||
fragments: tuple[TensorFragment, ...]
|
||||
compression: str = "identity"
|
||||
extensions: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@classmethod
|
||||
def from_bytes(
|
||||
cls,
|
||||
*,
|
||||
name: str,
|
||||
body: bytes,
|
||||
shape: list[int] | tuple[int, ...],
|
||||
dtype: str,
|
||||
byte_order: str = "little",
|
||||
compression: str = "identity",
|
||||
max_fragment_bytes: int = DEFAULT_FRAGMENT_BYTES,
|
||||
extensions: dict[str, Any] | None = None,
|
||||
) -> "NamedTensor":
|
||||
normalized_shape = _normalize_shape(shape)
|
||||
fragments = tuple(
|
||||
TensorFragment.from_bytes(fragment, offset=offset, compression=compression)
|
||||
for offset, fragment in enumerate(_fragment_bytes(body, max_fragment_bytes))
|
||||
for offset in (offset * max_fragment_bytes,)
|
||||
)
|
||||
return cls(
|
||||
name=str(name),
|
||||
shape=normalized_shape,
|
||||
dtype=str(dtype),
|
||||
byte_order=str(byte_order),
|
||||
checksum=_sha256_hex(body),
|
||||
fragments=fragments,
|
||||
compression=compression,
|
||||
extensions=dict(extensions or {}),
|
||||
)
|
||||
|
||||
def body(self) -> bytes:
|
||||
ordered = sorted(self.fragments, key=lambda frag: frag.offset)
|
||||
body = b"".join(fragment.body for fragment in ordered)
|
||||
if _sha256_hex(body) != self.checksum:
|
||||
raise ValueError(f"tensor {self.name!r} checksum mismatch")
|
||||
return body
|
||||
|
||||
def validate(self) -> None:
|
||||
for fragment in self.fragments:
|
||||
fragment.validate()
|
||||
self.body()
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
data = {
|
||||
"name": self.name,
|
||||
"shape": list(self.shape),
|
||||
"dtype": self.dtype,
|
||||
"byte_order": self.byte_order,
|
||||
"compression": self.compression,
|
||||
"checksum": self.checksum,
|
||||
"fragments": [fragment.to_dict() for fragment in self.fragments],
|
||||
}
|
||||
data.update(self.extensions)
|
||||
return data
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any]) -> "NamedTensor":
|
||||
known = {
|
||||
"name",
|
||||
"shape",
|
||||
"dtype",
|
||||
"byte_order",
|
||||
"compression",
|
||||
"checksum",
|
||||
"fragments",
|
||||
}
|
||||
tensor = cls(
|
||||
name=str(data["name"]),
|
||||
shape=_normalize_shape(list(data["shape"])),
|
||||
dtype=str(data["dtype"]),
|
||||
byte_order=str(data.get("byte_order", "little")),
|
||||
compression=str(data.get("compression", "identity")),
|
||||
checksum=str(data["checksum"]),
|
||||
fragments=tuple(TensorFragment.from_dict(fragment) for fragment in data.get("fragments", [])),
|
||||
extensions={k: v for k, v in data.items() if k not in known},
|
||||
)
|
||||
tensor.validate()
|
||||
return tensor
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ActivationEnvelope:
|
||||
"""Versioned envelope for shard activation traffic."""
|
||||
|
||||
request_id: str
|
||||
work_id: str
|
||||
route_session: str
|
||||
route_epoch: int
|
||||
shard_start: int
|
||||
effective_start: int
|
||||
phase: str
|
||||
position: int
|
||||
idempotency_step: int
|
||||
tensors: tuple[NamedTensor, ...]
|
||||
version: int = SCHEMA_VERSION
|
||||
schema: str = SCHEMA_NAME
|
||||
token_id_sideband: list[int] | None = None
|
||||
architecture_state: dict[str, Any] | None = None
|
||||
recurrent_state: dict[str, Any] | None = None
|
||||
mtp: dict[str, Any] | None = None
|
||||
extensions: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@classmethod
|
||||
def from_tensor_payload(
|
||||
cls,
|
||||
*,
|
||||
payload: Any,
|
||||
name: str,
|
||||
request_id: str,
|
||||
work_id: str,
|
||||
route_session: str,
|
||||
route_epoch: int,
|
||||
shard_start: int,
|
||||
effective_start: int,
|
||||
phase: str,
|
||||
position: int,
|
||||
idempotency_step: int,
|
||||
byte_order: str = "little",
|
||||
compression: str = "identity",
|
||||
max_fragment_bytes: int = DEFAULT_FRAGMENT_BYTES,
|
||||
token_id_sideband: list[int] | None = None,
|
||||
architecture_state: dict[str, Any] | None = None,
|
||||
recurrent_state: dict[str, Any] | None = None,
|
||||
mtp: dict[str, Any] | None = None,
|
||||
extensions: dict[str, Any] | None = None,
|
||||
) -> "ActivationEnvelope":
|
||||
tensor = NamedTensor.from_bytes(
|
||||
name=name,
|
||||
body=payload.body,
|
||||
shape=payload.shape,
|
||||
dtype="bfloat16",
|
||||
byte_order=byte_order,
|
||||
compression=compression,
|
||||
max_fragment_bytes=max_fragment_bytes,
|
||||
extensions={
|
||||
"attention_mask_header": payload.attention_mask_header,
|
||||
"position_ids_header": payload.position_ids_header,
|
||||
**({"past_len": payload.past_len} if payload.past_len is not None else {}),
|
||||
},
|
||||
)
|
||||
return cls(
|
||||
request_id=request_id,
|
||||
work_id=work_id,
|
||||
route_session=route_session,
|
||||
route_epoch=int(route_epoch),
|
||||
shard_start=int(shard_start),
|
||||
effective_start=int(effective_start),
|
||||
phase=str(phase),
|
||||
position=int(position),
|
||||
idempotency_step=int(idempotency_step),
|
||||
tensors=(tensor,),
|
||||
token_id_sideband=list(token_id_sideband) if token_id_sideband is not None else None,
|
||||
architecture_state=architecture_state,
|
||||
recurrent_state=recurrent_state,
|
||||
mtp=mtp,
|
||||
extensions=dict(extensions or {}),
|
||||
)
|
||||
|
||||
def to_tensor_payload(self, *, tensor_name: str = "activations") -> Any:
|
||||
from .model_backend import TensorPayload
|
||||
|
||||
tensor = self.tensor(tensor_name)
|
||||
return TensorPayload(
|
||||
body=tensor.body(),
|
||||
shape=list(tensor.shape),
|
||||
attention_mask_header=tensor.extensions.get("attention_mask_header"),
|
||||
position_ids_header=tensor.extensions.get("position_ids_header"),
|
||||
past_len=tensor.extensions.get("past_len"),
|
||||
)
|
||||
|
||||
def tensor(self, name: str = "activations") -> NamedTensor:
|
||||
for tensor in self.tensors:
|
||||
if tensor.name == name:
|
||||
return tensor
|
||||
raise KeyError(name)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
data = {
|
||||
"schema": self.schema,
|
||||
"version": self.version,
|
||||
"request_id": self.request_id,
|
||||
"work_id": self.work_id,
|
||||
"route_session": self.route_session,
|
||||
"route_epoch": self.route_epoch,
|
||||
"shard_start": self.shard_start,
|
||||
"effective_start": self.effective_start,
|
||||
"phase": self.phase,
|
||||
"position": self.position,
|
||||
"idempotency_step": self.idempotency_step,
|
||||
"tensors": [tensor.to_dict() for tensor in self.tensors],
|
||||
}
|
||||
if self.token_id_sideband is not None:
|
||||
data["token_id_sideband"] = list(self.token_id_sideband)
|
||||
if self.architecture_state is not None:
|
||||
data["architecture_state"] = self.architecture_state
|
||||
if self.recurrent_state is not None:
|
||||
data["recurrent_state"] = self.recurrent_state
|
||||
if self.mtp is not None:
|
||||
data["mtp"] = self.mtp
|
||||
data.update(self.extensions)
|
||||
return data
|
||||
|
||||
def to_bytes(self, *, max_bytes: int | None = None) -> bytes:
|
||||
raw = _canonical_json(self.to_dict())
|
||||
if max_bytes is not None and len(raw) > max_bytes:
|
||||
raise ValueError("activation envelope exceeds the size limit")
|
||||
return raw
|
||||
|
||||
@classmethod
|
||||
def from_bytes(cls, data: bytes) -> "ActivationEnvelope":
|
||||
payload = json.loads(data)
|
||||
if not isinstance(payload, dict):
|
||||
raise ValueError("activation envelope must be a JSON object")
|
||||
known = {
|
||||
"schema",
|
||||
"version",
|
||||
"request_id",
|
||||
"work_id",
|
||||
"route_session",
|
||||
"route_epoch",
|
||||
"shard_start",
|
||||
"effective_start",
|
||||
"phase",
|
||||
"position",
|
||||
"idempotency_step",
|
||||
"tensors",
|
||||
"token_id_sideband",
|
||||
"architecture_state",
|
||||
"recurrent_state",
|
||||
"mtp",
|
||||
}
|
||||
envelope = cls(
|
||||
schema=str(payload.get("schema", SCHEMA_NAME)),
|
||||
version=int(payload.get("version", SCHEMA_VERSION)),
|
||||
request_id=str(payload["request_id"]),
|
||||
work_id=str(payload["work_id"]),
|
||||
route_session=str(payload["route_session"]),
|
||||
route_epoch=int(payload["route_epoch"]),
|
||||
shard_start=int(payload["shard_start"]),
|
||||
effective_start=int(payload["effective_start"]),
|
||||
phase=str(payload["phase"]),
|
||||
position=int(payload["position"]),
|
||||
idempotency_step=int(payload["idempotency_step"]),
|
||||
tensors=tuple(NamedTensor.from_dict(item) for item in payload.get("tensors", [])),
|
||||
token_id_sideband=payload.get("token_id_sideband"),
|
||||
architecture_state=payload.get("architecture_state"),
|
||||
recurrent_state=payload.get("recurrent_state"),
|
||||
mtp=payload.get("mtp"),
|
||||
extensions={k: v for k, v in payload.items() if k not in known},
|
||||
)
|
||||
envelope.validate()
|
||||
return envelope
|
||||
|
||||
def validate(self) -> None:
|
||||
if self.version != SCHEMA_VERSION:
|
||||
raise ValueError("unsupported activation envelope version")
|
||||
if self.schema != SCHEMA_NAME:
|
||||
raise ValueError("unsupported activation envelope schema")
|
||||
if self.phase not in {"prefill", "decode"}:
|
||||
raise ValueError("phase must be prefill or decode")
|
||||
for tensor in self.tensors:
|
||||
tensor.validate()
|
||||
Reference in New Issue
Block a user