from __future__ import annotations import json import pyarrow as pa import pyarrow.parquet as pq from libs.export.merged_snapshot import export_merged_snapshot def _write_snapshot(root, name: str, rows_by_split: dict[str, list[dict]]) -> None: path = root / name path.mkdir(parents=True, exist_ok=True) for split, rows in rows_by_split.items(): if rows: table = pa.Table.from_pylist(rows) else: table = pa.table({}) pq.write_table(table, path / f"{split}.parquet") (path / "manifest.json").write_text(json.dumps({"snapshot_id": name})) def test_export_merged_snapshot_reassigns_signal_origin_and_resplits(tmp_path): _write_snapshot( tmp_path, "base_snapshot", { "train": [{"event_id": "E1", "event_date": "2026-01-01", "entry_convention": "next_open_after_reaction_close"}], "valid": [], "test": [], }, ) _write_snapshot( tmp_path, "cont_snapshot", { "train": [{"event_id": "E1::cont_d3", "event_date": "2026-01-04", "entry_convention": "next_open_after_continuation_signal"}], "valid": [], "test": [], }, ) manifest = export_merged_snapshot( source_snapshot_dirs=[tmp_path / "base_snapshot", tmp_path / "cont_snapshot"], output_dir=tmp_path, snapshot_id="merged", ) assert manifest["total_rows"] == 2 merged_rows = [] for split in ("train", "valid", "test"): table = pq.read_table(tmp_path / "merged" / f"{split}.parquet") merged_rows.extend(table.to_pylist()) origins = {row["event_id"]: row["signal_origin"] for row in merged_rows} assert origins["E1"] == "event_day" assert origins["E1::cont_d3"] == "continuation"