CoolFace
Apppublic

Qionk/a-share-quant

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
stock_data_store.py460 linesDownload Raw Back to predict
1"""2股票日线数据 MySQL 存取3store: AKShare → MySQL4load:  MySQL → DataFrame5list:  所有已存储股票概览6"""7 8import os9import time as _time10import pandas as pd11import numpy as np12from datetime import datetime13 14BATCH_SIZE = 50015_TABLE_ENSURED = False16 17 18def _get_conn():19    try:20        import streamlit as st21        host = st.secrets.get("MYSQL_HOST", os.environ.get("MYSQL_HOST", ""))22        port = int(st.secrets.get("MYSQL_PORT", os.environ.get("MYSQL_PORT", "3306")))23        user = st.secrets.get("MYSQL_USER", os.environ.get("MYSQL_USER", ""))24        password = st.secrets.get("MYSQL_PASSWORD", os.environ.get("MYSQL_PASSWORD", ""))25        database = st.secrets.get("MYSQL_DATABASE", os.environ.get("MYSQL_DATABASE", ""))26    except Exception:27        host = os.environ.get("MYSQL_HOST", "")28        port = int(os.environ.get("MYSQL_PORT", "3306"))29        user = os.environ.get("MYSQL_USER", "")30        password = os.environ.get("MYSQL_PASSWORD", "")31        database = os.environ.get("MYSQL_DATABASE", "")32 33    if not host or not user or not database:34        return None35 36    import pymysql37    return pymysql.connect(38        host=host, port=port, user=user, password=password,39        database=database, charset="utf8mb4",40        connect_timeout=10, read_timeout=30, write_timeout=30,41        autocommit=True,42    )43 44 45def _ensure_unique_index(conn):46    """确保 stock_daily_data 有 (stock_code, trade_date) 唯一索引,只执行一次"""47    global _TABLE_ENSURED48    if _TABLE_ENSURED:49        return50    cur = conn.cursor()51    try:52        cur.execute("""53            CREATE TABLE IF NOT EXISTS stock_daily_data (54                id INT AUTO_INCREMENT PRIMARY KEY,55                stock_code VARCHAR(20) NOT NULL,56                stock_name VARCHAR(100),57                trade_date DATE NOT NULL,58                open DOUBLE, high DOUBLE, low DOUBLE, close DOUBLE,59                volume DOUBLE, amount DOUBLE, pct_change DOUBLE, turnover DOUBLE,60                UNIQUE KEY uk_stock_date (stock_code, trade_date),61                INDEX idx_stock_code (stock_code)62            ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb463        """)64    except Exception:65        pass66    # 检查唯一索引是否已存在67    cur.execute("SHOW INDEX FROM stock_daily_data WHERE Key_name='uk_stock_date'")68    if not cur.fetchall():69        # 索引不存在,先去重再加70        try:71            cur.execute("""72                DELETE d1 FROM stock_daily_data d173                INNER JOIN stock_daily_data d274                WHERE d1.id < d2.id75                  AND d1.stock_code = d2.stock_code76                  AND d1.trade_date = d2.trade_date77            """)78        except Exception:79            pass80        try:81            cur.execute("ALTER TABLE stock_daily_data ADD UNIQUE KEY uk_stock_date (stock_code, trade_date)")82        except Exception:83            pass84    cur.close()85    _TABLE_ENSURED = True86 87 88_DEDUP_DONE = False89 90 91def _dedup_stock_table():92    """清理 stock_daily_data 中的重复行并确保唯一索引,整个进程只执行一次"""93    global _DEDUP_DONE94    if _DEDUP_DONE:95        return96    conn = _get_conn()97    if not conn:98        return99    _ensure_unique_index(conn)100    conn.close()101    _DEDUP_DONE = True102 103 104def store_stock_data(stock_code: str, stock_name: str, df: pd.DataFrame) -> int:105    """将 DataFrame 写入 stock_daily_data,返回写入行数"""106    conn = _get_conn()107    if not conn:108        return 0109 110    _ensure_unique_index(conn)111    cur = conn.cursor()112    sql = """113        INSERT INTO stock_daily_data114        (stock_code, stock_name, trade_date, open, high, low, close, volume, amount, pct_change, turnover)115        VALUES (%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s)116        ON DUPLICATE KEY UPDATE117        open=VALUES(open), high=VALUES(high), low=VALUES(low), close=VALUES(close),118        volume=VALUES(volume), amount=VALUES(amount), pct_change=VALUES(pct_change),119        turnover=VALUES(turnover), stock_name=VALUES(stock_name)120    """121 122    rows = []123    for idx, row in df.iterrows():124        trade_date = idx.date() if hasattr(idx, "date") else pd.Timestamp(idx).date()125        rows.append((126            stock_code, stock_name, trade_date,127            _safe_float(row.get("open")),128            _safe_float(row.get("high")),129            _safe_float(row.get("low")),130            _safe_float(row.get("close")),131            _safe_float(row.get("volume")),132            _safe_float(row.get("amount")),133            _safe_float(row.get("pct_change")),134            _safe_float(row.get("turnover")),135        ))136 137    for i in range(0, len(rows), BATCH_SIZE):138        batch = rows[i:i + BATCH_SIZE]139        cur.executemany(sql, batch)140 141    cur.close()142    conn.close()143    return len(rows)144 145 146def _trim_stock_data(stock_code: str, keep_days: int):147    """删除某股票超过 keep_days 天的旧数据"""148    conn = _get_conn()149    if not conn:150        return151    cur = conn.cursor()152    cur.execute(153        "SELECT trade_date FROM stock_daily_data WHERE stock_code=%s "154        "ORDER BY trade_date DESC LIMIT 1 OFFSET %s",155        (stock_code, keep_days),156    )157    row = cur.fetchone()158    if row:159        cutoff = row[0]160        cur.execute(161            "DELETE FROM stock_daily_data WHERE stock_code=%s AND trade_date < %s",162            (stock_code, cutoff),163        )164    cur.close()165    conn.close()166 167 168def delete_stock_data(stock_code: str):169    """删除某股票的全部数据(用于重新获取全量数据)"""170    conn = _get_conn()171    if not conn:172        return173    cur = conn.cursor()174    cur.execute("DELETE FROM stock_daily_data WHERE stock_code=%s", (stock_code,))175    cur.close()176    conn.close()177 178 179def _safe_float(v):180    if v is None:181        return None182    try:183        f = float(v)184        return None if (np.isnan(f) or np.isinf(f)) else f185    except (TypeError, ValueError):186        return None187 188 189def load_stock_from_db(stock_code: str) -> pd.DataFrame:190    """从 MySQL 加载股票日线数据,返回 DataFrame(date 索引,去重)"""191    conn = _get_conn()192    if not conn:193        return pd.DataFrame()194 195    cur = conn.cursor()196    cur.execute(197        "SELECT trade_date, open, high, low, close, volume, amount, pct_change, turnover "198        "FROM stock_daily_data WHERE stock_code=%s ORDER BY trade_date",199        (stock_code,))200    cols = [d[0] for d in cur.description]201    rows = cur.fetchall()202    cur.close()203    conn.close()204 205    if not rows:206        return pd.DataFrame()207 208    df = pd.DataFrame(rows, columns=cols)209    df["trade_date"] = pd.to_datetime(df["trade_date"])210    df = df.set_index("trade_date")211    df = df[~df.index.duplicated(keep="last")]212    df.index.name = "date"213    return df214 215 216def list_db_stocks() -> list:217    """返回已存储的股票列表 [{code, name, rows, start_date, end_date}],数据量受 max_days 上限约束"""218    conn = _get_conn()219    if not conn:220        return []221    cur = conn.cursor()222    cur.execute("""223        SELECT stock_code, stock_name, COUNT(DISTINCT trade_date) AS data_rows,224               MIN(trade_date) AS start_date, MAX(trade_date) AS end_date225        FROM stock_daily_data226        GROUP BY stock_code, stock_name227        ORDER BY stock_code228    """)229    cols = [d[0] for d in cur.description]230    rows = [dict(zip(cols, r)) for r in cur.fetchall()]231    for r in rows:232        r["rows"] = r.pop("data_rows")233    for r in rows:234        r["code"] = r["stock_code"]235        r["name"] = r["stock_name"]236        r["start_date"] = r["start_date"].strftime("%Y-%m-%d") if hasattr(r["start_date"], "strftime") else str(r["start_date"])237        r["end_date"] = r["end_date"].strftime("%Y-%m-%d") if hasattr(r["end_date"], "strftime") else str(r["end_date"])238    cur.close()239    conn.close()240    return rows241 242 243def list_stocks_with_status() -> list:244    """统一视图: 股票数据 + 训练状态 [{code, name, data_rows, start_date, end_date,245       trained, trained_at, trained_models, session_id}]"""246    _dedup_stock_table()247    conn = _get_conn()248    if not conn:249        return []250    cur = conn.cursor()251    cur.execute("""252        SELECT d.stock_code, d.stock_name,253               COUNT(DISTINCT d.trade_date) AS data_rows,254               MIN(d.trade_date) AS start_date,255               MAX(d.trade_date) AS end_date,256               MAX(t.trained_at) AS trained_at,257               MAX(t.selected_models) AS trained_models,258               MAX(t.id) AS session_id259        FROM stock_daily_data d260        LEFT JOIN training_sessions t ON d.stock_code = t.stock_code261        GROUP BY d.stock_code, d.stock_name262        ORDER BY d.stock_code263    """)264    cols = [d[0] for d in cur.description]265    rows = [dict(zip(cols, r)) for r in cur.fetchall()]266    for r in rows:267        r["code"] = r["stock_code"]268        r["name"] = r["stock_name"]269        r["rows"] = r.get("data_rows", 0)270        r["trained"] = r["trained_at"] is not None271        r["start_date"] = r["start_date"].strftime("%Y-%m-%d") if hasattr(r["start_date"], "strftime") else str(r["start_date"])272        r["end_date"] = r["end_date"].strftime("%Y-%m-%d") if hasattr(r["end_date"], "strftime") else str(r["end_date"])273        if r["trained_at"] and hasattr(r["trained_at"], "strftime"):274            r["trained_at"] = r["trained_at"].strftime("%Y-%m-%d %H:%M")275        if r.get("trained_models"):276            import json277            r["trained_models"] = json.loads(r["trained_models"]) if isinstance(r["trained_models"], str) else r["trained_models"]278    cur.close()279    conn.close()280    return rows281 282 283def has_stock_data(stock_code: str) -> bool:284    conn = _get_conn()285    if not conn:286        return False287    cur = conn.cursor()288    cur.execute("SELECT 1 FROM stock_daily_data WHERE stock_code=%s LIMIT 1", (stock_code,))289    exists = cur.fetchone() is not None290    cur.close()291    conn.close()292    return exists293 294 295def get_stock_name_from_db(stock_code: str) -> str:296    conn = _get_conn()297    if not conn:298        return ""299    cur = conn.cursor()300    cur.execute("SELECT stock_name FROM stock_daily_data WHERE stock_code=%s LIMIT 1", (stock_code,))301    row = cur.fetchone()302    cur.close()303    conn.close()304    return row[0] if row else ""305 306 307def list_stock_sessions(stock_code: str) -> list:308    """返回某股票的所有训练记录 [{session_id, trained_at, trained_models, forecast_days}]"""309    conn = _get_conn()310    if not conn:311        return []312    cur = conn.cursor()313    cur.execute(314        "SELECT id, trained_at, selected_models, forecast_days "315        "FROM training_sessions WHERE stock_code=%s ORDER BY trained_at DESC",316        (stock_code,))317    cols = [d[0] for d in cur.description]318    rows = [dict(zip(cols, r)) for r in cur.fetchall()]319    for r in rows:320        r["session_id"] = r["id"]321        if hasattr(r["trained_at"], "strftime"):322            r["trained_at"] = r["trained_at"].strftime("%Y-%m-%d %H:%M")323        if r.get("selected_models"):324            import json325            r["trained_models"] = json.loads(r["selected_models"]) if isinstance(r["selected_models"], str) else r["selected_models"]326    cur.close()327    conn.close()328    return rows329 330 331def _fill_turnover_if_missing(stock_code: str, stock_name: str, df: pd.DataFrame,332                              start_date: str, end_date: str) -> pd.DataFrame:333    """如果 turnover 缺失超过 50%,从网易源补充并批量更新 DB"""334    if 'turnover' not in df.columns or df.empty:335        return df336    valid_ratio = df['turnover'].notna().mean()337    if valid_ratio >= 0.5:338        return df339    try:340        from src.predict.data_input import _fetch_turnover_netease341        turnover_s = _fetch_turnover_netease(stock_code, start_date, end_date)342        if turnover_s is not None and not turnover_s.empty:343            df['turnover'] = turnover_s.reindex(df.index)344            # 批量写回 DB345            conn = _get_conn()346            if conn:347                cur = conn.cursor()348                updates = []349                for idx in df.index:350                    t_val = _safe_float(df.loc[idx, 'turnover'])351                    if t_val is not None:352                        trade_date = idx.date() if hasattr(idx, 'date') else pd.Timestamp(idx).date()353                        updates.append((t_val, stock_code, trade_date))354                if updates:355                    cur.executemany(356                        "UPDATE stock_daily_data SET turnover=%s "357                        "WHERE stock_code=%s AND trade_date=%s",358                        updates)359                cur.close()360                conn.close()361    except Exception:362        pass363    return df364 365 366def fetch_and_store(stock_code: str, start_date: str = "20200101",367                    end_date: str = None, max_days: int = 500,368                    progress_callback=None) -> tuple:369    """370    从 AKShare 获取数据并存入 MySQL371    start_date/end_date: YYYYMMDD 格式372    max_days: 最多保留最近多少交易日(None=不限制,建议500)373    返回: (DataFrame, stock_name, 是否已有数据)374    """375    if end_date is None:376        end_date = datetime.now().strftime("%Y%m%d")377 378    # 首次调用时确保唯一索引 + 清理历史重复数据379    _dedup_stock_table()380 381    if progress_callback:382        progress_callback("checking")383 384    # 已存在则检查是否需要重新获取385    if has_stock_data(stock_code):386        df = load_stock_from_db(stock_code)387        name = get_stock_name_from_db(stock_code)388 389        if df.empty:390            # 空数据,走重新获取流程391            pass392        else:393            last_date = df.index[-1]394            today = pd.Timestamp.now().normalize()395            db_start = df.index[0].strftime("%Y%m%d") if hasattr(df.index[0], 'strftime') else str(df.index[0])[:10].replace("-", "")396 397            # 若请求的起始日期早于DB中最早日期,需要重新获取全量数据398            if start_date < db_start:399                if progress_callback:400                    progress_callback("refetching")401                delete_stock_data(stock_code)402                # 跳出 if 块,走下面的全量获取403            else:404                # 已有数据在请求范围内,仅增量更新405                if last_date < today:406                    if progress_callback:407                        progress_callback("updating")408                    from src.predict.data_input import load_from_akshare409                    new_start = (last_date + pd.Timedelta(days=1)).strftime("%Y%m%d")410                    new_end = end_date or datetime.now().strftime("%Y%m%d")411                    try:412                        df_new = load_from_akshare(stock_code, new_start, new_end)413                        if df_new is not None and len(df_new) > 0:414                            store_stock_data(stock_code, name, df_new)415                            df = pd.concat([df, df_new]).sort_index()416                            df = df[~df.index.duplicated(keep="last")]417                    except Exception as e:418                        if progress_callback:419                            progress_callback("update_failed")420 421                # 按请求的起始日期截取422                req_start = pd.Timestamp(start_date)423                df = df[df.index >= req_start]424 425                # 补充换手率:如果 DB 数据中 turnover 大部分为空,尝试网易源补齐426                df = _fill_turnover_if_missing(stock_code, name, df, start_date, end_date)427 428                # 按 max_days 截取429                if max_days and len(df) > max_days:430                    df = df.tail(max_days)431                    _trim_stock_data(stock_code, max_days)432                if progress_callback:433                    progress_callback("done")434                return df, name, True435 436    if progress_callback:437        progress_callback("fetching")438 439    # AKShare 获取440    from src.predict.data_input import load_from_akshare, get_stock_name441    df = load_from_akshare(stock_code, start_date, end_date)442    name = get_stock_name(stock_code)443 444    # 截取最近 max_days 天445    if max_days and len(df) > max_days:446        df = df.tail(max_days)447 448    if progress_callback:449        progress_callback("storing")450 451    store_stock_data(stock_code, name, df)452 453    # 首次获取后也清理超出限制的旧数据454    if max_days:455        _trim_stock_data(stock_code, max_days)456 457    if progress_callback:458        progress_callback("done")459 460    return df, name, False