import numpy as np
import pandas as pd
import ta


def compute_technical_features(df: pd.DataFrame) -> pd.DataFrame:
    if df.empty or len(df) < 30:
        return df
    df = df.copy()
    n = len(df)

    df["ret_1"] = df["close"].pct_change(1)
    df["ret_5"] = df["close"].pct_change(5)
    df["ret_10"] = df["close"].pct_change(10)
    df["ret_20"] = df["close"].pct_change(20)

    df["log_ret_1"] = np.log(df["close"] / df["close"].shift(1))

    df["vol_5"] = df["log_ret_1"].rolling(5).std()
    df["vol_10"] = df["log_ret_1"].rolling(10).std()
    df["vol_20"] = df["log_ret_1"].rolling(20).std()

    df["rsi_6"] = ta.momentum.RSIIndicator(df["close"], window=6).rsi()
    df["rsi_14"] = ta.momentum.RSIIndicator(df["close"], window=14).rsi()

    macd = ta.trend.MACD(df["close"])
    df["macd"] = macd.macd()
    df["macd_signal"] = macd.macd_signal()
    df["macd_diff"] = macd.macd_diff()

    if n >= 20:
        bb = ta.volatility.BollingerBands(df["close"], window=20)
        df["bb_upper"] = bb.bollinger_hband()
        df["bb_lower"] = bb.bollinger_lband()
        df["bb_width"] = (df["bb_upper"] - df["bb_lower"]) / df["close"]
        df["bb_pct"] = bb.bollinger_pband()

    if n >= 14:
        df["atr_14"] = ta.volatility.AverageTrueRange(
            df["high"], df["low"], df["close"], window=14
        ).average_true_range()
        df["atr_ratio"] = df["atr_14"] / df["close"]

        df["adx_14"] = ta.trend.ADXIndicator(
            df["high"], df["low"], df["close"], window=14
        ).adx()

        df["mfi_14"] = ta.volume.MFIIndicator(
            df["high"], df["low"], df["close"], df["volume"], window=14
        ).money_flow_index()

    df["obv"] = ta.volume.OnBalanceVolumeIndicator(df["close"], df["volume"]).on_balance_volume()
    df["obv_ret"] = df["obv"].pct_change(5)

    df["vwma_5"] = (df["close"] * df["volume"]).rolling(5).sum() / df["volume"].rolling(5).sum()
    df["vwma_20"] = (df["close"] * df["volume"]).rolling(20).sum() / df["volume"].rolling(20).sum()
    df["price_to_vwma5"] = df["close"] / df["vwma_5"]
    df["price_to_vwma20"] = df["close"] / df["vwma_20"]

    df["ma5"] = df["close"].rolling(5).mean()
    df["ma10"] = df["close"].rolling(10).mean()
    df["ma20"] = df["close"].rolling(20).mean()
    if n >= 60:
        df["ma60"] = df["close"].rolling(60).mean()
    df["price_to_ma5"] = df["close"] / df["ma5"]
    df["price_to_ma20"] = df["close"] / df["ma20"]
    df["ma5_cross_ma20"] = (df["ma5"] > df["ma20"]).astype(int)

    df["vol_ratio_5"] = df["volume"] / df["volume"].rolling(5).mean()
    df["vol_ratio_10"] = df["volume"] / df["volume"].rolling(10).mean()
    df["vol_ratio_20"] = df["volume"] / df["volume"].rolling(20).mean()

    df["amihud"] = df["ret_1"].abs() / (df["amount"] + 1)
    df["amihud_5"] = df["amihud"].rolling(5).mean()

    if "turnover" in df.columns:
        df["turnover_5"] = df["turnover"].rolling(5).mean()
        df["turnover_ratio"] = df["turnover"] / df["turnover_5"]

    df["high_low_ratio"] = df["high"] / df["low"]
    df["open_close_ratio"] = df["open"] / df["close"]
    df["upper_shadow"] = (df["high"] - df[["open", "close"]].max(axis=1)) / df["close"]
    df["lower_shadow"] = (df[["open", "close"]].min(axis=1) - df["low"]) / df["close"]

    df["stoch_k"] = ta.momentum.StochasticOscillator(
        df["high"], df["low"], df["close"]
    ).stoch()

    df["cci_20"] = ta.trend.CCIIndicator(
        df["high"], df["low"], df["close"], window=20
    ).cci()

    df["willr_14"] = ta.momentum.WilliamsRIndicator(
        df["high"], df["low"], df["close"]
    ).williams_r()

    df["close"] = df["close"]
    return df


def get_latest_features(df: pd.DataFrame) -> dict:
    if df.empty:
        return {}
    feat_df = compute_technical_features(df)
    if feat_df.empty:
        return {}
    last = feat_df.iloc[-1]
    feature_cols = [
        "ret_1", "ret_5", "ret_10", "ret_20", "log_ret_1",
        "vol_5", "vol_10", "vol_20",
        "rsi_6", "rsi_14",
        "macd", "macd_signal", "macd_diff",
        "bb_upper", "bb_lower", "bb_width", "bb_pct",
        "atr_14", "atr_ratio", "adx_14",
        "obv", "obv_ret",
        "vwma_5", "vwma_20", "price_to_vwma5", "price_to_vwma20",
        "ma5", "ma10", "ma20", "ma60", "price_to_ma5", "price_to_ma20", "ma5_cross_ma20",
        "vol_ratio_5", "vol_ratio_10", "vol_ratio_20",
        "amihud", "amihud_5",
        "high_low_ratio", "open_close_ratio", "upper_shadow", "lower_shadow",
        "mfi_14", "stoch_k", "cci_20", "willr_14",
    ]
    features = {}
    for col in feature_cols:
        if col in last.index:
            val = last[col]
            features[col] = 0.0 if pd.isna(val) else float(val)
    return features
