import json
import sys
import tempfile
import unittest
import zipfile
from pathlib import Path

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

import paired_backtest_attribution as mod  # noqa: E402


def trade(
    pair,
    opened,
    closed,
    pnl,
    *,
    is_short=True,
    open_rate=100.0,
    close_rate=100.0,
    exit_reason="exit_signal",
):
    return {
        "pair": pair,
        "is_short": is_short,
        "open_timestamp": opened,
        "close_timestamp": closed,
        "open_date": "2026-01-01 00:00:00+00:00",
        "close_date": "2026-01-01 01:00:00+00:00",
        "open_rate": open_rate,
        "close_rate": close_rate,
        "amount": 1.0,
        "stake_amount": 100.0,
        "leverage": 2.0,
        "initial_stop_loss_abs": 105.0,
        "stop_loss_abs": 105.0,
        "profit_abs": pnl,
        "profit_ratio": pnl / 100.0,
        "fee_open": 0.00035,
        "fee_close": 0.00035,
        "funding_fees": 0.0,
        "exit_reason": exit_reason,
        "enter_tag": "",
        "orders": [],
    }


def result(trades):
    gains = sum(row["profit_abs"] for row in trades if row["profit_abs"] > 0)
    losses = -sum(row["profit_abs"] for row in trades if row["profit_abs"] < 0)
    pnl = sum(row["profit_abs"] for row in trades)
    return {
        "trades": trades,
        "profit_total_abs": pnl,
        "profit_total": pnl / 1000.0,
        "profit_factor": gains / losses if losses else 99.0,
        "max_drawdown_account": 0.01,
        "starting_balance": 1000.0,
        "final_balance": 1000.0 + pnl,
        "backtest_start": "2026-01-01 00:00:00",
        "backtest_end": "2026-05-01 00:00:00",
        "timeframe": "5m",
        "max_open_trades": 4,
        "trading_mode": "futures",
        "margin_mode": "isolated",
        "enable_protections": True,
    }


def arm(name, trades, wallet_sha="control-wallet"):
    return mod.LoadedArm(
        name=name,
        path=Path(f"/{name}.json"),
        sha256=name * 4,
        strategy=name,
        result_member=None,
        wallet_member=f"{name}_wallet.feather",
        wallet_sha256=wallet_sha,
        result=result(trades),
    )


class TradeComparisonTest(unittest.TestCase):
    def test_retention_distinguishes_modified_and_added_trades(self):
        shared = trade("BTC/USDT:USDT", 1000, 2000, 5.0)
        modified_left = trade("ETH/USDT:USDT", 3000, 4000, -2.0)
        modified_right = trade("ETH/USDT:USDT", 3000, 4500, 3.0)
        added = trade("SOL/USDT:USDT", 5000, 6000, 1.0)
        metrics, rows = mod.compare_trades(
            arm("left", [shared, modified_left]),
            arm("right", [shared, modified_right, added]),
        )
        self.assertEqual(metrics["entry_retention_ratio"], 1.0)
        self.assertEqual(metrics["exact_trade_retention_ratio"], 0.5)
        self.assertEqual(metrics["modified_trades"], 1)
        self.assertEqual(metrics["right_only_trades"], 1)
        classes = {row["classification"] for row in rows}
        self.assertEqual(classes, {"identical", "modified", "right_only"})

    def test_duplicate_entry_key_fails_closed(self):
        duplicate = trade("BTC/USDT:USDT", 1000, 2000, 1.0)
        with self.assertRaisesRegex(ValueError, "duplicate trade entry key"):
            mod.compare_trades(
                arm("left", [duplicate, dict(duplicate)]),
                arm("right", []),
            )


class EpisodeTest(unittest.TestCase):
    def test_overlapping_trade_differences_merge_into_one_diagnostic_episode(self):
        left = arm(
            "baseline_a",
            [
                trade("EDGE/USDT:USDT", 1000, 4000, -5.0),
                trade("BTC/USDT:USDT", 3000, 6000, 2.0),
            ],
        )
        right = arm(
            "candidate",
            [trade("POWER/USDT:USDT", 2000, 5000, 4.0)],
        )
        _, differences = mod.compare_trades(left, right)
        episodes = mod.divergence_episodes(left, right, differences)
        self.assertEqual(len(episodes), 1)
        self.assertEqual(episodes[0]["start_timestamp"], 1000)
        self.assertEqual(episodes[0]["end_timestamp"], 6000)
        self.assertEqual(episodes[0]["attributed_profit_delta"], 7.0)
        self.assertFalse(episodes[0]["formal_path_divergence_eligible"])


