#!/usr/bin/env python3
"""No-PnL integration parity check for D48 XSMOM V2 strategies."""

from __future__ import annotations

import argparse
import csv
import json
from pathlib import Path
from typing import Any

import pandas as pd

import authentic_raw_signal_coverage as common


def signal_keys(
    frames: dict[str, pd.DataFrame],
    *,
    columns: tuple[tuple[str, str], ...],
    window_start: pd.Timestamp,
    window_end: pd.Timestamp,
) -> set[tuple[str, str, str]]:
    keys: set[tuple[str, str, str]] = set()
    for pair, frame in frames.items():
        for side, column in columns:
            if column not in frame:
                continue
            selected = frame.loc[frame[column].fillna(0).astype(bool), "date"]
            for candle_time in selected:
                eligible = pd.Timestamp(candle_time) + pd.Timedelta(minutes=5)
                if window_start <= eligible < window_end:
                    key = (eligible.isoformat(), pair, side)
                    if key in keys:
                        raise ValueError(f"duplicate signal key: {key}")
                    keys.add(key)
    return keys


def coverage_sets(
    path: Path,
    *,
    window_start: pd.Timestamp | None = None,
    window_end: pd.Timestamp | None = None,
) -> tuple[set[tuple[str, str, str]], set[tuple[str, str, str]]]:
    passed: set[tuple[str, str, str]] = set()
    unscored: set[tuple[str, str, str]] = set()
    with path.open(encoding="utf-8", newline="") as handle:
        reader = csv.DictReader(handle)
        required = {
            "signal_eligible_time",
            "pair",
            "side",
            "candidate_decision",
            "candidate_reason",
        }
        if reader.fieldnames is None or required - set(reader.fieldnames):
            raise ValueError("coverage CSV lacks V2 decision fields")
        for row in reader:
            eligible_time = common.parse_timestamp(row["signal_eligible_time"])
            if (
                window_start is not None
                and window_end is not None
                and not (window_start <= eligible_time < window_end)
            ):
                continue
            key = (
                eligible_time.isoformat(),
                row["pair"],
                row["side"],
            )
            if row["candidate_decision"] == "pass":
                if key in passed:
                    raise ValueError(f"duplicate coverage key: {key}")
                passed.add(key)
            if row["candidate_reason"] == "xsmom_unscored_passthrough":
                unscored.add(key)
    return passed, unscored


def compare_sets(
    *,
    baseline_raw: set[tuple[str, str, str]],
    noop_raw: set[tuple[str, str, str]],
    noop_retained: set[tuple[str, str, str]],
    candidate_raw: set[tuple[str, str, str]],
    candidate_retained: set[tuple[str, str, str]],
    coverage_passed: set[tuple[str, str, str]],
    coverage_unscored: set[tuple[str, str, str]],
    expected_raw_count: int | None,
) -> tuple[dict[str, Any], list[dict[str, Any]]]:
    universe = (
        baseline_raw
        | noop_raw
        | noop_retained
        | candidate_raw
        | candidate_retained
        | coverage_passed
        | coverage_unscored
    )
    mismatches: list[dict[str, Any]] = []
    for eligible_time, pair, side in sorted(universe):
        flags = {
            "baseline_raw": (eligible_time, pair, side) in baseline_raw,
            "noop_raw": (eligible_time, pair, side) in noop_raw,
            "noop_retained": (eligible_time, pair, side) in noop_retained,
            "candidate_raw": (eligible_time, pair, side) in candidate_raw,
            "candidate_retained": (eligible_time, pair, side)
            in candidate_retained,
            "coverage_passed": (eligible_time, pair, side) in coverage_passed,
            "coverage_unscored": (eligible_time, pair, side)
            in coverage_unscored,
        }
        if not (
            flags["baseline_raw"] == flags["noop_raw"] == flags["candidate_raw"]
            and flags["noop_retained"] == flags["baseline_raw"]
            and flags["candidate_retained"] == flags["coverage_passed"]
            and (
                not flags["coverage_unscored"]
                or flags["candidate_retained"]
            )
        ):
            mismatches.append(
                {
                    "signal_eligible_time": eligible_time,
                    "pair": pair,
                    "side": side,
                    **flags,
                }
            )

    summary = {
        "baseline_raw_count": len(baseline_raw),
        "noop_raw_count": len(noop_raw),
        "noop_retained_count": len(noop_retained),
        "candidate_raw_count": len(candidate_raw),
        "candidate_retained_count": len(candidate_retained),
        "coverage_passed_count": len(coverage_passed),
        "coverage_unscored_count": len(coverage_unscored),
        "expected_raw_count": expected_raw_count,
        "expected_raw_count_matches": (
            None
            if expected_raw_count is None
            else len(baseline_raw) == expected_raw_count
        ),
        "baseline_noop_candidate_raw_exact": (
            baseline_raw == noop_raw == candidate_raw
        ),
        "noop_retains_baseline_exact": noop_retained == baseline_raw,
        "candidate_retains_coverage_pass_exact": (
            candidate_retained == coverage_passed
        ),
        "unscored_all_retained": coverage_unscored <= candidate_retained,
        "mismatch_count": len(mismatches),
    }
    summary["all_strategy_parity_gates"] = bool(
        summary["baseline_noop_candidate_raw_exact"]
        and summary["noop_retains_baseline_exact"]
        and summary["candidate_retains_coverage_pass_exact"]
        and summary["unscored_all_retained"]
        and summary["mismatch_count"] == 0
        and summary["expected_raw_count_matches"] is not False
    )
    return summary, mismatches


