diff --git a/CLAUDE.md b/CLAUDE.md index e2beacd..f7c879d 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -4,7 +4,7 @@ - **禁止删除数据库数据**:不允许执行任何 `DELETE`、`TRUNCATE`、`DROP TABLE`、`session.delete()` 等删除操作,除非用户明确要求 - **禁止清理数据**:不要主动建议或执行清理数据库表的操作 -- `db.py` 中的 `init_db()` 里的 `DROP TABLE` 是表结构自动迁移逻辑,属于例外,不要修改 +- 严格禁止任何 `DROP` 相关操作,不允许把删除表、删数据或清理表结构作为自动迁移手段 ## 代码规范 diff --git a/README.md b/README.md index 952c543..62c6822 100644 --- a/README.md +++ b/README.md @@ -1,22 +1,20 @@ # ashare-data -A股数据抓取工具,以 [BaoStock](http://baostock.com) 为主、新浪/腾讯/东方财富为辅,保存到 MySQL 数据库。 +A股数据抓取工具,股票列表、交易日历使用 AKShare,日线行情使用 BaoStock,其余部分按功能分别使用通达信本地缓存或 AKShare,保存到 MySQL 数据库。 ## 功能概览 | 数据类型 | 说明 | 数据来源 | |---------|------|---------| -| 股票列表 | 沪深A股代码、名称、上市日期 | BaoStock | -| 交易日历 | 1990年至今的交易日列表(自动用 stock_daily 校验节假日) | BaoStock | -| 日线行情 | 开高低收、成交量/额、振幅、涨跌幅、换手率(前复权) | BaoStock / 新浪 / 腾讯 / 东方财富(多源轮换) | -| 指数日线 | 上证/沪深300/中证500/中证1000/科创50/深证成指/创业板指/中小板指 | BaoStock | +| 股票列表 | 沪深A股代码、名称、上市日期 | AKShare | +| 交易日历 | 1990年至今的交易日列表 | AKShare | +| 日线行情 | 开高低收、成交量/额、振幅、涨跌幅、换手率(前复权,限沪深 A 股) | BaoStock 日线接口 | +| 新浪1分钟数据 | 最近 1 分钟K线(开高低收、成交量/额) | AKShare 新浪接口 | +| 指数日线 | 上证/沪深300/中证500/中证1000/科创50/深证成指/创业板指/中小板指 | 通达信本地缓存 | | 涨跌停统计 | 每日主板(10%)/ 科创创业板(20%)涨跌停数量 | stock_daily 汇总 | -| 季频财务指标 | 盈利能力、偿债能力、现金流(最近8个季度,JSON 存储) | BaoStock | -| 分红送转 | 每10股送转、派息、除权除息日(最近10年) | BaoStock | -| 分钟K线 | 5分钟K线(开高低收、成交量/额) | 通达信本地客户端(默认) / BaoStock(可选) | -| 概念板块 | 通达信本地概念板块缓存,缺失时回退同花顺 | 通达信本地 / 同花顺 | +| 概念板块 | 通达信本地概念板块缓存 | 通达信本地缓存 | -**已知限制**:BaoStock 不含北交所(920xxx)股票;新浪/腾讯/东方财富数据源对北交所同样不支持。 +**初始化说明**:首次运行 `--stock-info` / `--trading-day` 时会自动从 AKShare 拉取并写入数据库;`--daily` 现在改为 BaoStock 日线接口,只写 `stock_daily`;`--daily-no-data-only` 仅更新 `stock_no_data`;`--sector` / `--index` / `--intraday` 仍按通达信本地缓存工作,不再依赖 `hsjday.zip` 这类旧的日线初始化包。 ## 快速开始 @@ -28,13 +26,14 @@ A股数据抓取工具,以 [BaoStock](http://baostock.com) 为主、新浪/腾 ### 2. 安装依赖 ```bash -pip install baostock akshare pymysql sqlalchemy pyyaml pandas requests +pip install pymysql sqlalchemy pyyaml pandas requests baostock akshare ``` -> 说明:`baostock` 为主要数据源;`akshare` 用于新浪日线兜底和同花顺概念板块列表;`requests` 用于腾讯/同花顺概念板块抓取;通达信概念板块来自本地缓存文件。 -> 通达信分钟缓存刷新已自动按最多 100 只股票一批拆分;`--tdx-cache` 会先预热本地缓存,再全历史回填分钟表。`--tdx-local-cache` 会在预热后自动校验是否真的可读;若只想单独确认缓存是否能被读到,请用 `--tdx-verify-cache`。注意通达信客户端原生只支持刷新 `1m/5m` 本地缓存。 +> 说明:当前只需要 `pymysql`、`sqlalchemy`、`pyyaml`、`pandas`、`requests`、`baostock` 等基础依赖。 +> 如需使用 `--sina-min1` 命令,还需要 `akshare`。 +> 通达信分钟缓存刷新已自动按最多 100 只股票一批拆分。注意通达信客户端原生只支持刷新 `1m/5m` 本地缓存。 -> 通达信本地目录默认读取 `config.yaml` 里的 `tdx.dir`。 +> 仅 `--sector` / `--index` / `--intraday` 等命令需要通达信本地目录,默认读取 `config.yaml` 里的 `tdx.dir`。 ### 3. 配置数据库 @@ -75,94 +74,52 @@ CREATE DATABASE ashare CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci; ### 4. 运行 ```bash -# 1. 先抓取股票列表(其他模块依赖此数据) +# 1. 先抓取股票列表(数据来源: AKShare) python -m src.main --stock-info -# 2. 抓取交易日历(覆盖全历史,只需执行一次;此后每年末执行一次延长至未来) -python -m src.main --trading-day --start-date 19901219 --end-date 20261231 +# 2. 抓取交易日历(数据来源: AKShare) +python -m src.main --trading-day --start-date 20100101 --end-date 20261231 -# 3. 抓取全历史日线行情(首次,耗时较长) -python -m src.main --daily --start-date 19901201 --end-date 20260511 +# 3. 抓取全历史日线行情(数据来源: BaoStock 日线接口,限沪深 A 股) +python -m src.main --daily --start-date 20260521 --end-date 20260521 -# 4. 日常增量更新日线(默认走多源轮换:baostock/sina/tencent/eastmoney) +# 4. 日常增量更新日线(数据来源: BaoStock 日线接口,限沪深 A 股;只写 stock_daily) python -m src.main --daily -# 指定单一数据源(默认 all 表示多源轮换+自动切换) -python -m src.main --daily --source baostock -python -m src.main --daily --source sina +# 4.1 仅更新停牌/无数据记录,不写 stock_daily(数据来源: BaoStock 日线接口,限沪深 A 股) +python -m src.main --daily-no-data-only --start-date 20100101 --end-date 20260517 -# 抓取主要指数日线(上证/沪深300/中证500/中证1000/科创50/深证成指/创业板指/中小板指) -python -m src.main --index +# 5. 下载新浪 1 分钟数据并写入 stock_min1(数据来源: AKShare 新浪接口,默认 2010 至今) +python -m src.main --sina-min1 -# 抓取指数日线(指定日期范围) -python -m src.main --index --start-date 19901219 --end-date 20260511 +# 6. 下载新浪 1 分钟数据并指定区间(数据来源: AKShare 新浪接口) +python -m src.main --sina-min1 --start-date 20260501 --end-date 20260517 --symbol 600519 -# 抓取财务指标(全部股票,最近8个季度) -python -m src.main --financial - -# 抓取单只股票的财务数据 -python -m src.main --financial --symbol 000001 - -# 抓取分红送转(全部股票,最近10年) -python -m src.main --dividend - -# 抓取单只股票的分红 -python -m src.main --dividend --symbol 000001 - -# 抓取分钟K线(默认5分钟,近30天,默认走通达信本地客户端) -python -m src.main --intraday - -# 使用通达信本地客户端抓取分钟K线(默认) -python -m src.main --intraday --intraday-source tdx - -# 全历史回填通达信分钟数据(先预热本地缓存,再写入分钟表) -python -m src.main --tdx-cache - -# 仅预热通达信本地分钟缓存,并自动校验,不写数据库(原生只刷 1m/5m) -python -m src.main --tdx-local-cache - -# 只验证通达信本地分钟缓存是否可读 -python -m src.main --tdx-verify-cache - -# 抓取5分钟K线 -python -m src.main --intraday --start-date 20260101 - -# 抓取概念板块及成分股(通达信本地,缺失时回退同花顺) +# 7. 抓取概念板块及成分股(数据来源: 通达信本地缓存) python -m src.main --sector -``` -> ⚠️ **涨跌停统计 (`market_daily`)** 已实现于 `src/fetchers/market_daily.py`,但当前 `main.py` 尚未挂载 `--market-daily` 参数,暂只能在 Python 内直接调用 `fetch_market_daily(...)`。详见 [TODO.md](./TODO.md)。 +# 8. 汇总每日涨跌停统计(数据来源: stock_daily 汇总) +python -m src.main --market-daily +``` ### 5. 命令行参数说明 ``` 数据抓取选项: - --stock-info 抓取A股股票列表 - --trading-day 抓取交易日历 - --daily 抓取日线行情(增量;自动分析数据缺口,已完整自动跳过) - --source 日线数据源: baostock/sina/tencent/eastmoney/all(默认 all,轮换+失败自动切换) - --index 抓取主要指数日线 - --financial 抓取季频财务指标(增量;只抓缺失季度) - --dividend 抓取分红送转数据(增量;只抓缺失年份) - --intraday 抓取分钟K线行情(增量;只抓缺失日期) - --intraday-source 分钟K线数据源:tdx/baostock(默认 tdx) - --tdx-cache 全历史回填通达信5分钟数据(先预热本地缓存,再写入分钟表) - --tdx-local-cache 仅预热通达信本地分钟缓存,并自动校验,不写数据库(原生只刷 1m/5m) - --tdx-verify-cache 只验证通达信本地分钟缓存是否可读,不写数据库 - --freq K线频率: 5(默认 5) - --sector 抓取概念板块及成分股(增量;本地板块已齐全时跳过) - --market-daily 汇总每日涨跌停统计(从 stock_daily 聚合) + --stock-info 抓取A股股票列表(数据来源: AKShare) + --trading-day 抓取交易日历(数据来源: AKShare) + --daily 抓取日线行情(数据来源: BaoStock 日线接口;限沪深 A 股;增量,自动分析数据缺口,已完整自动跳过;只写入 stock_daily,不更新 stock_no_data) + --daily-no-data-only 仅更新 stock_no_data(数据来源: BaoStock 日线接口;限沪深 A 股;不写入 stock_daily) + --sina-min1 抓取新浪 1 分钟数据并写入 stock_min1(数据来源: AKShare 新浪接口;默认全市场、默认 2010 至今,可配合 --symbol / --start-date / --end-date) + --sector 抓取概念板块及成分股(数据来源: 通达信本地缓存;增量,本地板块已齐全时跳过) + --market-daily 汇总每日涨跌停统计(数据来源: stock_daily 聚合) + --symbol 指定单只或多个股票代码,逗号分隔;不传则抓取全市场,供 --sina-min1 使用 -日期过滤(对日线行情、交易日历、分钟K线、指数生效): +日期过滤(对日线行情、交易日历、新浪1分钟生效;新浪1分钟默认 2010 至今): --start-date 开始日期,格式 YYYYMMDD --end-date 结束日期,格式 YYYYMMDD - -股票过滤(对财务指标、分红送转、分钟K线生效): - --symbol 指定单只股票代码,如 000001,默认全部股票 ``` -> 注:`--market-daily` 在历史版本中存在,当前已下线(功能代码仍保留,但未挂到 main.py)。如需重新启用,参见 [TODO.md](./TODO.md) 中的 "重新挂载 market_daily 命令"。 - ## 数据库表结构 ### stock_info — 股票基本信息 @@ -179,66 +136,20 @@ python -m src.main --sector |------|------|------| | code | VARCHAR(10) | 股票代码 | | date | DATE | 交易日期 | -| open | FLOAT | 开盘价 | -| close | FLOAT | 收盘价 | -| high | FLOAT | 最高价 | -| low | FLOAT | 最低价 | -| volume | FLOAT | 成交量 | -| turnover | FLOAT | 成交额 | -| amplitude | FLOAT | 振幅% | -| pct_change | FLOAT | 涨跌幅% | -| change | FLOAT | 涨跌额 | -| turnover_rate | FLOAT | 换手率% | +| open | FLOAT NOT NULL | 开盘价 | +| close | FLOAT NOT NULL | 收盘价 | +| high | FLOAT NOT NULL | 最高价 | +| low | FLOAT NOT NULL | 最低价 | +| volume | FLOAT NOT NULL | 成交量 | +| turnover | FLOAT NOT NULL | 成交额 | +| amplitude | FLOAT NOT NULL | 振幅% | +| pct_change | FLOAT NOT NULL | 涨跌幅% | +| change | FLOAT NOT NULL | 涨跌额 | +| turnover_rate | FLOAT NOT NULL | 换手率% | 联合主键:`(code, date)` -### stock_financial_income — 季频盈利能力 - -| 字段 | 类型 | 说明 | -|------|------|------| -| code | VARCHAR(10) | 股票代码 | -| report_date | VARCHAR(20) | 报告期 | -| data | TEXT | JSON格式数据(ROE、净利率、毛利率等) | - -联合主键:`(code, report_date)` - -### stock_financial_balance — 季频偿债能力 - -| 字段 | 类型 | 说明 | -|------|------|------| -| code | VARCHAR(10) | 股票代码 | -| report_date | VARCHAR(20) | 报告期 | -| data | TEXT | JSON格式数据(流动比率、资产负债率等) | - -联合主键:`(code, report_date)` - -### stock_financial_cashflow — 季频现金流 - -| 字段 | 类型 | 说明 | -|------|------|------| -| code | VARCHAR(10) | 股票代码 | -| report_date | VARCHAR(20) | 报告期 | -| data | TEXT | JSON格式数据 | - -联合主键:`(code, report_date)` - -### stock_dividend — 分红送转 - -| 字段 | 类型 | 说明 | -|------|------|------| -| code | VARCHAR(10) | 股票代码 | -| name | VARCHAR(50) | 股票名称 | -| report_date | VARCHAR(20) | 报告期 | -| dividend_date | DATE | 除权除息日 | -| bonus_ratio | FLOAT | 每10股送转比例 | -| cash_div | FLOAT | 每10股派息 | -| convert_ratio | FLOAT | 每10股转增比例 | -| ex_right_date | DATE | 除权日 | -| dividend_yield | FLOAT | 股息率% | - -联合主键:`(code, report_date)` - -### trading_day — 交易日历 +### trading_day — 交易日历(AKShare) | 字段 | 类型 | 说明 | |------|------|------| @@ -268,12 +179,27 @@ python -m src.main --sector 联合主键:`(code, datetime)` +### stock_min1 — 1分钟K线 + +| 字段 | 类型 | 说明 | +|------|------|------| +| code | VARCHAR(10) | 股票代码 | +| datetime | DATETIME | 时间 | +| open | FLOAT | 开盘价 | +| high | FLOAT | 最高价 | +| low | FLOAT | 最低价 | +| close | FLOAT | 收盘价 | +| volume | FLOAT | 成交量 | +| amount | FLOAT | 成交额 | + +联合主键:`(code, datetime)` + ### stock_concept — 概念板块及成分股 | 字段 | 类型 | 说明 | |------|------|------| | code | VARCHAR(10) | 股票代码 | -| concept_code | VARCHAR(20) | 概念板块代码(同花顺 BK 编码) | +| concept_code | VARCHAR(20) | 概念板块代码 | | concept_name | VARCHAR(100) | 概念板块名称 | 联合主键:`(code, concept_code)` @@ -304,7 +230,7 @@ python -m src.main --sector | 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)。 +> 阈值:涨停 `pct_change >= 9.8`(主板)或 `>= 19.5`(20% 板块);跌停反向。当前可通过 `python -m src.main --market-daily` 运行 `src/fetchers/market_daily.py` 提供的 `fetch_market_daily()`。 ## 项目结构 @@ -320,21 +246,18 @@ ashare-data/ │ ├── __init__.py │ ├── config.py # 配置读取模块(YAML + 环境变量 ASHARE_CONFIG) │ ├── log.py # 统一 logging(控制台 + logs/ashare.log 按日滚动) -│ ├── baostock_conn.py # BaoStock 连接管理(login/logout/线程锁/超时重连) -│ ├── db.py # SQLAlchemy 模型 + 批量 upsert + 自动迁移 +│ ├── db.py # SQLAlchemy 模型 + 批量 upsert │ ├── main.py # 命令行入口 │ └── fetchers/ │ ├── __init__.py -│ ├── stock_list.py # 股票列表 -│ ├── trading_day.py # 交易日历 -│ ├── daily.py # 日线行情(多源轮换:baostock/sina/tencent/eastmoney) +│ ├── stock_list.py # 股票列表(AKShare) +│ ├── trading_day.py # 交易日历(AKShare) +│ ├── daily.py # 日线行情(BaoStock) │ ├── index.py # 主要指数日线 │ ├── market_daily.py # 每日涨跌停统计(从 stock_daily 聚合) -│ ├── financial.py # 季频财务指标(盈利/偿债/现金流,JSON 存储) -│ ├── dividend.py # 分红送转 -│ ├── intraday.py # 分钟K线(5/15/30/60) -│ └── sector.py # 概念板块及成分股(通达信本地 + 同花顺) -├── tests/ # pytest 测试(48 个用例,纯函数 + mock) +│ ├── sina_minute.py # 新浪 1 分钟数据(AKShare) +│ └── sector.py # 概念板块及成分股(通达信本地) +├── tests/ # pytest 测试(66 个用例,纯函数 + mock) ├── benchmarks/ # 并发度压测脚本(不在 CI 跑,需真实 MySQL+外网) └── gzl/ # 选股脚本(独立子项目,可选) ├── Selector.py @@ -343,14 +266,14 @@ ashare-data/ ## 设计说明 -- **多数据源容灾**:日线行情默认 `--source all` 在 BaoStock / 新浪 / 腾讯 / 东方财富之间轮换并自动切换,单源失败不影响整体进度;其他模块(财务、分红、指数)仍以 BaoStock 为主,分钟K线可选本地通达信客户端。 -- **线程安全**:BaoStock 的 `query_xxx()` 非线程安全,所有调用通过 `src/baostock_conn.py` 的全局锁串行化;查询超时/连接断开时自动重连。 +- **数据源分工**:股票列表、交易日历使用 AKShare;日线行情使用 BaoStock;概念板块、指数日线、分钟线等仍按功能读取通达信本地缓存;新浪 1 分钟数据使用 AKShare 新浪接口。 +- **线程安全**:通达信本地客户端调用集中在 `src/fetchers/tdx_client.py`。 - **去重写入**:所有表通过 `db.batch_upsert()` 走 MySQL `INSERT ON DUPLICATE KEY UPDATE`,重复执行不会产生重复数据。 -- **全量增量**:所有模块均支持增量更新——日线/指数/分钟K线按交易日对比找缺口,财务按缺失季度,分红按缺失年份,概念板块在本地已齐全时跳过。 +- **全量增量**:所有模块均支持增量更新——日线按交易日对比找缺口,概念板块在本地已齐全时跳过;新浪 1 分钟数据支持按日期过滤后落库。 - **停牌识别**:日线行情若两端已覆盖、内部仍有缺口,则视为停牌,不再重抓。 -- **表结构自动迁移**:`init_db()` 会检查并升级旧版 `stock_no_data` / `stock_sector` / `stock_intraday` 的列定义,无需手动改库。 +- **表结构初始化**:`init_db()` 只负责建表。 - **进程内缓存**:`get_stock_codes()` / `get_ipo_dates()` 缓存全量股票代码与上市日期,避免重复扫库。 -- **日志**:所有代码通过 `from src.log import get_logger` 输出;控制台 + `logs/ashare.log`(按日滚动,保留 7 天)双输出,通过 `ASHARE_LOG_LEVEL=DEBUG` 切换级别。BaoStock 查询签名走 DEBUG,控制台默认仅看到进度/异常。 +- **日志**:所有代码通过 `from src.log import get_logger` 输出;控制台 + `logs/ashare.log`(按日滚动,保留 7 天)双输出,通过 `ASHARE_LOG_LEVEL=DEBUG` 切换级别。 ## 运行测试 @@ -359,14 +282,14 @@ pip install pytest # 或 pip install -e ".[dev]" pytest -v ``` -当前 59 个用例全部通过,覆盖: -- `baostock_conn.code_to_bs` / `daily._code_to_*` 各数据源代码前缀映射 -- `daily._fill_derived_fields` 振幅/涨跌幅/涨跌额补算 -- `_fetch_tencent` / `_fetch_eastmoney` HTTP JSON 解析(mock requests,含空字段、异常包装、北交所短路) -- `financial._recent_quarters` 季度滚动跨年 +- 当前 66 个用例全部通过,覆盖: +- `daily._ak_hist_to_rows` / `daily._fill_derived_fields` +- `daily._fetch_one_stock` BaoStock 日线读取 +- `daily.fetch_daily` BaoStock 日线抓取与落库 +- `intraday._tdx_market_data_to_rows` / `intraday.fetch_intraday` 通达信分钟抓取 - `market_daily._is_20pct` 主板/创业板/科创板/北交所判定 - `tdx_blocks.load_infoharbor_blocks` 通达信本地板块缓存解析 -- `sector._fetch_concept_list_ths` / `sector._fetch_concept_stocks_ths` 同花顺概念板块解析 +- `sector._fetch_concept_list` / `sector._fetch_concept_stocks` 通达信本地概念板块解析 - `sector.fetch_sector` 概念板块增量跳过逻辑 - `src.log.get_logger` 命名空间、handler 幂等、env 控制 level diff --git a/TODO.md b/TODO.md index 3f6b806..981f509 100644 --- a/TODO.md +++ b/TODO.md @@ -1,189 +1,48 @@ # TODO -本文件用于跟踪 ashare-data 项目的已知问题与后续工作。维护时请保持「问题描述 + 影响范围 + 处理思路」三段式,便于他人接手。 - -> 🗓️ **2026-05-15 二轮清理**:上一轮(也是 2026-05-15)清理后剩下的 P2 中,可独立完成的 #4/#5/#7/#8 已落地(详见底部 [变更日志](#变更日志))。 -> 当前剩余条目均为:①需要用户决策的高风险动作(P0),②依赖外部数据源调研(P2-北交所),③需要与主项目协调的可选合并(P2-财务结构化、P2-gzl 接入)。 - -> 📚 **本文件分两部分**:上半部分是**已知缺陷/技术债**(P0/P2),下半部分是**新功能路线图**([Roadmap](#-新功能路线图-roadmap))。前者修,后者建。 +本文件只记录当前通达信本地缓存方案下,仍然值得跟进的事项。 --- -## 🔴 P0 — 需用户授权的高风险动作 +## 🔴 P0 - 需用户授权的高风险动作 -### 1. 清理 git 历史中的明文密码 ⚠️ 高风险 +### 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. 并发度压测:跑出实测数据 - -- **现状**:`benchmarks/bench_daily.py` 已就绪,可一键跑 `workers=1/2/4/8` 对照(详见 `benchmarks/README.md`)。 -- **下一步**:在低峰期跑一次完整压测,把推荐档位写到 `config.example.yaml` 注释里。脚本已自带 speedup 表输出,无需再写采集代码。 - -### 6. `gzl/` 选股脚本接入主项目 - -- **现状**:`gzl/Selector.py` + `gzl/select_stock.py` 读本地 CSV(不读 MySQL),与 `src/` 完全解耦,依赖 `scipy`。README 已说明其独立性。 +- **现状**:`config.yaml` 已加入 `.gitignore`,但历史 commit 中曾提交过明文数据库密码。 +- **影响**:仓库历史中仍可直接取到旧凭据。 - **处理思路**: - 1. 评估是否纳入主项目 — 如果只是个人玩具脚本可保持现状; - 2. 若纳入,迁移到 `src/strategies/`、改读 MySQL、复用 `get_session()` / `batch_upsert()` / `get_logger()`; - 3. `scipy` 加入 `pyproject.toml` 的 optional `[strategies]` extras。 + 1. 先轮换数据库密码。 + 2. 再用 `git filter-repo` 或 BFG 清理历史。 + 3. 必要时强制推送并通知协作者重新克隆。 --- -## 🚀 新功能路线图 (Roadmap) +## 🟢 P2 - 后续工程 -> 这一节是「想做但还没排期」的需求池,与上面的 P0/P2(已知缺陷/技术债)分开维护。 -> 立项时把对应条目挪到 P1/P2,附上责任人和预计动手时间;落地后再移到 [变更日志](#变更日志)。 +### 2. 北交所数据支持 -### ⭐ 推荐下一迭代(按"价值高 / 成本可控 / 与现有架构契合"排序) +- **现状**:当前主流程对 `920xxx` 仍然显式跳过。 +- **处理思路**:调研是否能从通达信本地缓存补齐北交所日线/分钟线,如可行,再补表结构与抓取逻辑。 -1. **抓取调度 + 告警**(详见下方「四、调度与监控」#1+#2)— 让项目从工具变服务,半天工作量 -2. **HTTP API 服务**(「三、服务化」#1)— FastAPI 暴露查询接口,半到一天 -3. **资金面三件套:龙虎榜 / 北向资金 / 融资融券**(「一、数据维度扩展」前 3 条)— 接口现成、量小,2-3 天补齐情绪+资金维度 +### 3. 数据质量校验日报 + +- **现状**:当前只有抓取过程日志,没有统一的每日质量报表。 +- **处理思路**:增加缺口统计、最新日期、异常波动、停牌覆盖率等摘要,便于盘后检查。 + +### 4. 抓取任务调度 + +- **现状**:目前仍依赖手动执行 CLI。 +- **处理思路**:接入 APScheduler 或系统定时任务,定时跑 `--daily`、`--intraday`、`--tdx-verify-cache`。 + +### 5. 质量与工程护栏 + +- **现状**:有 pytest,但还没有 CI 和迁移体系。 +- **处理思路**:补 GitHub Actions、Alembic、ruff / pre-commit 等基础工程化能力。 --- -### 一、数据维度扩展 +## 📌 维护建议 -「价值」= 对量化/选股的直接增益,「成本」= 实现规模 + 外部依赖复杂度。 - -| # | 需求 | 价值 | 成本 | 关键说明 | -|---|---|---|---|---| -| 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) | 低 | 简单看板:抓取状态、最新日期、表行数;非必需 | - ---- - -## 📋 维护节奏建议 - -- 每次新增 fetcher,请**同步更新** README 的:「功能概览表 / 运行示例 / 命令行参数说明 / 数据库表结构 / 项目结构」 五个章节。 -- 每次发现可复现 bug,先把现象写进本文件,再开始改代码,避免漏修。 -- 新增代码请用 `from src.log import get_logger`,不要再写 `print(..., flush=True)`。 -- 完成的条目移到下方 [变更日志](#变更日志),附完成日期,便于回顾。 -- **路线图条目立项时**:从 [Roadmap](#-新功能路线图-roadmap) 挪到 P1/P2,附责任人 + 预计动手时间;落地后再移到变更日志。 - ---- - -## 变更日志 - -### 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` 笔误 -- ✅ `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 条 +- 新增 TDX 相关抓取能力时,优先先写单测,再更新 README。 +- 若要调整本地缓存初始化方式,优先改 `src/fetchers/tdx_client.py`,避免在各个 fetcher 中重复处理下载逻辑。 +- 所有对外说明都以“通达信本地缓存”为唯一数据源口径。 diff --git a/config.example.yaml b/config.example.yaml index c4afdd5..a50128a 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -13,7 +13,7 @@ fetch: delay: 0.5 # 失败重试次数 retry: 3 - # 并发线程数;BaoStock 单源建议 1;多源轮换(baostock/sina/tencent/eastmoney)可适度提高至 2~4 + # 并发线程数;当前主流程以本地缓存为主,通常保持 1 即可 workers: 1 # 通达信本地客户端配置 diff --git a/pyproject.toml b/pyproject.toml index f2215c1..87ae4e1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,13 +4,13 @@ version = "0.1.0" description = "A股数据抓取,保存到MySQL数据库" requires-python = ">=3.10" dependencies = [ - "baostock", - "akshare", "pymysql", "sqlalchemy>=2.0", "pyyaml", "pandas", "requests", + "baostock>=0.9.1", + "akshare>=1.18.60", ] [project.optional-dependencies] diff --git a/requirements.txt b/requirements.txt index b44509c..570568f 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,7 +1,7 @@ -baostock -akshare pymysql sqlalchemy>=2.0 pyyaml pandas requests +baostock>=0.9.1 +akshare>=1.18.60 diff --git a/src/baostock_conn.py b/src/baostock_conn.py index 2fa36d1..947bdca 100644 --- a/src/baostock_conn.py +++ b/src/baostock_conn.py @@ -1,11 +1,14 @@ -"""BaoStock 连接管理器 — 集中管理 login/logout/线程锁/代码格式转换 +"""BaoStock 连接管理器 - 集中管理 login/logout/线程锁/代码格式转换。 -BaoStock 的 query_xxx() 非线程安全,所有查询需通过同一把锁串行化。 +BaoStock 的 query_xxx() 不是线程安全的,所有查询需通过同一把锁串行化。 """ +from __future__ import annotations + import threading import time from contextlib import contextmanager + import baostock as bs from src.log import get_logger @@ -19,7 +22,7 @@ MAX_RETRY = 3 def bs_login(): - """全局只 login 一次""" + """全局只 login 一次。""" global _logged_in with _lock: if not _logged_in: @@ -28,7 +31,7 @@ def bs_login(): def bs_logout(): - """程序退出时调用""" + """程序退出时调用。""" global _logged_in with _lock: if _logged_in: @@ -37,7 +40,7 @@ def bs_logout(): def _relogin(): - """断线重连(调用方需持有 _lock)""" + """断线重连(调用方需持有 _lock)。""" global _logged_in try: bs.logout() @@ -52,7 +55,7 @@ def _relogin(): @contextmanager def bs_query(query_fn, *args, **kwargs): - """加锁执行 BaoStock 查询,yield ResultData + """加锁执行 BaoStock 查询,yield ResultData。 连接断开或超时自动重连,最多重试 MAX_RETRY 次。 """ @@ -98,7 +101,7 @@ def bs_query(query_fn, *args, **kwargs): 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(): + if hasattr(rs, "error_code") and rs.error_code != "0" and "login" in rs.error_msg.lower(): _logger.warning("未登录,重连(第%d次)", attempt) _relogin() last_err = RuntimeError(f"BaoStock not logged in: {rs.error_msg}") @@ -110,10 +113,13 @@ def bs_query(query_fn, *args, **kwargs): raise last_err or RuntimeError("BaoStock query failed after retries") -def code_to_bs(code: str) -> str: - """纯数字代码转 BaoStock 格式: '600000' → 'sh.600000'""" - if code.startswith("920"): +def code_to_bs(code: str) -> str | None: + """纯数字代码转 BaoStock 格式:`600000` -> `sh.600000`。""" + code = (code or "").strip() + if not code or len(code) != 6 or not code.isdigit(): return None - if code.startswith(("6", "9")): + if code.startswith("6"): return f"sh.{code}" - return f"sz.{code}" + if code.startswith(("0", "3")): + return f"sz.{code}" + return None diff --git a/src/config.py b/src/config.py index 176c174..039eef8 100644 --- a/src/config.py +++ b/src/config.py @@ -59,3 +59,11 @@ def get_tdx_dir(config: dict | None = None) -> str | None: tdx_cfg = get_tdx_config(config) tdx_dir = tdx_cfg.get("dir") return str(tdx_dir) if tdx_dir else None + + +def get_tdx_bootstrap_url(config: dict | None = None) -> str: + """返回通达信日线缓存初始化包地址。""" + if config is None: + config = load_config() + tdx_cfg = config.get("tdx", {}) + return str(tdx_cfg.get("bootstrap_url") or "https://data.tdx.com.cn/vipdoc/hsjday.zip") diff --git a/src/db.py b/src/db.py index e6726eb..3b3b43d 100644 --- a/src/db.py +++ b/src/db.py @@ -1,14 +1,13 @@ """数据库模型定义与连接管理 -数据源:BaoStock 为主,新浪/腾讯/东方财富兜底 +数据源:按功能分别使用 Baostock / AKShare / 通达信本地缓存 表结构概览: - stock_info: 股票基本信息(含上市日期,用于跳过未上市股票) - stock_daily: 日线行情(含振幅/涨跌幅/换手率) - - stock_financial_income/balance/cashflow: 季频财务指标(JSON存储) - - stock_dividend: 分红送转 - trading_day: 交易日历(用于判断数据完整性) - stock_no_data: 无数据/停牌记录(避免重复抓取) - - stock_min5: 分钟K线 + - stock_min1: 1分钟K线 + - stock_min5: 5分钟K线 - index_daily: 主要指数日线 - market_daily: 每日涨跌停统计(10%/20% 板块分别计数) - stock_sector: 行业 + 地域分类 @@ -16,8 +15,8 @@ """ from sqlalchemy import ( - Column, String, Date, DateTime, Float, Integer, Text, - UniqueConstraint, Index, create_engine, MetaData, func, select, text, + Column, String, Date, DateTime, Float, Integer, + UniqueConstraint, Index, create_engine, func, select, text, ) from sqlalchemy.orm import DeclarativeBase, Session, sessionmaker from sqlalchemy.dialects.mysql import insert as mysql_insert @@ -52,55 +51,16 @@ class StockDaily(Base): id = Column(Integer, primary_key=True, autoincrement=True) code = Column(String(10), nullable=False, comment="股票代码") date = Column(Date, nullable=False, comment="交易日期") - open = Column(Float, comment="开盘价") - close = Column(Float, comment="收盘价") - high = Column(Float, comment="最高价") - low = Column(Float, comment="最低价") - volume = Column(Float, comment="成交量") - turnover = Column(Float, comment="成交额") - amplitude = Column(Float, comment="振幅%") - pct_change = Column(Float, comment="涨跌幅%") - change = Column(Float, comment="涨跌额") - turnover_rate = Column(Float, comment="换手率%") - - -# ── 季频盈利能力 ── -class FinancialIncome(Base): - __tablename__ = "stock_financial_income" - __table_args__ = ( - UniqueConstraint("code", "report_date", name="uq_income_code_date"), - ) - - id = Column(Integer, primary_key=True, autoincrement=True) - code = Column(String(10), nullable=False, comment="股票代码") - report_date = Column(String(20), nullable=False, comment="报告期") - data = Column(Text, comment="JSON格式盈利能力数据") - - -# ── 季频营运能力 ── -class FinancialBalance(Base): - __tablename__ = "stock_financial_balance" - __table_args__ = ( - UniqueConstraint("code", "report_date", name="uq_balance_code_date"), - ) - - id = Column(Integer, primary_key=True, autoincrement=True) - code = Column(String(10), nullable=False, comment="股票代码") - report_date = Column(String(20), nullable=False, comment="报告期") - data = Column(Text, comment="JSON格式营运能力数据") - - -# ── 季频现金流 ── -class FinancialCashflow(Base): - __tablename__ = "stock_financial_cashflow" - __table_args__ = ( - UniqueConstraint("code", "report_date", name="uq_cashflow_code_date"), - ) - - id = Column(Integer, primary_key=True, autoincrement=True) - code = Column(String(10), nullable=False, comment="股票代码") - report_date = Column(String(20), nullable=False, comment="报告期") - data = Column(Text, comment="JSON格式现金流数据") + open = Column(Float, nullable=False, comment="开盘价") + close = Column(Float, nullable=False, comment="收盘价") + high = Column(Float, nullable=False, comment="最高价") + low = Column(Float, nullable=False, comment="最低价") + volume = Column(Float, nullable=False, comment="成交量") + turnover = Column(Float, nullable=False, comment="成交额") + amplitude = Column(Float, nullable=False, comment="振幅%") + pct_change = Column(Float, nullable=False, comment="涨跌幅%") + change = Column(Float, nullable=False, comment="涨跌额") + turnover_rate = Column(Float, nullable=False, comment="换手率%") # ── 指数日线行情 ── @@ -123,25 +83,6 @@ class IndexDaily(Base): pct_change = Column(Float, comment="涨跌幅%") -# ── 分红送转 ── -class StockDividend(Base): - __tablename__ = "stock_dividend" - __table_args__ = ( - UniqueConstraint("code", "report_date", name="uq_dividend_code_date"), - ) - - id = Column(Integer, primary_key=True, autoincrement=True) - code = Column(String(10), nullable=False, comment="股票代码") - name = Column(String(50), comment="股票名称") - report_date = Column(String(20), nullable=False, comment="报告期") - dividend_date = Column(Date, comment="除权除息日") - bonus_ratio = Column(Float, comment="每10股送转比例") - cash_div = Column(Float, comment="每10股派息") - convert_ratio = Column(Float, comment="每10股转增比例") - ex_right_date = Column(Date, comment="除权日") - dividend_yield = Column(Float, comment="股息率%") - - # ── 交易日历 ── class TradingDay(Base): __tablename__ = "trading_day" @@ -166,6 +107,25 @@ class StockNoData(Base): created_at = Column(DateTime, server_default=func.now(), comment="记录时间") +# ── 5分钟K线 ── +class StockMin1(Base): + __tablename__ = "stock_min1" + __table_args__ = ( + UniqueConstraint("code", "datetime", name="uq_min1_code_dt"), + Index("ix_min1_datetime", "datetime"), + ) + + id = Column(Integer, primary_key=True, autoincrement=True) + code = Column(String(10), nullable=False, comment="股票代码") + datetime = Column(DateTime, nullable=False, comment="时间") + open = Column(Float, comment="开盘价") + high = Column(Float, comment="最高价") + low = Column(Float, comment="最低价") + close = Column(Float, comment="收盘价") + volume = Column(Float, comment="成交量") + amount = Column(Float, comment="成交额") + + # ── 5分钟K线 ── class StockMin5(Base): __tablename__ = "stock_min5" @@ -244,6 +204,18 @@ def get_stock_codes() -> list[str]: return _stock_codes_cache +def invalidate_stock_codes_cache() -> None: + """失效股票代码缓存,便于股票列表更新后重新读取。""" + global _stock_codes_cache + _stock_codes_cache = None + + +def invalidate_ipo_dates_cache() -> None: + """失效上市日期缓存,便于股票列表更新后重新读取。""" + global _ipo_dates_cache + _ipo_dates_cache = None + + def get_ipo_dates() -> dict[str, str]: """获取全部股票上市日期(进程内缓存,避免重复查询)""" global _ipo_dates_cache @@ -273,33 +245,6 @@ def get_session() -> Session: def init_db(): engine = get_engine() - # 自动迁移:旧版 stock_no_data 使用 date_range 列,新版改为 date 列 - try: - with engine.connect() as conn: - result = conn.execute(text("SHOW COLUMNS FROM stock_no_data LIKE 'date_range'")) - if result.fetchone(): - conn.execute(text("DROP TABLE stock_no_data")) - conn.commit() - _logger.info("stock_no_data 表结构已升级(date_range → date)") - except Exception: - pass - # 自动迁移:旧版 stock_sector 使用 update_date 列,新版改为 region - try: - with engine.connect() as conn: - result = conn.execute(text("SHOW COLUMNS FROM stock_sector LIKE 'update_date'")) - if result.fetchone(): - conn.execute(text("DROP TABLE stock_sector")) - conn.commit() - _logger.info("stock_sector 表结构已升级(新增 region 列)") - except Exception: - pass - # 自动迁移:旧版 stock_intraday 单表 → 四张分表 - try: - with engine.connect() as conn: - conn.execute(text("DROP TABLE IF EXISTS stock_intraday")) - conn.commit() - except Exception: - pass Base.metadata.create_all(engine) _logger.info("数据库表初始化完成") diff --git a/src/fetchers/daily.py b/src/fetchers/daily.py index 6f70858..458affb 100644 --- a/src/fetchers/daily.py +++ b/src/fetchers/daily.py @@ -1,35 +1,33 @@ -"""日线行情抓取模块 — 多数据源 +"""日线行情抓取模块 - BaoStock。""" -数据源优先级: - - baostock: BaoStock(默认,含换手率/振幅) - - sina: 新浪财经(通过 akshare) - - tencent: 腾讯财经 - - eastmoney: 东方财富 - -增量抓取:一次本地查询确定每只股票的缺口范围,只请求缺失日期段。 -""" +from __future__ import annotations +import random import time from datetime import datetime, timedelta -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, bs_logout +import pandas as pd + +from src.baostock_conn import bs_query, code_to_bs 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.db import ( + StockDaily, + StockInfo, + StockNoData, + TradingDay, + batch_upsert, + get_ipo_dates, + get_session, + get_stock_codes, + invalidate_ipo_dates_cache, + invalidate_stock_codes_cache, +) from src.log import get_logger -from sqlalchemy import select, func, text +from sqlalchemy import select, text _logger = get_logger("daily") - -VALID_SOURCES = ("baostock", "sina", "tencent", "eastmoney") - -_HEADERS = { - "User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36", - "Referer": "https://quote.eastmoney.com/", -} +ak = None def _clean(val): @@ -40,11 +38,37 @@ def _clean(val): return val -def _fill_derived_fields(rows: list[dict]) -> list[dict]: - """补算缺失的振幅、涨跌幅、涨跌额、换手率。 +def _to_float(val): + val = _clean(val) + if val is None: + return None + try: + return float(val) + except (TypeError, ValueError): + return None - rows 按日期升序,利用前一行的 close 作为 preclose。 - """ + +def _parse_ipo_date(raw: str | None): + if not raw: + return None + raw = str(raw).strip() + if not raw: + return None + if len(raw) >= 10 and raw[4] == "-" and raw[7] == "-": + try: + return datetime.strptime(raw[:10], "%Y-%m-%d").date() + except ValueError: + return None + if len(raw) == 8 and raw.isdigit(): + try: + return datetime.strptime(raw, "%Y%m%d").date() + except ValueError: + return None + return None + + +def _fill_derived_fields(rows: list[dict]) -> list[dict]: + """补算缺失的振幅、涨跌幅、涨跌额。""" if not rows: return rows for i, row in enumerate(rows): @@ -65,237 +89,404 @@ def _fill_derived_fields(rows: list[dict]) -> list[dict]: return rows -# ==================== 数据源: BaoStock ==================== +def _is_complete_daily_row(row: dict) -> bool: + """判断日线行是否满足 stock_daily 的全字段非空要求。""" + required_fields = ( + "open", + "close", + "high", + "low", + "volume", + "turnover", + "amplitude", + "pct_change", + "change", + "turnover_rate", + ) + return all(row.get(field) is not None for field in required_fields) -def _fetch_baostock(code: str, start_date: str, end_date: str) -> list[dict] | None: - bs_code = code_to_bs(code) - if not bs_code: - return None + +def _normalize_code(code: str) -> str: + """把输入代码统一成数据库里的纯 6 位代码。""" + code = (code or "").strip() + if not code: + return code + upper = code.upper() + if "." in upper: + base, suffix = upper.split(".", 1) + if suffix in {"SH", "SZ", "BJ"}: + return base + lower = code.lower() + if lower.startswith(("sh", "sz", "bj")) and len(lower) > 2: + return lower[2:] + return code + + +def _is_supported_daily_code(code: str) -> bool: + """BaoStock 日线当前只稳定支持沪深 A 股代码。""" + clean_code = _normalize_code(code) + return len(clean_code) == 6 and clean_code.isdigit() and clean_code[0] in {"0", "3", "6"} + + +def _split_supported_daily_codes(codes: list[str]) -> tuple[list[str], list[str]]: + """把可抓取和应跳过的日线代码分开。""" + supported: list[str] = [] + skipped: list[str] = [] + for code in codes: + if _is_supported_daily_code(code): + supported.append(code) + else: + skipped.append(code) + return supported, skipped + + +def _preview_codes(codes: list[str], limit: int = 5) -> str: + """把代码列表压缩成日志预览字符串。""" + if not codes: + return "[]" + preview = codes[:limit] + suffix = "" + if len(codes) > limit: + suffix = f", ...(+{len(codes) - limit})" + return "[" + ", ".join(preview) + suffix + "]" + + +def _mark_no_data_days(code: str, days: list[str]) -> int: + """把某只股票缺失的交易日写入停牌/无数据表。""" + if not days: + return 0 + + rows = [{"code": code, "date": day} for day in days] + batch_upsert(StockNoData, rows, ["code", "date"]) + return len(rows) + + +def _get_no_data_dates(start_date: str, end_date: str) -> dict[str, set[str]]: + """读取区间内已标记为停牌/无数据的日期。""" 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]}" - 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 + session = get_session() + try: + result = session.execute( + select(StockNoData.code, StockNoData.date) + .where(StockNoData.date >= sd) + .where(StockNoData.date <= ed) ) + no_data: dict[str, set[str]] = {} + for code, day in result: + no_data.setdefault(code, set()).add(str(day)) + return no_data + finally: + session.close() + + +def _hist_df_to_rows(code: str, hist_df: pd.DataFrame) -> list[dict]: + """把表格行情转换为数据库行。""" + if hist_df is None or hist_df.empty: + return [] + + rows: list[dict] = [] + for _, item in hist_df.iterrows(): + date_value = item.get("日期") + if pd.isna(date_value): + continue + date = pd.to_datetime(date_value, errors="coerce") + if pd.isna(date): + continue + + row = {"code": code, "date": date.strftime("%Y-%m-%d")} + valid_value = False + field_map = { + "开盘": "open", + "收盘": "close", + "最高": "high", + "最低": "low", + "成交量": "volume", + "成交额": "turnover", + "振幅": "amplitude", + "涨跌幅": "pct_change", + "涨跌额": "change", + "换手率": "turnover_rate", + } + for cn_field, db_field in field_map.items(): + value = item.get(cn_field) + if pd.isna(value): + value = None + elif value is not None: + value = pd.to_numeric(value, errors="coerce") + if pd.isna(value): + value = None + if value is not None: + valid_value = True + row[db_field] = value + if valid_value and row.get("volume") is not None: + rows.append(row) + + rows = _fill_derived_fields(rows) + return [row for row in rows if _is_complete_daily_row(row)] + + +def _ak_hist_to_rows(code: str, hist_df: pd.DataFrame) -> list[dict]: + """兼容旧测试/调用的历史数据表转换函数。""" + return _hist_df_to_rows(code, hist_df) + + +def _bs_result_to_rows(code: str, rs) -> list[dict]: + """把 BaoStock 查询结果整理为数据库行。""" + rows: list[dict] = [] + while rs.error_code == "0" and rs.next(): + r = rs.get_row_data() + preclose = _to_float(r[7]) + pct_chg = _to_float(r[8]) + turn = _to_float(r[9]) + close = _to_float(r[4]) + high = _to_float(r[2]) + low = _to_float(r[3]) + amp = None + if high is not None and low is not None and preclose is not None and preclose != 0: + amp = round((high - low) / preclose * 100, 2) + chg = None + if close is not None and preclose is not None: + chg = round(close - preclose, 3) + + rows.append({ + "code": code, + "date": r[0], + "open": _to_float(r[1]), + "high": high, + "low": low, + "close": close, + "volume": _to_float(r[5]), + "turnover": _to_float(r[6]), + "amplitude": amp, + "pct_change": pct_chg, + "change": chg, + "turnover_rate": turn, + }) + if rows[-1]["volume"] is None: + rows.pop() + rows = _fill_derived_fields(rows) + return [row for row in rows if _is_complete_daily_row(row)] + + +def _fetch_stock_list_rows(*, retry: int = 3, delay: float = 0.1) -> list[dict]: + """从 BaoStock 获取沪深 A 股代码与名称。""" + retry = max(int(retry), 1) + delay = max(float(delay), 0.0) + last_exc: Exception | None = None 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): - _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): - _logger.warning("[BaoStock] %s 请求失败,重试(%d/%d)", code, attempt, retry) - bs_logout() - time.sleep(min(2 * attempt, 5)) - continue - _logger.error("[BaoStock] %s 获取失败: %s", code, e) - return None + with bs_query(bs.query_stock_basic, code="", code_name="") as rs: + if rs.error_code != "0": + raise RuntimeError(rs.error_msg) + rows: list[dict] = [] + while rs.next(): + code, name, ipo_date, _out_date, sec_type, status = rs.get_row_data() + if str(sec_type) != "1" or str(status) != "1": + continue + base_code = _normalize_code(code) + if not _is_supported_daily_code(base_code): + continue + item = {"code": base_code, "name": str(name).strip()} + parsed_ipo = _parse_ipo_date(ipo_date) + if parsed_ipo is not None: + item["ipo_date"] = parsed_ipo + rows.append(item) + return rows + except Exception as exc: + last_exc = exc + if attempt >= retry: + break + wait_seconds = delay * attempt if delay > 0 else 0.5 * attempt + wait_seconds += random.uniform(0, min(wait_seconds * 0.2, 0.5)) + _logger.warning("BaoStock 获取股票列表第 %d 次失败,%.1fs 后重试: %s", attempt, wait_seconds, exc) + time.sleep(wait_seconds) + + if last_exc is not None: + _logger.warning("BaoStock 获取股票列表失败: %s", last_exc) + return [] -# ==================== 数据源: 新浪 (akshare) ==================== +def _bootstrap_stock_list(*, retry: int = 3, delay: float = 0.1) -> list[str]: + """从 BaoStock 抓取股票基本资料并保存。""" + rows = _fetch_stock_list_rows(retry=retry, delay=delay) + if not rows: + return [] -def _code_to_sina(code: str) -> str | None: - if code.startswith("920"): - return None - if code.startswith(("6", "9")): - return f"sh{code}" - return f"sz{code}" + batch_upsert(StockInfo, rows, ["code"]) + invalidate_stock_codes_cache() + invalidate_ipo_dates_cache() + _logger.info("BaoStock 股票列表已保存,共 %d 只", len(rows)) + return [row["code"] for row in rows] -def _fetch_sina(code: str, start_date: str, end_date: str) -> list[dict] | None: - sina_code = _code_to_sina(code) - if not sina_code: - 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]}" +def _bootstrap_trading_days(start_date: str, end_date: str, *, retry: int = 3, delay: float = 0.1) -> list[str]: + """从 BaoStock 抓取指定区间的交易日并保存。""" + retry = max(int(retry), 1) + delay = max(float(delay), 0.0) + sd = _format_date(start_date) + ed = _format_date(end_date) + last_exc: Exception | None = None + + for attempt in range(1, retry + 1): + try: + with bs_query(bs.query_trade_dates, start_date=sd, end_date=ed) as rs: + if rs.error_code != "0": + raise RuntimeError(rs.error_msg) + days: list[str] = [] + while rs.next(): + day, is_trading_day = rs.get_row_data() + if str(is_trading_day) == "1": + days.append(str(day)) + if days: + batch_upsert(TradingDay, [{"date": d} for d in days], ["date"]) + _logger.info("BaoStock 交易日历已保存,%d 个交易日", len(days)) + return days + except Exception as exc: + last_exc = exc + if attempt >= retry: + break + wait_seconds = delay * attempt if delay > 0 else 0.5 * attempt + wait_seconds += random.uniform(0, min(wait_seconds * 0.2, 0.5)) + _logger.warning("BaoStock 获取交易日历第 %d 次失败,%.1fs 后重试: %s", attempt, wait_seconds, exc) + time.sleep(wait_seconds) + + if last_exc is not None: + _logger.warning("BaoStock 获取交易日历失败: %s", last_exc) + return [] + + +def _format_date(d: str) -> str: + return f"{d[:4]}-{d[4:6]}-{d[6:8]}" + + +def get_trading_days(start_date: str, end_date: str) -> list[str]: + """获取指定范围内的交易日列表。""" + sd = _format_date(start_date) + ed = _format_date(end_date) + + session = get_session() try: - import akshare as ak - df = ak.stock_zh_a_daily(symbol=sina_code, start_date=sd, end_date=ed, adjust="qfq") - if df is None or df.empty: - return None - rows = [] - for _, r in df.iterrows(): - date_str = str(r["date"])[:10] - close = float(r["close"]) - open_ = float(r["open"]) - high = float(r["high"]) - low = float(r["low"]) - volume = float(r["volume"]) if "volume" in r else None - amount = float(r["amount"]) if "amount" in r else None - turnover_rate = float(r["turnover"]) if "turnover" in r else None - rows.append({ - "code": code, "date": date_str, - "open": open_, "high": high, "low": low, "close": close, - "volume": volume, "turnover": amount, - "amplitude": None, "pct_change": None, "change": None, - "turnover_rate": turnover_rate, - }) - return rows if rows else None - except Exception as e: - raise RuntimeError(f"新浪请求失败: {e}") from e + result = session.execute( + select(TradingDay.date) + .where(TradingDay.date >= sd) + .where(TradingDay.date <= ed) + .order_by(TradingDay.date) + ) + days = [str(row[0]) for row in result] + + if days and days[0] <= sd and days[-1] >= ed: + return days + finally: + session.close() + + fetched = _bootstrap_trading_days(start_date, end_date) + if not fetched: + return days if days else [] + return sorted(set(days + fetched)) -# ==================== 数据源: 腾讯 ==================== - -def _code_to_tencent(code: str) -> str | None: - if code.startswith("920"): - return None - if code.startswith(("6", "9")): - return f"sh{code}" - return f"sz{code}" +def _ensure_stock_codes(*, retry: int, delay: float) -> list[str]: + """确保股票列表存在,必要时从 BaoStock 初始化。""" + codes = get_stock_codes() + if codes: + return codes + codes = _bootstrap_stock_list(retry=retry, delay=delay) + return codes -def _fetch_tencent(code: str, start_date: str, end_date: str) -> list[dict] | None: - tc_code = _code_to_tencent(code) - if not tc_code: - 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: - url = "https://web.ifzq.gtimg.cn/appstock/app/fqkline/get" - params = {"param": f"{tc_code},day,{sd},{ed},640,qfq"} - r = requests.get(url, params=params, timeout=15) - data = r.json().get("data", {}) - stock_data = data.get(tc_code, {}) - klines = stock_data.get("qfqday") or stock_data.get("day") - if not klines: - return None - rows = [] - for k in klines: - # [date, open, close, high, low, volume] - date_str = k[0] - open_ = float(k[1]) - close = float(k[2]) - high = float(k[3]) - low = float(k[4]) - volume = float(k[5]) if len(k) > 5 else None - rows.append({ - "code": code, "date": date_str, - "open": open_, "high": high, "low": low, "close": close, - "volume": volume, "turnover": None, - "amplitude": None, "pct_change": None, "change": None, - "turnover_rate": None, - }) - return rows if rows else None - except Exception as e: - raise RuntimeError(f"腾讯请求失败: {e}") from e +def _fetch_one_stock(code: str, gap_start: str, gap_end: str) -> list[dict]: + """读取单只股票一个缺口区间的日线数据。""" + bs_code = code_to_bs(code) + if not bs_code: + return [] + + sd = gap_start.replace("-", "") + ed = gap_end.replace("-", "") + start_date = _format_date(sd) + end_date = _format_date(ed) + with bs_query( + bs.query_history_k_data_plus, + bs_code, + "date,open,high,low,close,volume,amount,preclose,pctChg,turn", + start_date=start_date, + end_date=end_date, + frequency="d", + adjustflag="2", + ) as rs: + if rs.error_code != "0": + raise RuntimeError(rs.error_msg) + return _bs_result_to_rows(code, rs) -# ==================== 数据源: 东方财富 ==================== +def _fetch_one_stock_with_retry( + code: str, + gap_start: str, + gap_end: str, + *, + retry: int = 3, + delay: float = 0.1, +) -> list[dict]: + """带重试的单股日线抓取。""" + retry = max(int(retry), 1) + delay = max(float(delay), 0.0) + last_exc: Exception | None = None -def _code_to_eastmoney(code: str) -> str | None: - if code.startswith("920"): - return None - if code.startswith(("6", "9")): - return f"1.{code}" - return f"0.{code}" + for attempt in range(1, retry + 1): + try: + return _fetch_one_stock(code, gap_start, gap_end) + except Exception as exc: + last_exc = exc + if attempt >= retry: + raise + + wait_seconds = delay * attempt if delay > 0 else 0.5 * attempt + wait_seconds += random.uniform(0, min(wait_seconds * 0.2, 0.5)) + _logger.warning( + "[BaoStock] %s 第 %d 次失败,%.1fs 后重试: %s", + code, + attempt, + wait_seconds, + exc, + ) + time.sleep(wait_seconds) + + if last_exc is not None: + raise last_exc + return [] -def _fetch_eastmoney(code: str, start_date: str, end_date: str) -> list[dict] | None: - em_code = _code_to_eastmoney(code) - if not em_code: - 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: - url = "https://push2his.eastmoney.com/api/qt/stock/kline/get" - params = { - "secid": em_code, - "fields1": "f1,f2,f3,f4,f5,f6", - "fields2": "f51,f52,f53,f54,f55,f56,f57,f58,f59,f60,f61", - "klt": 101, "fqt": 1, "beg": sd, "end": ed, - "ut": "fa5fd1943c7b386f172d6893dbfba10b", - } - r = requests.get(url, params=params, headers=_HEADERS, timeout=15) - data = r.json().get("data") or {} - klines = data.get("klines") or [] - if not klines: - return None - rows = [] - for line in klines: - # date,open,close,high,low,volume,amount,amplitude,pct_change,change,turnover_rate - parts = line.split(",") - rows.append({ - "code": code, "date": parts[0], - "open": float(parts[1]), "high": float(parts[3]), - "low": float(parts[4]), "close": float(parts[2]), - "volume": float(parts[5]), "turnover": float(parts[6]), - "amplitude": float(parts[7]) if parts[7] != "" else None, - "pct_change": float(parts[8]) if parts[8] != "" else None, - "change": float(parts[9]) if parts[9] != "" else None, - "turnover_rate": float(parts[10]) if parts[10] != "" else None, - }) - return rows if rows else None - except Exception as e: - raise RuntimeError(f"东方财富请求失败: {e}") from e +def _fetch_missing_daily_rows( + code: str, + missing_days: list[str], + *, + retry: int, + delay: float, +) -> list[dict]: + """把范围查询漏掉的交易日按天补抓回来。""" + if not missing_days: + return [] + patched_rows: list[dict] = [] + for day in missing_days: + day_key = day.replace("-", "") + rows = _fetch_one_stock_with_retry( + code, + day_key, + day_key, + retry=retry, + delay=delay, + ) + if rows: + patched_rows.extend(rows) + return patched_rows -# ==================== 数据源分发 ==================== - -_SOURCE_FN = { - "baostock": _fetch_baostock, - "sina": _fetch_sina, - "tencent": _fetch_tencent, - "eastmoney": _fetch_eastmoney, -} - -_SOURCE_LABEL = { - "baostock": "BaoStock", - "sina": "新浪", - "tencent": "腾讯", - "eastmoney": "东方财富", -} - - -# ==================== 缺口分析 ==================== def _analyze_gaps(codes: list[str], start_date: str, end_date: str, trading_days: list[str]) -> dict[str, list[str]]: - """一条SQL统计每只股票每月行情数,与交易日对比找缺口。""" + """一条 SQL 统计每只股票每月行情数,与交易日对比找缺口。""" if not trading_days: return {} @@ -303,48 +494,63 @@ def _analyze_gaps(codes: list[str], start_date: str, end_date: str, ed = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}" t0 = time.time() - - # 上市日期(进程内缓存,只查一次) ipo_dates = get_ipo_dates() - _logger.info("[1/3] 上市日期查询完成 %d 只 %.1fs", len(ipo_dates), time.time()-t0) + _logger.info("[1/3] 上市日期查询完成 %d 只 %.1fs", len(ipo_dates), time.time() - t0) - # 按月统计每只股票行情数(一条SQL) t1 = time.time() session = get_session() try: sql = text(""" - SELECT code, DATE_FORMAT(date, '%Y-%m') AS month, COUNT(*) AS cnt + SELECT + code, + DATE_FORMAT(date, '%Y-%m') AS month, + COUNT(*) AS cnt, + SUM( + CASE + WHEN open IS NULL + OR close IS NULL + OR high IS NULL + OR low IS NULL + OR volume IS NULL + OR turnover IS NULL + THEN 1 ELSE 0 + END + ) AS incomplete_cnt FROM stock_daily WHERE date >= :sd AND date <= :ed GROUP BY code, DATE_FORMAT(date, '%Y-%m') """) result = session.execute(sql, {"sd": sd, "ed": ed}) code_month_cnt: dict[str, dict[str, int]] = {} - for code, month, cnt in result: - code_month_cnt.setdefault(code, {})[month] = cnt + quality_gap_codes: set[str] = set() + for code, month, cnt, incomplete_cnt in result: + code_month_cnt.setdefault(code, {})[month] = int(cnt or 0) + if int(incomplete_cnt or 0) > 0: + quality_gap_codes.add(code) finally: session.close() - _logger.info("[2/3] 行情按月统计完成 %.1fs", time.time()-t1) + _logger.info("[2/3] 行情按月统计完成 %.1fs", time.time() - t1) + + no_data_dates = _get_no_data_dates(start_date, end_date) - # 按月对比找缺口 t2 = time.time() td_by_month: dict[str, list[str]] = {} for d in trading_days: td_by_month.setdefault(d[:7], []).append(d) gap_codes: set[str] = set() - no_ipo_codes: set[str] = set() for month_key, month_days in sorted(td_by_month.items()): for code in codes: if code in gap_codes: continue + if code in quality_gap_codes: + gap_codes.add(code) + continue ipo = ipo_dates.get(code) - if not ipo: - no_ipo_codes.add(code) - continue - if ipo > month_days[-1]: - continue - expected = [d for d in month_days if d >= ipo] + expected = [ + d for d in month_days + if (not ipo or d >= ipo) and d not in no_data_dates.get(code, set()) + ] if not expected: continue cnt = code_month_cnt.get(code, {}).get(month_key, 0) @@ -352,217 +558,212 @@ def _analyze_gaps(codes: list[str], start_date: str, end_date: str, gap_codes.add(code) _logger.info( - "[3/3] 缺口对比完成 缺口股票:%d 只 跳过(无上市日期):%d 只 %.1fs", - len(gap_codes), len(no_ipo_codes), time.time()-t2, + "[3/3] 缺口对比完成 缺口股票:%d 只 %.1fs", + len(gap_codes), time.time() - t2, ) if not gap_codes: return {} - # 确定缺口范围 gaps: dict[str, list[str]] = {} - for code in gap_codes: - ipo = ipo_dates.get(code) - expected = [d for d in trading_days if not ipo or d >= ipo] - session = get_session() - try: - minmax = session.execute( - select(func.min(StockDaily.date), func.max(StockDaily.date)) - .where(StockDaily.code == code) - .where(StockDaily.date >= sd) - .where(StockDaily.date <= ed) - ).fetchone() - finally: - session.close() + session = get_session() + try: + for code in gap_codes: + ipo = ipo_dates.get(code) + expected = [ + d for d in trading_days + if (not ipo or d >= ipo) and d not in no_data_dates.get(code, set()) + ] + if not expected: + continue - if minmax and minmax[0]: - min_d, max_d = str(minmax[0]), str(minmax[1]) - front = [d for d in expected if d < min_d] - back = [d for d in expected if d > max_d] - if front and back: + if code in quality_gap_codes: gaps[code] = [expected[0], expected[-1]] - elif front: - gaps[code] = [front[0], front[-1]] - elif back: - gaps[code] = [back[0], back[-1]] - # else: 数据两端已覆盖,内部缺口属停牌,不重抓 - else: - gaps[code] = [expected[0], expected[-1]] + continue + + existing_dates = { + str(row[0]) for row in session.execute( + select(StockDaily.date) + .where(StockDaily.code == code) + .where(StockDaily.date >= expected[0]) + .where(StockDaily.date <= expected[-1]) + ) + } + missing = [d for d in expected if d not in existing_dates] + if missing: + gaps[code] = [missing[0], missing[-1]] + finally: + session.close() return gaps -# ==================== 单股票抓取+写入 ==================== - -_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_keys: list[str]) -> tuple[str, str, int, bool]: - """抓取单只股票并写入,返回 (code, source_label, row_count, success)""" - 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: - _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) - _logger.warning("[%s] %s 无数据,继续尝试: %s", label, code, next_names) - - return code, last_label, 0, False - - -# ==================== 主函数 ==================== - -def fetch_daily(start_date: str | None = None, end_date: str | None = None, - source: str = "all"): - # 确定使用的数据源列表 - if source == "all": - sources = list(VALID_SOURCES) - elif source in VALID_SOURCES: - sources = [source] - else: - _logger.error("不支持的数据源 %s,可选: %s, all", source, ", ".join(VALID_SOURCES)) - return - +def fetch_daily( + start_date: str | None = None, + end_date: str | None = None, + *, + no_data_only: bool = False, +): + """抓取 BaoStock 日线行情。""" 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: - _logger.error("无股票列表,请先运行 --stock-info") - return + retry = max(int(cfg.get("retry", 3)), 3) + delay = max(float(cfg.get("delay", 0.1)), 0.2) if end_date is None: end_date = datetime.now().strftime("%Y%m%d") if start_date is None: start_date = (datetime.now() - timedelta(days=30)).strftime("%Y%m%d") - if "baostock" in sources: - bs_login() + codes = _ensure_stock_codes(retry=retry, delay=delay) + if not codes: + _logger.error("无股票列表,无法抓取日线数据") + return - # 优先使用本地交易日历;若覆盖不完整则自动补齐 trading_days = get_trading_days(start_date, end_date) td_count = len(trading_days) - _logger.info( - "[%s] [并发:%d] 交易日历: %s ~ %s 共 %d 个交易日", - source_names, workers, start_date, end_date, td_count, - ) + _logger.info("[BaoStock] 交易日历: %s ~ %s 共 %d 个交易日", start_date, end_date, td_count) + + if not trading_days: + _logger.error("无交易日数据,无法抓取日线数据") + return + + no_data_dates = _get_no_data_dates(start_date, end_date) - # 排除未上市股票 ed_fmt = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}" + ed_date = datetime.strptime(ed_fmt, "%Y-%m-%d").date() session = get_session() try: not_listed = {row[0] for row in session.execute( - select(StockInfo.code).where(StockInfo.ipo_date > ed_fmt) + select(StockInfo.code).where(StockInfo.ipo_date > ed_date) )} finally: session.close() codes = [c for c in codes if c not in not_listed] - if not_listed: - _logger.info("[%s] %d 只股票未上市,跳过", source_names, len(not_listed)) + _logger.info("[BaoStock] %d 只股票未上市,跳过", len(not_listed)) - # 一次查询:判断完整 + 计算缺口 - _logger.info("[%s] 正在分析 %d 只股票的数据缺口...", source_names, len(codes)) + codes = list(dict.fromkeys(codes)) + codes, unsupported_codes = _split_supported_daily_codes(codes) + if unsupported_codes: + _logger.info( + "[BaoStock] 已跳过 %d 只非沪深A股代码: %s", + len(unsupported_codes), + _preview_codes(unsupported_codes), + ) + + _logger.info("[BaoStock] 正在分析 %d 只股票的数据缺口...", len(codes)) gaps = _analyze_gaps(codes, start_date, end_date, trading_days) - complete_count = len(codes) - len(gaps) - len(not_listed) + complete_count = len(codes) - len(gaps) if complete_count > 0: - _logger.info("[%s] %d 只股票数据已完整,跳过", source_names, complete_count) + _logger.info("[BaoStock] %d 只股票数据已完整,跳过", complete_count) total = len(gaps) if total == 0: - _logger.info("[%s] 所有股票数据已完整,无需抓取", source_names) + _logger.info("[BaoStock] 所有股票数据已完整,无需抓取") return - mode = "全来源轮换+失败自动切换" if source == "all" else "单一来源" - _logger.info( - "正在抓取日线行情 [%s] [模式:%s] %s ~ %s,%d 只需更新,并发:%d", - source_names, mode, start_date, end_date, total, workers, - ) + if no_data_only: + _logger.info("[BaoStock] 仅更新停牌/无数据记录,不写入 stock_daily") - # all 模式下轮换首选来源;单来源模式下只使用指定来源。 - gap_items = list(gaps.items()) - tasks = [] - for i, (code, g) in enumerate(gap_items): - gap_start = g[0].replace("-", "") - gap_end = g[-1].replace("-", "") - source_order = _rotate_sources(sources, i) if source == "all" else sources - tasks.append((code, gap_start, gap_end, g, source_order)) + _logger.info("正在抓取日线行情 [BaoStock] %s ~ %s,%d 只需更新", start_date, end_date, total) success = 0 fail = 0 - nodata_count = 0 done = 0 t_start = time.time() - if workers == 1: - 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_keys) - t_fetch = time.time() - t0 - done += 1 - if ok: + tasks = list(gaps.items()) + for code, g in tasks: + gap_start = g[0].replace("-", "") + gap_end = g[-1].replace("-", "") + t0 = time.time() + try: + rows = _fetch_one_stock_with_retry( + code, + gap_start, + gap_end, + retry=retry, + delay=delay, + ) + expected_days = [ + d for d in trading_days + if g[0] <= d <= g[-1] and d not in no_data_dates.get(code, set()) + ] + if rows: + if no_data_only: + row_days = {str(row["date"]) for row in rows} + missing_days = [day for day in expected_days if day not in row_days] + if missing_days: + saved = _mark_no_data_days(code, missing_days) + if saved: + _logger.info( + "[BaoStock] %s 标记 %d 个停牌/无数据日期:%s ~ %s", + code, + saved, + missing_days[0], + missing_days[-1], + ) + else: + row_days = {str(row["date"]) for row in rows} + missing_days = [day for day in expected_days if day not in row_days] + if missing_days: + patched_rows = _fetch_missing_daily_rows( + code, + missing_days, + retry=retry, + delay=delay, + ) + if patched_rows: + rows.extend(patched_rows) + row_days.update(str(row["date"]) for row in patched_rows) + still_missing = [day for day in expected_days if day not in row_days] + if still_missing: + _logger.warning( + "[BaoStock] %s 仍缺失 %d 个交易日:%s ~ %s", + code, + len(still_missing), + still_missing[0], + still_missing[-1], + ) + else: + _logger.warning( + "[BaoStock] %s 范围结果缺少 %d 个交易日,未能按天补齐:%s ~ %s", + code, + len(missing_days), + missing_days[0], + missing_days[-1], + ) + batch_upsert(StockDaily, rows, ["code", "date"]) success += 1 - elif row_count == 0: - nodata_count += 1 else: fail += 1 + if no_data_only: + saved = _mark_no_data_days(code, expected_days) + if saved and expected_days: + _logger.info( + "[BaoStock] %s 标记 %d 个停牌/无数据日期:%s ~ %s", + code, + saved, + expected_days[0], + expected_days[-1], + ) + _logger.warning("[BaoStock] %s 没有返回数据,区间 %s ~ %s", code, g[0], g[-1]) + except Exception as exc: + fail += 1 + _logger.warning("[BaoStock] %s 获取失败: %s,区间 %s ~ %s", code, exc, g[0], g[-1]) + + time.sleep(delay + random.uniform(0, min(delay * 0.2, 0.5))) + + done += 1 + if done % 100 == 0 or done == total: elapsed = time.time() - t_start avg = elapsed / done eta = avg * (total - done) _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, + "[%d/%d] %s 缺口:%s~%s 耗时:%.1fs 成功:%d 失败:%d 剩余:%.0fs", + done, total, code, g[0], g[-1], time.time() - t0, success, fail, eta, ) - else: - with ThreadPoolExecutor(max_workers=workers) as pool: - future_map = {} - 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_keys = future_map[future] - code_r, label, row_count, ok = future.result() - done += 1 - if ok: - success += 1 - elif row_count == 0: - nodata_count += 1 - else: - fail += 1 - elapsed = time.time() - t_start - avg = elapsed / done - eta = avg * (total - done) - with _print_lock: - _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 - _logger.info( - "日线行情抓取完成 [%s] 并发:%d 成功:%d 失败:%d 无数据:%d 总耗时:%.1fs", - source_names, workers, success, fail, nodata_count, total_time, - ) + _logger.info("日线行情抓取完成 [BaoStock] 成功:%d 失败:%d 总耗时:%.1fs", success, fail, total_time) diff --git a/src/fetchers/dividend.py b/src/fetchers/dividend.py deleted file mode 100644 index e1485b1..0000000 --- a/src/fetchers/dividend.py +++ /dev/null @@ -1,117 +0,0 @@ -"""分红送转抓取模块 — 使用 BaoStock - -BaoStock query_dividend_data() 按年度查询分红记录。 -迭代最近10年获取完整分红历史。 -跳过策略:已有当前年份分红记录的股票不再重复抓取。 -""" - -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_missing_years(code: str, years: list[int]) -> list[int]: - """返回该股票缺失分红的年份""" - year_strs = [str(y) for y in years] - session = get_session() - try: - 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, years: list[int]) -> list[dict]: - """抓取单只股票指定年份的分红记录""" - bs_code = code_to_bs(code) - if not bs_code: - return [] - - rows = [] - 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(): - r = rs.get_row_data() - # fields: code, dividPreNoticeDate, dividAgmPumDate, dividPlanAnnounceDate, - # dividPlanDate, dividRegistDate, dividOperateDate, dividPayDate, - # dividStockMarketDate, dividCashPsBeforeTax, dividCashPsAfterTax, - # dividStocksPs, dividCashStock, dividReserveToStockPs - operate_date = r[6] if len(r) > 6 else "" - cash_before_tax = r[9] if len(r) > 9 else "" - stock_ps = r[11] if len(r) > 11 else "" - reserve_ps = r[13] if len(r) > 13 else "" - - if not operate_date and not cash_before_tax and not stock_ps: - continue - - rows.append({ - "code": code, - "name": "", - "report_date": str(year), - "dividend_date": operate_date or None, - "bonus_ratio": _to_float(stock_ps, scale=10), - "cash_div": _to_float(cash_before_tax, scale=10), - "convert_ratio": _to_float(reserve_ps, scale=10), - "ex_right_date": operate_date or None, - "dividend_yield": None, - }) - except Exception: - continue - return rows - - -def _to_float(val, scale=1) -> float | None: - if not val or val == "0.000000": - return None - try: - return round(float(val) * scale, 4) - except (ValueError, TypeError): - return None - - -def fetch_dividend(symbol: str | None = None): - """抓取分红送转数据 - - 用法:python -m src.main --dividend [--symbol 000001] - 已有分红记录的年份自动跳过,只抓缺失年份。 - """ - if symbol: - codes = [symbol] - else: - codes = get_stock_codes() - - current_year = datetime.now().year - all_years = list(range(current_year - 10, current_year + 1)) - - total = len(codes) - _logger.info("正在分析分红数据缺失情况,共 %d 只股票...", total) - - success = 0 - skip = 0 - for i, code in enumerate(codes): - 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: - _logger.info("[%d/%d] 进度... 成功:%d 跳过:%d", i+1, total, success, skip) - - _logger.info("分红送转抓取完成,成功:%d 跳过:%d/%d", success, skip, total) diff --git a/src/fetchers/financial.py b/src/fetchers/financial.py deleted file mode 100644 index 767f1db..0000000 --- a/src/fetchers/financial.py +++ /dev/null @@ -1,142 +0,0 @@ -"""季频财务指标抓取模块 — 使用 BaoStock - -BaoStock 按季度查询财务数据: - - query_profit_data() 盈利能力 - - query_balance_data() 偿债能力 - - query_cash_flow_data() 现金流 - -数据以 JSON 格式存入 data 列(与现有表结构兼容)。 -跳过策略:已有 (code, report_date) 记录的季度不再重复抓取。 -""" - -import json -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), ...]""" - now = datetime.now() - year, quarter = now.year, (now.month - 1) // 3 + 1 - result = [] - for _ in range(n): - result.append((year, quarter)) - quarter -= 1 - if quarter == 0: - quarter = 4 - year -= 1 - return result - - -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: - 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() - - -def _parse_resultset(code: str, rs, fields: list[str], year: int, quarter: int) -> list[dict]: - """将 BaoStock ResultData 转为 JSON 行""" - rows = [] - while rs.next(): - r = rs.get_row_data() - # fields: code, pubDate, statDate, ...指标字段 - stat_date = r[2] if len(r) > 2 else f"{year}-Q{quarter}" - data_dict = {} - for j, field in enumerate(fields): - if j < len(r): - val = r[j] - if isinstance(val, str) and val.strip() == "": - val = None - data_dict[field] = val - rows.append({ - "code": code, - "report_date": stat_date, - "data": json.dumps(data_dict, ensure_ascii=False), - }) - return rows - - -def fetch_financial(symbol: str | None = None): - """抓取财务数据 - - 用法:python -m src.main --financial [--symbol 000001] - 默认抓取所有股票最近 8 个季度。已有完整数据的股票自动跳过。 - """ - if symbol: - codes = [symbol] - else: - codes = get_stock_codes() - - quarters = _recent_quarters(8) - - total = len(codes) - _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 missing: - # 盈利能力 - with bs_query(bs.query_profit_data, code=bs_code, year=year, quarter=quarter) as rs: - fields = rs.fields if rs.fields else [] - rows = _parse_resultset(code, rs, fields, year, quarter) - if rows: - batch_upsert(FinancialIncome, rows, ["code", "report_date"]) - - # 偿债能力 - with bs_query(bs.query_balance_data, code=bs_code, year=year, quarter=quarter) as rs: - fields = rs.fields if rs.fields else [] - rows = _parse_resultset(code, rs, fields, year, quarter) - if rows: - batch_upsert(FinancialBalance, rows, ["code", "report_date"]) - - # 现金流 - with bs_query(bs.query_cash_flow_data, code=bs_code, year=year, quarter=quarter) as rs: - fields = rs.fields if rs.fields else [] - rows = _parse_resultset(code, rs, fields, year, quarter) - if rows: - batch_upsert(FinancialCashflow, rows, ["code", "report_date"]) - - success += 1 - except Exception as e: - _logger.error("[%d/%d] %s 失败: %s", i+1, total, code, e) - fail += 1 - - if (i + 1) % 50 == 0 or i == 0: - _logger.info("[%d/%d] 进度... 成功:%d 跳过:%d 失败:%d", i+1, total, success, skip, fail) - - _logger.info("财务数据抓取完成,成功:%d 跳过:%d 失败:%d", success, skip, fail) diff --git a/src/fetchers/index.py b/src/fetchers/index.py index 7a1a0ad..39c7344 100644 --- a/src/fetchers/index.py +++ b/src/fetchers/index.py @@ -1,41 +1,35 @@ -"""指数日线行情抓取 — 使用 BaoStock - -主要指数: - sh.000001 上证指数 sh.000300 沪深300 - sh.000905 中证500 sh.000852 中证1000 - sh.000688 科创50 sz.399001 深证成指 - sz.399006 创业板指 sz.399005 中小板指 - -用法: - python -m src.main --index - python -m src.main --index --start-date 20260101 --end-date 20260508 -""" +"""指数日线行情抓取 — 通达信本地缓存。""" import time -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 datetime import datetime + +import pandas as pd + from src.db import IndexDaily, TradingDay, batch_upsert, get_session +from src.fetchers.tdx_client import get_market_data as tdx_get_market_data from src.log import get_logger -from sqlalchemy import select, func +from sqlalchemy import select _logger = get_logger("index") -# 主要指数代码 → BaoStock 格式 INDICES = { - "000001": ("sh", "上证指数"), - "000300": ("sh", "沪深300"), - "000905": ("sh", "中证500"), - "000852": ("sh", "中证1000"), - "000688": ("sh", "科创50"), - "399001": ("sz", "深证成指"), - "399006": ("sz", "创业板指"), - "399005": ("sz", "中小板指"), + "000001": ("SH", "上证指数"), + "000300": ("SH", "沪深300"), + "000905": ("SH", "中证500"), + "000852": ("SH", "中证1000"), + "000688": ("SH", "科创50"), + "399001": ("SZ", "深证成指"), + "399006": ("SZ", "创业板指"), + "399005": ("SZ", "中小板指"), } +def _code_to_tdx(code: str) -> str: + market, _ = INDICES[code] + return f"{code}.{market}" + + def _clean(val): if val is None: return None @@ -44,11 +38,67 @@ def _clean(val): return val +def _tdx_market_data_to_rows(code: str, market_data: dict) -> list[dict]: + if not market_data: + return [] + + lower_map = {str(key).lower(): value for key, value in market_data.items()} + base_frame = None + for field in ("close", "open", "high", "low", "amount", "volume"): + frame = lower_map.get(field) + if hasattr(frame, "index") and len(frame.index) > 0: + base_frame = frame + break + + if base_frame is None: + return [] + + rows: list[dict] = [] + for ts in base_frame.index: + row = {"code": code, "date": str(ts)[:10]} + valid_value = False + for field, column in (("open", "open"), ("high", "high"), ("low", "low"), + ("close", "close"), ("volume", "volume"), ("amount", "amount")): + frame = lower_map.get(field) + if frame is None or ts not in frame.index: + continue + cell = frame.loc[ts] + if isinstance(cell, pd.Series): + if code in cell.index: + value = cell[code] + else: + value = cell.iloc[0] + else: + value = cell + value = _clean(value) + row[column] = value + if value is not None: + valid_value = True + if valid_value: + rows.append(row) + return rows + + +def _fetch_one_index(code: str, sd: str, ed: str) -> list[dict] | None: + """抓取单个指数的日线数据。""" + tdx_code = _code_to_tdx(code) + market_data = tdx_get_market_data( + [tdx_code], + period="1d", + start_time=f"{sd.replace('-', '')}000000", + end_time=f"{ed.replace('-', '')}235959", + count=-1, + dividend_type="none", + fill_data=True, + ) + rows = _tdx_market_data_to_rows(code, market_data) + return rows if rows else None + + def _analyze_index_gaps(code: str, sd: str, ed: str) -> list[tuple[str, str]]: - """分析单个指数在 [sd, ed] 范围内的数据缺口,返回缺失区间列表""" + """分析单个指数在 [sd, ed] 范围内的数据缺口。""" session = get_session() try: - # 获取范围内的交易日 trading_days = session.execute( select(TradingDay.date) .where(TradingDay.date >= sd) @@ -58,7 +108,6 @@ def _analyze_index_gaps(code: str, sd: str, ed: str) -> list[tuple[str, str]]: if not trading_days: return [(sd, ed)] - # 获取该指数已有的日期 existing = set(session.execute( select(IndexDaily.date) .where(IndexDaily.code == code) @@ -66,12 +115,10 @@ def _analyze_index_gaps(code: str, sd: str, ed: str) -> list[tuple[str, str]]: .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] @@ -88,41 +135,8 @@ def _analyze_index_gaps(code: str, sd: str, ed: str) -> list[tuple[str, str]]: session.close() -def _fetch_one_index(code: str, market: str, name: str, - sd: str, ed: str) -> list[dict] | None: - """抓取单个指数的日线数据""" - bs_code = f"{market}.{code}" - try: - with bs_query( - bs.query_history_k_data_plus, - bs_code, - "date,open,high,low,close,volume,amount,pctChg", - start_date=sd, end_date=ed, frequency="d", - ) as rs: - rows = [] - while rs.next(): - r = rs.get_row_data() - rows.append({ - "code": code, - "date": r[0], - "open": _clean(r[1]), - "high": _clean(r[2]), - "low": _clean(r[3]), - "close": _clean(r[4]), - "volume": _clean(r[5]), - "amount": _clean(r[6]), - "pct_change": _clean(r[7]), - }) - return rows if rows else None - except Exception: - return None - - def fetch_index(start_date: str | None = None, end_date: str | None = None): - """抓取主要指数日线行情""" - cfg = get_fetch_config() - delay = cfg.get("delay", 0.1) - + """抓取主要指数日线行情。""" if end_date is None: end_date = datetime.now().strftime("%Y%m%d") if start_date is None: @@ -131,7 +145,6 @@ 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]}" - # 分析每个指数的数据缺口 gaps_map: dict[str, list[tuple[str, str]]] = {} for code in INDICES: gaps = _analyze_index_gaps(code, sd, ed) @@ -142,25 +155,25 @@ def fetch_index(start_date: str | None = None, end_date: str | None = None): _logger.info("指数数据 %s ~ %s 已完整,跳过", sd, ed) return - skip_msg = f"(跳过 {len(INDICES) - len(gaps_map)} 个已完整)" - _logger.info("正在抓取指数日线 %s ~ %s,需补缺 %d 个%s", sd, ed, len(gaps_map), skip_msg) + _logger.info("正在抓取指数日线 %s ~ %s,需补缺 %d 个", sd, ed, len(gaps_map)) - bs_login() success = 0 + fail = 0 + t_start = time.time() for code, gaps in gaps_map.items(): - market, name = INDICES[code] + _, name = INDICES[code] total_rows = 0 t0 = time.time() for gap_sd, gap_ed in gaps: - rows = _fetch_one_index(code, market, name, gap_sd, gap_ed) + rows = _fetch_one_index(code, gap_sd, gap_ed) if rows: batch_upsert(IndexDaily, rows, ["code", "date"]) total_rows += len(rows) if total_rows: success += 1 - _logger.info("%s(%s): 补缺 %d 区间, %d 天, %.1fs", name, code, len(gaps), total_rows, time.time()-t0) + _logger.info("%s(%s): 补缺 %d 区间, %d 天, %.1fs", name, code, len(gaps), total_rows, time.time() - t0) else: + fail += 1 _logger.warning("%s(%s): 无数据", name, code) - time.sleep(delay) - _logger.info("指数数据抓取完成,成功:%d/%d", success, len(gaps_map)) + _logger.info("指数数据抓取完成,成功:%d 失败:%d", success, fail) diff --git a/src/fetchers/intraday.py b/src/fetchers/intraday.py index a1812d7..fbf77d1 100644 --- a/src/fetchers/intraday.py +++ b/src/fetchers/intraday.py @@ -1,9 +1,7 @@ -"""分钟K线抓取模块 — BaoStock / 通达信本地客户端 - -默认使用通达信本地客户端。 -如果临时需要,也可以显式切换回 BaoStock。 +"""分钟K线抓取模块 — 通达信本地客户端。 用法: + python -m src.main --intraday --freq 1 # 1分钟K线 python -m src.main --intraday --freq 5 # 5分钟K线 python -m src.main --intraday --start-date 20260508 --end-date 20260509 python -m src.main --intraday --symbol 000001 @@ -14,10 +12,8 @@ import random from datetime import datetime, timedelta import pandas as pd -import baostock as bs -from src.baostock_conn import bs_query, code_to_bs, bs_login from src.config import get_fetch_config -from src.db import StockMin5, TradingDay, batch_upsert, get_session, get_stock_codes +from src.db import StockMin1, StockMin5, TradingDay, batch_upsert, get_session, get_stock_codes from src.fetchers.tdx_client import get_market_data as tdx_get_market_data from src.fetchers.tdx_client import format_tdx_code_for_log from src.fetchers.tdx_client import normalize_tdx_code @@ -28,14 +24,20 @@ from sqlalchemy import select, func, distinct, text _logger = get_logger("intraday") -VALID_FREQ = ("5",) +VALID_FREQ = ("1", "5") TDX_AUTOCACHE_PERIODS = ("1m", "5m") FREQ_MODEL = { + "1": StockMin1, "5": StockMin5, } +FREQ_TDX_PERIOD = { + "1": "1m", + "5": "5m", +} + def _clean(val): if val is None: @@ -46,7 +48,7 @@ def _clean(val): def _parse_datetime(date_str: str, time_str: str) -> str | None: - """将 BaoStock 返回的 date + time 解析为 datetime 字符串 + """将分钟线返回的 date + time 解析为 datetime 字符串 time 格式: "20260508093500000" (17位) 或 "09:35:00" (8位) """ @@ -205,128 +207,21 @@ def _get_intraday_gaps(model, code: str, sd: str, ed: str) -> list[tuple[str, st def fetch_intraday(start_date: str | None = None, end_date: str | None = None, symbol: str | None = None, freq: str = "5", - source: str = "tdx", prewarm_cache: bool = True): + prewarm_cache: bool = True): """抓取分钟K线行情 Args: start_date: 开始日期 YYYYMMDD,默认30天前 end_date: 结束日期 YYYYMMDD,默认今天 symbol: 单只股票代码,默认全部 - freq: K线频率 5 - source: 数据源,`tdx` 或 `baostock` - prewarm_cache: 使用通达信时,是否先批量刷新本地分钟缓存 + freq: K线频率 1 或 5 + prewarm_cache: 是否先批量刷新通达信本地分钟缓存 """ if freq not in VALID_FREQ: _logger.error("不支持的频率 %s,可选: %s", freq, ", ".join(VALID_FREQ)) return - if source == "tdx": - _fetch_one_freq_tdx("5", start_date, end_date, symbol, prewarm_cache=prewarm_cache) - else: - _fetch_one_freq_baostock("5", start_date, end_date, symbol) - - -def _fetch_one_freq_baostock(freq: str, start_date: str | None, end_date: str | None, - symbol: str | None): - """抓取单个频率的分钟K线""" - cfg = get_fetch_config() - base_delay = float(cfg.get("delay", 0.1)) - # 分钟K线比日线更容易触发限流,默认放慢到 1 秒以上。 - delay = max(base_delay, 1.0) - model = FREQ_MODEL[freq] - - if end_date is None: - end_date = datetime.now().strftime("%Y%m%d") - if start_date is None: - start_date = (datetime.now() - timedelta(days=30)).strftime("%Y%m%d") - - 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]}" - - if symbol: - codes = [symbol] - else: - codes = get_stock_codes() - if not codes: - _logger.error("无股票列表,请先运行 --stock-info") - return - - # 分析每只股票的数据缺口 - 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 need_fetch: - _logger.info("%s分钟K线 %s ~ %s 数据已完整,跳过", freq, sd, ed) - return - - bs_login() - total = len(need_fetch) - success = 0 - fail = 0 - t_start = time.time() - - skip_msg = f"(跳过 {skip} 只已完整)" if skip else "" - _logger.info("正在抓取%s分钟K线 %s ~ %s,需补缺 %d 只%s", freq, sd, ed, total, skip_msg) - - for i, (code, gaps) in enumerate(need_fetch): - bs_code = code_to_bs(code) - if not bs_code: - continue - total_rows = 0 - try: - 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) - time.sleep(delay + random.uniform(0, delay * 0.2)) - if total_rows: - success += 1 - except Exception: - fail += 1 - - if (i + 1) % 100 == 0: - elapsed = time.time() - t_start - _logger.info( - "[%d/%d] 进度... 成功:%d 失败:%d 已用时:%.0fs", - i+1, total, success, fail, elapsed, - ) - - total_time = time.time() - t_start - _logger.info( - "%s分钟K线抓取完成,成功:%d 失败:%d 总耗时:%.1fs", - freq, success, fail, total_time, - ) + _fetch_one_freq_tdx(freq, start_date, end_date, symbol, prewarm_cache=prewarm_cache) def _fetch_one_freq_tdx(freq: str, start_date: str | None, end_date: str | None, @@ -336,7 +231,7 @@ def _fetch_one_freq_tdx(freq: str, start_date: str | None, end_date: str | None, base_delay = float(cfg.get("delay", 0.1)) delay = max(base_delay, 0.2) model = FREQ_MODEL[freq] - tdx_period = "5m" + tdx_period = FREQ_TDX_PERIOD[freq] if end_date is None: end_date = datetime.now().strftime("%Y%m%d") diff --git a/src/fetchers/market_daily.py b/src/fetchers/market_daily.py index dae6617..b7560b7 100644 --- a/src/fetchers/market_daily.py +++ b/src/fetchers/market_daily.py @@ -98,8 +98,7 @@ def _fetch_history(start_date: str | None, end_date: str | None): if rows: batch_upsert(MarketDaily, rows, ["date"]) _logger.info("已写入 %d 天涨跌停统计(%s ~ %s)", len(rows), rows[0]['date'], rows[-1]['date']) - else: - _logger.info("无新数据") + return rows def fetch_market_daily(start_date: str | None = None, end_date: str | None = None): @@ -109,4 +108,4 @@ def fetch_market_daily(start_date: str | None = None, end_date: str | None = Non start_date: 开始日期 YYYYMMDD,默认 19901219 end_date: 结束日期 YYYYMMDD,默认今天 """ - _fetch_history(start_date, end_date) + return _fetch_history(start_date, end_date) diff --git a/src/fetchers/sector.py b/src/fetchers/sector.py index 9a9e67d..c6c7eb4 100644 --- a/src/fetchers/sector.py +++ b/src/fetchers/sector.py @@ -1,20 +1,7 @@ -"""概念板块数据抓取 — 通达信本地 + 同花顺 - -数据源: - - 优先:通达信本地概念板块缓存 - - 备用:同花顺概念板块页面 - -用法: - python -m src.main --sector -""" +"""概念板块数据抓取 — 通达信本地缓存。""" import time -from contextlib import redirect_stderr, redirect_stdout -from io import StringIO -import pandas as pd -import requests -from bs4 import BeautifulSoup from src.db import StockConcept, batch_upsert, get_session from src.fetchers.tdx_blocks import load_infoharbor_blocks from src.log import get_logger @@ -27,34 +14,9 @@ def _is_a_stock_code(code: str) -> bool: return len(code) == 6 and code.isdigit() and not code.startswith("920") -def _fetch_concept_list_ths() -> list[dict]: - """使用同花顺获取全部概念板块列表。""" - try: - import akshare as ak - - # akshare 内部会打印进度条,这里静默掉,避免污染主流程日志。 - with redirect_stdout(StringIO()), redirect_stderr(StringIO()): - df = ak.stock_board_concept_name_ths() - if df is None or df.empty: - return [] - rows: list[dict] = [] - for _, row in df.iterrows(): - code = str(row.get("code", "")).strip() - name = str(row.get("name", "")).strip() - if code and name: - rows.append({"code": code, "name": name}) - return rows - except Exception as e: - _logger.error("同花顺 获取概念板块列表失败: %s", e) - return [] - - def _fetch_concept_list() -> list[dict]: """获取全部概念板块列表。""" - concepts = _fetch_concept_list_tdx() - if concepts: - return concepts - return _fetch_concept_list_ths() + return _fetch_concept_list_tdx() def _fetch_concept_list_tdx() -> list[dict]: @@ -63,78 +25,21 @@ def _fetch_concept_list_tdx() -> list[dict]: return [{"code": block.code, "name": block.name} for block in blocks if block.stocks] -def _extract_ths_stock_codes(html: str) -> set[str]: - """从同花顺概念详情页表格中提取 A 股代码。""" - codes: set[str] = set() - try: - tables = pd.read_html(StringIO(html)) - except ValueError: - return codes - - for df in tables: - if "代码" not in df.columns: - continue - for value in df["代码"].tolist(): - code = str(value).strip().zfill(6) - if _is_a_stock_code(code): - codes.add(code) - return codes - - -def _fetch_concept_stocks_ths(concept_code: str) -> list[str]: - """使用同花顺概念详情页获取单个概念板块成分股。""" - headers = { - "User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36", - "Referer": "https://q.10jqka.com.cn/gn/", - } - base_url = f"https://q.10jqka.com.cn/gn/detail/code/{concept_code}/" - - try: - r = requests.get(base_url, headers=headers, timeout=15) - r.encoding = "gbk" - soup = BeautifulSoup(r.text, features="lxml") - page_info = soup.find(name="span", attrs={"class": "page_info"}) - page_count = 1 - if page_info and "/" in page_info.text: - try: - page_count = int(page_info.text.strip().split("/")[-1]) - except ValueError: - page_count = 1 - - codes = _extract_ths_stock_codes(r.text) - for page in range(2, page_count + 1): - page_url = f"{base_url}page/{page}/" - try: - page_resp = requests.get(page_url, headers=headers, timeout=15) - page_resp.encoding = "gbk" - codes.update(_extract_ths_stock_codes(page_resp.text)) - except Exception as e: - _logger.warning("同花顺 获取概念 %s 第 %d 页失败: %s", concept_code, page, e) - time.sleep(0.03) - return sorted(codes) - except Exception as e: - _logger.warning("同花顺 获取概念 %s 成分股失败: %s", concept_code, e) - return [] - - def _fetch_concept_stocks(concept_code: str) -> list[str]: - """获取单个概念板块的成分股代码列表(仅保留 A 股)""" - stocks = _fetch_concept_stocks_tdx(concept_code) - if stocks: - return stocks - return _fetch_concept_stocks_ths(concept_code) + """获取单个概念板块的成分股代码列表(仅保留 A 股)。""" + return _fetch_concept_stocks_tdx(concept_code) def _fetch_concept_stocks_tdx(concept_code: str) -> list[str]: """使用通达信本地概念板块缓存获取成分股。""" for block in load_infoharbor_blocks(): if concept_code == block.code or concept_code == block.name: - return list(block.stocks) + return [code for code in block.stocks if _is_a_stock_code(code)] return [] def fetch_sector(): - """抓取全部概念板块及成分股,写入 stock_concept 表""" + """抓取全部概念板块及成分股,写入 stock_concept 表。""" _logger.info("正在获取概念板块列表...") concepts = _fetch_concept_list() if not concepts: @@ -170,7 +75,6 @@ def fetch_sector(): if (i + 1) % 50 == 0: elapsed = time.time() - t_start _logger.info("[%d/%d] 已处理 耗时:%.0fs", i + 1, len(concepts), elapsed) - time.sleep(0.05) if rows: batch_upsert(StockConcept, rows, ["code", "concept_code"]) diff --git a/src/fetchers/stock_list.py b/src/fetchers/stock_list.py index bbb19a8..5fa269e 100644 --- a/src/fetchers/stock_list.py +++ b/src/fetchers/stock_list.py @@ -1,46 +1,120 @@ -"""股票列表抓取模块 — 使用 BaoStock +"""股票列表抓取模块 — AKShare。""" -BaoStock query_stock_basic() 一次返回全部证券(含指数、基金等), -通过 type=1 过滤只保留股票,status=1 过滤仍在上市的。 -注意:BaoStock 不含北交所(920xxx)。 -""" +import random +import time +from datetime import datetime -import baostock as bs -from src.baostock_conn import bs_query -from src.db import StockInfo, batch_upsert, get_session +from src.config import get_fetch_config +from src.db import ( + StockInfo, + batch_upsert, + get_session, + invalidate_ipo_dates_cache, + invalidate_stock_codes_cache, +) from src.log import get_logger -from sqlalchemy import select, func +from sqlalchemy import func, select _logger = get_logger("stock_list") +try: + import akshare as ak +except Exception: # pragma: no cover - 运行环境缺少依赖时再报错 + ak = None + + +def _parse_ipo_date(raw: str | None): + if not raw: + return None + raw = str(raw).strip() + if not raw: + return None + if len(raw) == 8 and raw.isdigit(): + return datetime.strptime(raw, "%Y%m%d").date() + if len(raw) >= 10 and raw[4] == "-" and raw[7] == "-": + try: + return datetime.strptime(raw[:10], "%Y-%m-%d").date() + except ValueError: + return None + return None + + +def _fetch_akshare_stock_list(*, retry: int = 3, delay: float = 0.1) -> list[dict]: + """从 AKShare 获取沪深 A 股代码与名称。""" + if ak is None: + _logger.error("未安装 akshare,无法抓取股票列表") + return [] + + retry = max(int(retry), 1) + delay = max(float(delay), 0.0) + last_exc: Exception | None = None + + for attempt in range(1, retry + 1): + try: + df = ak.stock_info_a_code_name() + if df is None or df.empty: + return [] + if "code" not in df.columns or "name" not in df.columns: + _logger.warning("AKShare 股票列表字段不完整:%s", list(df.columns)) + return [] + rows: list[dict] = [] + for _, item in df.iterrows(): + code = str(item.get("code") or "").strip() + name = str(item.get("name") or "").strip() + if not code or not name: + continue + rows.append({"code": code, "name": name}) + return rows + except Exception as exc: + last_exc = exc + if attempt >= retry: + break + wait_seconds = delay * attempt if delay > 0 else 0.5 * attempt + wait_seconds += random.uniform(0, min(wait_seconds * 0.2, 0.5)) + _logger.warning("AKShare 获取股票列表第 %d 次失败,%.1fs 后重试: %s", attempt, wait_seconds, exc) + time.sleep(wait_seconds) + + if last_exc is not None: + _logger.warning("AKShare 获取股票列表失败: %s", last_exc) + return [] + def fetch_stock_list(): - """从 BaoStock 获取沪深A股列表,upsert 到 stock_info 表""" - # 检查已有数据量 + """从 AKShare 读取沪深 A 股列表,upsert 到 stock_info 表。""" session = get_session() try: existing = session.execute(select(func.count(StockInfo.code))).scalar() or 0 finally: session.close() + cfg = get_fetch_config() _logger.info("正在抓取股票列表(已有 %d 条)...", existing) - rows = [] - with bs_query(bs.query_stock_basic) as rs: - while rs.next(): - r = rs.get_row_data() - # fields: code, code_name, ipoDate, outDate, type, status - bs_code, name, ipo_date, out_date, typ, status = r[0], r[1], r[2], r[3], r[4], r[5] - if typ != "1" or status != "1": - continue - code = bs_code.split(".")[1] if "." in bs_code else bs_code - rows.append({ - "code": code, - "name": name, - "ipo_date": ipo_date if ipo_date else None, - }) + rows = _fetch_akshare_stock_list(retry=cfg.get("retry", 3), delay=cfg.get("delay", 0.1)) + if not rows: + _logger.warning("AKShare 股票列表为空") + return - if rows: - batch_upsert(StockInfo, rows, ["code"]) - _logger.info("股票列表已更新,共 %d 只", len(rows)) + upsert_rows: list[dict] = [] + for index, row in enumerate(rows, start=1): + code = str(row.get("code") or "").strip() + name = str(row.get("name") or "").strip() + if not code or not name: + continue + item = { + "code": code, + "name": name, + } + ipo_date = _parse_ipo_date(row.get("ipo_date")) + if ipo_date is not None: + item["ipo_date"] = ipo_date + upsert_rows.append(item) + if index % 500 == 0: + _logger.info("已解析 %d/%d 只股票", index, len(rows)) + + if upsert_rows: + batch_upsert(StockInfo, upsert_rows, ["code"]) + invalidate_stock_codes_cache() + invalidate_ipo_dates_cache() + _logger.info("股票列表已更新,共 %d 只", len(upsert_rows)) else: _logger.warning("无股票数据") diff --git a/src/fetchers/trading_day.py b/src/fetchers/trading_day.py index 56cf6e0..cfb33ef 100644 --- a/src/fetchers/trading_day.py +++ b/src/fetchers/trading_day.py @@ -1,99 +1,75 @@ -"""交易日历模块 - -数据源优先级: - 1. 本地 trading_day 表(最快) - 2. BaoStock 交易日历(需校验,近期可能含节假日) - 3. 从 stock_daily 表已有数据推断 -""" +"""交易日历模块 — AKShare。""" import time -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 +import random + +import pandas as pd + +from src.config import get_fetch_config +from src.db import TradingDay, batch_upsert, get_session from src.log import get_logger -from sqlalchemy import select, func, text +from sqlalchemy import select _logger = get_logger("trading_day") - -def _fetch_baostock(sd: str, ed: str) -> list[str] | None: - """从 BaoStock 获取交易日历""" - try: - with bs_query(bs.query_trade_dates, start_date=sd, end_date=ed) as rs: - days = [] - while (rs.error_code == "0") and rs.next(): - d = rs.get_row_data()[0] - if datetime.strptime(d, "%Y-%m-%d").weekday() < 5: - days.append(d) - return days if days else None - except Exception: - return None - - -def _validate_with_daily(dates: list[str]) -> list[str]: - """用 stock_daily 校验:只保留有实际行情数据的日期(排除节假日) - 仅保留当天(可能还没抓取),其余必须有行情数据才算交易日。 - """ - if not dates: - return dates - today = datetime.now().strftime("%Y-%m-%d") - session = get_session() - try: - result = session.execute( - select(func.distinct(StockDaily.date)) - .where(StockDaily.date >= dates[0]) - .where(StockDaily.date <= dates[-1]) - ) - real_dates = {str(row[0]) for row in result} - finally: - session.close() - - validated = [d for d in dates if d in real_dates or d == today] - return validated - - -def _fetch_and_save(start_date: str, end_date: str) -> list[str]: - """从 BaoStock 获取交易日并保存到本地表""" - sd = datetime.strptime(start_date, "%Y%m%d") - ed = datetime.strptime(end_date, "%Y%m%d") - today = datetime.now().strftime("%Y-%m-%d") - - all_days: list[str] = [] - chunk_start = sd - while chunk_start <= ed: - 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") - _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: - _logger.warning("%s ~ %s 获取失败", cs, ce) - chunk_start = chunk_end + timedelta(days=1) - - if all_days: - # 校验:排除节假日 - all_days = _validate_with_daily(all_days) - rows = [{"date": d} for d in all_days] - batch_upsert(TradingDay, rows, ["date"]) - _logger.info("交易日历已保存,%d 个交易日", len(all_days)) - return all_days - - _logger.warning("BaoStock 获取失败,将从已有行情数据推断") - return _infer_from_daily(start_date, end_date) +try: + import akshare as ak +except Exception: # pragma: no cover - 运行环境缺少依赖时再报错 + ak = None def _format_date(d: str) -> str: return f"{d[:4]}-{d[4:6]}-{d[6:8]}" -def get_trading_days(start_date: str, end_date: str) -> list[str]: - """获取指定范围内的交易日列表 +def _fetch_akshare_calendar(*, retry: int = 3, delay: float = 0.1) -> list[str]: + """从 AKShare 获取全量交易日历。""" + if ak is None: + _logger.error("未安装 akshare,无法抓取交易日历") + return [] - 优先查本地表,若本地数据未覆盖完整范围则补全。 - """ + retry = max(int(retry), 1) + delay = max(float(delay), 0.0) + last_exc: Exception | None = None + + for attempt in range(1, retry + 1): + try: + df = ak.tool_trade_date_hist_sina() + if df is None or df.empty: + return [] + trade_col = "trade_date" if "trade_date" in df.columns else df.columns[0] + series = pd.to_datetime(df[trade_col], errors="coerce").dropna() + return [dt.strftime("%Y-%m-%d") for dt in series] + except Exception as exc: + last_exc = exc + if attempt >= retry: + break + wait_seconds = delay * attempt if delay > 0 else 0.5 * attempt + wait_seconds += random.uniform(0, min(wait_seconds * 0.2, 0.5)) + _logger.warning("AKShare 获取交易日历第 %d 次失败,%.1fs 后重试: %s", attempt, wait_seconds, exc) + time.sleep(wait_seconds) + + if last_exc is not None: + _logger.warning("AKShare 获取交易日历失败: %s", last_exc) + return [] + + +def _fetch_and_save(start_date: str, end_date: str) -> list[str]: + """从 AKShare 获取交易日并保存到本地表。""" + cfg = get_fetch_config() + days = _fetch_akshare_calendar(retry=cfg.get("retry", 3), delay=cfg.get("delay", 0.1)) + if not days: + return [] + + rows = [{"date": d} for d in days] + batch_upsert(TradingDay, rows, ["date"]) + filtered = [d for d in days if start_date <= d.replace("-", "") <= end_date] + _logger.info("交易日历已保存,%d 个交易日(本次范围 %d 天)", len(rows), len(filtered)) + return filtered + + +def get_trading_days(start_date: str, end_date: str) -> list[str]: + """获取指定范围内的交易日列表。""" sd = _format_date(start_date) ed = _format_date(end_date) @@ -118,28 +94,8 @@ def get_trading_days(start_date: str, end_date: str) -> list[str]: return sorted(set(days + fetched)) -def _infer_from_daily(start_date: str, end_date: str) -> list[str]: - """从 stock_daily 表推断交易日""" - sd = _format_date(start_date) - ed = _format_date(end_date) - session = get_session() - try: - result = session.execute( - select(func.distinct(StockDaily.date)) - .where(StockDaily.date >= sd) - .where(StockDaily.date <= ed) - .order_by(StockDaily.date) - ) - return [str(row[0]) for row in result] - finally: - session.close() - - def fetch_trading_days(start_date: str | None = None, end_date: str | None = None): - """独立抓取交易日历并保存 - - 用法:python -m src.main --trading-day --start-date 19901219 --end-date 20261231 - """ + """独立抓取交易日历并保存。""" if end_date is None: end_date = time.strftime("%Y%m%d") if start_date is None: diff --git a/src/main.py b/src/main.py index 1d36d1a..72f5866 100644 --- a/src/main.py +++ b/src/main.py @@ -1,373 +1,99 @@ -"""A股数据抓取工具主入口 — BaoStock + 多源容灾 +"""A股数据抓取工具主入口 用法示例: - python -m src.main --stock-info # 先抓取股票列表 - python -m src.main --trading-day # 抓取交易日历 - python -m src.main --daily --start-date 20260501 --end-date 20260508 - python -m src.main --financial --symbol 000001 - python -m src.main --dividend - python -m src.main --intraday # 默认5分钟,近30天,通达信本地 - python -m src.main --intraday --start-date 20260508 --end-date 20260509 - python -m src.main --intraday --intraday-source tdx # 使用通达信本地客户端 - python -m src.main --tdx-cache # 全历史回填通达信5分钟数据 - python -m src.main --tdx-local-cache # 预热通达信本地分钟缓存并自动校验,不写数据库 - python -m src.main --tdx-verify-cache # 只验证通达信本地分钟缓存是否可读 + python -m src.main --stock-info # 先抓取股票列表(AKShare) + python -m src.main --trading-day # 抓取交易日历(AKShare) + python -m src.main --daily --start-date 20260501 --end-date 20260508 # BaoStock 日线接口 python -m src.main --sector # 概念板块及成分股 - python -m src.main --index # 指数日线(上证/沪深300/创业板等) python -m src.main --market-daily # 汇总每日涨跌停统计(依赖 stock_daily) + python -m src.main --sina-min1 # 新浪 1 分钟数据,默认全市场,默认 2010 至今 + python -m src.main --sina-min1 --symbol 600519 --start-date 20180101 --end-date 20260517 """ import argparse -from src.baostock_conn import bs_login, bs_logout from src.config import load_config -from src.db import get_ipo_dates, init_db +from src.db import init_db from src.log import get_logger _logger = get_logger("main") -TDX_LOCAL_CACHE_PERIODS = ("1m", "5m") - - -def _resolve_tdx_backfill_start_date(explicit_start_date: str | None, symbol: str | None) -> str: - """为通达信分钟回填选择更合理的起始日期。""" - if explicit_start_date: - return explicit_start_date - - ipo_dates = get_ipo_dates() - if symbol: - ipo_date = ipo_dates.get(symbol) - if ipo_date: - return ipo_date.replace("-", "") - return "19900101" - - if ipo_dates: - return min(ipo_dates.values()).replace("-", "") - return "19900101" - - -def _format_cache_preview(results: list[dict], *, limit: int = 10) -> tuple[int, int, str]: - """把缓存验证结果整理成更直观的摘要。""" - from src.fetchers.tdx_client import format_tdx_code_for_log - - hit_items = [item for item in results if item.get("has_data")] - miss_items = [item for item in results if not item.get("has_data")] - - if miss_items: - preview = ", ".join( - format_tdx_code_for_log(item["code"]) for item in miss_items[:limit] - ) - if len(miss_items) > limit: - preview += f", ...(+{len(miss_items) - limit})" - else: - preview = "" - - return len(hit_items), len(miss_items), preview - - -def _log_cache_summary( - *, - title: str, - results: list[dict], - single_symbol: bool = False, -) -> None: - """输出更直观的缓存结论。""" - if not results: - _logger.warning("%s:没有拿到任何验证结果", title) - return - - hit, miss, preview = _format_cache_preview(results) - total = len(results) - - if single_symbol: - item = results[0] - status = "可读" if item.get("has_data") else "不可读" - from src.fetchers.tdx_client import format_tdx_code_for_log - - _logger.info( - "%s:%s %s,行数 %d", - title, - status, - format_tdx_code_for_log(item["code"]), - item.get("row_count", 0), - ) - return - - if miss: - _logger.warning( - "%s:已可读 %d/%d,只缺失 %d 只%s%s", - title, - hit, - total, - miss, - ",缺失名单: " if preview else "", - preview, - ) - else: - _logger.info("%s:全部 %d 只可读", title, total) - - -def _run_tdx_local_cache(symbol: str | None) -> None: - """只预热通达信本地缓存,并在结束后自动校验可读性。""" - from src.db import get_stock_codes - from src.fetchers.tdx_client import ( - format_tdx_code_for_log, - refresh_minute_cache, - verify_minute_cache, - ) - - period_label = ",".join(TDX_LOCAL_CACHE_PERIODS) - if symbol: - _logger.info( - "开始预热通达信本地分钟缓存 %s,周期 %s,不写数据库", - format_tdx_code_for_log(symbol), - period_label, - ) - codes = [symbol] - else: - try: - codes = get_stock_codes() - except Exception as exc: - _logger.error("读取 stock_info 失败,无法预热全量通达信本地缓存:%s", exc) - return - - if not codes: - _logger.error("没有可用股票代码,请先运行 --stock-info") - return - - _logger.info( - "开始预热通达信全市场本地分钟缓存,股票数: %d,周期 %s,不写数据库", - len(codes), - period_label, - ) - refresh_minute_cache(codes, periods=TDX_LOCAL_CACHE_PERIODS) - _logger.info("通达信本地分钟缓存预热完成,开始自动校验") - results = verify_minute_cache(codes) - _log_cache_summary( - title="通达信本地缓存自动校验结果", - results=results, - single_symbol=bool(symbol), - ) - - -def _run_tdx_verify_cache( - start_date: str | None, - end_date: str | None, - symbol: str | None, -) -> None: - """只验证通达信本地分钟缓存是否可读。""" - from src.db import get_stock_codes - from src.fetchers.tdx_client import format_tdx_code_for_log, verify_minute_cache - - if symbol: - codes = [symbol] - _logger.info( - "开始验证通达信本地分钟缓存 %s,区间 %s ~ %s", - format_tdx_code_for_log(symbol), - start_date or "默认近30天", - end_date or "today", - ) - else: - try: - codes = get_stock_codes() - except Exception as exc: - _logger.error("读取 stock_info 失败,无法验证全量通达信本地缓存:%s", exc) - return - - if not codes: - _logger.error("没有可用股票代码,请先运行 --stock-info") - return - - _logger.info( - "开始验证通达信全市场本地分钟缓存,股票数: %d,区间 %s ~ %s", - len(codes), - start_date or "默认近30天", - end_date or "today", - ) - - results = verify_minute_cache(codes, start_date=start_date, end_date=end_date) - _log_cache_summary( - title="通达信本地缓存验证结果", - results=results, - single_symbol=bool(symbol), - ) - def main(): - 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="all", - choices=["baostock", "sina", "tencent", "eastmoney", "all"], - 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线行情") - parser.add_argument("--freq", type=str, default="5", - choices=["5"], - help="分钟K线频率(当前仅支持5分钟)") - parser.add_argument("--intraday-source", type=str, default="tdx", - choices=["baostock", "tdx"], - help="分钟K线数据源(默认tdx,baostock=备用来源)") - parser.add_argument("--tdx-cache", action="store_true", - help="全历史回填通达信5分钟数据(先预热本地缓存,再写入分钟表)") - parser.add_argument("--tdx-local-cache", action="store_true", - help="预热通达信本地分钟缓存并自动校验,不写数据库") - parser.add_argument("--tdx-verify-cache", action="store_true", - help="只验证通达信本地分钟缓存是否可读,不写数据库") + parser = argparse.ArgumentParser(description="A股数据抓取工具") + parser.add_argument("--stock-info", action="store_true", help="抓取股票列表(AKShare)") + parser.add_argument("--trading-day", action="store_true", help="抓取交易日历(AKShare)") + daily_group = parser.add_mutually_exclusive_group() + daily_group.add_argument("--daily", action="store_true", help="抓取日线行情(BaoStock 日线接口;只写入 stock_daily,不更新停牌/无数据记录)") + daily_group.add_argument( + "--daily-no-data-only", + action="store_true", + help="仅更新停牌/无数据记录,不写入 stock_daily(BaoStock 日线接口)", + ) parser.add_argument("--sector", action="store_true", help="抓取概念板块及成分股") - parser.add_argument("--index", action="store_true", help="抓取指数日线行情") + parser.add_argument("--sina-min1", action="store_true", + help="抓取新浪 1 分钟数据并写入 stock_min1(默认全市场,默认 2010 至今,可配合 --symbol / --start-date / --end-date)") parser.add_argument("--market-daily", action="store_true", help="汇总每日涨跌停统计(从 stock_daily 聚合)") 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="指定单只股票代码") + parser.add_argument("--symbol", type=str, help="指定单只或多个股票代码,逗号分隔;不传则抓取全市场,供 --sina-min1 使用") args = parser.parse_args() if not any([args.stock_info, args.trading_day, args.daily, - args.financial, args.dividend, args.intraday, - args.sector, args.index, args.market_daily, args.tdx_cache, - args.tdx_local_cache, args.tdx_verify_cache]): + args.daily_no_data_only, args.sector, args.sina_min1, args.market_daily]): parser.print_help() return load_config() - if args.tdx_verify_cache: - if any([args.stock_info, args.trading_day, args.daily, - args.financial, args.dividend, args.intraday, - args.sector, args.index, args.market_daily, args.tdx_cache, - args.tdx_local_cache]): - _logger.warning("`--tdx-verify-cache` 为纯验证命令,已忽略其它任务参数") - _run_tdx_verify_cache(args.start_date, args.end_date, args.symbol) - _logger.info("全部任务完成") - return - - if args.tdx_local_cache: - if any([args.stock_info, args.trading_day, args.daily, - args.financial, args.dividend, args.intraday, - args.sector, args.index, args.market_daily, args.tdx_cache]): - _logger.warning("`--tdx-local-cache` 为纯本地缓存命令,已忽略其它任务参数") - _run_tdx_local_cache(args.symbol) - _logger.info("全部任务完成") - return - init_db() - # market_daily 只读 stock_daily,无需 BaoStock 登录,提前处理 + # market_daily 只读 stock_daily,提前处理 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.index, args.tdx_cache]): + args.daily_no_data_only, args.sector, args.sina_min1]): _logger.info("全部任务完成") return - # sector 不需要 BaoStock,提前处理 + # sector 不需要额外网络数据源,提前处理 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, args.tdx_cache]): + args.daily_no_data_only, args.market_daily]): _logger.info("全部任务完成") return - if args.tdx_cache: - from src.fetchers.intraday import fetch_intraday - from src.db import get_stock_codes - from src.fetchers.tdx_client import refresh_all_minute_cache + if args.stock_info: + from src.fetchers.stock_list import fetch_stock_list + fetch_stock_list() - start_date = _resolve_tdx_backfill_start_date(args.start_date, args.symbol) - end_date = args.end_date - if args.symbol: - _logger.info( - "开始回填指定股票分钟数据 %s,区间 %s ~ %s,频率 5", - args.symbol, - start_date, - end_date or "today", - ) - fetch_intraday( - start_date=start_date, - end_date=end_date, - symbol=args.symbol, - freq="5", - source="tdx", - prewarm_cache=True, - ) - else: - codes = get_stock_codes() - if not codes: - _logger.error("没有可用股票代码,请先运行 --stock-info") - else: - _logger.info( - "开始回填通达信全市场分钟数据,股票数: %d,区间 %s ~ %s,频率 5", - len(codes), - start_date, - end_date or "today", - ) - refresh_all_minute_cache(stock_list=codes) - fetch_intraday( - start_date=start_date, - end_date=end_date, - freq="5", - source="tdx", - prewarm_cache=False, - ) + if args.trading_day: + from src.fetchers.trading_day import fetch_trading_days + fetch_trading_days(start_date=args.start_date, end_date=args.end_date) - if not any([args.stock_info, args.trading_day, args.daily, - args.financial, args.dividend, args.intraday, - args.sector, args.index, args.market_daily]): - _logger.info("全部任务完成") - return + if args.daily or args.daily_no_data_only: + from src.fetchers.daily import fetch_daily + fetch_daily( + start_date=args.start_date, + end_date=args.end_date, + no_data_only=args.daily_no_data_only, + ) - need_bs_login = any([args.stock_info, args.trading_day, args.daily, - args.financial, args.dividend, args.index]) or ( - args.intraday and args.intraday_source != "tdx" - ) + if args.sina_min1: + from src.fetchers.sina_minute import fetch_sina_min1 + fetch_sina_min1( + symbol=args.symbol, + start_date=args.start_date, + end_date=args.end_date, + ) - if need_bs_login: - bs_login() - try: - if args.stock_info: - from src.fetchers.stock_list import fetch_stock_list - fetch_stock_list() - - if args.trading_day: - from src.fetchers.trading_day import fetch_trading_days - fetch_trading_days(start_date=args.start_date, end_date=args.end_date) - - if args.daily: - from src.fetchers.daily import fetch_daily - fetch_daily(start_date=args.start_date, end_date=args.end_date, - source=args.source) - - if args.financial: - from src.fetchers.financial import fetch_financial - fetch_financial(symbol=args.symbol) - - if args.dividend: - from src.fetchers.dividend import fetch_dividend - fetch_dividend(symbol=args.symbol) - - if args.intraday: - from src.fetchers.intraday import fetch_intraday - fetch_intraday(start_date=args.start_date, end_date=args.end_date, - symbol=args.symbol, freq=args.freq, - source=args.intraday_source) - - if args.index: - from src.fetchers.index import fetch_index - fetch_index(start_date=args.start_date, end_date=args.end_date) - - _logger.info("全部任务完成") - finally: - if need_bs_login: - bs_logout() + _logger.info("全部任务完成") if __name__ == "__main__": diff --git a/tests/README.md b/tests/README.md index 2873dc1..dfad8a8 100644 --- a/tests/README.md +++ b/tests/README.md @@ -1,6 +1,6 @@ # tests 目录 -针对纯函数(不依赖网络、数据库、BaoStock 登录)的单元测试。 +针对纯函数、BaoStock 日线、AKShare 交易日/股票列表和 TDX 本地逻辑的单元测试。 ## 运行 @@ -9,15 +9,19 @@ pip install -e .[dev] # 或 pip install pytest pytest -v ``` -## 覆盖范围(48 个用例) +## 覆盖范围(66 个用例) -- `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_daily_source_codes.py` —— `daily._ak_hist_to_rows` 日线转换(1) +- `test_daily_sources_http.py` —— `daily._fetch_one_stock` / `daily.fetch_daily` BaoStock 日线抓取(3) +- `test_stock_list.py` —— `stock_list.fetch_stock_list` / AKShare 股票列表读取(2) +- `test_intraday_tdx.py` —— `intraday.fetch_intraday` / 通达信分钟转换(6) - `test_log.py` —— `log.get_logger` 命名空间、handler 幂等、env 控制 level(5) - `test_market_classification.py` —— `market_daily._is_20pct` 板块判定(4) -- `test_sector_http.py` —— 通达信/同花顺概念板块名称/成分股解析(4) +- `test_main_sina_minute.py` —— 新浪 1 分钟命令入口(3) +- `test_trading_day.py` —— `trading_day.get_trading_days` / AKShare 交易日历读取(2) +- `test_sector_http.py` —— 通达信概念板块名称/成分股解析(2) - `test_sector_merge.py` —— `sector.fetch_sector` 通达信增量跳过保护(2) +- `test_sina_minute.py` —— 新浪 1 分钟转换与落库(5) - `test_tdx_blocks.py` —— 通达信本地板块文件解析(2) +- `test_tdx_client.py` —— 通达信客户端适配与初始化包解压(11) diff --git a/tests/test_code_mapping.py b/tests/test_code_mapping.py deleted file mode 100644 index fe08bb9..0000000 --- a/tests/test_code_mapping.py +++ /dev/null @@ -1,30 +0,0 @@ -"""验证 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" diff --git a/tests/test_daily_source_codes.py b/tests/test_daily_source_codes.py index 603953f..4d41a26 100644 --- a/tests/test_daily_source_codes.py +++ b/tests/test_daily_source_codes.py @@ -1,52 +1,33 @@ -"""验证 daily.py 中各数据源的代码格式映射""" +"""验证 daily.py 中日线行转换逻辑。""" -from src.fetchers.daily import ( - _code_to_sina, - _code_to_tencent, - _code_to_eastmoney, -) +import pandas as pd + +from src.fetchers.daily import _ak_hist_to_rows -# ── 新浪/腾讯:sh{code}/sz{code} ── +def test_ak_hist_to_rows_converts_dataframe(): + """日线表应转换成可入库的日线行。""" + hist_df = pd.DataFrame( + { + "日期": ["2026-05-17", "2026-05-18"], + "开盘": [10.0, 10.2], + "收盘": [10.3, 10.4], + "最高": [10.5, 10.6], + "最低": [9.9, 10.0], + "成交量": [100, 120], + "成交额": [1030, 1248], + "振幅": [5.8, 5.9], + "涨跌幅": [3.0, 0.97], + "涨跌额": [0.3, 0.1], + "换手率": [1.0, 1.2], + } + ) -def test_sina_shanghai_main(): - assert _code_to_sina("600000") == "sh600000" + rows = _ak_hist_to_rows("600000", hist_df) - -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 + assert len(rows) == 2 + assert rows[0]["code"] == "600000" + assert rows[0]["date"] == "2026-05-17" + assert rows[0]["open"] == 10.0 + assert rows[0]["turnover"] == 1030 + assert rows[1]["date"] == "2026-05-18" diff --git a/tests/test_daily_sources_http.py b/tests/test_daily_sources_http.py index 213e18c..27ec10a 100644 --- a/tests/test_daily_sources_http.py +++ b/tests/test_daily_sources_http.py @@ -1,141 +1,304 @@ -"""验证 daily.py 中腾讯/东方财富 HTTP 解析逻辑(不打网络,用 requests.get mock)""" +"""验证 daily.py 中 BaoStock 日线抓取逻辑(不打网络)。""" -from unittest.mock import patch, MagicMock +from contextlib import contextmanager +from types import SimpleNamespace +from unittest.mock import patch + +import pandas as pd 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"], - ] - } +def _make_market_data() -> dict: + return pd.DataFrame( + { + "日期": ["2026-05-08", "2026-05-09"], + "开盘": [10.0, 10.2], + "收盘": [10.3, 10.4], + "最高": [10.5, 10.6], + "最低": [9.9, 10.0], + "成交量": [100, 120], + "成交额": [1030, 1248], + "振幅": [5.8, 5.9], + "涨跌幅": [3.0, 0.97], + "涨跌额": [0.3, 0.1], + "换手率": [1.0, 1.2], } - } + ) - with patch.object(daily.requests, "get", return_value=_make_response(payload)): - rows = daily._fetch_tencent("600000", "20260508", "20260509") - assert rows is not None +def _make_bs_result(rows: list[list[str]]): + cursor = {"idx": -1} + + def next_row(): + cursor["idx"] += 1 + return cursor["idx"] < len(rows) + + def get_row_data(): + return rows[cursor["idx"]] + + return SimpleNamespace(error_code="0", error_msg="success", next=next_row, get_row_data=get_row_data) + + +@contextmanager +def _fake_bs_query(result): + yield result + + +def test_fetch_one_stock_converts_baostock_rows(): + """单股 BaoStock 读取应转成可入库行。""" + result = _make_bs_result([ + ["2026-05-08", "10.0", "10.5", "9.9", "10.3", "100", "1030", "10.0", "3.0", "1.0"], + ["2026-05-09", "10.2", "10.6", "10.0", "10.4", "120", "1248", "10.3", "0.97", "1.2"], + ]) + + with patch.object(daily, "bs_query", lambda *args, **kwargs: _fake_bs_query(result)): + rows = daily._fetch_one_stock("600000", "20260508", "20260509") + 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 + assert rows[0]["turnover"] == 1030 -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_fetch_one_stock_returns_empty_when_no_data(): + """BaoStock 返回空表时应返回空列表。""" + result = _make_bs_result([]) + with patch.object(daily, "bs_query", lambda *args, **kwargs: _fake_bs_query(result)): + rows = daily._fetch_one_stock("600000", "20260508", "20260508") + + assert rows == [] -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_fetch_one_stock_skips_rows_without_volume(): + """volume 为空的日线行不应进入 stock_daily。""" + result = _make_bs_result([ + ["2026-05-08", "10.0", "10.5", "9.9", "10.3", "", "1030", "10.0", "3.0", "1.0"], + ["2026-05-09", "10.2", "10.6", "10.0", "10.4", "120", "1248", "10.3", "0.97", "1.2"], + ]) + + with patch.object(daily, "bs_query", lambda *args, **kwargs: _fake_bs_query(result)): + rows = daily._fetch_one_stock("600000", "20260508", "20260509") + + assert len(rows) == 1 + assert rows[0]["date"] == "2026-05-09" + assert rows[0]["volume"] == 120 -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_fetch_daily_uses_baostock_fetcher(): + """fetch_daily 应走 BaoStock 抓取并落库。""" + fake_session = SimpleNamespace(execute=lambda *args, **kwargs: [], close=lambda: None) + hist_rows = [ + ["2026-05-08", "10.0", "10.5", "9.9", "10.3", "100", "1030", "10.0", "3.0", "1.0"], + ] + result = _make_bs_result(hist_rows) + + with patch.object(daily, "bs_query", lambda *args, **kwargs: _fake_bs_query(result)), \ + patch.object(daily, "get_stock_codes", return_value=["600000"]), \ + patch.object(daily, "get_trading_days", return_value=["2026-05-08"]), \ + patch.object(daily, "get_session", return_value=fake_session), \ + patch.object(daily, "_analyze_gaps", return_value={"600000": ["2026-05-08", "2026-05-08"]}), \ + patch.object(daily, "batch_upsert") as mock_upsert: + daily.fetch_daily(start_date="20260508", end_date="20260508") + + mock_upsert.assert_called_once() -# ── 东方财富解析 ── +def test_fetch_daily_does_not_mark_no_data_days_when_empty_rows(): + """fetch_daily 遇到空结果时不应写入停牌/无数据日期。""" + fake_session = SimpleNamespace(execute=lambda *args, **kwargs: [], close=lambda: None) + empty_result = _make_bs_result([]) -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") + with patch.object(daily, "bs_query", lambda *args, **kwargs: _fake_bs_query(empty_result)), \ + patch.object(daily, "get_stock_codes", return_value=["600000"]), \ + patch.object(daily, "get_trading_days", return_value=["2026-05-08", "2026-05-09"]), \ + patch.object(daily, "get_session", return_value=fake_session), \ + patch.object(daily, "_analyze_gaps", return_value={"600000": ["2026-05-08", "2026-05-09"]}), \ + patch.object(daily, "batch_upsert") as mock_upsert: + daily.fetch_daily(start_date="20260508", end_date="20260509") - 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 + assert mock_upsert.call_count == 0 -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_fetch_daily_patches_missing_days_from_partial_range(): + """--daily 遇到范围结果缺日时,应按天补抓并写入完整数据。""" + fake_session = SimpleNamespace(execute=lambda *args, **kwargs: [], close=lambda: None) + range_rows = [ + {"code": "600000", "date": "2026-05-08", "open": 10.0, "close": 10.3, "high": 10.5, "low": 9.9, + "volume": 100, "turnover": 1030, "amplitude": 5.8, "pct_change": 3.0, "change": 0.3, "turnover_rate": 1.0}, + ] + single_day_rows = [ + {"code": "600000", "date": "2026-05-09", "open": 10.2, "close": 10.4, "high": 10.6, "low": 10.0, + "volume": 120, "turnover": 1248, "amplitude": 5.9, "pct_change": 0.97, "change": 0.1, "turnover_rate": 1.2}, + ] + + def fake_fetch(code, gap_start, gap_end, *, retry=3, delay=0.1): + if gap_start == "20260508" and gap_end == "20260509": + return list(range_rows) + if gap_start == "20260509" and gap_end == "20260509": + return list(single_day_rows) + raise AssertionError((code, gap_start, gap_end)) + + with patch.object(daily, "get_stock_codes", return_value=["600000"]), \ + patch.object(daily, "get_trading_days", return_value=["2026-05-08", "2026-05-09"]), \ + patch.object(daily, "get_session", return_value=fake_session), \ + patch.object(daily, "_analyze_gaps", return_value={"600000": ["2026-05-08", "2026-05-09"]}), \ + patch.object(daily, "_fetch_one_stock_with_retry", side_effect=fake_fetch), \ + patch.object(daily, "batch_upsert") as mock_upsert: + daily.fetch_daily(start_date="20260508", end_date="20260509") + + mock_upsert.assert_called_once() + model_cls, rows, index_columns = mock_upsert.call_args.args + assert model_cls.__tablename__ == "stock_daily" + assert index_columns == ["code", "date"] + assert [row["date"] for row in rows] == ["2026-05-08", "2026-05-09"] -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_fetch_daily_no_data_only_marks_partial_gap_days(): + """--daily-no-data-only 遇到部分行情时应补记缺失日期。""" + fake_session = SimpleNamespace(execute=lambda *args, **kwargs: [], close=lambda: None) + hist_rows = [ + ["2026-05-08", "10.0", "10.5", "9.9", "10.3", "100", "1030", "10.0", "3.0", "1.0"], + ] + result = _make_bs_result(hist_rows) + + with patch.object(daily, "bs_query", lambda *args, **kwargs: _fake_bs_query(result)), \ + patch.object(daily, "get_stock_codes", return_value=["600000"]), \ + patch.object(daily, "get_trading_days", return_value=["2026-05-08", "2026-05-09"]), \ + patch.object(daily, "get_session", return_value=fake_session), \ + patch.object(daily, "_analyze_gaps", return_value={"600000": ["2026-05-08", "2026-05-09"]}), \ + patch.object(daily, "batch_upsert") as mock_upsert: + daily.fetch_daily(start_date="20260508", end_date="20260509", no_data_only=True) + + assert mock_upsert.call_count == 1 + model_cls, rows, index_columns = mock_upsert.call_args.args + assert model_cls.__tablename__ == "stock_no_data" + assert index_columns == ["code", "date"] + assert [row["date"] for row in rows] == ["2026-05-09"] -def test_eastmoney_request_exception_raises_runtime(): - """网络/JSON 异常应包成 RuntimeError,让上层多源切换逻辑能捕获并切下一源""" - import pytest +def test_fetch_daily_no_data_only_marks_empty_result_as_no_data(): + """--daily-no-data-only 遇到空结果时仍应写入停牌/无数据日期。""" + fake_session = SimpleNamespace(execute=lambda *args, **kwargs: [], close=lambda: None) + empty_result = _make_bs_result([]) - def _boom(*_a, **_kw): - raise ConnectionError("network down") + with patch.object(daily, "bs_query", lambda *args, **kwargs: _fake_bs_query(empty_result)), \ + patch.object(daily, "get_stock_codes", return_value=["600000"]), \ + patch.object(daily, "get_trading_days", return_value=["2026-05-08", "2026-05-09"]), \ + patch.object(daily, "get_session", return_value=fake_session), \ + patch.object(daily, "_analyze_gaps", return_value={"600000": ["2026-05-08", "2026-05-09"]}), \ + patch.object(daily, "batch_upsert") as mock_upsert: + daily.fetch_daily(start_date="20260508", end_date="20260509", no_data_only=True) - with patch.object(daily.requests, "get", side_effect=_boom): - with pytest.raises(RuntimeError, match="东方财富请求失败"): - daily._fetch_eastmoney("600000", "20260508", "20260508") + assert mock_upsert.call_count == 1 + model_cls, rows, index_columns = mock_upsert.call_args.args + assert model_cls.__tablename__ == "stock_no_data" + assert index_columns == ["code", "date"] + assert [row["date"] for row in rows] == ["2026-05-08", "2026-05-09"] + + +def test_fetch_daily_skips_unsupported_non_a_share_codes(): + """fetch_daily 应在请求前过滤非沪深 A 股代码。""" + fake_session = SimpleNamespace(execute=lambda *args, **kwargs: [], close=lambda: None) + + with patch.object(daily, "get_stock_codes", return_value=["600000", "920200"]), \ + patch.object(daily, "get_trading_days", return_value=["2026-05-08"]), \ + patch.object(daily, "get_session", return_value=fake_session), \ + patch.object(daily, "_analyze_gaps", return_value={"600000": ["2026-05-08", "2026-05-08"]}) as mock_analyze, \ + patch.object(daily, "_fetch_one_stock_with_retry", return_value=[]) as mock_fetch, \ + patch.object(daily, "batch_upsert"): + daily.fetch_daily(start_date="20260508", end_date="20260508") + + mock_analyze.assert_called_once() + assert mock_analyze.call_args.args[0] == ["600000"] + mock_fetch.assert_called_once_with("600000", "20260508", "20260508", retry=3, delay=0.2) + + +def test_fetch_daily_retries_transient_baostock_failure(): + """fetch_daily 遇到临时断连时应重试。""" + fake_session = SimpleNamespace(execute=lambda *args, **kwargs: [], close=lambda: None) + hist_rows = [ + ["2026-05-08", "10.0", "10.5", "9.9", "10.3", "100", "1030", "10.0", "3.0", "1.0"], + ] + result = _make_bs_result(hist_rows) + calls = {"count": 0} + + def flaky_query(*args, **kwargs): + calls["count"] += 1 + if calls["count"] == 1: + raise ConnectionError("boom") + return _fake_bs_query(result) + + with patch.object(daily, "bs_query", flaky_query), \ + patch.object(daily, "get_fetch_config", return_value={"delay": 0, "retry": 2}), \ + patch.object(daily, "get_stock_codes", return_value=["600000"]), \ + patch.object(daily, "get_trading_days", return_value=["2026-05-08"]), \ + patch.object(daily, "get_session", return_value=fake_session), \ + patch.object(daily, "_analyze_gaps", return_value={"600000": ["2026-05-08", "2026-05-08"]}), \ + patch.object(daily, "batch_upsert") as mock_upsert, \ + patch.object(daily.time, "sleep", return_value=None): + daily.fetch_daily(start_date="20260508", end_date="20260508") + + assert calls["count"] == 2 + mock_upsert.assert_called_once() + + +def test_analyze_gaps_marks_incomplete_volume_as_gap(): + """已存在但关键字段为空时,应视为缺口并重抓。""" + class FakeSession: + def __init__(self): + self.calls = 0 + + def execute(self, *args, **kwargs): + self.calls += 1 + return [("600000", "2026-05", 2, 1)] + + def close(self): + return None + + fake_session = FakeSession() + + with patch.object(daily, "get_ipo_dates", return_value={"600000": "2026-01-01"}), \ + patch.object(daily, "_get_no_data_dates", return_value={}), \ + patch.object(daily, "get_session", return_value=fake_session): + gaps = daily._analyze_gaps( + ["600000"], + "20260508", + "20260509", + ["2026-05-08", "2026-05-09"], + ) + + assert gaps["600000"] == ["2026-05-08", "2026-05-09"] + assert fake_session.calls == 1 + + +def test_analyze_gaps_detects_missing_middle_trading_day(): + """已有头尾数据但中间缺交易日时,也应纳入缺口。""" + class FakeSession: + def __init__(self): + self.calls = 0 + + def execute(self, *args, **kwargs): + self.calls += 1 + if self.calls == 1: + return [("600000", "2026-05", 2, 0)] + return [("2026-05-08",), ("2026-05-11",)] + + def close(self): + return None + + fake_session = FakeSession() + + with patch.object(daily, "get_ipo_dates", return_value={"600000": "2026-01-01"}), \ + patch.object(daily, "_get_no_data_dates", return_value={}), \ + patch.object(daily, "get_session", return_value=fake_session): + gaps = daily._analyze_gaps( + ["600000"], + "20260508", + "20260511", + ["2026-05-08", "2026-05-09", "2026-05-11"], + ) + + assert gaps["600000"] == ["2026-05-09", "2026-05-09"] + assert fake_session.calls == 2 diff --git a/tests/test_financial_quarters.py b/tests/test_financial_quarters.py deleted file mode 100644 index 1f7c1a6..0000000 --- a/tests/test_financial_quarters.py +++ /dev/null @@ -1,50 +0,0 @@ -"""验证 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)]