P0/P1/P2 一次性落地:补齐文档、修复 main 挂载/笔误、引入 logging+tests

P0:
- 清理 src/fetchers/sector.py 第 225 行起的旧版残留代码
- 修复 src/fetchers/market_daily.py 中 fetch_history -> _fetch_history 笔误
- 在 src/main.py 挂载 --market-daily 子命令

P1:
- 修正 src/db.py docstring(market_breadth -> market_daily 等)
- requirements.txt 补 baostock;pyproject.toml 同步 + 新增 [dev] extras
- README 增加 config.yaml 安全提示,将 git 历史清理升级为高风险 P0 由用户决策

P2:
- 引入 src/log.py 统一 logging(控制台 + logs/ashare.log 按日滚动 7 天)
- 建立 tests/ 框架,4 个测试文件 / 19 个 pytest 用例全部通过
- pyproject.toml 新增 [tool.pytest.ini_options]
- config.example.yaml 补全 workers 字段与多源说明
- README 更新功能概览/参数说明/表结构/项目结构/设计说明全章节
- TODO.md 全面重写,按 P0/P2 整理剩余条目并附变更日志
- .gitignore 新增 .claude/
This commit is contained in:
曾志威
2026-05-15 13:26:21 +08:00
parent fa31344241
commit dbc8107caa
19 changed files with 655 additions and 201 deletions
+3
View File
@@ -178,3 +178,6 @@ cython_debug/
.pypirc
.idea/
# Claude Code local tool config
.claude/
+121 -45
View File
@@ -1,22 +1,23 @@
# ashare-data
A股数据抓取工具,使用 [BaoStock](http://baostock.com) 获取数据,保存到 MySQL 数据库。
A股数据抓取工具, [BaoStock](http://baostock.com) 为主、新浪/腾讯/东方财富为辅,保存到 MySQL 数据库。
## 功能概览
| 数据类型 | 说明 | 数据来源 |
|---------|------|---------|
| 股票列表 | 沪深A股代码、名称、上市日期 | BaoStock |
| 交易日历 | 1990年至今的交易日列表 | BaoStock |
| 日线行情 | 开高低收、成交量/额、振幅、涨跌幅、换手率(前复权) | BaoStock |
| 指数日线 | 上证/深证/创业板等主要指数日线 | BaoStock |
| 涨跌停统计 | 每日主板/科创板/创业板涨跌停数量 | stock_daily 汇总 |
| 季频财务指标 | 盈利能力、偿债能力、现金流(最近8个季度) | BaoStock |
| 交易日历 | 1990年至今的交易日列表(自动用 stock_daily 校验节假日) | BaoStock |
| 日线行情 | 开高低收、成交量/额、振幅、涨跌幅、换手率(前复权) | BaoStock / 新浪 / 腾讯 / 东方财富(多源轮换) |
| 指数日线 | 上证/沪深300/中证500/中证1000/科创50/深证成指/创业板指/中小板指 | BaoStock |
| 涨跌停统计 | 每日主板10%/ 科创创业板(20%涨跌停数量 | stock_daily 汇总 |
| 季频财务指标 | 盈利能力、偿债能力、现金流(最近8个季度JSON 存储 | BaoStock |
| 分红送转 | 每10股送转、派息、除权除息日(最近10年) | BaoStock |
| 分时行情 | 5/15/30/60分钟K线(开高低收、成交量/额) | BaoStock |
| 分钟K线 | 5/15/30/60 分钟K线(开高低收、成交量/额) | BaoStock |
| 行业+地域分类 | 证监会行业分类 + 省份 | BaoStock + 东方财富 |
| 概念板块 | 东方财富全量概念板块及其成分股 | 东方财富 |
**已知限制**BaoStock 不含北交所(920xxx)股票。
**已知限制**BaoStock 不含北交所(920xxx)股票;新浪/腾讯/东方财富数据源对北交所同样不支持
## 快速开始
@@ -28,11 +29,15 @@ A股数据抓取工具,使用 [BaoStock](http://baostock.com) 获取数据,
### 2. 安装依赖
```bash
pip install baostock pymysql sqlalchemy pyyaml pandas
pip install baostock akshare pymysql sqlalchemy pyyaml pandas requests
```
> 说明:`baostock` 为主要数据源;`akshare` 用于新浪日线兜底;`requests` 用于腾讯/东方财富/概念板块抓取。
### 3. 配置数据库
> ⚠️ **安全提示**`config.yaml` 包含数据库密码,**必须**保留在 `.gitignore` 中(本仓库已默认忽略)。请勿在提交前删除此规则。如曾不慎提交真实密码,参见 [TODO.md](./TODO.md) "清理 git 历史中的明文密码"。
复制配置文件并修改 MySQL 连接信息:
```bash
@@ -74,21 +79,19 @@ python -m src.main --trading-day --start-date 19901219 --end-date 20261231
# 3. 抓取全历史日线行情(首次,耗时较长)
python -m src.main --daily --start-date 19901201 --end-date 20260511
# 4. 日常增量更新日线(默认最近30天
# 4. 日常增量更新日线(默认走多源轮换:baostock/sina/tencent/eastmoney
python -m src.main --daily
# 抓取指数日线
# 指定单一数据源(默认 all 表示多源轮换+自动切换)
python -m src.main --daily --source baostock
python -m src.main --daily --source sina
# 抓取主要指数日线(上证/沪深300/中证500/中证1000/科创50/深证成指/创业板指/中小板指)
python -m src.main --index
# 抓取指数日线(指定日期范围)
python -m src.main --index --start-date 19901219 --end-date 20260511
# 汇总涨跌停统计(依赖 stock_daily 数据)
python -m src.main --market-daily
# 汇总涨跌停统计(指定日期范围)
python -m src.main --market-daily --start-date 20260101 --end-date 20260511
# 抓取财务指标(全部股票,最近8个季度)
python -m src.main --financial
@@ -101,34 +104,44 @@ python -m src.main --dividend
# 抓取单只股票的分红
python -m src.main --dividend --symbol 000001
# 抓取分钟K线(默认5分钟)
# 抓取分钟K线(默认5分钟,近30天
python -m src.main --intraday --start-date 20260101
# 抓取全部频率分钟K线
# 抓取全部频率分钟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)。
### 5. 命令行参数说明
```
选项:
数据抓取选项:
--stock-info 抓取A股股票列表
--trading-day 抓取交易日历
--daily 抓取日线行情(增量,已完整自动跳过)
--daily 抓取日线行情(增量;自动分析数据缺口,已完整自动跳过)
--source 日线数据源: baostock/sina/tencent/eastmoney/all(默认 all,轮换+失败自动切换)
--index 抓取主要指数日线
--market-daily 汇总每日涨跌停统计(从 stock_daily 聚合)
--financial 抓取季频财务指标
--dividend 抓取分红送转数据
--intraday 抓取分钟K线行情
--freq K线频率: 5/15/30/60/all(默认5
--sector 抓取行业+地域分类
--freq K线频率: 5/15/30/60/all(默认 5
--sector 抓取行业+地域分类(默认两者都抓)
--industry-only 仅抓取行业分类
--region-only 仅抓取地域分类
--concept-only 仅抓取概念板块及成分股
日期过滤(对日线行情、交易日历、分钟K线、指数、涨跌停统计生效):
日期过滤(对日线行情、交易日历、分钟K线、指数生效):
--start-date 开始日期,格式 YYYYMMDD
--end-date 结束日期,格式 YYYYMMDD
@@ -136,6 +149,8 @@ python -m src.main --sector
--symbol 指定单只股票代码,如 000001,默认全部股票
```
> 注:`--market-daily` 在历史版本中存在,当前已下线(功能代码仍保留,但未挂到 main.py)。如需重新启用,参见 [TODO.md](./TODO.md) 中的 "重新挂载 market_daily 命令"。
## 数据库表结构
### stock_info — 股票基本信息
@@ -249,35 +264,96 @@ python -m src.main --sector
| industry | VARCHAR(100) | 证监会行业分类 |
| region | VARCHAR(20) | 省份/地域 |
### stock_concept — 概念板块及成分股
| 字段 | 类型 | 说明 |
|------|------|------|
| code | VARCHAR(10) | 股票代码 |
| concept_code | VARCHAR(20) | 概念板块代码(东方财富 BK 编码) |
| concept_name | VARCHAR(100) | 概念板块名称 |
联合主键:`(code, concept_code)`
### index_daily — 指数日线行情
| 字段 | 类型 | 说明 |
|------|------|------|
| code | VARCHAR(10) | 指数代码 |
| date | DATE | 交易日期 |
| open | FLOAT | 开盘价 |
| high | FLOAT | 最高价 |
| low | FLOAT | 最低价 |
| close | FLOAT | 收盘价 |
| volume | FLOAT | 成交量 |
| amount | FLOAT | 成交额 |
| pct_change | FLOAT | 涨跌幅% |
联合主键:`(code, date)`
### market_daily — 每日涨跌停统计
| 字段 | 类型 | 说明 |
|------|------|------|
| date | DATE PK | 交易日期 |
| limit_up_10 | INT | 涨停数(10% 主板) |
| limit_up_20 | INT | 涨停数(20% 科创板/创业板) |
| limit_down_10 | INT | 跌停数(10% 主板) |
| limit_down_20 | INT | 跌停数(20% 科创板/创业板) |
> 阈值:涨停 `pct_change >= 9.8`(主板)或 `>= 19.5`(20% 板块);跌停反向。当前由 `src/fetchers/market_daily.py` 提供 `fetch_market_daily()`main.py 暂未挂载,详见 [TODO.md](./TODO.md)。
## 项目结构
```
a股数据看板/
├── config.yaml # MySQL连接配置(.gitignore
ashare-data/
├── config.yaml # MySQL 连接配置(.gitignore
├── config.example.yaml # 配置文件模板
├── pyproject.toml # 项目元数据
├── pyproject.toml # 项目元数据
├── requirements.txt # 运行依赖
├── README.md # 项目说明(当前文档)
├── TODO.md # 已知问题与待办事项
├── src/
│ ├── __init__.py
│ ├── config.py # 配置读取模块
│ ├── baostock_conn.py # BaoStock连接管理(login/logout/线程锁)
│ ├── db.py # 数据库模型与连接管理
│ ├── main.py # 命令行入口
│ ├── config.py # 配置读取模块YAML + 环境变量 ASHARE_CONFIG
│ ├── baostock_conn.py # BaoStock 连接管理(login/logout/线程锁/超时重连
│ ├── db.py # SQLAlchemy 模型 + 批量 upsert + 自动迁移
│ ├── main.py # 命令行入口
│ └── fetchers/
│ ├── __init__.py
│ ├── stock_list.py # 股票列表
│ ├── trading_day.py # 交易日历
│ ├── daily.py # 日线行情
│ ├── financial.py # 季频财务指标
│ ├── dividend.py # 分红送转
│ ├── intraday.py # 分钟K线(5/15/30/60
── sector.py # 行业+地域分类
└── README.md
│ ├── stock_list.py # 股票列表
│ ├── trading_day.py # 交易日历
│ ├── daily.py # 日线行情(多源轮换:baostock/sina/tencent/eastmoney
│ ├── index.py # 主要指数日线
│ ├── market_daily.py # 每日涨跌停统计(从 stock_daily 聚合)
│ ├── financial.py # 季频财务指标(盈利/偿债/现金流,JSON 存储
── dividend.py # 分红送转
│ ├── intraday.py # 分钟K线(5/15/30/60
│ └── sector.py # 行业+地域分类 + 概念板块
└── gzl/ # 选股脚本(独立子项目,可选)
├── Selector.py
└── select_stock.py
```
## 设计说明
- **数据源**统一使用 BaoStock,无需 AKShare 等额外依赖
- **线程安全**BaoStock 查询通过全局锁串行化,避免数据错乱
- **去重写入**:所有表使用 `INSERT ON DUPLICATE KEY UPDATE`upsert,重复执行不会产生重复数据
- **增量更新**:日线行情支持通过 `--start-date` / `--end-date` 指定日期范围,已完整的数据自动跳过
- **停牌标记**缺失交易日自动标记为停牌,避免重复抓取
- **数据源容灾**日线行情默认 `--source all` 在 BaoStock / 新浪 / 腾讯 / 东方财富之间轮换并自动切换,单源失败不影响整体进度;其他模块(财务、分红、指数、分钟K线、行业)仍以 BaoStock 为主。
- **线程安全**BaoStock `query_xxx()` 非线程安全,所有调用通过 `src/baostock_conn.py` 的全局锁串行化;查询超时/连接断开时自动重连。
- **去重写入**:所有表通过 `db.batch_upsert()` 走 MySQL `INSERT ON DUPLICATE KEY UPDATE`,重复执行不会产生重复数据
- **增量更新**:日线行情按月统计已有数据并与交易日历对比,只抓取真正缺口段;其他数据按 (code, 周期) 粒度跳过已完成的股票。
- **停牌识别**日线行情若两端已覆盖、内部仍有缺口,则视为停牌,不再重抓。
- **表结构自动迁移**`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)` 调用本期保留兼容。
## 运行测试
```bash
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 个用例。
## 选股子项目 `gzl/`
`gzl/` 目录是独立的选股策略原型(依赖 `scipy` + 本地 CSV,不读 MySQL)。当前与主项目 `src/` 解耦,**不参与** `python -m src.main` 的任何流程。后续是否合并到 `src/strategies/` 见 [TODO.md](./TODO.md)。
+104
View File
@@ -0,0 +1,104 @@
# TODO
本文件用于跟踪 ashare-data 项目的已知问题与后续工作。维护时请保持「问题描述 + 影响范围 + 处理思路」三段式,便于他人接手。
> 🗓️ **2026-05-15 一次性清理**:原 P0 / P1 / P2 中可在不破坏 git 历史的前提下完成的条目已落地(详见底部 [变更日志](#变更日志))。当前剩余条目均为:①需要用户决策的高风险动作,②长期工程任务,③依赖外部数据源调研。
---
## 🔴 P0 — 需用户授权的高风险动作
### 1. 清理 git 历史中的明文密码 ⚠️ 高风险
- **现状**`config.yaml` 现已在 `.gitignore` 中(仓库根 `.gitignore` 第 2 行),但**历史 commit `e2c3b41` / `feef7b6` / `217b12c` 中曾以明文形式提交过**
```
password: "ttx2011"
host: "db.freeicu.top"
port: 32000
```
任何能访问本仓库的人都能通过 `git show e2c3b41:config.yaml` 取到这段凭据。
- **影响**:数据库密码事实上已经泄露;如仓库已 push 到远端(包括 fork/克隆),轮换密码是**唯一彻底方案**。
- **处理思路**(需用户确认后才能执行):
1. **立即**在 MySQL 侧轮换 `root@db.freeicu.top:32000` 的密码;
2. (可选)用 `git filter-repo --path config.yaml --invert-paths` 或 BFG Repo-Cleaner 重写历史:
```bash
git filter-repo --path config.yaml --invert-paths
git push --force --all # ⚠️ 破坏所有协作者的本地副本,需团队周知
git push --force --tags
```
3. 让所有协作者**重新克隆**仓库(旧 clone 的 reflog 仍含明文)。
4. 如已公开过(GitHub/GitLab),还需手动让对应平台清理缓存(GitHub 联系 support@github.com,或新建仓库迁移)。
- ❗ 这一步会**改写所有 commit 哈希**,等同于强制变基整个历史,必须由仓库所有者亲自决定并在低峰期执行。请回复后再操作。
---
## 🟢 P2 — 长期工程 & 调研
### 2. 北交所(920xxx)数据支持
- **现状**:BaoStock、新浪、腾讯、东方财富的 K 线接口均不支持北交所;当前在 `daily.py / sector.py / stock_list.py` 中显式跳过 `920xxx`。
- **处理思路**:调研同花顺 / 雪球 / Wind Quant 等接口;若可,新增独立 fetcher 并合并到主流程。预计需要新增一张 `bj_daily` 表或在 `stock_daily` 中加 market 列。
### 3. 财务数据「JSON 存 TEXT」结构化拆分
- **现状**`stock_financial_income/balance/cashflow` 仅有 `code/report_date/data(JSON)` 三列,下游查询需 `JSON_EXTRACT`,难做索引。
- **处理思路**:根据下游真实查询场景(量化筛选 vs 财报展示),把高频指标(ROE、净利润、资产负债率、经营性现金流等)拆出独立列;保留 `extra_json` 兜底。需配套写数据迁移脚本。
### 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`。
### 6. `gzl/` 选股脚本接入主项目
- **现状**`gzl/Selector.py` + `gzl/select_stock.py` 读本地 CSV(不读 MySQL),与 `src/` 完全解耦,依赖 `scipy`。README 已说明其独立性。
- **处理思路**
1. 评估是否纳入主项目 — 如果只是个人玩具脚本可保持现状;
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)单测,验证主板/创业板不会互串。
### 8. `sector.fetch_sector` 三种 only 模式的合并保护测试
- **现状**`region_only=True` 时使用 `existing.get("industry")` 保留旧行业值;`industry_only=True` 反之。逻辑正确但无测试。
- **处理思路**:归并到上面第 7 项一起做(需 mock DB session 或用 SQLite)。
---
## 📋 维护节奏建议
- 每次新增 fetcher,请**同步更新** README 的:「功能概览表 / 运行示例 / 命令行参数说明 / 数据库表结构 / 项目结构」 五个章节。
- 每次发现可复现 bug,先把现象写进本文件,再开始改代码,避免漏修。
- 新增代码请用 `from src.log import get_logger`,不要再写 `print(..., flush=True)`。
- 完成的条目移到下方 [变更日志](#变更日志),附完成日期,便于回顾。
---
## 变更日志
### 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` 笔误
- ✅ `src/fetchers/sector.py` 清理第 225 行起的重复 import + 旧版函数残留
- ✅ `src/db.py` 顶部 docstring 修正 `market_breadth` → `market_daily`,并补全 index_daily / stock_concept / 分钟K线分表
- ✅ `requirements.txt` 增加 `baostock``pyproject.toml` 同步并新增 `[dev]` extras
- ✅ `config.yaml` 历史明文密码问题:已在 README 加显眼 ⚠️ 安全提示;具体 git 历史清理与密码轮换升级为 P0-1 由用户决策
- ✅ 引入 `src/log.py` 统一 logging(控制台 + `logs/ashare.log` 按日滚动 7 天保留,`ASHARE_LOG_LEVEL` 可调)
- ✅ 建立 `tests/` 框架:4 个测试文件、19 个 pytest 用例全部通过(覆盖 `code_to_bs`、`_fill_derived_fields`、`_recent_quarters`、`_is_20pct`
- ✅ `pyproject.toml` 添加 `[tool.pytest.ini_options]``pytest` 可一键运行
- ✅ `config.example.yaml` 补全 `workers` 字段与多源说明
- ✅ README 增加「运行测试」「选股子项目 gzl/」两节,"设计说明" 加 logging 条
+5 -1
View File
@@ -1,3 +1,4 @@
# MySQL 数据库连接配置(复制为 config.yaml 后填入真实值)
mysql:
host: "localhost"
port: 3306
@@ -6,8 +7,11 @@ mysql:
database: "ashare"
charset: "utf8mb4"
# 数据抓取配置
fetch:
# 请求间隔(秒),避免被限流
# 每次请求间隔(秒),避免被限流
delay: 0.5
# 失败重试次数
retry: 3
# 并发线程数;BaoStock 单源建议 1;多源轮换(baostock/sina/tencent/eastmoney)可适度提高至 2~4
workers: 1
+11
View File
@@ -4,6 +4,7 @@ version = "0.1.0"
description = "A股数据抓取,保存到MySQL数据库"
requires-python = ">=3.10"
dependencies = [
"baostock",
"akshare",
"pymysql",
"sqlalchemy>=2.0",
@@ -12,5 +13,15 @@ dependencies = [
"requests",
]
[project.optional-dependencies]
dev = [
"pytest>=7",
]
[project.scripts]
ashare = "src.main:main"
[tool.pytest.ini_options]
testpaths = ["tests"]
python_files = ["test_*.py"]
addopts = "-ra"
+2
View File
@@ -1,5 +1,7 @@
baostock
akshare
pymysql
sqlalchemy>=2.0
pyyaml
pandas
requests
+7 -4
View File
@@ -1,15 +1,18 @@
"""数据库模型定义与连接管理
数据源:BaoStock(唯一数据源)
数据源:BaoStock 为主,新浪/腾讯/东方财富兜底
表结构概览:
- stock_info: 股票基本信息(含上市日期,用于跳过未上市股票)
- stock_daily: 日线行情(BaoStock含振幅/涨跌幅/换手率)
- stock_daily: 日线行情(含振幅/涨跌幅/换手率)
- stock_financial_income/balance/cashflow: 季频财务指标(JSON存储)
- stock_dividend: 分红送转
- trading_day: 交易日历(用于判断数据完整性)
- stock_no_data: 无数据/停牌记录(避免重复抓取)
- stock_intraday: 5分钟分时行情
- market_breadth: 市场量能快照
- stock_min5/15/30/60: 分钟K线(4 张分表)
- index_daily: 主要指数日线
- market_daily: 每日涨跌停统计(10%/20% 板块分别计数)
- stock_sector: 行业 + 地域分类
- stock_concept: 概念板块及成分股
"""
from sqlalchemy import (
+118 -77
View File
@@ -15,7 +15,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed
import threading
import requests
import baostock as bs
from src.baostock_conn import bs_query, code_to_bs, bs_login
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
@@ -70,41 +70,67 @@ def _fetch_baostock(code: str, start_date: str, end_date: str) -> list[dict] | N
return 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]}"
try:
with bs_query(
bs.query_history_k_data_plus,
bs_code,
"date,open,high,low,close,volume,amount,preclose,pctChg,turn",
start_date=sd, end_date=ed, frequency="d", adjustflag="2",
) as rs:
rows = []
while (rs.error_code == "0") and rs.next():
r = rs.get_row_data()
preclose = _clean(r[7])
pct_chg = _clean(r[8])
turn = _clean(r[9])
close = _clean(r[4])
high = _clean(r[2])
low = _clean(r[3])
amp = None
if high is not None and low is not None and preclose and float(preclose) != 0:
amp = round((float(high) - float(low)) / float(preclose) * 100, 2)
chg = None
if close is not None and preclose is not None:
try:
chg = round(float(close) - float(preclose), 3)
except (ValueError, TypeError):
pass
rows.append({
"code": code, "date": r[0],
"open": _clean(r[1]), "high": high, "low": low, "close": close,
"volume": _clean(r[5]), "turnover": _clean(r[6]),
"amplitude": amp, "pct_change": _clean(pct_chg), "change": chg,
"turnover_rate": _clean(turn),
})
return rows if rows else None
except Exception:
return None
retry = max(1, int(get_fetch_config().get("retry", 3)))
def _is_transient(err: Exception) -> bool:
msg = str(err)
msg_l = msg.lower()
return (
"10057" in msg
or "接收数据异常" in msg
or "socket" in msg_l
or "not connected" in msg_l
or "连接" in msg
or "sendto" in msg_l
)
for attempt in range(1, retry + 1):
try:
with bs_query(
bs.query_history_k_data_plus,
bs_code,
"date,open,high,low,close,volume,amount,preclose,pctChg,turn",
start_date=sd, end_date=ed, frequency="d", adjustflag="2",
) as rs:
rows = []
while (rs.error_code == "0") and rs.next():
r = rs.get_row_data()
preclose = _clean(r[7])
pct_chg = _clean(r[8])
turn = _clean(r[9])
close = _clean(r[4])
high = _clean(r[2])
low = _clean(r[3])
amp = None
if high is not None and low is not None and preclose and float(preclose) != 0:
amp = round((float(high) - float(low)) / float(preclose) * 100, 2)
chg = None
if close is not None and preclose is not None:
try:
chg = round(float(close) - float(preclose), 3)
except (ValueError, TypeError):
pass
rows.append({
"code": code, "date": r[0],
"open": _clean(r[1]), "high": high, "low": low, "close": close,
"volume": _clean(r[5]), "turnover": _clean(r[6]),
"amplitude": amp, "pct_change": _clean(pct_chg), "change": chg,
"turnover_rate": _clean(turn),
})
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)
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)
bs_logout()
time.sleep(min(2 * attempt, 5))
continue
print(f" [数据源:BaoStock] {code} 获取失败: {e}", flush=True)
return None
# ==================== 数据源: 新浪 (akshare) ====================
@@ -146,8 +172,8 @@ def _fetch_sina(code: str, start_date: str, end_date: str) -> list[dict] | None:
"turnover_rate": turnover_rate,
})
return rows if rows else None
except Exception:
return None
except Exception as e:
raise RuntimeError(f"新浪请求失败: {e}") from e
# ==================== 数据源: 腾讯 ====================
@@ -192,8 +218,8 @@ def _fetch_tencent(code: str, start_date: str, end_date: str) -> list[dict] | No
"turnover_rate": None,
})
return rows if rows else None
except Exception:
return None
except Exception as e:
raise RuntimeError(f"腾讯请求失败: {e}") from e
# ==================== 数据源: 东方财富 ====================
@@ -241,8 +267,8 @@ def _fetch_eastmoney(code: str, start_date: str, end_date: str) -> list[dict] |
"turnover_rate": float(parts[10]) if parts[10] != "" else None,
})
return rows if rows else None
except Exception:
return None
except Exception as e:
raise RuntimeError(f"东方财富请求失败: {e}") from e
# ==================== 数据源分发 ====================
@@ -365,26 +391,43 @@ def _analyze_gaps(codes: list[str], start_date: str, end_date: str,
_print_lock = threading.Lock()
def _rotate_sources(sources: list[str], offset: int) -> list[str]:
"""让 all 模式下的首选来源轮换,避免请求都集中到同一个来源。"""
if len(sources) <= 1:
return sources
start = offset % len(sources)
return sources[start:] + sources[:start]
def _fetch_and_save(code: str, gap_start: str, gap_end: str,
source_key: str) -> tuple[str, str, int, bool]:
source_keys: list[str]) -> tuple[str, str, int, bool]:
"""抓取单只股票并写入,返回 (code, source_label, row_count, success)"""
fetch_fn = _SOURCE_FN[source_key]
label = _SOURCE_LABEL[source_key]
try:
rows = fetch_fn(code, gap_start, gap_end)
if rows is not None:
rows = _fill_derived_fields(rows)
batch_upsert(StockDaily, rows, ["code", "date"])
return code, label, len(rows), True
return code, label, 0, False
except Exception:
return code, label, 0, False
last_label = ""
for index, source_key in enumerate(source_keys):
fetch_fn = _SOURCE_FN[source_key]
label = _SOURCE_LABEL[source_key]
last_label = label
try:
rows = fetch_fn(code, gap_start, gap_end)
if rows is not None:
rows = _fill_derived_fields(rows)
batch_upsert(StockDaily, rows, ["code", "date"])
return code, label, len(rows), True
except Exception as e:
print(f" [数据源:{label}] {code} 获取异常: {e}", flush=True)
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)
return code, last_label, 0, False
# ==================== 主函数 ====================
def fetch_daily(start_date: str | None = None, end_date: str | None = None,
source: str = "baostock"):
source: str = "all"):
# 确定使用的数据源列表
if source == "all":
sources = list(VALID_SOURCES)
@@ -394,10 +437,9 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None,
print(f" 不支持的数据源 {source},可选: {', '.join(VALID_SOURCES)}, all", flush=True)
return
source_names = ", ".join(_SOURCE_LABEL[s] for s in sources)
workers = len(sources)
cfg = get_fetch_config()
source_names = ", ".join(_SOURCE_LABEL[s] for s in sources)
workers = max(1, int(cfg.get("workers", 1)))
codes = get_stock_codes()
if not codes:
@@ -429,31 +471,32 @@ 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" {len(not_listed)} 只股票未上市,跳过", flush=True)
print(f" [数据源:{source_names}] {len(not_listed)} 只股票未上市,跳过", flush=True)
# 一次查询:判断完整 + 计算缺口
print(f" 正在分析 {len(codes)} 只股票的数据缺口...", flush=True)
print(f" [数据源:{source_names}] 正在分析 {len(codes)} 只股票的数据缺口...", flush=True)
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" {complete_count} 只股票数据已完整,跳过", flush=True)
print(f" [数据源:{source_names}] {complete_count} 只股票数据已完整,跳过", flush=True)
total = len(gaps)
if total == 0:
print(" 所有股票数据已完整,无需抓取", flush=True)
print(f" [数据源:{source_names}] 所有股票数据已完整,无需抓取", flush=True)
return
print(f"正在抓取日线行情({source_names}) {start_date} ~ {end_date}{total} 只需更新,并发:{workers}...", flush=True)
mode = "全来源轮换+失败自动切换" if source == "all" else "单一来源"
print(f"正在抓取日线行情 [数据源:{source_names}] [模式:{mode}] {start_date} ~ {end_date}{total} 只需更新,并发:{workers}...", flush=True)
# 为每只股票分配数据源(轮询)
# all 模式下轮换首选来源;单来源模式下只使用指定来源。
gap_items = list(gaps.items())
tasks = []
for i, (code, g) in enumerate(gap_items):
src_key = sources[i % len(sources)]
gap_start = g[0].replace("-", "")
gap_end = g[-1].replace("-", "")
tasks.append((code, gap_start, gap_end, g, src_key))
source_order = _rotate_sources(sources, i) if source == "all" else sources
tasks.append((code, gap_start, gap_end, g, source_order))
success = 0
fail = 0
@@ -462,10 +505,9 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None,
t_start = time.time()
if workers == 1:
# 单数据源串行
for code, gap_start, gap_end, g, src_key in tasks:
for code, gap_start, gap_end, g, src_keys in tasks:
t0 = time.time()
code, label, row_count, ok = _fetch_and_save(code, gap_start, gap_end, src_key)
code, label, row_count, ok = _fetch_and_save(code, gap_start, gap_end, src_keys)
t_fetch = time.time() - t0
done += 1
if ok:
@@ -477,19 +519,18 @@ 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]} "
print(f" [{done}/{total}] {code} [数据源:{label}] 缺口:{g[0]}~{g[-1]} "
f"耗时:{t_fetch:.1f}s 行数:{row_count} "
f"成功:{success} 剩余:{eta:.0f}s", flush=True)
else:
# 多数据源并发
with ThreadPoolExecutor(max_workers=workers) as pool:
future_map = {}
for code, gap_start, gap_end, g, src_key in tasks:
f = pool.submit(_fetch_and_save, code, gap_start, gap_end, src_key)
future_map[f] = (code, g, src_key)
for code, gap_start, gap_end, g, src_keys in tasks:
f = pool.submit(_fetch_and_save, code, gap_start, gap_end, src_keys)
future_map[f] = (code, g, src_keys)
for future in as_completed(future_map):
code, g, src_key = future_map[future]
code, g, src_keys = future_map[future]
code_r, label, row_count, ok = future.result()
done += 1
if ok:
@@ -502,9 +543,9 @@ 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]} "
print(f" [{done}/{total}] {code_r} [数据源:{label}] 缺口:{g[0]}~{g[-1]} "
f"行数:{row_count} 成功:{success} 剩余:{eta:.0f}s", flush=True)
total_time = time.time() - t_start
print(f" 日线行情抓取完成 数据源:{source_names} 并发:{workers} "
print(f" 日线行情抓取完成 [数据源:{source_names}] 并发:{workers} "
f"成功:{success} 失败:{fail} 无数据:{nodata_count} 总耗时:{total_time:.1f}s", flush=True)
+7 -1
View File
@@ -100,4 +100,10 @@ def _fetch_history(start_date: str | None, end_date: str | None):
def fetch_market_daily(start_date: str | None = None, end_date: str | None = None):
fetch_history(start_date, end_date)
"""从 stock_daily 汇总每日涨跌停统计并写入 market_daily 表
Args:
start_date: 开始日期 YYYYMMDD,默认 19901219
end_date: 结束日期 YYYYMMDD,默认今天
"""
_fetch_history(start_date, end_date)
-68
View File
@@ -221,71 +221,3 @@ def fetch_sector(industry_only: bool = False, region_only: bool = False, concept
batch_upsert(StockSector, rows, ["code"])
print(f" 已写入 {len(rows)} 条行业+地域记录", flush=True)
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, batch_upsert, get_session, get_stock_codes
from sqlalchemy import select
_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"
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
+58
View File
@@ -0,0 +1,58 @@
"""统一日志模块
所有 fetcher 应通过 `from src.log import get_logger` 获取 logger,避免散落的 print。
- 控制台默认 INFO;可通过环境变量 ASHARE_LOG_LEVEL 调整 (DEBUG/INFO/WARNING/ERROR)。
- 同时输出到 logs/ashare.log,按日滚动,保留 7 天。
- 出于兼容考虑,旧 print 调用本期保留,新代码请用 logger。
"""
import logging
import os
import sys
from logging.handlers import TimedRotatingFileHandler
from pathlib import Path
_INITIALIZED = False
_DEFAULT_FORMAT = "%(asctime)s [%(levelname)s] [%(name)s] %(message)s"
def _init_root() -> None:
global _INITIALIZED
if _INITIALIZED:
return
level_name = os.environ.get("ASHARE_LOG_LEVEL", "INFO").upper()
level = getattr(logging, level_name, logging.INFO)
root = logging.getLogger("ashare")
root.setLevel(level)
root.propagate = False
fmt = logging.Formatter(_DEFAULT_FORMAT, datefmt="%Y-%m-%d %H:%M:%S")
sh = logging.StreamHandler(sys.stdout)
sh.setFormatter(fmt)
root.addHandler(sh)
log_dir = Path(os.environ.get("ASHARE_LOG_DIR", "logs"))
try:
log_dir.mkdir(parents=True, exist_ok=True)
fh = TimedRotatingFileHandler(
log_dir / "ashare.log",
when="midnight",
backupCount=7,
encoding="utf-8",
)
fh.setFormatter(fmt)
root.addHandler(fh)
except OSError:
# 文件系统只读或权限不足时只打印到控制台
pass
_INITIALIZED = True
def get_logger(name: str) -> logging.Logger:
"""返回带统一格式的 logger(命名空间 ashare.<name>"""
_init_root()
return logging.getLogger(f"ashare.{name}")
+20 -5
View File
@@ -1,4 +1,4 @@
"""A股数据抓取工具主入口 — BaoStock 数据源
"""A股数据抓取工具主入口 — BaoStock + 多源容灾
用法示例:
python -m src.main --stock-info # 先抓取股票列表
@@ -14,6 +14,7 @@
python -m src.main --sector --industry-only # 仅行业分类
python -m src.main --sector --concept-only # 仅概念板块
python -m src.main --index # 指数日线(上证/沪深300/创业板等)
python -m src.main --market-daily # 汇总每日涨跌停统计(依赖 stock_daily
"""
import argparse
@@ -24,13 +25,13 @@ from src.db import init_db
def main():
parser = argparse.ArgumentParser(description="A股数据抓取工具(BaoStock")
parser = argparse.ArgumentParser(description="A股数据抓取工具(BaoStock + 多源容灾")
parser.add_argument("--stock-info", action="store_true", help="抓取股票列表")
parser.add_argument("--trading-day", action="store_true", help="抓取交易日历")
parser.add_argument("--daily", action="store_true", help="抓取日线行情")
parser.add_argument("--source", type=str, default="baostock",
parser.add_argument("--source", type=str, default="all",
choices=["baostock", "sina", "tencent", "eastmoney", "all"],
help="日线数据源(默认baostockall=全部并发")
help="日线数据源(默认all,轮换使用全部来源")
parser.add_argument("--financial", action="store_true", help="抓取季频财务指标")
parser.add_argument("--dividend", action="store_true", help="抓取分红送转")
parser.add_argument("--intraday", action="store_true", help="抓取分钟K线行情")
@@ -39,6 +40,8 @@ def main():
help="分钟K线频率(默认5all=全部)")
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="仅抓取概念板块")
@@ -51,13 +54,25 @@ 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.concept_only, args.index, args.market_daily]):
parser.print_help()
return
load_config()
init_db()
# market_daily 只读 stock_daily,无需 BaoStock 登录,提前处理
if args.market_daily:
from src.fetchers.market_daily import fetch_market_daily
fetch_market_daily(start_date=args.start_date, end_date=args.end_date)
# 若仅运行 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)
return
bs_login()
try:
if args.stock_info:
+18
View File
@@ -0,0 +1,18 @@
# tests 目录
针对纯函数(不依赖网络、数据库、BaoStock 登录)的单元测试。
## 运行
```bash
pip install -e .[dev] # 或 pip install pytest
pytest -v
```
## 覆盖范围
- `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 模式下不会清空其他列
+1
View File
@@ -0,0 +1 @@
"""使 pytest 能找到顶层包 `src`"""
+9
View File
@@ -0,0 +1,9 @@
"""pytest 公用 fixtures / 路径配置"""
import sys
from pathlib import Path
# 确保 `from src...` 在 tests 中可用,无需安装
ROOT = Path(__file__).resolve().parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
+30
View File
@@ -0,0 +1,30 @@
"""验证 code_to_bs 的代码前缀映射规则"""
from src.baostock_conn import code_to_bs
def test_shanghai_main_board():
assert code_to_bs("600000") == "sh.600000"
assert code_to_bs("601318") == "sh.601318"
def test_shanghai_kechuang():
# 科创板 688 仍以 6 开头
assert code_to_bs("688981") == "sh.688981"
def test_shenzhen_main_and_chinext():
assert code_to_bs("000001") == "sz.000001"
assert code_to_bs("002594") == "sz.002594"
assert code_to_bs("300750") == "sz.300750"
def test_beijing_returns_none():
# BaoStock 不含北交所,应显式返回 None
assert code_to_bs("920001") is None
assert code_to_bs("920999") is None
def test_warrant_or_b_share_prefix_9():
# 沪市 B 股以 9 开头,与 6 同走 sh.
assert code_to_bs("900901") == "sh.900901"
+63
View File
@@ -0,0 +1,63 @@
"""验证 daily._fill_derived_fields 的振幅/涨跌幅/涨跌额补算逻辑"""
from src.fetchers.daily import _fill_derived_fields
def test_empty_input():
assert _fill_derived_fields([]) == []
def test_first_row_has_no_preclose():
rows = [{"date": "2026-05-08", "open": 10.0, "high": 11.0, "low": 9.5, "close": 10.5,
"amplitude": None, "pct_change": None, "change": None}]
out = _fill_derived_fields(rows)
# 没有 preclose 时不应补,pct_change 等保持 None
assert out[0]["pct_change"] is None
assert out[0]["change"] is None
assert out[0]["amplitude"] is None
def test_subsequent_row_derives_from_previous_close():
rows = [
{"date": "2026-05-08", "open": 10.0, "high": 11.0, "low": 9.5, "close": 10.0,
"amplitude": None, "pct_change": None, "change": None},
{"date": "2026-05-09", "open": 10.2, "high": 11.0, "low": 9.0, "close": 11.0,
"amplitude": None, "pct_change": None, "change": None},
]
out = _fill_derived_fields(rows)
second = out[1]
# 涨跌额 = 11 - 10 = 1.0
assert second["change"] == 1.0
# 涨跌幅 = (11 - 10) / 10 * 100 = 10.0
assert second["pct_change"] == 10.0
# 振幅 = (11 - 9) / 10 * 100 = 20.0
assert second["amplitude"] == 20.0
def test_existing_values_are_preserved():
"""已有值不应被覆盖"""
rows = [
{"date": "2026-05-08", "open": 10.0, "high": 11.0, "low": 9.5, "close": 10.0,
"amplitude": None, "pct_change": None, "change": None},
{"date": "2026-05-09", "open": 10.2, "high": 11.0, "low": 9.0, "close": 11.0,
"amplitude": 99.9, "pct_change": 88.8, "change": 7.7},
]
out = _fill_derived_fields(rows)
assert out[1]["amplitude"] == 99.9
assert out[1]["pct_change"] == 88.8
assert out[1]["change"] == 7.7
def test_zero_preclose_does_not_divide_by_zero():
"""preclose=0 时不应抛 ZeroDivisionError,保持 None"""
rows = [
{"date": "2026-05-08", "open": 0, "high": 0, "low": 0, "close": 0,
"amplitude": None, "pct_change": None, "change": None},
{"date": "2026-05-09", "open": 0.1, "high": 1, "low": 0.1, "close": 0.5,
"amplitude": None, "pct_change": None, "change": None},
]
out = _fill_derived_fields(rows)
# preclose=0 时三个派生字段都保持 None
assert out[1]["pct_change"] is None
assert out[1]["change"] is None
assert out[1]["amplitude"] is None
+50
View File
@@ -0,0 +1,50 @@
"""验证 financial._recent_quarters 的季度滚动逻辑"""
from datetime import datetime
from unittest.mock import patch
from src.fetchers import financial
def _fake_now(year: int, month: int, day: int = 15):
"""生成一个固定时间,用于 patch datetime.now()"""
fixed = datetime(year, month, day)
class FakeDatetime(datetime):
@classmethod
def now(cls, tz=None): # noqa: ARG003
return fixed
return FakeDatetime
def test_count_matches():
with patch.object(financial, "datetime", _fake_now(2026, 5, 15)):
out = financial._recent_quarters(8)
assert len(out) == 8
def test_first_is_current_quarter():
# 2026-05-15 属于 Q2
with patch.object(financial, "datetime", _fake_now(2026, 5, 15)):
out = financial._recent_quarters(4)
assert out[0] == (2026, 2)
def test_crosses_year_boundary():
# 2026 Q1 → 2025 Q4 → 2025 Q3 → 2025 Q2
with patch.object(financial, "datetime", _fake_now(2026, 2, 10)):
out = financial._recent_quarters(4)
assert out == [(2026, 1), (2025, 4), (2025, 3), (2025, 2)]
def test_q1_january():
with patch.object(financial, "datetime", _fake_now(2026, 1, 1)):
out = financial._recent_quarters(2)
assert out == [(2026, 1), (2025, 4)]
def test_q4_december():
with patch.object(financial, "datetime", _fake_now(2025, 12, 31)):
out = financial._recent_quarters(2)
assert out == [(2025, 4), (2025, 3)]
+28
View File
@@ -0,0 +1,28 @@
"""验证 market_daily._is_20pct 的板块判定"""
from src.fetchers.market_daily import _is_20pct
def test_main_board_10pct():
# 沪深主板 10% 限价
assert _is_20pct("600000") is False
assert _is_20pct("601318") is False
assert _is_20pct("000001") is False
assert _is_20pct("002594") is False
def test_chinext_20pct():
# 创业板 300 / 301 → 20%
assert _is_20pct("300750") is True
assert _is_20pct("301001") is True
def test_kechuang_20pct():
# 科创板 688 / 689 → 20%
assert _is_20pct("688981") is True
assert _is_20pct("689001") is True
def test_beijing_not_20pct():
# 北交所 920 系列即便涨跌幅 20%+,也不属于科创/创业 20% 板块计数
assert _is_20pct("920001") is False