"""日线行情多源轮换 + 多 worker 并发压测 用法: python -m benchmarks.bench_daily --codes 200 --days 30 --workers 1,2,4,8 行为: - 从 stock_info 随机/顺序取若干只股票(必须先 --stock-info); - 估算它们最近 N 天的真实缺口; - 顺次以指定的 workers 配置跑一遍真实抓取,记录耗时、成功率、无数据数。 - 不会清空已有数据,跑完即可被认为是常规增量。 说明: - 此脚本会真实命中外部数据源(baostock / 新浪 / 腾讯 / 东方财富), 请在配置允许的窗口内运行,避免被限流。 - 脚本不修改 config.yaml;要切换 workers 是通过传参,并临时覆盖 fetch.workers。 """ from __future__ import annotations import argparse import random import time from datetime import datetime, timedelta from src.config import load_config, get_fetch_config from src.db import init_db, get_stock_codes from src.fetchers.daily import fetch_daily from src.log import get_logger _logger = get_logger("bench") def _parse_workers(arg: str) -> list[int]: return [max(1, int(x.strip())) for x in arg.split(",") if x.strip()] def _pick_codes(n: int, seed: int | None) -> list[str]: all_codes = get_stock_codes() if not all_codes: raise SystemExit("stock_info 为空,请先运行:python -m src.main --stock-info") if n >= len(all_codes): return all_codes rng = random.Random(seed) return rng.sample(all_codes, n) def main(): parser = argparse.ArgumentParser(description="日线抓取并发度压测") parser.add_argument("--codes", type=int, default=200, help="参与压测的股票数(默认 200)") parser.add_argument("--days", type=int, default=30, help="时间窗口天数(默认近 30 天)") parser.add_argument( "--workers", type=str, default="1,2,4,8", help="并发档位列表,逗号分隔(默认 1,2,4,8)", ) parser.add_argument( "--source", type=str, default="all", choices=["baostock", "sina", "tencent", "eastmoney", "all"], help="数据源(默认 all 多源轮换)", ) parser.add_argument("--seed", type=int, default=42, help="随机种子,固定后多次跑可重复") args = parser.parse_args() load_config() init_db() workers_list = _parse_workers(args.workers) end_date = datetime.now().strftime("%Y%m%d") start_date = (datetime.now() - timedelta(days=args.days)).strftime("%Y%m%d") _logger.info( "压测计划: codes=%d days=%d source=%s workers=%s seed=%s", args.codes, args.days, args.source, workers_list, args.seed, ) _logger.info("窗口: %s ~ %s", start_date, end_date) sample_codes = _pick_codes(args.codes, args.seed) _logger.info("已选股票样本(前10): %s ...", sample_codes[:10]) cfg = get_fetch_config() results = [] for w in workers_list: _logger.info("==== workers=%d 开始 ====", w) cfg["workers"] = w from src.fetchers import daily as daily_mod original_get_stock_codes = daily_mod.get_stock_codes daily_mod.get_stock_codes = lambda: sample_codes try: t0 = time.time() fetch_daily(start_date=start_date, end_date=end_date, source=args.source) elapsed = time.time() - t0 finally: daily_mod.get_stock_codes = original_get_stock_codes _logger.info("==== workers=%d 完成: %.1fs ====", w, elapsed) results.append((w, elapsed)) _logger.info("==== 汇总 ====") base = results[0][1] if results else 0.0 for w, sec in results: speedup = base / sec if sec > 0 else 0.0 _logger.info(" workers=%d %.1fs speedup=%.2fx", w, sec, speedup) if __name__ == "__main__": main()