From 3921f94457b6f77bf77ecda5dfd171a999884509 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 3 Sep 2026 15:00:23 -0400 Subject: [PATCH 1/8] 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/8] 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/8] 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/8] 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/8] 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/8] 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 From 4a1207b34854f737d5b2b81761464051719d49fd Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 3 Sep 2026 18:53:31 -0400 Subject: [PATCH 7/8] Build the iterators on ezmsg-baseproc's producer classes This was the only ezmsg source package still on `ez.Unit` + `GenState`; blackrock, lsl, neo and nwb all use `BaseStatefulProducer` and `BaseProducerUnit`. Four things follow from joining them. **The file open no longer blocks the event loop.** `pyxdf.load_xdf` reads and decodes the entire file, and it ran inside a synchronous `initialize()`, so graph startup stalled every other unit in the process for as long as that took. It is now a state reset, and `_areset_state` puts it on a worker thread. This deliberately diverges from ezmsg-neo and ezmsg-nwb, which call `_reset_state()` eagerly from `__init__` and therefore still pay the first open on the loop -- nwb's own test says as much ("Discard the eager sync invocation from __init__"), so its `_areset_state` only covers reopens after a settings change. Nothing in this package's public surface reads stream metadata before the first chunk, so there is nothing to lose by waiting, and a test asserts both halves: construction touches no file, and the load lands off the loop. **Settings can change at runtime.** `BaseProducerUnit` brings `INPUT_SETTINGS` and delegates to `update_settings`; previously settings were read once in `initialize` and a running unit could not be retargeted. `playback_rate` and `self_terminating` are listed in `NONRESET_SETTINGS_FIELDS` because they belong to the unit's pacing rather than to the reader, so changing them must not throw away a loaded file. **The two near-identical unit bodies collapse** into `_XDFUnitBase`, which owns the playback clock and end-of-file handling. Both publishers are named `produce`, which matters: ezmsg collects publishers per attribute, so keeping the old `pub_chunk`/`pub_multi` names left `BaseProducerUnit.produce` running alongside them -- two publishers draining one producer, neither stopping. That hung the new graph tests until the names matched. **Templates are built in one place.** `_build_template` and `_with_time` replace the copy in each iterator, so the chunk_dim declaration and the fingerprint priming exist once rather than twice. `XDFIterator` itself is unchanged except for a public `exhausted` property and `print` becoming `ez.logger.info` -- a library writing to stdout is a defect regardless. The keyword API is preserved: `BaseProducer.__init__` forwards `**kwargs` to the settings type, so `XDFAxisArrayIterator(filepath=..., select=...)` still works, alongside `settings=`. `XDFMultiIteratorUnitSettings` is kept as an alias of the settings type that moved to `iter.py`. 27 tests pass, up from 20. The seven new ones cover the producer contract (lazy open, off-loop load, both construction forms, which settings reopen the file) and both units end to end in a real `ez.run` graph -- which is what caught the publisher-name collision. The async test is driven with `asyncio.run` rather than written `async def`: the pytest config names `asyncio_mode` but pytest-asyncio is not installed, so an async test would be silently never awaited and would pass without running. --- pyproject.toml | 3 +- src/ezmsg/xdf/iter.py | 374 +++++++++++++++++++++++++--------------- src/ezmsg/xdf/source.py | 178 +++++++++---------- tests/test_iter.py | 121 ++++++++++++- 4 files changed, 435 insertions(+), 241 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 02d0879..45243a7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -10,6 +10,7 @@ dynamic = ["version"] dependencies = [ # 3.10.0b2 for AxisArray.chunk_dim and CoordinateAxis.fingerprint. "ezmsg>=3.10.0b2", + "ezmsg-baseproc>=1.11.0", "numpy>=2.0.2", "pyxdf>=1.16.8", ] @@ -63,7 +64,7 @@ select = ["E", "F", "I", "W"] [tool.ruff.lint.isort] known-first-party = ["ezmsg.xdf"] -known-third-party = ["ezmsg"] +known-third-party = ["ezmsg", "ezmsg.baseproc"] [tool.uv.sources] # Uncomment to use development version of ezmsg from git diff --git a/src/ezmsg/xdf/iter.py b/src/ezmsg/xdf/iter.py index 3d84766..e204351 100644 --- a/src/ezmsg/xdf/iter.py +++ b/src/ezmsg/xdf/iter.py @@ -1,9 +1,16 @@ +import asyncio +import os import queue +import typing +from dataclasses import field from pathlib import Path +import ezmsg.core as ez import numpy as np import numpy.typing as npt import pyxdf +from ezmsg.baseproc.protocols import processor_state +from ezmsg.baseproc.stateful import BaseStatefulProducer from ezmsg.util.messages.axisarray import AxisArray from ezmsg.util.messages.util import replace @@ -143,7 +150,7 @@ def _scan_file(self): self._streams = [self._streams[stream_names.index(_)] for _ in self._select] self._metadata = {k: self._metadata[k] for k in self._select} - print( + ez.logger.info( f"Imported {len(self._streams)} streams from {self._filepath} " f"spanning {xdf_dur:.2f} s beginning at t={xdf_t0:.2f}." ) @@ -156,12 +163,17 @@ def stream_meta(self) -> list[dict] | dict: def n_chunks(self) -> int: return self._n_chunks + @property + def exhausted(self) -> bool: + """True once every chunk boundary has been handed out.""" + return self._chunk_ix >= self._n_chunks + def __iter__(self): self._chunk_ix = 0 return self def __next__(self) -> dict[str, tuple[npt.NDArray, npt.NDArray]]: - if self._chunk_ix >= self.n_chunks: + if self.exhausted: raise StopIteration else: out_dict = {} @@ -190,156 +202,236 @@ def labels_from_strm(strm: dict) -> list[str]: return labels -class XDFAxisArrayIterator(XDFIterator): - def __init__(self, *args, select: str, **kwargs): - """ - This Iterator loads only a single stream and yields a single :obj:`AxisArray` object per chunk. +def _build_template(stream: dict, name: str, n_ch: int, fs: float) -> AxisArray: + """The message every chunk of *stream* is a `replace` of. - Args: - *args: - select: Unlike :obj:`XDFIterator`, this must be a single string, the name of the stream to select. - **kwargs: + Built once per stream so the `ch` axis object -- and the fingerprint cached + on it -- is shared by every message, which is what makes priming cheap. + """ + labels = labels_from_strm(stream) + time_ax = ( + AxisArray.TimeAxis(fs=fs, offset=0.0) + if fs + else 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 + return AxisArray( + data=np.zeros((0, n_ch), dtype=stream["time_series"].dtype), + dims=["time", "ch"], + axes={"time": time_ax, "ch": ch_ax}, + key=name, + # Messages accumulate along `time`, whether the stream is regular or + # carries per-sample timestamps; `ch` describes the stream itself. + chunk_dim="time", + ) + + +def _with_time(template: AxisArray, data: npt.NDArray, tvec: npt.NDArray, fallback_t: float) -> AxisArray: + """A chunk message: the template's data replaced, and its time axis advanced. + + An irregular stream carries every timestamp; a regular one carries only where + the chunk starts, since its gain says the rest. + """ + time_ax = template.axes["time"] + if isinstance(time_ax, AxisArray.CoordinateAxis): + t_kwargs = {"data": tvec if len(tvec) else np.array([])} + else: + t_kwargs = {"offset": tvec[0] if len(tvec) else fallback_t} + return replace( + template, + data=data, + axes={**template.axes, "time": replace(time_ax, **t_kwargs)}, + ) + + +class XDFIteratorSettings(ez.Settings): + """Settings shared by both AxisArray iterators. + + ``playback_rate`` and ``self_terminating`` belong to the unit rather than to + the reader, and are listed in :attr:`NONRESET_SETTINGS_FIELDS` so changing + either does not reopen the file. + """ + + filepath: typing.Union[os.PathLike, str] + select: str = "" + chunk_dur: float = 1.0 + start_time: float | None = None + stop_time: float | None = None + rezero: bool = True + playback_rate: float | None = None + self_terminating: bool = False + """ + If True, the unit will raise a :obj:`ez.NormalTermination` exception when the file is exhausted. + 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. + """ + + +class XDFMultiIteratorSettings(XDFIteratorSettings): + select: set[str] | None = None + force_single_sample: set = field(default_factory=set) + + +@processor_state +class XDFIteratorState: + reader: XDFIterator | None = None + template: AxisArray | None = None + + +@processor_state +class XDFMultiIteratorState: + reader: XDFIterator | None = None + templates: dict | None = None + pubqueue: queue.SimpleQueue | None = None + + +class _XDFProducerBase: + """Shared plumbing for the two AxisArray producers. + + The file load is deliberately *not* run from ``__init__``. ezmsg-neo and + ezmsg-nwb both reset eagerly there and so pay the whole open on the event + loop during ``initialize``; here the first ``__acall__`` triggers + ``_areset_state``, which puts it on a worker thread. Nothing in this + package's public surface reads stream metadata before the first chunk, so + there is nothing to lose by waiting. + """ + + NONRESET_SETTINGS_FIELDS = frozenset({"playback_rate", "self_terminating"}) + + async def _areset_state(self) -> None: + """Offload the sync open onto a worker thread. + + ``pyxdf.load_xdf`` reads and decodes the entire file, which for a + several-hundred-megabyte recording is seconds of pure CPU and I/O. On the + event loop that stalls every other unit in the process. """ - kwargs["select"] = set((select,)) - super().__init__(*args, **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) - 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": 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", + await asyncio.to_thread(self._reset_state) + + def _build_reader(self, select: set[str] | None) -> XDFIterator: + return XDFIterator( + filepath=self.settings.filepath, + select=select, + chunk_dur=self.settings.chunk_dur, + start_time=self.settings.start_time, + stop_time=self.settings.stop_time, + rezero=self.settings.rezero, ) + +class XDFAxisArrayIterator( + _XDFProducerBase, + BaseStatefulProducer[XDFIteratorSettings, AxisArray, XDFIteratorState], +): + """Loads a single stream and produces one :obj:`AxisArray` per chunk. + + ``select`` must be a single stream name, unlike :obj:`XDFIterator`. + """ + + @property + def exhausted(self) -> bool: + reader = self._state.reader + return reader is not None and reader.exhausted + + def _reset_state(self) -> None: + reader = self._build_reader({self.settings.select}) + meta = reader.stream_meta[self.settings.select] + self._state.reader = reader + self._state.template = _build_template( + reader._streams[0], + name=reader._streams[0]["info"]["name"][0], + n_ch=meta["channel_count"], + fs=meta["nominal_srate"], + ) + + async def _produce(self) -> AxisArray | None: + reader = self._state.reader + try: + chunk_dict = next(reader) + except StopIteration: + return None + data, tvec = chunk_dict.get(self.settings.select, (None, None)) + if data is None: + return None + return _with_time(self._state.template, data, tvec, reader._last_time) + def __next__(self) -> AxisArray: - result: AxisArray | None = None - chunk_dict = super().__next__() - # Should only be 1 in self._select. If there are more then we overwrite with the last. - for strm_name in self._select: - if strm_name in chunk_dict: - data, tvec = chunk_dict[strm_name] - if isinstance(self._template.axes["time"], AxisArray.CoordinateAxis): - t_kwargs = {"data": tvec} - else: - t_kwargs = {"offset": tvec[0] if len(tvec) else self._last_time} - result = replace( - self._template, - data=data, - axes={ - **self._template.axes, - "time": replace( - self._template.axes["time"], - **t_kwargs, - ), - }, - ) + result = self() + if result is None: + raise StopIteration return result -class XDFMultiAxArrIterator(XDFIterator): - def __init__(self, *args, force_single_sample: set = set(), **kwargs): - """ - This Iterator loads multiple streams and yields a :obj:`AxisArray` object per iteration, - but the stream source might different between chunks. +class XDFMultiAxArrIterator( + _XDFProducerBase, + BaseStatefulProducer[XDFMultiIteratorSettings, AxisArray, XDFMultiIteratorState], +): + """Loads multiple streams and produces one :obj:`AxisArray` per iteration. - Args: - *args: - force_single_sample: Use this to identify irregular-rate streams that might conceivably have more than one - event within the defined chunk_dur, for which :obj:`AxisArray` cannot represent timestamps properly. - **kwargs: - """ - super().__init__(*args, **kwargs) - self._force_single_sample = force_single_sample - stream_names = [_["info"]["name"][0] for _ in self._streams] - - # Create template messages for each stream - self._templates = {} - for stream_name, stream_meta in self._metadata.items(): - stream = self._streams[stream_names.index(stream_name)] - labels = labels_from_strm(stream) - fs = stream_meta["nominal_srate"] - time_ax = ( - AxisArray.TimeAxis(fs=fs, offset=0.0) - 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": ch_ax, - }, - key=stream_name, - chunk_dim="time", + Which stream a given message came from varies; read ``.key``. Returns + ``None`` when a chunk held nothing for any stream, and raises + ``StopIteration`` only once the file is done. + + ``force_single_sample`` names irregular-rate streams that may carry more than + one event within ``chunk_dur``, which :obj:`AxisArray` cannot represent as a + single message with correct timestamps; those are split one event per + message. + """ + + @property + def exhausted(self) -> bool: + reader = self._state.reader + if reader is None: + return False + return reader.exhausted and self._state.pubqueue.empty() + + def _reset_state(self) -> None: + reader = self._build_reader(self.settings.select) + stream_names = [_["info"]["name"][0] for _ in reader._streams] + self._state.reader = reader + self._state.pubqueue = queue.SimpleQueue() + self._state.templates = { + name: _build_template( + reader._streams[stream_names.index(name)], + name=name, + n_ch=meta["channel_count"], + fs=meta["nominal_srate"], ) - self._pubqueue: queue.SimpleQueue[AxisArray] = queue.SimpleQueue() + for name, meta in reader.stream_meta.items() + } - def __next__(self) -> AxisArray | None: - if self._pubqueue.empty(): - chunk_dict = super().__next__() - for k, template in self._templates.items(): - if k in chunk_dict and len(chunk_dict[k][1]) > 0: - data, tvec = chunk_dict[k] - if k in self._force_single_sample: - if isinstance(template.axes["time"], AxisArray.CoordinateAxis): - t_kwargs = {"data": np.array([])} - else: - t_kwargs = {"offset": 0.0} - for ix, _t in enumerate(tvec): - if "data" in t_kwargs: - t_kwargs["data"] = np.array([_t]) - else: - t_kwargs["offset"] = _t - self._pubqueue.put_nowait( - replace( - template, - data=data[ix : ix + 1], - axes={ - **template.axes, - "time": replace(template.axes["time"], **t_kwargs), - }, - ) - ) - else: - 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} - self._pubqueue.put_nowait( - replace( - template, - data=data, - axes={ - **template.axes, - "time": replace( - template.axes["time"], - **t_kwargs, - ), - }, - ) - ) + def _enqueue_chunk(self, chunk_dict: dict) -> None: + reader = self._state.reader + for name, template in self._state.templates.items(): + if name not in chunk_dict or len(chunk_dict[name][1]) == 0: + continue + data, tvec = chunk_dict[name] + if name in self.settings.force_single_sample: + for ix, stamp in enumerate(tvec): + self._state.pubqueue.put_nowait(_with_time(template, data[ix : ix + 1], np.array([stamp]), stamp)) + else: + self._state.pubqueue.put_nowait(_with_time(template, data, tvec, reader._last_time)) + + async def _produce(self) -> AxisArray | None: + if self._state.pubqueue.empty(): + try: + self._enqueue_chunk(next(self._state.reader)) + except StopIteration: + return None try: - return self._pubqueue.get_nowait() + return self._state.pubqueue.get_nowait() except queue.Empty: return None + + def __next__(self) -> AxisArray | None: + if self.exhausted: + raise StopIteration + return self() diff --git a/src/ezmsg/xdf/source.py b/src/ezmsg/xdf/source.py index db2bd7e..bc4a12e 100644 --- a/src/ezmsg/xdf/source.py +++ b/src/ezmsg/xdf/source.py @@ -1,14 +1,30 @@ import asyncio -import os import time import typing -from dataclasses import field import ezmsg.core as ez -from ezmsg.util.generator import GenState +from ezmsg.baseproc.units import BaseProducerUnit from ezmsg.util.messages.axisarray import AxisArray -from .iter import XDFAxisArrayIterator, XDFMultiAxArrIterator +from .iter import ( + XDFAxisArrayIterator, + XDFIteratorSettings, + XDFMultiAxArrIterator, + XDFMultiIteratorSettings, +) + +# The settings types moved to `iter.py` so the producers can own them, but they +# were importable from here first. +XDFMultiIteratorUnitSettings = XDFMultiIteratorSettings + +__all__ = [ + "PlaybackClock", + "XDFIteratorSettings", + "XDFIteratorUnit", + "XDFMultiIteratorSettings", + "XDFMultiIteratorUnit", + "XDFMultiIteratorUnitSettings", +] class PlaybackClock: @@ -48,109 +64,75 @@ def step(self) -> None: time.sleep(self._get_duration()) -class XDFIteratorSettings(ez.Settings): - filepath: typing.Union[os.PathLike, str] - select: str - chunk_dur: float = 1.0 - start_time: float | None = None - stop_time: float | None = None - rezero: bool = True - playback_rate: float | None = None - self_terminating: bool = False - """ - If True, the unit will raise a :obj:`ez.NormalTermination` exception when the file is exhausted. - 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. - """ +class _XDFUnitBase: + """Playback pacing and end-of-file handling, shared by both units. + Note both subclasses name their publisher ``produce``: that is the name + ``BaseProducerUnit`` uses, and ezmsg collects publishers per attribute, so a + differently named one would run *alongside* the base class\'s rather than + replacing it -- two publishers draining one producer, neither stopping. -class XDFIteratorUnit(ez.Unit): - STATE = GenState - SETTINGS = XDFIteratorSettings + The producer supplies chunks as fast as they can be sliced out of memory; + ``playback_rate`` is what turns that into a paced stream, and it is a + property of the unit rather than of the reader. + """ - OUTPUT_SIGNAL = ez.OutputStream(AxisArray) OUTPUT_TERM = ez.OutputStream(typing.Any) - def initialize(self) -> None: - self.construct_generator() - - def construct_generator(self): - self.STATE.gen = XDFAxisArrayIterator( - filepath=self.SETTINGS.filepath, - select=self.SETTINGS.select, - chunk_dur=self.SETTINGS.chunk_dur, - start_time=self.SETTINGS.start_time, - stop_time=self.SETTINGS.stop_time, - rezero=self.SETTINGS.rezero, + async def initialize(self) -> None: + await super().initialize() + self._clock = ( + PlaybackClock(rate=self.SETTINGS.playback_rate, step_dur=self.SETTINGS.chunk_dur) + if self.SETTINGS.playback_rate is not None + else None ) - if self.SETTINGS.playback_rate is not None: - self._clock = PlaybackClock(rate=self.SETTINGS.playback_rate, step_dur=self.SETTINGS.chunk_dur) - else: - self._clock = None - @ez.publisher(OUTPUT_SIGNAL) - async def pub_chunk(self) -> typing.AsyncGenerator: - try: - while True: - if self._clock is not None: - await self._clock.astep() - msg = next(self.STATE.gen) - if msg.data.size > 0: - yield self.OUTPUT_SIGNAL, msg - else: - await asyncio.sleep(0) - except StopIteration: - 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 - - -class XDFMultiIteratorUnitSettings(XDFIteratorSettings): - select: set[str] | None = None # Override with a default - force_single_sample: set = field(default_factory=set) - - -class XDFMultiIteratorUnit(ez.Unit): - STATE = GenState - SETTINGS = XDFMultiIteratorUnitSettings + async def _finish(self) -> typing.AsyncGenerator: + 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 + + +class XDFIteratorUnit( + _XDFUnitBase, + BaseProducerUnit[XDFIteratorSettings, AxisArray, XDFAxisArrayIterator], +): + SETTINGS = XDFIteratorSettings OUTPUT_SIGNAL = ez.OutputStream(AxisArray) - OUTPUT_TERM = ez.OutputStream(typing.Any) - def initialize(self) -> None: - self.construct_generator() - - def construct_generator(self): - self.STATE.gen = XDFMultiAxArrIterator( - filepath=self.SETTINGS.filepath, - select=self.SETTINGS.select, - chunk_dur=self.SETTINGS.chunk_dur, - start_time=self.SETTINGS.start_time, - stop_time=self.SETTINGS.stop_time, - rezero=self.SETTINGS.rezero, - force_single_sample=self.SETTINGS.force_single_sample, - ) - if self.SETTINGS.playback_rate is not None: - self._clock = PlaybackClock(rate=self.SETTINGS.playback_rate, step_dur=self.SETTINGS.chunk_dur) - else: - self._clock = None + @ez.publisher(OUTPUT_SIGNAL) + async def produce(self) -> typing.AsyncGenerator: + while not self.producer.exhausted: + if self._clock is not None: + await self._clock.astep() + msg = await self.producer.__acall__() + if msg is not None and msg.data.size > 0: + yield self.OUTPUT_SIGNAL, msg + else: + await asyncio.sleep(0) + async for out in self._finish(): + yield out + + +class XDFMultiIteratorUnit( + _XDFUnitBase, + BaseProducerUnit[XDFMultiIteratorSettings, AxisArray, XDFMultiAxArrIterator], +): + SETTINGS = XDFMultiIteratorSettings + + OUTPUT_SIGNAL = ez.OutputStream(AxisArray) @ez.publisher(OUTPUT_SIGNAL) - async def pub_multi(self) -> typing.AsyncGenerator: - try: - while True: - if self._clock is not None: - await self._clock.astep() - msg = next(self.STATE.gen) - if msg is not None: - yield self.OUTPUT_SIGNAL, msg - else: - await asyncio.sleep(0) - except StopIteration: - 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 + async def produce(self) -> typing.AsyncGenerator: + while not self.producer.exhausted: + if self._clock is not None: + await self._clock.astep() + msg = await self.producer.__acall__() + if msg is not None: + yield self.OUTPUT_SIGNAL, msg + else: + await asyncio.sleep(0) + async for out in self._finish(): + yield out diff --git a/tests/test_iter.py b/tests/test_iter.py index 708942c..deb67c3 100644 --- a/tests/test_iter.py +++ b/tests/test_iter.py @@ -10,12 +10,21 @@ import math import pickle +import ezmsg.core as ez import numpy as np import pytest from conftest import EEG_STREAM, MARKER_STREAM from ezmsg.util.messages.axisarray import AxisArray +from ezmsg.util.messages.util import replace as replace_settings -from ezmsg.xdf.iter import XDFAxisArrayIterator, XDFIterator, XDFMultiAxArrIterator +from ezmsg.xdf.iter import ( + XDFAxisArrayIterator, + XDFIterator, + XDFIteratorSettings, + XDFMultiAxArrIterator, + XDFMultiIteratorSettings, +) +from ezmsg.xdf.source import XDFIteratorUnit, XDFMultiIteratorUnit def eeg_messages(path, **kwargs) -> list[AxisArray]: @@ -185,3 +194,113 @@ 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 + + +class TestTheProducerContract: + """The producers are `BaseStatefulProducer`s, which means the file open is a + state reset rather than construction work.""" + + def test_construction_does_not_read_the_file(self, test_xdf_path): + """Deliberately unlike ezmsg-neo and ezmsg-nwb, which reset eagerly in + ``__init__`` and so pay the whole open on the event loop during + ``initialize``. Nothing here reads stream metadata before the first + chunk, so there is nothing to lose by waiting.""" + it = XDFAxisArrayIterator(filepath=test_xdf_path, select=EEG_STREAM.name) + assert it._state.reader is None + assert next(it) is not None + assert it._state.reader is not None + + def test_the_file_load_runs_off_the_event_loop(self, test_xdf_path): + """``pyxdf.load_xdf`` reads and decodes the whole file. On the event loop + that stalls every other unit in the process. + + Driven with ``asyncio.run`` rather than an async test, because this repo + has no pytest-asyncio -- the pytest config names ``asyncio_mode`` but the + plugin is not installed, so an ``async def`` test is silently never + awaited and passes without running. + """ + import asyncio + import threading + + seen: list[int] = [] + + class Spy(XDFAxisArrayIterator): + def _reset_state(self): + seen.append(threading.get_ident()) + super()._reset_state() + + producer = Spy(filepath=test_xdf_path, select=EEG_STREAM.name) + assert not seen, "construction should not have opened anything" + + loop_tid: list[int] = [] + + async def drive(): + loop_tid.append(threading.get_ident()) + await producer.__acall__() + + asyncio.run(drive()) + + assert len(seen) == 1 + assert seen[0] != loop_tid[0], "_reset_state ran on the event-loop thread" + + def test_settings_arrive_as_keywords_or_as_a_settings_object(self, test_xdf_path): + by_kwargs = XDFAxisArrayIterator(filepath=test_xdf_path, select=EEG_STREAM.name, chunk_dur=0.5) + by_settings = XDFAxisArrayIterator( + settings=XDFIteratorSettings(filepath=test_xdf_path, select=EEG_STREAM.name, chunk_dur=0.5) + ) + assert by_kwargs.settings == by_settings.settings + + def test_pacing_settings_do_not_reopen_the_file(self, test_xdf_path): + """``playback_rate`` and ``self_terminating`` belong to the unit, so + changing either must not throw away a loaded file.""" + it = XDFAxisArrayIterator(filepath=test_xdf_path, select=EEG_STREAM.name) + next(it) + reader = it._state.reader + it.update_settings(replace_settings(it.settings, playback_rate=2.0, self_terminating=True)) + next(it) + assert it._state.reader is reader + + def test_changing_the_file_does_reopen(self, test_xdf_path): + it = XDFAxisArrayIterator(filepath=test_xdf_path, select=EEG_STREAM.name) + next(it) + reader = it._state.reader + it.update_settings(replace_settings(it.settings, chunk_dur=0.25)) + next(it) + assert it._state.reader is not reader + + +class TestTheUnitsInAGraph: + @staticmethod + def _run(unit_cls, settings) -> list[AxisArray]: + collected: list[AxisArray] = [] + + class Collector(ez.Unit): + INPUT_SIGNAL = ez.InputStream(AxisArray) + + @ez.subscriber(INPUT_SIGNAL) + async def on_msg(self, msg: AxisArray) -> None: + collected.append(msg) + + src, sink = unit_cls(settings), Collector() + ez.run(SRC=src, SINK=sink, connections=((src.OUTPUT_SIGNAL, sink.INPUT_SIGNAL),)) + return collected + + def test_the_single_stream_unit_publishes_the_whole_file(self, test_xdf_path): + msgs = self._run( + XDFIteratorUnit, + XDFIteratorSettings(filepath=test_xdf_path, select=EEG_STREAM.name, self_terminating=True), + ) + assert msgs, "no messages published" + assert sum(m.data.shape[0] for m in msgs) == EEG_STREAM.n_samples + assert all(m.chunk_dim == "time" for m in msgs) + assert all("_fingerprint" in m.axes["ch"].__dict__ for m in msgs) + + def test_the_multi_stream_unit_publishes_both_streams(self, test_xdf_path): + msgs = self._run( + XDFMultiIteratorUnit, + XDFMultiIteratorSettings(filepath=test_xdf_path, self_terminating=True), + ) + assert {m.key for m in msgs} == {EEG_STREAM.name, MARKER_STREAM.name} + eeg = [m for m in msgs if m.key == EEG_STREAM.name] + assert sum(m.data.shape[0] for m in eeg) == EEG_STREAM.n_samples + assert all(m.chunk_dim == "time" for m in msgs) From 1bfde0597dbc9e1eba7586691d5c94df0063ec1f Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Fri, 4 Sep 2026 06:09:21 +0000 Subject: [PATCH 8/8] Initial plan