def run_strategy_signal_sets(
    strategy: Any,
    main_frames: dict[str, pd.DataFrame],
    *,
    raw_columns: tuple[tuple[str, str], ...],
    retained_columns: tuple[tuple[str, str], ...],
    window_start: pd.Timestamp,
    window_end: pd.Timestamp,
) -> tuple[set[tuple[str, str, str]], set[tuple[str, str, str]]]:
    raw: set[tuple[str, str, str]] = set()
    retained: set[tuple[str, str, str]] = set()
    for pair, frame in main_frames.items():
        analyzed = strategy.advise_indicators(frame.copy(), {"pair": pair})
        analyzed = strategy.advise_entry(analyzed, {"pair": pair})
        raw.update(
            signal_keys(
                {pair: analyzed},
                columns=raw_columns,
                window_start=window_start,
                window_end=window_end,
            )
        )
        retained.update(
            signal_keys(
                {pair: analyzed},
                columns=retained_columns,
                window_start=window_start,
                window_end=window_end,
            )
        )
    return raw, retained


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--config", type=Path, required=True)
    parser.add_argument("--baseline-strategy-file", type=Path, required=True)
    parser.add_argument("--v2-strategy-file", type=Path, required=True)
    parser.add_argument("--data-dir", type=Path, required=True)
    parser.add_argument("--coverage-csv", type=Path, required=True)
    parser.add_argument("--window-start", required=True)
    parser.add_argument("--window-end", required=True)
    parser.add_argument("--expected-raw-count", type=int)
    parser.add_argument("--output-dir", type=Path, required=True)
    return parser


def main() -> int:
    args = build_parser().parse_args()
    config = common.read_config(args.config)
    pairs = common.load_pairs(config, [])
    window_start = common.parse_timestamp(args.window_start)
    window_end = common.parse_timestamp(args.window_end)

    daily_frames = {
        pair: common.load_feather(
            common.pair_file(args.data_dir, pair, "1d-futures")
        )
        for pair in pairs
    }
    main_frames = {
        pair: common.load_feather(
            common.pair_file(args.data_dir, pair, "5m-futures")
        )
        for pair in pairs
    }
    provider = common.HistoricalDataProvider(
        {(pair, "1d"): frame for pair, frame in daily_frames.items()}
    )

    strategy_config = dict(config)
    from freqtrade.enums import RunMode

    strategy_config["runmode"] = RunMode.BACKTEST
    strategy_config["timeframe"] = "5m"
    baseline_class = common.load_strategy_class(
        args.baseline_strategy_file, "VolatilityBreakoutRiskCap"
    )
    noop_class = common.load_strategy_class(
        args.v2_strategy_file, "VolatilityBreakoutXsmom28dV2NoOp"
    )
    candidate_class = common.load_strategy_class(
        args.v2_strategy_file, "VolatilityBreakoutXsmom28dV2"
    )
    strategies = [
        baseline_class(strategy_config),
        noop_class(strategy_config),
        candidate_class(strategy_config),
    ]
    for strategy in strategies:
        strategy.dp = provider

    raw_columns = (("long", "enter_long"), ("short", "enter_short"))
    xsmom_raw_columns = (
        ("long", "xsmom_raw_long"),
        ("short", "xsmom_raw_short"),
    )
    baseline_raw, _ = run_strategy_signal_sets(
        strategies[0],
        main_frames,
        raw_columns=raw_columns,
        retained_columns=raw_columns,
        window_start=window_start,
        window_end=window_end,
    )
    noop_raw, noop_retained = run_strategy_signal_sets(
        strategies[1],
        main_frames,
        raw_columns=xsmom_raw_columns,
        retained_columns=raw_columns,
        window_start=window_start,
        window_end=window_end,
    )
    candidate_raw, candidate_retained = run_strategy_signal_sets(
        strategies[2],
        main_frames,
        raw_columns=xsmom_raw_columns,
        retained_columns=raw_columns,
        window_start=window_start,
        window_end=window_end,
    )
    coverage_passed, coverage_unscored = coverage_sets(
        args.coverage_csv,
        window_start=window_start,
        window_end=window_end,
    )
    summary, mismatches = compare_sets(
        baseline_raw=baseline_raw,
        noop_raw=noop_raw,
        noop_retained=noop_retained,
        candidate_raw=candidate_raw,
        candidate_retained=candidate_retained,
        coverage_passed=coverage_passed,
        coverage_unscored=coverage_unscored,
        expected_raw_count=args.expected_raw_count,
    )
    summary.update(
        {
            "window_start": window_start.isoformat(),
            "window_end": window_end.isoformat(),
            "config_sha256": common.sha256_file(args.config),
            "baseline_source_sha256": common.sha256_file(
                args.baseline_strategy_file
            ),
            "v2_source_sha256": common.sha256_file(args.v2_strategy_file),
            "coverage_csv_sha256": common.sha256_file(args.coverage_csv),
        }
    )

    args.output_dir.mkdir(parents=True, exist_ok=True)
    (args.output_dir / "summary.json").write_text(
        json.dumps(summary, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
        encoding="utf-8",
    )
    mismatch_fields = (
        "signal_eligible_time",
        "pair",
        "side",
        "baseline_raw",
        "noop_raw",
        "noop_retained",
        "candidate_raw",
        "candidate_retained",
        "coverage_passed",
        "coverage_unscored",
    )
    with (args.output_dir / "mismatches.csv").open(
        "w", encoding="utf-8", newline=""
    ) as handle:
        writer = csv.DictWriter(handle, fieldnames=mismatch_fields)
        writer.writeheader()
        writer.writerows(mismatches)
    print(json.dumps(summary, sort_keys=True))
    return 0


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