from __future__ import annotations

from historical_replay import Candle, run_coin_slice_replay, run_parameter_sweep, render_parameter_sweep_report
from strategy_registry import StrategyPreset


def _candle(i: int, close: float) -> Candle:
    return Candle(ts=1_700_000_000 + i * 60, open=close, high=close, low=close, close=close, volume=1_000_000)


def _base_preset() -> StrategyPreset:
    return StrategyPreset(
        strategy_id="base",
        label="Base",
        source="unit",
        parameters={
            "flash_crash_trigger_pct": 4.5,
            "break_even_activation_pct": 0.4,
            "v_shape_activation_pct": 1.0,
            "v_shape_trail_dist_pct": 0.5,
            "dead_fish_time_limit_mins": 30,
            "atr_sl_multiplier": 2.0,
            "max_hard_stop_pct": 3.0,
            "trade_size_usd": 50.0,
            "default_leverage": 5,
            "min_volume_24h": 1,
        },
    )


def test_coin_slice_replay_reports_without_single_coin_leakage() -> None:
    candles = {
        "WLD": [_candle(0, 100), _candle(1, 100), _candle(2, 95), _candle(3, 97), _candle(4, 96.5)],
        "BTC": [_candle(0, 100), _candle(1, 100), _candle(2, 95), _candle(3, 90), _candle(4, 89)],
    }

    report = run_coin_slice_replay([_base_preset()], candles, volumes={"WLD": 9_000_000, "BTC": 9_000_000}, exclude_coins=["WLD"])

    assert report.excluded_coins == ("WLD",)
    assert report.results[0].strategy_id == "base"
    assert all(trade.coin != "WLD" for trade in report.results[0].aggregate.trades)


def test_parameter_sweep_clones_preset_with_stable_strategy_ids() -> None:
    candles = {
        "BTC": [_candle(0, 100), _candle(1, 100), _candle(2, 95), _candle(3, 97), _candle(4, 96.5)],
    }

    results = run_parameter_sweep(
        _base_preset(),
        candles,
        volumes={"BTC": 9_000_000},
        grid={"flash_crash_trigger_pct": [4.0, 4.5], "dead_fish_time_limit_mins": [20, 40]},
        segments=1,
    )

    assert len(results) == 4
    assert results[0].strategy_id.startswith("base__")
    assert "flash_crash_trigger_pct" in results[0].parameters
    assert "dead_fish_time_limit_mins" in results[0].parameters


def test_parameter_sweep_report_contains_top_variants() -> None:
    candles = {"BTC": [_candle(0, 100), _candle(1, 100), _candle(2, 95), _candle(3, 97), _candle(4, 96.5)]}
    results = run_parameter_sweep(_base_preset(), candles, volumes={"BTC": 9_000_000}, grid={"flash_crash_trigger_pct": [4.5]}, segments=1)

    report = render_parameter_sweep_report(results, top_n=1)

    assert "Parameter Sweep" in report
    assert "base__" in report
    assert "params=" in report
