import numpy as np
import pandas as pd


def add_curated_context(df: pd.DataFrame) -> pd.DataFrame:
    df = df.copy()
    o, c, h, l = df['Open'], df['Close'], df['High'], df['Low']
    body = (c - o).abs()
    candle_range = h - l

    df['doji'] = (body < 0.1 * candle_range).astype(float)

    df['engulfing'] = 0
    bullish_eng = (c.shift(1) < o.shift(1)) & (c > o) & (c >= o.shift(1)) & (o <= c.shift(1))
    bearish_eng = (c.shift(1) > o.shift(1)) & (c < o) & (c <= o.shift(1)) & (o >= c.shift(1))
    df.loc[bullish_eng, 'engulfing'] = 1
    df.loc[bearish_eng, 'engulfing'] = -1

    higher_high = (h > h.shift(1)).astype(int)
    lower_low = (l < l.shift(1)).astype(int)
    df['hh_streak'] = 0
    df['ll_streak'] = 0
    streak = 0
    for i in range(1, len(df)):
        if higher_high.iloc[i]:
            streak = streak + 1 if higher_high.iloc[i-1] else 1
        else:
            streak = 0
        df.loc[df.index[i], 'hh_streak'] = streak
    streak = 0
    for i in range(1, len(df)):
        if lower_low.iloc[i]:
            streak = streak + 1 if lower_low.iloc[i-1] else 1
        else:
            streak = 0
        df.loc[df.index[i], 'll_streak'] = streak

    for col in ['rsi_signal', 'macd_signal', 'directional_bias']:
        if col in df.columns:
            df[f'{col}_accel'] = df[col] - df[col].shift(3)

    if 'signal_strength' in df.columns:
        df['strength_trend'] = df['signal_strength'].rolling(5).mean()

    # ── Расстояние от последнего HH/LL ─────────────────────────────────
    # Чем дольше рынок не обновлял хаи/лои, тем выше вероятность
    # скорого сильного движения (сжатая пружина).
    h = df['High'].values
    l = df['Low'].values
    n = len(df)

    last_hh_bars = np.zeros(n, dtype=np.int32)
    last_ll_bars = np.zeros(n, dtype=np.int32)
    hh_val = np.full(n, np.nan)
    ll_val = np.full(n, np.nan)

    last_hh_idx = -1
    last_hh_high = -1.0
    last_ll_idx = -1
    last_ll_low = 1e10

    for i in range(n):
        # Обновляем максимум (скользящий highest high за всё время)
        if h[i] > last_hh_high or last_hh_idx == -1:
            last_hh_idx = i
            last_hh_high = h[i]
        # Обновляем минимум
        if l[i] < last_ll_low or last_ll_idx == -1:
            last_ll_idx = i
            last_ll_low = l[i]

        last_hh_bars[i] = i - last_hh_idx
        last_ll_bars[i] = i - last_ll_idx
        hh_val[i] = float(last_hh_high)
        ll_val[i] = float(last_ll_low)

    df['bars_since_hh'] = last_hh_bars
    df['bars_since_ll'] = last_ll_bars
    df['dist_from_hh'] = (df['Close'] - hh_val) / np.clip(hh_val, 1e-10, None)
    df['dist_from_ll'] = (df['Close'] - ll_val) / np.clip(ll_val, 1e-10, None)

    return df


CONTEXT_COLS = [
    'doji', 'engulfing', 'hh_streak', 'll_streak',
    'rsi_signal_accel', 'macd_signal_accel',
    'directional_bias_accel', 'strength_trend',
    'bars_since_hh', 'bars_since_ll',
    'dist_from_hh', 'dist_from_ll',
]
