Coverage for tests/io/test_prefetch.py: 100%
106 statements
« prev ^ index » next coverage.py v7.15.1, created at 2026-07-18 05:24 +0000
« prev ^ index » next coverage.py v7.15.1, created at 2026-07-18 05:24 +0000
1"""Tests for asynchronous (prefetching) iteration: EventReader(async_read=True).
3The iterator must be byte-identical to synchronous iteration, propagate
4worker exceptions, survive early exits without deadlocking, and guard the
5reader against concurrent direct reads.
6"""
7import numpy as np
8import pytest
10from evutils.io import EventReader, EventWriter
11from evutils.io._prefetch import PrefetchIterator
12from evutils.types import Event_dtype
15from typing import Any
16@pytest.fixture()
17def raw_file(tmp_path: Any) -> Any:
18 rng = np.random.default_rng(7)
19 n = 200_000
20 ev = np.zeros(n, dtype=Event_dtype)
21 ev["t"] = np.sort(rng.integers(0, 1_000_000, n))
22 ev["x"] = rng.integers(0, 1280, n)
23 ev["y"] = rng.integers(0, 720, n)
24 ev["p"] = rng.integers(0, 2, n)
25 p = tmp_path / "events.raw"
26 with EventWriter(p) as w:
27 w.write(ev)
28 return p, ev
31@pytest.mark.parametrize("suffix", [".raw", ".npz"])
32def test_async_matches_sync(tmp_path: Any, raw_file: Any, suffix: str) -> None:
33 """Async iteration yields the same windows as sync, for a native (C)
34 decoder and a pure-Python one."""
35 p, ev = raw_file
36 if suffix == ".npz":
37 p2 = tmp_path / "events.npz"
38 with EventWriter(p2) as w:
39 w.write(ev)
40 p = p2
42 with EventReader(p, n_events=30_000) as r:
43 sync_chunks = [np.asarray(c).copy() for c in r]
44 with EventReader(p, n_events=30_000, async_read=True) as r:
45 async_chunks = [np.asarray(c) for c in r]
47 assert len(sync_chunks) == len(async_chunks)
48 for s, a in zip(sync_chunks, async_chunks):
49 assert np.array_equal(s, a)
52def test_async_ext_triggers(raw_file: Any) -> None:
53 """(events, triggers) tuples pass through the prefetch queue unchanged."""
54 p, _ = raw_file
55 with EventReader(p, n_events=50_000, ext_trigger=True, async_read=True) as r:
56 for ev, tr in r:
57 assert len(ev) > 0
58 assert hasattr(tr, "id")
61def test_direct_read_guarded_while_iterating(raw_file: Any) -> None:
62 p, _ = raw_file
63 with EventReader(p, n_events=50_000, async_read=True) as r:
64 it = iter(r)
65 next(it)
66 with pytest.raises(RuntimeError, match="asynchronous iterator is active"):
67 r.read()
68 with pytest.raises(RuntimeError, match="asynchronous iterator is active"):
69 r.read_all()
70 # After exhausting the iterator, direct reads are allowed again.
71 for _ in it:
72 pass
73 assert len(r.read_all()) == 0 # EOF, but no guard error
76def test_early_break_and_reset(raw_file: Any) -> None:
77 """Breaking out of async iteration must not deadlock; reset() cancels the
78 worker and a fresh (async) iteration sees the whole file again."""
79 p, ev = raw_file
80 with EventReader(p, n_events=10_000, async_read=True) as r:
81 for i, _ in enumerate(r):
82 if i == 1:
83 break
84 r.reset() # cancels the active iterator
85 total = sum(len(c) for c in r)
86 assert total == len(ev)
89def test_close_with_active_iterator(raw_file: Any) -> None:
90 p, _ = raw_file
91 r = EventReader(p, n_events=10_000, async_read=True)
92 it = iter(r)
93 next(it)
94 r.close() # must join the worker and not raise
95 assert r._active_prefetch is None
98def test_only_one_active_iterator(raw_file: Any) -> None:
99 p, _ = raw_file
100 with EventReader(p, n_events=50_000, async_read=True) as r:
101 it = iter(r)
102 next(it)
103 with pytest.raises(RuntimeError, match="asynchronous iterator is active"):
104 iter(r)
105 it.close()
106 r.seek(n=0)
107 assert sum(len(c) for c in r) > 0 # a new iterator is fine now
110def test_worker_exception_propagates() -> None:
111 """An exception raised by the source surfaces in the consumer thread."""
112 def broken() -> Any:
113 yield 1
114 yield 2
115 raise ValueError("decoder blew up")
117 it = PrefetchIterator(broken())
118 assert next(it) == 1
119 assert next(it) == 2
120 with pytest.raises(ValueError, match="decoder blew up"):
121 next(it)
124def test_prefetch_iterator_bounded() -> None:
125 """The queue never buffers more than `depth` chunks ahead."""
126 import time
128 produced = []
130 def source() -> Any:
131 for i in range(100):
132 produced.append(i)
133 yield i
135 it = PrefetchIterator(source(), depth=2)
136 time.sleep(0.3) # give the worker every chance to run ahead
137 # depth chunks buffered + one in the worker's hand at most
138 assert len(produced) <= 2 + 1
139 assert list(it) == list(range(100))
142def test_prefetch_iterator_close_idempotent() -> None:
143 it = PrefetchIterator(iter(range(10)), depth=1)
144 next(it)
145 it.close()
146 it.close()
147 with pytest.raises(StopIteration):
148 next(it)