From 9e071f314ca1a6ec8a677223767a97e7d3991f7c Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Wed, 30 Sep 2026 20:13:22 -0400 Subject: [PATCH 1/3] Compute CoordinateAxis.fingerprint when pickling The cached fingerprint already rode along when a producer had touched it, but a producer that never did shipped axes without one, so the first consumer past a process boundary recomputed it for every message. __getstate__ now materializes it first. A reused axis costs a dict lookup; a per-message axis moves the one computation to the publisher. Returns self.__dict__ rather than super().__getstate__(): object only gained __getstate__ in Python 3.11, and ezmsg still supports 3.10. --- src/ezmsg/util/messages/axisarray.py | 22 ++++++++++++++++++++-- tests/messages/test_axisarray.py | 23 +++++++++++++++++++++++ 2 files changed, 43 insertions(+), 2 deletions(-) diff --git a/src/ezmsg/util/messages/axisarray.py b/src/ezmsg/util/messages/axisarray.py index c52d1293..e2c77a7e 100644 --- a/src/ezmsg/util/messages/axisarray.py +++ b/src/ezmsg/util/messages/axisarray.py @@ -253,6 +253,22 @@ def __eq__(self, other): return ArrayWithNamedDims.__eq__(self, other) return NotImplemented + def __getstate__(self) -> dict[str, typing.Any]: + """ + Materialize :attr:`fingerprint` before pickling, so it rides along. + + On the far side of a process boundary every message unpickles into a + new axis object; without a precomputed fingerprint, the first consumer + there would recompute it for every message. Computing it here costs a + dict lookup when the producer reuses its axis objects, and otherwise + moves the one computation to the publisher, whichever unit built the + axis and whether or not it touched the fingerprint. + """ + self.fingerprint + + # Note: should return super().__getstate__() after Python 3.10 support is dropped. + return self.__dict__ + @property def fingerprint(self) -> tuple | None: """ @@ -274,8 +290,10 @@ def fingerprint(self) -> tuple | None: Computed on first access and cached on the instance, so the cost is paid once per axis object rather than once per consumer per message. The - cached value is part of ``__dict__``, so it survives pickling and - arrives already computed on the far side of a process boundary. + cached value is part of ``__dict__``, and pickling computes it if it + has not been already (see :meth:`__getstate__`), so it always arrives + precomputed on the far side of a process boundary. Producers therefore + need not touch it themselves. ``None`` when the contents cannot be digested (a non-numpy backing array, or an object dtype holding values with no string form). Callers diff --git a/tests/messages/test_axisarray.py b/tests/messages/test_axisarray.py index 82c63671..ffabcd5c 100644 --- a/tests/messages/test_axisarray.py +++ b/tests/messages/test_axisarray.py @@ -545,6 +545,29 @@ def test_survives_pickling(self): assert restored.__dict__.get("_fingerprint") is not None # arrived precomputed assert restored.fingerprint == expected + def test_pickling_computes_it_if_untouched(self): + """A producer that never touched the fingerprint still ships it, so the + first consumer past a process boundary does not recompute it per + message.""" + import pickle + + from ezmsg.core.messagemarshal import MessageMarshal + + axis = self._axis(["A", "B", "C"]) + assert "_fingerprint" not in axis.__dict__ + restored = pickle.loads(pickle.dumps(axis)) + assert restored.__dict__.get("_fingerprint") is not None + assert restored.fingerprint == self._axis(["A", "B", "C"]).fingerprint + + # The same through ezmsg's own marshal (protocol 5, out-of-band buffers). + msg = AxisArray( + np.zeros((2, 3)), + dims=["time", "ch"], + axes={"time": AxisArray.TimeAxis(fs=10.0), "ch": self._axis(["A", "B", "C"])}, + ) + restored_msg = MessageMarshal.load(MessageMarshal.dump(msg)) + assert restored_msg.axes["ch"].__dict__.get("_fingerprint") is not None + def test_replace_yields_a_fresh_fingerprint(self): """``replace`` builds a new axis, so the cache cannot leak across.""" axis = self._axis(["A", "B"]) From ee0136e2c309f37e6942293935341b4e31d21cef Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 1 Oct 2026 02:15:58 -0400 Subject: [PATCH 2/3] Widen CoordinateAxis.fingerprint's digest to 64 bits crc32 alone was enough while the fingerprint only decided state resets, where a collision means missing one reconfiguration. The transport is about to send axes a receiver already holds as a token derived from the fingerprint, and a collision there would silently swap one axis's values for another's. Fold adler32 in beside crc32: a structurally different checksum, so equal-length contents must collide in both. The digest stays a single int, so the fingerprint's shape is unchanged. Cost on a structured (label, x, y, z) axis: 1.0 vs 0.36 us at 256 channels, 3.8 vs 1.3 us at 1024 -- once per axis object, and ~10x cheaper than blake2b. The new test constructs a genuine crc32 collision and checks the fingerprints still differ. --- src/ezmsg/util/messages/axisarray.py | 11 ++++++- tests/messages/test_axisarray.py | 44 ++++++++++++++++++++++++++++ 2 files changed, 54 insertions(+), 1 deletion(-) diff --git a/src/ezmsg/util/messages/axisarray.py b/src/ezmsg/util/messages/axisarray.py index e2c77a7e..1710e3c9 100644 --- a/src/ezmsg/util/messages/axisarray.py +++ b/src/ezmsg/util/messages/axisarray.py @@ -338,9 +338,18 @@ def _compute_fingerprint(self) -> tuple | None: # CPython's siphash over the result at ~5.5 GB/s, while crc32 reads the # array's buffer directly at ~29 GB/s. # + # 64 bits, not 32: the fingerprint keys more than state resets -- the + # transport sends an axis a receiver already holds as a token derived + # from it -- and a collision there would silently swap one axis's + # values for another's. crc32 and adler32 are structurally different + # checksums, so equal contents must collide in both; together they cost + # ~3x crc32 alone (1.0 vs 0.36 us on a 256-channel struct axis), still + # ~10x cheaper than a cryptographic digest. + # # dtype goes in as the object, not str(dtype): numpy builds a structured # dtype's repr field by field, which costs ~10x the checksum it annotates. - return (self.unit, tuple(self.dims), data.dtype, data.shape, zlib.crc32(data)) + digest = (zlib.crc32(data) << 32) | zlib.adler32(data) + return (self.unit, tuple(self.dims), data.dtype, data.shape, digest) @dataclass(eq=False) diff --git a/tests/messages/test_axisarray.py b/tests/messages/test_axisarray.py index ffabcd5c..4bad4a44 100644 --- a/tests/messages/test_axisarray.py +++ b/tests/messages/test_axisarray.py @@ -421,6 +421,38 @@ def test_to_xr_dataarray(): ) +def _crc32_collision(n: int) -> tuple[bytes, bytes]: + """Two different n-byte strings with the same crc32. + + For equal lengths crc32(a) ^ crc32(b) is linear in a ^ b over GF(2), so 33 + single-bit flips (in a 32-bit space) must contain a combination whose crc32 + effects cancel; Gaussian elimination finds it. + """ + import zlib + + base = bytes(n) + c0 = zlib.crc32(base) + basis: dict[int, tuple[int, int]] = {} + for bit in range(33): + flipped = bytearray(base) + flipped[bit // 8] ^= 1 << (bit % 8) + value, mask = zlib.crc32(bytes(flipped)) ^ c0, 1 << bit + while value: + pivot = value.bit_length() - 1 + if pivot not in basis: + basis[pivot] = (value, mask) + break + value ^= basis[pivot][0] + mask ^= basis[pivot][1] + if value == 0: + other = bytearray(base) + for b in range(33): + if mask >> b & 1: + other[b // 8] ^= 1 << (b % 8) + return base, bytes(other) + raise AssertionError("unreachable: 33 vectors in 32 dimensions are dependent") + + class TestCoordinateAxisFingerprint: """``CoordinateAxis.fingerprint`` is derived from the contents, not assigned. @@ -447,6 +479,18 @@ def test_different_contents_differ(self): != self._axis(["X", "Y", "Z"]).fingerprint ) + def test_a_crc32_collision_is_still_told_apart(self): + """The digest is 64 bits (crc32 + adler32): two contents built to share + a crc32 must still fingerprint differently, because the transport keys + axes a receiver already holds on the fingerprint.""" + import zlib + + a, b = _crc32_collision(64) + assert a != b and zlib.crc32(a) == zlib.crc32(b) + fa = CoordinateAxis(data=np.frombuffer(a, np.uint8), dims=["ch"]).fingerprint + fb = CoordinateAxis(data=np.frombuffer(b, np.uint8), dims=["ch"]).fingerprint + assert fa != fb + def test_reorder_is_detected(self): assert ( self._axis(["A", "B", "C"]).fingerprint From 8f439975ea6654283b9e218730c7f8f74ce9ad30 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 1 Oct 2026 12:16:35 -0400 Subject: [PATCH 3/3] Chain CoordinateAxis.__getstate__ to ArrayWithNamedDims Stacked on #269, ArrayWithNamedDims now defines __getstate__ (making strided data contiguous before pickling), and returning self.__dict__ here bypassed it for coordinate axes. super().__getstate__() reaches it through the MRO, so it is also safe on Python 3.10, where object has no __getstate__. --- src/ezmsg/util/messages/axisarray.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/src/ezmsg/util/messages/axisarray.py b/src/ezmsg/util/messages/axisarray.py index 1710e3c9..a6acc78c 100644 --- a/src/ezmsg/util/messages/axisarray.py +++ b/src/ezmsg/util/messages/axisarray.py @@ -265,9 +265,10 @@ def __getstate__(self) -> dict[str, typing.Any]: axis and whether or not it touched the fingerprint. """ self.fingerprint - - # Note: should return super().__getstate__() after Python 3.10 support is dropped. - return self.__dict__ + # ArrayWithNamedDims.__getstate__ (reached through the MRO, so this + # works on 3.10 too) makes strided coordinate data contiguous; the + # fingerprint is computed above from the same values either way. + return super().__getstate__() @property def fingerprint(self) -> tuple | None: