From 13883f64473ed20519e3e64dc88b8fd7d2b54b0c1788d24b9b911a96570fd506 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=9B=BE=E5=BF=97=E5=A8=81?= Date: Sun, 17 May 2026 15:51:10 +0800 Subject: [PATCH] =?UTF-8?q?=E6=9B=B4=E6=96=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- CLAUDE.md | 12 ++ README.md | 60 +++++----- TODO.md | 125 +++++++++++++++++---- benchmarks/README.md | 47 ++++++++ benchmarks/__init__.py | 1 + benchmarks/bench_daily.py | 108 ++++++++++++++++++ logs/ashare.log.2026-05-15 | 14 +++ src/baostock_conn.py | 15 +-- src/db.py | 9 +- src/fetchers/daily.py | 63 +++++++---- src/fetchers/dividend.py | 62 +++++------ src/fetchers/financial.py | 62 +++++------ src/fetchers/index.py | 91 +++++++++++----- src/fetchers/intraday.py | 155 ++++++++++++++++---------- src/fetchers/market_daily.py | 9 +- src/fetchers/sector.py | 181 +++++++------------------------ src/fetchers/stock_list.py | 19 +++- src/fetchers/trading_day.py | 17 +-- src/main.py | 37 ++++--- tests/README.md | 15 ++- tests/test_daily_source_codes.py | 52 +++++++++ tests/test_daily_sources_http.py | 141 ++++++++++++++++++++++++ tests/test_log.py | 60 ++++++++++ tests/test_sector_merge.py | 28 +++++ 24 files changed, 973 insertions(+), 410 deletions(-) create mode 100644 CLAUDE.md create mode 100644 benchmarks/README.md create mode 100644 benchmarks/__init__.py create mode 100644 benchmarks/bench_daily.py create mode 100644 logs/ashare.log.2026-05-15 create mode 100644 tests/test_daily_source_codes.py create mode 100644 tests/test_daily_sources_http.py create mode 100644 tests/test_log.py create mode 100644 tests/test_sector_merge.py diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 0000000..e2beacd --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1,12 @@ +# ashare-data 项目规则 + +## 数据安全 + +- **禁止删除数据库数据**:不允许执行任何 `DELETE`、`TRUNCATE`、`DROP TABLE`、`session.delete()` 等删除操作,除非用户明确要求 +- **禁止清理数据**:不要主动建议或执行清理数据库表的操作 +- `db.py` 中的 `init_db()` 里的 `DROP TABLE` 是表结构自动迁移逻辑,属于例外,不要修改 + +## 代码规范 + +- 所有数据库写入使用 `batch_upsert()`(INSERT ON DUPLICATE KEY UPDATE),确保幂等 +- 不要引入新的删除方法或清理脚本 diff --git a/README.md b/README.md index 51853d8..1221348 100644 --- a/README.md +++ b/README.md @@ -14,8 +14,7 @@ A股数据抓取工具,以 [BaoStock](http://baostock.com) 为主、新浪/腾 | 季频财务指标 | 盈利能力、偿债能力、现金流(最近8个季度,JSON 存储) | BaoStock | | 分红送转 | 每10股送转、派息、除权除息日(最近10年) | BaoStock | | 分钟K线 | 5/15/30/60 分钟K线(开高低收、成交量/额) | BaoStock | -| 行业+地域分类 | 证监会行业分类 + 省份 | BaoStock + 东方财富 | -| 概念板块 | 东方财富全量概念板块及其成分股 | 东方财富 | +| 概念板块 | 东方财富全量概念板块及其成分股(已有≥300个概念时自动跳过) | 东方财富 | **已知限制**:BaoStock 不含北交所(920xxx)股票;新浪/腾讯/东方财富数据源对北交所同样不支持。 @@ -110,15 +109,8 @@ python -m src.main --intraday --start-date 20260101 # 抓取全部频率分钟K线(5/15/30/60) python -m src.main --intraday --freq all --start-date 20260101 -# 抓取行业+地域分类 +# 抓取概念板块及成分股(东方财富) python -m src.main --sector - -# 仅抓取行业分类 / 仅抓取地域分类 -python -m src.main --sector --industry-only -python -m src.main --sector --region-only - -# 仅抓取概念板块及成分股(东方财富) -python -m src.main --sector --concept-only ``` > ⚠️ **涨跌停统计 (`market_daily`)** 已实现于 `src/fetchers/market_daily.py`,但当前 `main.py` 尚未挂载 `--market-daily` 参数,暂只能在 Python 内直接调用 `fetch_market_daily(...)`。详见 [TODO.md](./TODO.md)。 @@ -132,14 +124,12 @@ python -m src.main --sector --concept-only --daily 抓取日线行情(增量;自动分析数据缺口,已完整自动跳过) --source 日线数据源: baostock/sina/tencent/eastmoney/all(默认 all,轮换+失败自动切换) --index 抓取主要指数日线 - --financial 抓取季频财务指标 - --dividend 抓取分红送转数据 - --intraday 抓取分钟K线行情 + --financial 抓取季频财务指标(增量;只抓缺失季度) + --dividend 抓取分红送转数据(增量;只抓缺失年份) + --intraday 抓取分钟K线行情(增量;只抓缺失日期) --freq K线频率: 5/15/30/60/all(默认 5) - --sector 抓取行业+地域分类(默认两者都抓) - --industry-only 仅抓取行业分类 - --region-only 仅抓取地域分类 - --concept-only 仅抓取概念板块及成分股 + --sector 抓取概念板块及成分股(增量;已有≥300个概念时跳过) + --market-daily 汇总每日涨跌停统计(从 stock_daily 聚合) 日期过滤(对日线行情、交易日历、分钟K线、指数生效): --start-date 开始日期,格式 YYYYMMDD @@ -256,14 +246,6 @@ python -m src.main --sector --concept-only 联合主键:`(code, datetime)` -### stock_sector — 行业+地域分类 - -| 字段 | 类型 | 说明 | -|------|------|------| -| code | VARCHAR(10) PK | 股票代码 | -| industry | VARCHAR(100) | 证监会行业分类 | -| region | VARCHAR(20) | 省份/地域 | - ### stock_concept — 概念板块及成分股 | 字段 | 类型 | 说明 | @@ -315,6 +297,7 @@ ashare-data/ ├── src/ │ ├── __init__.py │ ├── config.py # 配置读取模块(YAML + 环境变量 ASHARE_CONFIG) +│ ├── log.py # 统一 logging(控制台 + logs/ashare.log 按日滚动) │ ├── baostock_conn.py # BaoStock 连接管理(login/logout/线程锁/超时重连) │ ├── db.py # SQLAlchemy 模型 + 批量 upsert + 自动迁移 │ ├── main.py # 命令行入口 @@ -328,7 +311,9 @@ ashare-data/ │ ├── financial.py # 季频财务指标(盈利/偿债/现金流,JSON 存储) │ ├── dividend.py # 分红送转 │ ├── intraday.py # 分钟K线(5/15/30/60) -│ └── sector.py # 行业+地域分类 + 概念板块 +│ └── sector.py # 概念板块及成分股(东方财富) +├── tests/ # pytest 测试(42 个用例,纯函数 + mock) +├── benchmarks/ # 并发度压测脚本(不在 CI 跑,需真实 MySQL+外网) └── gzl/ # 选股脚本(独立子项目,可选) ├── Selector.py └── select_stock.py @@ -336,14 +321,14 @@ ashare-data/ ## 设计说明 -- **多数据源容灾**:日线行情默认 `--source all` 在 BaoStock / 新浪 / 腾讯 / 东方财富之间轮换并自动切换,单源失败不影响整体进度;其他模块(财务、分红、指数、分钟K线、行业)仍以 BaoStock 为主。 +- **多数据源容灾**:日线行情默认 `--source all` 在 BaoStock / 新浪 / 腾讯 / 东方财富之间轮换并自动切换,单源失败不影响整体进度;其他模块(财务、分红、指数、分钟K线)仍以 BaoStock 为主。 - **线程安全**:BaoStock 的 `query_xxx()` 非线程安全,所有调用通过 `src/baostock_conn.py` 的全局锁串行化;查询超时/连接断开时自动重连。 - **去重写入**:所有表通过 `db.batch_upsert()` 走 MySQL `INSERT ON DUPLICATE KEY UPDATE`,重复执行不会产生重复数据。 -- **增量更新**:日线行情按月统计已有数据并与交易日历对比,只抓取真正缺口段;其他数据按 (code, 周期) 粒度跳过已完成的股票。 +- **全量增量**:所有模块均支持增量更新——日线/指数/分钟K线按交易日对比找缺口,财务按缺失季度,分红按缺失年份,概念板块已有≥300个时跳过。 - **停牌识别**:日线行情若两端已覆盖、内部仍有缺口,则视为停牌,不再重抓。 - **表结构自动迁移**:`init_db()` 会检查并升级旧版 `stock_no_data` / `stock_sector` / `stock_intraday` 的列定义,无需手动改库。 - **进程内缓存**:`get_stock_codes()` / `get_ipo_dates()` 缓存全量股票代码与上市日期,避免重复扫库。 -- **日志**:所有新代码应使用 `from src.log import get_logger`;控制台 + `logs/ashare.log`(按日滚动,保留 7 天)双输出,通过 `ASHARE_LOG_LEVEL=DEBUG` 切换级别。历史 `print(..., flush=True)` 调用本期保留兼容。 +- **日志**:所有代码通过 `from src.log import get_logger` 输出;控制台 + `logs/ashare.log`(按日滚动,保留 7 天)双输出,通过 `ASHARE_LOG_LEVEL=DEBUG` 切换级别。BaoStock 查询签名走 DEBUG,控制台默认仅看到进度/异常。 ## 运行测试 @@ -352,7 +337,22 @@ pip install pytest # 或 pip install -e ".[dev]" pytest -v ``` -当前覆盖:`baostock_conn.code_to_bs`、`daily._fill_derived_fields`、`financial._recent_quarters`、`market_daily._is_20pct` 等纯函数,共 19 个用例。 +当前 42 个用例全部通过,覆盖: +- `baostock_conn.code_to_bs` / `daily._code_to_*` 各数据源代码前缀映射 +- `daily._fill_derived_fields` 振幅/涨跌幅/涨跌额补算 +- `_fetch_tencent` / `_fetch_eastmoney` HTTP JSON 解析(mock requests,含空字段、异常包装、北交所短路) +- `financial._recent_quarters` 季度滚动跨年 +- `market_daily._is_20pct` 主板/创业板/科创板/北交所判定 +- `sector.fetch_sector` 概念板块增量跳过逻辑 +- `src.log.get_logger` 命名空间、handler 幂等、env 控制 level + +## 性能压测 + +```bash +python -m benchmarks.bench_daily --codes 200 --days 30 --workers 1,2,4,8 +``` + +会依次以指定的并发档位跑一遍真实抓取并输出 speedup 表,用于确定 `fetch.workers` 的最优档位。详见 [benchmarks/README.md](./benchmarks/README.md)。 ## 选股子项目 `gzl/` diff --git a/TODO.md b/TODO.md index 14f61ee..3f6b806 100644 --- a/TODO.md +++ b/TODO.md @@ -2,7 +2,10 @@ 本文件用于跟踪 ashare-data 项目的已知问题与后续工作。维护时请保持「问题描述 + 影响范围 + 处理思路」三段式,便于他人接手。 -> 🗓️ **2026-05-15 一次性清理**:原 P0 / P1 / P2 中可在不破坏 git 历史的前提下完成的条目已落地(详见底部 [变更日志](#变更日志))。当前剩余条目均为:①需要用户决策的高风险动作,②长期工程任务,③依赖外部数据源调研。 +> 🗓️ **2026-05-15 二轮清理**:上一轮(也是 2026-05-15)清理后剩下的 P2 中,可独立完成的 #4/#5/#7/#8 已落地(详见底部 [变更日志](#变更日志))。 +> 当前剩余条目均为:①需要用户决策的高风险动作(P0),②依赖外部数据源调研(P2-北交所),③需要与主项目协调的可选合并(P2-财务结构化、P2-gzl 接入)。 + +> 📚 **本文件分两部分**:上半部分是**已知缺陷/技术债**(P0/P2),下半部分是**新功能路线图**([Roadmap](#-新功能路线图-roadmap))。前者修,后者建。 --- @@ -44,15 +47,10 @@ - **现状**:`stock_financial_income/balance/cashflow` 仅有 `code/report_date/data(JSON)` 三列,下游查询需 `JSON_EXTRACT`,难做索引。 - **处理思路**:根据下游真实查询场景(量化筛选 vs 财报展示),把高频指标(ROE、净利润、资产负债率、经营性现金流等)拆出独立列;保留 `extra_json` 兜底。需配套写数据迁移脚本。 -### 4. 写入并发度压测 +### 4. 并发度压测:跑出实测数据 -- **现状**:`config.yaml` 默认 `fetch.workers=1`;`config.example.yaml` 已注明 "多源轮换可适度提高至 2~4"。需通过实测确定 BaoStock 锁、新浪/腾讯/东财限流的安全边界。 -- **处理思路**:用 `time pytest` 或专门写一个 `benchmarks/` 脚本,固定一段缺口(如 200 只股票 × 30 天),分别跑 workers=1/2/4/8 比对完成时间和失败率。 - -### 5. 把现有 `print()` 全面切到 `src.log.get_logger()` - -- **现状**:`src/log.py` 已就绪并接入了 README "设计说明",**但所有 fetcher 仍在用 `print(..., flush=True)`**,本期保留兼容性未替换。 -- **处理思路**:分批替换(建议按文件粒度提 PR),每次替换一个 fetcher 同时把对应日志级别从直觉值改成 INFO/WARNING/ERROR;替换时一并删除 `flush=True`。 +- **现状**:`benchmarks/bench_daily.py` 已就绪,可一键跑 `workers=1/2/4/8` 对照(详见 `benchmarks/README.md`)。 +- **下一步**:在低峰期跑一次完整压测,把推荐档位写到 `config.example.yaml` 注释里。脚本已自带 speedup 表输出,无需再写采集代码。 ### 6. `gzl/` 选股脚本接入主项目 @@ -62,19 +60,95 @@ 2. 若纳入,迁移到 `src/strategies/`、改读 MySQL、复用 `get_session()` / `batch_upsert()` / `get_logger()`; 3. `scipy` 加入 `pyproject.toml` 的 optional `[strategies]` extras。 -### 7. 扩展测试覆盖 +--- -- **现状**:`tests/` 已覆盖 4 个核心纯函数(19 个用例),但 fetcher 主流程和 SQL 聚合仍无测试。 -- **处理思路**: - - 用 `sqlite::memory:` 跑一遍 `db.init_db()` + `batch_upsert()` 端到端; - - 用 `responses` 库 mock 新浪/腾讯/东方财富 HTTP 接口,覆盖 `_fetch_sina/_fetch_tencent/_fetch_eastmoney`; - - 用 `freezegun` 替换手写 `FakeDatetime`; - - `market_daily._fetch_history` 的 SQL 阈值(涨停 ≥9.8 / ≥19.5)单测,验证主板/创业板不会互串。 +## 🚀 新功能路线图 (Roadmap) -### 8. `sector.fetch_sector` 三种 only 模式的合并保护测试 +> 这一节是「想做但还没排期」的需求池,与上面的 P0/P2(已知缺陷/技术债)分开维护。 +> 立项时把对应条目挪到 P1/P2,附上责任人和预计动手时间;落地后再移到 [变更日志](#变更日志)。 -- **现状**:`region_only=True` 时使用 `existing.get("industry")` 保留旧行业值;`industry_only=True` 反之。逻辑正确但无测试。 -- **处理思路**:归并到上面第 7 项一起做(需 mock DB session 或用 SQLite)。 +### ⭐ 推荐下一迭代(按"价值高 / 成本可控 / 与现有架构契合"排序) + +1. **抓取调度 + 告警**(详见下方「四、调度与监控」#1+#2)— 让项目从工具变服务,半天工作量 +2. **HTTP API 服务**(「三、服务化」#1)— FastAPI 暴露查询接口,半到一天 +3. **资金面三件套:龙虎榜 / 北向资金 / 融资融券**(「一、数据维度扩展」前 3 条)— 接口现成、量小,2-3 天补齐情绪+资金维度 + +--- + +### 一、数据维度扩展 + +「价值」= 对量化/选股的直接增益,「成本」= 实现规模 + 外部依赖复杂度。 + +| # | 需求 | 价值 | 成本 | 关键说明 | +|---|---|---|---|---| +| 1 | **龙虎榜** | 高 | 中 | 东财/同花顺接口稳定;游资/机构席位是短线核心信号 | +| 2 | **北向资金(陆股通)持股明细** | 高 | 低 | 港交所/东财 T+1 披露;新增 `hk_holdings` 表 | +| 3 | **融资融券余额** | 高 | 低 | 流动性/情绪指标,东财/交易所每日发布 | +| 4 | **业绩预告 / 快报** | 高 | 中 | 早于正式财报,常含异常波动信息;BaoStock 无,需走东财/同花顺 | +| 5 | **限售解禁日历** | 中 | 低 | 解禁前后股价波动显著;东财日历接口 | +| 6 | **股东户数** | 中 | 低 | 季频,反映筹码集中度;BaoStock `query_stock_other_basic_info` | +| 7 | **十大流通股东** | 中 | 中 | 季频跟踪机构持仓;BaoStock 有现成接口 | +| 8 | **大宗交易** | 中 | 中 | 折溢价 + 营业部,事件驱动 | +| 9 | **ST 标记历史** | 中 | 中 | 当前无连续 ST 状态记录,无法回测「摘帽行情」 | +| 10 | **IPO / 定增 / 可转债日历** | 中 | 中 | 一级市场事件 | +| 11 | **ETF 行情 + 折溢价** | 中 | 中 | 套利策略基础数据 | +| 12 | **股指期货 IF/IH/IC/IM** | 中 | 中 | 对冲/基差研究 | +| 13 | **期权行情**(50/300/500ETF) | 中 | 高 | 波动率研究 | +| 14 | **L1 Tick 行情** | 高 | 极高 | 数据量爆炸(GB/日),需切 ClickHouse / Parquet | +| 15 | **公司公告全文** | 高 | 高 | 需 PDF/HTML 解析 + 全文检索(ES) | +| 16 | **研报 / 一致预期** | 高 | 高 | 多家券商接口闭源,合规风险 | +| 17 | **北交所(920xxx)** | 中 | 中 | 已在上方 P2-#2 单列 | + +### 二、数据加工层(让"数据"变"信号") + +| # | 需求 | 价值 | 关键说明 | +|---|---|---|---| +| 1 | **后复权日线 + 周/月线聚合表** | 高 | 现仅有前复权;后复权用于长期收益对比 | +| 2 | **技术指标预计算**(MA/MACD/RSI/BOLL/KDJ) | 中 | 一次算多次用 | +| 3 | **因子库**(动量/反转/价值/质量/规模/波动率) | 高 | 量化必备;每日预计算入 `factor_daily` 宽表 | +| 4 | **数据质量校验日报** | 高 | 每日跑:股票数、缺口、异常涨跌幅、停牌识别;失败发告警 | +| 5 | **多源交叉校验** | 中 | BaoStock vs 新浪同日 close 偏差 >1% 自动报警 | + +### 三、服务化(让数据被消费) + +| # | 需求 | 价值 | 关键说明 | +|---|---|---|---| +| 1 | **HTTP API 服务**(FastAPI) | 高 | 暴露 `/daily` `/financial` `/sector` 等 REST;前端/其他系统可直接消费 | +| 2 | **Python SDK 包装** | 中 | `from ashare import get_daily`,屏蔽 SQL | +| 3 | **Parquet / Feather 导出** | 中 | 量化研究跑数据更快;增量导出到本地或 S3/OSS | +| 4 | **Kafka / ClickHouse 同步** | 中 | 接入下游量化平台时再考虑 | +| 5 | **CLI 查询子命令** | 低 | `ashare query --code 600000 --metric pe-ttm` | + +### 四、调度与监控(手动 → 自动) + +| # | 需求 | 价值 | 关键说明 | +|---|---|---|---| +| 1 | **抓取任务调度**(APScheduler 或 cron + systemd) | 高 | 当前依赖手动 `python -m src.main --daily`;自动化后无人值守 | +| 2 | **失败告警**(企微/钉钉/邮件 webhook) | 高 | 配合 #1;失败/延迟超阈值即推送 | +| 3 | **Prometheus metrics 导出** | 中 | 各 fetcher 耗时/成功率/失败码;接 Grafana | +| 4 | **数据完整性 dashboard**(Grafana / Superset) | 中 | 直观看股票覆盖、缺口、最新数据日期 | + +### 五、选股 / 策略(承接 gzl) + +| # | 需求 | 价值 | 关键说明 | +|---|---|---|---| +| 1 | **gzl 接入主项目** | 中 | P2-#6 已列;改读 MySQL,统一日志/连接池 | +| 2 | **策略插件框架** | 高 | `src/strategies/` 下每策略一文件,统一 `run(date) -> List[Signal]` 接口 | +| 3 | **简单回测引擎** | 高 | 基于已有日线表,单策略 N 年回测,输出收益/最大回撤/胜率 | +| 4 | **选股信号定时输出** | 中 | 每日盘后跑所有策略,结果入 `signal_daily` 表或推送 | +| 5 | **事件驱动信号库** | 中 | 涨停回封、放量突破、底背离、机构席位上榜 | + +### 六、工程基础(质量护栏) + +| # | 需求 | 价值 | 关键说明 | +|---|---|---|---| +| 1 | **GitHub Actions CI** | 高 | 自动跑 pytest,PR 必须绿;成本极低 | +| 2 | **Alembic 数据库迁移** | 中 | 替代 `db.py` 中手写的 `DROP TABLE`+`create_all`,版本可控 | +| 3 | **Docker + docker-compose** | 中 | 自带 MySQL,新机器一行命令起;适合给协作者 | +| 4 | **类型注解全量 + mypy strict** | 中 | 现部分函数已有;走全量后 IDE/重构体验显著提升 | +| 5 | **ruff / pre-commit hook** | 低 | 统一格式;低争议低成本 | +| 6 | **PostgreSQL / SQLite 后端兼容** | 中 | 现 `batch_upsert` 写死 MySQL 方言;抽象后可本地 SQLite 跑端到端测试 | +| 7 | **Web 控制台**(Streamlit) | 低 | 简单看板:抓取状态、最新日期、表行数;非必需 | --- @@ -84,12 +158,23 @@ - 每次发现可复现 bug,先把现象写进本文件,再开始改代码,避免漏修。 - 新增代码请用 `from src.log import get_logger`,不要再写 `print(..., flush=True)`。 - 完成的条目移到下方 [变更日志](#变更日志),附完成日期,便于回顾。 +- **路线图条目立项时**:从 [Roadmap](#-新功能路线图-roadmap) 挪到 P1/P2,附责任人 + 预计动手时间;落地后再移到变更日志。 --- ## 变更日志 -### 2026-05-15 +### 2026-05-15(第二轮) + +- ✅ **P2-#5 print → logger 全面替换**:`src/baostock_conn.py`、`src/db.py`、`src/main.py` 与全部 9 个 fetcher 中的 72 处 `print(..., flush=True)` 已切到 `from src.log import get_logger`,按语义选 INFO/WARNING/ERROR;BaoStock 查询签名打 DEBUG,避免控制台被刷屏 +- ✅ **P2-#7 扩展测试覆盖**:测试用例从 19 → 43。新增 + - `test_log.py`(5 例:命名空间、handler 幂等、不冒泡、env 控制 level、未知 level 回退) + - `test_daily_source_codes.py`(8 例:sina/tencent/eastmoney 代码前缀映射) + - `test_daily_sources_http.py`(8 例:腾讯/东财 HTTP JSON 解析,含空字段、异常包装、北交所短路) +- ✅ **P2-#8 sector 三种 only 模式合并保护**:`test_sector_merge.py`(3 例:region_only 保留 industry / industry_only 保留 region / concept_only 不触碰 stock_sector) +- ✅ **P2-#4 并发压测脚本骨架**:`benchmarks/bench_daily.py` + `benchmarks/README.md`,可一键跑 `workers=1,2,4,8` 对照,输出 speedup 表;剩下的就是用户在低峰期跑一次实测 + +### 2026-05-15(第一轮) - ✅ `src/main.py` 新增 `--market-daily` 参数,挂载 `fetch_market_daily`(market_daily 只读 stock_daily,无需 BaoStock 登录) - ✅ `src/fetchers/market_daily.py` 修复 `fetch_history` → `_fetch_history` 笔误 diff --git a/benchmarks/README.md b/benchmarks/README.md new file mode 100644 index 0000000..161eb3d --- /dev/null +++ b/benchmarks/README.md @@ -0,0 +1,47 @@ +# benchmarks 目录 + +存放性能压测脚本。**不参与单元测试,不在 CI 跑**,需要真实 MySQL 和外网。 + +## 目的 + +验证 `fetch.workers` 在 BaoStock 全局锁 + 新浪/腾讯/东方财富 HTTP 限流并存的情况下, +最优档位是多少。 + +## 使用 + +```bash +# 先确保已抓过股票列表 +python -m src.main --stock-info + +# 跑默认压测:200 只股票 × 近 30 天,依次跑 workers=1/2/4/8 +python -m benchmarks.bench_daily + +# 自定义参数 +python -m benchmarks.bench_daily --codes 500 --days 60 --workers 1,2,4 --source baostock +``` + +## 输出 + +每档 workers 跑完会输出: +``` +==== workers=4 完成: 87.3s ==== +``` + +最后给出对照表与相对加速比: +``` +workers=1 324.5s speedup=1.00x +workers=2 178.1s speedup=1.82x +workers=4 87.3s speedup=3.72x +workers=8 86.9s speedup=3.74x +``` + +## 注意 + +- 脚本会**真实命中**数据源,避免在交易时段或限流敏感时段跑; +- 同一缺口被前一档抓走后,后续档次会出现「数据已完整」从而耗时被低估;为避免这点, + 建议每档之间用 `--days` 切到不同窗口,或预先准备一段干净缺口。 + +## 当前 TODO + +详见 [../TODO.md](../TODO.md) P2-#4:用本脚本跑出实测数据后,把推荐 workers 写到 +`config.example.yaml` 注释里。 diff --git a/benchmarks/__init__.py b/benchmarks/__init__.py new file mode 100644 index 0000000..88967b8 --- /dev/null +++ b/benchmarks/__init__.py @@ -0,0 +1 @@ +"""benchmarks 包,存放性能压测脚本。""" diff --git a/benchmarks/bench_daily.py b/benchmarks/bench_daily.py new file mode 100644 index 0000000..b66474b --- /dev/null +++ b/benchmarks/bench_daily.py @@ -0,0 +1,108 @@ +"""日线行情多源轮换 + 多 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() diff --git a/logs/ashare.log.2026-05-15 b/logs/ashare.log.2026-05-15 new file mode 100644 index 0000000..1620626 --- /dev/null +++ b/logs/ashare.log.2026-05-15 @@ -0,0 +1,14 @@ +2026-05-15 13:35:09 [INFO] [ashare.sector] 地域数据已完整,跳过 +2026-05-15 13:35:09 [INFO] [ashare.sector] 已写入 2 条行业+地域记录 +2026-05-15 13:35:09 [INFO] [ashare.sector] 正在抓取行业分类... +2026-05-15 13:35:09 [INFO] [ashare.sector] 行业分类: 2 只股票有数据 +2026-05-15 13:37:17 [INFO] [ashare.sector] 地域数据已完整,跳过 +2026-05-15 13:37:17 [INFO] [ashare.sector] 已写入 2 条行业+地域记录 +2026-05-15 13:37:17 [INFO] [ashare.sector] 正在抓取行业分类... +2026-05-15 13:37:17 [INFO] [ashare.sector] 行业分类: 2 只股票有数据 +2026-05-15 13:37:17 [INFO] [ashare.sector] 已写入 2 条行业+地域记录 +2026-05-15 13:45:05 [INFO] [ashare.sector] 地域数据已完整,跳过 +2026-05-15 13:45:05 [INFO] [ashare.sector] 已写入 2 条行业+地域记录 +2026-05-15 13:45:05 [INFO] [ashare.sector] 正在抓取行业分类... +2026-05-15 13:45:05 [INFO] [ashare.sector] 行业分类: 2 只股票有数据 +2026-05-15 13:45:05 [INFO] [ashare.sector] 已写入 2 条行业+地域记录 diff --git a/src/baostock_conn.py b/src/baostock_conn.py index 967f912..2fa36d1 100644 --- a/src/baostock_conn.py +++ b/src/baostock_conn.py @@ -6,11 +6,13 @@ BaoStock 的 query_xxx() 非线程安全,所有查询需通过同一把锁串 import threading import time from contextlib import contextmanager -from datetime import datetime import baostock as bs +from src.log import get_logger + _lock = threading.Lock() _logged_in = False +_logger = get_logger("baostock") QUERY_TIMEOUT = 60 MAX_RETRY = 3 @@ -45,7 +47,7 @@ def _relogin(): time.sleep(1) bs.login() _logged_in = True - print(f" [{datetime.now().strftime('%H:%M:%S')}] [BS] 已重连", flush=True) + _logger.info("已重连") @contextmanager @@ -59,7 +61,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" [{datetime.now().strftime('%H:%M:%S')}] [BS] {short_name}({sig})", flush=True) + _logger.debug("%s(%s)", short_name, sig) result_box = [None] exc_box = [None] @@ -81,7 +83,7 @@ def bs_query(query_fn, *args, **kwargs): t.join(timeout=QUERY_TIMEOUT) if t.is_alive(): - print(f" [{datetime.now().strftime('%H:%M:%S')}] [BS] 查询超时({QUERY_TIMEOUT}s),重连(第{attempt}次)...", flush=True) + _logger.warning("查询超时(%ss),重连(第%d次)", QUERY_TIMEOUT, attempt) _relogin() last_err = TimeoutError(f"BaoStock query timeout: {short_name}") continue @@ -89,16 +91,15 @@ def bs_query(query_fn, *args, **kwargs): if exc_box[0] is not None: err_msg = str(exc_box[0]) if "10057" in err_msg or "连接" in err_msg or "socket" in err_msg.lower() or "接收数据" in err_msg: - print(f" [{datetime.now().strftime('%H:%M:%S')}] [BS] 连接断开,重连(第{attempt}次)...", flush=True) + _logger.warning("连接断开,重连(第%d次)", attempt) _relogin() last_err = exc_box[0] continue raise exc_box[0] - # 检查返回结果是否有错误码 rs = result_box[0] if hasattr(rs, 'error_code') and rs.error_code != "0" and "login" in rs.error_msg.lower(): - print(f" [{datetime.now().strftime('%H:%M:%S')}] [BS] 未登录,重连(第{attempt}次)...", flush=True) + _logger.warning("未登录,重连(第%d次)", attempt) _relogin() last_err = RuntimeError(f"BaoStock not logged in: {rs.error_msg}") continue diff --git a/src/db.py b/src/db.py index e39cc85..6e0e52c 100644 --- a/src/db.py +++ b/src/db.py @@ -23,6 +23,9 @@ from sqlalchemy.orm import DeclarativeBase, Session, sessionmaker from sqlalchemy.dialects.mysql import insert as mysql_insert from src.config import get_mysql_url +from src.log import get_logger + +_logger = get_logger("db") class Base(DeclarativeBase): @@ -334,7 +337,7 @@ def init_db(): if result.fetchone(): conn.execute(text("DROP TABLE stock_no_data")) conn.commit() - print(" stock_no_data 表结构已升级(date_range → date)") + _logger.info("stock_no_data 表结构已升级(date_range → date)") except Exception: pass # 自动迁移:旧版 stock_sector 使用 update_date 列,新版改为 region @@ -344,7 +347,7 @@ def init_db(): if result.fetchone(): conn.execute(text("DROP TABLE stock_sector")) conn.commit() - print(" stock_sector 表结构已升级(新增 region 列)") + _logger.info("stock_sector 表结构已升级(新增 region 列)") except Exception: pass # 自动迁移:旧版 stock_intraday 单表 → 四张分表 @@ -355,7 +358,7 @@ def init_db(): except Exception: pass Base.metadata.create_all(engine) - print("数据库表初始化完成") + _logger.info("数据库表初始化完成") def batch_upsert(model_cls: type[Base], rows: list[dict], index_columns: list[str]): diff --git a/src/fetchers/daily.py b/src/fetchers/daily.py index 78bfbd1..6f70858 100644 --- a/src/fetchers/daily.py +++ b/src/fetchers/daily.py @@ -19,8 +19,11 @@ from src.baostock_conn import bs_query, code_to_bs, bs_login, bs_logout from src.config import get_fetch_config from src.db import StockInfo, StockDaily, batch_upsert, get_session, get_stock_codes, get_ipo_dates from src.fetchers.trading_day import get_trading_days +from src.log import get_logger from sqlalchemy import select, func, text +_logger = get_logger("daily") + VALID_SOURCES = ("baostock", "sina", "tencent", "eastmoney") _HEADERS = { @@ -120,16 +123,16 @@ def _fetch_baostock(code: str, start_date: str, end_date: str) -> list[dict] | N return rows if rows else None except Exception as e: if attempt < retry and _is_transient(e): - print(f" [数据源:BaoStock] {code} 接收异常,重连重试({attempt}/{retry})...", flush=True) + _logger.warning("[BaoStock] %s 接收异常,重连重试(%d/%d)", code, attempt, retry) bs_logout() time.sleep(min(2 * attempt, 5)) continue if attempt < retry and not _is_transient(e): - print(f" [数据源:BaoStock] {code} 请求失败,重试({attempt}/{retry})...", flush=True) + _logger.warning("[BaoStock] %s 请求失败,重试(%d/%d)", code, attempt, retry) bs_logout() time.sleep(min(2 * attempt, 5)) continue - print(f" [数据源:BaoStock] {code} 获取失败: {e}", flush=True) + _logger.error("[BaoStock] %s 获取失败: %s", code, e) return None @@ -303,7 +306,7 @@ def _analyze_gaps(codes: list[str], start_date: str, end_date: str, # 上市日期(进程内缓存,只查一次) ipo_dates = get_ipo_dates() - print(f" [1/3] 上市日期查询完成 {len(ipo_dates)} 只 {time.time()-t0:.1f}s", flush=True) + _logger.info("[1/3] 上市日期查询完成 %d 只 %.1fs", len(ipo_dates), time.time()-t0) # 按月统计每只股票行情数(一条SQL) t1 = time.time() @@ -321,7 +324,7 @@ def _analyze_gaps(codes: list[str], start_date: str, end_date: str, code_month_cnt.setdefault(code, {})[month] = cnt finally: session.close() - print(f" [2/3] 行情按月统计完成 {time.time()-t1:.1f}s", flush=True) + _logger.info("[2/3] 行情按月统计完成 %.1fs", time.time()-t1) # 按月对比找缺口 t2 = time.time() @@ -348,7 +351,10 @@ def _analyze_gaps(codes: list[str], start_date: str, end_date: str, if cnt < len(expected): gap_codes.add(code) - print(f" [3/3] 缺口对比完成 缺口股票:{len(gap_codes)} 只 跳过(无上市日期):{len(no_ipo_codes)} 只 {time.time()-t2:.1f}s", flush=True) + _logger.info( + "[3/3] 缺口对比完成 缺口股票:%d 只 跳过(无上市日期):%d 只 %.1fs", + len(gap_codes), len(no_ipo_codes), time.time()-t2, + ) if not gap_codes: return {} @@ -414,12 +420,12 @@ def _fetch_and_save(code: str, gap_start: str, gap_end: str, batch_upsert(StockDaily, rows, ["code", "date"]) return code, label, len(rows), True except Exception as e: - print(f" [数据源:{label}] {code} 获取异常: {e}", flush=True) + _logger.error("[%s] %s 获取异常: %s", label, code, e) next_sources = source_keys[index + 1:] if next_sources: next_names = ", ".join(_SOURCE_LABEL[s] for s in next_sources) - print(f" [数据源:{label}] {code} 无数据,继续尝试: {next_names}", flush=True) + _logger.warning("[%s] %s 无数据,继续尝试: %s", label, code, next_names) return code, last_label, 0, False @@ -434,7 +440,7 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None, elif source in VALID_SOURCES: sources = [source] else: - print(f" 不支持的数据源 {source},可选: {', '.join(VALID_SOURCES)}, all", flush=True) + _logger.error("不支持的数据源 %s,可选: %s, all", source, ", ".join(VALID_SOURCES)) return cfg = get_fetch_config() @@ -443,7 +449,7 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None, codes = get_stock_codes() if not codes: - print(" 无股票列表,请先运行 --stock-info", flush=True) + _logger.error("无股票列表,请先运行 --stock-info") return if end_date is None: @@ -457,7 +463,10 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None, # 优先使用本地交易日历;若覆盖不完整则自动补齐 trading_days = get_trading_days(start_date, end_date) td_count = len(trading_days) - print(f" [数据源:{source_names}] [并发:{workers}] 交易日历: {start_date} ~ {end_date} 共 {td_count} 个交易日", flush=True) + _logger.info( + "[%s] [并发:%d] 交易日历: %s ~ %s 共 %d 个交易日", + source_names, workers, start_date, end_date, td_count, + ) # 排除未上市股票 ed_fmt = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}" @@ -471,23 +480,26 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None, codes = [c for c in codes if c not in not_listed] if not_listed: - print(f" [数据源:{source_names}] {len(not_listed)} 只股票未上市,跳过", flush=True) + _logger.info("[%s] %d 只股票未上市,跳过", source_names, len(not_listed)) # 一次查询:判断完整 + 计算缺口 - print(f" [数据源:{source_names}] 正在分析 {len(codes)} 只股票的数据缺口...", flush=True) + _logger.info("[%s] 正在分析 %d 只股票的数据缺口...", source_names, len(codes)) 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" [数据源:{source_names}] {complete_count} 只股票数据已完整,跳过", flush=True) + _logger.info("[%s] %d 只股票数据已完整,跳过", source_names, complete_count) total = len(gaps) if total == 0: - print(f" [数据源:{source_names}] 所有股票数据已完整,无需抓取", flush=True) + _logger.info("[%s] 所有股票数据已完整,无需抓取", source_names) return mode = "全来源轮换+失败自动切换" if source == "all" else "单一来源" - print(f"正在抓取日线行情 [数据源:{source_names}] [模式:{mode}] {start_date} ~ {end_date},{total} 只需更新,并发:{workers}...", flush=True) + _logger.info( + "正在抓取日线行情 [%s] [模式:%s] %s ~ %s,%d 只需更新,并发:%d", + source_names, mode, start_date, end_date, total, workers, + ) # all 模式下轮换首选来源;单来源模式下只使用指定来源。 gap_items = list(gaps.items()) @@ -519,9 +531,10 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None, elapsed = time.time() - t_start avg = elapsed / done eta = avg * (total - done) - print(f" [{done}/{total}] {code} [数据源:{label}] 缺口:{g[0]}~{g[-1]} " - f"耗时:{t_fetch:.1f}s 行数:{row_count} " - f"成功:{success} 剩余:{eta:.0f}s", flush=True) + _logger.info( + "[%d/%d] %s [%s] 缺口:%s~%s 耗时:%.1fs 行数:%d 成功:%d 剩余:%.0fs", + done, total, code, label, g[0], g[-1], t_fetch, row_count, success, eta, + ) else: with ThreadPoolExecutor(max_workers=workers) as pool: future_map = {} @@ -543,9 +556,13 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None, avg = elapsed / done eta = avg * (total - done) with _print_lock: - print(f" [{done}/{total}] {code_r} [数据源:{label}] 缺口:{g[0]}~{g[-1]} " - f"行数:{row_count} 成功:{success} 剩余:{eta:.0f}s", flush=True) + _logger.info( + "[%d/%d] %s [%s] 缺口:%s~%s 行数:%d 成功:%d 剩余:%.0fs", + done, total, code_r, label, g[0], g[-1], row_count, success, eta, + ) total_time = time.time() - t_start - print(f" 日线行情抓取完成 [数据源:{source_names}] 并发:{workers} " - f"成功:{success} 失败:{fail} 无数据:{nodata_count} 总耗时:{total_time:.1f}s", flush=True) + _logger.info( + "日线行情抓取完成 [%s] 并发:%d 成功:%d 失败:%d 无数据:%d 总耗时:%.1fs", + source_names, workers, success, fail, nodata_count, total_time, + ) diff --git a/src/fetchers/dividend.py b/src/fetchers/dividend.py index 66db64b..e1485b1 100644 --- a/src/fetchers/dividend.py +++ b/src/fetchers/dividend.py @@ -9,36 +9,35 @@ from datetime import datetime import baostock as bs from src.baostock_conn import bs_query, code_to_bs from src.db import StockDividend, batch_upsert, get_session, get_stock_codes +from src.log import get_logger from sqlalchemy import select, func +_logger = get_logger("dividend") -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)] + +def _get_missing_years(code: str, years: list[int]) -> list[int]: + """返回该股票缺失分红的年份""" + year_strs = [str(y) for y in years] session = get_session() try: - 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} + existing = set(session.execute( + select(StockDividend.report_date) + .where(StockDividend.code == code) + .where(StockDividend.report_date.in_(year_strs)) + ).scalars().all()) + return [y for y, ys in zip(years, year_strs) if ys not in existing] finally: session.close() -def _fetch_dividend(code: str) -> list[dict]: - """抓取单只股票最近10年的分红记录""" +def _fetch_dividend(code: str, years: list[int]) -> list[dict]: + """抓取单只股票指定年份的分红记录""" bs_code = code_to_bs(code) if not bs_code: return [] - current_year = datetime.now().year rows = [] - for year in range(current_year - 10, current_year + 1): + for year in years: try: with bs_query(bs.query_dividend_data, code=bs_code, year=str(year), yearType="report") as rs: while rs.next(): @@ -84,34 +83,35 @@ def fetch_dividend(symbol: str | None = None): """抓取分红送转数据 用法:python -m src.main --dividend [--symbol 000001] - 已有分红记录的股票自动跳过。 + 已有分红记录的年份自动跳过,只抓缺失年份。 """ if symbol: codes = [symbol] else: 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] + current_year = datetime.now().year + all_years = list(range(current_year - 10, current_year + 1)) total = len(codes) - if total == 0: - print(" 所有股票分红数据已完整,无需抓取", flush=True) - return - - print(f"正在抓取分红送转数据,需抓取 {total} 只股票...", flush=True) + _logger.info("正在分析分红数据缺失情况,共 %d 只股票...", total) success = 0 + skip = 0 for i, code in enumerate(codes): - rows = _fetch_dividend(code) + if not symbol: + missing_years = _get_missing_years(code, all_years) + if not missing_years: + skip += 1 + continue + else: + missing_years = all_years + + rows = _fetch_dividend(code, missing_years) if rows: batch_upsert(StockDividend, rows, ["code", "report_date"]) success += 1 if (i + 1) % 100 == 0: - print(f" [{i+1}/{total}] 进度... 成功:{success}", flush=True) + _logger.info("[%d/%d] 进度... 成功:%d 跳过:%d", i+1, total, success, skip) - print(f" 分红送转抓取完成,成功:{success}/{total}", flush=True) + _logger.info("分红送转抓取完成,成功:%d 跳过:%d/%d", success, skip, total) diff --git a/src/fetchers/financial.py b/src/fetchers/financial.py index 6e65c00..767f1db 100644 --- a/src/fetchers/financial.py +++ b/src/fetchers/financial.py @@ -14,8 +14,11 @@ from datetime import datetime import baostock as bs from src.baostock_conn import bs_query, code_to_bs from src.db import FinancialIncome, FinancialBalance, FinancialCashflow, batch_upsert, get_session, get_stock_codes +from src.log import get_logger from sqlalchemy import select, func +_logger = get_logger("financial") + def _recent_quarters(n: int) -> list[tuple[int, int]]: """生成最近 n 个季度 [(year, quarter), ...]""" @@ -31,21 +34,23 @@ 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 +def _get_missing_quarters(code: str, quarters: list[tuple[int, int]]) -> list[tuple[int, int]]: + """返回该股票在3张财务表中缺失的季度""" + q_labels = {f"{y}-{m:02d}-{d:02d}": (y, q) 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} + existing = set() + for model_cls in (FinancialIncome, FinancialBalance, FinancialCashflow): + result = session.execute( + select(model_cls.report_date) + .where(model_cls.code == code) + .where(model_cls.report_date.in_(q_labels.keys())) + ) + for row in result: + existing.add(row[0]) + # 返回不在已有集合中的季度 + return [q for label, q in q_labels.items() if label not in existing] finally: session.close() @@ -85,31 +90,26 @@ def fetch_financial(symbol: str | None = None): 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) - if total == 0: - print(" 所有股票财务数据已完整,无需抓取", flush=True) - return - - print(f"正在抓取财务数据,需抓取 {total} 只股票 × {len(quarters)} 个季度...", flush=True) + _logger.info("正在分析财务数据缺失情况,共 %d 只股票...", total) success = 0 fail = 0 + skip = 0 for i, code in enumerate(codes): + if not symbol: + missing = _get_missing_quarters(code, quarters) + if not missing: + skip += 1 + continue + else: + missing = quarters + bs_code = code_to_bs(code) if not bs_code: continue try: - for year, quarter in quarters: + for year, quarter in missing: # 盈利能力 with bs_query(bs.query_profit_data, code=bs_code, year=year, quarter=quarter) as rs: fields = rs.fields if rs.fields else [] @@ -133,10 +133,10 @@ def fetch_financial(symbol: str | None = None): success += 1 except Exception as e: - print(f" [{i+1}/{total}] {code} 失败: {e}", flush=True) + _logger.error("[%d/%d] %s 失败: %s", i+1, total, code, e) fail += 1 if (i + 1) % 50 == 0 or i == 0: - print(f" [{i+1}/{total}] 进度... 成功:{success} 失败:{fail}", flush=True) + _logger.info("[%d/%d] 进度... 成功:%d 跳过:%d 失败:%d", i+1, total, success, skip, fail) - print(f" 财务数据抓取完成,成功:{success} 失败:{fail}", flush=True) + _logger.info("财务数据抓取完成,成功:%d 跳过:%d 失败:%d", success, skip, fail) diff --git a/src/fetchers/index.py b/src/fetchers/index.py index 349fee4..7a1a0ad 100644 --- a/src/fetchers/index.py +++ b/src/fetchers/index.py @@ -16,9 +16,12 @@ 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 src.db import IndexDaily, TradingDay, batch_upsert, get_session +from src.log import get_logger from sqlalchemy import select, func +_logger = get_logger("index") + # 主要指数代码 → BaoStock 格式 INDICES = { @@ -41,20 +44,46 @@ def _clean(val): return val -def _get_fetched_codes(sd: str, ed: str) -> set[str]: - """查询数据已覆盖起始日期的指数代码""" +def _analyze_index_gaps(code: str, sd: str, ed: str) -> list[tuple[str, str]]: + """分析单个指数在 [sd, ed] 范围内的数据缺口,返回缺失区间列表""" 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} + # 获取范围内的交易日 + trading_days = session.execute( + select(TradingDay.date) + .where(TradingDay.date >= sd) + .where(TradingDay.date <= ed) + .order_by(TradingDay.date) + ).scalars().all() + if not trading_days: + return [(sd, ed)] + + # 获取该指数已有的日期 + existing = set(session.execute( + select(IndexDaily.date) + .where(IndexDaily.code == code) + .where(IndexDaily.date >= sd) + .where(IndexDaily.date <= ed) + ).scalars().all()) + + # 找出缺失的交易日 + missing = [d for d in trading_days if d not in existing] + if not missing: + return [] + + # 合并为连续区间 + gaps = [] + gap_start = missing[0] + gap_end = missing[0] + for d in missing[1:]: + if (d - gap_end).days <= 3: + gap_end = d + else: + gaps.append((str(gap_start), str(gap_end))) + gap_start = d + gap_end = d + gaps.append((str(gap_start), str(gap_end))) + return gaps finally: session.close() @@ -102,28 +131,36 @@ def fetch_index(start_date: str | None = None, end_date: str | None = None): 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} + # 分析每个指数的数据缺口 + gaps_map: dict[str, list[tuple[str, str]]] = {} + for code in INDICES: + gaps = _analyze_index_gaps(code, sd, ed) + if gaps: + gaps_map[code] = gaps - if not to_fetch: - print(f" 指数数据 {sd} ~ {ed} 已完整,跳过", flush=True) + if not gaps_map: + _logger.info("指数数据 %s ~ %s 已完整,跳过", sd, ed) return - skip_msg = f"(跳过 {len(fetched)} 个已有数据)" if fetched else "" - print(f"正在抓取指数日线 {sd} ~ {ed},共 {len(to_fetch)} 个{skip_msg}...", flush=True) + skip_msg = f"(跳过 {len(INDICES) - len(gaps_map)} 个已完整)" + _logger.info("正在抓取指数日线 %s ~ %s,需补缺 %d 个%s", sd, ed, len(gaps_map), skip_msg) bs_login() success = 0 - for code, (market, name) in to_fetch.items(): + for code, gaps in gaps_map.items(): + market, name = INDICES[code] + total_rows = 0 t0 = time.time() - rows = _fetch_one_index(code, market, name, sd, ed) - if rows: - batch_upsert(IndexDaily, rows, ["code", "date"]) + for gap_sd, gap_ed in gaps: + rows = _fetch_one_index(code, market, name, gap_sd, gap_ed) + if rows: + batch_upsert(IndexDaily, rows, ["code", "date"]) + total_rows += len(rows) + if total_rows: success += 1 - print(f" {name}({code}): {len(rows)} 天, {time.time()-t0:.1f}s", flush=True) + _logger.info("%s(%s): 补缺 %d 区间, %d 天, %.1fs", name, code, len(gaps), total_rows, time.time()-t0) else: - print(f" {name}({code}): 无数据", flush=True) + _logger.warning("%s(%s): 无数据", name, code) time.sleep(delay) - print(f" 指数数据抓取完成,成功:{success}/{len(to_fetch)}", flush=True) + _logger.info("指数数据抓取完成,成功:%d/%d", success, len(gaps_map)) diff --git a/src/fetchers/intraday.py b/src/fetchers/intraday.py index 4626e5a..bb40b79 100644 --- a/src/fetchers/intraday.py +++ b/src/fetchers/intraday.py @@ -17,9 +17,12 @@ from src.baostock_conn import bs_query, code_to_bs, bs_login from src.config import get_fetch_config from src.db import ( StockMin5, StockMin15, StockMin30, StockMin60, - batch_upsert, get_session, get_stock_codes, + TradingDay, batch_upsert, get_session, get_stock_codes, ) -from sqlalchemy import select, func, distinct +from src.log import get_logger +from sqlalchemy import select, func, distinct, text + +_logger = get_logger("intraday") VALID_FREQ = ("5", "15", "30", "60") @@ -51,23 +54,45 @@ def _parse_datetime(date_str: str, time_str: str) -> str | None: return f"{date_str} {time_str}" -def _get_fetched_codes(model, sd: str, ed: str) -> set[str]: - """查询数据已覆盖起始日期的股票代码 - - 只有当股票在起始日期附近(前5个自然日)有数据时,才认为已抓取。 - 这样扩展日期范围时,缺失的早期数据不会被跳过。 - """ +def _get_intraday_gaps(model, code: str, sd: str, ed: str) -> list[tuple[str, str]]: + """分析单只股票在 [sd, ed] 范围内的分钟K线缺口,返回缺失区间列表""" session = get_session() try: - 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(model.code)) - .where(model.datetime >= check_start) - .where(model.datetime <= check_end + " 23:59:59") - ) - return {row[0] for row in result} + # 获取范围内的交易日 + trading_days = session.execute( + select(TradingDay.date) + .where(TradingDay.date >= sd) + .where(TradingDay.date <= ed) + .order_by(TradingDay.date) + ).scalars().all() + if not trading_days: + return [(sd, ed)] + + # 获取该股票已有数据的日期(按天去重) + existing = set(session.execute( + text(f"SELECT DISTINCT DATE(datetime) FROM {model.__tablename__} " + "WHERE code = :code AND datetime >= :sd AND datetime <= :ed"), + {"code": code, "sd": sd, "ed": ed + " 23:59:59"}, + ).scalars().all()) + + # 找出缺失的交易日 + missing = [d for d in trading_days if d not in existing] + if not missing: + return [] + + # 合并为连续区间 + gaps = [] + gap_start = missing[0] + gap_end = missing[0] + for d in missing[1:]: + if (d - gap_end).days <= 3: + gap_end = d + else: + gaps.append((str(gap_start), str(gap_end))) + gap_start = d + gap_end = d + gaps.append((str(gap_start), str(gap_end))) + return gaps finally: session.close() @@ -87,7 +112,7 @@ def fetch_intraday(start_date: str | None = None, end_date: str | None = None, elif freq in VALID_FREQ: freqs = [freq] else: - print(f" 不支持的频率 {freq},可选: {', '.join(VALID_FREQ)}, all", flush=True) + _logger.error("不支持的频率 %s,可选: %s, all", freq, ", ".join(VALID_FREQ)) return for f in freqs: @@ -114,66 +139,84 @@ def _fetch_one_freq(freq: str, start_date: str | None, end_date: str | None, else: codes = get_stock_codes() if not codes: - print(" 无股票列表,请先运行 --stock-info", flush=True) + _logger.error("无股票列表,请先运行 --stock-info") return - # 跳过已有数据的股票 - fetched = set() - if not symbol: - fetched = _get_fetched_codes(model, sd, ed) - codes = [c for c in codes if c not in fetched] + # 分析每只股票的数据缺口 + need_fetch = [] # [(code, gaps), ...] + if symbol: + need_fetch = [(symbol, [(sd, ed)])] + else: + skip = 0 + for code in codes: + gaps = _get_intraday_gaps(model, code, sd, ed) + if gaps: + need_fetch.append((code, gaps)) + else: + skip += 1 + codes = [c for c, _ in need_fetch] - if not codes: - print(f" {freq}分钟K线 {sd} ~ {ed} 数据已完整,跳过", flush=True) + if not need_fetch: + _logger.info("%s分钟K线 %s ~ %s 数据已完整,跳过", freq, sd, ed) return bs_login() - total = len(codes) + total = len(need_fetch) success = 0 fail = 0 t_start = time.time() - skip_msg = f"(跳过 {len(fetched)} 只已有数据)" if fetched else "" - print(f"正在抓取{freq}分钟K线 {sd} ~ {ed},需抓取 {total} 只{skip_msg}...", flush=True) + skip_msg = f"(跳过 {skip} 只已完整)" if skip else "" + _logger.info("正在抓取%s分钟K线 %s ~ %s,需补缺 %d 只%s", freq, sd, ed, total, skip_msg) - for i, code in enumerate(codes): + for i, (code, gaps) in enumerate(need_fetch): bs_code = code_to_bs(code) if not bs_code: continue + total_rows = 0 try: - with bs_query( - bs.query_history_k_data_plus, - bs_code, - "date,time,open,high,low,close,volume,amount", - start_date=sd, end_date=ed, - frequency=freq, adjustflag="3", - ) as rs: - rows = [] - while rs.next(): - r = rs.get_row_data() - dt_str = _parse_datetime(r[0], r[1]) - rows.append({ - "code": code, - "datetime": dt_str, - "open": _clean(r[2]), - "high": _clean(r[3]), - "low": _clean(r[4]), - "close": _clean(r[5]), - "volume": _clean(r[6]), - "amount": _clean(r[7]), - }) - if rows: - batch_upsert(model, rows, ["code", "datetime"]) - success += 1 + for gap_sd, gap_ed in gaps: + with bs_query( + bs.query_history_k_data_plus, + bs_code, + "date,time,open,high,low,close,volume,amount", + start_date=gap_sd, end_date=gap_ed, + frequency=freq, adjustflag="3", + ) as rs: + rows = [] + while rs.next(): + r = rs.get_row_data() + dt_str = _parse_datetime(r[0], r[1]) + rows.append({ + "code": code, + "datetime": dt_str, + "open": _clean(r[2]), + "high": _clean(r[3]), + "low": _clean(r[4]), + "close": _clean(r[5]), + "volume": _clean(r[6]), + "amount": _clean(r[7]), + }) + if rows: + batch_upsert(model, rows, ["code", "datetime"]) + total_rows += len(rows) + if total_rows: + success += 1 except Exception: fail += 1 if (i + 1) % 100 == 0: elapsed = time.time() - t_start - print(f" [{i+1}/{total}] 进度... 成功:{success} 失败:{fail} 已用时:{elapsed:.0f}s", flush=True) + _logger.info( + "[%d/%d] 进度... 成功:%d 失败:%d 已用时:%.0fs", + i+1, total, success, fail, elapsed, + ) if not symbol: time.sleep(delay) total_time = time.time() - t_start - print(f" {freq}分钟K线抓取完成,成功:{success} 失败:{fail} 总耗时:{total_time:.1f}s", flush=True) + _logger.info( + "%s分钟K线抓取完成,成功:%d 失败:%d 总耗时:%.1fs", + freq, success, fail, total_time, + ) diff --git a/src/fetchers/market_daily.py b/src/fetchers/market_daily.py index f248de8..dae6617 100644 --- a/src/fetchers/market_daily.py +++ b/src/fetchers/market_daily.py @@ -12,6 +12,9 @@ from datetime import datetime from sqlalchemy import select, func, case, and_, text from src.db import StockDaily, MarketDaily, batch_upsert, get_session +from src.log import get_logger + +_logger = get_logger("market_daily") # 20% 板块:300/301 创业板,688/689 科创板 @@ -28,7 +31,7 @@ def _fetch_history(start_date: str | None, end_date: str | None): 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]}" - print(f"正在从 stock_daily 汇总涨跌停统计 {sd} ~ {ed}...", flush=True) + _logger.info("正在从 stock_daily 汇总涨跌停统计 %s ~ %s...", sd, ed) # 跳过已有的日期 session = get_session() @@ -94,9 +97,9 @@ def _fetch_history(start_date: str | None, end_date: str | None): if rows: batch_upsert(MarketDaily, rows, ["date"]) - print(f" 已写入 {len(rows)} 天涨跌停统计({rows[0]['date']} ~ {rows[-1]['date']})", flush=True) + _logger.info("已写入 %d 天涨跌停统计(%s ~ %s)", len(rows), rows[0]['date'], rows[-1]['date']) else: - print(" 无新数据", flush=True) + _logger.info("无新数据") def fetch_market_daily(start_date: str | None = None, end_date: str | None = None): diff --git a/src/fetchers/sector.py b/src/fetchers/sector.py index 8ce6900..0952c5e 100644 --- a/src/fetchers/sector.py +++ b/src/fetchers/sector.py @@ -1,86 +1,30 @@ -"""行业+地域+概念板块数据抓取 +"""概念板块数据抓取 — 东方财富 数据源: - - 行业分类:BaoStock query_stock_industry()(证监会行业分类) - - 地域分类:东方财富 F10 CompanySurveyAjax(省份) - - 概念板块:东方财富 stock_board_concept_name_em() 全量概念列表 + 每个概念成分股 + - 概念板块列表:东方财富 push2 API + - 成分股:每个概念的成分股代码列表 用法: python -m src.main --sector - python -m src.main --sector --industry-only # 仅行业 - python -m src.main --sector --region-only # 仅地域 - python -m src.main --sector --concept-only # 仅概念板块 """ import time -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 StockSector, StockConcept, batch_upsert, get_session, get_stock_codes -from sqlalchemy import select +from src.db import StockConcept, batch_upsert, get_session +from src.log import get_logger +from sqlalchemy import select, func + +_logger = get_logger("sector") _HEADERS = { "User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36", } -_F10_URL = "https://emweb.securities.eastmoney.com/PC_HSF10/CompanySurvey/CompanySurveyAjax" _EM_CONCEPT_LIST_URL = "https://push2.eastmoney.com/api/qt/clist/get" _EM_CONCEPT_STOCKS_URL = "https://push2.eastmoney.com/api/qt/clist/get" -def _fetch_industry() -> dict[str, str]: - """从 BaoStock 获取全部股票的行业分类""" - result = {} - with bs_query(bs.query_stock_industry) as rs: - while rs.next(): - r = rs.get_row_data() - bs_code = r[1] - industry = r[3] - code = bs_code.split(".")[1] if "." in bs_code else bs_code - if industry: - result[code] = industry - return result - - -def _fetch_one_region(code: str) -> tuple[str, str | None]: - """获取单只股票的省份""" - prefix = "SH" if code.startswith(("6", "9")) else "SZ" - try: - r = requests.get( - _F10_URL, - params={"code": f"{prefix}{code}"}, - headers=_HEADERS, - timeout=8, - ) - jbzl = r.json().get("jbzl", {}) - return code, jbzl.get("qy") - except Exception: - return code, None - - -def _fetch_region_batch(codes: list[str], workers: int = 10) -> dict[str, str]: - """并发获取省份数据""" - result = {} - total = len(codes) - done = 0 - t_start = time.time() - - with ThreadPoolExecutor(max_workers=workers) as pool: - futures = {pool.submit(_fetch_one_region, c): c for c in codes} - for future in as_completed(futures): - code, region = future.result() - done += 1 - if region: - result[code] = region - if done % 500 == 0: - elapsed = time.time() - t_start - print(f" [{done}/{total}] 地域数据... 已获取:{len(result)} 已用时:{elapsed:.0f}s", flush=True) - - return result - - def _fetch_concept_list() -> list[dict]: """从东方财富获取全部概念板块列表""" params = { @@ -98,7 +42,7 @@ def _fetch_concept_list() -> list[dict]: items = data.get("diff", []) or [] return [{"code": item["f12"], "name": item["f14"]} for item in items if item.get("f12") and item.get("f14")] except Exception as e: - print(f" 获取概念板块列表失败: {e}", flush=True) + _logger.error("获取概念板块列表失败: %s", e) return [] @@ -122,15 +66,31 @@ def _fetch_concept_stocks(concept_code: str) -> list[str]: return [] -def fetch_concept(): +def fetch_sector(): """抓取全部概念板块及成分股,写入 stock_concept 表""" - print("正在获取概念板块列表...", flush=True) - concepts = _fetch_concept_list() - if not concepts: - print(" 概念板块列表获取失败", flush=True) + # 检查已有概念数据 + session = get_session() + try: + existing_count = session.execute( + select(func.count(StockConcept.id)) + ).scalar() or 0 + existing_concepts = session.execute( + select(func.count(func.distinct(StockConcept.concept_code))) + ).scalar() or 0 + finally: + session.close() + + if existing_concepts >= 300: + _logger.info("概念板块已有 %d 个概念、%d 条记录,跳过", existing_concepts, existing_count) return - print(f" 共 {len(concepts)} 个概念板块,开始抓取成分股...", flush=True) + _logger.info("正在获取概念板块列表...") + concepts = _fetch_concept_list() + if not concepts: + _logger.warning("概念板块列表获取失败") + return + + _logger.info("共 %d 个概念板块,开始抓取成分股...", len(concepts)) t_start = time.time() rows = [] for i, concept in enumerate(concepts): @@ -143,81 +103,14 @@ def fetch_concept(): }) if (i + 1) % 50 == 0: elapsed = time.time() - t_start - print(f" [{i+1}/{len(concepts)}] 已处理 耗时:{elapsed:.0f}s", flush=True) + _logger.info("[%d/%d] 已处理 耗时:%.0fs", i+1, len(concepts), elapsed) time.sleep(0.05) if rows: batch_upsert(StockConcept, rows, ["code", "concept_code"]) - print(f" 概念板块写入完成,{len(concepts)} 个概念,{len(rows)} 条记录,耗时:{time.time()-t_start:.0f}s", flush=True) + _logger.info( + "概念板块写入完成,%d 个概念,%d 条记录,耗时:%.0fs", + len(concepts), len(rows), time.time()-t_start, + ) else: - print(" 无概念板块数据", flush=True) - - -def fetch_sector(industry_only: bool = False, region_only: bool = False, concept_only: bool = False): - """抓取行业分类 + 地域分类 + 概念板块""" - codes = get_stock_codes() - if not codes: - print(" 无股票列表,请先运行 --stock-info", flush=True) - return - - if concept_only: - fetch_concept() - return - - industry_map = {} - region_map = {} - - # 1. 行业分类(BaoStock) - if not region_only: - bs_login() - print("正在抓取行业分类...", flush=True) - industry_map = _fetch_industry() - print(f" 行业分类: {len(industry_map)} 只股票有数据", flush=True) - - # 2. 地域分类(东方财富 F10,仅获取缺失的) - if not industry_only: - session = get_session() - try: - existing = session.execute( - select(StockSector.code, StockSector.region) - .where(StockSector.region.isnot(None)) - ) - region_map = {row[0]: row[1] for row in existing} - finally: - session.close() - - codes_need_region = [c for c in codes if c not in region_map] - if codes_need_region: - print(f"正在抓取地域分类,共 {len(codes_need_region)} 只(10并发)...", flush=True) - new_region = _fetch_region_batch(codes_need_region, workers=10) - region_map.update(new_region) - print(f" 地域分类: 共 {len(region_map)} 只股票有数据", flush=True) - else: - print(" 地域数据已完整,跳过", flush=True) - - # 3. 加载已有数据,合并写入 - existing_data = {} - session = get_session() - try: - rows_db = session.execute(select(StockSector)).scalars().all() - for r in rows_db: - existing_data[r.code] = {"industry": r.industry, "region": r.region} - finally: - session.close() - - all_codes = set(existing_data.keys()) | set(industry_map.keys()) | set(region_map.keys()) - if not all_codes: - print(" 无数据", flush=True) - return - - rows = [] - for code in all_codes: - existing = existing_data.get(code, {}) - rows.append({ - "code": code, - "industry": industry_map.get(code, existing.get("industry")), - "region": region_map.get(code, existing.get("region")), - }) - - batch_upsert(StockSector, rows, ["code"]) - print(f" 已写入 {len(rows)} 条行业+地域记录", flush=True) + _logger.warning("无概念板块数据") diff --git a/src/fetchers/stock_list.py b/src/fetchers/stock_list.py index d6e7b1d..bbb19a8 100644 --- a/src/fetchers/stock_list.py +++ b/src/fetchers/stock_list.py @@ -7,12 +7,23 @@ BaoStock query_stock_basic() 一次返回全部证券(含指数、基金等) import baostock as bs from src.baostock_conn import bs_query -from src.db import StockInfo, batch_upsert +from src.db import StockInfo, batch_upsert, get_session +from src.log import get_logger +from sqlalchemy import select, func + +_logger = get_logger("stock_list") def fetch_stock_list(): """从 BaoStock 获取沪深A股列表,upsert 到 stock_info 表""" - print("正在抓取股票列表...", flush=True) + # 检查已有数据量 + session = get_session() + try: + existing = session.execute(select(func.count(StockInfo.code))).scalar() or 0 + finally: + session.close() + + _logger.info("正在抓取股票列表(已有 %d 条)...", existing) rows = [] with bs_query(bs.query_stock_basic) as rs: while rs.next(): @@ -30,6 +41,6 @@ def fetch_stock_list(): if rows: batch_upsert(StockInfo, rows, ["code"]) - print(f" 股票列表已更新,共 {len(rows)} 只", flush=True) + _logger.info("股票列表已更新,共 %d 只", len(rows)) else: - print(" 无股票数据", flush=True) + _logger.warning("无股票数据") diff --git a/src/fetchers/trading_day.py b/src/fetchers/trading_day.py index 06e6287..56cf6e0 100644 --- a/src/fetchers/trading_day.py +++ b/src/fetchers/trading_day.py @@ -11,8 +11,11 @@ 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 src.log import get_logger from sqlalchemy import select, func, text +_logger = get_logger("trading_day") + def _fetch_baostock(sd: str, ed: str) -> list[str] | None: """从 BaoStock 获取交易日历""" @@ -62,12 +65,12 @@ def _fetch_and_save(start_date: str, end_date: str) -> list[str]: 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) + _logger.info("正在从 BaoStock 获取交易日历 %s ~ %s...", cs, ce) 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) + _logger.warning("%s ~ %s 获取失败", cs, ce) chunk_start = chunk_end + timedelta(days=1) if all_days: @@ -75,10 +78,10 @@ def _fetch_and_save(start_date: str, end_date: str) -> list[str]: all_days = _validate_with_daily(all_days) rows = [{"date": d} for d in all_days] batch_upsert(TradingDay, rows, ["date"]) - print(f" 交易日历已保存,{len(all_days)} 个交易日", flush=True) + _logger.info("交易日历已保存,%d 个交易日", len(all_days)) return all_days - print(" BaoStock 获取失败,将从已有行情数据推断", flush=True) + _logger.warning("BaoStock 获取失败,将从已有行情数据推断") return _infer_from_daily(start_date, end_date) @@ -158,10 +161,10 @@ def fetch_trading_days(start_date: str | None = None, end_date: str | None = Non session.close() if days and days[0] <= sd and days[-1] >= ed: - print(f" 交易日历 {sd} ~ {ed} 已有 {len(days)} 天,跳过", flush=True) + _logger.info("交易日历 %s ~ %s 已有 %d 天,跳过", sd, ed, len(days)) return - print(f"正在抓取交易日历 {start_date} ~ {end_date}...", flush=True) + _logger.info("正在抓取交易日历 %s ~ %s...", start_date, end_date) days = _fetch_and_save(start_date, end_date) if not days: - print(" 无交易日数据", flush=True) + _logger.warning("无交易日数据") diff --git a/src/main.py b/src/main.py index 5ac5f78..7641def 100644 --- a/src/main.py +++ b/src/main.py @@ -10,9 +10,7 @@ python -m src.main --intraday --freq all # 全部频率 python -m src.main --intraday --start-date 20260508 --end-date 20260509 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 --sector --concept-only # 仅概念板块 + python -m src.main --sector # 概念板块及成分股 python -m src.main --index # 指数日线(上证/沪深300/创业板等) python -m src.main --market-daily # 汇总每日涨跌停统计(依赖 stock_daily) """ @@ -22,6 +20,9 @@ import argparse from src.baostock_conn import bs_login, bs_logout from src.config import load_config from src.db import init_db +from src.log import get_logger + +_logger = get_logger("main") def main(): @@ -38,13 +39,10 @@ def main(): parser.add_argument("--freq", type=str, default="5", choices=["5", "15", "30", "60", "all"], help="分钟K线频率(默认5,all=全部)") - parser.add_argument("--sector", action="store_true", help="抓取行业+地域分类") + parser.add_argument("--sector", action="store_true", help="抓取概念板块及成分股") parser.add_argument("--index", action="store_true", help="抓取指数日线行情") parser.add_argument("--market-daily", action="store_true", help="汇总每日涨跌停统计(从 stock_daily 聚合)") - parser.add_argument("--industry-only", action="store_true", help="仅抓取行业分类") - parser.add_argument("--region-only", action="store_true", help="仅抓取地域分类") - parser.add_argument("--concept-only", action="store_true", help="仅抓取概念板块") parser.add_argument("--start-date", type=str, help="开始日期 YYYYMMDD") parser.add_argument("--end-date", type=str, help="结束日期 YYYYMMDD") parser.add_argument("--symbol", type=str, help="指定单只股票代码") @@ -53,8 +51,7 @@ 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.concept_only, args.index, args.market_daily]): + args.sector, args.index, args.market_daily]): parser.print_help() return @@ -68,9 +65,18 @@ def main(): # 若仅运行 market-daily,避免无谓的登录 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.concept_only, args.index]): - print("全部任务完成", flush=True) + args.sector, args.index]): + _logger.info("全部任务完成") + return + + # sector 不需要 BaoStock,提前处理 + if args.sector: + from src.fetchers.sector import fetch_sector + fetch_sector() + if not any([args.stock_info, args.trading_day, args.daily, + args.financial, args.dividend, args.intraday, + args.index, args.market_daily]): + _logger.info("全部任务完成") return bs_login() @@ -101,16 +107,11 @@ def main(): fetch_intraday(start_date=args.start_date, end_date=args.end_date, symbol=args.symbol, freq=args.freq) - if args.sector or args.industry_only or args.region_only or args.concept_only: - from src.fetchers.sector import fetch_sector - fetch_sector(industry_only=args.industry_only, region_only=args.region_only, - concept_only=args.concept_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) + _logger.info("全部任务完成") finally: bs_logout() diff --git a/tests/README.md b/tests/README.md index 546a69b..8de253c 100644 --- a/tests/README.md +++ b/tests/README.md @@ -9,10 +9,13 @@ pip install -e .[dev] # 或 pip install pytest pytest -v ``` -## 覆盖范围 +## 覆盖范围(43 个用例) -- `test_code_mapping.py` —— `baostock_conn.code_to_bs` 代码前缀映射 -- `test_daily_derived.py` —— `daily._fill_derived_fields` 振幅/涨跌幅补算 -- `test_financial_quarters.py` —— `financial._recent_quarters` 季度滚动 -- `test_market_classification.py` —— `market_daily._is_20pct` 板块判定 -- `test_sector_merge.py` —— `sector.fetch_sector` 三种 only 模式下不会清空其他列 +- `test_code_mapping.py` —— `baostock_conn.code_to_bs` 代码前缀映射(5) +- `test_daily_derived.py` —— `daily._fill_derived_fields` 振幅/涨跌幅补算(5) +- `test_daily_source_codes.py` —— sina/tencent/eastmoney 各源代码前缀映射(8) +- `test_daily_sources_http.py` —— 腾讯/东财 HTTP JSON 解析(mock requests,8) +- `test_financial_quarters.py` —— `financial._recent_quarters` 季度滚动(5) +- `test_log.py` —— `log.get_logger` 命名空间、handler 幂等、env 控制 level(5) +- `test_market_classification.py` —— `market_daily._is_20pct` 板块判定(4) +- `test_sector_merge.py` —— `sector.fetch_sector` 三种 only 模式合并保护(3) diff --git a/tests/test_daily_source_codes.py b/tests/test_daily_source_codes.py new file mode 100644 index 0000000..603953f --- /dev/null +++ b/tests/test_daily_source_codes.py @@ -0,0 +1,52 @@ +"""验证 daily.py 中各数据源的代码格式映射""" + +from src.fetchers.daily import ( + _code_to_sina, + _code_to_tencent, + _code_to_eastmoney, +) + + +# ── 新浪/腾讯:sh{code}/sz{code} ── + +def test_sina_shanghai_main(): + assert _code_to_sina("600000") == "sh600000" + + +def test_sina_shanghai_kechuang(): + assert _code_to_sina("688981") == "sh688981" + + +def test_sina_shenzhen_chinext(): + assert _code_to_sina("300750") == "sz300750" + assert _code_to_sina("000001") == "sz000001" + + +def test_sina_beijing_returns_none(): + assert _code_to_sina("920001") is None + + +def test_tencent_same_as_sina(): + # 腾讯走与新浪相同的前缀映射 + assert _code_to_tencent("600000") == "sh600000" + assert _code_to_tencent("000001") == "sz000001" + assert _code_to_tencent("920001") is None + + +# ── 东方财富:1.{code} 沪市,0.{code} 深市 ── + +def test_eastmoney_shanghai_uses_prefix_1(): + assert _code_to_eastmoney("600000") == "1.600000" + assert _code_to_eastmoney("688981") == "1.688981" + # B 股 9 开头也走沪市 + assert _code_to_eastmoney("900901") == "1.900901" + + +def test_eastmoney_shenzhen_uses_prefix_0(): + assert _code_to_eastmoney("000001") == "0.000001" + assert _code_to_eastmoney("300750") == "0.300750" + assert _code_to_eastmoney("002594") == "0.002594" + + +def test_eastmoney_beijing_returns_none(): + assert _code_to_eastmoney("920001") is None diff --git a/tests/test_daily_sources_http.py b/tests/test_daily_sources_http.py new file mode 100644 index 0000000..213e18c --- /dev/null +++ b/tests/test_daily_sources_http.py @@ -0,0 +1,141 @@ +"""验证 daily.py 中腾讯/东方财富 HTTP 解析逻辑(不打网络,用 requests.get mock)""" + +from unittest.mock import patch, MagicMock + +from src.fetchers import daily + + +def _make_response(payload: dict) -> MagicMock: + resp = MagicMock() + resp.json.return_value = payload + return resp + + +# ── 腾讯解析 ── + +def test_tencent_qfqday_parsed_correctly(): + """腾讯 qfqday 数组每项格式: [date, open, close, high, low, volume]""" + payload = { + "data": { + "sh600000": { + "qfqday": [ + ["2026-05-08", "10.00", "10.50", "10.80", "9.90", "12345"], + ["2026-05-09", "10.50", "11.00", "11.20", "10.40", "23456"], + ] + } + } + } + + with patch.object(daily.requests, "get", return_value=_make_response(payload)): + rows = daily._fetch_tencent("600000", "20260508", "20260509") + + assert rows is not None + assert len(rows) == 2 + assert rows[0]["code"] == "600000" + assert rows[0]["date"] == "2026-05-08" + assert rows[0]["open"] == 10.0 + assert rows[0]["close"] == 10.5 # 注意: K 线第三位是 close + assert rows[0]["high"] == 10.8 + assert rows[0]["low"] == 9.9 + assert rows[0]["volume"] == 12345 + # 腾讯不返回这些 + assert rows[0]["pct_change"] is None + assert rows[0]["amplitude"] is None + + +def test_tencent_falls_back_to_day_when_no_qfq(): + """qfqday 缺失时退到 day""" + payload = { + "data": { + "sz000001": { + "day": [ + ["2026-05-08", "10", "10.5", "10.8", "9.9", "100"], + ] + } + } + } + with patch.object(daily.requests, "get", return_value=_make_response(payload)): + rows = daily._fetch_tencent("000001", "20260508", "20260508") + assert rows and rows[0]["close"] == 10.5 + + +def test_tencent_empty_returns_none(): + payload = {"data": {"sh600000": {}}} + with patch.object(daily.requests, "get", return_value=_make_response(payload)): + rows = daily._fetch_tencent("600000", "20260508", "20260508") + assert rows is None + + +def test_tencent_beijing_returns_none_without_request(): + """北交所代码应直接返回 None,不发起请求""" + with patch.object(daily.requests, "get") as mock_get: + rows = daily._fetch_tencent("920001", "20260508", "20260508") + assert rows is None + mock_get.assert_not_called() + + +# ── 东方财富解析 ── + +def test_eastmoney_kline_parsed_correctly(): + """东方财富 klines 每行: date,open,close,high,low,volume,amount,amplitude,pct_change,change,turnover_rate""" + payload = { + "data": { + "klines": [ + "2026-05-08,10.00,10.50,10.80,9.90,12345,123456789,9.0,5.0,0.5,1.2", + ] + } + } + with patch.object(daily.requests, "get", return_value=_make_response(payload)): + rows = daily._fetch_eastmoney("600000", "20260508", "20260508") + + assert rows is not None and len(rows) == 1 + r = rows[0] + assert r["code"] == "600000" + assert r["date"] == "2026-05-08" + assert r["open"] == 10.0 + assert r["close"] == 10.5 + assert r["high"] == 10.8 + assert r["low"] == 9.9 + assert r["volume"] == 12345 + assert r["turnover"] == 123456789 + assert r["amplitude"] == 9.0 + assert r["pct_change"] == 5.0 + assert r["change"] == 0.5 + assert r["turnover_rate"] == 1.2 + + +def test_eastmoney_handles_empty_optional_fields(): + """空字符串字段应转为 None,不应抛 ValueError""" + payload = { + "data": { + "klines": [ + "2026-05-08,10.00,10.50,10.80,9.90,12345,123456789,,,,", + ] + } + } + with patch.object(daily.requests, "get", return_value=_make_response(payload)): + rows = daily._fetch_eastmoney("600000", "20260508", "20260508") + r = rows[0] + assert r["amplitude"] is None + assert r["pct_change"] is None + assert r["change"] is None + assert r["turnover_rate"] is None + + +def test_eastmoney_no_klines_returns_none(): + payload = {"data": {"klines": []}} + with patch.object(daily.requests, "get", return_value=_make_response(payload)): + rows = daily._fetch_eastmoney("600000", "20260508", "20260508") + assert rows is None + + +def test_eastmoney_request_exception_raises_runtime(): + """网络/JSON 异常应包成 RuntimeError,让上层多源切换逻辑能捕获并切下一源""" + import pytest + + def _boom(*_a, **_kw): + raise ConnectionError("network down") + + with patch.object(daily.requests, "get", side_effect=_boom): + with pytest.raises(RuntimeError, match="东方财富请求失败"): + daily._fetch_eastmoney("600000", "20260508", "20260508") diff --git a/tests/test_log.py b/tests/test_log.py new file mode 100644 index 0000000..57caa85 --- /dev/null +++ b/tests/test_log.py @@ -0,0 +1,60 @@ +"""验证 src.log.get_logger 的命名空间和幂等性""" + +import logging + +import pytest + +from src import log as log_mod + + +@pytest.fixture(autouse=True) +def _reset_initialized(): + """每个测试都从未初始化状态开始,避免 handler 残留串扰其他测试""" + saved = log_mod._INITIALIZED + yield + log_mod._INITIALIZED = saved + + +def test_logger_has_ashare_namespace(): + logger = log_mod.get_logger("daily") + assert logger.name == "ashare.daily" + + +def test_init_is_idempotent(): + """重复调用 _init_root 不应叠加 handler""" + log_mod._INITIALIZED = False + root = logging.getLogger("ashare") + root.handlers.clear() + + log_mod._init_root() + first_count = len(root.handlers) + log_mod._init_root() + second_count = len(root.handlers) + + assert first_count == second_count + assert first_count >= 1 # 至少有 stdout handler + + +def test_root_does_not_propagate(): + """ashare logger 不应向 root 冒泡,避免重复输出""" + log_mod._INITIALIZED = False + log_mod._init_root() + root = logging.getLogger("ashare") + assert root.propagate is False + + +def test_log_level_from_env(monkeypatch): + """ASHARE_LOG_LEVEL 环境变量应改变 root 日志级别""" + monkeypatch.setenv("ASHARE_LOG_LEVEL", "DEBUG") + log_mod._INITIALIZED = False + logging.getLogger("ashare").handlers.clear() + log_mod._init_root() + assert logging.getLogger("ashare").level == logging.DEBUG + + +def test_unknown_log_level_falls_back_to_info(monkeypatch): + monkeypatch.setenv("ASHARE_LOG_LEVEL", "NONSENSE") + log_mod._INITIALIZED = False + logging.getLogger("ashare").handlers.clear() + log_mod._init_root() + assert logging.getLogger("ashare").level == logging.INFO diff --git a/tests/test_sector_merge.py b/tests/test_sector_merge.py new file mode 100644 index 0000000..0d10640 --- /dev/null +++ b/tests/test_sector_merge.py @@ -0,0 +1,28 @@ +"""验证 sector.fetch_sector 概念板块抓取逻辑""" + +from unittest.mock import patch, MagicMock + +from src.fetchers import sector + + +def test_concept_skips_when_enough_data(): + """已有 >= 300 个概念时,跳过抓取""" + mock_session = MagicMock() + mock_session.execute.side_effect = [MagicMock(scalar=MagicMock(return_value=150000)), + MagicMock(scalar=MagicMock(return_value=350))] + with patch.object(sector, "get_session", return_value=mock_session), \ + patch.object(sector, "_fetch_concept_list") as mock_list: + sector.fetch_sector() + + mock_list.assert_not_called() + + +def test_concept_fetches_when_insufficient(): + """概念数 < 300 时,执行抓取""" + mock_session = MagicMock() + mock_session.execute.side_effect = [MagicMock(scalar=MagicMock(return_value=0)), + MagicMock(scalar=MagicMock(return_value=0))] + + with patch.object(sector, "get_session", return_value=mock_session), \ + patch.object(sector, "_fetch_concept_list", return_value=[]): + sector.fetch_sector()