import csv
import sys
import tempfile
import unittest
from pathlib import Path

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

import xsmom_28d_v2_coverage as mod  # noqa: E402


def raw_row(
    *,
    pair="BTC/USDT:USDT",
    side="short",
    signal_time="2026-01-10T12:00:00Z",
    listing_time="2019-09-08T17:55:00Z",
    snapshot=True,
    scoreable=True,
    count=11,
    rank="11",
    decision="pass",
    reason="xsmom_bottom4",
):
    return {
        "signal_candle_time": "2026-01-10T11:55:00Z",
        "signal_eligible_time": signal_time,
        "pair": pair,
        "side": side,
        "listing_time": listing_time,
        "rank_snapshot_valid": str(snapshot),
        "pair_rank_eligible": str(scoreable),
        "eligible_universe_count": str(count),
        "xsmom_rank": rank if scoreable else "",
        "xsmom_decision": decision,
        "xsmom_reason": reason,
    }


class TransformTest(unittest.TestCase):
    def test_cold_start_is_noop_faithful_passthrough(self):
        row = raw_row(
            pair="POWER/USDT:USDT",
            signal_time="2026-01-04T16:15:00Z",
            listing_time="2025-12-06T09:00:00Z",
            scoreable=False,
            rank="",
            decision="block",
            reason="xsmom_missing_pair",
        )
        result = mod.transform_row("training", row)
        self.assertEqual(result["maturity_time"], "2026-01-05T00:00:00+00:00")
        self.assertEqual(result["cohort"], "edge_power_cold_start")
        self.assertEqual(result["noop_decision"], "pass")
        self.assertEqual(result["candidate_decision"], "pass")
        self.assertEqual(
            result["candidate_reason"], "xsmom_unscored_passthrough"
        )
        self.assertTrue(result["raw_decision_passthrough_fidelity"])
        self.assertFalse(result["xsmom_changed"])
        self.assertFalse(result["xsmom_episode_trigger_eligible"])

    def test_snapshot_unavailable_takes_precedence(self):
        row = raw_row(
            snapshot=False,
            scoreable=False,
            count=9,
            rank="",
            decision="block",
            reason="xsmom_snapshot_unavailable",
        )
        result = mod.transform_row("training", row)
        self.assertEqual(result["candidate_decision"], "block")
        self.assertEqual(
            result["candidate_reason"], "xsmom_snapshot_unavailable"
        )
        self.assertFalse(result["raw_decision_passthrough_fidelity"])

    def test_scoreable_decision_must_match_registered_v1(self):
        result = mod.transform_row("training", raw_row())
        self.assertTrue(result["candidate_matches_registered_v1_when_scoreable"])
        self.assertEqual(result["candidate_decision"], "pass")

    def test_post_maturity_unscoreable_is_explicit(self):
        result = mod.transform_row(
            "training",
            raw_row(
                pair="EDGE/USDT:USDT",
                signal_time="2026-04-18T06:35:00Z",
                listing_time="2026-03-19T14:00:00Z",
                scoreable=False,
                rank="",
                decision="block",
                reason="xsmom_missing_pair",
            ),
        )
        self.assertTrue(result["post_maturity"])
        self.assertTrue(result["unexpected_post_maturity_unscoreable"])


class SummaryTest(unittest.TestCase):
    def test_all_historical_gates_pass_for_valid_small_fixture(self):
        rows = [
            mod.transform_row("training", raw_row()),
            mod.transform_row(
                "training",
                raw_row(
                    pair="POWER/USDT:USDT",
                    signal_time="2026-01-04T16:15:00Z",
                    listing_time="2025-12-06T09:00:00Z",
                    scoreable=False,
                    rank="",
                    decision="block",
                    reason="xsmom_missing_pair",
                ),
            ),
        ]
        summary = mod.summarize_rows(
            rows, historical=True, expected_cold_start_count=1
        )
        self.assertEqual(summary["cold_start_signal_count"], 1)
        self.assertTrue(summary["cold_start_count_matches_expected"])
        self.assertTrue(
            summary["raw_decision_passthrough_fidelity_gate_100pct"]
        )
        self.assertTrue(summary["passthrough_zero_xsmom_change_gate"])
        self.assertTrue(summary["scoreable_v1_decision_match_gate_100pct"])

    def test_csv_reader_rejects_missing_required_columns(self):
        with tempfile.TemporaryDirectory() as directory:
            path = Path(directory) / "bad.csv"
            with path.open("w", newline="", encoding="utf-8") as handle:
                writer = csv.DictWriter(handle, fieldnames=["pair"])
                writer.writeheader()
                writer.writerow({"pair": "BTC/USDT:USDT"})
            with self.assertRaises(ValueError):
                mod.read_input("training", path)


if __name__ == "__main__":
    unittest.main()
