This commit is contained in:
曾志威
2026-05-11 21:03:32 +08:00
parent 4fcbab887c
commit 947ed6dc78
4 changed files with 226 additions and 166 deletions
+3 -1
View File
@@ -73,7 +73,9 @@ python -m src.main --trading-day --start-date 19901219 --end-date 20261231
python -m src.main --daily
# 指定日期范围抓取日线
python -m src.main --daily --start-date 19901201 --end-date 20260508
python -m src.main --daily --start-date 19901201 --end-date 20260511
python -m src.main --daily --start-date 19920101 --end-date 20260511
# 抓取财务指标(全部股票,最近8个季度)
python -m src.main --financial
+2 -1
View File
@@ -5,6 +5,7 @@ BaoStock 的 query_xxx() 非线程安全,所有查询需通过同一把锁串
import threading
from contextlib import contextmanager
from datetime import datetime
import baostock as bs
_lock = threading.Lock()
@@ -43,7 +44,7 @@ def bs_query(query_fn, *args, **kwargs):
sig_parts = [repr(a) for a in args]
sig_parts += [f"{k}={v!r}" for k, v in kwargs.items()]
sig = ", ".join(sig_parts)
print(f" [BS] {short_name}({sig})", flush=True)
print(f" [{datetime.now().strftime('%H:%M:%S')}] [BS] {short_name}({sig})", flush=True)
with _lock:
rs = query_fn(*args, **kwargs)
yield rs
+148 -140
View File
@@ -1,10 +1,6 @@
"""日线行情抓取模块 — 使用 BaoStock
单数据源架构,代码大幅简化
跳过策略:
- 数据完整 = 行情记录数 + 已标记停牌天数 >= 交易日总数
- 未上市股票(ipo_date > 查询结束日期)
- 增量抓取:只请求缺失的日期段,不重复抓已有数据
增量抓取:一次本地查询确定每只股票的缺口范围,只请求缺失日期段
"""
import time
@@ -24,132 +20,122 @@ def _clean(val):
return val
def _get_not_listed(end_date: str) -> set[str]:
session = get_session()
try:
result = session.execute(
select(StockInfo.code).where(StockInfo.ipo_date > end_date)
)
return {row[0] for row in result}
finally:
session.close()
def _get_complete_codes(start_date: str, end_date: str, trading_days: list[str]) -> set[str]:
"""数据完整的判断:行情记录数 + 已标记停牌天数 >= 交易日总数"""
def _analyze_gaps(codes: list[str], start_date: str, end_date: str,
trading_days: list[str]) -> dict[str, list[str]]:
"""按月分段分析缺口。用 COUNT 对比交易日数,不逐条加载。"""
if not trading_days:
return set()
td_count = len(trading_days)
session = get_session()
try:
rec_rows = session.execute(
select(StockDaily.code, func.count(StockDaily.id))
.where(StockDaily.date >= start_date)
.where(StockDaily.date <= end_date)
.group_by(StockDaily.code)
)
rec_counts = {row[0]: row[1] for row in rec_rows}
return {}
susp_rows = session.execute(
select(StockNoData.code, func.count(StockNoData.id))
.where(StockNoData.date >= start_date)
.where(StockNoData.date <= end_date)
.group_by(StockNoData.code)
)
susp_counts = {row[0]: row[1] for row in susp_rows}
complete = set()
for code in set(rec_counts) | set(susp_counts):
if rec_counts.get(code, 0) + susp_counts.get(code, 0) >= td_count:
complete.add(code)
return complete
finally:
session.close()
def _get_gaps(codes: list[str], start_date: str, end_date: str,
trading_days: list[str]) -> dict[str, list[str]]:
"""查询每只股票缺失的日期段,返回 {code: [gap_start, gap_end]}"""
sd = f"{start_date[:4]}-{start_date[4:6]}-{start_date[6:8]}"
ed = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}"
td_set = set(trading_days)
# 上市日期
session = get_session()
try:
# 已有行情日期
result = session.execute(
select(StockDaily.code, StockDaily.date)
.where(StockDaily.code.in_(codes))
.where(StockDaily.date >= sd)
.where(StockDaily.date <= ed)
ipo_result = session.execute(
select(StockInfo.code, StockInfo.ipo_date)
.where(StockInfo.code.in_(codes))
)
code_dates: dict[str, set[str]] = {}
for code, date in result:
code_dates.setdefault(code, set()).add(str(date))
# 已标记停牌日期
nodata_result = session.execute(
select(StockNoData.code, StockNoData.date)
.where(StockNoData.code.in_(codes))
.where(StockNoData.date >= sd)
.where(StockNoData.date <= ed)
)
for code, date in nodata_result:
code_dates.setdefault(code, set()).add(str(date))
ipo_dates: dict[str, str] = {}
for code, ipo in ipo_result:
if ipo:
ipo_dates[code] = str(ipo)
finally:
session.close()
# 按月分段
td_by_month: dict[str, list[str]] = {}
for d in trading_days:
key = d[:7] # "2026-05"
td_by_month.setdefault(key, []).append(d)
gap_codes: set[str] = set()
total_months = len(td_by_month)
t0 = time.time()
for idx, (month_key, month_days) in enumerate(sorted(td_by_month.items())):
m_start = month_days[0]
m_end = month_days[-1]
session = get_session()
try:
cnt_result = session.execute(
select(StockDaily.code, func.count(StockDaily.id))
.where(StockDaily.code.in_(codes))
.where(StockDaily.date >= m_start)
.where(StockDaily.date <= m_end)
.group_by(StockDaily.code)
)
code_cnt = {row[0]: row[1] for row in cnt_result}
finally:
session.close()
new_gaps = 0
for code in codes:
if code in gap_codes:
continue
ipo = ipo_dates.get(code)
if ipo and ipo > m_end:
continue
expected = [d for d in month_days if not ipo or d >= ipo]
if not expected:
continue
if code_cnt.get(code, 0) < len(expected):
gap_codes.add(code)
new_gaps += 1
print(f" [{idx+1}/{total_months}] {month_key} 交易日:{len(month_days)} 新增缺口:{new_gaps} "
f"累计:{len(gap_codes)} 已用时:{time.time()-t0:.1f}s", flush=True)
if not gap_codes:
return {}
# 确定缺口范围
gaps: dict[str, list[str]] = {}
for code in codes:
existing = code_dates.get(code, set())
missing = [d for d in trading_days if d not in existing]
if not missing:
continue
gaps[code] = [missing[0], missing[-1]]
for code in gap_codes:
ipo = ipo_dates.get(code)
expected = [d for d in trading_days if not ipo or d >= ipo]
session = get_session()
try:
minmax = session.execute(
select(func.min(StockDaily.date), func.max(StockDaily.date))
.where(StockDaily.code == code)
.where(StockDaily.date >= sd)
.where(StockDaily.date <= ed)
).fetchone()
finally:
session.close()
if minmax and minmax[0]:
min_d, max_d = str(minmax[0]), str(minmax[1])
front = [d for d in expected if d < min_d]
back = [d for d in expected if d > max_d]
if front and back:
gaps[code] = [expected[0], expected[-1]]
elif front:
gaps[code] = [front[0], front[-1]]
elif back:
gaps[code] = [back[0], back[-1]]
else:
gaps[code] = [expected[0], expected[-1]]
else:
gaps[code] = [expected[0], expected[-1]]
return gaps
def _mark_suspensions_local(codes: list[str], start_date: str, end_date: str,
trading_days: list[str]):
"""纯本地数据库批量标记停牌天:比较 stock_daily 与 trading_day"""
if not codes or not trading_days:
def _mark_suspensions(nodata_map: dict[str, list[str]]):
if not nodata_map:
return
sd = f"{start_date[:4]}-{start_date[4:6]}-{start_date[6:8]}"
ed = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}"
session = get_session()
try:
result = session.execute(
select(StockDaily.code, StockDaily.date)
.where(StockDaily.code.in_(codes))
.where(StockDaily.date >= sd)
.where(StockDaily.date <= ed)
)
code_dates: dict[str, set[str]] = {}
for code, date in result:
code_dates.setdefault(code, set()).add(str(date))
nodata_result = session.execute(
select(StockNoData.code, StockNoData.date)
.where(StockNoData.code.in_(codes))
.where(StockNoData.date >= sd)
.where(StockNoData.date <= ed)
)
nodata_existing = {(code, str(date)) for code, date in nodata_result}
rows = []
for code in codes:
existing = code_dates.get(code, set())
for td in trading_days:
if td not in existing and (code, td) not in nodata_existing:
rows.append({"code": code, "date": td})
if rows:
batch_upsert(StockNoData, rows, ["code", "date"])
print(f" 本地批量标记停牌: {len(rows)}", flush=True)
finally:
session.close()
rows = []
for code, dates in nodata_map.items():
rows.extend({"code": code, "date": d} for d in dates)
if rows:
batch_upsert(StockNoData, rows, ["code", "date"])
print(f" 标记停牌: {len(rows)} 条({len(nodata_map)} 只股票)", flush=True)
def _fetch_baostock(code: str, start_date: str, end_date: str) -> list[dict] | None:
"""从 BaoStock 获取日线行情,含振幅/涨跌幅/换手率"""
bs_code = code_to_bs(code)
if not bs_code:
return None
@@ -208,42 +194,60 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None):
bs_login()
from src.fetchers.trading_day import get_trading_days
trading_days = get_trading_days(start_date, end_date)
# 直接查本地 trading_day 表,不调 BaoStock
sd_fmt = f"{start_date[:4]}-{start_date[4:6]}-{start_date[6:8]}"
ed_fmt = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}"
from src.db import TradingDay
session = get_session()
try:
td_result = session.execute(
select(TradingDay.date)
.where(TradingDay.date >= sd_fmt)
.where(TradingDay.date <= ed_fmt)
.order_by(TradingDay.date)
)
trading_days = [str(row[0]) for row in td_result]
finally:
session.close()
td_count = len(trading_days)
print(f" 交易日历: {start_date} ~ {end_date}{td_count} 个交易日", flush=True)
complete = _get_complete_codes(start_date, end_date, trading_days)
not_listed = _get_not_listed(end_date)
skip_set = complete | not_listed
if complete:
print(f" {len(complete)} 只股票数据已完整(含停牌天),跳过...", flush=True)
if not_listed:
print(f" {len(not_listed)} 只股票未上市,跳过...", flush=True)
codes = [c for c in codes if c not in skip_set]
# 排除未上市股票
session = get_session()
try:
not_listed = {row[0] for row in session.execute(
select(StockInfo.code).where(StockInfo.ipo_date > ed_fmt)
)}
finally:
session.close()
codes = [c for c in codes if c not in not_listed]
total = len(codes)
if not_listed:
print(f" {len(not_listed)} 只股票未上市,跳过", flush=True)
# 一次查询:判断完整 + 计算缺口
print(f" 正在分析 {len(codes)} 只股票的数据缺口...", flush=True)
gaps = _analyze_gaps(codes, start_date, end_date, trading_days)
complete_count = len(codes) - len(gaps) - len(not_listed)
if complete_count > 0:
print(f" {complete_count} 只股票数据已完整,跳过", flush=True)
total = len(gaps)
if total == 0:
print(" 所有股票数据已完整,无需抓取", flush=True)
return
# 查询每只股票实际缺失的日期范围
print(f" 正在分析 {total} 只股票的数据缺口...", flush=True)
gaps = _get_gaps(codes, start_date, end_date, trading_days)
print(f"正在抓取日线行情 {start_date} ~ {end_date}{len(gaps)} 只需更新...", flush=True)
print(f"正在抓取日线行情 {start_date} ~ {end_date}{total} 只需更新...", flush=True)
success = 0
fail = 0
nodata_count = 0
fetched_codes = []
nodata_map: dict[str, list[str]] = {}
t_start = time.time()
for i, code in enumerate(codes):
g = gaps.get(code)
if not g:
continue
# 增量:只请求缺失的日期段
for i, code in enumerate(gaps):
g = gaps[code]
gap_start = g[0].replace("-", "")
gap_end = g[-1].replace("-", "")
@@ -259,9 +263,16 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None):
print(f" {code} 写入失败: {e}", flush=True)
fail += 1
continue
returned_dates = {row["date"] for row in rows}
missing = [d for d in trading_days if d not in returned_dates
and g[0] <= d <= g[-1]]
if missing:
nodata_map[code] = missing
else:
missing = [d for d in trading_days if g[0] <= d <= g[-1]]
if missing:
nodata_map[code] = missing
nodata_count += 1
fetched_codes.append(code)
t_write = time.time() - t1
elapsed = time.time() - t_start
@@ -272,10 +283,7 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None):
f"成功:{success} 剩余:{eta:.0f}s", flush=True)
time.sleep(delay)
# 批量本地标记停牌
if fetched_codes:
print(f" 正在本地批量标记停牌天...", flush=True)
_mark_suspensions_local(fetched_codes, start_date, end_date, trading_days)
_mark_suspensions(nodata_map)
total_time = time.time() - t_start
print(f" 日线行情抓取完成,成功:{success} 失败:{fail} 停牌:{nodata_count} 总耗时:{total_time:.1f}s", flush=True)
+73 -24
View File
@@ -1,21 +1,21 @@
"""交易日历模块 — 使用 BaoStock
"""交易日历模块
数据源优先级:
1. 本地 trading_day 表(最快,之前已缓存
2. BaoStock 交易日历(需过滤掉周末
3. 从 stock_daily 表已有数据推断(最后手段)
1. 本地 trading_day 表(最快)
2. BaoStock 交易日历(需校验,近期可能含节假日
3. 从 stock_daily 表已有数据推断
"""
import time
from datetime import datetime
from datetime import datetime, timedelta
import baostock as bs
from src.db import TradingDay, StockDaily, batch_upsert, get_session
from src.baostock_conn import bs_query
from sqlalchemy import select, func
from sqlalchemy import select, func, text
def _fetch_baostock(sd: str, ed: str) -> list[str] | None:
"""从 BaoStock 获取交易日历,过滤周末"""
"""从 BaoStock 获取交易日历"""
try:
with bs_query(bs.query_trade_dates, start_date=sd, end_date=ed) as rs:
days = []
@@ -28,32 +28,68 @@ def _fetch_baostock(sd: str, ed: str) -> list[str] | None:
return None
def _validate_with_daily(dates: list[str]) -> list[str]:
"""用 stock_daily 校验:只保留有实际行情数据的日期(排除节假日)
仅保留当天(可能还没抓取),其余必须有行情数据才算交易日。
"""
if not dates:
return dates
today = datetime.now().strftime("%Y-%m-%d")
session = get_session()
try:
result = session.execute(
select(func.distinct(StockDaily.date))
.where(StockDaily.date >= dates[0])
.where(StockDaily.date <= dates[-1])
)
real_dates = {str(row[0]) for row in result}
finally:
session.close()
validated = [d for d in dates if d in real_dates or d == today]
return validated
def _fetch_and_save(start_date: str, end_date: str) -> list[str]:
"""从 BaoStock 获取交易日并保存到本地表"""
sd = f"{start_date[:4]}-{start_date[4:6]}-{start_date[6:8]}"
ed = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}"
sd = datetime.strptime(start_date, "%Y%m%d")
ed = datetime.strptime(end_date, "%Y%m%d")
today = datetime.now().strftime("%Y-%m-%d")
print(f" 正在从 BaoStock 获取交易日历 {sd} ~ {ed}...", flush=True)
days = _fetch_baostock(sd, ed)
if days:
rows = [{"date": d} for d in days]
all_days: list[str] = []
chunk_start = sd
while chunk_start <= ed:
chunk_end = min(chunk_start.replace(year=chunk_start.year + 5), ed)
cs = chunk_start.strftime("%Y-%m-%d")
ce = chunk_end.strftime("%Y-%m-%d")
print(f" 正在从 BaoStock 获取交易日历 {cs} ~ {ce}...", flush=True)
days = _fetch_baostock(cs, ce)
if days:
all_days.extend(d for d in days if d <= today)
else:
print(f" {cs} ~ {ce} 获取失败", flush=True)
chunk_start = chunk_end + timedelta(days=1)
if all_days:
# 校验:排除节假日
all_days = _validate_with_daily(all_days)
rows = [{"date": d} for d in all_days]
batch_upsert(TradingDay, rows, ["date"])
print(f" 交易日历已保存,{len(days)} 个交易日", flush=True)
return days
print(f" 交易日历已保存,{len(all_days)} 个交易日", flush=True)
return all_days
print(" BaoStock 获取失败,将从已有行情数据推断", flush=True)
return _infer_from_daily(start_date, end_date)
def _format_date(d: str) -> str:
"""YYYYMMDD → YYYY-MM-DD"""
return f"{d[:4]}-{d[4:6]}-{d[6:8]}"
def get_trading_days(start_date: str, end_date: str) -> list[str]:
"""获取指定范围内的交易日列表
优先查本地表,若本地数据未覆盖完整范围则从 BaoStock 补全。
优先查本地表,若本地数据未覆盖完整范围则补全。
"""
sd = _format_date(start_date)
ed = _format_date(end_date)
@@ -73,15 +109,10 @@ def get_trading_days(start_date: str, end_date: str) -> list[str]:
finally:
session.close()
# 本地数据不完整,从 BaoStock 获取
fetched = _fetch_and_save(start_date, end_date)
# 合并本地 + 新获取的数据去重
if not fetched:
return days if days else []
all_days = sorted(set(days + fetched))
return all_days
return sorted(set(days + fetched))
def _infer_from_daily(start_date: str, end_date: str) -> list[str]:
@@ -105,13 +136,31 @@ def fetch_trading_days(start_date: str | None = None, end_date: str | None = Non
"""独立抓取交易日历并保存
用法:python -m src.main --trading-day --start-date 19901219 --end-date 20261231
默认从 1990-12-19(沪市开市日)到今天。
"""
if end_date is None:
end_date = time.strftime("%Y%m%d")
if start_date is None:
start_date = "19901219"
sd = _format_date(start_date)
ed = _format_date(end_date)
session = get_session()
try:
result = session.execute(
select(TradingDay.date)
.where(TradingDay.date >= sd)
.where(TradingDay.date <= ed)
.order_by(TradingDay.date)
)
days = [str(row[0]) for row in result]
finally:
session.close()
if days and days[0] <= sd and days[-1] >= ed:
print(f" 交易日历 {sd} ~ {ed} 已有 {len(days)} 天,跳过", flush=True)
return
print(f"正在抓取交易日历 {start_date} ~ {end_date}...", flush=True)
days = _fetch_and_save(start_date, end_date)
if not days: