This commit is contained in:
曾志威
2026-05-11 10:20:09 +08:00
parent 55de951959
commit 4fcbab887c
9 changed files with 365 additions and 84 deletions
+5
View File
@@ -39,6 +39,11 @@ def bs_query(query_fn, *args, **kwargs):
row = rs.get_row_data()
"""
bs_login()
short_name = query_fn.__name__.replace("query_", "")
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)
with _lock:
rs = query_fn(*args, **kwargs)
yield rs
+35 -1
View File
@@ -14,7 +14,7 @@
from sqlalchemy import (
Column, String, Date, DateTime, Float, Integer, Text,
UniqueConstraint, Index, create_engine, MetaData, func, text,
UniqueConstraint, Index, create_engine, MetaData, func, select, text,
)
from sqlalchemy.orm import DeclarativeBase, Session, sessionmaker
from sqlalchemy.dialects.mysql import insert as mysql_insert
@@ -97,6 +97,26 @@ class FinancialCashflow(Base):
data = Column(Text, comment="JSON格式现金流数据")
# ── 指数日线行情 ──
class IndexDaily(Base):
__tablename__ = "index_daily"
__table_args__ = (
UniqueConstraint("code", "date", name="uq_index_code_date"),
Index("ix_index_date", "date"),
)
id = Column(Integer, primary_key=True, autoincrement=True)
code = Column(String(10), nullable=False, comment="指数代码")
date = Column(Date, nullable=False, comment="交易日期")
open = Column(Float, comment="开盘价")
high = Column(Float, comment="最高价")
low = Column(Float, comment="最低价")
close = Column(Float, comment="收盘价")
volume = Column(Float, comment="成交量")
amount = Column(Float, comment="成交额")
pct_change = Column(Float, comment="涨跌幅%")
# ── 分红送转 ──
class StockDividend(Base):
__tablename__ = "stock_dividend"
@@ -245,6 +265,20 @@ class StockSector(Base):
# ── 数据库连接管理 ──
_engine = None
_SessionFactory = None
_stock_codes_cache: list[str] | None = None
def get_stock_codes() -> list[str]:
"""获取全部股票代码(进程内缓存,避免重复查询)"""
global _stock_codes_cache
if _stock_codes_cache is None:
session = get_session()
try:
result = session.execute(select(StockInfo.code))
_stock_codes_cache = [row[0] for row in result]
finally:
session.close()
return _stock_codes_cache
def get_engine():
+113 -37
View File
@@ -1,25 +1,22 @@
"""日线行情抓取模块 — 使用 BaoStock
单数据源架构,代码大幅简化。
跳过策略(停牌天按天记录)
跳过策略:
- 数据完整 = 行情记录数 + 已标记停牌天数 >= 交易日总数
- 未上市股票(ipo_date > 查询结束日期)
- 抓取成功后自动识别缺失交易日并标记为停牌
- 增量抓取:只请求缺失的日期段,不重复抓已有数据
"""
import json
import re
import time
from datetime import datetime, timedelta
import baostock as bs
from src.baostock_conn import bs_query, code_to_bs, bs_login
from src.config import get_fetch_config
from src.db import StockInfo, StockDaily, StockNoData, batch_upsert, get_session
from src.db import StockInfo, StockDaily, StockNoData, batch_upsert, get_session, get_stock_codes
from sqlalchemy import select, func
def _clean(val):
"""将空字符串、无效值转为 None"""
if val is None:
return None
if isinstance(val, str) and val.strip() == "":
@@ -27,17 +24,7 @@ def _clean(val):
return val
def _get_stock_codes() -> list[str]:
session = get_session()
try:
result = session.execute(select(StockInfo.code))
return [row[0] for row in result]
finally:
session.close()
def _get_not_listed(end_date: str) -> set[str]:
"""查询在end_date之后上市的股票(未上市,需跳过)"""
session = get_session()
try:
result = session.execute(
@@ -80,12 +67,85 @@ def _get_complete_codes(start_date: str, end_date: str, trading_days: list[str])
session.close()
def _record_nodata_days(code: str, days: list[str]):
"""记录个股的停牌/无数据日期(按天粒度)"""
if not days:
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)
)
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))
finally:
session.close()
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]]
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:
return
rows = [{"code": code, "date": d} for d in days]
batch_upsert(StockNoData, rows, ["code", "date"])
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()
def _fetch_baostock(code: str, start_date: str, end_date: str) -> list[dict] | None:
@@ -136,7 +196,7 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None):
cfg = get_fetch_config()
delay = cfg.get("delay", 0.1)
codes = _get_stock_codes()
codes = get_stock_codes()
if not codes:
print(" 无股票列表,请先运行 --stock-info", flush=True)
return
@@ -167,39 +227,55 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None):
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)
success = 0
fail = 0
nodata_count = 0
skipped = len(skip_set)
fetched_codes = []
t_start = time.time()
print(f"正在抓取日线行情 {start_date} ~ {end_date},需抓取 {total} 只(跳过 {skipped} 只)...", flush=True)
for i, code in enumerate(codes):
rows = _fetch_baostock(code, start_date, end_date)
g = gaps.get(code)
if not g:
continue
# 增量:只请求缺失的日期段
gap_start = g[0].replace("-", "")
gap_end = g[-1].replace("-", "")
t0 = time.time()
rows = _fetch_baostock(code, gap_start, gap_end)
t_fetch = time.time() - t0
t1 = time.time()
if rows is not None:
try:
batch_upsert(StockDaily, rows, ["code", "date"])
success += 1
returned_dates = {row["date"] for row in rows}
missing_days = [d for d in trading_days if d not in returned_dates]
if missing_days:
_record_nodata_days(code, missing_days)
except Exception as e:
print(f" {code} 写入失败: {e}", flush=True)
fail += 1
continue
else:
# 无数据:标记所有交易日为停牌
_record_nodata_days(code, trading_days)
nodata_count += 1
fetched_codes.append(code)
t_write = time.time() - t1
total_elapsed = time.time() - t_start
avg = total_elapsed / (i + 1)
elapsed = time.time() - t_start
avg = elapsed / (i + 1)
eta = avg * (total - i - 1)
print(f" [{i+1}/{total}] {code} 成功:{success} 失败:{fail} 停牌:{nodata_count} "
f"已用时:{total_elapsed:.0f}s 预计剩余:{eta:.0f}s", flush=True)
print(f" [{i+1}/{total}] {code} 缺口:{g[0]}~{g[-1]} "
f"网络:{t_fetch:.1f}s 写入:{t_write:.1f}s 行数:{len(rows) if rows else 0} "
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)
total_time = time.time() - t_start
print(f" 日线行情抓取完成,成功:{success} 失败:{fail} 停牌:{nodata_count} 总耗时:{total_time:.1f}s", flush=True)
+29 -7
View File
@@ -2,20 +2,30 @@
BaoStock query_dividend_data() 按年度查询分红记录。
迭代最近10年获取完整分红历史。
跳过策略:已有当前年份分红记录的股票不再重复抓取。
"""
from datetime import datetime
import baostock as bs
from src.baostock_conn import bs_query, code_to_bs
from src.db import StockInfo, StockDividend, batch_upsert, get_session
from sqlalchemy import select
from src.db import StockDividend, batch_upsert, get_session, get_stock_codes
from sqlalchemy import select, func
def _get_stock_codes() -> list[str]:
def _get_dividend_done(codes: list[str]) -> set[str]:
"""查询已有完整10年分红记录的股票代码"""
current_year = datetime.now().year
years = [str(y) for y in range(current_year - 10, current_year + 1)]
session = get_session()
try:
result = session.execute(select(StockInfo.code))
return [row[0] for row in result]
result = session.execute(
select(StockDividend.code, func.count(StockDividend.id))
.where(StockDividend.code.in_(codes))
.where(StockDividend.report_date.in_(years))
.group_by(StockDividend.code)
)
# 有记录即可跳过(分红非每年都有,只要有任意年份数据就说明已抓过)
return {row[0] for row in result if row[1] > 0}
finally:
session.close()
@@ -74,14 +84,26 @@ def fetch_dividend(symbol: str | None = None):
"""抓取分红送转数据
用法:python -m src.main --dividend [--symbol 000001]
已有分红记录的股票自动跳过。
"""
if symbol:
codes = [symbol]
else:
codes = _get_stock_codes()
codes = get_stock_codes()
# 跳过已有分红数据的股票
if not symbol:
done = _get_dividend_done(codes)
if done:
print(f" {len(done)} 只股票已有分红数据,跳过", flush=True)
codes = [c for c in codes if c not in done]
total = len(codes)
print(f"正在抓取分红送转数据,共 {total} 只股票...", flush=True)
if total == 0:
print(" 所有股票分红数据已完整,无需抓取", flush=True)
return
print(f"正在抓取分红送转数据,需抓取 {total} 只股票...", flush=True)
success = 0
for i, code in enumerate(codes):
+41 -15
View File
@@ -6,23 +6,15 @@ BaoStock 按季度查询财务数据:
- query_cash_flow_data() 现金流
数据以 JSON 格式存入 data 列(与现有表结构兼容)。
跳过策略:已有 (code, report_date) 记录的季度不再重复抓取。
"""
import json
from datetime import datetime
import baostock as bs
from src.baostock_conn import bs_query, code_to_bs
from src.db import StockInfo, FinancialIncome, FinancialBalance, FinancialCashflow, batch_upsert, get_session
from sqlalchemy import select
def _get_stock_codes() -> list[str]:
session = get_session()
try:
result = session.execute(select(StockInfo.code))
return [row[0] for row in result]
finally:
session.close()
from src.db import FinancialIncome, FinancialBalance, FinancialCashflow, batch_upsert, get_session, get_stock_codes
from sqlalchemy import select, func
def _recent_quarters(n: int) -> list[tuple[int, int]]:
@@ -39,6 +31,25 @@ def _recent_quarters(n: int) -> list[tuple[int, int]]:
return result
def _get_existing_financial(codes: list[str], quarters: list[tuple[int, int]],
model_cls) -> set[str]:
"""查询已有财务数据的 (code) 集合:三张表都有完整8季度数据的股票"""
q_labels = {f"{y}-{m:02d}-{d:02d}" for y, _q in quarters
for m, d in [(3, 31), (6, 30), (9, 30), (12, 31)]}
session = get_session()
try:
result = session.execute(
select(model_cls.code, func.count(model_cls.id))
.where(model_cls.code.in_(codes))
.where(model_cls.report_date.in_(q_labels))
.group_by(model_cls.code)
)
q_count = len(q_labels)
return {row[0] for row in result if row[1] >= q_count}
finally:
session.close()
def _parse_resultset(code: str, rs, fields: list[str], year: int, quarter: int) -> list[dict]:
"""将 BaoStock ResultData 转为 JSON 行"""
rows = []
@@ -65,16 +76,31 @@ def fetch_financial(symbol: str | None = None):
"""抓取财务数据
用法:python -m src.main --financial [--symbol 000001]
默认抓取所有股票最近 8 个季度。
默认抓取所有股票最近 8 个季度。已有完整数据的股票自动跳过。
"""
if symbol:
codes = [symbol]
else:
codes = _get_stock_codes()
codes = get_stock_codes()
quarters = _recent_quarters(8)
# 跳过三张表都有完整季度数据的股票
if not symbol:
inc_done = _get_existing_financial(codes, quarters, FinancialIncome)
bal_done = _get_existing_financial(codes, quarters, FinancialBalance)
cf_done = _get_existing_financial(codes, quarters, FinancialCashflow)
done = inc_done & bal_done & cf_done
if done:
print(f" {len(done)} 只股票财务数据已完整,跳过", flush=True)
codes = [c for c in codes if c not in done]
total = len(codes)
quarters = _recent_quarters(8)
print(f"正在抓取财务数据,共 {total} 只股票 × {len(quarters)} 个季度...", flush=True)
if total == 0:
print(" 所有股票财务数据已完整,无需抓取", flush=True)
return
print(f"正在抓取财务数据,需抓取 {total} 只股票 × {len(quarters)} 个季度...", flush=True)
success = 0
fail = 0
+129
View File
@@ -0,0 +1,129 @@
"""指数日线行情抓取 — 使用 BaoStock
主要指数:
sh.000001 上证指数 sh.000300 沪深300
sh.000905 中证500 sh.000852 中证1000
sh.000688 科创50 sz.399001 深证成指
sz.399006 创业板指 sz.399005 中小板指
用法:
python -m src.main --index
python -m src.main --index --start-date 20260101 --end-date 20260508
"""
import time
from datetime import datetime, timedelta
import baostock as bs
from src.baostock_conn import bs_query, bs_login
from src.config import get_fetch_config
from src.db import IndexDaily, batch_upsert, get_session
from sqlalchemy import select, func
# 主要指数代码 → BaoStock 格式
INDICES = {
"000001": ("sh", "上证指数"),
"000300": ("sh", "沪深300"),
"000905": ("sh", "中证500"),
"000852": ("sh", "中证1000"),
"000688": ("sh", "科创50"),
"399001": ("sz", "深证成指"),
"399006": ("sz", "创业板指"),
"399005": ("sz", "中小板指"),
}
def _clean(val):
if val is None:
return None
if isinstance(val, str) and val.strip() == "":
return None
return val
def _get_fetched_codes(sd: str, ed: str) -> set[str]:
"""查询数据已覆盖起始日期的指数代码"""
session = get_session()
try:
from sqlalchemy import distinct
sd_date = datetime.strptime(sd, "%Y-%m-%d").date()
check_start = str(sd_date - timedelta(days=5))
check_end = str(sd_date + timedelta(days=5))
result = session.execute(
select(distinct(IndexDaily.code))
.where(IndexDaily.date >= check_start)
.where(IndexDaily.date <= check_end)
)
return {row[0] for row in result}
finally:
session.close()
def _fetch_one_index(code: str, market: str, name: str,
sd: str, ed: str) -> list[dict] | None:
"""抓取单个指数的日线数据"""
bs_code = f"{market}.{code}"
try:
with bs_query(
bs.query_history_k_data_plus,
bs_code,
"date,open,high,low,close,volume,amount,pctChg",
start_date=sd, end_date=ed, frequency="d",
) as rs:
rows = []
while rs.next():
r = rs.get_row_data()
rows.append({
"code": code,
"date": r[0],
"open": _clean(r[1]),
"high": _clean(r[2]),
"low": _clean(r[3]),
"close": _clean(r[4]),
"volume": _clean(r[5]),
"amount": _clean(r[6]),
"pct_change": _clean(r[7]),
})
return rows if rows else None
except Exception:
return None
def fetch_index(start_date: str | None = None, end_date: str | None = None):
"""抓取主要指数日线行情"""
cfg = get_fetch_config()
delay = cfg.get("delay", 0.1)
if end_date is None:
end_date = datetime.now().strftime("%Y%m%d")
if start_date is None:
start_date = "19901219"
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]}"
# 跳过已有数据的指数
fetched = _get_fetched_codes(sd, ed)
to_fetch = {k: v for k, v in INDICES.items() if k not in fetched}
if not to_fetch:
print(f" 指数数据 {sd} ~ {ed} 已完整,跳过", flush=True)
return
skip_msg = f"(跳过 {len(fetched)} 个已有数据)" if fetched else ""
print(f"正在抓取指数日线 {sd} ~ {ed},共 {len(to_fetch)}{skip_msg}...", flush=True)
bs_login()
success = 0
for code, (market, name) in to_fetch.items():
t0 = time.time()
rows = _fetch_one_index(code, market, name, sd, ed)
if rows:
batch_upsert(IndexDaily, rows, ["code", "date"])
success += 1
print(f" {name}({code}): {len(rows)} 天, {time.time()-t0:.1f}s", flush=True)
else:
print(f" {name}({code}): 无数据", flush=True)
time.sleep(delay)
print(f" 指数数据抓取完成,成功:{success}/{len(to_fetch)}", flush=True)
+3 -12
View File
@@ -16,8 +16,8 @@ import baostock as bs
from src.baostock_conn import bs_query, code_to_bs, bs_login
from src.config import get_fetch_config
from src.db import (
StockInfo, StockMin5, StockMin15, StockMin30, StockMin60,
batch_upsert, get_session,
StockMin5, StockMin15, StockMin30, StockMin60,
batch_upsert, get_session, get_stock_codes,
)
from sqlalchemy import select, func, distinct
@@ -32,15 +32,6 @@ FREQ_MODEL = {
}
def _get_stock_codes() -> list[str]:
session = get_session()
try:
result = session.execute(select(StockInfo.code))
return [row[0] for row in result]
finally:
session.close()
def _clean(val):
if val is None:
return None
@@ -121,7 +112,7 @@ def _fetch_one_freq(freq: str, start_date: str | None, end_date: str | None,
if symbol:
codes = [symbol]
else:
codes = _get_stock_codes()
codes = get_stock_codes()
if not codes:
print(" 无股票列表,请先运行 --stock-info", flush=True)
return
+2 -11
View File
@@ -15,7 +15,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed
import requests
import baostock as bs
from src.baostock_conn import bs_query, bs_login
from src.db import StockInfo, StockSector, batch_upsert, get_session
from src.db import StockSector, batch_upsert, get_session, get_stock_codes
from sqlalchemy import select
@@ -26,15 +26,6 @@ _HEADERS = {
_F10_URL = "https://emweb.securities.eastmoney.com/PC_HSF10/CompanySurvey/CompanySurveyAjax"
def _get_stock_codes() -> list[str]:
session = get_session()
try:
result = session.execute(select(StockInfo.code))
return [row[0] for row in result]
finally:
session.close()
def _fetch_industry() -> dict[str, str]:
"""从 BaoStock 获取全部股票的行业分类"""
result = {}
@@ -88,7 +79,7 @@ def _fetch_region_batch(codes: list[str], workers: int = 10) -> dict[str, str]:
def fetch_sector(industry_only: bool = False, region_only: bool = False):
"""抓取行业分类 + 地域分类"""
codes = _get_stock_codes()
codes = get_stock_codes()
if not codes:
print(" 无股票列表,请先运行 --stock-info", flush=True)
return
+8 -1
View File
@@ -12,6 +12,7 @@
python -m src.main --intraday --symbol 000001 --freq 30
python -m src.main --sector # 行业+地域分类
python -m src.main --sector --industry-only # 仅行业分类
python -m src.main --index # 指数日线(上证/沪深300/创业板等)
"""
import argparse
@@ -32,6 +33,7 @@ def main():
choices=["5", "15", "30", "60", "all"],
help="分钟K线频率(默认5all=全部)")
parser.add_argument("--sector", action="store_true", help="抓取行业+地域分类")
parser.add_argument("--index", action="store_true", help="抓取指数日线行情")
parser.add_argument("--industry-only", action="store_true", help="仅抓取行业分类")
parser.add_argument("--region-only", action="store_true", help="仅抓取地域分类")
parser.add_argument("--start-date", type=str, help="开始日期 YYYYMMDD")
@@ -42,7 +44,8 @@ def main():
if not any([args.stock_info, args.trading_day, args.daily,
args.financial, args.dividend, args.intraday,
args.sector, args.industry_only, args.region_only]):
args.sector, args.industry_only, args.region_only,
args.index]):
parser.print_help()
return
@@ -80,6 +83,10 @@ def main():
from src.fetchers.sector import fetch_sector
fetch_sector(industry_only=args.industry_only, region_only=args.region_only)
if args.index:
from src.fetchers.index import fetch_index
fetch_index(start_date=args.start_date, end_date=args.end_date)
print("全部任务完成", flush=True)
finally:
bs_logout()