#!/usr/bin/env python3
"""Print causal adaptive-lag diagnostics for the archived M2/BTC export."""

from __future__ import annotations

import argparse
import sys
from pathlib import Path

import pandas as pd

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

from strategies.adaptive_lag_correlation import (
    AdaptiveLagCorrelationEngine,
    LagCorrelationConfig,
    Transform,
    align_to_regular_grid,
)


DEFAULT_CSV = Path("data/COINBASE_BTCUSD, 360-IndicatorM2LI_DebugArchive.csv")


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--csv", type=Path, default=DEFAULT_CSV)
    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=180)
    parser.add_argument("--min-observations", type=int, default=120)
    parser.add_argument("--distinct-peak-exclusion-radius", type=int, default=7)
    parser.add_argument("--transform", choices=[item.value for item in Transform], default=Transform.LEVELS.value)
    parser.add_argument("--top", type=int, default=5)
    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 = pd.Series(frame["M2 US EU CN"].to_numpy(), index=timestamps, name="global_m2")
    target = pd.Series(frame["close"].to_numpy(), index=timestamps, name="btc_close")
    source, target = align_to_regular_grid(source, target, frequency="1D")
    engine = AdaptiveLagCorrelationEngine(
        LagCorrelationConfig(
            min_lag=args.min_lag,
            max_lag=args.max_lag,
            window=args.window,
            min_observations=args.min_observations,
            transform=Transform(args.transform),
            distinct_peak_exclusion_radius=args.distinct_peak_exclusion_radius,
        )
    )
    evaluation_positions = [len(target) // 2, (3 * len(target)) // 4, len(target) - 1]
    print(f"grid=1D source=M2 US EU CN target=close transform={args.transform}")
    print(f"lags={args.min_lag}..{args.max_lag} window={args.window} min_observations={args.min_observations}")
    for position in evaluation_positions:
        estimate = engine.estimate_at(source, target, target.index[position])
        print(f"\n{estimate.evaluated_at.date()} best_lag={estimate.best_lag} correlation={estimate.best_correlation} observations={estimate.observations}")
        print(f"nearby quality: separation={estimate.confidence.relative_peak_separation} margin={estimate.confidence.margin} runner_up_lag={estimate.confidence.runner_up_lag}")
        print(f"distinct peak (radius={args.distinct_peak_exclusion_radius}): lag={estimate.confidence.distinct_peak_lag} correlation={estimate.confidence.distinct_peak_correlation} separation={estimate.confidence.distinct_peak_relative_separation}")
        for candidate in estimate.top_candidates(args.top):
            print(f"  lag={candidate.lag:3d} correlation={candidate.correlation:+.6f} observations={candidate.observations}")


if __name__ == "__main__":
    main()
