"""Technical indicators engine — builds feature matrix from OHLCV data."""
import logging
from typing import Dict, List, Optional

import numpy as np
import pandas as pd
import ta

from backend.config import config

logger = logging.getLogger(__name__)


def compute_features(df: pd.DataFrame) -> pd.DataFrame:
    """Compute all technical indicators for OHLCV DataFrame.
    
    Input: DataFrame with columns [Open, High, Low, Close, Volume]
    Output: Same DataFrame with added indicator columns.
    """
    fc = config.features
    
    # Ensure numeric
    for col in ["Open", "High", "Low", "Close", "Volume"]:
        if col in df.columns:
            df[col] = pd.to_numeric(df[col], errors="coerce")
    
    close = df["Close"]
    high = df["High"]
    low = df["Low"]
    volume = df["Volume"]
    
    # --- SMA ---
    for p in fc.get("sma_periods", [10, 20, 50]):
        df[f"SMA_{p}"] = ta.trend.sma_indicator(close, window=p)
        df[f"SMA_{p}_ratio"] = close / df[f"SMA_{p}"] - 1  # % away from SMA
    
    # --- EMA ---
    for p in fc.get("ema_periods", [12, 26, 50]):
        df[f"EMA_{p}"] = ta.trend.ema_indicator(close, window=p)
        df[f"EMA_{p}_ratio"] = close / df[f"EMA_{p}"] - 1
    
    # --- RSI ---
    rsi_period = fc.get("rsi_period", 14)
    df["RSI"] = ta.momentum.rsi(close, window=rsi_period)
    df["RSI_diff"] = df["RSI"].diff()
    
    # --- MACD ---
    macd = ta.trend.MACD(
        close,
        window_fast=fc.get("macd_fast", 12),
        window_slow=fc.get("macd_slow", 26),
        window_sign=fc.get("macd_signal", 9),
    )
    df["MACD"] = macd.macd()
    df["MACD_signal"] = macd.macd_signal()
    df["MACD_histogram"] = macd.macd_diff()
    
    # --- Bollinger Bands ---
    bb = ta.volatility.BollingerBands(
        close,
        window=fc.get("bb_period", 20),
        window_dev=fc.get("bb_std", 2),
    )
    df["BB_upper"] = bb.bollinger_hband()
    df["BB_middle"] = bb.bollinger_mavg()
    df["BB_lower"] = bb.bollinger_lband()
    df["BB_width"] = (df["BB_upper"] - df["BB_lower"]) / df["BB_middle"]
    df["BB_pct"] = (close - df["BB_lower"]) / (df["BB_upper"] - df["BB_lower"])
    
    # --- Volume ---
    vol_sma = fc.get("volume_sma", 20)
    df["Volume_SMA"] = ta.trend.sma_indicator(volume, window=vol_sma)
    df["Volume_ratio"] = volume / df["Volume_SMA"]
    
    # --- ATR ---
    atr_period = fc.get("atr_period", 14)
    df["ATR"] = ta.volatility.average_true_range(high, low, close, window=atr_period)
    df["ATR_pct"] = df["ATR"] / close * 100
    
    # --- Stochastic ---
    stoch_k = fc.get("stoch_k", 14)
    stoch_d = fc.get("stoch_d", 3)
    stoch = ta.momentum.StochasticOscillator(high, low, close, window=stoch_k, smooth_window=stoch_d)
    df["Stoch_K"] = stoch.stoch()
    df["Stoch_D"] = stoch.stoch_signal()
    
    # --- CCI (in ta.trend, not momentum) ---
    cci = ta.trend.CCIIndicator(high, low, close, window=20)
    df["CCI"] = cci.cci()
    
    # --- OBV ---
    df["OBV"] = ta.volume.on_balance_volume(close, volume)
    df["OBV_SMA"] = ta.trend.sma_indicator(df["OBV"], window=20)
    
    # --- Price change features ---
    df["Return_1"] = close.pct_change(1)
    df["Return_3"] = close.pct_change(3)
    df["Return_5"] = close.pct_change(5)
    df["Volatility_10"] = df["Return_1"].rolling(10).std()
    
    # --- EMA Crossover signals ---
    if "EMA_12" in df.columns and "EMA_26" in df.columns:
        df["EMA_cross"] = (df["EMA_12"] > df["EMA_26"]).astype(int)
    
    # Drop NaN rows created by indicators (keep last rows for prediction)
    initial_len = len(df)
    df_clean = df.dropna()
    dropped = initial_len - len(df_clean)
    if dropped > 0:
        logger.debug(f"Dropped {dropped} rows with NaN indicators")
    
    return df_clean


def engineer_features(df: pd.DataFrame, lookahead: int = 5) -> pd.DataFrame:
    """Add target columns for prediction.
    
    Creates:
      - target_1h: % change in next bar (regression target)
      - target_direction: 1=up, 0=down (classification target)
      - target_pct: actual % change ahead
    
    Args:
        df: DataFrame with computed indicators
        lookahead: bars ahead to predict
    """
    close = df["Close"]
    
    # Future price
    df["future_close"] = close.shift(-lookahead)
    df["target_pct"] = (df["future_close"] - close) / close * 100
    
    # Classification target
    df["target_direction"] = (df["future_close"] > close).astype(int)
    
    # Drop rows where target is unknown (last N rows)
    df_predictable = df.dropna(subset=["future_close", "target_pct"])
    
    return df_predictable


def get_feature_columns(df: pd.DataFrame) -> List[str]:
    """Get list of feature columns (exclude metadata and targets)."""
    exclude = {"Date", "future_close", "target_pct", "target_direction",
               "Open", "High", "Low", "Close", "Volume"}
    return [c for c in df.columns if c not in exclude]


def prepare_training_data(df: pd.DataFrame) -> tuple:
    """Prepare X, y for model training.
    
    Returns:
        (X, y_classification, y_regression, feature_names)
    """
    features = get_feature_columns(df)
    X = df[features].values.astype(np.float32)
    y_class = df["target_direction"].values.astype(np.int32)
    y_reg = df["target_pct"].values.astype(np.float32)
    
    # Handle any remaining NaN/inf
    X = np.nan_to_num(X, nan=0.0, posinf=0.0, neginf=0.0)
    
    return X, y_class, y_reg, features
