SHA256
更新
This commit is contained in:
@@ -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),确保幂等
|
||||
- 不要引入新的删除方法或清理脚本
|
||||
@@ -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/`
|
||||
|
||||
|
||||
@@ -2,7 +2,10 @@
|
||||
|
||||
本文件用于跟踪 ashare-data 项目的已知问题与后续工作。维护时请保持「问题描述 + 影响范围 + 处理思路」三段式,便于他人接手。
|
||||
|
||||
> 🗓️ **2026-05-15 一次性清理**:原 P0 / P1 / P2 中可在不破坏 git 历史的前提下完成的条目已落地(详见底部 [变更日志](#变更日志))。当前剩余条目均为:①需要用户决策的高风险动作,②长期工程任务,③依赖外部数据源调研。
|
||||
> 🗓️ **2026-05-15 二轮清理**:上一轮(也是 2026-05-15)清理后剩下的 P2 中,可独立完成的 #4/#5/#7/#8 已落地(详见底部 [变更日志](#变更日志))。
|
||||
> 当前剩余条目均为:①需要用户决策的高风险动作(P0),②依赖外部数据源调研(P2-北交所),③需要与主项目协调的可选合并(P2-财务结构化、P2-gzl 接入)。
|
||||
|
||||
> 📚 **本文件分两部分**:上半部分是**已知缺陷/技术债**(P0/P2),下半部分是**新功能路线图**([Roadmap](#-新功能路线图-roadmap))。前者修,后者建。
|
||||
|
||||
---
|
||||
|
||||
@@ -44,15 +47,10 @@
|
||||
- **现状**:`stock_financial_income/balance/cashflow` 仅有 `code/report_date/data(JSON)` 三列,下游查询需 `JSON_EXTRACT`,难做索引。
|
||||
- **处理思路**:根据下游真实查询场景(量化筛选 vs 财报展示),把高频指标(ROE、净利润、资产负债率、经营性现金流等)拆出独立列;保留 `extra_json` 兜底。需配套写数据迁移脚本。
|
||||
|
||||
### 4. 写入并发度压测
|
||||
### 4. 并发度压测:跑出实测数据
|
||||
|
||||
- **现状**:`config.yaml` 默认 `fetch.workers=1`;`config.example.yaml` 已注明 "多源轮换可适度提高至 2~4"。需通过实测确定 BaoStock 锁、新浪/腾讯/东财限流的安全边界。
|
||||
- **处理思路**:用 `time pytest` 或专门写一个 `benchmarks/` 脚本,固定一段缺口(如 200 只股票 × 30 天),分别跑 workers=1/2/4/8 比对完成时间和失败率。
|
||||
|
||||
### 5. 把现有 `print()` 全面切到 `src.log.get_logger()`
|
||||
|
||||
- **现状**:`src/log.py` 已就绪并接入了 README "设计说明",**但所有 fetcher 仍在用 `print(..., flush=True)`**,本期保留兼容性未替换。
|
||||
- **处理思路**:分批替换(建议按文件粒度提 PR),每次替换一个 fetcher 同时把对应日志级别从直觉值改成 INFO/WARNING/ERROR;替换时一并删除 `flush=True`。
|
||||
- **现状**:`benchmarks/bench_daily.py` 已就绪,可一键跑 `workers=1/2/4/8` 对照(详见 `benchmarks/README.md`)。
|
||||
- **下一步**:在低峰期跑一次完整压测,把推荐档位写到 `config.example.yaml` 注释里。脚本已自带 speedup 表输出,无需再写采集代码。
|
||||
|
||||
### 6. `gzl/` 选股脚本接入主项目
|
||||
|
||||
@@ -62,19 +60,95 @@
|
||||
2. 若纳入,迁移到 `src/strategies/`、改读 MySQL、复用 `get_session()` / `batch_upsert()` / `get_logger()`;
|
||||
3. `scipy` 加入 `pyproject.toml` 的 optional `[strategies]` extras。
|
||||
|
||||
### 7. 扩展测试覆盖
|
||||
---
|
||||
|
||||
- **现状**:`tests/` 已覆盖 4 个核心纯函数(19 个用例),但 fetcher 主流程和 SQL 聚合仍无测试。
|
||||
- **处理思路**:
|
||||
- 用 `sqlite::memory:` 跑一遍 `db.init_db()` + `batch_upsert()` 端到端;
|
||||
- 用 `responses` 库 mock 新浪/腾讯/东方财富 HTTP 接口,覆盖 `_fetch_sina/_fetch_tencent/_fetch_eastmoney`;
|
||||
- 用 `freezegun` 替换手写 `FakeDatetime`;
|
||||
- `market_daily._fetch_history` 的 SQL 阈值(涨停 ≥9.8 / ≥19.5)单测,验证主板/创业板不会互串。
|
||||
## 🚀 新功能路线图 (Roadmap)
|
||||
|
||||
### 8. `sector.fetch_sector` 三种 only 模式的合并保护测试
|
||||
> 这一节是「想做但还没排期」的需求池,与上面的 P0/P2(已知缺陷/技术债)分开维护。
|
||||
> 立项时把对应条目挪到 P1/P2,附上责任人和预计动手时间;落地后再移到 [变更日志](#变更日志)。
|
||||
|
||||
- **现状**:`region_only=True` 时使用 `existing.get("industry")` 保留旧行业值;`industry_only=True` 反之。逻辑正确但无测试。
|
||||
- **处理思路**:归并到上面第 7 项一起做(需 mock DB session 或用 SQLite)。
|
||||
### ⭐ 推荐下一迭代(按"价值高 / 成本可控 / 与现有架构契合"排序)
|
||||
|
||||
1. **抓取调度 + 告警**(详见下方「四、调度与监控」#1+#2)— 让项目从工具变服务,半天工作量
|
||||
2. **HTTP API 服务**(「三、服务化」#1)— FastAPI 暴露查询接口,半到一天
|
||||
3. **资金面三件套:龙虎榜 / 北向资金 / 融资融券**(「一、数据维度扩展」前 3 条)— 接口现成、量小,2-3 天补齐情绪+资金维度
|
||||
|
||||
---
|
||||
|
||||
### 一、数据维度扩展
|
||||
|
||||
「价值」= 对量化/选股的直接增益,「成本」= 实现规模 + 外部依赖复杂度。
|
||||
|
||||
| # | 需求 | 价值 | 成本 | 关键说明 |
|
||||
|---|---|---|---|---|
|
||||
| 1 | **龙虎榜** | 高 | 中 | 东财/同花顺接口稳定;游资/机构席位是短线核心信号 |
|
||||
| 2 | **北向资金(陆股通)持股明细** | 高 | 低 | 港交所/东财 T+1 披露;新增 `hk_holdings` 表 |
|
||||
| 3 | **融资融券余额** | 高 | 低 | 流动性/情绪指标,东财/交易所每日发布 |
|
||||
| 4 | **业绩预告 / 快报** | 高 | 中 | 早于正式财报,常含异常波动信息;BaoStock 无,需走东财/同花顺 |
|
||||
| 5 | **限售解禁日历** | 中 | 低 | 解禁前后股价波动显著;东财日历接口 |
|
||||
| 6 | **股东户数** | 中 | 低 | 季频,反映筹码集中度;BaoStock `query_stock_other_basic_info` |
|
||||
| 7 | **十大流通股东** | 中 | 中 | 季频跟踪机构持仓;BaoStock 有现成接口 |
|
||||
| 8 | **大宗交易** | 中 | 中 | 折溢价 + 营业部,事件驱动 |
|
||||
| 9 | **ST 标记历史** | 中 | 中 | 当前无连续 ST 状态记录,无法回测「摘帽行情」 |
|
||||
| 10 | **IPO / 定增 / 可转债日历** | 中 | 中 | 一级市场事件 |
|
||||
| 11 | **ETF 行情 + 折溢价** | 中 | 中 | 套利策略基础数据 |
|
||||
| 12 | **股指期货 IF/IH/IC/IM** | 中 | 中 | 对冲/基差研究 |
|
||||
| 13 | **期权行情**(50/300/500ETF) | 中 | 高 | 波动率研究 |
|
||||
| 14 | **L1 Tick 行情** | 高 | 极高 | 数据量爆炸(GB/日),需切 ClickHouse / Parquet |
|
||||
| 15 | **公司公告全文** | 高 | 高 | 需 PDF/HTML 解析 + 全文检索(ES) |
|
||||
| 16 | **研报 / 一致预期** | 高 | 高 | 多家券商接口闭源,合规风险 |
|
||||
| 17 | **北交所(920xxx)** | 中 | 中 | 已在上方 P2-#2 单列 |
|
||||
|
||||
### 二、数据加工层(让"数据"变"信号")
|
||||
|
||||
| # | 需求 | 价值 | 关键说明 |
|
||||
|---|---|---|---|
|
||||
| 1 | **后复权日线 + 周/月线聚合表** | 高 | 现仅有前复权;后复权用于长期收益对比 |
|
||||
| 2 | **技术指标预计算**(MA/MACD/RSI/BOLL/KDJ) | 中 | 一次算多次用 |
|
||||
| 3 | **因子库**(动量/反转/价值/质量/规模/波动率) | 高 | 量化必备;每日预计算入 `factor_daily` 宽表 |
|
||||
| 4 | **数据质量校验日报** | 高 | 每日跑:股票数、缺口、异常涨跌幅、停牌识别;失败发告警 |
|
||||
| 5 | **多源交叉校验** | 中 | BaoStock vs 新浪同日 close 偏差 >1% 自动报警 |
|
||||
|
||||
### 三、服务化(让数据被消费)
|
||||
|
||||
| # | 需求 | 价值 | 关键说明 |
|
||||
|---|---|---|---|
|
||||
| 1 | **HTTP API 服务**(FastAPI) | 高 | 暴露 `/daily` `/financial` `/sector` 等 REST;前端/其他系统可直接消费 |
|
||||
| 2 | **Python SDK 包装** | 中 | `from ashare import get_daily`,屏蔽 SQL |
|
||||
| 3 | **Parquet / Feather 导出** | 中 | 量化研究跑数据更快;增量导出到本地或 S3/OSS |
|
||||
| 4 | **Kafka / ClickHouse 同步** | 中 | 接入下游量化平台时再考虑 |
|
||||
| 5 | **CLI 查询子命令** | 低 | `ashare query --code 600000 --metric pe-ttm` |
|
||||
|
||||
### 四、调度与监控(手动 → 自动)
|
||||
|
||||
| # | 需求 | 价值 | 关键说明 |
|
||||
|---|---|---|---|
|
||||
| 1 | **抓取任务调度**(APScheduler 或 cron + systemd) | 高 | 当前依赖手动 `python -m src.main --daily`;自动化后无人值守 |
|
||||
| 2 | **失败告警**(企微/钉钉/邮件 webhook) | 高 | 配合 #1;失败/延迟超阈值即推送 |
|
||||
| 3 | **Prometheus metrics 导出** | 中 | 各 fetcher 耗时/成功率/失败码;接 Grafana |
|
||||
| 4 | **数据完整性 dashboard**(Grafana / Superset) | 中 | 直观看股票覆盖、缺口、最新数据日期 |
|
||||
|
||||
### 五、选股 / 策略(承接 gzl)
|
||||
|
||||
| # | 需求 | 价值 | 关键说明 |
|
||||
|---|---|---|---|
|
||||
| 1 | **gzl 接入主项目** | 中 | P2-#6 已列;改读 MySQL,统一日志/连接池 |
|
||||
| 2 | **策略插件框架** | 高 | `src/strategies/` 下每策略一文件,统一 `run(date) -> List[Signal]` 接口 |
|
||||
| 3 | **简单回测引擎** | 高 | 基于已有日线表,单策略 N 年回测,输出收益/最大回撤/胜率 |
|
||||
| 4 | **选股信号定时输出** | 中 | 每日盘后跑所有策略,结果入 `signal_daily` 表或推送 |
|
||||
| 5 | **事件驱动信号库** | 中 | 涨停回封、放量突破、底背离、机构席位上榜 |
|
||||
|
||||
### 六、工程基础(质量护栏)
|
||||
|
||||
| # | 需求 | 价值 | 关键说明 |
|
||||
|---|---|---|---|
|
||||
| 1 | **GitHub Actions CI** | 高 | 自动跑 pytest,PR 必须绿;成本极低 |
|
||||
| 2 | **Alembic 数据库迁移** | 中 | 替代 `db.py` 中手写的 `DROP TABLE`+`create_all`,版本可控 |
|
||||
| 3 | **Docker + docker-compose** | 中 | 自带 MySQL,新机器一行命令起;适合给协作者 |
|
||||
| 4 | **类型注解全量 + mypy strict** | 中 | 现部分函数已有;走全量后 IDE/重构体验显著提升 |
|
||||
| 5 | **ruff / pre-commit hook** | 低 | 统一格式;低争议低成本 |
|
||||
| 6 | **PostgreSQL / SQLite 后端兼容** | 中 | 现 `batch_upsert` 写死 MySQL 方言;抽象后可本地 SQLite 跑端到端测试 |
|
||||
| 7 | **Web 控制台**(Streamlit) | 低 | 简单看板:抓取状态、最新日期、表行数;非必需 |
|
||||
|
||||
---
|
||||
|
||||
@@ -84,12 +158,23 @@
|
||||
- 每次发现可复现 bug,先把现象写进本文件,再开始改代码,避免漏修。
|
||||
- 新增代码请用 `from src.log import get_logger`,不要再写 `print(..., flush=True)`。
|
||||
- 完成的条目移到下方 [变更日志](#变更日志),附完成日期,便于回顾。
|
||||
- **路线图条目立项时**:从 [Roadmap](#-新功能路线图-roadmap) 挪到 P1/P2,附责任人 + 预计动手时间;落地后再移到变更日志。
|
||||
|
||||
---
|
||||
|
||||
## 变更日志
|
||||
|
||||
### 2026-05-15
|
||||
### 2026-05-15(第二轮)
|
||||
|
||||
- ✅ **P2-#5 print → logger 全面替换**:`src/baostock_conn.py`、`src/db.py`、`src/main.py` 与全部 9 个 fetcher 中的 72 处 `print(..., flush=True)` 已切到 `from src.log import get_logger`,按语义选 INFO/WARNING/ERROR;BaoStock 查询签名打 DEBUG,避免控制台被刷屏
|
||||
- ✅ **P2-#7 扩展测试覆盖**:测试用例从 19 → 43。新增
|
||||
- `test_log.py`(5 例:命名空间、handler 幂等、不冒泡、env 控制 level、未知 level 回退)
|
||||
- `test_daily_source_codes.py`(8 例:sina/tencent/eastmoney 代码前缀映射)
|
||||
- `test_daily_sources_http.py`(8 例:腾讯/东财 HTTP JSON 解析,含空字段、异常包装、北交所短路)
|
||||
- ✅ **P2-#8 sector 三种 only 模式合并保护**:`test_sector_merge.py`(3 例:region_only 保留 industry / industry_only 保留 region / concept_only 不触碰 stock_sector)
|
||||
- ✅ **P2-#4 并发压测脚本骨架**:`benchmarks/bench_daily.py` + `benchmarks/README.md`,可一键跑 `workers=1,2,4,8` 对照,输出 speedup 表;剩下的就是用户在低峰期跑一次实测
|
||||
|
||||
### 2026-05-15(第一轮)
|
||||
|
||||
- ✅ `src/main.py` 新增 `--market-daily` 参数,挂载 `fetch_market_daily`(market_daily 只读 stock_daily,无需 BaoStock 登录)
|
||||
- ✅ `src/fetchers/market_daily.py` 修复 `fetch_history` → `_fetch_history` 笔误
|
||||
|
||||
@@ -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` 注释里。
|
||||
@@ -0,0 +1 @@
|
||||
"""benchmarks 包,存放性能压测脚本。"""
|
||||
@@ -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()
|
||||
@@ -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 条行业+地域记录
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
@@ -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("无概念板块数据")
|
||||
|
||||
@@ -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("无股票数据")
|
||||
|
||||
@@ -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
@@ -10,9 +10,7 @@
|
||||
python -m src.main --intraday --freq all # 全部频率
|
||||
python -m src.main --intraday --start-date 20260508 --end-date 20260509
|
||||
python -m src.main --intraday --symbol 000001 --freq 30
|
||||
python -m src.main --sector # 行业+地域分类
|
||||
python -m src.main --sector --industry-only # 仅行业分类
|
||||
python -m src.main --sector --concept-only # 仅概念板块
|
||||
python -m src.main --sector # 概念板块及成分股
|
||||
python -m src.main --index # 指数日线(上证/沪深300/创业板等)
|
||||
python -m src.main --market-daily # 汇总每日涨跌停统计(依赖 stock_daily)
|
||||
"""
|
||||
@@ -22,6 +20,9 @@ import argparse
|
||||
from src.baostock_conn import bs_login, bs_logout
|
||||
from src.config import load_config
|
||||
from src.db import init_db
|
||||
from src.log import get_logger
|
||||
|
||||
_logger = get_logger("main")
|
||||
|
||||
|
||||
def main():
|
||||
@@ -38,13 +39,10 @@ def main():
|
||||
parser.add_argument("--freq", type=str, default="5",
|
||||
choices=["5", "15", "30", "60", "all"],
|
||||
help="分钟K线频率(默认5,all=全部)")
|
||||
parser.add_argument("--sector", action="store_true", help="抓取行业+地域分类")
|
||||
parser.add_argument("--sector", action="store_true", help="抓取概念板块及成分股")
|
||||
parser.add_argument("--index", action="store_true", help="抓取指数日线行情")
|
||||
parser.add_argument("--market-daily", action="store_true",
|
||||
help="汇总每日涨跌停统计(从 stock_daily 聚合)")
|
||||
parser.add_argument("--industry-only", action="store_true", help="仅抓取行业分类")
|
||||
parser.add_argument("--region-only", action="store_true", help="仅抓取地域分类")
|
||||
parser.add_argument("--concept-only", action="store_true", help="仅抓取概念板块")
|
||||
parser.add_argument("--start-date", type=str, help="开始日期 YYYYMMDD")
|
||||
parser.add_argument("--end-date", type=str, help="结束日期 YYYYMMDD")
|
||||
parser.add_argument("--symbol", type=str, help="指定单只股票代码")
|
||||
@@ -53,8 +51,7 @@ def main():
|
||||
|
||||
if not any([args.stock_info, args.trading_day, args.daily,
|
||||
args.financial, args.dividend, args.intraday,
|
||||
args.sector, args.industry_only, args.region_only,
|
||||
args.concept_only, args.index, args.market_daily]):
|
||||
args.sector, args.index, args.market_daily]):
|
||||
parser.print_help()
|
||||
return
|
||||
|
||||
@@ -68,9 +65,18 @@ def main():
|
||||
# 若仅运行 market-daily,避免无谓的登录
|
||||
if not any([args.stock_info, args.trading_day, args.daily,
|
||||
args.financial, args.dividend, args.intraday,
|
||||
args.sector, args.industry_only, args.region_only,
|
||||
args.concept_only, args.index]):
|
||||
print("全部任务完成", flush=True)
|
||||
args.sector, args.index]):
|
||||
_logger.info("全部任务完成")
|
||||
return
|
||||
|
||||
# sector 不需要 BaoStock,提前处理
|
||||
if args.sector:
|
||||
from src.fetchers.sector import fetch_sector
|
||||
fetch_sector()
|
||||
if not any([args.stock_info, args.trading_day, args.daily,
|
||||
args.financial, args.dividend, args.intraday,
|
||||
args.index, args.market_daily]):
|
||||
_logger.info("全部任务完成")
|
||||
return
|
||||
|
||||
bs_login()
|
||||
@@ -101,16 +107,11 @@ def main():
|
||||
fetch_intraday(start_date=args.start_date, end_date=args.end_date,
|
||||
symbol=args.symbol, freq=args.freq)
|
||||
|
||||
if args.sector or args.industry_only or args.region_only or args.concept_only:
|
||||
from src.fetchers.sector import fetch_sector
|
||||
fetch_sector(industry_only=args.industry_only, region_only=args.region_only,
|
||||
concept_only=args.concept_only)
|
||||
|
||||
if args.index:
|
||||
from src.fetchers.index import fetch_index
|
||||
fetch_index(start_date=args.start_date, end_date=args.end_date)
|
||||
|
||||
print("全部任务完成", flush=True)
|
||||
_logger.info("全部任务完成")
|
||||
finally:
|
||||
bs_logout()
|
||||
|
||||
|
||||
+9
-6
@@ -9,10 +9,13 @@ pip install -e .[dev] # 或 pip install pytest
|
||||
pytest -v
|
||||
```
|
||||
|
||||
## 覆盖范围
|
||||
## 覆盖范围(43 个用例)
|
||||
|
||||
- `test_code_mapping.py` —— `baostock_conn.code_to_bs` 代码前缀映射
|
||||
- `test_daily_derived.py` —— `daily._fill_derived_fields` 振幅/涨跌幅补算
|
||||
- `test_financial_quarters.py` —— `financial._recent_quarters` 季度滚动
|
||||
- `test_market_classification.py` —— `market_daily._is_20pct` 板块判定
|
||||
- `test_sector_merge.py` —— `sector.fetch_sector` 三种 only 模式下不会清空其他列
|
||||
- `test_code_mapping.py` —— `baostock_conn.code_to_bs` 代码前缀映射(5)
|
||||
- `test_daily_derived.py` —— `daily._fill_derived_fields` 振幅/涨跌幅补算(5)
|
||||
- `test_daily_source_codes.py` —— sina/tencent/eastmoney 各源代码前缀映射(8)
|
||||
- `test_daily_sources_http.py` —— 腾讯/东财 HTTP JSON 解析(mock requests,8)
|
||||
- `test_financial_quarters.py` —— `financial._recent_quarters` 季度滚动(5)
|
||||
- `test_log.py` —— `log.get_logger` 命名空间、handler 幂等、env 控制 level(5)
|
||||
- `test_market_classification.py` —— `market_daily._is_20pct` 板块判定(4)
|
||||
- `test_sector_merge.py` —— `sector.fetch_sector` 三种 only 模式合并保护(3)
|
||||
|
||||
@@ -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
|
||||
@@ -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")
|
||||
@@ -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
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user