#!/usr/bin/env python3
"""Coverage-only acceptance audit for D48-XSMOM-28D-V2.

The inputs must be authentic baseline raw-signal CSV files produced before any
portfolio, order, fill, or PnL path. This script reuses the exact V2 decision
function imported by both research strategies. It does not run a candidate
backtest and does not inspect subsequent returns.
"""

from __future__ import annotations

import argparse
import csv
import hashlib
import json
import math
import sys
from collections import Counter
from datetime import datetime, timezone
from pathlib import Path
from typing import Any

import pandas as pd

RESEARCH_STRATEGY_DIR = (
    Path(__file__).resolve().parents[1] / "strategies" / "research"
)
if str(RESEARCH_STRATEGY_DIR) not in sys.path:
    sys.path.insert(0, str(RESEARCH_STRATEGY_DIR))

from xsmom_28d_v2_math import (  # noqa: E402
    REASON_NOOP,
    REASON_UNSCORED_PASSTHROUGH,
    decision_for_side,
    first_full_listing_day,
)

SCRIPT_VERSION = "1.0"
SNAPSHOT_COVERAGE_GATE = 0.98
COMMON_SUPPORT_GATE = 0.95
FRESH_POST_MATURITY_GATE = 0.98
FRESH_POST_MATURITY_PAIR_GATE = 0.95
HISTORICAL_POST_MATURITY_GATE = 1.0

REQUIRED_FIELDS = {
    "signal_candle_time",
    "signal_eligible_time",
    "pair",
    "side",
    "listing_time",
    "rank_snapshot_valid",
    "pair_rank_eligible",
    "eligible_universe_count",
    "xsmom_rank",
}

OUTPUT_FIELDS = (
    "window",
    "signal_candle_time",
    "signal_eligible_time",
    "pair",
    "side",
    "listing_time",
    "maturity_time",
    "cohort",
    "rank_snapshot_valid",
    "pair_rank_eligible",
    "eligible_universe_count",
    "xsmom_rank",
    "noop_decision",
    "noop_reason",
    "candidate_decision",
    "candidate_reason",
    "candidate_bucket",
    "candidate_matches_registered_v1_when_scoreable",
    "post_maturity",
    "unexpected_post_maturity_unscoreable",
    "raw_decision_passthrough_fidelity",
    "xsmom_changed",
    "xsmom_episode_trigger_eligible",
)


def sha256_file(path: Path, chunk_size: int = 1024 * 1024) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as handle:
        for chunk in iter(lambda: handle.read(chunk_size), b""):
            digest.update(chunk)
    return digest.hexdigest()


def parse_bool(value: str, field: str) -> bool:
    if value == "True":
        return True
    if value == "False":
        return False
    raise ValueError(f"{field} must be True or False, got {value!r}")


def maturity_time(listing_time: object) -> pd.Timestamp:
    """First signal time whose prior 29 daily endpoints can all be full days."""
    return first_full_listing_day(listing_time) + pd.Timedelta(days=29)


def cohort_for(pair: str, post_maturity: bool) -> str:
    if pair in {"EDGE/USDT:USDT", "POWER/USDT:USDT"}:
        return "edge_power_mature" if post_maturity else "edge_power_cold_start"
    return "other_mature" if post_maturity else "other_cold_start"


