Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 31 additions & 3 deletions src/ezmsg/util/messages/axisarray.py
Original file line number Diff line number Diff line change
Expand Up @@ -253,6 +253,23 @@ 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
# 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:
"""
Expand All @@ -274,8 +291,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
Expand Down Expand Up @@ -320,9 +339,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)
Expand Down
67 changes: 67 additions & 0 deletions tests/messages/test_axisarray.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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
Expand Down Expand Up @@ -545,6 +589,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"])
Expand Down