import pandas as pd
import numpy as np


def compute_market_context(snapshot_df: pd.DataFrame) -> dict:
    if snapshot_df is None or snapshot_df.empty:
        return {}

    features = {}
    chg_col = _find(snapshot_df.columns.tolist(), ["涨跌幅", "涨跌幅(%)"])
    vol_col = _find(snapshot_df.columns.tolist(), ["成交额", "成交额(元)"])
    amp_col = _find(snapshot_df.columns.tolist(), ["振幅", "振幅(%)"])

    if chg_col:
        chg = pd.to_numeric(snapshot_df[chg_col], errors="coerce").dropna()
        features["market_advance_count"] = (chg > 0).sum()
        features["market_decline_count"] = (chg < 0).sum()
        features["market_flat_count"] = (chg == 0).sum()
        total = len(chg)
        features["market_advance_ratio"] = features["market_advance_count"] / total if total > 0 else 0.5
        features["market_mean_return"] = chg.mean()
        features["market_median_return"] = chg.median()
        features["market_return_std"] = chg.std()
        features["market_top5_return"] = chg.nlargest(5).mean()
        features["market_bottom5_return"] = chg.nsmallest(5).mean()

    if vol_col:
        vol = pd.to_numeric(snapshot_df[vol_col], errors="coerce").dropna()
        features["market_total_volume"] = vol.sum()
        features["market_avg_volume"] = vol.mean()

    if amp_col:
        amp = pd.to_numeric(snapshot_df[amp_col], errors="coerce").dropna()
        features["market_mean_amplitude"] = amp.mean()

    return features


def compute_sector_features(board_df: pd.DataFrame, stock_code: str = None) -> dict:
    if board_df is None or board_df.empty:
        return {}

    features = {}
    chg_col = _find(board_df.columns.tolist(), ["涨跌幅", "涨跌幅(%)"])
    name_col = _find(board_df.columns.tolist(), ["板块名称"])
    leader_col = _find(board_df.columns.tolist(), ["领涨股票"])

    if chg_col:
        chg = pd.to_numeric(board_df[chg_col], errors="coerce").dropna()
        features["sector_count"] = len(chg)
        features["sector_mean_return"] = chg.mean()
        features["sector_top3_return"] = chg.nlargest(3).mean()
        features["sector_bottom3_return"] = chg.nsmallest(3).mean()
        features["sector_return_std"] = chg.std()

        if name_col:
            top_sectors = board_df.nlargest(5, chg_col)[name_col].tolist()
            features["top5_sectors"] = ",".join(top_sectors)

    return features


def compute_stock_sector_rank(board_df: pd.DataFrame, stock_sectors: list = None) -> dict:
    if board_df is None or board_df.empty or not stock_sectors:
        return {"stock_sector_rank": 0, "stock_in_hot_sector": 0}

    chg_col = _find(board_df.columns.tolist(), ["涨跌幅", "涨跌幅(%)"])
    name_col = _find(board_df.columns.tolist(), ["板块名称"])

    if not chg_col or not name_col:
        return {"stock_sector_rank": 0, "stock_in_hot_sector": 0}

    board_df = board_df.copy()
    board_df["_chg"] = pd.to_numeric(board_df[chg_col], errors="coerce")
    board_df = board_df.sort_values("_chg", ascending=False).reset_index(drop=True)
    board_df["_rank"] = range(1, len(board_df) + 1)

    sector_ranks = []
    for sector in stock_sectors:
        match = board_df[board_df[name_col] == sector]
        if not match.empty:
            sector_ranks.append(match["_rank"].iloc[0])

    if sector_ranks:
        avg_rank = np.mean(sector_ranks)
        total = len(board_df)
        return {
            "stock_sector_rank": avg_rank,
            "stock_sector_rank_pct": avg_rank / total,
            "stock_in_hot_sector": 1 if avg_rank <= total * 0.2 else 0,
        }
    return {"stock_sector_rank": 0, "stock_in_hot_sector": 0}


def _find(cols: list, candidates: list) -> str:
    for c in candidates:
        if c in cols:
            return c
    return ""
