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" }, diff --git a/src/ezmsg/core/axiselision.py b/src/ezmsg/core/axiselision.py new file mode 100644 index 00000000..405cbe65 --- /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 + +# 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 +# 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, _default_stream_dim + if _AxisArray is None: + from ..util.messages.axisarray import AxisArray, CoordinateAxis, DEFAULT_STREAM_DIM + + _AxisArray, _CoordinateAxis, _default_stream_dim = AxisArray, CoordinateAxis, DEFAULT_STREAM_DIM + 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 _default_stream_dim in d["dims"]: + stream = _default_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/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/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..d7a2c2f4 --- /dev/null +++ b/tests/test_axiselision.py @@ -0,0 +1,206 @@ +"""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_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) + 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, 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: + 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 +@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, f"/ELIDE/E2E/{force_tcp}", force_tcp) + 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 + 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" + + +@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