class SensitivityTest(unittest.TestCase):
    def test_edge_power_arithmetic_removal_is_explicit(self):
        left = arm(
            "baseline_a",
            [
                trade("EDGE/USDT:USDT", 1000, 2000, -10.0),
                trade("POWER/USDT:USDT", 3000, 4000, 3.0),
                trade("BTC/USDT:USDT", 5000, 6000, 2.0),
            ],
        )
        right = arm(
            "candidate",
            [
                trade("EDGE/USDT:USDT", 1000, 2000, 1.0),
                trade("POWER/USDT:USDT", 3000, 4000, 5.0),
                trade("BTC/USDT:USDT", 5000, 6000, 4.0),
            ],
        )
        rows = {row["scenario"]: row for row in mod.arithmetic_sensitivity(left, right)}
        self.assertEqual(rows["none"]["profit_delta"], 15.0)
        self.assertEqual(rows["exclude_EDGE_POWER"]["profit_delta"], 2.0)
        self.assertTrue(rows["exclude_EDGE_POWER"]["drawdown_not_recomputed"])


class FourArmReportTest(unittest.TestCase):
    def test_noop_mismatch_blocks_candidate_interpretation(self):
        baseline_trade = trade("BTC/USDT:USDT", 1000, 2000, 2.0)
        arms = {
            "baseline_a": arm("baseline_a", [baseline_trade]),
            "noop": arm(
                "noop",
                [trade("BTC/USDT:USDT", 1000, 2000, 3.0)],
            ),
            "candidate": arm("candidate", [baseline_trade]),
            "baseline_b_repeat": arm("baseline_b_repeat", [baseline_trade]),
        }
        report, _ = mod.build_report(arms)
        self.assertTrue(report["protocol"]["valid"])
        self.assertFalse(report["sandwich_isolation"]["passed"])
        self.assertFalse(report["summary"]["safe_to_interpret_candidate_contrast"])

    def test_baseline_b_repeat_mismatch_blocks_candidate_interpretation(self):
        baseline_trade = trade("BTC/USDT:USDT", 1000, 2000, 2.0)
        arms = {
            "baseline_a": arm("baseline_a", [baseline_trade]),
            "noop": arm("noop", [baseline_trade]),
            "candidate": arm("candidate", [baseline_trade]),
            "baseline_b_repeat": arm(
                "baseline_b_repeat",
                [baseline_trade],
                wallet_sha="different-wallet",
            ),
        }
        report, _ = mod.build_report(arms)
        check = report["sandwich_isolation"]["checks"][
            "baseline_a_vs_baseline_b_repeat"
        ]
        self.assertFalse(check["wallet_path_identical"])
        self.assertFalse(report["summary"]["safe_to_interpret_candidate_contrast"])

    def test_zip_loader_and_cli_write_all_contract_outputs(self):
        with tempfile.TemporaryDirectory() as raw:
            root = Path(raw)
            archive_paths = {}
            base_trade = trade("BTC/USDT:USDT", 1000, 2000, 2.0)
            for name in mod.ARM_NAMES:
                archive = root / f"{name}.zip"
                payload = {"strategy": {name: result([base_trade])}}
                with zipfile.ZipFile(archive, "w") as zipped:
                    zipped.writestr("backtest-result.json", json.dumps(payload))
                    wallet_bytes = (
                        b"candidate-wallet"
                        if name == "candidate"
                        else b"identical-control-wallet"
                    )
                    zipped.writestr(
                        f"backtest-result_{name}_wallet.feather", wallet_bytes
                    )
                archive_paths[name] = archive
            output = root / "output"
            args = []
            for name in mod.ARM_NAMES:
                args.extend([f"--{name.replace('_', '-')}", str(archive_paths[name])])
            args.extend(["--output-dir", str(output)])
            code = mod.main(args)
            self.assertEqual(code, 0)
            self.assertTrue((output / "paired_backtest_attribution.json").is_file())
            self.assertTrue((output / "metrics.csv").is_file())
            self.assertTrue((output / "trade_differences.csv").is_file())
            self.assertTrue((output / "divergence_episodes.csv").is_file())
            self.assertTrue((output / "edge_power_sensitivity.csv").is_file())
            self.assertTrue((output / "realized_equity_path.csv").is_file())
            payload = json.loads(
                (output / "paired_backtest_attribution.json").read_text(encoding="utf-8")
            )
            self.assertTrue(payload["summary"]["safe_to_interpret_candidate_contrast"])
            self.assertEqual(
                payload["formal_contrast"]["contrast"],
                "baseline_a_vs_candidate",
            )
            self.assertNotIn("contrasts", payload)

    def test_result_total_must_equal_trade_sum(self):
        bad = result([trade("BTC/USDT:USDT", 1000, 2000, 2.0)])
        bad["profit_total_abs"] = 99.0
        with tempfile.TemporaryDirectory() as raw:
            path = Path(raw) / "bad.json"
            path.write_text(json.dumps({"strategy": {"bad": bad}}), encoding="utf-8")
            with self.assertRaisesRegex(ValueError, "does not equal summed trade profit"):
                mod.load_arm("baseline_a", path, None)


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