#!/usr/bin/env python3 """Behavioural tests for the sras-viewer refactor. Covers what the golden-hash harness can't: the v6->v7 cache round-trip (including block carry-forward), parallel-vs-serial identity, the no-mask fast path, and the ROI bounding-box mask optimisation. Usage: python tools/test_refactor.py [--scratch DIR] """ import argparse import shutil import sys import tempfile from pathlib import Path import numpy as np sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) import sras_compute as compute # noqa: E402 from sras_compute import ( # noqa: E402 cache_file, compute_dc_image, compute_rf_image, dc_image_mv, ) from sras_format import CH3_IDX, CH4_IDX, SrasFile, adc_to_mv # noqa: E402 import tools.make_test_sras as gen # noqa: E402 _failures: list[str] = [] def check(name: str, ok: bool, detail: str = ""): print(f" {'PASS' if ok else 'FAIL'} {name}" + (f" — {detail}" if detail else "")) if not ok: _failures.append(name) def test_cache_roundtrip(scratch: Path): """v6 -> v7 for DC, then FFT, asserting the first block survives the second write (the carry-forward path in write_v7_cache).""" print("\ncache round-trip (v6 -> v7, both blocks)") path = scratch / "roundtrip.sras" gen.write(path, n_angles=3, seed=1, samples_per_frame=64) src = SrasFile(str(path)) check("source is v6", src.version == 6, f"got v{src.version}") expect_dc3 = [dc_image_mv(src, a, CH3_IDX) for a in range(src.n_angles)] expect_dc4 = [dc_image_mv(src, a, CH4_IDX) for a in range(src.n_angles)] expect_fft = [compute_rf_image(src, a, dc_threshold_mv=None, apply_bg_sub=True) for a in range(src.n_angles)] err = cache_file(str(path), "dc", True) check("dc cache_file succeeded", err == "", err) after_dc = SrasFile(str(path)) check("version flipped to 7", after_dc.version == 7, f"got v{after_dc.version}") check("dc3 stored for every angle", all(x is not None for x in after_dc.precomputed_dc3_mv)) check("dc3 values round-trip", all(np.allclose(after_dc.precomputed_dc3_mv[a], expect_dc3[a], atol=1e-4) for a in range(after_dc.n_angles))) check("dc4 values round-trip", all(np.allclose(after_dc.precomputed_dc4_mv[a], expect_dc4[a], atol=1e-4) for a in range(after_dc.n_angles))) check("no fft block yet", all(x is None for x in after_dc.precomputed_freq_mhz)) check("cached images are native float32", after_dc.precomputed_dc3_mv[0].dtype == np.float32 and after_dc.precomputed_dc3_mv[0].dtype.byteorder in ("=", "|"), str(after_dc.precomputed_dc3_mv[0].dtype.byteorder)) check("cached images are writable", after_dc.precomputed_dc3_mv[0].flags.writeable) err = cache_file(str(path), "fft", True) check("fft cache_file succeeded", err == "", err) both = SrasFile(str(path)) check("fft stored for every angle", all(x is not None for x in both.precomputed_freq_mhz)) check("fft values round-trip", all(np.allclose(both.precomputed_freq_mhz[a], expect_fft[a], atol=1e-3) for a in range(both.n_angles))) check("DC block carried forward through the FFT write", all(np.allclose(both.precomputed_dc3_mv[a], expect_dc3[a], atol=1e-4) for a in range(both.n_angles))) check("bg_sub flag persisted", both.precomputed_bg_sub is True) # The fast path must reproduce a fresh compute, and masking must still # apply on top of a cached (unmasked) image. fresh = SrasFile(str(path)) fresh.precomputed_freq_mhz = [None] * fresh.n_angles dc4 = dc_image_mv(both, 0, CH4_IDX) thr = float(np.median(dc4)) check("cached fast path == fresh compute (unmasked)", np.allclose(compute_rf_image(both, 0, dc_threshold_mv=None, apply_bg_sub=True), compute_rf_image(fresh, 0, dc_threshold_mv=None, apply_bg_sub=True), atol=1e-3)) check("cached fast path == fresh compute (masked)", np.allclose(compute_rf_image(both, 0, dc_threshold_mv=thr, apply_bg_sub=True), compute_rf_image(fresh, 0, dc_threshold_mv=thr, apply_bg_sub=True), atol=1e-3)) # Waveform data must be byte-identical to the pre-cache file. orig = scratch / "roundtrip_orig.sras" gen.write(orig, n_angles=3, seed=1, samples_per_frame=64) o, n = SrasFile(str(orig)), SrasFile(str(path)) check("waveform data untouched by the cache write", all(np.array_equal(np.asarray(o.data[a]), np.asarray(n.data[a])) for a in range(o.n_angles))) def test_partial_v7_cache(scratch: Path): """Only some angles cached: uncached angles must compute, not read zeros. This is the v5 bug the ragged normalisation fixed, checked via v7.""" print("\npartial cache (only some angles stored)") path = scratch / "partial.sras" gen.write(path, n_angles=3, seed=2, samples_per_frame=64) src = SrasFile(str(path)) expected = [compute_rf_image(src, a, dc_threshold_mv=None, apply_bg_sub=True) for a in range(src.n_angles)] partial = [expected[0], None, expected[2]] # angle 1 deliberately absent src.write_v7_cache(new_freq_mhz=partial, new_bg_sub=True) reread = SrasFile(str(path)) check("angle 1 is not cached", reread.precomputed_freq_mhz[1] is None) check("angles 0 and 2 are cached", reread.precomputed_freq_mhz[0] is not None and reread.precomputed_freq_mhz[2] is not None) img1 = compute_rf_image(reread, 1, dc_threshold_mv=None, apply_bg_sub=True) check("uncached angle computes rather than returning zeros", np.any(img1 != 0) and np.allclose(img1, expected[1], atol=1e-3)) def test_parallel_identity(scratch: Path): """Forcing 1 worker vs many must give identical output — catches chunk-boundary and race bugs.""" print("\nparallel vs serial identity") path = scratch / "parallel.sras" # Many rows, so the row loop actually splits into several chunks. n_rows, n_frames, spf = 48, 9, 256 gen.write(path, n_angles=1, seed=3, samples_per_frame=spf, geometry=[(n_rows, n_frames)]) sras = SrasFile(str(path)) saved_budget, saved_workers = compute._TOTAL_BYTES_BUDGET, compute._MAX_WORKERS try: # Shrink the budget so chunk_rows collapses to 1 and every row is # its own chunk — the worst case for boundary bugs. compute._TOTAL_BYTES_BUDGET = 8 * n_frames * spf * 4 compute._MAX_WORKERS = 1 dc_serial = compute_dc_image(sras, 0, CH4_IDX) rf_serial = compute_rf_image(sras, 0, dc_threshold_mv=None, apply_bg_sub=True) dc4 = adc_to_mv(dc_serial, *sras.cal(CH4_IDX)) thr = float(np.median(dc4)) rf_masked_serial = compute_rf_image(sras, 0, dc_threshold_mv=thr, apply_bg_sub=True) chunk_rows, n_workers = compute._plan_chunks( n_rows, n_frames, spf, live_multiplier=compute._FFT_LIVE_MULTIPLIER) check("serial plan uses 1 worker", n_workers == 1, f"chunk_rows={chunk_rows}") check("work actually splits into multiple chunks", chunk_rows < n_rows, f"chunk_rows={chunk_rows} of {n_rows} rows") compute._MAX_WORKERS = 8 chunk_rows, n_workers = compute._plan_chunks( n_rows, n_frames, spf, live_multiplier=compute._FFT_LIVE_MULTIPLIER) check("parallel plan uses >1 worker", n_workers > 1, f"chunk_rows={chunk_rows} workers={n_workers}") dc_par = compute_dc_image(sras, 0, CH4_IDX) rf_par = compute_rf_image(sras, 0, dc_threshold_mv=None, apply_bg_sub=True) rf_masked_par = compute_rf_image(sras, 0, dc_threshold_mv=thr, apply_bg_sub=True) check("dc image identical", np.array_equal(dc_serial, dc_par)) check("rf image identical (unmasked)", np.array_equal(rf_serial, rf_par)) check("rf image identical (masked)", np.array_equal(rf_masked_serial, rf_masked_par)) finally: compute._TOTAL_BYTES_BUDGET, compute._MAX_WORKERS = saved_budget, saved_workers def test_nomask_equals_low_threshold(scratch: Path): """dc_threshold_mv=None must equal a threshold below every pixel, while skipping the CH4 read.""" print("\nno-mask path") path = scratch / "nomask.sras" gen.write(path, n_angles=2, seed=4, samples_per_frame=128) sras = SrasFile(str(path)) for a in range(sras.n_angles): none_img = compute_rf_image(sras, a, dc_threshold_mv=None, apply_bg_sub=True) low_img = compute_rf_image(sras, a, dc_threshold_mv=-1e9, apply_bg_sub=True) check(f"angle {a}: None == -1e9 threshold", np.array_equal(none_img, low_img)) check(f"angle {a}: image is non-degenerate", len(np.unique(none_img)) > 1, f"{len(np.unique(none_img))} unique") def test_roi_mask(): """The bbox-restricted mask must equal a full-grid point-in-polygon test.""" print("\nROI mask (bbox fast path vs full grid)") from matplotlib.path import Path as MplPath from sras_viewer import RoiQuad rng = np.random.default_rng(0) x = np.linspace(-2.0, 3.0, 137) y = np.linspace(1.0, 4.0, 91) cases = { "axis-aligned rect": np.array([[0.0, 1.5], [1.0, 1.5], [1.0, 3.0], [0.0, 3.0]]), "skewed quad": np.array([[-0.5, 1.2], [1.7, 1.9], [1.2, 3.4], [-1.0, 2.6]]), "entirely outside": np.array([[8.0, 8.0], [9.0, 8.0], [9.0, 9.0], [8.0, 9.0]]), "covers whole grid": np.array([[-9.0, -9.0], [9.0, -9.0], [9.0, 9.0], [-9.0, 9.0]]), "straddles left edge": np.array([[-4.0, 2.0], [0.5, 2.0], [0.5, 3.0], [-4.0, 3.0]]), } for _ in range(5): cases[f"random {_}"] = rng.uniform([-2.5, 0.5], [3.5, 4.5], size=(4, 2)) for name, pts in cases.items(): roi = RoiQuad(pts) fast = roi.mask_for_grid(x, y) X, Y = np.meshgrid(x.astype(np.float64), y.astype(np.float64)) slow = MplPath(pts).contains_points( np.column_stack([X.ravel(), Y.ravel()])).reshape(X.shape) check(f"{name} ({int(slow.sum())} px inside)", np.array_equal(fast, slow)) # Descending y axis (images are stored top-down in some scans). roi = RoiQuad(cases["skewed quad"]) y_desc = y[::-1] fast = roi.mask_for_grid(x, y_desc) X, Y = np.meshgrid(x.astype(np.float64), y_desc.astype(np.float64)) slow = MplPath(cases["skewed quad"]).contains_points( np.column_stack([X.ravel(), Y.ravel()])).reshape(X.shape) check("descending y axis", np.array_equal(fast, slow)) def test_legacy_parse_and_average(scratch: Path): """v2-v4 parsing plus the sras_average.py rewrite (which now streams via SrasFile rather than slurping the whole file).""" import subprocess print("\nlegacy formats (v2-v4) and sras_average") repo = Path(__file__).resolve().parent.parent for version in (2, 3, 4): path = scratch / f"legacy_v{version}.sras" meta = gen.write_legacy(path, version=version, n_angles=2, n_rows=4, n_frames=12, samples_per_frame=32, seed=version) s = SrasFile(str(path)) check(f"v{version} parses", s.version == version, f"got v{s.version}") check(f"v{version} geometry uniform across angles", list(s.n_rows) == [4, 4] and list(s.n_frames) == [12, 12], f"rows={list(s.n_rows)} frames={list(s.n_frames)}") check(f"v{version} waveform data matches what was written", all(np.array_equal(np.asarray(s.data[a]), meta["data"][a]) for a in range(s.n_angles))) check(f"v{version} background {'present' if version >= 4 else 'absent'}", (s.background is not None) == (version >= 4)) check(f"v{version} precomputed stores are ragged lists", isinstance(s.precomputed_freq_mhz, list) and len(s.precomputed_freq_mhz) == s.n_angles) # DC image must equal a direct mean of the known input. expect = meta["data"][0][:, CH3_IDX, :, :].astype(np.float64).mean(axis=-1) check(f"v{version} DC image equals a direct mean", np.allclose(compute_dc_image(s, 0, CH3_IDX), expect, atol=1e-3)) src = scratch / "legacy_v4.sras" meta = gen.write_legacy(src, version=4, n_angles=2, n_rows=4, n_frames=12, samples_per_frame=32, seed=4) dst = scratch / "legacy_v4_avg.sras" if dst.exists(): dst.unlink() proc = subprocess.run( [sys.executable, str(repo / "sras_average.py"), str(src), str(dst), "--n", "4"], capture_output=True, text=True, cwd=repo) check("sras_average ran", proc.returncode == 0, (proc.stderr or proc.stdout).strip()[-200:]) if dst.exists(): avg = SrasFile(str(dst)) check("averaged file parses", avg.version == 4) check("frame count divided by 4", list(avg.n_frames) == [3, 3], f"{list(avg.n_frames)}") check("angles/rows/channels unchanged", avg.n_angles == 2 and list(avg.n_rows) == [4, 4] and avg.n_channels == meta["n_channels"]) check("calibration preserved", np.allclose(avg.ch_ymult_mv, SrasFile(str(src)).ch_ymult_mv)) check("background preserved", np.array_equal(avg.background, SrasFile(str(src)).background)) src_data = meta["data"] expect0 = src_data[0][:, :, 0:4, :].astype(np.float32).mean(axis=2).astype(np.int16) check("first averaged group equals the mean of its 4 source frames", np.array_equal(np.asarray(avg.data[0])[:, :, 0, :], expect0)) # Remainder handling: 12 frames / 5 -> 2 full groups + 1 partial. dst2 = scratch / "legacy_v4_avg5.sras" subprocess.run([sys.executable, str(repo / "sras_average.py"), str(src), str(dst2), "--n", "5"], capture_output=True, text=True, cwd=repo) if dst2.exists(): check("partial trailing group kept by default", list(SrasFile(str(dst2)).n_frames) == [3, 3], f"{list(SrasFile(str(dst2)).n_frames)}") dst3 = scratch / "legacy_v4_avg5d.sras" subprocess.run([sys.executable, str(repo / "sras_average.py"), str(src), str(dst3), "--n", "5", "--discard-remainder"], capture_output=True, text=True, cwd=repo) if dst3.exists(): check("--discard-remainder drops the partial group", list(SrasFile(str(dst3)).n_frames) == [2, 2], f"{list(SrasFile(str(dst3)).n_frames)}") def test_unsupported_version_reported(scratch: Path): """cache_file must report, not raise, for a file it can't handle.""" print("\nerror reporting") bogus = scratch / "bogus.sras" bogus.write_bytes(b"SRAS" + bytes([99]) + b"\x00" * 200) err = cache_file(str(bogus), "dc", True) check("bad version returns an error string", bool(err), err) missing = cache_file(str(scratch / "does_not_exist.sras"), "dc", True) check("missing file returns an error string", bool(missing), missing) def main(): p = argparse.ArgumentParser(description=__doc__) p.add_argument("--scratch") args = p.parse_args() tmp = None if args.scratch: scratch = Path(args.scratch) scratch.mkdir(parents=True, exist_ok=True) else: tmp = tempfile.mkdtemp(prefix="sras_test_") scratch = Path(tmp) try: test_cache_roundtrip(scratch) test_partial_v7_cache(scratch) test_parallel_identity(scratch) test_nomask_equals_low_threshold(scratch) test_roi_mask() test_legacy_parse_and_average(scratch) test_unsupported_version_reported(scratch) finally: if tmp: shutil.rmtree(tmp, ignore_errors=True) print() if _failures: print(f"{len(_failures)} FAILURE(S): " + ", ".join(_failures)) sys.exit(1) print("All checks passed.") if __name__ == "__main__": main()