import pandas as pd

from config import STOCK_FILTER


def filter_st(df: pd.DataFrame) -> pd.DataFrame:
    if df.empty:
        return df
    name_col = _find_col(df, ["名称", "股票名称", "name"])
    if name_col:
        mask = ~df[name_col].str.contains(r"ST|退", na=False)
        return df[mask].copy()
    return df


def filter_new_stocks(df: pd.DataFrame, min_days: int = 60) -> pd.DataFrame:
    if df.empty:
        return df
    ipo_col = _find_col(df, ["上市日期", "ipo_date", "上市天数"])
    if ipo_col and "上市天数" in df.columns:
        return df[df["上市天数"] >= min_days].copy()
    return df


def filter_low_turnover(df: pd.DataFrame, min_amount: float = 50_000_000) -> pd.DataFrame:
    if df.empty:
        return df
    amt_col = _find_col(df, ["成交额", "成交额(元)", "amount"])
    if amt_col:
        return df[df[amt_col] >= min_amount].copy()
    return df


def filter_suspended(df: pd.DataFrame) -> pd.DataFrame:
    if df.empty:
        return df
    price_col = _find_col(df, ["最新价", "收盘", "close", "price"])
    if price_col:
        return df[df[price_col] > 0].copy()
    return df


def filter_suspended_or_limit(df: pd.DataFrame) -> pd.DataFrame:
    if df.empty:
        return df
    df = filter_suspended(df)
    chg_col = _find_col(df, ["涨跌幅", "涨跌幅(%)", "change_pct"])
    if chg_col:
        df = df[(df[chg_col] > -9.9) & (df[chg_col] < 9.9)].copy()
    return df


def clean_snapshot(df: pd.DataFrame) -> pd.DataFrame:
    if df.empty:
        return df
    df = filter_st(df)
    df = filter_suspended_or_limit(df)
    amt_col = _find_col(df, ["成交额", "成交额(元)", "amount"])
    if amt_col and df[amt_col].sum() > 0:
        df = filter_low_turnover(df, STOCK_FILTER["min_turnover"])
    df = df.reset_index(drop=True)
    return df


def prepare_hist(df: pd.DataFrame) -> pd.DataFrame:
    if df.empty:
        return df
    col_map = {}
    for c in df.columns:
        cl = c.lower()
        if "日期" in c or "date" in cl:
            col_map[c] = "date"
        elif "开盘" in c or "open" in cl:
            col_map[c] = "open"
        elif "收盘" in c or "close" in cl:
            col_map[c] = "close"
        elif "最高" in c or "high" in cl:
            col_map[c] = "high"
        elif "最低" in c or "low" in cl:
            col_map[c] = "low"
        elif "成交量" in c or "volume" in cl:
            col_map[c] = "volume"
        elif "成交额" in c or "amount" in cl:
            col_map[c] = "amount"
        elif "换手率" in c or "turnover" in cl:
            col_map[c] = "turnover"
        elif "振幅" in c or "amplitude" in cl:
            col_map[c] = "amplitude"
        elif "涨跌幅" in c or "change" in cl:
            col_map[c] = "change_pct"
    df = df.rename(columns=col_map)
    if "date" in df.columns:
        df["date"] = pd.to_datetime(df["date"])
        df = df.sort_values("date").reset_index(drop=True)
    num_cols = ["open", "close", "high", "low", "volume", "amount", "turnover", "change_pct"]
    for c in num_cols:
        if c in df.columns:
            df[c] = pd.to_numeric(df[c], errors="coerce")
    return df


def _find_col(df: pd.DataFrame, candidates: list) -> str:
    for c in candidates:
        if c in df.columns:
            return c
    return ""
