"""Compare the MQL5 fractal export against the Python reference.

Run FractalValidation.mq5 first; it writes fractal_validation.csv into
MQL5\\Files. Point --csv at that file and pass the same parameters the script
was run with.

    python validate_fractals.py --csv fractal_validation.csv --n 2

The script reports, per feature column, the maximum absolute deviation and the
coefficient of determination between the two implementations.
"""

from __future__ import annotations

import argparse
from pathlib import Path

import numpy as np
import pandas as pd

from fractals_reference import FRAC_EMPTY, fractal_features

FEATURES = [
    "fractal_high",
    "fractal_low",
    "fractal_high_strength",
    "fractal_low_strength",
    "valid_fractal_high",
    "valid_fractal_low",
    "fractal_breakout_up",
    "fractal_breakout_down",
    "resistance_level",
    "support_level",
    "distance_to_resistance",
    "distance_to_support",
    "fractal_trend_strength",
    "fractal_trend_direction",
    "fractal_ma_ratio",
    "fractal_buy_signal",
    "fractal_sell_signal",
    "signal_strength",
]


def r_squared(a: np.ndarray, b: np.ndarray) -> float:
    """Coefficient of determination of a against b; 1.0 for identical inputs."""
    if a.size == 0:
        return float("nan")
    ss_res = float(np.sum((a - b) ** 2))
    ss_tot = float(np.sum((b - b.mean()) ** 2))
    if ss_tot == 0.0:
        return 1.0 if ss_res == 0.0 else 0.0
    return 1.0 - ss_res / ss_tot


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--csv", type=Path, required=True)
    parser.add_argument("--n", type=int, default=2)
    parser.add_argument("--lookback", type=int, default=20)
    parser.add_argument("--ma-period", type=int, default=20)
    parser.add_argument("--threshold", type=float, default=0.001)
    parser.add_argument("--dynamic-threshold", action="store_true")
    parser.add_argument("--tolerance", type=float, default=1e-10)
    args = parser.parse_args()

    mt5 = pd.read_csv(args.csv)
    index = pd.RangeIndex(len(mt5))

    reference = fractal_features(
        pd.Series(mt5["high"].to_numpy(), index=index),
        pd.Series(mt5["low"].to_numpy(), index=index),
        pd.Series(mt5["close"].to_numpy(), index=index),
        pd.Series(mt5["volatility"].to_numpy(), index=index),
        n=args.n,
        lookback_period=args.lookback,
        ma_period=args.ma_period,
        threshold=args.threshold,
        dynamic_threshold=args.dynamic_threshold,
    )

    print(f"{'feature':<26}{'max abs diff':>16}{'r2':>12}{'status':>10}")
    print("-" * 64)

    failures = 0
    for column in FEATURES:
        got = mt5[column].to_numpy(dtype=np.float64)
        want = reference[column].to_numpy(dtype=np.float64)

        live = (got > FRAC_EMPTY + 1.0) & (want > FRAC_EMPTY + 1.0)
        if (got <= FRAC_EMPTY + 1.0).sum() != (want <= FRAC_EMPTY + 1.0).sum():
            print(f"{column:<26}{'sentinel count differs':>38}")
            failures += 1
            continue

        deviation = float(np.max(np.abs(got[live] - want[live]))) if live.any() else 0.0
        r2 = r_squared(got[live], want[live])
        status = "PASS" if deviation <= args.tolerance else "FAIL"
        if status == "FAIL":
            failures += 1
        print(f"{column:<26}{deviation:>16.3e}{r2:>12.6f}{status:>10}")

    print()
    if failures == 0:
        print(f"ALL CHECKS PASSED  ({len(FEATURES)} columns, {len(mt5)} bars)")
    else:
        print(f"{failures} column(s) failed")
    return 1 if failures else 0


if __name__ == "__main__":
    raise SystemExit(main())
