"""Dataset Export: snapshot features + labels to Parquet.""" from __future__ import annotations import argparse import asyncio import json import yaml from libs.common.config import get_settings from libs.common.logging import configure_logging, get_logger from libs.db.session import get_session from libs.export.snapshot_export import export_dataset_snapshot logger = get_logger(__name__) async def run_dataset_export( snapshot_id: str | None, split_policy: str, output_dir: str, feature_versions: list[str] | None = None, label_version: str = "label-2.0.0", symbols: list[str] | None = None, ) -> dict: async with get_session() as session: manifest = await export_dataset_snapshot( session=session, snapshot_id=snapshot_id, split_policy=split_policy, output_dir=output_dir, feature_versions=feature_versions, label_version=label_version, symbols=symbols, ) return manifest def main() -> None: parser = argparse.ArgumentParser(description="Dataset Export") parser.add_argument("--snapshot-id", default=None, help="Snapshot ID (UUID, auto-generated if omitted)") parser.add_argument( "--split-policy", default="temporal_70_15_15", help="Split policy string (default: temporal_70_15_15)", ) parser.add_argument( "--output-dir", default="./data/datasets/snapshots", help="Output directory for Parquet files", ) parser.add_argument( "--feature-versions", nargs="+", default=None, help="Feature versions to merge (e.g. market_v1 event_v1 financial_v1)", ) parser.add_argument( "--label-version", default="label-2.0.0", help="Label version to export (default: label-2.0.0)", ) parser.add_argument( "--symbols-file", default=None, help="YAML file with 'symbols' list to filter export (e.g. configs/symbols_midcap.yaml)", ) parser.add_argument("--json", action="store_true", help="Print manifest JSON to stdout") args = parser.parse_args() settings = get_settings() configure_logging(settings.log_level) symbols = None if args.symbols_file: with open(args.symbols_file) as f: cfg = yaml.safe_load(f) or {} symbols = [s for s in cfg.get("symbols", []) if isinstance(s, str)] print(f"Filtering to {len(symbols)} symbols from {args.symbols_file}") manifest = asyncio.run( run_dataset_export( snapshot_id=args.snapshot_id, split_policy=args.split_policy, output_dir=args.output_dir, feature_versions=args.feature_versions, label_version=args.label_version, symbols=symbols, ) ) if args.json: print(json.dumps(manifest, indent=2)) else: print(f"Snapshot exported: {manifest['snapshot_id']}") print(f" Output: {manifest['output_dir']}") print(f" Total rows: {manifest['total_rows']}") for split, count in manifest["row_counts"].items(): print(f" {split}: {count} rows") if __name__ == "__main__": main()