def transform_row(window: str, raw: dict[str, str]) -> dict[str, Any]:
    missing = REQUIRED_FIELDS - raw.keys()
    if missing:
        raise ValueError(f"input row missing fields: {', '.join(sorted(missing))}")

    signal_time = pd.Timestamp(raw["signal_eligible_time"])
    signal_time = (
        signal_time.tz_localize("UTC")
        if signal_time.tz is None
        else signal_time.tz_convert("UTC")
    )
    mature_at = maturity_time(raw["listing_time"])
    post_maturity = bool(signal_time >= mature_at)
    snapshot_valid = parse_bool(
        raw["rank_snapshot_valid"], "rank_snapshot_valid"
    )
    pair_scoreable = parse_bool(
        raw["pair_rank_eligible"], "pair_rank_eligible"
    )
    rank = float(raw["xsmom_rank"]) if raw["xsmom_rank"] else None
    eligible_count = int(raw["eligible_universe_count"])
    allowed, reason, bucket = decision_for_side(
        side=raw["side"],
        snapshot_valid=snapshot_valid,
        pair_scoreable=pair_scoreable,
        rank=rank,
        eligible_count=eligible_count,
    )
    candidate_decision = "pass" if allowed else "block"

    if pair_scoreable:
        old_decision = raw.get("xsmom_decision")
        old_reason = raw.get("xsmom_reason")
        matches_v1 = bool(
            old_decision == candidate_decision and old_reason == reason
        )
    else:
        matches_v1 = None

    raw_decision_passthrough_fidelity = (
        bool(
            snapshot_valid
            and candidate_decision == "pass"
            and reason == REASON_UNSCORED_PASSTHROUGH
        )
        if not pair_scoreable
        else None
    )
    xsmom_changed = bool(pair_scoreable and candidate_decision == "block")
    return {
        "window": window,
        "signal_candle_time": raw["signal_candle_time"],
        "signal_eligible_time": signal_time.isoformat(),
        "pair": raw["pair"],
        "side": raw["side"],
        "listing_time": raw["listing_time"],
        "maturity_time": mature_at.isoformat(),
        "cohort": cohort_for(raw["pair"], post_maturity),
        "rank_snapshot_valid": snapshot_valid,
        "pair_rank_eligible": pair_scoreable,
        "eligible_universe_count": eligible_count,
        "xsmom_rank": rank,
        "noop_decision": "pass",
        "noop_reason": REASON_NOOP,
        "candidate_decision": candidate_decision,
        "candidate_reason": reason,
        "candidate_bucket": bucket,
        "candidate_matches_registered_v1_when_scoreable": matches_v1,
        "post_maturity": post_maturity,
        "unexpected_post_maturity_unscoreable": bool(
            post_maturity and not pair_scoreable
        ),
        "raw_decision_passthrough_fidelity": raw_decision_passthrough_fidelity,
        "xsmom_changed": xsmom_changed,
        # Actual PATH_DIVERGENCE requires portfolio state. This field only
        # establishes whether the XSMOM decision is allowed to trigger one.
        "xsmom_episode_trigger_eligible": xsmom_changed,
    }


def read_input(label: str, path: Path) -> list[dict[str, Any]]:
    with path.open(encoding="utf-8", newline="") as handle:
        reader = csv.DictReader(handle)
        if reader.fieldnames is None:
            raise ValueError(f"{path}: missing CSV header")
        missing = REQUIRED_FIELDS - set(reader.fieldnames)
        if missing:
            raise ValueError(
                f"{path}: missing fields: {', '.join(sorted(missing))}"
            )
        return [transform_row(label, row) for row in reader]


def ratio(numerator: int, denominator: int) -> float | None:
    return numerator / denominator if denominator else None


def row_set_sha256(rows: list[dict[str, Any]]) -> str:
    payload = "\n".join(
        sorted(
            "|".join(
                (
                    str(row["window"]),
                    str(row["signal_eligible_time"]),
                    str(row["pair"]),
                    str(row["side"]),
                )
            )
            for row in rows
        )
    )
    return hashlib.sha256(payload.encode("utf-8")).hexdigest()


