Qionk/a-share-quant
0
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()