From 3921f94457b6f77bf77ecda5dfd171a999884509 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 3 Sep 2026 15:00:23 -0400 Subject: [PATCH 1/6] Apply ruff-format to iter.py The file predates the repo's current ruff config and had never been run through it, so the pre-commit hook reflows it wholesale on first touch. Separated here so the change that follows is readable. --- src/ezmsg/xdf/iter.py | 43 +++++++++++-------------------------------- 1 file changed, 11 insertions(+), 32 deletions(-) diff --git a/src/ezmsg/xdf/iter.py b/src/ezmsg/xdf/iter.py index d683698..7652c0d 100644 --- a/src/ezmsg/xdf/iter.py +++ b/src/ezmsg/xdf/iter.py @@ -1,5 +1,5 @@ -from pathlib import Path import queue +from pathlib import Path import numpy as np import numpy.typing as npt @@ -12,8 +12,7 @@ class XDFIterator: def __init__( self, filepath: Path | str, - select: set[str] - | None = None, # If set, then the iterator yields only AxisArray of selected stream(s). + select: set[str] | None = None, # If set, then the iterator yields only AxisArray of selected stream(s). # If None (default), then the iterator yields dicts with keys for each stream chunk_dur: float = 1.0, # Attempt to chunk data into chunks of this duration. start_time: float | None = None, @@ -60,9 +59,7 @@ def __init__( self._chunk_ix = 0 self._last_time = 0.0 self._metadata = {} - self._prev_file_read_s: float = ( - 0 # File read header in seconds for previous iteration - ) + self._prev_file_read_s: float = 0 # File read header in seconds for previous iteration self._time_range: tuple[float | None, float | None] = (start_time, stop_time) self._scan_file() @@ -79,9 +76,7 @@ def _scan_file(self): # Load xdf self._streams, fileheader = pyxdf.load_xdf( self._filepath, - select_streams=None - if (self._select is None or self._rezero) - else [{"name": _} for _ in self._select], + select_streams=None if (self._select is None or self._rezero) else [{"name": _} for _ in self._select], ) self._metadata = {} self._file_read_s = 0 @@ -175,9 +170,7 @@ def __next__(self) -> dict[str, tuple[npt.NDArray, npt.NDArray]]: (self._chunk_ix + 1) * self._chunk_dur + self._t0, ) for strm in self._streams: - b_chunk = np.logical_and( - strm["time_stamps"] >= t_start, strm["time_stamps"] < t_stop - ) + b_chunk = np.logical_and(strm["time_stamps"] >= t_start, strm["time_stamps"] < t_stop) out_tvec = strm["time_stamps"][b_chunk] out_data = strm["time_series"][b_chunk] out_dict[strm["info"]["name"][0]] = (out_data, out_tvec) @@ -212,19 +205,11 @@ def __init__(self, *args, select: str, **kwargs): _sel = [_ for _ in self._select][0] labels = labels_from_strm(self._streams[0]) if self._metadata[_sel].get("nominal_srate", None): - time_ax = AxisArray.TimeAxis( - fs=self._metadata[_sel]["nominal_srate"], offset=0 - ) + time_ax = AxisArray.TimeAxis(fs=self._metadata[_sel]["nominal_srate"], offset=0) else: - time_ax = AxisArray.CoordinateAxis( - data=np.array([]), - dims=["time"], - unit="s" - ) + time_ax = AxisArray.CoordinateAxis(data=np.array([]), dims=["time"], unit="s") self._template = AxisArray( - data=np.zeros( - (0, len(labels)), dtype=self._streams[0]["time_series"].dtype - ), + data=np.zeros((0, len(labels)), dtype=self._streams[0]["time_series"].dtype), dims=["time", "ch"], axes={ "time": time_ax, @@ -286,9 +271,7 @@ def __init__(self, *args, force_single_sample: set = set(), **kwargs): else AxisArray.CoordinateAxis(data=np.array([]), dims=["time"], unit="s") ) self._templates[stream_name] = AxisArray( - data=np.zeros( - (0, stream_meta["channel_count"]), dtype=stream["time_series"].dtype - ), + data=np.zeros((0, stream_meta["channel_count"]), dtype=stream["time_series"].dtype), dims=["time", "ch"], axes={ "time": time_ax, @@ -320,9 +303,7 @@ def __next__(self) -> AxisArray | None: data=data[ix : ix + 1], axes={ **template.axes, - "time": replace( - template.axes["time"], **t_kwargs - ), + "time": replace(template.axes["time"], **t_kwargs), }, ) ) @@ -330,9 +311,7 @@ def __next__(self) -> AxisArray | None: if isinstance(template.axes["time"], AxisArray.CoordinateAxis): t_kwargs = {"data": tvec if len(tvec) else np.array([])} else: - t_kwargs = { - "offset": tvec[0] if len(tvec) else self._last_time - } + t_kwargs = {"offset": tvec[0] if len(tvec) else self._last_time} self._pubqueue.put_nowait( replace( template, From 5b9bcf22a5a802c281cb6ad02844a637153eff9d Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 3 Sep 2026 15:00:35 -0400 Subject: [PATCH 2/6] Declare chunk_dim and prime the channel fingerprint Two things every consumer of a replayed file needs, and only the reader can supply. `chunk_dim` names the dimension messages accumulate along. A consumer caches state against the stream's configuration -- channel count, labels, sample rate -- and must exclude the one dimension whose length is just however much of the file this chunk covered. It cannot reliably infer which that is: it is `time` here but `win` downstream of a windowing stage, so a guess either thrashes on chunk-size jitter or stops noticing real changes. Both iterators declare it, and it holds whether the stream is regular or carries per-sample timestamps. `CoordinateAxis.fingerprint` is a content digest, computed on first access and cached on the instance. Each template builds its `ch` axis once and every message reuses that object, so priming costs one checksum per stream. Left cold it is computed by the first stateful consumer in this process -- and, because unpickling builds a new axis object per message, by the first consumer in every other process, on every message. Unverified by tests: this repo has no XDF fixture, only the placeholder in tests/test_iter.py. The change mirrors ezmsg-neo and ezmsg-nwb, where it is covered. Requires ezmsg 3.10.0b2 for both fields. --- pyproject.toml | 3 ++- src/ezmsg/xdf/iter.py | 18 ++++++++++++++++-- 2 files changed, 18 insertions(+), 3 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 150815a..02d0879 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,7 +8,8 @@ authors = [ requires-python = ">=3.10.15" dynamic = ["version"] dependencies = [ - "ezmsg>=3.6.0", + # 3.10.0b2 for AxisArray.chunk_dim and CoordinateAxis.fingerprint. + "ezmsg>=3.10.0b2", "numpy>=2.0.2", "pyxdf>=1.16.8", ] diff --git a/src/ezmsg/xdf/iter.py b/src/ezmsg/xdf/iter.py index 7652c0d..3d84766 100644 --- a/src/ezmsg/xdf/iter.py +++ b/src/ezmsg/xdf/iter.py @@ -208,14 +208,25 @@ def __init__(self, *args, select: str, **kwargs): time_ax = AxisArray.TimeAxis(fs=self._metadata[_sel]["nominal_srate"], offset=0) else: time_ax = AxisArray.CoordinateAxis(data=np.array([]), dims=["time"], unit="s") + ch_ax = AxisArray.CoordinateAxis(data=np.array(labels), dims=["ch"]) + # Compute the channel fingerprint once, now. It is cached on the axis and + # pickled with it, and every message from this stream reuses this same axis + # object, so one checksum covers the whole file. Left cold it would be + # computed by the first stateful consumer in this process -- and, since + # unpickling builds a new axis object per message, by the first consumer in + # every other process, on every message. + ch_ax.fingerprint self._template = AxisArray( data=np.zeros((0, len(labels)), dtype=self._streams[0]["time_series"].dtype), dims=["time", "ch"], axes={ "time": time_ax, - "ch": AxisArray.CoordinateAxis(data=np.array(labels), dims=["ch"]), + "ch": ch_ax, }, key=self._streams[0]["info"]["name"][0], + # Messages accumulate along `time`, whether the stream is regular or + # carries per-sample timestamps; `ch` describes the stream itself. + chunk_dim="time", ) def __next__(self) -> AxisArray: @@ -270,14 +281,17 @@ def __init__(self, *args, force_single_sample: set = set(), **kwargs): if fs else AxisArray.CoordinateAxis(data=np.array([]), dims=["time"], unit="s") ) + ch_ax = AxisArray.CoordinateAxis(data=np.array(labels), dims=["ch"]) + ch_ax.fingerprint # primed once per stream -- see the single-stream iterator above self._templates[stream_name] = AxisArray( data=np.zeros((0, stream_meta["channel_count"]), dtype=stream["time_series"].dtype), dims=["time", "ch"], axes={ "time": time_ax, - "ch": AxisArray.CoordinateAxis(data=np.array(labels), dims=["ch"]), + "ch": ch_ax, }, key=stream_name, + chunk_dim="time", ) self._pubqueue: queue.SimpleQueue[AxisArray] = queue.SimpleQueue() From e603ee66a2dbd510195dc7d66d7ef239ef121a11 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 3 Sep 2026 18:36:38 -0400 Subject: [PATCH 3/6] Add a generated XDF fixture and test what the iterators produce `tests/test_iter.py` was a `test_dummy` placeholder with a TODO asking for a small XDF source, so nothing in this package was covered -- including the chunk_dim and fingerprint change in the previous commit. The fixture is generated rather than checked in as a binary, so it stays readable and adjustable: a test needing an irregular stream, a string stream or a different channel layout changes an argument instead of asking someone to produce a new recording. `create_test_xdf.py` writes the XDF chunked format directly, with each field pinned against what pyxdf's reader actually consumes -- `_read_varlen_int`, `_read_chunk3` and the tag dispatch in `load_xdf` -- since that is the only reader these files ever meet. Three choices in the generator are deliberate and would otherwise look arbitrary: * Every sample carries an explicit timestamp rather than relying on delta decompression. The compressed form encodes "same as last plus 1/srate", which would make the fixture silently agree with any reader that got the nominal rate wrong. * Streams start at t=10 s, not 0, so a reader honouring `rezero` is distinguishable from one ignoring it. * ClockOffset chunks are written even though there is no clock skew to model. Without them pyxdf logs "Segments and clock-segments differ" on every load, and a fixture that warns every time trains readers to ignore warnings. The default file has two streams -- a 100 Hz 4-channel float32 ramp and an irregular string marker stream -- split across several sample chunks so the reader's chunk stitching is exercised rather than arriving as one block. 20 tests cover sample values and ordering, chunk_dur, rezero, the nominal rate reaching the axis gain, per-sample timestamps on the irregular stream, force_single_sample, and both fields the previous commit added. A class at the end checks the fixture itself, since a wrong fixture would make every other assertion agree with the wrong thing. Verified by mutation rather than by the tests merely passing: dropping chunk_dim fails 3, dropping the fingerprint priming fails 3, and rebuilding the channel axis per message instead of reusing it fails 2. --- tests/conftest.py | 26 ++++ tests/create_test_xdf.py | 281 +++++++++++++++++++++++++++++++++++++++ tests/test_iter.py | 190 +++++++++++++++++++++++++- 3 files changed, 494 insertions(+), 3 deletions(-) create mode 100644 tests/conftest.py create mode 100644 tests/create_test_xdf.py diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..7e8d496 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,26 @@ +"""Shared fixtures for the ezmsg-xdf tests. + +The XDF file is generated rather than checked in as a binary -- see +``create_test_xdf.py`` for why, and for what it contains. +""" + +import sys +from pathlib import Path + +import pytest + +sys.path.insert(0, str(Path(__file__).parent)) + +from create_test_xdf import DEFAULT_STREAMS, StreamSpec, write_test_xdf # noqa: E402 + +# Kept alongside the fixture so tests can assert against them by name rather +# than by index into DEFAULT_STREAMS. +EEG_STREAM: StreamSpec = DEFAULT_STREAMS[0] +MARKER_STREAM: StreamSpec = DEFAULT_STREAMS[1] + + +@pytest.fixture(scope="session") +def test_xdf_path(tmp_path_factory) -> Path: + """A two-stream XDF: a 100 Hz 4-channel float32 stream and an irregular + string marker stream, both starting at t=10 s.""" + return write_test_xdf(tmp_path_factory.mktemp("xdf") / "test.xdf") diff --git a/tests/create_test_xdf.py b/tests/create_test_xdf.py new file mode 100644 index 0000000..cf9cd54 --- /dev/null +++ b/tests/create_test_xdf.py @@ -0,0 +1,281 @@ +"""Write a small, self-describing XDF file for the tests. + +Checked in as a generator rather than a binary blob so the fixture is readable +and adjustable: a test that needs an irregular stream, a string stream or a +particular channel layout changes an argument here instead of asking someone to +produce a new recording. + +The format is the XDF specification's chunked layout, and the field-by-field +choices below are pinned against what ``pyxdf``'s reader actually consumes -- +:func:`pyxdf.pyxdf._read_varlen_int`, :func:`~pyxdf.pyxdf._read_chunk3` and the +tag dispatch in :func:`~pyxdf.load_xdf` -- since that is the only reader these +files ever meet. + +Run standalone to drop a file next to this script:: + + python tests/create_test_xdf.py /tmp/test.xdf +""" + +from __future__ import annotations + +import struct +import typing +from pathlib import Path + +import numpy as np + +# XDF chunk tags. Boundary chunks (5) are omitted -- they only help a reader +# resynchronise after corruption, and these files are not corrupt. ClockOffset +# chunks (4) are written even though there is no clock skew to model, because +# pyxdf warns ("Segments and clock-segments differ") on a stream that has sample +# segments but no clock segments, and a fixture that logs a warning on every load +# trains readers to ignore warnings. +TAG_FILE_HEADER = 1 +TAG_STREAM_HEADER = 2 +TAG_SAMPLES = 3 +TAG_CLOCK_OFFSET = 4 +TAG_STREAM_FOOTER = 6 + +FORMAT_DTYPES: dict[str, np.dtype] = { + "int8": np.dtype(" bytes: + """XDF's length prefix: one byte saying how wide the count is, then the count. + + The reader accepts widths of 1, 4 and 8 only, so this picks the narrowest of + those rather than the narrowest possible. + """ + if value < 256: + return b"\x01" + bytes([value]) + if value < 2**32: + return b"\x04" + struct.pack(" bytes: + """One length-prefixed chunk. + + The declared length covers the tag and, where the tag has one, the stream id + -- the reader subtracts both back off before reading the payload. + """ + body = struct.pack(" str: + rows = "".join( + f"{unit}{ch_type}" for label in labels + ) + return f"{rows}" + + +def _stream_header_xml( + name: str, + stream_type: str, + labels: typing.Sequence[str], + srate: float, + channel_format: str, + unit: str, + ch_type: str, +) -> bytes: + return ( + '' + f"{name}" + f"{stream_type}" + f"{len(labels)}" + f"{srate:g}" + f"{channel_format}" + f"0.0" + f"{name}-uid" + f"default" + f"test" + f"{_xml_channels(labels, unit, ch_type)}" + "" + ).encode("utf-8") + + +def _stream_footer_xml(first: float, last: float, n_samples: int, srate: float) -> bytes: + return ( + '' + f"{first!r}" + f"{last!r}" + f"{n_samples}" + f"{srate:g}" + "" + ).encode("utf-8") + + +def _clock_offset_chunk(stream_id: int, collection_time: float, offset: float) -> bytes: + return _chunk( + TAG_CLOCK_OFFSET, + struct.pack(" bytes: + """A [Samples] chunk with an explicit timestamp on every sample. + + Every sample carries its stamp rather than relying on delta decompression. + That is more bytes than a real recorder would write, and deliberately so: + the alternative encodes "same as last plus 1/srate", which would make a + fixture silently agree with any reader that got the nominal rate wrong. + """ + out = [_varlen_int(len(stamps))] + if channel_format == "string": + for row, stamp in zip(values, stamps): + out.append(b"\x08" + struct.pack(" np.ndarray: + if spec.channel_format == "string": + return np.array( + [[f"{spec.name}-{i}-{ch}" for ch in range(len(spec.labels))] for i in range(spec.n_samples)], + dtype=object, + ) + # A per-channel ramp offset by channel index: every sample is identifiable, + # so a test can assert on ordering and chunk boundaries, not just on shape. + ramp = np.arange(spec.n_samples, dtype=np.float64)[:, None] + offsets = np.arange(len(spec.labels), dtype=np.float64)[None, :] * 1000.0 + return (ramp + offsets).astype(FORMAT_DTYPES[spec.channel_format]) + + +def _stamps_for(spec: StreamSpec, rng: np.random.Generator) -> np.ndarray: + if spec.srate > 0: + return spec.t0 + np.arange(spec.n_samples) / spec.srate + # Irregular: monotonic but unevenly spaced, which is the whole point of the + # format's per-sample timestamps. + gaps = rng.uniform(0.05, 0.25, size=spec.n_samples) + return spec.t0 + np.cumsum(gaps) + + +DEFAULT_STREAMS: tuple[StreamSpec, ...] = ( + StreamSpec(name="EEGSignal", labels=("Fz", "Cz", "Pz", "Oz"), srate=100.0, n_samples=250), + StreamSpec( + name="Markers", + labels=("marker",), + srate=0.0, + n_samples=12, + channel_format="string", + stream_type="Markers", + unit="none", + ch_type="Marker", + ), +) + + +def write_test_xdf( + path: Path | str, + streams: typing.Sequence[StreamSpec] = DEFAULT_STREAMS, + samples_per_chunk: int = 32, + seed: int = 0, +) -> Path: + """Write *streams* to *path* and return it. + + Samples are split across several [Samples] chunks so the file exercises the + reader's chunk stitching rather than arriving as one block. + """ + rng = np.random.default_rng(seed) + path = Path(path) + + out = [b"XDF:"] + out.append( + _chunk( + TAG_FILE_HEADER, + b'1.0', + ) + ) + + prepared = [] + for stream_id, spec in enumerate(streams, start=1): + values = _values_for(spec, rng) + stamps = _stamps_for(spec, rng) + prepared.append((stream_id, spec, values, stamps)) + out.append( + _chunk( + TAG_STREAM_HEADER, + _stream_header_xml( + spec.name, spec.stream_type, spec.labels, spec.srate, spec.channel_format, spec.unit, spec.ch_type + ), + stream_id=stream_id, + ) + ) + + # Interleave the streams' sample chunks, as a real recording would have them. + offset = 0 + while any(offset < spec.n_samples for _, spec, _, _ in prepared): + for stream_id, spec, values, stamps in prepared: + if offset >= spec.n_samples: + continue + stop = min(offset + samples_per_chunk, spec.n_samples) + out.append( + _chunk( + TAG_SAMPLES, + _samples_chunk(values[offset:stop], stamps[offset:stop], spec.channel_format), + stream_id=stream_id, + ) + ) + offset += samples_per_chunk + + # A pair of zero-offset clock measurements per stream, bracketing its + # samples. Zero because these timestamps are already on one clock; the + # chunks exist so the reader sees a clock segment covering the data. + for stream_id, spec, _, stamps in prepared: + out.append(_clock_offset_chunk(stream_id, float(stamps[0]) - 1.0, 0.0)) + out.append(_clock_offset_chunk(stream_id, float(stamps[-1]) + 1.0, 0.0)) + + for stream_id, spec, _, stamps in prepared: + out.append( + _chunk( + TAG_STREAM_FOOTER, + _stream_footer_xml(float(stamps[0]), float(stamps[-1]), spec.n_samples, spec.srate), + stream_id=stream_id, + ) + ) + + path.write_bytes(b"".join(out)) + return path + + +if __name__ == "__main__": + import sys + + target = Path(sys.argv[1] if len(sys.argv) > 1 else "test.xdf") + write_test_xdf(target) + print(f"wrote {target} ({target.stat().st_size} bytes)") diff --git a/tests/test_iter.py b/tests/test_iter.py index f1a04f4..708942c 100644 --- a/tests/test_iter.py +++ b/tests/test_iter.py @@ -1,3 +1,187 @@ -def test_dummy(): - # TODO: Add tests. Requires a small XDF source. - pass +"""Behaviour of the XDF iterators, against a generated fixture. + +These pin what the iterators produce -- sample values, ordering, chunking, and +the two message fields consumers key their cached state on -- so that the +producer's contract is checked rather than merely exercised. +""" + +from __future__ import annotations + +import math +import pickle + +import numpy as np +import pytest +from conftest import EEG_STREAM, MARKER_STREAM +from ezmsg.util.messages.axisarray import AxisArray + +from ezmsg.xdf.iter import XDFAxisArrayIterator, XDFIterator, XDFMultiAxArrIterator + + +def eeg_messages(path, **kwargs) -> list[AxisArray]: + return list(XDFAxisArrayIterator(filepath=path, select=EEG_STREAM.name, **kwargs)) + + +class TestTheRawIterator: + def test_it_finds_both_streams(self, test_xdf_path): + it = XDFIterator(filepath=test_xdf_path, chunk_dur=1.0) + seen: set[str] = set() + for chunk in it: + seen |= set(chunk) + assert seen == {EEG_STREAM.name, MARKER_STREAM.name} + + +class TestTheSingleStreamIterator: + def test_it_yields_every_sample_in_order(self, test_xdf_path): + msgs = eeg_messages(test_xdf_path, chunk_dur=0.5) + data = np.concatenate([m.data for m in msgs], axis=0) + assert data.shape == (EEG_STREAM.n_samples, len(EEG_STREAM.labels)) + # create_test_xdf writes a per-channel ramp: sample i, channel c == i + 1000c. + expected = np.arange(EEG_STREAM.n_samples)[:, None] + np.arange(len(EEG_STREAM.labels))[None, :] * 1000.0 + np.testing.assert_allclose(data, expected) + + def test_chunk_dur_controls_how_much_arrives_at_once(self, test_xdf_path): + long = eeg_messages(test_xdf_path, chunk_dur=1.0) + short = eeg_messages(test_xdf_path, chunk_dur=0.25) + assert len(short) > len(long) + assert max(m.data.shape[0] for m in long) > max(m.data.shape[0] for m in short) + + def test_it_carries_the_channel_labels(self, test_xdf_path): + msg = eeg_messages(test_xdf_path)[0] + assert list(msg.axes["ch"].data) == list(EEG_STREAM.labels) + + def test_rezero_moves_the_first_sample_to_zero(self, test_xdf_path): + """The fixture starts at t=10 s, so this distinguishes honouring the + setting from ignoring it.""" + rezeroed = eeg_messages(test_xdf_path, rezero=True)[0] + assert rezeroed.axes["time"].offset == pytest.approx(0.0, abs=1e-6) + + def test_without_rezero_the_file_clock_is_preserved(self, test_xdf_path): + raw = eeg_messages(test_xdf_path, rezero=False)[0] + assert raw.axes["time"].offset == pytest.approx(EEG_STREAM.t0, abs=1e-6) + + def test_the_nominal_rate_becomes_the_axis_gain(self, test_xdf_path): + msg = eeg_messages(test_xdf_path)[0] + assert msg.axes["time"].gain == pytest.approx(1.0 / EEG_STREAM.srate) + + +class TestTheMultiStreamIterator: + @staticmethod + def _messages(path, **kwargs) -> list[AxisArray]: + it = XDFMultiAxArrIterator(filepath=path, chunk_dur=1.0, **kwargs) + return [msg for msg in it if msg is not None] + + def test_both_streams_come_through(self, test_xdf_path): + keys = {m.key for m in self._messages(test_xdf_path)} + assert keys == {EEG_STREAM.name, MARKER_STREAM.name} + + def test_the_irregular_stream_keeps_per_sample_timestamps(self, test_xdf_path): + markers = [m for m in self._messages(test_xdf_path) if m.key == MARKER_STREAM.name] + assert markers, "no marker messages" + time_ax = markers[0].axes["time"] + assert isinstance(time_ax, AxisArray.CoordinateAxis) + stamps = np.concatenate([m.axes["time"].data for m in markers]) + assert len(stamps) == MARKER_STREAM.n_samples + assert np.all(np.diff(stamps) > 0), "timestamps must stay monotonic" + + def test_force_single_sample_splits_an_irregular_stream(self, test_xdf_path): + """Without it, several events inside one ``chunk_dur`` arrive together.""" + batched = [m for m in self._messages(test_xdf_path) if m.key == MARKER_STREAM.name] + split = [ + m + for m in self._messages(test_xdf_path, force_single_sample={MARKER_STREAM.name}) + if m.key == MARKER_STREAM.name + ] + assert len(split) == MARKER_STREAM.n_samples + assert len(split) > len(batched) + assert all(m.data.shape[0] == 1 for m in split) + + +class TestMessagesArriveReadyForConsumers: + """Two things only the source can supply, both set once per stream. + + ``chunk_dim`` names the dimension messages accumulate along -- the one whose + length is just however much of the file this chunk covered, and which a + consumer must leave out of the state it caches against the stream's + configuration. ``fingerprint`` is the channel axis's content digest, cached + on the axis and pickled with it; priming it at construction spares the first + consumer in every process from recomputing it on every message. + """ + + def test_the_single_stream_iterator_declares_its_chunk_dim(self, test_xdf_path): + assert all(m.chunk_dim == "time" for m in eeg_messages(test_xdf_path)) + + def test_the_multi_stream_iterator_declares_it_for_every_stream(self, test_xdf_path): + it = XDFMultiAxArrIterator(filepath=test_xdf_path, chunk_dur=1.0) + undeclared = sorted({m.key for m in it if m is not None and m.chunk_dim != "time"}) + assert not undeclared, f"streams not declaring chunk_dim='time': {undeclared}" + + def test_the_channel_axis_is_primed(self, test_xdf_path): + msg = eeg_messages(test_xdf_path)[0] + assert "_fingerprint" in msg.axes["ch"].__dict__ + assert msg.axes["ch"].fingerprint is not None + + def test_every_stream_of_the_multi_iterator_is_primed(self, test_xdf_path): + it = XDFMultiAxArrIterator(filepath=test_xdf_path, chunk_dur=1.0) + cold = sorted({m.key for m in it if m is not None and "_fingerprint" not in m.axes["ch"].__dict__}) + assert not cold, f"streams handing over a cold ch axis: {cold}" + + def test_one_axis_object_serves_the_whole_stream(self, test_xdf_path): + """What makes priming cheap: the checksum is paid once, not per message.""" + msgs = eeg_messages(test_xdf_path, chunk_dur=0.25) + assert len(msgs) > 1 + assert len({id(m.axes["ch"]) for m in msgs}) == 1 + + def test_the_chunk_axis_is_left_cold(self, test_xdf_path): + """Digesting per-message timestamps would be pure cost: no consumer reads + the chunk axis's fingerprint.""" + it = XDFMultiAxArrIterator(filepath=test_xdf_path, chunk_dur=1.0) + markers = [m for m in it if m is not None and m.key == MARKER_STREAM.name] + assert markers, "no marker messages" + assert all("_fingerprint" not in m.axes["time"].__dict__ for m in markers) + + def test_it_all_survives_the_transport(self, test_xdf_path): + msg = eeg_messages(test_xdf_path)[0] + landed = pickle.loads(pickle.dumps(msg)) + assert landed.chunk_dim == "time" + assert "_fingerprint" in landed.axes["ch"].__dict__ + assert landed.axes["ch"].__dict__["_fingerprint"] == msg.axes["ch"].fingerprint + + +class TestTheFixtureItself: + """The generator is test code, and a wrong fixture would make every + assertion above agree with the wrong thing.""" + + def test_pyxdf_reads_back_what_was_written(self, test_xdf_path): + import pyxdf + + streams, header = pyxdf.load_xdf(str(test_xdf_path)) + assert header["info"]["version"] == ["1.0"] + by_name = {s["info"]["name"][0]: s for s in streams} + assert set(by_name) == {EEG_STREAM.name, MARKER_STREAM.name} + + eeg = by_name[EEG_STREAM.name] + assert np.asarray(eeg["time_series"]).shape == (EEG_STREAM.n_samples, len(EEG_STREAM.labels)) + assert float(eeg["info"]["nominal_srate"][0]) == EEG_STREAM.srate + assert eeg["time_stamps"][0] == pytest.approx(EEG_STREAM.t0) + labels = [c["label"][0] for c in eeg["info"]["desc"][0]["channels"][0]["channel"]] + assert labels == list(EEG_STREAM.labels) + + markers = by_name[MARKER_STREAM.name] + assert float(markers["info"]["nominal_srate"][0]) == 0.0 + assert np.asarray(markers["time_series"]).shape == (MARKER_STREAM.n_samples, 1) + + def test_it_loads_without_pyxdf_warnings(self, test_xdf_path, caplog): + """A fixture that logs on every load trains readers to ignore warnings.""" + import logging + + import pyxdf + + with caplog.at_level(logging.WARNING, logger="pyxdf"): + pyxdf.load_xdf(str(test_xdf_path)) + assert not caplog.records, [r.getMessage() for r in caplog.records] + + def test_the_samples_span_more_than_one_chunk(self, test_xdf_path): + """Otherwise the reader's chunk stitching is never exercised.""" + assert EEG_STREAM.n_samples > 32 + assert math.ceil(EEG_STREAM.n_samples / 32) > 1 From f1ed07050fbe33daba3b61d6269bb5f56dd9a32f Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 3 Sep 2026 18:57:36 -0400 Subject: [PATCH 4/6] Run the tests in CI The push and pull_request triggers were commented out, leaving workflow_dispatch as the only way to run the suite. That was reasonable when the suite was a single `test_dummy` placeholder, and is not now that it covers the iterators, the producers and both units. Every other ezmsg source package triggers on push to main and on pull requests; blackrock, neo and nwb also include dev, which is where these PRs are based, so that is matched here. Without it a PR to dev reports only the publish workflow's build job, which does not run a single test. --- .github/workflows/python-tests.yml | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/.github/workflows/python-tests.yml b/.github/workflows/python-tests.yml index 22602eb..aff7219 100644 --- a/.github/workflows/python-tests.yml +++ b/.github/workflows/python-tests.yml @@ -1,10 +1,12 @@ name: Test package on: -# push: -# branches: [main] -# pull_request: -# branches: [main] + push: + branches: [main] + pull_request: + branches: + - main + - dev workflow_dispatch: jobs: From 0ef5d097ceefdf91e15098499e42b0d1208960dd Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 3 Sep 2026 19:01:34 -0400 Subject: [PATCH 5/6] Key the uv cache on pyproject.toml, not a lock file that isn't there setup-uv is configured with `cache-dependency-glob: "uv.lock"`, but no uv.lock is committed here -- nor in any sibling ezmsg package. The glob matches nothing and the action fails the job before a single dependency is installed: ##[error]No file in /home/runner/work/ezmsg-xdf/ezmsg-xdf matched to [uv.lock], make sure you have checked out the target repository Latent since the workflow was written, and invisible until the previous commit turned the triggers on. ezmsg-lsl already keys on pyproject.toml; this matches it. --- .github/workflows/python-tests.yml | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/.github/workflows/python-tests.yml b/.github/workflows/python-tests.yml index aff7219..57abf79 100644 --- a/.github/workflows/python-tests.yml +++ b/.github/workflows/python-tests.yml @@ -27,7 +27,9 @@ jobs: uses: astral-sh/setup-uv@v3 with: enable-cache: true - cache-dependency-glob: "uv.lock" + # No uv.lock is committed here (nor in any sibling ezmsg package), + # so keying the cache on it matches nothing and setup-uv fails outright. + cache-dependency-glob: "pyproject.toml" - name: Set up Python ${{ matrix.python-version }} run: uv python install ${{ matrix.python-version }} From c275ab13d2f343c5162efd120f55bed9b0773d1c Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 3 Sep 2026 19:03:07 -0400 Subject: [PATCH 6/6] Apply ruff-format to source.py and strip its trailing whitespace ruff has flagged W291 on line 65 since before any of this work. It went unnoticed because the test workflow, which runs the lint step, was never triggered; with the triggers on it fails every matrix entry. Fixing it means touching the file, and this one predates the repo's current ruff config just as iter.py did, so the pre-commit hook reflows it wholesale on first touch. Both are in one commit here because the whitespace fix cannot be staged without the reformat. --- src/ezmsg/xdf/source.py | 10 +++------- 1 file changed, 3 insertions(+), 7 deletions(-) diff --git a/src/ezmsg/xdf/source.py b/src/ezmsg/xdf/source.py index 50690a5..db2bd7e 100644 --- a/src/ezmsg/xdf/source.py +++ b/src/ezmsg/xdf/source.py @@ -62,7 +62,7 @@ class XDFIteratorSettings(ez.Settings): Note, however, that this will terminate the pipeline even if the data published by this unit are still in transit, which will lead to the pipeline output being truncated before it has finished processing the stream. `self_terminating` should only be used when it is not important that the pipeline finish processing data, such - as during prototyping and testing. + as during prototyping and testing. """ @@ -102,9 +102,7 @@ async def pub_chunk(self) -> typing.AsyncGenerator: else: await asyncio.sleep(0) except StopIteration: - ez.logger.debug( - f"File ({self.SETTINGS.filepath} :: {self.SETTINGS.select}) exhausted." - ) + ez.logger.debug(f"File ({self.SETTINGS.filepath} :: {self.SETTINGS.select}) exhausted.") if self.SETTINGS.self_terminating: raise ez.NormalTermination yield self.OUTPUT_TERM, True @@ -152,9 +150,7 @@ async def pub_multi(self) -> typing.AsyncGenerator: else: await asyncio.sleep(0) except StopIteration: - ez.logger.debug( - f"File ({self.SETTINGS.filepath} :: {self.SETTINGS.select}) exhausted." - ) + ez.logger.debug(f"File ({self.SETTINGS.filepath} :: {self.SETTINGS.select}) exhausted.") if self.SETTINGS.self_terminating: raise ez.NormalTermination yield self.OUTPUT_TERM, True