"""Tests for NeuralBacktester."""

import pandas as pd
import numpy as np

from backtest.engine import NeuralBacktester, BacktestPosition
from domain import TradeDirection, DEFAULT_RISK_PER_TRADE


def _make_price_bar(timestamp: int, open_p: float, high: float, low: float, close: float, volume: int = 1000):
    return {
        "timestamp": timestamp,
        "Open": open_p,
        "High": high,
        "Low": low,
        "Close": close,
        "Volume": volume,
    }


def _make_signal(timestamp: int, signal_type: str = "LONG", entry_price: float = None,
                  sl_price: float = None, tp_price: float = None,
                  ai_confidence: float = 0.8):
    sig = {
        "timestamp": timestamp,
        "signal_type": signal_type,
        "entry_price": entry_price,
        "ai_confidence": ai_confidence,
    }
    if sl_price is not None:
        sig["sl_price"] = sl_price
    if tp_price is not None:
        sig["tp_price"] = tp_price
    return sig


class TestNeuralBacktesterInit:
    def test_default_initialization(self):
        bt = NeuralBacktester()
        assert bt.capital == 1_000_000
        assert bt.risk_pct == DEFAULT_RISK_PER_TRADE
        assert bt.max_positions == 5
        assert bt.positions == []
        assert bt.trades == []

    def test_custom_capital(self):
        bt = NeuralBacktester(capital=500_000, risk_pct=0.02, max_positions=3)
        assert bt.capital == 500_000
        assert bt.risk_pct == 0.02
        assert bt.max_positions == 3


class TestNeuralBacktesterRun:
    def test_run_empty_signals(self):
        bt = NeuralBacktester()
        prices = pd.DataFrame([_make_price_bar(1, 100, 105, 95, 100)])
        trades = bt.run(pd.DataFrame(), prices)
        assert trades == []

    def test_run_empty_prices(self):
        bt = NeuralBacktester()
        signals = pd.DataFrame([_make_signal(1, "LONG", 100, 95, 110)])
        trades = bt.run(signals, pd.DataFrame())
        assert trades == []

    def test_run_long_profit(self):
        bt = NeuralBacktester(capital=1_000_000, risk_pct=0.01)
        ts0 = 1_000_000
        prices = pd.DataFrame([
            _make_price_bar(ts0, 100, 100, 100, 100),
            _make_price_bar(ts0 + 3600, 102, 112, 101, 111),
        ])
        signals = pd.DataFrame([_make_signal(ts0, "LONG", 100, 95, 110)])
        trades = bt.run(signals, prices, ticker="SBER")

        assert len(trades) == 1
        assert trades[0].direction == TradeDirection.BUY
        assert trades[0].status == "TP"
        assert trades[0].net_pnl > 0

    def test_run_long_loss(self):
        bt = NeuralBacktester(capital=1_000_000, risk_pct=0.01)
        ts0 = 1_000_000
        prices = pd.DataFrame([
            _make_price_bar(ts0, 100, 100, 100, 100),
            _make_price_bar(ts0 + 3600, 98, 99, 94, 94),
        ])
        signals = pd.DataFrame([_make_signal(ts0, "LONG", 100, 95, 110)])
        trades = bt.run(signals, prices, ticker="SBER")

        assert len(trades) == 1
        assert trades[0].status == "SL"
        assert trades[0].net_pnl < 0

    def test_run_short_profit(self):
        bt = NeuralBacktester(capital=1_000_000, risk_pct=0.01)
        ts0 = 1_000_000
        prices = pd.DataFrame([
            _make_price_bar(ts0, 100, 100, 100, 100),
            _make_price_bar(ts0 + 3600, 98, 99, 89, 89),
        ])
        signals = pd.DataFrame([_make_signal(ts0, "SHORT", 100, 105, 90)])
        trades = bt.run(signals, prices, ticker="SBER")

        assert len(trades) == 1
        assert trades[0].direction == TradeDirection.SELL
        assert trades[0].status == "TP"

    def test_run_short_loss(self):
        bt = NeuralBacktester(capital=1_000_000, risk_pct=0.01)
        ts0 = 1_000_000
        prices = pd.DataFrame([
            _make_price_bar(ts0, 100, 100, 100, 100),
            _make_price_bar(ts0 + 3600, 102, 106, 101, 106),
        ])
        signals = pd.DataFrame([_make_signal(ts0, "SHORT", 100, 105, 90)])
        trades = bt.run(signals, prices, ticker="SBER")

        assert len(trades) == 1
        assert trades[0].status == "SL"

    def test_run_multiple_signals(self):
        bt = NeuralBacktester(capital=1_000_000, risk_pct=0.01)
        ts0 = 1_000_000
        prices = pd.DataFrame([
            _make_price_bar(ts0, 100, 100, 100, 100),
            _make_price_bar(ts0 + 3600, 100, 105, 95, 100),
            _make_price_bar(ts0 + 7200, 100, 110, 95, 105),
        ])
        signals = pd.DataFrame([
            _make_signal(ts0, "LONG", 100, 95, 110),
            _make_signal(ts0 + 3600, "SHORT", 100, 105, 90),
        ])
        trades = bt.run(signals, prices, ticker="SBER")

        assert len(trades) == 2

    def test_equity_recorded(self):
        bt = NeuralBacktester(capital=1_000_000, risk_pct=0.01)
        ts0 = 1_000_000
        prices = pd.DataFrame([
            _make_price_bar(ts0, 100, 100, 100, 100),
            _make_price_bar(ts0 + 3600, 105, 110, 104, 109),
        ])
        signals = pd.DataFrame([_make_signal(ts0, "LONG", 100, 95, 110)])
        bt.run(signals, prices, ticker="SBER")

        assert len(bt.equity) > 0
        assert bt.equity[0]["equity"] == 1_000_000


class TestNeuralBacktesterSave:
    def test_save_trades(self, tmp_path):
        bt = NeuralBacktester(capital=1_000_000, risk_pct=0.01)
        ts0 = 1_000_000
        prices = pd.DataFrame([
            _make_price_bar(ts0, 100, 100, 100, 100),
            _make_price_bar(ts0 + 3600, 102, 112, 101, 111),
        ])
        signals = pd.DataFrame([_make_signal(ts0, "LONG", 100, 95, 110)])
        bt.run(signals, prices, ticker="SBER")

        bt.save_trades("SBER", "H1", output_dir=tmp_path)
        csv_files = list(tmp_path.glob("trades_SBER_H1.csv"))
        assert len(csv_files) == 1

    def test_save_trades_no_trades(self, tmp_path):
        bt = NeuralBacktester()
        bt.save_trades("SBER", "H1", output_dir=tmp_path)
        assert len(list(tmp_path.glob("*.csv"))) == 0


class TestBacktestPosition:
    def test_position_creation(self):
        pos = BacktestPosition(
            entry_time=1_000_000,
            entry=100.0,
            sl=95.0,
            tp=110.0,
            size=10,
            direction=TradeDirection.BUY,
        )
        assert pos.entry == 100.0
        assert pos.direction == TradeDirection.BUY
        assert pos.size == 10
