From b92916bcfe1fc7b83497953258e1df0192d2f7d7 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 1 Oct 2026 02:43:21 -0400 Subject: [PATCH 1/4] Elide coordinate axes a receiving channel already holds An AxisArray stream sends the same non-stream coordinate axes (channel labels, positions, ...) every message -- often more bytes than the data -- and every message unpickles into new axis objects, so consumers past a process hop can only compare them by value. The publisher now sends each such axis in full once (_AxisDef: token + axis) and thereafter as a 16-byte _AxisRef. The token is a blake2b digest of CoordinateAxis.fingerprint, computed once per axis object. The channel keeps an AxisTable of the axes it was given, copied out of the transport's memory, and substitutes them for references, so every message of a stream shares one axis object on the far side too. - Scope: the axes of a top-level AxisArray, never its stream axis (stream_dim, else "time"). Everything else pickles exactly as before. - Negotiation: a channel sends Command.ELIDE_OK after connecting (an older publisher ignores the byte); a publisher elides only while every non-local channel has said so. - Resets: a channel joining, or a channel that receives a reference it cannot resolve (it drops that message, as the stale-SHM path does, and sends Command.AXIS_RESEND), makes the publisher define every axis again. - EZMSG_DISABLE_AXIS_ELISION turns it off. - fast_replace drops the cached _wire_token with _fingerprint; unrolling the drop makes it slightly faster than the one-name loop it replaces. SHM round trip, (10, n) float32 + structured (label, x, y, z) channel axis: 256 ch 94-96 -> 86-88 us, 1024 ch 91-96 -> 87-88 us. Messages are 58% smaller. Byte payloads: unchanged within noise (perf ab, 12 rounds). --- src/ezmsg/core/axiselision.py | 205 +++++++++++++++++++++++++++++++ src/ezmsg/core/messagecache.py | 4 +- src/ezmsg/core/messagechannel.py | 42 ++++++- src/ezmsg/core/messagemarshal.py | 13 +- src/ezmsg/core/netprotocol.py | 7 ++ src/ezmsg/core/pubclient.py | 31 ++++- src/ezmsg/util/messages/util.py | 7 +- tests/messages/test_replace.py | 13 ++ tests/test_axiselision.py | 195 +++++++++++++++++++++++++++++ 9 files changed, 506 insertions(+), 11 deletions(-) create mode 100644 src/ezmsg/core/axiselision.py create mode 100644 tests/test_axiselision.py diff --git a/src/ezmsg/core/axiselision.py b/src/ezmsg/core/axiselision.py new file mode 100644 index 00000000..1293008e --- /dev/null +++ b/src/ezmsg/core/axiselision.py @@ -0,0 +1,205 @@ +""" +Wire elision of coordinate axes a receiver already holds. + +An :class:`~ezmsg.util.messages.axisarray.AxisArray` stream sends the same +non-stream coordinate axes (channel labels, positions, ...) in every message, +often more bytes than the data itself. Within a process that costs nothing -- +messages share the axis objects -- but across one, every message serializes, +copies and unpickles them again, and the receiver gets a new axis object per +message, so consumers can only compare it by value. + +Here the publisher sends each such axis in full once (an :class:`_AxisDef`: +token plus axis) and thereafter as a 16-byte :class:`_AxisRef`. The receiving +channel keeps the axes it was given, copied out of the transport's memory, and +substitutes them for the references, so every message of a stream shares one +axis object on the far side too. + +The token is derived from :attr:`CoordinateAxis.fingerprint` (a 64-bit content +digest), so equal axes share a token no matter which object carried them. + +Only the axes of a top-level ``AxisArray`` message are elided, and never its +stream axis, whose values change every message. Everything else pickles as +before. Set ``EZMSG_DISABLE_AXIS_ELISION`` to turn the feature off. +""" + +import hashlib +import os +import pickle +import typing + +from collections import OrderedDict + +ELISION_ENABLED = "EZMSG_DISABLE_AXIS_ELISION" not in os.environ + +# Stream dimension assumed for a message that does not declare one; mirrors +# ezmsg-baseproc's default, so a per-message "time" axis is never elided. +FALLBACK_STREAM_DIM = "time" + +# Publisher: forget what was announced (forcing definitions again) past this +# many distinct axes. Receiver: keep at most this many. +MAX_ANNOUNCED = 256 +MAX_TABLE = 256 + +_AxisArray: typing.Any = None +_CoordinateAxis: typing.Any = None + + +def _types() -> tuple[typing.Any, typing.Any]: + # Imported lazily: ezmsg.util.messages imports ezmsg.core. + global _AxisArray, _CoordinateAxis + if _AxisArray is None: + from ..util.messages.axisarray import AxisArray, CoordinateAxis + + _AxisArray, _CoordinateAxis = AxisArray, CoordinateAxis + return _AxisArray, _CoordinateAxis + + +class MissingAxis(Exception): + """A message referenced an axis this receiver does not hold.""" + + +class _AxisRef: + """Stands in, on the wire, for an axis the receiver already holds.""" + + __slots__ = ("token",) + + def __init__(self, token: bytes) -> None: + self.token = token + + def __reduce__(self): + return (_AxisRef, (self.token,)) + + +class _AxisDef: + """Carries an axis on the wire along with the token later messages will use.""" + + __slots__ = ("token", "axis") + + def __init__(self, token: bytes, axis: typing.Any) -> None: + self.token = token + self.axis = axis + + def __reduce__(self): + return (_AxisDef, (self.token, self.axis)) + + +def wire_token(axis: typing.Any) -> bytes | None: + """The axis's token, computed once per axis object; None if it has no + fingerprint (contents that cannot be digested are never elided).""" + d = axis.__dict__ + token = d.get("_wire_token") + if token is None: + fp = axis.fingerprint + if fp is None: + return None + # The fingerprint holds a numpy dtype, slow to unpickle; a fixed-size + # digest of it is what travels. + token = d["_wire_token"] = hashlib.blake2b(pickle.dumps(fp, protocol=5), digest_size=16).digest() + return token + + +def _stream_dim(d: dict) -> str | None: + stream = d.get("stream_dim") + if stream is None and FALLBACK_STREAM_DIM in d["dims"]: + stream = FALLBACK_STREAM_DIM + return stream + + +class AxisElision: + """Publisher side: which axes have been announced to the current channels.""" + + __slots__ = ("announced",) + + def __init__(self) -> None: + self.announced: set[bytes] = set() + + def reset(self) -> None: + """Send every axis in full again (a channel joined, or asked).""" + self.announced.clear() + + def wire(self, obj: typing.Any) -> typing.Any: + """``obj``, or a shallow stand-in whose eligible axes are elided.""" + AxisArray, CoordinateAxis = _types() + if not isinstance(obj, AxisArray): + return obj + d = obj.__dict__ + axes = d["axes"] + stream = _stream_dim(d) + new_axes = None + for dim, axis in axes.items(): + if dim == stream or type(axis) is not CoordinateAxis: + continue + token = wire_token(axis) + if token is None: + continue + if new_axes is None: + new_axes = dict(axes) + if token in self.announced: + new_axes[dim] = _AxisRef(token) + else: + if len(self.announced) >= MAX_ANNOUNCED: + self.announced.clear() + self.announced.add(token) + new_axes[dim] = _AxisDef(token, axis) + if new_axes is None: + return obj + # A bare copy of the message (no __init__, so no validation) differing + # only in its axes dict; the caller's message is untouched. + wired = object.__new__(type(obj)) + wired.__dict__.update(d) + wired.__dict__["axes"] = new_axes + return wired + + +def _owned(axis: typing.Any) -> typing.Any: + """A copy of ``axis`` whose data owns its memory (it may arrive as a view + into a channel's shared memory, which must not be pinned or read after the + slot is reused), keeping its cached fingerprint and token.""" + import numpy as np + + data = axis.data + if isinstance(data, np.ndarray) and data.flags.owndata: + return axis + out = object.__new__(type(axis)) + out.__dict__.update(axis.__dict__) + out.__dict__["data"] = np.array(data, copy=True) + return out + + +class AxisTable: + """Receiver side: the axes this channel has been given, by token.""" + + __slots__ = ("_axes",) + + def __init__(self) -> None: + self._axes: "OrderedDict[bytes, typing.Any]" = OrderedDict() + + def __len__(self) -> int: + return len(self._axes) + + def resolve(self, obj: typing.Any) -> typing.Any: + """Replace an ``AxisArray``'s elided axes in place; record new ones. + + :raises MissingAxis: if a reference names an axis not held here. + """ + AxisArray, _ = _types() + if not isinstance(obj, AxisArray): + return obj + axes = obj.__dict__["axes"] + for dim, axis in axes.items(): + kind = type(axis) + if kind is _AxisRef: + try: + axes[dim] = self._axes[axis.token] + except KeyError: + raise MissingAxis(axis.token.hex()) from None + elif kind is _AxisDef: + held = self._axes.get(axis.token) + if held is None: + held = self._axes[axis.token] = _owned(axis.axis) + if len(self._axes) > MAX_TABLE: + self._axes.popitem(last=False) + else: + self._axes.move_to_end(axis.token) + axes[dim] = held + return obj diff --git a/src/ezmsg/core/messagecache.py b/src/ezmsg/core/messagecache.py index f14cd2f7..641e95c9 100644 --- a/src/ezmsg/core/messagecache.py +++ b/src/ezmsg/core/messagecache.py @@ -78,7 +78,7 @@ def put_local(self, obj: typing.Any, msg_id: int) -> None: ) ) - def put_from_mem(self, mem: memoryview) -> None: + def put_from_mem(self, mem: memoryview, axis_table: typing.Any = None) -> None: """ Reconstitute a message in mem and keep it in cache, releasing and overwriting the existing slot in cache. @@ -89,7 +89,7 @@ def put_from_mem(self, mem: memoryview) -> None: :type from_mem: memoryview :raises UninitializedMemory: If mem buffer is not properly initialized. """ - ctx = MessageMarshal.obj_from_mem(mem) + ctx = MessageMarshal.obj_from_mem(mem, axis_table) self._put( CacheEntry( object=ctx.__enter__(), diff --git a/src/ezmsg/core/messagechannel.py b/src/ezmsg/core/messagechannel.py index 033ace55..08b02ebe 100644 --- a/src/ezmsg/core/messagechannel.py +++ b/src/ezmsg/core/messagechannel.py @@ -6,6 +6,7 @@ from uuid import UUID from contextlib import contextmanager, suppress +from .axiselision import AxisTable, MissingAxis, ELISION_ENABLED from .shm import SHMContext from .messagemarshal import MessageMarshal, UninitializedMemory from .backpressure import Backpressure @@ -149,7 +150,8 @@ async def _reattach_shm(self, shm_name: str, frame_end: int, msg_id: int) -> Non preserved = chan._snapshot_cached_messages() chan.cache.clear() for preserved_msg in preserved: - chan.cache.put_from_mem(preserved_msg) + with suppress(MissingAxis): + chan.cache.put_from_mem(preserved_msg, chan._axis_table) if chan.shm is not None: old_shm = chan.shm @@ -245,6 +247,9 @@ def __init__( self._graph_address = graph_address self._local_backpressure = None self._channel_kind = ProfileChannelType.UNKNOWN + # Axes the publisher sent in full, for resolving its elided references. + self._axis_table = AxisTable() + self._axes_requested = False @classmethod async def create( @@ -311,6 +316,10 @@ async def create( if num_buffers <= 0: proto.close() raise ValueError("publisher reports invalid num_buffers") + if ELISION_ENABLED: + # Tell the publisher we can resolve elided axes. An older publisher's + # read loop ignores the byte and keeps sending axes in full. + proto.write(Command.ELIDE_OK.value) chan = cls(UUID(id_str), pub_id, num_buffers, shm, graph_address, _guard=cls._SENTINEL) chan.topic = topic @@ -403,7 +412,11 @@ def _deliver_from_shm(self, msg_id: int) -> None: self._release_backpressure(msg_id, self.id) return - self.cache.put_from_mem(shm_buf) + try: + self.cache.put_from_mem(shm_buf, self._axis_table) + except MissingAxis as exc: + self._missing_axis(msg_id, exc) + return self._set_channel_kind(ProfileChannelType.SHM) self._finish_delivery(msg_id) @@ -414,11 +427,34 @@ def _deliver_from_tcp(self, msg_id: int, obj_bytes: bytes) -> None: Called inline from :meth:`ChannelProtocol.frames_available`. """ assert MessageMarshal.msg_id(obj_bytes) == msg_id - self.cache.put_from_mem(memoryview(obj_bytes).toreadonly()) + try: + self.cache.put_from_mem(memoryview(obj_bytes).toreadonly(), self._axis_table) + except MissingAxis as exc: + self._missing_axis(msg_id, exc) + return self._set_channel_kind(ProfileChannelType.TCP) self._finish_delivery(msg_id) + def _missing_axis(self, msg_id: int, exc: MissingAxis) -> None: + """A message referenced an axis we do not hold (we missed or evicted + its definition). Drop it, as the stale-SHM path does, and ask the + publisher to send its axes in full again -- once per episode, since + several messages already in flight may reference it too.""" + logger.warning( + "Channel %s dropping message %s from publisher %s: unknown axis %s", + self.id, + msg_id, + self.pub_id, + exc, + ) + if not self._axes_requested: + self._axes_requested = True + self._proto.write(Command.AXIS_RESEND.value) + self._release_backpressure(msg_id, self.id) + def _finish_delivery(self, msg_id: int) -> None: + # A message resolved, so any axes we asked for have arrived. + self._axes_requested = False if not self._notify_clients(msg_id): # Nobody is listening; need to ack! self.cache.release(msg_id) diff --git a/src/ezmsg/core/messagemarshal.py b/src/ezmsg/core/messagemarshal.py index 4621d642..4d1d02b4 100644 --- a/src/ezmsg/core/messagemarshal.py +++ b/src/ezmsg/core/messagemarshal.py @@ -106,7 +106,7 @@ def msg_id(cls, raw: memoryview | bytes) -> int: @classmethod @contextmanager - def obj_from_mem(cls, mem: memoryview) -> Generator[Any, None, None]: + def obj_from_mem(cls, mem: memoryview, axis_table: Any = None) -> Generator[Any, None, None]: """ Deserialize an object from a memory buffer. @@ -115,9 +115,12 @@ def obj_from_mem(cls, mem: memoryview) -> Generator[Any, None, None]: :param mem: Memory buffer containing serialized object. :type mem: memoryview + :param axis_table: The receiving channel's + :class:`~ezmsg.core.axiselision.AxisTable`, to resolve elided axes. :return: Context manager yielding the deserialized object. :rtype: Generator[Any, None, None] :raises UninitializedMemory: If memory buffer is not properly initialized. + :raises MissingAxis: If the message references an axis not in ``axis_table``. """ cls._assert_initialized(mem) @@ -138,6 +141,8 @@ def obj_from_mem(cls, mem: memoryview) -> Generator[Any, None, None]: sidx += bsz obj = cls.load(buffers) + if axis_table is not None: + axis_table.resolve(obj) try: yield obj @@ -149,7 +154,7 @@ def obj_from_mem(cls, mem: memoryview) -> Generator[Any, None, None]: @classmethod @contextmanager def serialize( - cls, msg_id: int, obj: Any + cls, msg_id: int, obj: Any, elision: Any = None ) -> Generator[tuple[int, bytes, list[memoryview]], None, None]: """ Serialize an object for network transmission. @@ -161,9 +166,13 @@ def serialize( :type msg_id: int :param obj: Object to serialize. :type obj: Any + :param elision: The publisher's :class:`~ezmsg.core.axiselision.AxisElision`, + when every receiver can resolve elided axes. :return: Context manager yielding (total_size, header, buffers) tuple. :rtype: Generator[tuple[int, bytes, list[memoryview]], None, None] """ + if elision is not None: + obj = elision.wire(obj) buffers = cls.dump(obj) header = uint64_to_bytes(len(buffers)) buf_lengths = [len(buf) for buf in buffers] diff --git a/src/ezmsg/core/netprotocol.py b/src/ezmsg/core/netprotocol.py index ff6bf242..a7b55803 100644 --- a/src/ezmsg/core/netprotocol.py +++ b/src/ezmsg/core/netprotocol.py @@ -347,6 +347,13 @@ def _generate_next_value_(name, start, count, last_values) -> bytes: PROCESS_ROUTE_RESPONSE = enum.auto() ERROR = enum.auto() + # Channel -> Publisher: axis elision (appended, so no existing value moves). + # A channel that can resolve elided axes says so once after connecting; an + # older publisher's read loop ignores the byte. + ELIDE_OK = enum.auto() + # A channel received a reference to an axis it does not hold. + AXIS_RESEND = enum.auto() + def create_socket( host: str | None = None, diff --git a/src/ezmsg/core/pubclient.py b/src/ezmsg/core/pubclient.py index 30b3c882..213dd20a 100644 --- a/src/ezmsg/core/pubclient.py +++ b/src/ezmsg/core/pubclient.py @@ -7,6 +7,7 @@ from contextlib import suppress from dataclasses import dataclass +from .axiselision import AxisElision, ELISION_ENABLED from .backpressure import Backpressure from .shm import SHMContext from .graphserver import GraphService @@ -75,6 +76,8 @@ def _resolve_allow_local(force_tcp: bool, allow_local: bool | None) -> bool: class PubChannelInfo(ChannelInfo): pid: int shm_ok: bool = False + # The channel can resolve elided axes (it sent Command.ELIDE_OK). + elide_ok: bool = False class Publisher: @@ -279,6 +282,10 @@ def __init__( self._retired_shms: list[tuple[int, SHMContext]] = [] self._msg_id = 0 self._channels = dict() + # Axis elision: on only while every channel this publisher serializes + # for can resolve it (see _update_elision). + self._elision = AxisElision() + self._elide = False self._channel_tasks = dict() self._running = asyncio.Event() if not start_paused: @@ -424,6 +431,9 @@ async def _handle_channel( :type reader: asyncio.StreamReader """ self._channels[info.id] = info + # A new channel holds no axes yet: announce them all again. + self._elision.reset() + self._update_elision() try: while True: @@ -437,6 +447,14 @@ async def _handle_channel( self._backpressure.free(info.id, msg_id % self._num_buffers) self._profile.sample_inflight(self._backpressure.pressure) + elif msg == Command.ELIDE_OK.value: + info.elide_ok = True + self._update_elision() + + elif msg == Command.AXIS_RESEND.value: + logger.debug(f"Publisher {self.id}: Channel {info.id} asked for axes again") + self._elision.reset() + except (ConnectionResetError, BrokenPipeError): logger.debug(f"Publisher {self.id}: Channel {info.id} connection fail") @@ -445,6 +463,15 @@ async def _handle_channel( self._profile.sample_inflight(self._backpressure.pressure) await close_stream_writer(self._channels[info.id].writer) del self._channels[info.id] + self._update_elision() + + def _update_elision(self) -> None: + remote = [ch for ch in self._channels.values() if not self._can_deliver_locally(ch)] + enable = ELISION_ENABLED and bool(remote) and all(ch.elide_ok for ch in remote) + if enable and not self._elide: + # What was announced while off may not have reached everyone. + self._elision.reset() + self._elide = enable async def sync(self) -> None: """ @@ -517,7 +544,9 @@ async def broadcast(self, obj: Any) -> None: self._local_channel.put_local(self._msg_id, obj) if any(not self._can_deliver_locally(ch) for ch in self._channels.values()): - with MessageMarshal.serialize(self._msg_id, obj) as ( + with MessageMarshal.serialize( + self._msg_id, obj, self._elision if self._elide else None + ) as ( total_size, header, buffers, diff --git a/src/ezmsg/util/messages/util.py b/src/ezmsg/util/messages/util.py index 61aca2da..9f73c5df 100644 --- a/src/ezmsg/util/messages/util.py +++ b/src/ezmsg/util/messages/util.py @@ -11,7 +11,7 @@ # raises TypeError), and a value derived from the *old* field values must not be # carried onto a modified copy. Dropping is always safe -- the copy recomputes # on next access. Costs ~0.01 us per replace. -_DERIVED_CACHE_ATTRS = ("_fingerprint",) +_DERIVED_CACHE_ATTRS = ("_fingerprint", "_wire_token") def fast_replace(arr: T, **kwargs: Any) -> T: @@ -38,8 +38,9 @@ def fast_replace(arr: T, **kwargs: Any) -> T: :rtype: T """ out_kwargs = arr.__dict__.copy() # Shallow copy - for name in _DERIVED_CACHE_ATTRS: - out_kwargs.pop(name, None) + # _DERIVED_CACHE_ATTRS, unrolled: cheaper than a loop over even one name. + out_kwargs.pop("_fingerprint", None) + out_kwargs.pop("_wire_token", None) out_kwargs.update(kwargs) return arr.__class__(**out_kwargs) diff --git a/tests/messages/test_replace.py b/tests/messages/test_replace.py index 475c64c8..9fe262ad 100644 --- a/tests/messages/test_replace.py +++ b/tests/messages/test_replace.py @@ -57,3 +57,16 @@ def test_axisarray_replace_is_unaffected(self, replace_fn): # The axis object is passed through by reference, cache intact. assert updated.axes["ch"] is axis assert updated.axes["ch"].fingerprint == axis.fingerprint + + +@pytest.mark.parametrize("name", __import__("ezmsg.util.messages.util", fromlist=["x"])._DERIVED_CACHE_ATTRS) +def test_fast_replace_drops_every_derived_cache_attr(name): + """fast_replace drops these by unrolled pops; keep it in step with the list.""" + import numpy as np + + from ezmsg.util.messages.axisarray import CoordinateAxis + from ezmsg.util.messages.util import fast_replace + + ax = CoordinateAxis(data=np.arange(3), dims=["ch"]) + ax.__dict__[name] = "stale" + assert name not in fast_replace(ax, data=np.arange(4)).__dict__ diff --git a/tests/test_axiselision.py b/tests/test_axiselision.py new file mode 100644 index 00000000..6f59e4e5 --- /dev/null +++ b/tests/test_axiselision.py @@ -0,0 +1,195 @@ +"""Wire elision of coordinate axes a receiver already holds.""" + +import asyncio +import pickle + +import numpy as np +import pytest + +import ezmsg.core as ez +from ezmsg.core.axiselision import AxisElision, AxisTable, MissingAxis, _AxisDef, _AxisRef, wire_token +from ezmsg.core.messagemarshal import MessageMarshal +from ezmsg.util.messages.axisarray import AxisArray, CoordinateAxis +from ezmsg.util.messages.util import replace + +STRUCT = np.dtype([("label", "U8"), ("x", "f8"), ("y", "f8"), ("z", "f8")]) + + +def ch_axis(n=16, prefix="e"): + a = np.zeros(n, STRUCT) + a["label"] = [f"{prefix}{i:03d}" for i in range(n)] + a["x"] = np.arange(n) + return CoordinateAxis(data=a, dims=["ch"]) + + +def msg(ch, i=0, stream_dim="time", time_axis=None): + return AxisArray( + np.full((4, len(ch.data)), float(i), np.float32), + dims=["time", "ch"], + axes={"time": time_axis if time_axis is not None else AxisArray.TimeAxis(fs=100.0, offset=i), "ch": ch}, + key="k", + **({"stream_dim": stream_dim} if stream_dim else {}), + ) + + +def roundtrip(obj, elision, table): + """Serialize as a publisher would and load as a channel would.""" + with MessageMarshal.serialize(0, obj, elision) as (total, header, buffers): + raw = bytearray(total) + MessageMarshal._write(memoryview(raw), header, buffers) + with MessageMarshal.obj_from_mem(memoryview(raw), table) as out: + return out + + +class TestWire: + def test_first_message_defines_then_references(self): + el, ch = AxisElision(), ch_axis() + first = el.wire(msg(ch, 0)) + second = el.wire(msg(ch, 1)) + assert type(first.axes["ch"]) is _AxisDef + assert type(second.axes["ch"]) is _AxisRef + assert first.axes["ch"].token == second.axes["ch"].token == wire_token(ch) + + def test_stream_axis_is_never_elided(self): + el = AxisElision() + events = CoordinateAxis(data=np.arange(4.0), dims=["time"], unit="s") + wired = el.wire(msg(ch_axis(), time_axis=events)) + assert wired.axes["time"] is events + + def test_undeclared_stream_dim_falls_back_to_time(self): + el = AxisElision() + events = CoordinateAxis(data=np.arange(4.0), dims=["time"], unit="s") + wired = el.wire(msg(ch_axis(), stream_dim=None, time_axis=events)) + assert wired.axes["time"] is events and type(wired.axes["ch"]) is _AxisDef + + def test_the_callers_message_is_untouched(self): + el, ch = AxisElision(), ch_axis() + m = msg(ch) + el.wire(m) + el.wire(m) + assert m.axes["ch"] is ch + + def test_other_objects_pass_through(self): + el = AxisElision() + for obj in (b"bytes", {"a": 1}, np.zeros(3)): + assert el.wire(obj) is obj + + def test_reset_defines_again(self): + el, ch = AxisElision(), ch_axis() + el.wire(msg(ch)) + el.reset() + assert type(el.wire(msg(ch)).axes["ch"]) is _AxisDef + + def test_equal_axes_share_a_token(self): + assert wire_token(ch_axis()) == wire_token(ch_axis()) + assert wire_token(ch_axis()) != wire_token(ch_axis(prefix="z")) + + def test_replace_drops_the_cached_token(self): + ch = ch_axis() + wire_token(ch) + assert "_wire_token" not in replace(ch, data=ch_axis(prefix="z").data).__dict__ + + +class TestResolve: + def test_messages_share_one_owned_axis(self): + el, table, ch = AxisElision(), AxisTable(), ch_axis() + outs = [roundtrip(msg(ch, i), el, table) for i in range(3)] + held = outs[0].axes["ch"] + assert all(o.axes["ch"] is held for o in outs) + assert held.data.flags.owndata + assert np.array_equal(held.data, ch.data) + assert [float(o.data[0, 0]) for o in outs] == [0.0, 1.0, 2.0] + + def test_a_relabel_reaches_the_receiver(self): + el, table = AxisElision(), AxisTable() + roundtrip(msg(ch_axis()), el, table) + out = roundtrip(msg(ch_axis(prefix="z")), el, table) + assert out.axes["ch"].data["label"][0] == "z000" + + def test_unknown_reference_raises(self): + el, ch = AxisElision(), ch_axis() + roundtrip(msg(ch), el, AxisTable()) # defined to a *different* table + with pytest.raises(MissingAxis): + roundtrip(msg(ch), el, AxisTable()) + + def test_without_elision_nothing_changes(self): + ch = ch_axis() + out = roundtrip(msg(ch), None, AxisTable()) + assert np.array_equal(out.axes["ch"].data, ch.data) + + def test_wired_messages_survive_a_publisher_side_copy(self): + """SHM grow re-serializes slots with copy_obj (no table): references + must survive that round trip unresolved.""" + el, table, ch = AxisElision(), AxisTable(), ch_axis() + roundtrip(msg(ch, 0), el, table) + with MessageMarshal.serialize(1, msg(ch, 1), el) as (total, header, buffers): + src = bytearray(total + 64) + MessageMarshal._write(memoryview(src), header, buffers) + dst = bytearray(total + 64) + MessageMarshal.copy_obj(memoryview(src), memoryview(dst)) + with MessageMarshal.obj_from_mem(memoryview(dst), table) as out: + assert out.axes["ch"] is roundtrip(msg(ch, 2), el, table).axes["ch"] + + +async def _pubsub(ctx, topic): + pub = await ctx.publisher(topic, host="127.0.0.1", num_buffers=4, allow_local=False) + sub = await ctx.subscriber(topic) + for _ in range(100): # the channel's ELIDE_OK arrives asynchronously + if pub._elide: + break + await asyncio.sleep(0.01) + return pub, sub + + +async def _recv(sub): + async with sub.recv_zero_copy() as m: + return m.axes["ch"], float(m.data[0, 0]), m.axes["ch"].data["label"][0] + + +@pytest.mark.asyncio +async def test_end_to_end_over_shm(): + async with ez.GraphContext(auto_start=True) as ctx: + pub, sub = await _pubsub(ctx, "/ELIDE/E2E") + assert pub._elide + ch = ch_axis(64) + got = [] + for i in range(5): + await pub.broadcast(msg(ch, i)) + got.append(await _recv(sub)) + assert [g[1] for g in got] == [0.0, 1.0, 2.0, 3.0, 4.0] + assert all(g[0] is got[0][0] for g in got) # one shared axis on the far side + await pub.broadcast(msg(ch_axis(64, prefix="z"), 5)) + assert (await _recv(sub))[2] == "z000" + + +@pytest.mark.asyncio +async def test_a_receiver_that_lost_its_axes_recovers(): + async with ez.GraphContext(auto_start=True) as ctx: + pub, sub = await _pubsub(ctx, "/ELIDE/RECOVER") + ch = ch_axis(32) + await pub.broadcast(msg(ch, 0)) + await _recv(sub) + channel = sub._channels[pub.id] + channel._axis_table = AxisTable() # as if its definition had been missed + await pub.broadcast(msg(ch, 1)) # a reference it cannot resolve: dropped + for _ in range(100): + if not pub._elision.announced: + break + await asyncio.sleep(0.01) + assert not pub._elision.announced # the channel asked for axes again + await pub.broadcast(msg(ch, 2)) + assert (await _recv(sub))[1] == 2.0 + + +@pytest.mark.asyncio +async def test_disabled_while_a_channel_cannot_resolve(monkeypatch): + """A channel that never says ELIDE_OK (an older ezmsg) keeps elision off.""" + from ezmsg.core import messagechannel + + monkeypatch.setattr(messagechannel, "ELISION_ENABLED", False) + async with ez.GraphContext(auto_start=True) as ctx: + pub, sub = await _pubsub(ctx, "/ELIDE/OLD") + assert not pub._elide + ch = ch_axis(8) + await pub.broadcast(msg(ch, 0)) + assert (await _recv(sub))[1] == 0.0 From c16bff63e59f84127db79d85e0ce1dc6a78863c9 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 1 Oct 2026 02:51:24 -0400 Subject: [PATCH 2/4] Test axis elision end to end over TCP as well as SHM --- tests/test_axiselision.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/tests/test_axiselision.py b/tests/test_axiselision.py index 6f59e4e5..b9d8e6d0 100644 --- a/tests/test_axiselision.py +++ b/tests/test_axiselision.py @@ -131,8 +131,8 @@ def test_wired_messages_survive_a_publisher_side_copy(self): assert out.axes["ch"] is roundtrip(msg(ch, 2), el, table).axes["ch"] -async def _pubsub(ctx, topic): - pub = await ctx.publisher(topic, host="127.0.0.1", num_buffers=4, allow_local=False) +async def _pubsub(ctx, topic, force_tcp=False): + pub = await ctx.publisher(topic, host="127.0.0.1", num_buffers=4, allow_local=False, force_tcp=force_tcp) sub = await ctx.subscriber(topic) for _ in range(100): # the channel's ELIDE_OK arrives asynchronously if pub._elide: @@ -147,9 +147,10 @@ async def _recv(sub): @pytest.mark.asyncio -async def test_end_to_end_over_shm(): +@pytest.mark.parametrize("force_tcp", [False, True], ids=["shm", "tcp"]) +async def test_end_to_end(force_tcp): async with ez.GraphContext(auto_start=True) as ctx: - pub, sub = await _pubsub(ctx, "/ELIDE/E2E") + pub, sub = await _pubsub(ctx, f"/ELIDE/E2E/{force_tcp}", force_tcp) assert pub._elide ch = ch_axis(64) got = [] @@ -158,6 +159,7 @@ async def test_end_to_end_over_shm(): got.append(await _recv(sub)) assert [g[1] for g in got] == [0.0, 1.0, 2.0, 3.0, 4.0] assert all(g[0] is got[0][0] for g in got) # one shared axis on the far side + assert sub._channels[pub.id].channel_kind.name == ("TCP" if force_tcp else "SHM") await pub.broadcast(msg(ch_axis(64, prefix="z"), 5)) assert (await _recv(sub))[2] == "z000" From d7174548244aafbf6f19af1df7bbe001b9054577 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 1 Oct 2026 12:32:56 -0400 Subject: [PATCH 3/4] version bump --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 9c784722..f874c659 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "ezmsg" -version = "3.10.0b4" +version = "3.10.0b5" description = "A simple DAG-based computation model" authors = [ { name = "Griffin Milsap", email = "griffin.milsap@gmail.com" }, From 2bb568a5db12387831ef90e11a2d8de8c1f23b8c Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 1 Oct 2026 13:19:56 -0400 Subject: [PATCH 4/4] Define the default stream dimension once, on AxisArray Elision kept its own FALLBACK_STREAM_DIM ("time") for messages that do not declare stream_dim, mirroring ezmsg-baseproc's STREAMING_DIMS by hand. Both must agree on which axis changes every message, so move the definition to ezmsg.util.messages.axisarray.DEFAULT_STREAM_DIM, next to stream_dim itself, and read it from there (lazily, like elision's other axisarray imports). Consumers such as ezmsg-baseproc can now import the same constant. --- src/ezmsg/core/axiselision.py | 18 +++++++++--------- src/ezmsg/util/messages/axisarray.py | 13 +++++++++++-- tests/test_axiselision.py | 9 +++++++++ 3 files changed, 29 insertions(+), 11 deletions(-) diff --git a/src/ezmsg/core/axiselision.py b/src/ezmsg/core/axiselision.py index 1293008e..405cbe65 100644 --- a/src/ezmsg/core/axiselision.py +++ b/src/ezmsg/core/axiselision.py @@ -31,10 +31,6 @@ ELISION_ENABLED = "EZMSG_DISABLE_AXIS_ELISION" not in os.environ -# Stream dimension assumed for a message that does not declare one; mirrors -# ezmsg-baseproc's default, so a per-message "time" axis is never elided. -FALLBACK_STREAM_DIM = "time" - # Publisher: forget what was announced (forcing definitions again) past this # many distinct axes. Receiver: keep at most this many. MAX_ANNOUNCED = 256 @@ -42,15 +38,19 @@ _AxisArray: typing.Any = None _CoordinateAxis: typing.Any = None +# Stream dimension assumed for a message that does not declare one, so that a +# per-message axis by that name is never elided; the shared definition is +# ezmsg.util.messages.axisarray.DEFAULT_STREAM_DIM. +_default_stream_dim: str | None = None def _types() -> tuple[typing.Any, typing.Any]: # Imported lazily: ezmsg.util.messages imports ezmsg.core. - global _AxisArray, _CoordinateAxis + global _AxisArray, _CoordinateAxis, _default_stream_dim if _AxisArray is None: - from ..util.messages.axisarray import AxisArray, CoordinateAxis + from ..util.messages.axisarray import AxisArray, CoordinateAxis, DEFAULT_STREAM_DIM - _AxisArray, _CoordinateAxis = AxisArray, CoordinateAxis + _AxisArray, _CoordinateAxis, _default_stream_dim = AxisArray, CoordinateAxis, DEFAULT_STREAM_DIM return _AxisArray, _CoordinateAxis @@ -100,8 +100,8 @@ def wire_token(axis: typing.Any) -> bytes | None: def _stream_dim(d: dict) -> str | None: stream = d.get("stream_dim") - if stream is None and FALLBACK_STREAM_DIM in d["dims"]: - stream = FALLBACK_STREAM_DIM + if stream is None and _default_stream_dim in d["dims"]: + stream = _default_stream_dim return stream diff --git a/src/ezmsg/util/messages/axisarray.py b/src/ezmsg/util/messages/axisarray.py index a6acc78c..1b1366a8 100644 --- a/src/ezmsg/util/messages/axisarray.py +++ b/src/ezmsg/util/messages/axisarray.py @@ -118,6 +118,15 @@ def create_time_axis(cls, fs: float, offset: float = 0.0) -> "LinearAxis": return cls(unit="s", gain=1.0 / fs, offset=offset) +DEFAULT_STREAM_DIM = "time" +"""The dimension assumed to be the stream dimension of an :class:`AxisArray` that +does not declare :attr:`~AxisArray.stream_dim` (when it has one by that name). + +Anything that must decide which axis changes every message -- ezmsg's transport, +which never elides the stream axis, and consumers that key cached state on the +rest -- reads it from here so they cannot disagree. +""" + # Distinguishes "no fingerprint cached yet" from "cached, and it is None". _UNSET = object() @@ -397,8 +406,8 @@ class AxisArray(ArrayWithNamedDims): under :meth:`transpose`. Declaring it here puts the answer where it is known -- in the producer -- instead of asking every consumer to guess. - ``None`` means "not declared", leaving consumers to fall back on their own - convention. Any operation that *renames* this dimension is responsible for + ``None`` means "not declared", leaving consumers to fall back on a + convention: :data:`DEFAULT_STREAM_DIM`, if the message has that dimension. Any operation that *renames* this dimension is responsible for updating it, exactly as it already updates ``dims``. """ diff --git a/tests/test_axiselision.py b/tests/test_axiselision.py index b9d8e6d0..d7a2c2f4 100644 --- a/tests/test_axiselision.py +++ b/tests/test_axiselision.py @@ -62,6 +62,15 @@ def test_undeclared_stream_dim_falls_back_to_time(self): wired = el.wire(msg(ch_axis(), stream_dim=None, time_axis=events)) assert wired.axes["time"] is events and type(wired.axes["ch"]) is _AxisDef + def test_the_fallback_is_axisarrays_default_stream_dim(self): + """Elision and consumers (e.g. ezmsg-baseproc) must agree on which axis + is per-message when stream_dim is undeclared: one shared definition.""" + from ezmsg.core import axiselision + from ezmsg.util.messages.axisarray import DEFAULT_STREAM_DIM + + AxisElision().wire(b"anything") # loads the lazy imports + assert axiselision._default_stream_dim == DEFAULT_STREAM_DIM == "time" + def test_the_callers_message_is_untouched(self): el, ch = AxisElision(), ch_axis() m = msg(ch)