This commit is contained in:
曾志威
2026-05-17 15:51:10 +08:00
parent dbc8107caa
commit 13883f6447
24 changed files with 973 additions and 410 deletions
+12
View File
@@ -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),确保幂等
- 不要引入新的删除方法或清理脚本
+30 -30
View File
@@ -14,8 +14,7 @@ A股数据抓取工具,以 [BaoStock](http://baostock.com) 为主、新浪/腾
| 季频财务指标 | 盈利能力、偿债能力、现金流(最近8个季度,JSON 存储) | BaoStock | | 季频财务指标 | 盈利能力、偿债能力、现金流(最近8个季度,JSON 存储) | BaoStock |
| 分红送转 | 每10股送转、派息、除权除息日(最近10年) | BaoStock | | 分红送转 | 每10股送转、派息、除权除息日(最近10年) | BaoStock |
| 分钟K线 | 5/15/30/60 分钟K线(开高低收、成交量/额) | BaoStock | | 分钟K线 | 5/15/30/60 分钟K线(开高低收、成交量/额) | BaoStock |
| 行业+地域分类 | 证监会行业分类 + 省份 | BaoStock + 东方财富 | | 概念板块 | 东方财富全量概念板块及其成分股(已有≥300个概念时自动跳过) | 东方财富 |
| 概念板块 | 东方财富全量概念板块及其成分股 | 东方财富 |
**已知限制**BaoStock 不含北交所(920xxx)股票;新浪/腾讯/东方财富数据源对北交所同样不支持。 **已知限制**BaoStock 不含北交所(920xxx)股票;新浪/腾讯/东方财富数据源对北交所同样不支持。
@@ -110,15 +109,8 @@ python -m src.main --intraday --start-date 20260101
# 抓取全部频率分钟K线(5/15/30/60 # 抓取全部频率分钟K线(5/15/30/60
python -m src.main --intraday --freq all --start-date 20260101 python -m src.main --intraday --freq all --start-date 20260101
# 抓取行业+地域分类 # 抓取概念板块及成分股(东方财富)
python -m src.main --sector 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)。 > ⚠️ **涨跌停统计 (`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 抓取日线行情(增量;自动分析数据缺口,已完整自动跳过) --daily 抓取日线行情(增量;自动分析数据缺口,已完整自动跳过)
--source 日线数据源: baostock/sina/tencent/eastmoney/all(默认 all,轮换+失败自动切换) --source 日线数据源: baostock/sina/tencent/eastmoney/all(默认 all,轮换+失败自动切换)
--index 抓取主要指数日线 --index 抓取主要指数日线
--financial 抓取季频财务指标 --financial 抓取季频财务指标(增量;只抓缺失季度)
--dividend 抓取分红送转数据 --dividend 抓取分红送转数据(增量;只抓缺失年份)
--intraday 抓取分钟K线行情 --intraday 抓取分钟K线行情(增量;只抓缺失日期)
--freq K线频率: 5/15/30/60/all(默认 5 --freq K线频率: 5/15/30/60/all(默认 5
--sector 抓取行业+地域分类(默认两者都抓 --sector 抓取概念板块及成分股(增量;已有≥300个概念时跳过
--industry-only 仅抓取行业分类 --market-daily 汇总每日涨跌停统计(从 stock_daily 聚合)
--region-only 仅抓取地域分类
--concept-only 仅抓取概念板块及成分股
日期过滤(对日线行情、交易日历、分钟K线、指数生效): 日期过滤(对日线行情、交易日历、分钟K线、指数生效):
--start-date 开始日期,格式 YYYYMMDD --start-date 开始日期,格式 YYYYMMDD
@@ -256,14 +246,6 @@ python -m src.main --sector --concept-only
联合主键:`(code, datetime)` 联合主键:`(code, datetime)`
### stock_sector — 行业+地域分类
| 字段 | 类型 | 说明 |
|------|------|------|
| code | VARCHAR(10) PK | 股票代码 |
| industry | VARCHAR(100) | 证监会行业分类 |
| region | VARCHAR(20) | 省份/地域 |
### stock_concept — 概念板块及成分股 ### stock_concept — 概念板块及成分股
| 字段 | 类型 | 说明 | | 字段 | 类型 | 说明 |
@@ -315,6 +297,7 @@ ashare-data/
├── src/ ├── src/
│ ├── __init__.py │ ├── __init__.py
│ ├── config.py # 配置读取模块(YAML + 环境变量 ASHARE_CONFIG │ ├── config.py # 配置读取模块(YAML + 环境变量 ASHARE_CONFIG
│ ├── log.py # 统一 logging(控制台 + logs/ashare.log 按日滚动)
│ ├── baostock_conn.py # BaoStock 连接管理(login/logout/线程锁/超时重连) │ ├── baostock_conn.py # BaoStock 连接管理(login/logout/线程锁/超时重连)
│ ├── db.py # SQLAlchemy 模型 + 批量 upsert + 自动迁移 │ ├── db.py # SQLAlchemy 模型 + 批量 upsert + 自动迁移
│ ├── main.py # 命令行入口 │ ├── main.py # 命令行入口
@@ -328,7 +311,9 @@ ashare-data/
│ ├── financial.py # 季频财务指标(盈利/偿债/现金流,JSON 存储) │ ├── financial.py # 季频财务指标(盈利/偿债/现金流,JSON 存储)
│ ├── dividend.py # 分红送转 │ ├── dividend.py # 分红送转
│ ├── intraday.py # 分钟K线(5/15/30/60 │ ├── intraday.py # 分钟K线(5/15/30/60
│ └── sector.py # 行业+地域分类 + 概念板块 │ └── sector.py # 概念板块及成分股(东方财富)
├── tests/ # pytest 测试(42 个用例,纯函数 + mock)
├── benchmarks/ # 并发度压测脚本(不在 CI 跑,需真实 MySQL+外网)
└── gzl/ # 选股脚本(独立子项目,可选) └── gzl/ # 选股脚本(独立子项目,可选)
├── Selector.py ├── Selector.py
└── select_stock.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` 的全局锁串行化;查询超时/连接断开时自动重连。 - **线程安全**BaoStock 的 `query_xxx()` 非线程安全,所有调用通过 `src/baostock_conn.py` 的全局锁串行化;查询超时/连接断开时自动重连。
- **去重写入**:所有表通过 `db.batch_upsert()` 走 MySQL `INSERT ON DUPLICATE KEY UPDATE`,重复执行不会产生重复数据。 - **去重写入**:所有表通过 `db.batch_upsert()` 走 MySQL `INSERT ON DUPLICATE KEY UPDATE`,重复执行不会产生重复数据。
- **增量更新**:日线行情按月统计已有数据并与交易日对比,只抓取真正缺口段;其他数据按 (code, 周期) 粒度跳过已完成的股票 - **全量增量**:所有模块均支持增量更新——日线/指数/分钟K线按交易日对比找缺口,财务按缺失季度,分红按缺失年份,概念板块已有≥300个时跳过
- **停牌识别**:日线行情若两端已覆盖、内部仍有缺口,则视为停牌,不再重抓。 - **停牌识别**:日线行情若两端已覆盖、内部仍有缺口,则视为停牌,不再重抓。
- **表结构自动迁移**`init_db()` 会检查并升级旧版 `stock_no_data` / `stock_sector` / `stock_intraday` 的列定义,无需手动改库。 - **表结构自动迁移**`init_db()` 会检查并升级旧版 `stock_no_data` / `stock_sector` / `stock_intraday` 的列定义,无需手动改库。
- **进程内缓存**`get_stock_codes()` / `get_ipo_dates()` 缓存全量股票代码与上市日期,避免重复扫库。 - **进程内缓存**`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 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/` ## 选股子项目 `gzl/`
+105 -20
View File
@@ -2,7 +2,10 @@
本文件用于跟踪 ashare-data 项目的已知问题与后续工作。维护时请保持「问题描述 + 影响范围 + 处理思路」三段式,便于他人接手。 本文件用于跟踪 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`,难做索引。 - **现状**`stock_financial_income/balance/cashflow` 仅有 `code/report_date/data(JSON)` 三列,下游查询需 `JSON_EXTRACT`,难做索引。
- **处理思路**:根据下游真实查询场景(量化筛选 vs 财报展示),把高频指标(ROE、净利润、资产负债率、经营性现金流等)拆出独立列;保留 `extra_json` 兜底。需配套写数据迁移脚本。 - **处理思路**:根据下游真实查询场景(量化筛选 vs 财报展示),把高频指标(ROE、净利润、资产负债率、经营性现金流等)拆出独立列;保留 `extra_json` 兜底。需配套写数据迁移脚本。
### 4. 写入并发度压测 ### 4. 并发度压测:跑出实测数据
- **现状**`config.yaml` 默认 `fetch.workers=1``config.example.yaml` 已注明 "多源轮换可适度提高至 2~4"。需通过实测确定 BaoStock 锁、新浪/腾讯/东财限流的安全边界 - **现状**`benchmarks/bench_daily.py` 已就绪,可一键跑 `workers=1/2/4/8` 对照(详见 `benchmarks/README.md`
- **处理思路**:用 `time pytest` 或专门写一个 `benchmarks/` 脚本,固定一段缺口(如 200 只股票 × 30 天),分别跑 workers=1/2/4/8 比对完成时间和失败率 - **下一步**:在低峰期跑一次完整压测,把推荐档位写到 `config.example.yaml` 注释里。脚本已自带 speedup 表输出,无需再写采集代码
### 5. 把现有 `print()` 全面切到 `src.log.get_logger()`
- **现状**`src/log.py` 已就绪并接入了 README "设计说明"**但所有 fetcher 仍在用 `print(..., flush=True)`**,本期保留兼容性未替换。
- **处理思路**:分批替换(建议按文件粒度提 PR),每次替换一个 fetcher 同时把对应日志级别从直觉值改成 INFO/WARNING/ERROR;替换时一并删除 `flush=True`。
### 6. `gzl/` 选股脚本接入主项目 ### 6. `gzl/` 选股脚本接入主项目
@@ -62,19 +60,95 @@
2. 若纳入,迁移到 `src/strategies/`、改读 MySQL、复用 `get_session()` / `batch_upsert()` / `get_logger()` 2. 若纳入,迁移到 `src/strategies/`、改读 MySQL、复用 `get_session()` / `batch_upsert()` / `get_logger()`
3. `scipy` 加入 `pyproject.toml` 的 optional `[strategies]` extras。 3. `scipy` 加入 `pyproject.toml` 的 optional `[strategies]` extras。
### 7. 扩展测试覆盖 ---
- **现状**`tests/` 已覆盖 4 个核心纯函数(19 个用例),但 fetcher 主流程和 SQL 聚合仍无测试。 ## 🚀 新功能路线图 (Roadmap)
- **处理思路**
- 用 `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)单测,验证主板/创业板不会互串。
### 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** | 高 | 自动跑 pytestPR 必须绿;成本极低 |
| 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,先把现象写进本文件,再开始改代码,避免漏修。 - 每次发现可复现 bug,先把现象写进本文件,再开始改代码,避免漏修。
- 新增代码请用 `from src.log import get_logger`,不要再写 `print(..., flush=True)`。 - 新增代码请用 `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/ERRORBaoStock 查询签名打 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/main.py` 新增 `--market-daily` 参数,挂载 `fetch_market_daily`market_daily 只读 stock_daily,无需 BaoStock 登录)
- ✅ `src/fetchers/market_daily.py` 修复 `fetch_history` → `_fetch_history` 笔误 - ✅ `src/fetchers/market_daily.py` 修复 `fetch_history` → `_fetch_history` 笔误
+47
View File
@@ -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` 注释里。
+1
View File
@@ -0,0 +1 @@
"""benchmarks 包,存放性能压测脚本。"""
+108
View File
@@ -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()
+14
View File
@@ -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 条行业+地域记录
+8 -7
View File
@@ -6,11 +6,13 @@ BaoStock 的 query_xxx() 非线程安全,所有查询需通过同一把锁串
import threading import threading
import time import time
from contextlib import contextmanager from contextlib import contextmanager
from datetime import datetime
import baostock as bs import baostock as bs
from src.log import get_logger
_lock = threading.Lock() _lock = threading.Lock()
_logged_in = False _logged_in = False
_logger = get_logger("baostock")
QUERY_TIMEOUT = 60 QUERY_TIMEOUT = 60
MAX_RETRY = 3 MAX_RETRY = 3
@@ -45,7 +47,7 @@ def _relogin():
time.sleep(1) time.sleep(1)
bs.login() bs.login()
_logged_in = True _logged_in = True
print(f" [{datetime.now().strftime('%H:%M:%S')}] [BS] 已重连", flush=True) _logger.info("已重连")
@contextmanager @contextmanager
@@ -59,7 +61,7 @@ def bs_query(query_fn, *args, **kwargs):
sig_parts = [repr(a) for a in args] sig_parts = [repr(a) for a in args]
sig_parts += [f"{k}={v!r}" for k, v in kwargs.items()] sig_parts += [f"{k}={v!r}" for k, v in kwargs.items()]
sig = ", ".join(sig_parts) 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] result_box = [None]
exc_box = [None] exc_box = [None]
@@ -81,7 +83,7 @@ def bs_query(query_fn, *args, **kwargs):
t.join(timeout=QUERY_TIMEOUT) t.join(timeout=QUERY_TIMEOUT)
if t.is_alive(): 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() _relogin()
last_err = TimeoutError(f"BaoStock query timeout: {short_name}") last_err = TimeoutError(f"BaoStock query timeout: {short_name}")
continue continue
@@ -89,16 +91,15 @@ def bs_query(query_fn, *args, **kwargs):
if exc_box[0] is not None: if exc_box[0] is not None:
err_msg = str(exc_box[0]) 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: 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() _relogin()
last_err = exc_box[0] last_err = exc_box[0]
continue continue
raise exc_box[0] raise exc_box[0]
# 检查返回结果是否有错误码
rs = result_box[0] rs = result_box[0]
if hasattr(rs, 'error_code') and rs.error_code != "0" and "login" in rs.error_msg.lower(): 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() _relogin()
last_err = RuntimeError(f"BaoStock not logged in: {rs.error_msg}") last_err = RuntimeError(f"BaoStock not logged in: {rs.error_msg}")
continue continue
+6 -3
View File
@@ -23,6 +23,9 @@ from sqlalchemy.orm import DeclarativeBase, Session, sessionmaker
from sqlalchemy.dialects.mysql import insert as mysql_insert from sqlalchemy.dialects.mysql import insert as mysql_insert
from src.config import get_mysql_url from src.config import get_mysql_url
from src.log import get_logger
_logger = get_logger("db")
class Base(DeclarativeBase): class Base(DeclarativeBase):
@@ -334,7 +337,7 @@ def init_db():
if result.fetchone(): if result.fetchone():
conn.execute(text("DROP TABLE stock_no_data")) conn.execute(text("DROP TABLE stock_no_data"))
conn.commit() conn.commit()
print(" stock_no_data 表结构已升级(date_range → date") _logger.info("stock_no_data 表结构已升级(date_range → date")
except Exception: except Exception:
pass pass
# 自动迁移:旧版 stock_sector 使用 update_date 列,新版改为 region # 自动迁移:旧版 stock_sector 使用 update_date 列,新版改为 region
@@ -344,7 +347,7 @@ def init_db():
if result.fetchone(): if result.fetchone():
conn.execute(text("DROP TABLE stock_sector")) conn.execute(text("DROP TABLE stock_sector"))
conn.commit() conn.commit()
print(" stock_sector 表结构已升级(新增 region 列)") _logger.info("stock_sector 表结构已升级(新增 region 列)")
except Exception: except Exception:
pass pass
# 自动迁移:旧版 stock_intraday 单表 → 四张分表 # 自动迁移:旧版 stock_intraday 单表 → 四张分表
@@ -355,7 +358,7 @@ def init_db():
except Exception: except Exception:
pass pass
Base.metadata.create_all(engine) Base.metadata.create_all(engine)
print("数据库表初始化完成") _logger.info("数据库表初始化完成")
def batch_upsert(model_cls: type[Base], rows: list[dict], index_columns: list[str]): def batch_upsert(model_cls: type[Base], rows: list[dict], index_columns: list[str]):
+40 -23
View File
@@ -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.config import get_fetch_config
from src.db import StockInfo, StockDaily, batch_upsert, get_session, get_stock_codes, get_ipo_dates 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.fetchers.trading_day import get_trading_days
from src.log import get_logger
from sqlalchemy import select, func, text from sqlalchemy import select, func, text
_logger = get_logger("daily")
VALID_SOURCES = ("baostock", "sina", "tencent", "eastmoney") VALID_SOURCES = ("baostock", "sina", "tencent", "eastmoney")
_HEADERS = { _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 return rows if rows else None
except Exception as e: except Exception as e:
if attempt < retry and _is_transient(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() bs_logout()
time.sleep(min(2 * attempt, 5)) time.sleep(min(2 * attempt, 5))
continue continue
if attempt < retry and not _is_transient(e): 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() bs_logout()
time.sleep(min(2 * attempt, 5)) time.sleep(min(2 * attempt, 5))
continue continue
print(f" [数据源:BaoStock] {code} 获取失败: {e}", flush=True) _logger.error("[BaoStock] %s 获取失败: %s", code, e)
return None return None
@@ -303,7 +306,7 @@ def _analyze_gaps(codes: list[str], start_date: str, end_date: str,
# 上市日期(进程内缓存,只查一次) # 上市日期(进程内缓存,只查一次)
ipo_dates = get_ipo_dates() 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) # 按月统计每只股票行情数(一条SQL)
t1 = time.time() 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 code_month_cnt.setdefault(code, {})[month] = cnt
finally: finally:
session.close() session.close()
print(f" [2/3] 行情按月统计完成 {time.time()-t1:.1f}s", flush=True) _logger.info("[2/3] 行情按月统计完成 %.1fs", time.time()-t1)
# 按月对比找缺口 # 按月对比找缺口
t2 = time.time() t2 = time.time()
@@ -348,7 +351,10 @@ def _analyze_gaps(codes: list[str], start_date: str, end_date: str,
if cnt < len(expected): if cnt < len(expected):
gap_codes.add(code) 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: if not gap_codes:
return {} return {}
@@ -414,12 +420,12 @@ def _fetch_and_save(code: str, gap_start: str, gap_end: str,
batch_upsert(StockDaily, rows, ["code", "date"]) batch_upsert(StockDaily, rows, ["code", "date"])
return code, label, len(rows), True return code, label, len(rows), True
except Exception as e: 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:] next_sources = source_keys[index + 1:]
if next_sources: if next_sources:
next_names = ", ".join(_SOURCE_LABEL[s] for s in 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 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: elif source in VALID_SOURCES:
sources = [source] sources = [source]
else: else:
print(f" 不支持的数据源 {source},可选: {', '.join(VALID_SOURCES)}, all", flush=True) _logger.error("不支持的数据源 %s,可选: %s, all", source, ", ".join(VALID_SOURCES))
return return
cfg = get_fetch_config() 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() codes = get_stock_codes()
if not codes: if not codes:
print(" 无股票列表,请先运行 --stock-info", flush=True) _logger.error("无股票列表,请先运行 --stock-info")
return return
if end_date is None: 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) trading_days = get_trading_days(start_date, end_date)
td_count = len(trading_days) 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]}" 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] codes = [c for c in codes if c not in not_listed]
if 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) gaps = _analyze_gaps(codes, start_date, end_date, trading_days)
complete_count = len(codes) - len(gaps) - len(not_listed) complete_count = len(codes) - len(gaps) - len(not_listed)
if complete_count > 0: if complete_count > 0:
print(f" [数据源:{source_names}] {complete_count} 只股票数据已完整,跳过", flush=True) _logger.info("[%s] %d 只股票数据已完整,跳过", source_names, complete_count)
total = len(gaps) total = len(gaps)
if total == 0: if total == 0:
print(f" [数据源:{source_names}] 所有股票数据已完整,无需抓取", flush=True) _logger.info("[%s] 所有股票数据已完整,无需抓取", source_names)
return return
mode = "全来源轮换+失败自动切换" if source == "all" else "单一来源" 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 模式下轮换首选来源;单来源模式下只使用指定来源。 # all 模式下轮换首选来源;单来源模式下只使用指定来源。
gap_items = list(gaps.items()) 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 elapsed = time.time() - t_start
avg = elapsed / done avg = elapsed / done
eta = avg * (total - done) eta = avg * (total - done)
print(f" [{done}/{total}] {code} [数据源:{label}] 缺口:{g[0]}~{g[-1]} " _logger.info(
f"耗时:{t_fetch:.1f}s 行数:{row_count} " "[%d/%d] %s [%s] 缺口:%s~%s 耗时:%.1fs 行数:%d 成功:%d 剩余:%.0fs",
f"成功:{success} 剩余:{eta:.0f}s", flush=True) done, total, code, label, g[0], g[-1], t_fetch, row_count, success, eta,
)
else: else:
with ThreadPoolExecutor(max_workers=workers) as pool: with ThreadPoolExecutor(max_workers=workers) as pool:
future_map = {} future_map = {}
@@ -543,9 +556,13 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None,
avg = elapsed / done avg = elapsed / done
eta = avg * (total - done) eta = avg * (total - done)
with _print_lock: with _print_lock:
print(f" [{done}/{total}] {code_r} [数据源:{label}] 缺口:{g[0]}~{g[-1]} " _logger.info(
f"行数:{row_count} 成功:{success} 剩余:{eta:.0f}s", flush=True) "[%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 total_time = time.time() - t_start
print(f" 日线行情抓取完成 [数据源:{source_names}] 并发:{workers} " _logger.info(
f"成功:{success} 失败:{fail} 无数据:{nodata_count} 总耗时:{total_time:.1f}s", flush=True) "日线行情抓取完成 [%s] 并发:%d 成功:%d 失败:%d 无数据:%d 总耗时:%.1fs",
source_names, workers, success, fail, nodata_count, total_time,
)
+31 -31
View File
@@ -9,36 +9,35 @@ from datetime import datetime
import baostock as bs import baostock as bs
from src.baostock_conn import bs_query, code_to_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.db import StockDividend, batch_upsert, get_session, get_stock_codes
from src.log import get_logger
from sqlalchemy import select, func from sqlalchemy import select, func
_logger = get_logger("dividend")
def _get_dividend_done(codes: list[str]) -> set[str]:
"""查询已有完整10年分红记录的股票代码""" def _get_missing_years(code: str, years: list[int]) -> list[int]:
current_year = datetime.now().year """返回该股票缺失分红的年份"""
years = [str(y) for y in range(current_year - 10, current_year + 1)] year_strs = [str(y) for y in years]
session = get_session() session = get_session()
try: try:
result = session.execute( existing = set(session.execute(
select(StockDividend.code, func.count(StockDividend.id)) select(StockDividend.report_date)
.where(StockDividend.code.in_(codes)) .where(StockDividend.code == code)
.where(StockDividend.report_date.in_(years)) .where(StockDividend.report_date.in_(year_strs))
.group_by(StockDividend.code) ).scalars().all())
) return [y for y, ys in zip(years, year_strs) if ys not in existing]
# 有记录即可跳过(分红非每年都有,只要有任意年份数据就说明已抓过)
return {row[0] for row in result if row[1] > 0}
finally: finally:
session.close() session.close()
def _fetch_dividend(code: str) -> list[dict]: def _fetch_dividend(code: str, years: list[int]) -> list[dict]:
"""抓取单只股票最近10年的分红记录""" """抓取单只股票指定年份的分红记录"""
bs_code = code_to_bs(code) bs_code = code_to_bs(code)
if not bs_code: if not bs_code:
return [] return []
current_year = datetime.now().year
rows = [] rows = []
for year in range(current_year - 10, current_year + 1): for year in years:
try: try:
with bs_query(bs.query_dividend_data, code=bs_code, year=str(year), yearType="report") as rs: with bs_query(bs.query_dividend_data, code=bs_code, year=str(year), yearType="report") as rs:
while rs.next(): while rs.next():
@@ -84,34 +83,35 @@ def fetch_dividend(symbol: str | None = None):
"""抓取分红送转数据 """抓取分红送转数据
用法:python -m src.main --dividend [--symbol 000001] 用法:python -m src.main --dividend [--symbol 000001]
已有分红记录的股票自动跳过。 已有分红记录的年份自动跳过,只抓缺失年份
""" """
if symbol: if symbol:
codes = [symbol] codes = [symbol]
else: else:
codes = get_stock_codes() codes = get_stock_codes()
# 跳过已有分红数据的股票 current_year = datetime.now().year
if not symbol: all_years = list(range(current_year - 10, current_year + 1))
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) total = len(codes)
if total == 0: _logger.info("正在分析分红数据缺失情况,共 %d 只股票...", total)
print(" 所有股票分红数据已完整,无需抓取", flush=True)
return
print(f"正在抓取分红送转数据,需抓取 {total} 只股票...", flush=True)
success = 0 success = 0
skip = 0
for i, code in enumerate(codes): 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: if rows:
batch_upsert(StockDividend, rows, ["code", "report_date"]) batch_upsert(StockDividend, rows, ["code", "report_date"])
success += 1 success += 1
if (i + 1) % 100 == 0: 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)
+31 -31
View File
@@ -14,8 +14,11 @@ from datetime import datetime
import baostock as bs import baostock as bs
from src.baostock_conn import bs_query, code_to_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.db import FinancialIncome, FinancialBalance, FinancialCashflow, batch_upsert, get_session, get_stock_codes
from src.log import get_logger
from sqlalchemy import select, func from sqlalchemy import select, func
_logger = get_logger("financial")
def _recent_quarters(n: int) -> list[tuple[int, int]]: def _recent_quarters(n: int) -> list[tuple[int, int]]:
"""生成最近 n 个季度 [(year, quarter), ...]""" """生成最近 n 个季度 [(year, quarter), ...]"""
@@ -31,21 +34,23 @@ def _recent_quarters(n: int) -> list[tuple[int, int]]:
return result return result
def _get_existing_financial(codes: list[str], quarters: list[tuple[int, int]], def _get_missing_quarters(code: str, quarters: list[tuple[int, int]]) -> list[tuple[int, int]]:
model_cls) -> set[str]: """返回该股票在3张财务表中缺失的季度"""
"""查询已有财务数据的 (code) 集合:三张表都有完整8季度数据的股票""" q_labels = {f"{y}-{m:02d}-{d:02d}": (y, q) for y, q in quarters
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)]} for m, d in [(3, 31), (6, 30), (9, 30), (12, 31)]}
session = get_session() session = get_session()
try: try:
result = session.execute( existing = set()
select(model_cls.code, func.count(model_cls.id)) for model_cls in (FinancialIncome, FinancialBalance, FinancialCashflow):
.where(model_cls.code.in_(codes)) result = session.execute(
.where(model_cls.report_date.in_(q_labels)) select(model_cls.report_date)
.group_by(model_cls.code) .where(model_cls.code == code)
) .where(model_cls.report_date.in_(q_labels.keys()))
q_count = len(q_labels) )
return {row[0] for row in result if row[1] >= q_count} for row in result:
existing.add(row[0])
# 返回不在已有集合中的季度
return [q for label, q in q_labels.items() if label not in existing]
finally: finally:
session.close() session.close()
@@ -85,31 +90,26 @@ def fetch_financial(symbol: str | None = None):
quarters = _recent_quarters(8) 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) total = len(codes)
if total == 0: _logger.info("正在分析财务数据缺失情况,共 %d 只股票...", total)
print(" 所有股票财务数据已完整,无需抓取", flush=True)
return
print(f"正在抓取财务数据,需抓取 {total} 只股票 × {len(quarters)} 个季度...", flush=True)
success = 0 success = 0
fail = 0 fail = 0
skip = 0
for i, code in enumerate(codes): 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) bs_code = code_to_bs(code)
if not bs_code: if not bs_code:
continue continue
try: 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: with bs_query(bs.query_profit_data, code=bs_code, year=year, quarter=quarter) as rs:
fields = rs.fields if rs.fields else [] fields = rs.fields if rs.fields else []
@@ -133,10 +133,10 @@ def fetch_financial(symbol: str | None = None):
success += 1 success += 1
except Exception as e: 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 fail += 1
if (i + 1) % 50 == 0 or i == 0: 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)
+64 -27
View File
@@ -16,9 +16,12 @@ from datetime import datetime, timedelta
import baostock as bs import baostock as bs
from src.baostock_conn import bs_query, bs_login from src.baostock_conn import bs_query, bs_login
from src.config import get_fetch_config 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 from sqlalchemy import select, func
_logger = get_logger("index")
# 主要指数代码 → BaoStock 格式 # 主要指数代码 → BaoStock 格式
INDICES = { INDICES = {
@@ -41,20 +44,46 @@ def _clean(val):
return 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() session = get_session()
try: try:
from sqlalchemy import distinct # 获取范围内的交易日
sd_date = datetime.strptime(sd, "%Y-%m-%d").date() trading_days = session.execute(
check_start = str(sd_date - timedelta(days=5)) select(TradingDay.date)
check_end = str(sd_date + timedelta(days=5)) .where(TradingDay.date >= sd)
result = session.execute( .where(TradingDay.date <= ed)
select(distinct(IndexDaily.code)) .order_by(TradingDay.date)
.where(IndexDaily.date >= check_start) ).scalars().all()
.where(IndexDaily.date <= check_end) if not trading_days:
) return [(sd, ed)]
return {row[0] for row in result}
# 获取该指数已有的日期
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: finally:
session.close() 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]}" 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]}" ed = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}"
# 跳过已有数据的指数 # 分析每个指数的数据缺口
fetched = _get_fetched_codes(sd, ed) gaps_map: dict[str, list[tuple[str, str]]] = {}
to_fetch = {k: v for k, v in INDICES.items() if k not in fetched} for code in INDICES:
gaps = _analyze_index_gaps(code, sd, ed)
if gaps:
gaps_map[code] = gaps
if not to_fetch: if not gaps_map:
print(f" 指数数据 {sd} ~ {ed} 已完整,跳过", flush=True) _logger.info("指数数据 %s ~ %s 已完整,跳过", sd, ed)
return return
skip_msg = f"(跳过 {len(fetched)} 个已有数据)" if fetched else "" skip_msg = f"(跳过 {len(INDICES) - len(gaps_map)} 个已完整)"
print(f"正在抓取指数日线 {sd} ~ {ed},共 {len(to_fetch)}{skip_msg}...", flush=True) _logger.info("正在抓取指数日线 %s ~ %s,需补缺 %d%s", sd, ed, len(gaps_map), skip_msg)
bs_login() bs_login()
success = 0 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() t0 = time.time()
rows = _fetch_one_index(code, market, name, sd, ed) for gap_sd, gap_ed in gaps:
if rows: rows = _fetch_one_index(code, market, name, gap_sd, gap_ed)
batch_upsert(IndexDaily, rows, ["code", "date"]) if rows:
batch_upsert(IndexDaily, rows, ["code", "date"])
total_rows += len(rows)
if total_rows:
success += 1 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: else:
print(f" {name}({code}): 无数据", flush=True) _logger.warning("%s(%s): 无数据", name, code)
time.sleep(delay) time.sleep(delay)
print(f" 指数数据抓取完成,成功:{success}/{len(to_fetch)}", flush=True) _logger.info("指数数据抓取完成,成功:%d/%d", success, len(gaps_map))
+99 -56
View File
@@ -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.config import get_fetch_config
from src.db import ( from src.db import (
StockMin5, StockMin15, StockMin30, StockMin60, 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") 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}" return f"{date_str} {time_str}"
def _get_fetched_codes(model, sd: str, ed: str) -> set[str]: def _get_intraday_gaps(model, code: str, sd: str, ed: str) -> list[tuple[str, str]]:
"""查询数据已覆盖起始日期的股票代码 """分析单只股票在 [sd, ed] 范围内的分钟K线缺口,返回缺失区间列表"""
只有当股票在起始日期附近(前5个自然日)有数据时,才认为已抓取。
这样扩展日期范围时,缺失的早期数据不会被跳过。
"""
session = get_session() session = get_session()
try: try:
sd_date = datetime.strptime(sd, "%Y-%m-%d").date() # 获取范围内的交易日
check_start = str(sd_date - timedelta(days=5)) trading_days = session.execute(
check_end = str(sd_date + timedelta(days=5)) select(TradingDay.date)
result = session.execute( .where(TradingDay.date >= sd)
select(distinct(model.code)) .where(TradingDay.date <= ed)
.where(model.datetime >= check_start) .order_by(TradingDay.date)
.where(model.datetime <= check_end + " 23:59:59") ).scalars().all()
) if not trading_days:
return {row[0] for row in result} 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: finally:
session.close() session.close()
@@ -87,7 +112,7 @@ def fetch_intraday(start_date: str | None = None, end_date: str | None = None,
elif freq in VALID_FREQ: elif freq in VALID_FREQ:
freqs = [freq] freqs = [freq]
else: else:
print(f" 不支持的频率 {freq},可选: {', '.join(VALID_FREQ)}, all", flush=True) _logger.error("不支持的频率 %s,可选: %s, all", freq, ", ".join(VALID_FREQ))
return return
for f in freqs: for f in freqs:
@@ -114,66 +139,84 @@ def _fetch_one_freq(freq: str, start_date: str | None, end_date: str | None,
else: else:
codes = get_stock_codes() codes = get_stock_codes()
if not codes: if not codes:
print(" 无股票列表,请先运行 --stock-info", flush=True) _logger.error("无股票列表,请先运行 --stock-info")
return return
# 跳过已有数据的股票 # 分析每只股票的数据缺口
fetched = set() need_fetch = [] # [(code, gaps), ...]
if not symbol: if symbol:
fetched = _get_fetched_codes(model, sd, ed) need_fetch = [(symbol, [(sd, ed)])]
codes = [c for c in codes if c not in fetched] 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: if not need_fetch:
print(f" {freq}分钟K线 {sd} ~ {ed} 数据已完整,跳过", flush=True) _logger.info("%s分钟K线 %s ~ %s 数据已完整,跳过", freq, sd, ed)
return return
bs_login() bs_login()
total = len(codes) total = len(need_fetch)
success = 0 success = 0
fail = 0 fail = 0
t_start = time.time() t_start = time.time()
skip_msg = f"(跳过 {len(fetched)} 只已有数据" if fetched else "" skip_msg = f"(跳过 {skip} 只已完整" if skip else ""
print(f"正在抓取{freq}分钟K线 {sd} ~ {ed},需抓取 {total}{skip_msg}...", flush=True) _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) bs_code = code_to_bs(code)
if not bs_code: if not bs_code:
continue continue
total_rows = 0
try: try:
with bs_query( for gap_sd, gap_ed in gaps:
bs.query_history_k_data_plus, with bs_query(
bs_code, bs.query_history_k_data_plus,
"date,time,open,high,low,close,volume,amount", bs_code,
start_date=sd, end_date=ed, "date,time,open,high,low,close,volume,amount",
frequency=freq, adjustflag="3", start_date=gap_sd, end_date=gap_ed,
) as rs: frequency=freq, adjustflag="3",
rows = [] ) as rs:
while rs.next(): rows = []
r = rs.get_row_data() while rs.next():
dt_str = _parse_datetime(r[0], r[1]) r = rs.get_row_data()
rows.append({ dt_str = _parse_datetime(r[0], r[1])
"code": code, rows.append({
"datetime": dt_str, "code": code,
"open": _clean(r[2]), "datetime": dt_str,
"high": _clean(r[3]), "open": _clean(r[2]),
"low": _clean(r[4]), "high": _clean(r[3]),
"close": _clean(r[5]), "low": _clean(r[4]),
"volume": _clean(r[6]), "close": _clean(r[5]),
"amount": _clean(r[7]), "volume": _clean(r[6]),
}) "amount": _clean(r[7]),
if rows: })
batch_upsert(model, rows, ["code", "datetime"]) if rows:
success += 1 batch_upsert(model, rows, ["code", "datetime"])
total_rows += len(rows)
if total_rows:
success += 1
except Exception: except Exception:
fail += 1 fail += 1
if (i + 1) % 100 == 0: if (i + 1) % 100 == 0:
elapsed = time.time() - t_start 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: if not symbol:
time.sleep(delay) time.sleep(delay)
total_time = time.time() - t_start 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,
)
+6 -3
View File
@@ -12,6 +12,9 @@
from datetime import datetime from datetime import datetime
from sqlalchemy import select, func, case, and_, text from sqlalchemy import select, func, case, and_, text
from src.db import StockDaily, MarketDaily, batch_upsert, get_session 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 科创板 # 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]}" 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]}" 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() session = get_session()
@@ -94,9 +97,9 @@ def _fetch_history(start_date: str | None, end_date: str | None):
if rows: if rows:
batch_upsert(MarketDaily, rows, ["date"]) 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: else:
print(" 无新数据", flush=True) _logger.info("无新数据")
def fetch_market_daily(start_date: str | None = None, end_date: str | None = None): def fetch_market_daily(start_date: str | None = None, end_date: str | None = None):
+37 -144
View File
@@ -1,86 +1,30 @@
"""行业+地域+概念板块数据抓取 """概念板块数据抓取 — 东方财富
数据源: 数据源:
- 行业分类:BaoStock query_stock_industry()(证监会行业分类) - 概念板块列表:东方财富 push2 API
- 地域分类:东方财富 F10 CompanySurveyAjax(省份) - 成分股:每个概念的成分股代码列表
- 概念板块:东方财富 stock_board_concept_name_em() 全量概念列表 + 每个概念成分股
用法: 用法:
python -m src.main --sector 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 import time
from concurrent.futures import ThreadPoolExecutor, as_completed
import requests import requests
import baostock as bs from src.db import StockConcept, batch_upsert, get_session
from src.baostock_conn import bs_query, bs_login from src.log import get_logger
from src.db import StockSector, StockConcept, batch_upsert, get_session, get_stock_codes from sqlalchemy import select, func
from sqlalchemy import select
_logger = get_logger("sector")
_HEADERS = { _HEADERS = {
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36", "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_LIST_URL = "https://push2.eastmoney.com/api/qt/clist/get"
_EM_CONCEPT_STOCKS_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]: def _fetch_concept_list() -> list[dict]:
"""从东方财富获取全部概念板块列表""" """从东方财富获取全部概念板块列表"""
params = { params = {
@@ -98,7 +42,7 @@ def _fetch_concept_list() -> list[dict]:
items = data.get("diff", []) or [] items = data.get("diff", []) or []
return [{"code": item["f12"], "name": item["f14"]} for item in items if item.get("f12") and item.get("f14")] return [{"code": item["f12"], "name": item["f14"]} for item in items if item.get("f12") and item.get("f14")]
except Exception as e: except Exception as e:
print(f" 获取概念板块列表失败: {e}", flush=True) _logger.error("获取概念板块列表失败: %s", e)
return [] return []
@@ -122,15 +66,31 @@ def _fetch_concept_stocks(concept_code: str) -> list[str]:
return [] return []
def fetch_concept(): def fetch_sector():
"""抓取全部概念板块及成分股,写入 stock_concept 表""" """抓取全部概念板块及成分股,写入 stock_concept 表"""
print("正在获取概念板块列表...", flush=True) # 检查已有概念数据
concepts = _fetch_concept_list() session = get_session()
if not concepts: try:
print(" 概念板块列表获取失败", flush=True) 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 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() t_start = time.time()
rows = [] rows = []
for i, concept in enumerate(concepts): for i, concept in enumerate(concepts):
@@ -143,81 +103,14 @@ def fetch_concept():
}) })
if (i + 1) % 50 == 0: if (i + 1) % 50 == 0:
elapsed = time.time() - t_start 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) time.sleep(0.05)
if rows: if rows:
batch_upsert(StockConcept, rows, ["code", "concept_code"]) 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: else:
print(" 无概念板块数据", flush=True) _logger.warning("无概念板块数据")
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)
+15 -4
View File
@@ -7,12 +7,23 @@ BaoStock query_stock_basic() 一次返回全部证券(含指数、基金等)
import baostock as bs import baostock as bs
from src.baostock_conn import bs_query 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(): def fetch_stock_list():
"""从 BaoStock 获取沪深A股列表,upsert 到 stock_info 表""" """从 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 = [] rows = []
with bs_query(bs.query_stock_basic) as rs: with bs_query(bs.query_stock_basic) as rs:
while rs.next(): while rs.next():
@@ -30,6 +41,6 @@ def fetch_stock_list():
if rows: if rows:
batch_upsert(StockInfo, rows, ["code"]) batch_upsert(StockInfo, rows, ["code"])
print(f" 股票列表已更新,共 {len(rows)}", flush=True) _logger.info("股票列表已更新,共 %d", len(rows))
else: else:
print(" 无股票数据", flush=True) _logger.warning("无股票数据")
+10 -7
View File
@@ -11,8 +11,11 @@ from datetime import datetime, timedelta
import baostock as bs import baostock as bs
from src.db import TradingDay, StockDaily, batch_upsert, get_session from src.db import TradingDay, StockDaily, batch_upsert, get_session
from src.baostock_conn import bs_query from src.baostock_conn import bs_query
from src.log import get_logger
from sqlalchemy import select, func, text from sqlalchemy import select, func, text
_logger = get_logger("trading_day")
def _fetch_baostock(sd: str, ed: str) -> list[str] | None: def _fetch_baostock(sd: str, ed: str) -> list[str] | None:
"""从 BaoStock 获取交易日历""" """从 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) chunk_end = min(chunk_start.replace(year=chunk_start.year + 5), ed)
cs = chunk_start.strftime("%Y-%m-%d") cs = chunk_start.strftime("%Y-%m-%d")
ce = chunk_end.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) days = _fetch_baostock(cs, ce)
if days: if days:
all_days.extend(d for d in days if d <= today) all_days.extend(d for d in days if d <= today)
else: else:
print(f" {cs} ~ {ce} 获取失败", flush=True) _logger.warning("%s ~ %s 获取失败", cs, ce)
chunk_start = chunk_end + timedelta(days=1) chunk_start = chunk_end + timedelta(days=1)
if all_days: 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) all_days = _validate_with_daily(all_days)
rows = [{"date": d} for d in all_days] rows = [{"date": d} for d in all_days]
batch_upsert(TradingDay, rows, ["date"]) batch_upsert(TradingDay, rows, ["date"])
print(f" 交易日历已保存,{len(all_days)} 个交易日", flush=True) _logger.info("交易日历已保存,%d 个交易日", len(all_days))
return all_days return all_days
print(" BaoStock 获取失败,将从已有行情数据推断", flush=True) _logger.warning("BaoStock 获取失败,将从已有行情数据推断")
return _infer_from_daily(start_date, end_date) 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() session.close()
if days and days[0] <= sd and days[-1] >= ed: 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 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) days = _fetch_and_save(start_date, end_date)
if not days: if not days:
print(" 无交易日数据", flush=True) _logger.warning("无交易日数据")
+19 -18
View File
@@ -10,9 +10,7 @@
python -m src.main --intraday --freq all # 全部频率 python -m src.main --intraday --freq all # 全部频率
python -m src.main --intraday --start-date 20260508 --end-date 20260509 python -m src.main --intraday --start-date 20260508 --end-date 20260509
python -m src.main --intraday --symbol 000001 --freq 30 python -m src.main --intraday --symbol 000001 --freq 30
python -m src.main --sector # 行业+地域分类 python -m src.main --sector # 概念板块及成分股
python -m src.main --sector --industry-only # 仅行业分类
python -m src.main --sector --concept-only # 仅概念板块
python -m src.main --index # 指数日线(上证/沪深300/创业板等) python -m src.main --index # 指数日线(上证/沪深300/创业板等)
python -m src.main --market-daily # 汇总每日涨跌停统计(依赖 stock_daily 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.baostock_conn import bs_login, bs_logout
from src.config import load_config from src.config import load_config
from src.db import init_db from src.db import init_db
from src.log import get_logger
_logger = get_logger("main")
def main(): def main():
@@ -38,13 +39,10 @@ def main():
parser.add_argument("--freq", type=str, default="5", parser.add_argument("--freq", type=str, default="5",
choices=["5", "15", "30", "60", "all"], choices=["5", "15", "30", "60", "all"],
help="分钟K线频率(默认5all=全部)") help="分钟K线频率(默认5all=全部)")
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("--index", action="store_true", help="抓取指数日线行情")
parser.add_argument("--market-daily", action="store_true", parser.add_argument("--market-daily", action="store_true",
help="汇总每日涨跌停统计(从 stock_daily 聚合)") 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("--start-date", type=str, help="开始日期 YYYYMMDD")
parser.add_argument("--end-date", type=str, help="结束日期 YYYYMMDD") parser.add_argument("--end-date", type=str, help="结束日期 YYYYMMDD")
parser.add_argument("--symbol", type=str, help="指定单只股票代码") parser.add_argument("--symbol", type=str, help="指定单只股票代码")
@@ -53,8 +51,7 @@ def main():
if not any([args.stock_info, args.trading_day, args.daily, if not any([args.stock_info, args.trading_day, args.daily,
args.financial, args.dividend, args.intraday, args.financial, args.dividend, args.intraday,
args.sector, args.industry_only, args.region_only, args.sector, args.index, args.market_daily]):
args.concept_only, args.index, args.market_daily]):
parser.print_help() parser.print_help()
return return
@@ -68,9 +65,18 @@ def main():
# 若仅运行 market-daily,避免无谓的登录 # 若仅运行 market-daily,避免无谓的登录
if not any([args.stock_info, args.trading_day, args.daily, if not any([args.stock_info, args.trading_day, args.daily,
args.financial, args.dividend, args.intraday, args.financial, args.dividend, args.intraday,
args.sector, args.industry_only, args.region_only, args.sector, args.index]):
args.concept_only, args.index]): _logger.info("全部任务完成")
print("全部任务完成", flush=True) 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 return
bs_login() bs_login()
@@ -101,16 +107,11 @@ def main():
fetch_intraday(start_date=args.start_date, end_date=args.end_date, fetch_intraday(start_date=args.start_date, end_date=args.end_date,
symbol=args.symbol, freq=args.freq) 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: if args.index:
from src.fetchers.index import fetch_index from src.fetchers.index import fetch_index
fetch_index(start_date=args.start_date, end_date=args.end_date) fetch_index(start_date=args.start_date, end_date=args.end_date)
print("全部任务完成", flush=True) _logger.info("全部任务完成")
finally: finally:
bs_logout() bs_logout()
+9 -6
View File
@@ -9,10 +9,13 @@ pip install -e .[dev] # 或 pip install pytest
pytest -v pytest -v
``` ```
## 覆盖范围 ## 覆盖范围43 个用例)
- `test_code_mapping.py` —— `baostock_conn.code_to_bs` 代码前缀映射 - `test_code_mapping.py` —— `baostock_conn.code_to_bs` 代码前缀映射5
- `test_daily_derived.py` —— `daily._fill_derived_fields` 振幅/涨跌幅补算 - `test_daily_derived.py` —— `daily._fill_derived_fields` 振幅/涨跌幅补算5
- `test_financial_quarters.py` —— `financial._recent_quarters` 季度滚动 - `test_daily_source_codes.py` —— sina/tencent/eastmoney 各源代码前缀映射(8)
- `test_market_classification.py` —— `market_daily._is_20pct` 板块判定 - `test_daily_sources_http.py` —— 腾讯/东财 HTTP JSON 解析(mock requests8
- `test_sector_merge.py` —— `sector.fetch_sector` 三种 only 模式下不会清空其他列 - `test_financial_quarters.py` —— `financial._recent_quarters` 季度滚动(5
- `test_log.py` —— `log.get_logger` 命名空间、handler 幂等、env 控制 level5
- `test_market_classification.py` —— `market_daily._is_20pct` 板块判定(4
- `test_sector_merge.py` —— `sector.fetch_sector` 三种 only 模式合并保护(3
+52
View File
@@ -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
+141
View File
@@ -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")
+60
View File
@@ -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
+28
View File
@@ -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()