"""Merge multiple snapshot directories and re-split them temporally.""" from __future__ import annotations import json from pathlib import Path from typing import Any import pyarrow.parquet as pq from libs.common.time_utils import utc_now from libs.export.snapshot_export import _rows_to_table, _temporal_split def export_merged_snapshot( *, source_snapshot_dirs: list[str | Path], output_dir: str | Path, snapshot_id: str, split_policy: str = "temporal_70_15_15", ) -> dict[str, Any]: rows: list[dict[str, Any]] = [] source_ids: list[str] = [] for source_dir in source_snapshot_dirs: source_path = Path(source_dir) source_ids.append(source_path.name) signal_origin = "continuation" if "cont_" in source_path.name else "event_day" for split in ("train", "valid", "test"): table = pq.read_table(source_path / f"{split}.parquet") for row in table.to_pylist(): merged_row = dict(row) merged_row.setdefault("signal_origin", signal_origin) rows.append(merged_row) splits = _temporal_split(rows, split_policy) out_path = Path(output_dir) / snapshot_id out_path.mkdir(parents=True, exist_ok=True) row_counts: dict[str, int] = {} for split_name, split_rows in splits.items(): pq.write_table(_rows_to_table(split_rows), out_path / f"{split_name}.parquet") row_counts[split_name] = len(split_rows) manifest = { "snapshot_id": snapshot_id, "created_at_utc": utc_now().isoformat(), "source_snapshot_ids": source_ids, "transform": "merged_snapshot", "split_policy": split_policy, "row_counts": row_counts, "total_rows": sum(row_counts.values()), "output_dir": str(out_path), } (out_path / "manifest.json").write_text(json.dumps(manifest, indent=2)) return manifest