from typing import Optional
import re
import numpy as np
import pandas as pd
from db.connection import get_connection
from config import TIMEFRAMES, MOEX_TICKERS

# Дополнительные тикеры (не MOEX), разрешённые для загрузки
EXTRA_TICKERS = {'BITCOIN', 'BITCOINC', 'EURUSD', 'ZCASH'}

# Белый список тикеров + таймфреймов для защиты от SQL-инъекции
VALID_TICKERS = {t.upper() for t in MOEX_TICKERS} | EXTRA_TICKERS
VALID_TFS = {'H1', 'D1', 'W1'}


def clean_spikes(df: pd.DataFrame, max_return: float = 0.90,
                 level_low: float = 0.30, level_high: float = 3.0) -> pd.DataFrame:
    """Обнаруживает и интерполирует аномальные свечи.
    
    Два прохода:
      Pass 1 — смежные возвраты > max_return (catch spike candles: DROP 35→2 + JUMP 2→35)
      Pass 2 — абсолютный уровень цены < level_low × median или > level_high × median
               (catch persistent corrupted candles: 6 свечей по ~2р между DROP и JUMP)
    
    SNGSP/PLZL: ошибочные данные — падение с ~35р → ~2р (94%) с возвратом через 6 свечей.
    Все O/H/L/C/V аномалий заменяются линейной интерполяцией между чистыми соседями.
    """
    if df is None or len(df) < 3:
        return df
    
    df = df.copy()
    close = df['Close'].values.astype(float)
    n = len(close)
    
    # --- Pass 1: аномалии по смежным возвратам ---
    returns = np.zeros(n)
    returns[1:] = np.abs(close[1:] / close[:-1] - 1.0)
    mask = returns > max_return
    
    # --- Pass 2: аномалии по абсолютному уровню ---
    clean_close = close[~mask]
    if mask.any() and len(clean_close) > 10:
        median_close = np.median(clean_close)
        level_mask = (close < level_low * median_close) | (close > level_high * median_close)
        if level_mask.any():
            new_flags = level_mask & (~mask)
            mask = mask | level_mask
    
    if not mask.any():
        return df
    
    total_spikes = mask.sum()
    spike_idx = np.where(mask)[0]
    
    # --- Интерполяция O/H/L/C ---
    for col in ['Open', 'High', 'Low', 'Close']:
        vals = df[col].values.astype(float)
        
        for idx in spike_idx:
            prev_idx = idx - 1
            # Ищем предыдущую чистую свечу
            while prev_idx >= 0 and mask[prev_idx]:
                prev_idx -= 1
            
            next_idx = idx + 1
            # Ищем следующую чистую свечу
            while next_idx < n and mask[next_idx]:
                next_idx += 1
            
            if prev_idx >= 0 and next_idx < n and next_idx > prev_idx:
                span = next_idx - prev_idx
                vals[idx] = vals[prev_idx] + (vals[next_idx] - vals[prev_idx]) * (idx - prev_idx) / span
            elif prev_idx >= 0:
                vals[idx] = vals[prev_idx]
            else:
                vals[idx] = vals[next_idx] if next_idx < n else vals[idx]
        
        df[col] = vals
    
    # --- Volume: среднее окружающих чистых свечей ---
    vol = df['Volume'].values.astype(float)
    for idx in spike_idx:
        prev_v = vol[idx - 1] if idx > 0 else 0
        next_v = vol[idx + 1] if idx + 1 < n else 0
        vol[idx] = (prev_v + next_v) / 2 if (prev_v > 0 or next_v > 0) else vol[idx]
    df['Volume'] = vol.astype(int)
    
    return df


def _validate_table_name(ticker: str, timeframe: str) -> str:
    """Валидирует тикер и таймфрейм перед подстановкой в SQL."""
    ticker = ticker.upper()
    tf = timeframe.upper()
    if ticker not in VALID_TICKERS:
        raise ValueError(f"Недопустимый тикер: {ticker}")
    if tf not in VALID_TFS:
        raise ValueError(f"Недопустимый таймфрейм: {tf}")
    return f"{ticker}_{tf}"


def load_dataframe(ticker: str, timeframe: str, limit: Optional[int] = None,
                   clean: bool = True) -> pd.DataFrame:
    """Загружает свечные данные. Если clean=True — удаляет ценовые аномалии."""
    table = _validate_table_name(ticker, timeframe)
    if limit:
        limit_int = int(limit)
        if limit_int <= 0:
            raise ValueError(f"LIMIT должен быть положительным: {limit_int}")
        query = """
            SELECT * FROM (
                SELECT timestamp, `Date`, `Time`, Open, High, Low, Close, Volume
                FROM {}
                ORDER BY timestamp DESC
                LIMIT %s
            ) sub
            ORDER BY timestamp ASC
        """.format(table)
    else:
        query = """
            SELECT timestamp, `Date`, `Time`, Open, High, Low, Close, Volume
            FROM {}
            ORDER BY timestamp
        """.format(table)

    with get_connection() as conn:
        if limit:
            df = pd.read_sql(query, conn, params=(limit_int,))
        else:
            df = pd.read_sql(query, conn)

    # Очистка ценовых аномалий (SNGSP, PLZL)
    if clean and df is not None and len(df) >= 3:
        df = clean_spikes(df)

    return df


def load_multiple_timeframes(ticker: str, timeframes: Optional[list[str]] = None, limit: Optional[int] = None) -> dict[str, pd.DataFrame]:
    if timeframes is None:
        timeframes = list(TIMEFRAMES.keys())
    return {
        tf: load_dataframe(ticker, tf, limit)
        for tf in timeframes
    }


def load_multitimeframe_features(
    ticker: str,
    main_tf: str = 'H1',
    higher_tfs: Optional[list[str]] = None,
    limit: Optional[int] = None,
) -> pd.DataFrame:
    from features.technical import engineer_features, FEATURE_COLS

    if higher_tfs is None:
        higher_tfs = ['D1', 'W1']

    all_tfs = [main_tf] + [tf for tf in higher_tfs if tf != main_tf]
    dfs = load_multiple_timeframes(ticker, all_tfs, limit)

    main_df = engineer_features(dfs[main_tf])

    for tf in higher_tfs:
        if tf == main_tf:
            continue
        df_h = engineer_features(dfs[tf], skip_temporal=True)
        rename = {col: f'{tf}_{col}' for col in FEATURE_COLS}
        df_h = df_h[['timestamp'] + FEATURE_COLS].rename(columns=rename)

        # Shift higher-TF features by 1 period to avoid lookahead
        # D1 bar at day D gets timestamp of day D+1, so H1 bar on day D uses day D-1's D1 data
        tf_seconds = {'D1': 86400, 'W1': 604800}
        shift_s = tf_seconds.get(tf, 86400)
        df_h['timestamp'] = df_h['timestamp'] + shift_s

        main_df = pd.merge_asof(
            main_df.sort_values('timestamp'),
            df_h.sort_values('timestamp'),
            on='timestamp',
            direction='backward',
        )

    return main_df