def summarize_rows(
    rows: list[dict[str, Any]],
    *,
    historical: bool,
    expected_cold_start_count: int | None,
) -> dict[str, Any]:
    total = len(rows)
    snapshot_count = sum(row["rank_snapshot_valid"] for row in rows)
    scoreable_count = sum(row["pair_rank_eligible"] for row in rows)
    post = [row for row in rows if row["post_maturity"]]
    post_scoreable = sum(row["pair_rank_eligible"] for row in post)
    unscored = [row for row in rows if not row["pair_rank_eligible"]]
    passthrough = [
        row
        for row in unscored
        if row["rank_snapshot_valid"]
    ]
    passthrough_ok = sum(
        row["raw_decision_passthrough_fidelity"] is True for row in passthrough
    )
    passthrough_changed = sum(row["xsmom_changed"] for row in passthrough)
    scoreable_v1_match = [
        row for row in rows if row["pair_rank_eligible"]
    ]
    match_count = sum(
        row["candidate_matches_registered_v1_when_scoreable"] is True
        for row in scoreable_v1_match
    )

    per_pair: list[dict[str, Any]] = []
    for pair in sorted({row["pair"] for row in rows}):
        pair_rows = [row for row in rows if row["pair"] == pair]
        pair_post = [row for row in pair_rows if row["post_maturity"]]
        pair_post_scoreable = sum(row["pair_rank_eligible"] for row in pair_post)
        per_pair.append(
            {
                "pair": pair,
                "signals": len(pair_rows),
                "scoreable_signals": sum(
                    row["pair_rank_eligible"] for row in pair_rows
                ),
                "common_support_ratio": ratio(
                    sum(row["pair_rank_eligible"] for row in pair_rows),
                    len(pair_rows),
                ),
                "post_maturity_signals": len(pair_post),
                "post_maturity_scoreable": pair_post_scoreable,
                "post_maturity_scoreability_ratio": ratio(
                    pair_post_scoreable, len(pair_post)
                ),
            }
        )

    post_gate = (
        HISTORICAL_POST_MATURITY_GATE
        if historical
        else FRESH_POST_MATURITY_GATE
    )
    per_pair_post_gate = (
        HISTORICAL_POST_MATURITY_GATE
        if historical
        else FRESH_POST_MATURITY_PAIR_GATE
    )
    return {
        "raw_signal_count": total,
        "raw_signal_row_set_sha256": row_set_sha256(rows),
        "snapshot_valid_count": snapshot_count,
        "snapshot_valid_ratio": ratio(snapshot_count, total),
        "snapshot_valid_gate_98pct": bool(
            total and snapshot_count / total >= SNAPSHOT_COVERAGE_GATE
        ),
        "common_support_count": scoreable_count,
        "common_support_ratio": ratio(scoreable_count, total),
        "common_support_gate_95pct": bool(
            total and scoreable_count / total >= COMMON_SUPPORT_GATE
        ),
        "post_maturity_signal_count": len(post),
        "post_maturity_scoreable_count": post_scoreable,
        "post_maturity_scoreability_ratio": ratio(post_scoreable, len(post)),
        "post_maturity_scoreability_gate": bool(
            post and post_scoreable / len(post) >= post_gate
        ),
        "post_maturity_every_pair_gate": bool(
            per_pair
            and all(
                item["post_maturity_signals"] > 0
                and item["post_maturity_scoreability_ratio"] >= per_pair_post_gate
                for item in per_pair
            )
        ),
        "unscored_signal_count": len(unscored),
        "snapshot_unavailable_unscored_count": len(unscored) - len(passthrough),
        "cold_start_signal_count": len(passthrough),
        "cold_start_row_set_sha256": row_set_sha256(passthrough),
        "cold_start_by_pair": {
            pair: {
                "signals": len(pair_rows),
                "first_signal_eligible_time": min(
                    row["signal_eligible_time"] for row in pair_rows
                ),
                "last_signal_eligible_time": max(
                    row["signal_eligible_time"] for row in pair_rows
                ),
                "row_set_sha256": row_set_sha256(pair_rows),
            }
            for pair in sorted({row["pair"] for row in passthrough})
            if (
                pair_rows := [
                    row for row in passthrough if row["pair"] == pair
                ]
            )
        },
        "expected_cold_start_signal_count": expected_cold_start_count,
        "cold_start_count_matches_expected": (
            None
            if expected_cold_start_count is None
            else len(passthrough) == expected_cold_start_count
        ),
        "raw_decision_passthrough_fidelity_count": passthrough_ok,
        "raw_decision_passthrough_fidelity_ratio": ratio(
            passthrough_ok, len(passthrough)
        ),
        "raw_decision_passthrough_fidelity_gate_100pct": (
            passthrough_ok == len(passthrough)
        ),
        "passthrough_xsmom_changed_count": passthrough_changed,
        "passthrough_zero_xsmom_change_gate": passthrough_changed == 0,
        "scoreable_v1_decision_match_count": match_count,
        "scoreable_v1_decision_match_gate_100pct": bool(
            scoreable_v1_match and match_count == len(scoreable_v1_match)
        ),
        "candidate_changed_scoreable_signal_count": sum(
            row["xsmom_changed"] for row in rows
        ),
        "candidate_reasons": dict(
            sorted(Counter(row["candidate_reason"] for row in rows).items())
        ),
        "cohorts": dict(sorted(Counter(row["cohort"] for row in rows).items())),
        "coverage_by_pair": per_pair,
    }


