import numpy as np
import pandas as pd

ACTION_HOLD, ACTION_BUY, ACTION_SELL = 0, 1, 2

ENTRY_PENALTY = 0.1
DENSE_SCALE = 1.0
TIMEOUT_PENALTY = 0.5


class TradingEnv:
    def __init__(
        self,
        df: pd.DataFrame,
        feature_cols: list[str],
        atr_mult_sl: float = 1.5,
        atr_mult_tp: float = 3.0,
    ):
        self.df = df.reset_index(drop=True)
        self.feature_cols = feature_cols
        self.atr_mult_sl = atr_mult_sl
        self.atr_mult_tp = atr_mult_tp

        self.opens = df['Open'].values.astype(np.float64)
        self.closes = df['Close'].values.astype(np.float64)
        self.highs = df['High'].values.astype(np.float64)
        self.lows = df['Low'].values.astype(np.float64)
        self.atr_vals = df['atr'].values.astype(np.float64)
        self.feature_data = df[feature_cols].values.astype(np.float32)
        self.n_features = len(feature_cols) + 3
        self.reset()

    def reset(self, start_idx: int = 0) -> np.ndarray:
        self._idx = start_idx
        self._position = 0
        self._entry_price = 0.0
        self._sl_price = 0.0
        self._tp_price = 0.0
        self._entry_atr = 0.0
        self._entry_direction = 0  # +1 long, -1 short
        self._position_bars = 0
        self._closed_pnl = 0.0
        self._total_reward = 0.0
        self._trade_count = 0
        self._prev_pnl = 0.0
        self._done = False
        return self._get_state()

    def _get_state(self) -> np.ndarray:
        i = self._idx
        feats = np.nan_to_num(self.feature_data[i], nan=0.0, posinf=0.0, neginf=0.0)
        pos_flag = np.array([float(self._position)], dtype=np.float32)
        pnl = np.array([float(self._closed_pnl)], dtype=np.float32)
        bars = np.array([float(self._position_bars)], dtype=np.float32)
        return np.concatenate([feats, pos_flag, pnl, bars])

    def _mark_to_market(self) -> float:
        i = self._idx
        if self._position == 1:
            return (self.closes[i] - self._entry_price) / self._entry_price
        elif self._position == -1:
            return (self._entry_price - self.closes[i]) / self._entry_price
        return 0.0

    def step(self, action: int) -> tuple[np.ndarray, float, bool, dict]:
        if self._done:
            return np.zeros(self.n_features, dtype=np.float32), 0.0, True, {}

        i = self._idx
        reward = 0.0
        done_signal = False

        # — Dense step reward (mark-to-market PnL change) —
        if self._position != 0:
            current_pnl = self._mark_to_market()
            step_pnl = current_pnl - self._prev_pnl
            reward += step_pnl * DENSE_SCALE
            self._prev_pnl = current_pnl

        # — Check SL/TP on current bar —
        if self._position != 0:
            self._position_bars += 1
            high, low = self.highs[i], self.lows[i]
            tp_hit, sl_hit = False, False

            if self._position == 1:
                tp_hit = high >= self._tp_price
                sl_hit = low <= self._sl_price
            else:
                tp_hit = low <= self._tp_price
                sl_hit = high >= self._sl_price

            if tp_hit:
                reward += 2.0
                self._closed_pnl += self.atr_mult_tp * self._entry_atr / self._entry_price
                self._position = 0
                self._trade_count += 1
                done_signal = True
            elif sl_hit:
                reward += -1.0
                self._closed_pnl += -self.atr_mult_sl * self._entry_atr / self._entry_price
                self._position = 0
                self._trade_count += 1
                done_signal = True

        # — Enter new position (на следующем баре по Open) —
        if self._position == 0:
            if action in (ACTION_BUY, ACTION_SELL):
                reward -= ENTRY_PENALTY
                direction = 1 if action == ACTION_BUY else -1
                self._position = direction
                self._entry_direction = direction
                # Вход на следующем баре по Open (исправление lookahead bias)
                next_idx = min(i + 1, len(self.opens) - 1)
                self._entry_price = self.opens[next_idx]
                self._entry_atr = self.atr_vals[i]

                if direction == 1:
                    self._sl_price = self._entry_price - self.atr_vals[i] * self.atr_mult_sl
                    self._tp_price = self._entry_price + self.atr_vals[i] * self.atr_mult_tp
                else:
                    self._sl_price = self._entry_price + self.atr_vals[i] * self.atr_mult_sl
                    self._tp_price = self._entry_price - self.atr_vals[i] * self.atr_mult_tp

                self._position_bars = 0
                self._prev_pnl = 0.0

        # — Force-close on timeout —
        if self._position != 0 and self._position_bars >= 100 and not done_signal:
            exit_pnl = self._mark_to_market()
            reward -= TIMEOUT_PENALTY
            self._closed_pnl += exit_pnl
            self._position = 0
            self._trade_count += 1

        self._idx += 1
        self._total_reward += reward

        if self._idx >= len(self.df):
            self._done = True

        info = {
            'position': self._position,
            'total_reward': self._total_reward,
            'trade_count': self._trade_count,
            'closed_pnl': self._closed_pnl,
        }

        if self._done:
            return np.zeros(self.n_features, dtype=np.float32), reward, True, info
        return self._get_state(), reward, False, info
