CoolFace
Apppublic

Qionk/a-share-quant

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
daily_update.py181 linesDownload Raw Back to scripts
1#!/usr/bin/env python32"""3每日增量更新:检查数据库中所有股票,自动拉取缺失的交易日数据4 5用法:6  # 本地运行(需配好 .streamlit/secrets.toml)7  python scripts/daily_update.py8 9  # 通过环境变量指定数据库10  MYSQL_HOST=xxx MYSQL_PORT=3308 MYSQL_USER=xxx MYSQL_PASSWORD=xxx MYSQL_DATABASE=xxx python scripts/daily_update.py11 12  # GitHub Actions 示例:13  # - cron: '0 10 * * 1-5'  # 工作日每天早上10点14  # - env: MYSQL_HOST/MYSQL_PORT/MYSQL_USER/MYSQL_PASSWORD/MYSQL_DATABASE via Secrets15"""16import sys, os, time17from datetime import datetime18 19sys.stdout.reconfigure(line_buffering=True) if hasattr(sys.stdout, "reconfigure") else None20 21ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))22sys.path.insert(0, ROOT)23 24import pymysql25import pandas as pd26import numpy as np27 28 29def get_conn():30    host = os.environ.get("MYSQL_HOST", "")31    port = int(os.environ.get("MYSQL_PORT", "3306"))32    user = os.environ.get("MYSQL_USER", "")33    password = os.environ.get("MYSQL_PASSWORD", "")34    database = os.environ.get("MYSQL_DATABASE", "")35 36    if not host or not user or not database:37        # fallback: read from .streamlit/secrets.toml38        try:39            import tomllib40            secrets_path = os.path.join(ROOT, ".streamlit", "secrets.toml")41            if os.path.exists(secrets_path):42                with open(secrets_path, "rb") as f:43                    secrets = tomllib.load(f)44                host = secrets.get("MYSQL_HOST", "")45                port = int(secrets.get("MYSQL_PORT", "3306"))46                user = secrets.get("MYSQL_USER", "")47                password = secrets.get("MYSQL_PASSWORD", "")48                database = secrets.get("MYSQL_DATABASE", "")49        except Exception:50            pass51 52    if not host or not user or not database:53        print("[错误] MySQL 连接信息未配置!")54        print("请设置环境变量: MYSQL_HOST, MYSQL_PORT, MYSQL_USER, MYSQL_PASSWORD, MYSQL_DATABASE")55        print("或在 .streamlit/secrets.toml 中配置")56        sys.exit(1)57 58    return pymysql.connect(59        host=host, port=port, user=user, password=password,60        database=database, charset="utf8mb4",61        connect_timeout=10, read_timeout=30, write_timeout=30,62        autocommit=True,63    )64 65 66def safe_float(v):67    if v is None:68        return None69    try:70        f = float(v)71        return None if (np.isnan(f) or np.isinf(f)) else f72    except (TypeError, ValueError):73        return None74 75 76def main():77    print(f"\n{'='*60}")78    print(f"  每日增量数据更新")79    print(f"  时间: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")80    print(f"{'='*60}")81 82    conn = get_conn()83    cur = conn.cursor()84    print("\n[OK] MySQL 连接成功")85 86    # 1. 列出所有股票及其最新日期87    cur.execute("""88        SELECT stock_code, stock_name, MAX(trade_date) AS last_date89        FROM stock_daily_data90        GROUP BY stock_code, stock_name91        ORDER BY stock_code92    """)93    stocks = [(r[0], r[1], r[2]) for r in cur.fetchall()]94 95    if not stocks:96        print("数据库中没有股票,无需更新")97        cur.close()98        conn.close()99        return100 101    print(f"\n共 {len(stocks)} 只股票待检查\n")102 103    # 2. 导入 AKShare104    from src.predict.data_input import load_from_akshare105 106    INSERT_SQL = """107        INSERT INTO stock_daily_data108        (stock_code, stock_name, trade_date, open, high, low, close, volume, amount, pct_change, turnover)109        VALUES (%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s)110        ON DUPLICATE KEY UPDATE111        open=VALUES(open), high=VALUES(high), low=VALUES(low), close=VALUES(close),112        volume=VALUES(volume), amount=VALUES(amount), pct_change=VALUES(pct_change),113        turnover=VALUES(turnover), stock_name=VALUES(stock_name)114    """115 116    today = datetime.now()117    updated = 0118    skipped = 0119    failed = 0120 121    for code, name, last_date in stocks:122        # 判断是否需要更新:最新日期 < 今天(允许1天延迟,A股数据次日更新)123        if isinstance(last_date, str):124            last_date = datetime.strptime(last_date, "%Y-%m-%d").date()125        elif hasattr(last_date, "date"):126            last_date = last_date.date()127 128        days_behind = (today.date() - last_date).days129        # 周末/假期没有交易数据,days_behind == 0 表示已有当天数据才跳过130        if days_behind == 0:131            skipped += 1132            continue133 134        try:135            start = (last_date + pd.Timedelta(days=1)).strftime("%Y%m%d")136            end = today.strftime("%Y%m%d")137            print(f"[更新] {code} {name}  缺 {days_behind} 天  ({last_date} → {today.date()})")138 139            df = load_from_akshare(code, start, end)140            if df is None or len(df) == 0:141                skipped += 1142                continue143 144            rows = []145            for idx, row in df.iterrows():146                trade_date = idx.date() if hasattr(idx, "date") else pd.Timestamp(idx).date()147                rows.append((148                    code, name, trade_date,149                    safe_float(row.get("open")),150                    safe_float(row.get("high")),151                    safe_float(row.get("low")),152                    safe_float(row.get("close")),153                    safe_float(row.get("volume")),154                    safe_float(row.get("amount")),155                    safe_float(row.get("pct_change")),156                    safe_float(row.get("turnover")),157                ))158 159            for i in range(0, len(rows), 500):160                cur.executemany(INSERT_SQL, rows[i:i + 500])161 162            updated += 1163            print(f"  [OK] +{len(rows)} 条")164 165            time.sleep(1)  # 反爬间隔166 167        except Exception as e:168            failed += 1169            print(f"  [失败] {e}")170 171    # 3. 汇总172    print(f"\n{'='*60}")173    print(f"  完成!  更新: {updated}  跳过(已有最新): {skipped}  失败: {failed}")174    print(f"{'='*60}\n")175 176    cur.close()177    conn.close()178 179 180if __name__ == "__main__":181    main()