import sys
import unittest
from pathlib import Path

import numpy as np
import pandas as pd

RESEARCH_DIR = (
    Path(__file__).resolve().parents[2] / "strategies" / "research"
)
sys.path.insert(0, str(RESEARCH_DIR))

import xsmom_28d_v2_math as mod  # noqa: E402


class DailyScoreTest(unittest.TestCase):
    def test_intraday_listing_excludes_partial_day(self):
        closes = pd.Series(
            range(1, 31),
            index=pd.date_range("2025-12-06", "2026-01-04", freq="1D", tz="UTC"),
            dtype=float,
        )
        score = mod.daily_score_series(closes, "2025-12-06T09:00:00Z")
        self.assertTrue(pd.isna(score.loc["2026-01-03"]))
        self.assertAlmostEqual(score.loc["2026-01-04"], np.log(30 / 2))

    def test_midnight_listing_keeps_listing_day(self):
        closes = pd.Series(
            range(1, 30),
            index=pd.date_range("2025-12-07", "2026-01-04", freq="1D", tz="UTC"),
            dtype=float,
        )
        score = mod.daily_score_series(closes, "2025-12-07T00:00:00Z")
        self.assertAlmostEqual(score.loc["2026-01-04"], np.log(29 / 1))

    def test_missing_or_nonpositive_middle_day_blocks_score(self):
        dates = pd.date_range("2025-12-07", "2026-01-04", freq="1D", tz="UTC")
        missing = pd.Series(range(1, 30), index=dates, dtype=float).drop(dates[10])
        nonpositive = pd.Series(range(1, 30), index=dates, dtype=float)
        nonpositive.iloc[10] = 0
        self.assertTrue(
            pd.isna(
                mod.daily_score_series(
                    missing, "2025-12-07T00:00:00Z"
                ).loc["2026-01-04"]
            )
        )
        self.assertTrue(
            pd.isna(
                mod.daily_score_series(
                    nonpositive, "2025-12-07T00:00:00Z"
                ).loc["2026-01-04"]
            )
        )


class RankingTest(unittest.TestCase):
    def test_exact_score_ties_use_canonical_pair_string(self):
        dates = pd.date_range("2025-12-03", "2025-12-31", freq="1D", tz="UTC")
        closes = {
            f"P{i:02d}/USDT:USDT": pd.Series(range(1, 30), index=dates, dtype=float)
            for i in reversed(range(10))
        }
        listings = {pair: "2020-01-01T00:00:00Z" for pair in closes}
        matrices = mod.build_rank_matrices(closes, listings)
        day = pd.Timestamp("2025-12-31T00:00:00Z")
        self.assertTrue(matrices.snapshot_valid.loc[day])
        self.assertEqual(matrices.ranks.loc[day, "P00/USDT:USDT"], 1)
        self.assertEqual(matrices.ranks.loc[day, "P09/USDT:USDT"], 10)

    def test_snapshot_requires_ten_scoreable_pairs(self):
        dates = pd.date_range("2025-12-03", "2025-12-31", freq="1D", tz="UTC")
        closes = {
            f"P{i}/USDT:USDT": pd.Series(range(1, 30), index=dates, dtype=float)
            for i in range(9)
        }
        matrices = mod.build_rank_matrices(
            closes, {pair: "2020-01-01T00:00:00Z" for pair in closes}
        )
        self.assertFalse(matrices.snapshot_valid.iloc[-1])


class CausalAlignmentTest(unittest.TestCase):
    def test_2355_signal_uses_daily_close_available_at_0000(self):
        days = pd.date_range("2025-12-03", "2025-12-31", freq="1D", tz="UTC")
        closes = {
            f"P{i:02d}/USDT:USDT": pd.Series(
                np.arange(1, 30) * (i + 1), index=days, dtype=float
            )
            for i in range(10)
        }
        matrices = mod.build_rank_matrices(
            closes, {pair: "2020-01-01T00:00:00Z" for pair in closes}
        )
        frame = pd.DataFrame(
            {"date": pd.to_datetime(["2025-12-31T23:50:00Z", "2025-12-31T23:55:00Z"])}
        )
        result = mod.attach_signal_time_rank_state(
            frame, pair="P00/USDT:USDT", matrices=matrices
        )
        self.assertFalse(result.loc[0, "xsmom_pair_rank_eligible"])
        self.assertTrue(result.loc[1, "xsmom_pair_rank_eligible"])
        self.assertEqual(
            result.loc[1, "xsmom_last_complete_utc_day"],
            pd.Timestamp("2025-12-31T00:00:00Z"),
        )


class V2DecisionTest(unittest.TestCase):
    def test_precedence_and_buckets(self):
        self.assertEqual(
            mod.decision_for_side(
                side="short",
                snapshot_valid=False,
                pair_scoreable=False,
                rank=None,
                eligible_count=9,
            )[:2],
            (False, mod.REASON_SNAPSHOT_UNAVAILABLE),
        )
        self.assertEqual(
            mod.decision_for_side(
                side="short",
                snapshot_valid=True,
                pair_scoreable=False,
                rank=None,
                eligible_count=11,
            )[:2],
            (True, mod.REASON_UNSCORED_PASSTHROUGH),
        )
        self.assertEqual(
            mod.decision_for_side(
                side="long",
                snapshot_valid=True,
                pair_scoreable=True,
                rank=4,
                eligible_count=12,
            )[:2],
            (True, mod.REASON_TOP4),
        )
        self.assertEqual(
            mod.decision_for_side(
                side="short",
                snapshot_valid=True,
                pair_scoreable=True,
                rank=9,
                eligible_count=12,
            )[:2],
            (True, mod.REASON_BOTTOM4),
        )

    @staticmethod
    def signal_frame():
        return pd.DataFrame(
            {
                "enter_long": [1, 1, 0, 0],
                "enter_short": [0, 0, 1, 1],
                "enter_tag": ["base", "base", "base", "base"],
                "xsmom_long_allowed": [True, False, True, True],
                "xsmom_short_allowed": [True, True, True, False],
                "xsmom_long_reason": [
                    mod.REASON_TOP4,
                    mod.REASON_NOT_TOP4,
                    mod.REASON_UNSCORED_PASSTHROUGH,
                    mod.REASON_TOP4,
                ],
                "xsmom_short_reason": [
                    mod.REASON_NOT_BOTTOM4,
                    mod.REASON_BOTTOM4,
                    mod.REASON_UNSCORED_PASSTHROUGH,
                    mod.REASON_NOT_BOTTOM4,
                ],
            }
        )

    def test_noop_preserves_every_raw_marker_and_tag(self):
        source = self.signal_frame()
        result = mod.apply_v2_gate(source, enforce_rank_gate=False)
        pd.testing.assert_series_equal(result["enter_long"], source["enter_long"])
        pd.testing.assert_series_equal(result["enter_short"], source["enter_short"])
        pd.testing.assert_series_equal(result["enter_tag"], source["enter_tag"])
        self.assertEqual(
            set(result.loc[result["xsmom_raw_long"] | result["xsmom_raw_short"],
                           "xsmom_applied_reason"]),
            {mod.REASON_NOOP},
        )

    def test_candidate_blocks_only_rank_denials(self):
        result = mod.apply_v2_gate(self.signal_frame(), enforce_rank_gate=True)
        self.assertEqual(result["enter_long"].tolist(), [1, 0, 0, 0])
        self.assertEqual(result["enter_short"].tolist(), [0, 0, 1, 0])
        self.assertEqual(
            result.loc[2, "xsmom_applied_reason"],
            mod.REASON_UNSCORED_PASSTHROUGH,
        )
        self.assertEqual(result["enter_tag"].tolist(), ["base"] * 4)


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