from __future__ import annotations

from historical_replay import Candle
from strategy_lab import LabStrategy, default_research_sampler_strategies, render_strategy_lab_report, run_strategy_lab


def _candle(i: int, close: float, *, high: float | None = None, low: float | None = None, volume: float = 100.0) -> Candle:
    return Candle(
        ts=1_700_000_000 + i * 60,
        open=close,
        high=high if high is not None else close,
        low=low if low is not None else close,
        close=close,
        volume=volume,
    )


def test_momentum_breakout_enters_after_trend_breakout_and_takes_profit() -> None:
    strategy = LabStrategy(
        strategy_id="lab_momentum_breakout",
        family="momentum_breakout",
        parameters={
            "lookback": 5,
            "breakout_pct": 1.0,
            "take_profit_pct": 1.5,
            "stop_loss_pct": 2.0,
            "max_hold_bars": 8,
            "trade_size_usd": 100.0,
        },
    )
    candles = {
        "BTC": [
            _candle(0, 100), _candle(1, 101), _candle(2, 102), _candle(3, 103), _candle(4, 104),
            _candle(5, 105.2), _candle(6, 107.0),
        ]
    }

    result = run_strategy_lab([strategy], candles, volumes={"BTC": 100_000_000})[0]

    assert result.strategy_id == "lab_momentum_breakout"
    assert result.closed_trades == 1
    assert result.trades[0].exit_reason == "take_profit"
    assert result.total_pnl_usd > 0


def test_mean_reversion_rsi_enters_oversold_dump_and_exits_on_rebound() -> None:
    strategy = LabStrategy(
        strategy_id="lab_rsi_reversion",
        family="rsi_mean_reversion",
        parameters={
            "rsi_period": 3,
            "oversold_rsi": 25.0,
            "exit_rsi": 55.0,
            "take_profit_pct": 4.0,
            "stop_loss_pct": 3.0,
            "max_hold_bars": 10,
            "trade_size_usd": 100.0,
        },
    )
    candles = {"ETH": [_candle(i, close) for i, close in enumerate([100, 98, 96, 94, 95, 96, 97, 98])]}

    result = run_strategy_lab([strategy], candles, volumes={"ETH": 100_000_000})[0]

    assert result.closed_trades == 1
    assert result.trades[0].exit_reason in {"rsi_recovered", "take_profit", "end_of_replay"}
    assert result.trades[0].coin == "ETH"


def test_liquidation_reversal_enters_volume_spike_dump_and_exits_on_snapback() -> None:
    strategy = LabStrategy(
        strategy_id="lab_liquidation_reversal",
        family="liquidation_reversal",
        parameters={
            "lookback": 5,
            "drop_pct": 3.0,
            "volume_spike_mult": 2.0,
            "take_profit_pct": 2.0,
            "stop_loss_pct": 1.8,
            "max_hold_bars": 8,
            "trade_size_usd": 100.0,
        },
    )
    candles = {
        "WLD": [
            _candle(0, 100, volume=100),
            _candle(1, 101, volume=100),
            _candle(2, 102, volume=100),
            _candle(3, 101, volume=100),
            _candle(4, 100, volume=100),
            _candle(5, 96, low=95, volume=400),
            _candle(6, 98.5, high=99, volume=240),
        ]
    }

    result = run_strategy_lab([strategy], candles, volumes={"WLD": 100_000_000})[0]

    assert result.closed_trades == 1
    assert result.trades[0].exit_reason == "take_profit"
    assert result.total_pnl_usd > 0


def test_liquidation_reversal_rejects_dump_without_volume_spike() -> None:
    strategy = LabStrategy(
        strategy_id="lab_liquidation_reversal",
        family="liquidation_reversal",
        parameters={"lookback": 5, "drop_pct": 3.0, "volume_spike_mult": 2.0, "max_hold_bars": 8},
    )
    candles = {
        "WLD": [
            _candle(0, 100, volume=100),
            _candle(1, 101, volume=100),
            _candle(2, 102, volume=100),
            _candle(3, 101, volume=100),
            _candle(4, 100, volume=100),
            _candle(5, 96, low=95, volume=110),
            _candle(6, 98.5, high=99, volume=120),
        ]
    }

    result = run_strategy_lab([strategy], candles, volumes={"WLD": 100_000_000})[0]

    assert result.closed_trades == 0