def json_csv_value(value: Any) -> Any:
    if isinstance(value, bool):
        return str(value)
    if value is None:
        return ""
    return value


def write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
    with path.open("w", encoding="utf-8", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=OUTPUT_FIELDS)
        writer.writeheader()
        for row in rows:
            writer.writerow(
                {field: json_csv_value(row.get(field)) for field in OUTPUT_FIELDS}
            )


def parse_input(value: str) -> tuple[str, Path]:
    label, separator, raw_path = value.partition("=")
    if not separator or not label or not raw_path:
        raise argparse.ArgumentTypeError("--input must be LABEL=CSV_PATH")
    return label, Path(raw_path)


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--input", action="append", type=parse_input, required=True)
    parser.add_argument("--output-dir", type=Path, required=True)
    parser.add_argument(
        "--expected-cold-start-count",
        type=int,
        default=111,
        help="Frozen combined historical expectation; use 0 for fresh-forward audits.",
    )
    parser.add_argument(
        "--fresh-forward",
        action="store_true",
        help="Apply fresh-forward post-maturity thresholds instead of historical 100%.",
    )
    parser.add_argument("--strategy-source", type=Path, action="append", default=[])
    return parser


def core_acceptance(summary: dict[str, Any]) -> bool:
    gates = (
        "snapshot_valid_gate_98pct",
        "common_support_gate_95pct",
        "post_maturity_scoreability_gate",
        "post_maturity_every_pair_gate",
        "raw_decision_passthrough_fidelity_gate_100pct",
        "passthrough_zero_xsmom_change_gate",
        "scoreable_v1_decision_match_gate_100pct",
    )
    return all(summary[key] for key in gates)


def main() -> int:
    args = build_parser().parse_args()
    rows: list[dict[str, Any]] = []
    inputs: dict[str, dict[str, Any]] = {}
    for label, path in args.input:
        if not path.is_file():
            raise SystemExit(f"input does not exist: {path}")
        input_rows = read_input(label, path)
        rows.extend(input_rows)
        inputs[label] = {
            "path": str(path),
            "sha256": sha256_file(path),
            "raw_signal_count": len(input_rows),
        }

    historical = not args.fresh_forward
    combined = summarize_rows(
        rows,
        historical=historical,
        expected_cold_start_count=args.expected_cold_start_count,
    )
    expected_by_window = {
        "training": args.expected_cold_start_count,
        "retrospective": 0,
    }
    coverage_by_window = {
        label: summarize_rows(
            [row for row in rows if row["window"] == label],
            historical=historical,
            expected_cold_start_count=expected_by_window.get(label),
        )
        for label in inputs
    }
    summary = {
        "script_version": SCRIPT_VERSION,
        "generated_at": datetime.now(timezone.utc).isoformat(),
        "mode": "fresh_forward" if args.fresh_forward else "historical",
        "inputs": inputs,
        "strategy_source_sha256": {
            str(path): sha256_file(path)
            for path in args.strategy_source
            if path.is_file()
        },
        **combined,
        "coverage_by_window": coverage_by_window,
    }
    summary["all_pre_pnl_acceptance_gates"] = bool(
        core_acceptance(summary)
        and summary["cold_start_count_matches_expected"]
        and all(
            core_acceptance(window)
            and window["cold_start_count_matches_expected"] is not False
            for window in coverage_by_window.values()
        )
    )

    args.output_dir.mkdir(parents=True, exist_ok=True)
    write_csv(args.output_dir / "v2_raw_signal_coverage.csv", rows)
    (args.output_dir / "summary.json").write_text(
        json.dumps(summary, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
        encoding="utf-8",
    )
    print(
        json.dumps(
            {
                "raw_signals": summary["raw_signal_count"],
                "cold_start_signals": summary["cold_start_signal_count"],
                "all_pre_pnl_acceptance_gates": summary[
                    "all_pre_pnl_acceptance_gates"
                ],
            },
            sort_keys=True,
        )
    )
    return 0


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