diff --git a/README.md b/README.md index 2d56510..ec0df40 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/src/baostock_conn.py b/src/baostock_conn.py index a87da20..e9fce53 100644 --- a/src/baostock_conn.py +++ b/src/baostock_conn.py @@ -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 diff --git a/src/fetchers/daily.py b/src/fetchers/daily.py index 162c817..7487016 100644 --- a/src/fetchers/daily.py +++ b/src/fetchers/daily.py @@ -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) diff --git a/src/fetchers/trading_day.py b/src/fetchers/trading_day.py index e137dde..06e6287 100644 --- a/src/fetchers/trading_day.py +++ b/src/fetchers/trading_day.py @@ -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: