import json
import sys
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch

import pandas as pd

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

import compare_research_data as mod  # noqa: E402


def frame(dates, close, **columns):
    size = len(dates)
    payload = {
        "date": pd.to_datetime(dates, utc=True),
        "open": columns.pop("open", close),
        "high": columns.pop("high", close),
        "low": columns.pop("low", close),
        "close": close,
        "volume": columns.pop("volume", [1.0] * size),
    }
    payload.update(columns)
    return pd.DataFrame(payload)


class ListingPolicyTest(unittest.TestCase):
    def test_partial_listing_day_is_excluded_from_daily_comparison(self):
        production = frame(
            ["2026-03-19T00:00:00Z", "2026-03-20T00:00:00Z"],
            [999.0, 2.0],
        )
        reference = frame(
            ["2026-03-19T00:00:00Z", "2026-03-20T00:00:00Z"],
            [1.0, 2.0],
        )
        result = mod.compare_frames(
            production,
            reference,
            pair="EDGE/USDT:USDT",
            spec=mod.DATASETS[1],
            window_start=mod.as_utc(mod.DEFAULT_START, "start"),
            window_end=mod.as_utc(mod.DEFAULT_END, "end"),
            listing_time=mod.as_utc("2026-03-19T14:00:00Z", "listing"),
        )
        self.assertEqual(result["effective_start"], "2026-03-20T00:00:00+00:00")
        self.assertEqual(result["compared_timestamps"], 1)
        self.assertEqual(result["status"], "match")

    def test_intraday_comparison_begins_at_exact_listing_timestamp(self):
        production = frame(
            ["2025-12-06T08:00:00Z", "2025-12-06T09:00:00Z"],
            [9.0, 2.0],
        )
        reference = frame(
            ["2025-12-06T08:00:00Z", "2025-12-06T09:00:00Z"],
            [1.0, 2.0],
        )
        result = mod.compare_frames(
            production,
            reference,
            pair="POWER/USDT:USDT",
            spec=mod.DATASETS[0],
            window_start=mod.as_utc(mod.DEFAULT_START, "start"),
            window_end=mod.as_utc(mod.DEFAULT_END, "end"),
            listing_time=mod.as_utc("2025-12-06T09:00:00Z", "listing"),
        )
        self.assertEqual(result["effective_start"], "2025-12-06T09:00:00+00:00")
        self.assertEqual(result["status"], "match")


class DifferenceTest(unittest.TestCase):
    def test_reports_timestamp_value_and_duplicate_differences_without_deduping(self):
        production = frame(
            [
                "2026-01-01T00:00:00Z",
                "2026-01-01T00:05:00Z",
                "2026-01-01T00:05:00Z",
                "2026-01-01T00:10:00Z",
            ],
            [10.0, 11.0, 11.5, 12.0],
        )
        reference = frame(
            [
                "2026-01-01T00:00:00Z",
                "2026-01-01T00:05:00Z",
                "2026-01-01T00:15:00Z",
            ],
            [10.5, 11.0, 13.0],
        )
        result = mod.compare_frames(
            production,
            reference,
            pair="BTC/USDT:USDT",
            spec=mod.DATASETS[0],
            window_start=mod.as_utc(mod.DEFAULT_START, "start"),
            window_end=mod.as_utc(mod.DEFAULT_END, "end"),
            listing_time=mod.as_utc("2019-09-08T17:55:00Z", "listing"),
        )
        self.assertEqual(result["status"], "structural_mismatch")
        self.assertEqual(result["production_only_timestamps"], 1)
        self.assertEqual(result["reference_only_timestamps"], 1)
        self.assertEqual(result["production_duplicate_timestamps"], 1)
        self.assertEqual(result["production_duplicate_extra_rows"], 1)
        self.assertEqual(result["ambiguous_overlap_timestamps"], 1)
        self.assertEqual(result["compared_timestamps"], 1)
        self.assertEqual(result["columns"]["close"]["mismatched_cells"], 1)
        self.assertAlmostEqual(result["max_absolute_difference"], 0.5)

    def test_nan_matches_nan_but_one_sided_nan_is_a_mismatch(self):
        production = frame(
            ["2026-01-01T00:00:00Z", "2026-01-01T00:05:00Z"],
            [float("nan"), float("nan")],
        )
        reference = frame(
            ["2026-01-01T00:00:00Z", "2026-01-01T00:05:00Z"],
            [float("nan"), 2.0],
        )
        result = mod.compare_frames(
            production,
            reference,
            pair="BTC/USDT:USDT",
            spec=mod.DATASETS[0],
            window_start=mod.as_utc(mod.DEFAULT_START, "start"),
            window_end=mod.as_utc(mod.DEFAULT_END, "end"),
            listing_time=None,
        )
        self.assertEqual(result["columns"]["close"]["nonfinite_mismatches"], 1)
        self.assertEqual(result["status"], "value_mismatch")


class CommandTest(unittest.TestCase):
    def test_main_writes_json_and_csv_and_returns_nonzero_on_mismatch(self):
        with tempfile.TemporaryDirectory() as raw:
            root = Path(raw)
            production_dir = root / "production"
            reference_dir = root / "reference"
            output_dir = root / "output"
            production_dir.mkdir()
            reference_dir.mkdir()
            listing_path = root / "listing.json"
            listing_path.write_text(
                json.dumps({"BTC/USDT:USDT": "2019-09-08T17:55:00Z"}),
                encoding="utf-8",
            )
            for spec in mod.DATASETS:
                mod.pair_file(production_dir, "BTC/USDT:USDT", spec.suffix).touch()
                mod.pair_file(reference_dir, "BTC/USDT:USDT", spec.suffix).touch()

            production = frame(["2026-01-01T00:00:00Z"], [1.0])
            reference = frame(["2026-01-01T00:00:00Z"], [2.0])
            with patch.object(
                mod.pd,
                "read_feather",
                side_effect=lambda path: production if Path(path).parent == production_dir else reference,
            ):
                code = mod.main(
                    [
                        "--production-data-dir",
                        str(production_dir),
                        "--reference-data-dir",
                        str(reference_dir),
                        "--listing-dates",
                        str(listing_path),
                        "--pairs",
                        "BTC/USDT:USDT",
                        "--output-dir",
                        str(output_dir),
                    ]
                )
            self.assertEqual(code, 1)
            payload = json.loads(
                (output_dir / "research_data_comparison.json").read_text(encoding="utf-8")
            )
            self.assertEqual(payload["summary"]["value_mismatch"], 4)
            self.assertFalse(payload["summary"]["all_match"])
            csv_text = (output_dir / "research_data_comparison.csv").read_text(
                encoding="utf-8"
            )
            self.assertIn("max_absolute_difference", csv_text)


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