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 |
| 分红送转 | 每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/`
+105 -20
View File
@@ -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** | 高 | 自动跑 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,先把现象写进本文件,再开始改代码,避免漏修。
- 新增代码请用 `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/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 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
+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 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]):
+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.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,
)
+31 -31
View File
@@ -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)
+31 -31
View File
@@ -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)
+64 -27
View File
@@ -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))
+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.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,
)
+6 -3
View File
@@ -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):
+37 -144
View File
@@ -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("无概念板块数据")
+15 -4
View File
@@ -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("无股票数据")
+10 -7
View File
@@ -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("无交易日数据")
+19 -18
View File
@@ -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线频率(默认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("--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()
+9 -6
View File
@@ -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 requests8
- `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()