def test_relative_strength_rotation_enters_only_top_ranked_coin() -> None:
    strategy = LabStrategy(
        strategy_id="lab_relative_strength_rotation",
        family="relative_strength_rotation",
        parameters={
            "lookback": 4,
            "top_n": 1,
            "min_momentum_pct": 3.0,
            "take_profit_pct": 2.0,
            "stop_loss_pct": 2.0,
            "max_hold_bars": 4,
            "trade_size_usd": 100.0,
            "max_total_trades": 1,
        },
    )
    candles = {
        "BTC": [_candle(i, close) for i, close in enumerate([100, 101, 102, 103, 104, 104.5])],
        "SOL": [_candle(i, close) for i, close in enumerate([100, 103, 106, 109, 112, 115])],
    }

    result = run_strategy_lab([strategy], candles, volumes={"BTC": 100_000_000, "SOL": 100_000_000})[0]

    assert result.closed_trades == 1
    assert result.trades[0].coin == "SOL"
    assert result.total_pnl_usd > 0


def test_risk_managed_trend_following_requires_aligned_fast_slow_and_long_trends() -> None:
    strategy = LabStrategy(
        strategy_id="lab_risk_managed_trend_following",
        family="risk_managed_trend_following",
        parameters={
            "fast_sma": 3,
            "slow_sma": 5,
            "long_lookback": 8,
            "min_long_trend_pct": 4.0,
            "take_profit_pct": 2.0,
            "stop_loss_pct": 2.0,
            "max_hold_bars": 4,
            "trade_size_usd": 100.0,
        },
    )
    candles = {
        "ETH": [_candle(i, close) for i, close in enumerate([100, 101, 102, 103, 104, 105, 106, 107, 108, 111])],
        "DOGE": [_candle(i, close) for i, close in enumerate([100, 101, 100, 99, 98, 97, 96, 95, 94, 93])],
    }

    result = run_strategy_lab([strategy], candles, volumes={"ETH": 100_000_000, "DOGE": 100_000_000})[0]

    assert result.closed_trades == 1
    assert result.trades[0].coin == "ETH"
    assert result.total_pnl_usd > 0


def test_risk_managed_trend_following_rejects_excessive_atr_chop() -> None:
    strategy = LabStrategy(
        strategy_id="lab_risk_managed_trend_following_atr_guard",
        family="risk_managed_trend_following",
        parameters={
            "fast_sma": 3,
            "slow_sma": 5,
            "long_lookback": 8,
            "min_long_trend_pct": 4.0,
            "atr_period": 4,
            "max_atr_pct": 2.0,
            "max_hold_bars": 4,
            "trade_size_usd": 100.0,
        },
    )
    candles = {
        "ETH": [
            _candle(0, 100, high=101, low=99),
            _candle(1, 101, high=106, low=96),
            _candle(2, 102, high=108, low=97),
            _candle(3, 103, high=109, low=98),
            _candle(4, 104, high=110, low=99),
            _candle(5, 105, high=112, low=100),
            _candle(6, 106, high=113, low=101),
            _candle(7, 107, high=114, low=102),
            _candle(8, 108, high=115, low=103),
            _candle(9, 111, high=119, low=104),
        ]
    }

    result = run_strategy_lab([strategy], candles, volumes={"ETH": 100_000_000})[0]

    assert result.closed_trades == 0


def test_relative_strength_rotation_goes_to_cash_when_market_breadth_is_too_thin() -> None:
    strategy = LabStrategy(
        strategy_id="lab_relative_strength_rotation_breadth_guard",
        family="relative_strength_rotation",
        parameters={
            "lookback": 4,
            "top_n": 2,
            "min_momentum_pct": 3.0,
            "min_positive_candidates": 2,
            "max_hold_bars": 4,
            "trade_size_usd": 100.0,
        },
    )
    candles = {
        "SOL": [_candle(i, close) for i, close in enumerate([100, 103, 106, 109, 112, 113])],
        "BTC": [_candle(i, close) for i, close in enumerate([100, 99, 100, 99, 100, 99])],
        "ETH": [_candle(i, close) for i, close in enumerate([100, 100, 99, 100, 99, 100])],
    }

    result = run_strategy_lab([strategy], candles, volumes={"SOL": 100_000_000, "BTC": 100_000_000, "ETH": 100_000_000})[0]

    assert result.closed_trades == 0


