"""Run with: python -m unittest discover -s tests -p test_ml_pipeline.py"""
import sys
import json
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch

import numpy as np
import pandas as pd

sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "ml"))
import train_phase2 as training
import score_live as scoring
from probability_calibration import apply_calibration


class MlPipelineTests(unittest.TestCase):
    def history(self):
        rng = np.random.default_rng(42)
        dates = pd.date_range("2020-01-01", periods=160).repeat(8)
        n = len(dates)
        target = np.arange(n) % 2
        frame = pd.DataFrame({
            "signal_date": dates, "evaluated_through": dates + pd.Timedelta(days=2),
            "signal_id": np.arange(n), "symbol": "TEST", "source": "BACKTEST",
            "leakage_audit_pass": 1, "eligible_success_model": 1,
            "label_target_before_stop": target,
            "outcome": np.where(target, "TARGET", "STOP"),
        })
        for name in training.FEATURE_CANDIDATES[:10]:
            frame[name] = target + rng.normal(size=n)
        return frame

    def test_labels_never_cross_split_boundaries(self):
        tr, va, te, _ = training.date_split(self.history())
        self.assertTrue((tr.evaluated_through < va.signal_date.min()).all())
        self.assertTrue((va.evaluated_through < te.signal_date.min()).all())
        with self.assertRaisesRegex(ValueError, "evaluated_through"):
            training.date_split(self.history().drop(columns="evaluated_through"))

    def test_train_and_inference_round_trip(self):
        frame = self.history()
        # Future-only feature must never enter the production schema.
        frame["rsi14"] = np.nan
        frame.loc[frame.index >= 896, "rsi14"] = np.arange(len(frame) - 896)
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp)
            with patch.object(training, "select_features", wraps=training.select_features) as select:
                result = training.train_one(frame, "success", root, 0.6, 200, 2)
                self.assertLess(select.call_args_list[0].args[0].signal_date.max(), frame.signal_date.max())
            self.assertNotIn("rsi14", result["features"])
            self.assertGreater(result["purged_rows"], 0)
            self.assertTrue(result["walk_forward"])
            self.assertEqual(result["tuning"]["status"], "tuned")
            self.assertEqual(result["tuning"]["folds"], 3)
            self.assertGreater(result["split_rows"]["calibration"], 0)
            self.assertLess(result["calibration_range"]["end"], result["selection_range"]["start"])
            registry = {"models": {"success": {
                "artifact": "success/production_champion.joblib",
                "features": "success/features.json", "metrics": "success/metrics.json",
                "champion": result["champion"],
            }}}
            live = frame.tail(8).copy()
            probs, ranks, _, features, _ = scoring.load_and_score(live, root, registry, "success")
            self.assertTrue(np.isfinite(probs).all())
            predictions = pd.read_csv(root / "success/test_predictions.csv")
            column = "logistic_probability" if result["champion"] == "logistic" else "xgb_probability"
            np.testing.assert_allclose(probs, predictions[column].tail(8), rtol=1e-6)
            for date in live.signal_date.unique():
                self.assertEqual(ranks[live.signal_date == date].max(), 100)
            # Explicitly exercise calibrated artifact inference even if the
            # synthetic training run chose the raw-score fallback.
            metadata_path = root / "success/metrics.json"
            metadata = json.loads(metadata_path.read_text())
            metadata["calibration"] = {"method": "sigmoid_logit", "slope": 0.8, "intercept": -0.2}
            metadata_path.write_text(json.dumps(metadata))
            raw_column = "logistic_raw_probability" if result["champion"] == "logistic" else "xgb_raw_probability"
            calibrated, *_ = scoring.load_and_score(live, root, registry, "success")
            np.testing.assert_allclose(calibrated, apply_calibration(predictions[raw_column].tail(8), metadata["calibration"]), rtol=1e-6)
            registry["models"]["success"]["calibration"] = {"method": "none"}
            with self.assertRaisesRegex(RuntimeError, "metadata mismatch"):
                scoring.load_and_score(live, root, registry, "success")
            del registry["models"]["success"]["calibration"]
            del metadata["calibration"]
            metadata_path.write_text(json.dumps(metadata))
            legacy, *_ = scoring.load_and_score(live, root, registry, "success")
            np.testing.assert_allclose(legacy, predictions[raw_column].tail(8), rtol=1e-6)
            live.loc[live.index[0], features] = np.inf
            with self.assertRaisesRegex(RuntimeError, "no usable"):
                scoring.load_and_score(live, root, registry, "success")

    def test_precise_resolution_dates_and_legacy_fallback(self):
        frame = self.history().head(4).copy()
        frame["evaluated_through"] = pd.Timestamp("2020-06-01")
        frame["entry_date"] = "2020-01-02"
        frame["exit_date"] = ["2020-01-03", None, "2019-12-01", "2021-01-01"]
        success = training.build_objective_df(frame, "success")
        self.assertEqual(success.label_resolved_date.iloc[0], pd.Timestamp("2020-01-03"))
        self.assertTrue(success.label_resolved_date.iloc[1:].eq(pd.Timestamp("2020-06-01")).all())
        self.assertEqual(len(training.purge_unresolved(success, pd.Timestamp("2020-02-01"))), 1)
        frame["eligible_entry_model"] = 1
        frame["label_entry_triggered"] = [1, 0, 1, 0]
        trigger = training.build_objective_df(frame, "trigger")
        self.assertEqual(trigger.label_resolved_date.iloc[0], pd.Timestamp("2020-01-02"))
        self.assertEqual(trigger.label_resolved_date.iloc[1], pd.Timestamp("2020-06-01"))

    def test_calibration_selection_and_fallback(self):
        actual = np.repeat([0.1, 0.3, 0.7, 0.9], 100)
        labels = np.concatenate([np.r_[np.ones(n), np.zeros(100-n)] for n in [10, 30, 70, 90]])
        raw = np.sqrt(actual)
        fitted = training.fit_calibration(labels, raw)
        selected, metrics = training.select_calibration(labels, raw, fitted)
        self.assertEqual(selected["method"], "sigmoid_logit")
        self.assertLess(metrics["candidate"]["brier"], metrics["raw"]["brier"])
        self.assertEqual(training.fit_calibration(np.ones(60), np.full(60, .8))["method"], "none")
        bad = {"method": "sigmoid_logit", "slope": 1, "intercept": 10}
        selected, _ = training.select_calibration(labels, actual, bad)
        self.assertEqual(selected["method"], "none")
        with self.assertRaisesRegex(ValueError, "parameters"):
            apply_calibration(raw, {"method": "sigmoid_logit", "slope": -1, "intercept": 0})
        np.testing.assert_array_equal(apply_calibration(raw, {}), raw)

    def test_not_triggered_expiry_preserves_negative_training_labels(self):
        frame = self.history().head(4).copy()
        frame["outcome"] = "NOT_TRIGGERED"
        frame["eligible_entry_model"] = 1
        frame["label_entry_triggered"] = 0
        frame["evaluated_through"] = "2026-09-25"
        frame["entry_window_end_date"] = ["2020-01-15", None, "2019-01-01", "2027-01-01"]
        for objective in ["trigger", "direct_target"]:
            built = training.build_objective_df(frame, objective)
            purged = training.purge_unresolved(built, pd.Timestamp("2020-02-01"))
            self.assertEqual(len(purged), 1)
            self.assertEqual(purged.target.iloc[0], 0)

    def test_calibration_partition_purges_overlapping_labels(self):
        objective = training.build_objective_df(self.history(), "success")
        _, validation, _, _ = training.date_split(objective)
        cal, selection = training.calibration_split(validation)
        self.assertFalse(cal.empty)
        self.assertTrue((cal.label_resolved_date < selection.signal_date.min()).all())
        self.assertTrue(set(cal.signal_id).isdisjoint(selection.signal_id))

    def test_live_validation_rejects_fractional_audit_and_bad_dates(self):
        live = self.history().head(2).assign(
            dataset_version="ml-dataset-v1b", stock_id=1, source="LIVE", vcp_status="READY")
        live["leakage_audit_pass"] = live["leakage_audit_pass"].astype(float)
        live.loc[0, "leakage_audit_pass"] = 1.5
        with self.assertRaisesRegex(RuntimeError, "leakage_audit_pass"):
            scoring.validate_live_dataset(live)
        live["leakage_audit_pass"] = 1
        live["signal_date"] = "invalid"
        with self.assertRaisesRegex(RuntimeError, "signal_date"):
            scoring.validate_live_dataset(live)

    def test_infinite_values_do_not_count_as_coverage(self):
        frame = pd.DataFrame({"rsi14": [np.inf, -np.inf, 1, 2]})
        features, coverage, _, _ = training.select_features(frame, 0.6)
        self.assertEqual(features, [])
        self.assertEqual(coverage["rsi14"], 0.5)

    def test_deployment_gate_rejects_weak_or_missing_results(self):
        good = {"rows": 100, "pr_auc": .7, "positive_rate": .3, "brier": .15, "top_10pct_lift": 2}
        baseline = {"brier": .21}
        self.assertTrue(training.deployment_check(good, baseline)["eligible"])
        for changes in [{"rows": 49}, {"pr_auc": .3}, {"brier": .25}, {"top_10pct_lift": 1}, {"pr_auc": float("nan")}]:
            result = training.deployment_check({**good, **changes}, baseline)
            self.assertFalse(result["eligible"])
            self.assertTrue(result["reasons"])

    def test_tuning_falls_back_on_insufficient_history(self):
        frame = training.build_objective_df(self.history().head(80), "success")
        params, report = training.tune_xgb(frame, .6)
        self.assertEqual(params, {})
        self.assertEqual(report["status"], "default")

    def test_scoring_cli_refuses_ineligible_run(self):
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp)
            csv = root / "live.csv"
            csv.write_text("signal_id\n1\n")
            (root / "model_registry.json").write_text(json.dumps({"deployment": {"eligible": False}}))
            with patch.object(sys, "argv", ["score_live.py", "--csv", str(csv), "--run-dir", str(root), "--out", str(root / "scores.csv")]):
                with self.assertRaisesRegex(RuntimeError, "deployment checks"):
                    scoring.main()
            self.assertFalse((root / "scores.csv").exists())


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