import numpy as np
import pandas as pd
import torch
from rl.environment import TradingEnv
from rl.dqn import DQNAgent


def get_rl_features(multi: bool = False) -> list[str]:
    if multi:
        from main import get_directional_feature_cols
        multi_feats = get_directional_feature_cols(multi)
        return [c for c in multi_feats if c not in ('outcome_long', 'outcome_short')]
    from features.technical import FEATURE_COLS
    from features.directional import DIRECTIONAL_COLS
    from features.context import CONTEXT_COLS
    return FEATURE_COLS + DIRECTIONAL_COLS + CONTEXT_COLS


def train_dqn_chronological(
    df_train: pd.DataFrame,
    df_test: pd.DataFrame,
    feature_cols: list[str],
    n_epochs: int = 20,
    episode_size: int = 2000,
    atr_mult_sl: float = 1.5,
    atr_mult_tp: float = 3.0,
    hidden_dim: int = 128,
    lr: float = 1e-3,
    gamma: float = 0.99,
    epsilon_start: float = 1.0,
    epsilon_end: float = 0.01,
    epsilon_decay_steps: int = 5000,
    buffer_capacity: int = 50000,
    batch_size: int = 128,
    target_update: int = 200,
    device: str = 'cpu',
    verbose: bool = True,
):
    state_dim = len(feature_cols) + 3
    agent = DQNAgent(
        state_dim=state_dim, action_dim=3, hidden_dim=hidden_dim,
        lr=lr, gamma=gamma,
        epsilon_start=epsilon_start, epsilon_end=epsilon_end,
        epsilon_decay=1.0 - (1.0 / epsilon_decay_steps) if epsilon_decay_steps > 1 else 0.99,
        buffer_capacity=buffer_capacity,
        batch_size=batch_size, target_update=target_update,
        device=device,
    )

    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
        agent.optimizer, mode='max', factor=0.5, patience=3, min_lr=1e-5,
    )

    total_steps = 0
    episode_counter = 0
    history = {'epoch': [], 'episode': [], 'total_reward': [], 'trade_count': [], 'epsilon': []}
    best_test_return = -float('inf')
    best_state = None

    for epoch in range(n_epochs):
        for start in range(0, len(df_train), episode_size):
            end = min(start + episode_size, len(df_train))
            chunk = df_train.iloc[start:end].reset_index(drop=True)
            if len(chunk) < 100:
                continue

            env = TradingEnv(chunk, feature_cols, atr_mult_sl, atr_mult_tp)
            state = env.reset()
            done = False
            ep_reward = 0.0

            while not done:
                action = agent.act(state)
                next_state, reward, done, info = env.step(action)
                agent.memory.push(state, action, reward, next_state, done)
                state = next_state
                ep_reward += reward

                if len(agent.memory) >= batch_size:
                    loss = agent.update()  # Эпсилон теперь падает здесь (step-based decay)
                    total_steps += 1

            ep_trades = env._trade_count
            history['epoch'].append(epoch)
            history['episode'].append(episode_counter)
            history['total_reward'].append(ep_reward)
            history['trade_count'].append(ep_trades)
            history['epsilon'].append(agent.epsilon)
            episode_counter += 1

        # End of epoch — test evaluation
        test_result = evaluate_dqn(agent, df_test, feature_cols, atr_mult_sl, atr_mult_tp)
        test_ret = test_result['total_return']

        scheduler.step(test_ret)

        if test_ret > best_test_return:
            best_test_return = test_ret
            best_state = agent.q_net.state_dict().copy()

        if verbose:
            lr_now = agent.optimizer.param_groups[0]['lr']
            print(f"Epoch {epoch:2d}: test_trades={test_result['trade_count']:3d} "
                  f"wr={test_result['win_rate']:.1%} ret={test_ret:+.2%} "
                  f"eps={agent.epsilon:.3f} lr={lr_now:.1e}")

    # Restore best model
    if best_state is not None:
        agent.q_net.load_state_dict(best_state)
        agent.target_net.load_state_dict(best_state)

    return agent, pd.DataFrame(history)


def evaluate_dqn(
    agent: DQNAgent,
    df_test: pd.DataFrame,
    feature_cols: list[str],
    atr_mult_sl: float = 1.5,
    atr_mult_tp: float = 3.0,
) -> dict:
    env = TradingEnv(df_test.reset_index(drop=True), feature_cols, atr_mult_sl, atr_mult_tp)
    state = env.reset()
    done = False

    entry_price = None
    entry_atr = None
    entry_idx = None
    trades = []

    while not done:
        action = agent.act(state, eval_mode=True)
        next_state, reward, done, info = env.step(action)

        # Запоминаем точку входа
        if env._position != 0 and entry_price is None:
            entry_price = env._entry_price
            entry_atr = env._entry_atr
            entry_dir = env._entry_direction  # 1 или -1
            entry_idx = env._idx - 1  # idx после step() уже увеличен
        # Сделка закрыта — вычисляем реальную доходность
        elif env._position == 0 and entry_price is not None:
            if entry_idx is not None and entry_idx < len(env.closes):
                exit_price = env.closes[min(env._idx - 1, len(env.closes) - 1)]
                # Корректный return с учётом направления: long=(exit-entry)/entry, short=(entry-exit)/entry
                ret = entry_dir * (exit_price - entry_price) / entry_price
            else:
                # Fallback на старый метод если нет индекса
                if reward > 1.0:
                    ret = atr_mult_tp * entry_atr / entry_price
                elif reward < 0:
                    ret = -atr_mult_sl * entry_atr / entry_price
                else:
                    ret = (env._closed_pnl / len(trades) * 0.01) if trades else 0.0

            trades.append({'ret': ret, 'reward': reward, 'entry': entry_price, 'exit': exit_price if entry_idx is not None else None})
            entry_price = None
            entry_atr = None
            entry_idx = None

        state = next_state

    if trades:
        returns = np.array([t['ret'] for t in trades])
        win_rate = (returns > 0).mean()
        cumulative = (1 + returns).prod() - 1
    else:
        returns = np.array([])
        win_rate = 0.0
        cumulative = 0.0

    return {
        'trade_count': len(trades),
        'win_rate': win_rate,
        'total_return': cumulative,
        'closed_pnl': float(returns.sum()) if len(returns) > 0 else 0.0,
        'trades': trades,
    }


def print_rl_results(results: dict, label: str = ''):
    print(f"\n  RL Results {label}:")
    print(f"    Trades:      {results['trade_count']}")
    print(f"    Win Rate:    {results['win_rate']:.1%}")
    print(f"    Total Ret:   {results['total_return']:+.2%}")
    print(f"    Closed PnL:  {results['closed_pnl']:.4f}")
