#!/usr/bin/env python3
"""Compare causal adaptive-lag forward pairs with fixed-lag M2/BTC baselines."""

from __future__ import annotations

import argparse
import sys
from pathlib import Path

import pandas as pd

REPO_ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(REPO_ROOT))

from strategies.adaptive_lag_correlation import LagCorrelationConfig, Transform, align_to_regular_grid
from strategies.adaptive_lag_walkforward import (
    WalkForwardConfig,
    evaluate_adaptive_lag,
    evaluate_fixed_lag,
    moving_block_bootstrap_correlation,
)


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--csv", type=Path, default=Path("data/COINBASE_BTCUSD, 360-IndicatorM2LI_DebugArchive.csv"))
    parser.add_argument("--transform", choices=[item.value for item in Transform], default="log_return")
    parser.add_argument("--min-lag", type=int, default=30)
    parser.add_argument("--max-lag", type=int, default=120)
    parser.add_argument("--window", type=int, default=365)
    parser.add_argument("--min-observations", type=int, default=240)
    parser.add_argument("--start-at", default="2020-01-01")
    parser.add_argument("--stride", type=int, default=7)
    parser.add_argument("--source-availability-delay", type=int, default=0)
    parser.add_argument("--fixed-lags", type=int, nargs="+", default=[0, 45, 72, 90])
    parser.add_argument("--bootstrap-block-length", type=int, default=14)
    parser.add_argument("--bootstrap-resamples", type=int, default=1_000)
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    frame = pd.read_csv(args.csv)
    timestamps = pd.to_datetime(frame["time"], unit="s", utc=True)
    source, target = align_to_regular_grid(
        pd.Series(frame["M2 US EU CN"].to_numpy(), index=timestamps),
        pd.Series(frame["close"].to_numpy(), index=timestamps),
    )
    transform = Transform(args.transform)
    estimator = LagCorrelationConfig(
        min_lag=args.min_lag, max_lag=args.max_lag, window=args.window,
        min_observations=args.min_observations, transform=transform,
    )
    adaptive = evaluate_adaptive_lag(source, target, WalkForwardConfig(
        estimator=estimator, evaluation_stride=args.stride, start_at=args.start_at,
        source_availability_delay=args.source_availability_delay,
    ))
    results = [("adaptive", adaptive)]
    results.extend((f"fixed-{lag}", evaluate_fixed_lag(
        source, target, lag=lag, transform=transform, stride=args.stride,
        start_at=args.start_at, source_availability_delay=args.source_availability_delay,
    )) for lag in args.fixed_lags)
    print(f"transform={transform.value} start={args.start_at} stride={args.stride} availability_delay={args.source_availability_delay}")
    print("method       n    correlation  95% block-bootstrap CI    hit-rate  mean-lag")
    for name, result in results:
        summary = result.summary
        interval = moving_block_bootstrap_correlation(
            result.records["source_value"], result.records["target_value"],
            block_length=args.bootstrap_block_length, resamples=args.bootstrap_resamples, seed=7,
        ) if summary.observations >= 2 else (float("nan"), float("nan"))
        print(f"{name:11} {summary.observations:4d} {summary.correlation:+.4f}      [{interval[0]:+.4f}, {interval[1]:+.4f}]   {summary.directional_hit_rate:.3f}    {summary.mean_lag:.1f}")


if __name__ == "__main__":
    main()
