"""Tests for multi-pair G10 pipeline (Option C). Tests the prepare_multipair.py merge logic and train.py multipair build(). Run: cd ~/dev/AI/jepa-fx-risk && .venv/bin/python -m pytest tests/test_multipair.py -v """ import importlib.util import numpy as np import pandas as pd import pytest import os def _import_mp(): spec = importlib.util.spec_from_file_location("prepare_multipair", "scripts/prepare_multipair.py") mod = importlib.util.module_from_spec(spec) spec.loader.exec_module(mod) return mod @pytest.fixture(scope="module") def mp(): return _import_mp() def _pair_df(start: str, n_hours: int, seed: int) -> pd.DataFrame: """Synthetic single-pair hourly parquet (same schema as prepare_hourly output).""" rng = np.random.default_rng(seed) dts = pd.date_range(start, periods=n_hours, freq="h") closes = 1.1 + np.cumsum(rng.normal(0, 0.001, n_hours)) return pd.DataFrame({ "datetime": dts, "close": closes, "ret": rng.normal(0, 0.001, n_hours), "realized_vol": np.abs(rng.normal(0.0005, 0.0001, n_hours)), }) # 1. merge_pair_parquets returns inner join on datetime def test_merge_inner_join(mp): eur = _pair_df("2020-01-01 00:00", 100, seed=1) # t0 to t0+99h gbp = _pair_df("2020-01-01 20:00", 60, seed=2) # t0+20 to t0+79h → 60 common result = mp.merge_pair_parquets({"eurusd": eur, "gbpusd": gbp}) assert len(result) == 60, f"expected 60 (inner join), got {len(result)}" # 2. merge_pair_parquets prefixes columns with pair name def test_merge_column_prefixes(mp): eur = _pair_df("2020-01-01 00:00", 50, seed=1) gbp = _pair_df("2020-01-01 00:00", 50, seed=2) result = mp.merge_pair_parquets({"eurusd": eur, "gbpusd": gbp}) assert "datetime" in result.columns, "datetime column missing" assert "eurusd_ret" in result.columns assert "eurusd_rv" in result.columns assert "gbpusd_ret" in result.columns assert "gbpusd_rv" in result.columns # raw pair columns should not leak through unprefixed assert "ret" not in result.columns assert "realized_vol" not in result.columns # 3. No NaN in merged output def test_merge_no_nan(mp): eur = _pair_df("2020-01-01 00:00", 50, seed=1) gbp = _pair_df("2020-01-01 00:00", 50, seed=2) result = mp.merge_pair_parquets({"eurusd": eur, "gbpusd": gbp}) nan_count = result.isnull().sum().sum() assert nan_count == 0, f"{nan_count} NaN values in merged output" # 4. PAIRS constant is a non-empty list starting with eurusd def test_pairs_constant(mp): assert hasattr(mp, "PAIRS"), "PAIRS constant missing from prepare_multipair.py" assert len(mp.PAIRS) >= 2, "PAIRS must have at least 2 pairs" assert mp.PAIRS[0] == "eurusd", "first pair must be eurusd (target pair)" # 5. merge target column is eurusd_rv (for build() target selection) def test_merge_has_eurusd_rv_as_target(mp): eur = _pair_df("2020-01-01 00:00", 50, seed=1) gbp = _pair_df("2020-01-01 00:00", 50, seed=2) result = mp.merge_pair_parquets({"eurusd": eur, "gbpusd": gbp}) assert "eurusd_rv" in result.columns, "eurusd_rv (target) missing from merged output" assert (result["eurusd_rv"] > 0).all(), "eurusd_rv should be positive" # 6. train.py recognises JEPA_USE_MULTIPAIR env var def test_use_multipair_knob(): import importlib.util as ilu spec = ilu.spec_from_file_location(f"train_mp_{id(None)}", "train.py") mod = ilu.module_from_spec(spec) saved = os.environ.get("JEPA_USE_MULTIPAIR") os.environ["JEPA_USE_MULTIPAIR"] = "1" try: spec.loader.exec_module(mod) finally: if saved is None: os.environ.pop("JEPA_USE_MULTIPAIR", None) else: os.environ["JEPA_USE_MULTIPAIR"] = saved assert hasattr(mod, "USE_MULTIPAIR"), "USE_MULTIPAIR knob missing from train.py" assert mod.USE_MULTIPAIR is True # 7. build() uses n_pairs*2 channels when multipair parquet present def test_build_uses_multipair_channels(): import importlib.util as ilu multipair_path = "data/processed/eurusd_multipair.parquet" if not os.path.exists(multipair_path): pytest.skip("eurusd_multipair.parquet not present — run data:prepare:multipair first") saved = os.environ.get("JEPA_USE_MULTIPAIR") os.environ["JEPA_USE_MULTIPAIR"] = "1" try: spec = ilu.spec_from_file_location(f"train_mp2_{id(None)}", "train.py") mod = ilu.module_from_spec(spec) spec.loader.exec_module(mod) (Xtr, _), _ = mod.build() finally: if saved is None: os.environ.pop("JEPA_USE_MULTIPAIR", None) else: os.environ["JEPA_USE_MULTIPAIR"] = saved mp = _import_mp() expected_ch = len(mp.PAIRS) * 2 assert Xtr.shape[2] == expected_ch, ( f"expected {expected_ch} channels (n_pairs={len(mp.PAIRS)}×2), got {Xtr.shape[2]}" )