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