import numpy as np
import pandas as pd


BACKTEST_CONFIG = {
    'long_prob_threshold': 0.55,
    'short_prob_threshold': 0.55,
    'atr_mult_sl': 1.5,
    'atr_mult_tp': 3.0,
    'risk_per_trade': 0.01,
    # Количество периодов в году для корректной годовой аннуализации Sharpe Ratio.
    # 252  = D1 (daily)
    # 2016 = H1  (252 × ~8 часов MOEX)
    # 52   = W1  (weekly)
    'periods_per_year': 252,
}


def simulate_trades(
    df: pd.DataFrame,
    long_probs: np.ndarray,
    short_probs: np.ndarray,
    config: dict | None = None,
) -> pd.DataFrame:
    cfg = {**BACKTEST_CONFIG, **(config or {})}
    threshold = cfg['long_prob_threshold']
    atr_mult_tp = cfg['atr_mult_tp']
    atr_mult_sl = cfg['atr_mult_sl']
    risk_pct = cfg['risk_per_trade']

    df = df.copy()
    df['long_prob'] = long_probs
    df['short_prob'] = short_probs
    df['trade_direction'] = 0
    df['trade_return'] = 0.0

    for i in range(len(df)):
        if long_probs[i] >= threshold and long_probs[i] > short_probs[i]:
            # LONG trade
            df.loc[df.index[i], 'trade_direction'] = 1
            if df.loc[df.index[i], 'outcome_long'] == 1:
                ret = atr_mult_tp * df.loc[df.index[i], 'atr'] / df.loc[df.index[i], 'Close']
            else:
                ret = -atr_mult_sl * df.loc[df.index[i], 'atr'] / df.loc[df.index[i], 'Close']
            df.loc[df.index[i], 'trade_return'] = ret

        elif short_probs[i] >= threshold and short_probs[i] > long_probs[i]:
            # SHORT trade
            df.loc[df.index[i], 'trade_direction'] = -1
            if df.loc[df.index[i], 'outcome_short'] == 1:
                ret = atr_mult_tp * df.loc[df.index[i], 'atr'] / df.loc[df.index[i], 'Close']
            else:
                ret = -atr_mult_sl * df.loc[df.index[i], 'atr'] / df.loc[df.index[i], 'Close']
            df.loc[df.index[i], 'trade_return'] = ret

    return df


def compute_returns(trades: pd.DataFrame, initial_capital: float = 100000.0) -> pd.DataFrame:
    cfg_equity = pd.DataFrame(index=trades.index)
    cfg_equity['equity'] = initial_capital
    cfg_equity['daily_return'] = 0.0

    capital = initial_capital
    for i in range(len(trades)):
        ret = trades['trade_return'].iloc[i]
        if ret != 0:
            capital *= (1 + ret)
        cfg_equity.loc[trades.index[i], 'equity'] = capital
        if i > 0:
            prev = cfg_equity['equity'].iloc[i - 1]
            cfg_equity.loc[trades.index[i], 'daily_return'] = (capital - prev) / prev if prev > 0 else 0

    cfg_equity['peak'] = cfg_equity['equity'].cummax()
    cfg_equity['drawdown'] = (cfg_equity['equity'] - cfg_equity['peak']) / cfg_equity['peak']
    return cfg_equity


def compute_metrics(trades: pd.DataFrame, equity_curve: pd.DataFrame,
                    periods_per_year: int = 252) -> dict:
    """
    Вычисляет метрики производительности стратегии.

    Args:
        trades: DataFrame с колонками trade_direction, trade_return, ...
        equity_curve: DataFrame с колонками equity, daily_return, drawdown
        periods_per_year: число периодов в году для аннуализации Sharpe Ratio.
                         252 = D1, 2016 = H1 (MOEX), 52 = W1.
    """
    trade_rows = trades[trades['trade_direction'] != 0]
    n_trades = len(trade_rows)
    if n_trades == 0:
        return {'n_trades': 0, 'error': 'no trades'}

    winners = trade_rows[trade_rows['trade_return'] > 0]
    losers = trade_rows[trade_rows['trade_return'] < 0]
    win_rate = len(winners) / n_trades

    total_return = (equity_curve['equity'].iloc[-1] / equity_curve['equity'].iloc[0]) - 1
    max_dd = equity_curve['drawdown'].min()

    bar_returns = equity_curve['daily_return']
    sharpe = np.nan
    if bar_returns.std() > 0:
        # Правильная аннуализация: sqrt(periods_per_year), а не sqrt(количество баров)
        sharpe = bar_returns.mean() / bar_returns.std() * np.sqrt(periods_per_year)

    avg_win = winners['trade_return'].mean() if len(winners) > 0 else 0
    avg_loss = losers['trade_return'].mean() if len(losers) > 0 else 0
    profit_factor = abs(winners['trade_return'].sum() / losers['trade_return'].sum()) if len(losers) > 0 and losers['trade_return'].sum() != 0 else float('inf')

    long_trades = trade_rows[trade_rows['trade_direction'] == 1]
    short_trades = trade_rows[trade_rows['trade_direction'] == -1]

    return {
        'n_trades': n_trades,
        'win_rate': win_rate,
        'total_return': total_return,
        'max_drawdown': max_dd,
        'sharpe_ratio': sharpe,
        'avg_win': avg_win,
        'avg_loss': avg_loss,
        'profit_factor': profit_factor,
        'n_long': len(long_trades),
        'long_win_rate': long_trades['trade_return'].gt(0).mean() if len(long_trades) > 0 else 0,
        'n_short': len(short_trades),
        'short_win_rate': short_trades['trade_return'].gt(0).mean() if len(short_trades) > 0 else 0,
    }


def print_metrics(metrics: dict, label: str = ''):
    print(f"\n{'='*50}")
    print(f"BACKTEST RESULTS {label}")
    print(f"{'='*50}")
    print(f"  Trades:          {metrics.get('n_trades', 'N/A')}")
    print(f"  Win Rate:        {metrics.get('win_rate', 0):.1%}")
    print(f"  Total Return:    {metrics.get('total_return', 0):.2%}")
    print(f"  Max Drawdown:    {metrics.get('max_drawdown', 0):.2%}")
    print(f"  Sharpe Ratio:    {metrics.get('sharpe_ratio', 0):.2f}")
    print(f"  Profit Factor:   {metrics.get('profit_factor', 0):.2f}")
    print(f"  Avg Win:         {metrics.get('avg_win', 0):.4f}")
    print(f"  Avg Loss:        {metrics.get('avg_loss', 0):.4f}")
    print(f"  Long Trades:     {metrics.get('n_long', 0)} (WR: {metrics.get('long_win_rate', 0):.1%})")
    print(f"  Short Trades:    {metrics.get('n_short', 0)} (WR: {metrics.get('short_win_rate', 0):.1%})")