def test_donchian_volume_breakout_requires_range_high_break_and_volume_confirmation() -> None:
    strategy = LabStrategy(
        strategy_id="lab_donchian_volume_breakout",
        family="donchian_volume_breakout",
        parameters={
            "lookback": 5,
            "breakout_pct": 0.5,
            "volume_spike_mult": 1.8,
            "take_profit_pct": 1.5,
            "stop_loss_pct": 1.2,
            "max_hold_bars": 4,
            "trade_size_usd": 100.0,
        },
    )
    candles = {
        "HYPE": [
            _candle(0, 100, high=101, volume=100),
            _candle(1, 100.2, high=101, volume=100),
            _candle(2, 99.8, high=101, volume=100),
            _candle(3, 100.1, high=101, volume=100),
            _candle(4, 100.0, high=101, volume=100),
            _candle(5, 102.0, high=102.2, volume=250),
            _candle(6, 104.0, high=104.2, volume=180),
        ],
        "SUI": [
            _candle(0, 100, high=101, volume=100),
            _candle(1, 100.2, high=101, volume=100),
            _candle(2, 99.8, high=101, volume=100),
            _candle(3, 100.1, high=101, volume=100),
            _candle(4, 100.0, high=101, volume=100),
            _candle(5, 102.0, high=102.2, volume=120),
            _candle(6, 104.0, high=104.2, volume=120),
        ],
    }

    result = run_strategy_lab([strategy], candles, volumes={"HYPE": 100_000_000, "SUI": 100_000_000})[0]

    assert result.closed_trades == 1
    assert result.trades[0].coin == "HYPE"
    assert result.trades[0].exit_reason == "take_profit"


def test_generate_lab_grid_strategies_covers_independent_archetypes() -> None:
    from strategy_lab import generate_lab_grid_strategies

    strategies = generate_lab_grid_strategies()

    families = {strategy.family for strategy in strategies}
    assert {
        "momentum_breakout",
        "rsi_mean_reversion",
        "trend_pullback",
        "volatility_squeeze_breakout",
        "liquidation_reversal",
        "relative_strength_rotation",
        "risk_managed_trend_following",
        "donchian_volume_breakout",
    }.issubset(families)
    assert len({strategy.strategy_id for strategy in strategies}) == len(strategies)
    assert len(strategies) >= 25


def test_strategy_lab_grid_includes_defensive_trend_and_rotation_guards() -> None:
    from strategy_lab import generate_lab_grid_strategies

    strategies = generate_lab_grid_strategies()

    assert any(strategy.family == "risk_managed_trend_following" and "max_atr_pct" in strategy.parameters for strategy in strategies)
    assert any(strategy.family == "relative_strength_rotation" and "min_positive_candidates" in strategy.parameters for strategy in strategies)


def test_research_sampler_strategies_are_labelled_not_champion_candidates() -> None:
    samplers = default_research_sampler_strategies()

    assert {s.strategy_id for s in samplers} == {"sampler_squeeze_breakout_research", "sampler_relative_strength_research"}
    assert all(s.parameters["research_sampler"] is True for s in samplers)
    assert all(s.parameters["not_champion"] is True for s in samplers)
    assert all(s.parameters["paper_only"] is True for s in samplers)

    report = render_strategy_lab_report([], sampler_strategies=samplers)
    assert "Research sampler — not a champion candidate" in report
    assert "sampler_squeeze_breakout_research" in report
    assert "sampler_relative_strength_research" in report


def test_strategy_lab_ranks_by_score_and_marks_observation_only() -> None:
    winner = LabStrategy(
        strategy_id="winner",
        family="momentum_breakout",
        parameters={"lookback": 3, "breakout_pct": 0.5, "take_profit_pct": 1.0, "stop_loss_pct": 3.0, "max_hold_bars": 8},
    )
    loser = LabStrategy(
        strategy_id="loser",
        family="momentum_breakout",
        parameters={"lookback": 3, "breakout_pct": 0.5, "take_profit_pct": 8.0, "stop_loss_pct": 0.5, "max_hold_bars": 8},
    )
    candles = {"SOL": [_candle(i, close) for i, close in enumerate([100, 101, 102, 103, 105, 104, 103, 102])]}

    results = run_strategy_lab([loser, winner], candles, volumes={"SOL": 100_000_000})
    report = render_strategy_lab_report(results, min_samples=5)

    assert [result.strategy_id for result in results] == ["winner", "loser"]
    assert "Strategy Lab Tournament" in report
    assert "observation-only" in report
