feat:修改获取数据方式为按日期获取,而不是按股票循环,加快数据获取速度,增加按日期检查数据库缺失数据并补充的功能。

This commit is contained in:
2026-08-02 11:17:16 +08:00
parent 063f790650
commit e3089c25a6
2 changed files with 294 additions and 189 deletions
+223 -89
View File
@@ -307,7 +307,7 @@ def import_trade_cal(start_date: Optional[str] = None, end_date: Optional[str] =
def import_daily_for_stock(ts_code: str, start_date: str, end_date: str, conn) -> int:
"""
导入单只股票的日线行情
导入单只股票的日线行情(保留用于单只股票补充/重试)
返回导入的记录数
"""
pro = get_ts_pro()
@@ -325,11 +325,9 @@ def import_daily_for_stock(ts_code: str, start_date: str, end_date: str, conn) -
df = normalize_columns(df)
# 转换日期
if "trade_date" in df.columns:
df["trade_date"] = pd.to_datetime(df["trade_date"], format="%Y%m%d", errors="coerce")
# 数值列处理 NaN
numeric_cols = [
"open", "high", "low", "close", "pre_close", "change", "pct_chg",
"vol", "amount",
@@ -338,7 +336,6 @@ def import_daily_for_stock(ts_code: str, start_date: str, end_date: str, conn) -
if col in df.columns:
df[col] = pd.to_numeric(df[col], errors="coerce")
# 额外列 (Tushare Pro 不同版本返回字段可能不同)
for col in ["turnover_rate", "volume_ratio", "ma5", "ma10", "ma20", "ma_v_5", "ma_v_10", "ma_v_20"]:
if col not in df.columns:
df[col] = None
@@ -347,85 +344,145 @@ def import_daily_for_stock(ts_code: str, start_date: str, end_date: str, conn) -
return batch_insert("daily", df, conn, conflict_cols)
def import_daily_batch(
stock_list: List[str],
def _normalize_daily_df(df: pd.DataFrame) -> pd.DataFrame:
"""标准化 daily DataFrame 的列和类型 (供 import_daily_by_date 复用)"""
df = normalize_columns(df)
if "trade_date" in df.columns:
df["trade_date"] = pd.to_datetime(df["trade_date"], format="%Y%m%d", errors="coerce")
numeric_cols = [
"open", "high", "low", "close", "pre_close", "change", "pct_chg",
"vol", "amount",
]
for col in numeric_cols:
if col in df.columns:
df[col] = pd.to_numeric(df[col], errors="coerce")
for col in ["turnover_rate", "volume_ratio", "ma5", "ma10", "ma20", "ma_v_5", "ma_v_10", "ma_v_20"]:
if col not in df.columns:
df[col] = None
return df
def import_daily_by_date(
start_date: Optional[str] = None,
end_date: Optional[str] = None,
conn=None,
sleep_interval: float = 0.3,
):
"""
批量导入多只股票的日线行情
- stock_list: 股票代码列表
- sleep_interval: API 调用间隔 (避免频率限制)
按交易日批量导入日线行情 (高效模式)
使用 pro.daily(trade_date='YYYYMMDD') 一次性拉取全市场当日数据
大幅减少 API 调用次数: 约250交易日/年 × 16年 ≈ 4000次 (原来需要 5000股票 × 16年 = 80000次)
返回: 失败的交易日列表
"""
if start_date is None:
start_date = START_DATE
if end_date is None:
end_date = END_DATE
total = len(stock_list)
logger.info("=" * 60)
logger.info(
f"[3/7] 导入日线行情 (daily): {start_date} ~ {end_date}, "
f"{total} 只股票"
)
own_conn = conn is None
if own_conn:
conn = get_pg_connection()
conn = get_pg_connection()
# 从 trade_cal 获取交易日列表
try:
cursor = conn.cursor()
cursor.execute(
"""
SELECT DISTINCT cal_date FROM trade_cal
WHERE is_open = 1
AND cal_date >= %s AND cal_date <= %s
ORDER BY cal_date
""",
(start_date, end_date),
)
trade_dates = [row[0] for row in cursor.fetchall()]
cursor.close()
finally:
if own_conn:
conn.close()
conn = get_pg_connection()
total = len(trade_dates)
if total == 0:
logger.warning(f" 日期范围 {start_date} ~ {end_date} 内无交易日")
if own_conn:
conn.close()
return []
logger.info(f" 日期范围 {start_date} ~ {end_date}: 共 {total} 个交易日")
pro = get_ts_pro()
success_count = 0
fail_list = []
for i, ts_code in enumerate(stock_list, 1):
for i, td in enumerate(trade_dates, 1):
try:
n = import_daily_for_stock(ts_code, start_date, end_date, conn)
if n > 0:
success_count += 1
if i % 50 == 0 or i == total:
logger.info(f" 进度: {i}/{total} 成功={success_count} 失败={len(fail_list)}")
td_str = td.strftime("%Y%m%d") if hasattr(td, "strftime") else str(td).replace("-", "")
def fetch():
return pro.daily(trade_date=td_str)
df = fetch_with_retry(fetch, max_retries=3)
if df is None or df.empty:
logger.warning(f" [{td_str}] 返回空数据 (可能非交易日或API限制)")
continue
df = _normalize_daily_df(df)
conflict_cols = ["ts_code", "trade_date"]
batch_insert("daily", df, conn, conflict_cols)
success_count += 1
except Exception as e:
logger.error(f" [{ts_code}] 导入失败: {e}")
fail_list.append(ts_code)
conn.rollback()
logger.error(f" [{td}] 导入失败: {e}")
fail_list.append(str(td))
try:
conn.rollback()
except Exception:
pass
time.sleep(sleep_interval) # API 频率控制
if i % 50 == 0 or i == total:
logger.info(f" 进度: {i}/{total} 成功={success_count} 失败={len(fail_list)}")
conn.close()
logger.info(f" 日线行情导入完成: 成功 {success_count}/{total}")
time.sleep(sleep_interval)
if own_conn:
conn.close()
logger.info(
f" 日线行情按日期导入完成: 成功 {success_count}/{total} 个交易日"
)
if fail_list:
logger.warning(f" 失败列表({len(fail_list)}): {fail_list[:20]}...")
logger.warning(f" 失败日期({len(fail_list)}): {fail_list[:20]}...")
return fail_list
def import_daily_by_year(
stock_list: List[str],
start_year: int = 2010,
end_year: int = 2025,
sleep_interval: float = 0.3,
):
"""
按年份逐批导入日线行情 (断点续传友好)
适合大数据量导入,每年每只股票可单独重试
按年份逐批导入日线行情 (按交易日循环拉取全市场数据)
不再需要 stock_list 参数 — 每次 API 调用拉取当日全市场数据
"""
logger.info("=" * 60)
logger.info(
f"[3/7] 按年导入日线行情: {start_year} ~ {end_year}, "
f"{len(stock_list)} 只股票"
f"[3/7] 按年导入日线行情 (按交易日): {start_year} ~ {end_year}"
)
total_imported = 0
for year in range(start_year, end_year + 1):
year_start = f"{year}-01-01"
year_end = f"{year}-12-31"
logger.info(f"--- 导入 {year} 年日线行情 ---")
fail_list = import_daily_batch(
stock_list,
import_daily_by_date(
start_date=year_start,
end_date=year_end,
sleep_interval=0.2,
sleep_interval=sleep_interval,
)
total_imported += 1
logger.info(f" {year} 年完成\n")
logger.info(f" 所有年份日线行情导入完成!")
logger.info(" 所有年份日线行情导入完成!")
# ============================================================
@@ -956,9 +1013,8 @@ def full_import(
logger.error("无法获取股票列表,请先导入 stock_basic")
return
# Step 3: 日线行情 (按年导入)
# Step 3: 日线行情 (按年导入,按交易日循环拉取全市场数据)
import_daily_by_year(
stock_codes,
start_year=int(start_date[:4]),
end_year=int(end_date[:4]),
)
@@ -1071,58 +1127,136 @@ def check_table_summary(conn=None):
conn.close()
def resume_daily_by_year(stock_list, start_year=2010, end_year=2025):
def get_missing_daily_dates(
start_date: Optional[str] = None,
end_date: Optional[str] = None,
conn=None,
) -> List[str]:
"""
从中断点恢复按年导入日线行情
自动跳过数据库已有的年份,只导入缺失年份的数据
获取 daily 表中缺失的交易日列表
对比 trade_cal 中 is_open=1 的日期和 daily 表已有的 trade_date
返回未导入的交易日列表。
返回: 缺失交易日字符串列表 (YYYY-MM-DD 格式)
"""
conn = get_pg_connection()
if start_date is None:
start_date = START_DATE
if end_date is None:
end_date = END_DATE
own_conn = conn is None
if own_conn:
conn = get_pg_connection()
try:
cursor = conn.cursor()
cursor.execute("""
SELECT DISTINCT EXTRACT(YEAR FROM trade_date)::int AS year
FROM daily
ORDER BY year
""")
completed_years = set(row[0] for row in cursor.fetchall())
cursor.close()
finally:
conn.close()
logger.info(f"已完成年份: {sorted(completed_years)}")
logger.info(f"待导入年份: {[y for y in range(start_year, end_year+1) if y not in completed_years]}")
for year in range(start_year, end_year + 1):
if year in completed_years:
# 检查该年的股票覆盖是否完整
conn = get_pg_connection()
try:
cursor = conn.cursor()
cursor.execute(
sql.SQL("""
SELECT COUNT(DISTINCT ts_code)
FROM {}
WHERE EXTRACT(YEAR FROM trade_date) = %s
""").format(sql.Identifier("daily")),
(year,),
)
stock_count = cursor.fetchone()[0]
cursor.close()
finally:
conn.close()
logger.info(f" {year} 年: 已有 {stock_count} 只股票, 跳过")
continue
year_start = f"{year}-01-01"
year_end = f"{year}-12-31"
logger.info(f"--- 导入 {year} 年日线行情 ---")
import_daily_batch(
stock_list,
start_date=year_start,
end_date=year_end,
sleep_interval=0.2,
cursor.execute(
"""
SELECT tc.cal_date
FROM trade_cal tc
WHERE tc.is_open = 1
AND tc.cal_date >= %s
AND tc.cal_date <= %s
AND NOT EXISTS (
SELECT 1 FROM daily d
WHERE d.trade_date = tc.cal_date
)
ORDER BY tc.cal_date
""",
(start_date, end_date),
)
logger.info(f" {year} 年完成\n")
missing_dates = [row[0].strftime("%Y-%m-%d") if hasattr(row[0], "strftime") else str(row[0])[:10]
for row in cursor.fetchall()]
cursor.close()
return missing_dates
finally:
if own_conn:
conn.close()
def resume_daily_by_date(
start_date: Optional[str] = None,
end_date: Optional[str] = None,
sleep_interval: float = 0.3,
):
"""
按缺失日期断点续传日线行情
自动查询 daily 表已有的 trade_date 与 trade_cal 对比,
只导入缺失日期的全市场数据。
用法:
resume_daily_by_date(start_date="2010-01-01", end_date="2025-12-31")
"""
if start_date is None:
start_date = START_DATE
if end_date is None:
end_date = END_DATE
logger.info("=" * 60)
logger.info(f"[断点续传] 检测缺失日期: {start_date} ~ {end_date}")
missing_dates = get_missing_daily_dates(start_date, end_date)
if not missing_dates:
logger.info(" 所有交易日数据已完整,无需续传!")
return
total = len(missing_dates)
logger.info(f" 发现 {total} 个缺失交易日待导入")
if total <= 20:
logger.info(f" 缺失日期: {missing_dates}")
else:
logger.info(f" 缺失日期 (前20): {missing_dates[:20]}")
# 按年份分组统计
years_map: Dict[int, List[str]] = {}
for d in missing_dates:
y = int(d[:4])
years_map.setdefault(y, []).append(d)
for y in sorted(years_map):
logger.info(f" {y} 年: {len(years_map[y])} 个缺失交易日")
# 逐日期导入
conn = get_pg_connection()
pro = get_ts_pro()
success_count = 0
fail_list = []
for i, td_str in enumerate(missing_dates, 1):
try:
td_compact = td_str.replace("-", "")
def fetch():
return pro.daily(trade_date=td_compact)
df = fetch_with_retry(fetch, max_retries=3)
if df is None or df.empty:
logger.warning(f" [{td_str}] 返回空数据,跳过")
continue
df = _normalize_daily_df(df)
conflict_cols = ["ts_code", "trade_date"]
batch_insert("daily", df, conn, conflict_cols)
success_count += 1
except Exception as e:
logger.error(f" [{td_str}] 导入失败: {e}")
fail_list.append(td_str)
try:
conn.rollback()
except Exception:
pass
if i % 50 == 0 or i == total:
logger.info(f" 续传进度: {i}/{total} 成功={success_count} 失败={len(fail_list)}")
time.sleep(sleep_interval)
conn.close()
logger.info(f" 断点续传完成: 成功 {success_count}/{total}")
if fail_list:
logger.warning(f" 失败日期({len(fail_list)}): {fail_list[:20]}...")
return fail_list
if __name__ == "__main__":