diff --git a/src/ezmsg/util/messages/axisarray.py b/src/ezmsg/util/messages/axisarray.py index c52d1293..a6acc78c 100644 --- a/src/ezmsg/util/messages/axisarray.py +++ b/src/ezmsg/util/messages/axisarray.py @@ -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: """ @@ -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 @@ -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) diff --git a/tests/messages/test_axisarray.py b/tests/messages/test_axisarray.py index 82c63671..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 @@ -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"])