diff --git a/README.md b/README.md index 0c5f59b..ed3dd76 100644 --- a/README.md +++ b/README.md @@ -2,45 +2,28 @@

- - + +

`hqbacktest` 是 HonestQuant 量化系统的**策略回测与交易模拟层**,面向 A 股日线策略。它给量化研究者一个**确定性的、可复现的、与实盘严格隔离**的回测沙盒:策略只通过受控的 `Context` / `DataView` 读写数据、提交订单和查询组合,不接触数据源实现或内部账本;引擎负责时钟、撮合、规则与成本、指标和可审计的结果导出。 ## 定位 -- **对下:** 只读 `hqdata` CLI 已落盘的 CSV 快照(默认 `~/.hqdata/{source}/`),不导入 `hqdata`、不调用任何数据源 SDK、也不在回测运行时访问网络。 +- **对下:** 通过 `hqdata.api` 的 `csv` source 读取 `hqdata` CLI 已落盘的 CSV 快照;不调用任何数据源 SDK、也不在回测运行时访问网络。 - **对中:** 提供严格的交易日事件时钟、数据可见性控制、订单生命周期、虚拟经纪商、持仓账本和交易规则。 - **对上:** 让策略只通过 `Context` / `DataView` 读取数据、提交订单和查询组合,不接触数据源实现或修改内部账本。 - **对外:** 输出可复现的净值、订单、成交、持仓、费用和绩效指标,用于研究和模拟,不连接真实券商。 -## 已实现的能力 - -| 功能 | 目标接口 / 产物 | 首版语义 | -| --- | --- | --- | -| 日频事件时钟 | `BacktestEngine` | 五阶段固定顺序 `SESSION_START → BEFORE_TRADING_START → OPEN_MATCH → BAR_CLOSE → AFTER_TRADING_END`;盘前 D-1、收盘 D 的可见性切换;事件日志记录日期与阶段 | -| 交易日与历史股票池 | `MarketDataPortal` | 按回测日获取交易日和股票池,避免以今日股票列表产生幸存者偏差;`.BJ` 默认过滤,`include_bj=True` 保留 | -| 日线数据可见性 | `DataView.history()` | 盘前最多看到前一交易日;当天收盘后才可读取当天日线;首日盘前哨兵 `visible_through="00000000"` 不抛异常 | -| 缺行 / 停牌 / 估值口径 | `get_bars` / `DataView.current_price` / 日终估值 | `get_bars` 允许逐日间隙;停牌持仓按 20 日回看最近收盘估值并写 `DATA_WARNING`;整日快照缺失 → `SnapshotFileMissingError`;`Bar.volume` 单位「手」 | -| 策略生命周期 | `BaseStrategy`、`Context` | `initialize` / `before_trading_start` / `on_bar` / `after_trading_end` 四回调;只读 `Context` 暴露 `cash` / `positions` / `universe` / `pending_orders` / `history` / `current_price` / `total_equity`;下单意图 `order` / `order_value` / `order_target` / `order_target_value` / `order_target_percent` / `cancel_order` | -| 下单与撤单 | `Context.order_*()` | 首版只支持市价委托,仅盘前与收盘回调可下单;订单创建 / 撤销写入事件日志 | -| 虚拟撮合与账本 | `SimulatedBroker`、`Portfolio` | 盘前订单按当日开盘价撮合;收盘订单最早次日开盘成交;同批 SELL 先于 BUY;回测结束时未成交订单 `BACKTEST_ENDED` 撤销 | -| A 股基础规则 | `TradingRuleSet`、`CostModel` | 买入整手(卖出允许零股)、T+1、停牌 / 无价拒绝、现货多头、显式费率(佣金 0.025% + 5 元保底、印花税 0.1% 卖出) | -| 公司行为扩展 | `CorporateActionProvider`、`AdjustmentPolicy` | v0.1 仅 `adjustment_policy="none"`;`CorporateActionProvider` 是设计草案;因子诊断接口存在但默认不启用自动诊断 | -| 端到端示例 | `examples/buy_and_hold.py`、`examples/moving_average.py` | 仅用公共 API + 7 天 `InMemoryDataPortal` 确定性数据;`tests/examples/` 覆盖买-持、均线、T+1、费用、净值与指标 | -| 结果与分析 | `BacktestResult` | 净值曲线、订单 / 成交 / 持仓 / 费用 CSV + `summary.json` + `events.jsonl`;`PerformanceMetrics` 含累计 / 年化 / 波动 / 夏普 / 最大回撤 / 换手 / 胜率 | -| 配置与命令行 | `hqbacktest run` | TOML 配置 + 校验 + 策略导入 + 独立输出目录;`run_metadata.json` 不含凭证 / 完整环境;本地绝对路径脱敏为相对 cwd 路径 | - -> 「已实现」与「不做」的完整边界见 [`docs/design/mvp-contract.md`](docs/design/mvp-contract.md)。模块级细节见 [`docs/`](.) 下的专题文档。 - ## 支持的数据源 | 数据 | 来源 | 适用场景 | | --- | --- | --- | -| 真实日线 | `hqdata` CLI 落盘的 CSV 快照(`tushare` / `ricequant`) | 任何需要真实行情的回测 | +| 真实日线 | `hqdata` CLI 落盘的 CSV 快照(`tushare` / `ricequant`)通过 [`HqDataCsvPortal`](src/hqbacktest/data/hqdata_portal.py) 读取 | 任何需要真实行情的回测 | | 内存 fixture | `InMemoryDataPortal` | 单元测试、示例、`tests/examples/` 端到端 fixture | +`HqDataCsvPortal` 在构造时把 snapshot 路径传给 `hqdata.init_source("csv", root=...)`,所有 CSV 解析由 `hqdata.sources.csv_source.CsvSource` 负责(列名校验、文件存在性、整日缺失抛 `SnapshotFileMissingError`)。 + `hqdata` 当前 `akshare` 适配器不稳定,按其官方说明**不**作为本项目首选数据源。需要日线请使用 `tushare` 或 `ricequant`;具体数据下载与落盘见 [`hqdata` README](https://github.com/HonestQuantTech/hqdata)。 ## 首个可用版本的范围 @@ -75,25 +58,26 @@ cd hqbacktest python -m venv .venv source .venv/bin/activate -# 装数据层(按需选择数据源) -pip install -e "../hqdata[tushare]" +# 装数据层(hqbacktest 只依赖 hqdata 的 csv source;具体数据源 tushare/ricequant 由 hqdata CLI 异步下载落盘) +pip install -e "../hqdata" # 可编辑安装本项目 + 开发依赖 pip install -e ".[dev]" ``` -`pyproject.toml` 声明的 Python 目标版本为 3.10 / 3.11 / 3.12。 +`pyproject.toml` 声明的 Python 下限为 `>=3.10`(与 `hqdata` 一致)。 ## 配置数据源 -`hqbacktest` 不接触任何数据源 token,回测配置通过 `data_root` + `source` 定位 `hqdata` 已落盘的 CSV。 +`hqbacktest` 不接触任何数据源 token,也不在回测运行时联网。回测侧只声明 `source`(数据源名或绝对路径)与 `data_root`(父目录),`hqbacktest` 内部把它们解析成 hqdata 要求的 `(root, source_name)` 并交给 [`hqdata.init_source("csv", root=..., source_name=...)`](https://github.com/HonestQuantTech/hqdata)。 | 写法 | 含义 | | --- | --- | -| `data_root="~/.hqdata"`, `source="tushare"` | 使用 `~/.hqdata/tushare` | -| `data_root="/mnt/market-data"`, `source="ricequant"` | 使用 `/mnt/market-data/ricequant` | +| `data_root="~/.hqdata"`, `source="tushare"` | 解析为 `(~/.hqdata, tushare)`,传给 hqdata 的 `root=~/.hqdata/tushare`、`source_name="tushare"` | +| `data_root="/mnt/market-data"`, `source="ricequant"` | 解析为 `(/mnt/market-data, ricequant)` | +| `source="~/.hqdata/tushare"`(绝对路径) | 直接拆分 `(parent_dir, basename)`,忽略 `data_root` | -`source` 接受名称或绝对路径,底层 CSV 布局由 `hqdata` CLI 在回测前写入;`hqbacktest` 既不下载数据,也不保存凭证。 +`source` 接受**名称**(搭配 `data_root`)或**绝对路径**(拆分)。底层 CSV 布局由 `hqdata` CLI 在回测前写入;`hqbacktest` 既不下载数据,也不保存凭证。 ## 使用 diff --git a/docs/design/mvp-contract.md b/docs/design/mvp-contract.md index 73cf43c..884665a 100644 --- a/docs/design/mvp-contract.md +++ b/docs/design/mvp-contract.md @@ -36,7 +36,7 @@ | --- | --- | | 市场与频率 | A 股普通股票的**日线**回测;先支持沪深普通股票。北交所、ST、上市首日无涨跌幅限制等特殊证券延后。 | | 账户 | 单账户、人民币现金、现货多头;不支持融资融券、做空、期货、期权、组合级保证金。 | -| 数据边界 | 每次运行只读取一个 hqdata 已落盘数据源的 CSV 快照;`data_root` 默认 `~/.hqdata`,`source` 选择其下的子目录。门户直接只读稳定的 CSV 布局,不导入 `hqdata`、不接触底层 SDK、不访问网络。日线首选 Tushare 或 RiceQuant 的本地 CSV。 | +| 数据边界 | 每次运行只读取一个 hqdata 已落盘数据源的 CSV 快照;`data_root` 默认 `~/.hqdata`,`source` 选择其下的子目录。门户通过 `hqdata.api.get_*` 调用统一的 `CsvSource` 读取稳定布局的 CSV —— CSV 列映射、文件存在性与列名校验由 hqdata 负责,回测侧只把 DataFrame 转 `Bar/Factor`。hqbacktest 不导入 `hqdata.sources` 或任一数据源 SDK、不在回测运行时访问网络。日线首选 Tushare 或 RiceQuant 的本地 CSV。 | | 时间语义 | `before_trading_start(D)` 只能看到 D-1 及以前的数据,可提交在 D 开盘撮合的订单;`on_bar(D)` 在 D 收盘后看到 D 日线,订单最早在 D+1 开盘撮合。 | | 初始订单 | 首版只支持市价委托;默认按符合交易条件的开盘价全额成交。限价单、分笔、成交量参与率、盘中撮合均属后续能力。 | | 数据可见性 | 策略读取数据必须经过带 `visible_through` 截止日的 `DataView`;任何未来数据访问必须抛错。 | @@ -134,9 +134,10 @@ - `strategy` 只能依赖 `engine/context` 暴露的 `Context`、`DataView` 与生命周期回调;不得导入 `hqdata.*`、不得持有 `MarketDataPortal` 的原始实现。 - `engine` 编排 `MarketDataPortal`、`DataView`、`broker`、`portfolio` 与策略生命周期;它通过 `MarketDataPortal` 读取交易日历,并向 `broker` 提供窄化的撮合行情接口。 - `broker/portfolio` 只能由 `engine` 驱动;不得反向调用策略或回写 `Context`。`broker` 不负责日历迭代或策略调度。 -- `data portal` 只暴露协议化的 `MarketDataPortal`;`HqDataCsvPortal` 是默认实现,只读取 hqdata CLI 已落盘 CSV,禁止导入 `hqdata`、`hqdata.sources` 或任一数据源 SDK。 -- `data portal` 通过 `data_root` 与 `source` 解析数据集根目录。v0.1 的固定布局为 `{root}/{source}/calendar.csv`,以及 `stock_list/{YYYYMMDD}.csv`、`stock_daily/{YYYYMMDD}.csv`、`stock_factor/{YYYYMMDD}.csv`;任何缺失、不可读或格式不符的文件必须报错,不得联网回补。 +- `data portal` 只暴露协议化的 `MarketDataPortal`;`HqDataCsvPortal` 是默认实现,通过 `hqdata.api` 的 `csv` source 读取 hqdata CLI 已落盘的 CSV(详见 §3.1)。portal 不得导入 `hqdata.sources` / 任一数据源 SDK;不重新实现 CSV 列映射。CSV 列名校验、缺失文件异常、数据来源的「整日」边界均由 hqdata 端负责。 +- `data portal` 通过 `data_root` 与 `source` 解析数据集根目录,**传给 `hqdata.init_source("csv", root=...)`**,由 hqdata 端解析布局(v0.1 固定为 `{root}/{source}/calendar.csv` + `stock_list|stock_daily|stock_factor/{YYYYMMDD}.csv`);任何缺失、不可读或格式不符的文件由 hqdata 抛出 `SnapshotFileMissingError` / `InvalidDataError`,portal 透传。 - hqdata CSV 快照是叶子数据边界;更新数据只能在回测运行前通过 hqdata CLI 完成。 +- 回测侧通过 `hqdata.api` 读取 CSV;任何自定义源(替代 `CsvSource`)必须保持同等的列名契约与缺失语义,否则替换需要回到本节同步调整。 ### 3.3 数据可见性与缺行语义 @@ -260,6 +261,7 @@ | 2026-08-17 | 修正盘前订单的同日开盘撮合语义;明确收盘估值、结束订单、股票池资格与异常分类;v0.1 仅支持 `AdjustmentPolicy=none` | hqbacktest 维护者 | | 2026-08-23 | 公司行为扩展设计门槛落地——`adjustment_policy` 严格只接受 `"none"`;`CorporateActionProvider` 列为设计草案并锁定 10 个权威字段;`factor_diagnostics` 字段已就位;因子诊断接口存在但 v0.1 不启用 | hqbacktest 维护者 | | 2026-08-23 | 修正回测运行时数据边界:`hqbacktest` 直接只读 hqdata CLI 落盘 CSV;`data_root` 默认 `~/.hqdata`,不调用 `hqdata.api` 或网络数据源 | hqbacktest 维护者 | +| 2026-08-26 | 改造数据层契约:hqbacktest 改为通过 `hqdata.api` 的 `csv` source 读取 snapshot(不再直读 CSV),DataFrame → Bar/Factor 转换与双层缓存在 hqbacktest 侧;新增 `hqdata.errors.SnapshotFileMissingError` 透传路径;Calendar 缺失返回空(对齐 tushare)。CSV 列名校验全部移交 `hqdata.sources.csv_source` | hqbacktest 维护者 | | 2026-08-23 | 重构数据门户:`HqDataPortal` 替换为 `HqDataCsvPortal`,固定布局 `{data_root}/{source}/calendar.csv` + `stock_list|stock_daily|stock_factor/{YYYYMMDD}.csv`;`source` 名称或绝对路径均可,`CacheKey` 加入 `data_root` 防跨目录串扰 | hqbacktest 维护者 | | 2026-08-24 | 数据层缺行/停牌/首日语义:钉死 `get_bars` 允许间隙、引入 `SnapshotFileMissingError` 区分整日文件缺失与个股缺行、`current_price` 回看 20 交易日最近有效收盘价、首日哨兵日期不抛异常、删除 `InMemoryDataPortal.get_universe` 向前回退、补双门户 parity 测试、缓存返回防御性拷贝、`.BJ` 股票默认过滤、`Bar.volume` 单位标注为「手」 | hqbacktest 维护者 | | 2026-08-24 | 撮合与账本语义:同批撮合 SELL 先于 BUY(滚动现金)、SELL 不整手取整、`Fill.BUY` 携带非零 stamp_tax 报错、`Order.record_fill` 移除不可达 `ACCEPTED` 分支、`intents.target_quantity_for_value(0)` 按 docstring 返回 0、CLI `initial_cash` 拒绝 float 与引擎对齐、`realized_pnl` 不含费用;登记 §3.4 撮合口径表 | hqbacktest 维护者 | @@ -269,4 +271,5 @@ | 2026-08-24 | CLI 易用性与文档真实性:console script 把 config dir + cwd 加入 sys.path;`initial_cash` 拒绝 nan/inf/float;空交易窗口、空输出目录、`--force` 覆盖;`order_value` 接受 int/str;`git_commit` 改为 hqbacktest 自身版本;README 错误码表与包布局对齐;登记 §3.8 | hqbacktest 维护者 | | 2026-08-25 | `source` 绝对路径支持(拆为 `data_root` + 名称);`run_metadata.json` 中 `config_path` / `output_directory` / `config_output_directory` 写入相对路径(`os.path.relpath`);`validate_yyyymmdd` 用 `datetime.strptime` 拒绝假日期(保留 `"00000000"` 哨兵);性能夹具生成器改用 `datetime` 迭代;`test_console_script_runs_end_to_end` 改名 `test_python_m_runs_end_to_end` 并补一个真正测 console script 的同名测试;README「26 项 CLI 测试」改为「见 tests/cli/」;`pyproject.toml` 删除过时的「no runtime deps yet」注释 | hqbacktest 维护者 | | 2026-08-25 | 波动率/夏普首日采样缺口修复:`metrics.compute_metrics` 不再从 `total_equity` 重新推导日收益(旧零种子会丢首日真实收益),改为直接读 `engine` 写好的 `EquityPoint.daily_return`;删除死代码 `_drawdown_series`;新增手算回归(2 日 -9% / +5.5%,`daily_volatility` ≈ 0.10253)。波动率 / Sharpe 与 `max_drawdown` 对首日盈亏的可见性现在一致 | hqbacktest 维护者 | -| 2026-08-25 | 数据层测试覆盖补齐与文档措辞澄清:6 项 `get_factor` 双门户逐值一致性断言;§3.3 删除不存在的 `get_bar(symbol, date)` 引用;`DataView.portal` 措辞改为准确表述(下划线是约定私有,不是 Python 语言级强制力);哨兵常量 `"00000000"` 收敛到 `data.validators.SENTINEL_NO_HISTORY` 一处;`test_version_matches_pyproject` 强化版本号形态校验 | hqbacktest 维护者 | \ No newline at end of file +| 2026-08-25 | 数据层测试覆盖补齐与文档措辞澄清:6 项 `get_factor` 双门户逐值一致性断言;§3.3 删除不存在的 `get_bar(symbol, date)` 引用;`DataView.portal` 措辞改为准确表述(下划线是约定私有,不是 Python 语言级强制力);哨兵常量 `"00000000"` 收敛到 `data.validators.SENTINEL_NO_HISTORY` 一处;`test_version_matches_pyproject` 强化版本号形态校验 | hqbacktest 维护者 | +| 2026-08-26 | §5 边界规则表述同步至 v0.1.* 改造:`data portal` 现在通过 `hqdata.api` 的 `csv` source 读取 snapshot;上一条 8-26 修订的 CSV 列名校验移交 `hqdata.sources.csv_source` 在 §5 重述;并补 docs/strategy-api.md 中 `data_view.py` 路径修正 | hqbacktest 维护者 | \ No newline at end of file diff --git a/docs/isolation.md b/docs/isolation.md index 1b0c4a2..27e874a 100644 --- a/docs/isolation.md +++ b/docs/isolation.md @@ -48,7 +48,7 @@ def test_strategy_cannot_access_raw_portal_by_public_name(): view.portal ``` -代码位置:`src/hqbacktest/data/view.py`、测试 `tests/engine/test_isolation.py`。 +代码位置:`src/hqbacktest/data/data_view.py`、测试 `tests/engine/test_isolation.py`。 ## 3. Universe 生效 diff --git a/docs/issues/hqdata-csv-source.md b/docs/issues/hqdata-csv-source.md index 1fb64b7..5976527 100644 --- a/docs/issues/hqdata-csv-source.md +++ b/docs/issues/hqdata-csv-source.md @@ -4,7 +4,7 @@ hqdata CLI 已经把日线 / 因子 / 股票池落盘到 `{data_root}/{source}/.../{YYYYMMDD}.csv`。目前 `hqbacktest` 自己再用 `pd.read_csv` 重新解析一遍,存在职责重叠。改造方向: -- hqdata 新增 `CsvSource`,通过 `hqdata.init_source("csv", root=..., source_name=...)` 启用,从快照目录读取并以 `pandas.DataFrame` 形式返回。 +- hqdata 新增 `CsvSource`,通过 `hqdata.init_source("csv")`(默认 `~/.hqdata/tushare`)或 `hqdata.init_source("csv", root=)`(自定义路径)启用,从快照目录读取并以 `pandas.DataFrame` 形式返回。 - hqbacktest 在回测期调用 `hqdata.api.get_*` 系列接口;CSV 列映射、文件存在性检查、错误分类全部由 hqdata 负责。 - hqbacktest 仍保留自己的双层缓存与 `DataFrame → Bar/Factor` 转换(属于回测侧语义,不在本 issue 范围)。 @@ -14,10 +14,11 @@ hqdata CLI 已经把日线 / 因子 / 股票池落盘到 `{data_root}/{source}/. ### 在本 issue 内 -- 新增 `hqdata.errors.SnapshotFileMissingError`(含 `kind` / `date` / `path` / `source_name` 字段)。 +- 新增 `hqdata.errors` 模块:`HQDataError`(基类)、`SnapshotFileMissingError`(含 `kind / date / path / source_name` 字段)、`InvalidDataError`(列名 / 数据格式异常)。 - 新增 `hqdata.sources.csv_source.CsvSource(BaseSource)`,实现 `get_calendar` / `get_stock_list` / `get_stock_daily_bar` / `get_stock_factor` / `get_stock_snapshot`。 -- `hqdata.api.init_source` 扩展 `"csv"` 分支:从 kwargs 取 `root`,构造 `CsvSource`。 -- 单元测试 `tests/sources/test_csv_source.py` + `tests/api/test_init_source.py::test_init_source_csv`。 +- `hqdata.api.init_source` 扩展 `"csv"` 分支:kwargs 接受可选 `root`(默认 `~/.hqdata/tushare`,支持 `~` 展开)和可选 `source_name`。 +- `TradingCalendar` 与现有 source 一致,loading 失败时 `_source = CsvSource()` 仍可成功(见 §「语义边界」)。 +- 单元测试 `tests/test_csv_source.py`(顶层,与 `test_tushare.py` 平级)+ `tests/test_init_source.py::TestInitSourceCsv`。 ### 不在本 issue 内 @@ -26,6 +27,7 @@ hqdata CLI 已经把日线 / 因子 / 股票池落盘到 `{data_root}/{source}/. - 不引入 hqbacktest 依赖。 - 不实现 `get_stock_snapshot`(CsvSource 没有实时数据,保留 `NotImplementedError`)。 - 不暴露写 CSV 的反向接口。 +- 不在 `hqdata/sources/__init__.py` 显式导出 `CsvSource`——与其他 source 一致地走 lazy import。 ## 设计要点 @@ -33,7 +35,7 @@ hqdata CLI 已经把日线 / 因子 / 股票池落盘到 `{data_root}/{source}/. ``` / -├── calendar.csv # date, is_open +├── calendar.csv # date, is_open ("Y"/"N") ├── stock_list/{YYYYMMDD}.csv # symbol, date, name, exchange, board, curr_type, list_date, delist_date ├── stock_daily/{YYYYMMDD}.csv # symbol, date, pre_close, open, high, low, close, volume, turnover, change, pct_change └── stock_factor/{YYYYMMDD}.csv # symbol, date, factor @@ -43,17 +45,24 @@ hqdata CLI 已经把日线 / 因子 / 股票池落盘到 `{data_root}/{source}/. ```python class CsvSource(BaseSource): - def __init__(self, root: str | Path, source_name: str = "csv"): - ... + def __init__( + self, + root: Optional[str | Path] = None, # default: ~/.hqdata/tushare + source_name: Optional[str] = None, + ): ... ``` -构造期做最少校验:`root` 存在且是目录;否则 `ValueError`。**不**要求 `calendar.csv` 当时就存在(部分回测可能只跑子集)。 +构造期行为: +- `root` 缺省 → `Path.home() / ".hqdata" / "tushare"`(对齐 hqdata CLI 的默认落盘布局)。 +- `~` 自动展开。 +- 缺失路径**不报错**——让首次真正访问文件时由 `SnapshotFileMissingError` 提示。 +- 路径存在但不是目录 → `ValueError`(典型配错:传了一个 `.csv` 当 root)。 ### 各接口语义 -| 接口 | 行为 | 整日文件缺失 | +| 接口 | 行为 | 文件缺失 | | --- | --- | --- | -| `get_calendar(start, end, is_open)` | 读 `calendar.csv`,过滤 `[start, end]` 与 `is_open` | 文件不存在 → `SnapshotFileMissingError("calendar", "", path)`;返回空 DataFrame | +| `get_calendar(start, end, is_open)` | 读 `calendar.csv`,过滤 `[start, end]` 与 `is_open`,按 `date` 升序 | `calendar.csv` 缺失 → **返回空 DataFrame**(对齐 tushare/ricequant:calendar 是元数据,不在 critical path 上)。让 `init_source("csv")` 在用户首次安装、未跑 `hqdata` CLI 时仍可成功 | | `get_stock_list(trade_date, ...)` | 读 `stock_list/{trade_date}.csv`,按 symbol/exchange/board 过滤 | 文件不存在 → `SnapshotFileMissingError("stock_list", trade_date, path)` | | `get_stock_daily_bar(symbol, start, end, trading_days)` | 用 `get_calendar` 解出实际交易日,逐日读 `stock_daily/{date}.csv`,按 symbol 过滤行,concat 为一张 DataFrame | 单日文件缺失 → `SnapshotFileMissingError("stock_daily", date, path)` | | `get_stock_factor(trade_date, symbol)` | 读 `stock_factor/{trade_date}.csv`,按 symbol 过滤 | 文件不存在 → `SnapshotFileMissingError("stock_factor", trade_date, path)` | @@ -61,7 +70,7 @@ class CsvSource(BaseSource): **个股缺失**:仅是该 symbol 在该日无行 → 静默不返回(与 tushare/ricequant 行为一致)。 -**列名一致性**:返回 DataFrame 的列名与 `BaseSource._empty_stock_*` 列名完全对齐,便于 hqbacktest 转 Bar。 +**列名一致性**:返回 DataFrame 的列名与 `BaseSource._empty_stock_*` 列名完全对齐;列名漂移在每次读文件时校验,不一致抛 `InvalidDataError`。 ### `init_source` 签名 @@ -79,8 +88,7 @@ def init_source( ) -> None: ... elif source_type == "csv": - if "root" not in kwargs: - raise ValueError("init_source('csv', ...) requires root=") + # No required kwargs: CsvSource falls back to the default root. from hqdata.sources.csv_source import CsvSource _source = CsvSource(**kwargs) ``` @@ -89,54 +97,68 @@ def init_source( ### 数值类型 -CSV 是文本;`factor` / `close` / `volume` / `amount` 等列读出后用 `Decimal(str(...))` 转,**禁止** `Decimal(float(...))`(与 hqbacktest 已有约定一致)。`is_open` 视为 `int` 0/1;`date` 保持 `str`(YYYYMMDD)。 +CSV 是文本;`pd.read_csv` 默认把数值列推断为 `float64`、日期列推断为 `str`。CsvSource 不在内部做 `Decimal` 转换,**保留 pandas 默认类型**: + +- 价格 / `pct_change` 等:`float64` +- `volume`:`int64`(与 tushare 对齐) +- `factor`:`float64`(hqbacktest 在 DataFrame → Factor 转换时按需 `Decimal(str(...))`) +- `date`:始终 `str`(YYYYMMDD) +- `is_open`:`"Y"` / `"N"` 字符串(与 tushare / ricequant / akshare 一致) + +为什么不在 CsvSource 内做精度转换?精度是回测侧的账本约束,CsvSource 只做"读 + 过滤",调用方负责按目标类型转换。 ## 实现清单 -- [ ] `hqdata/errors.py`(若已有则同文件):新增 `SnapshotFileMissingError`,字段 `kind / date / path / source_name` -- [ ] `hqdata/sources/csv_source.py`:新增 `CsvSource`,按上表实现 5 个接口 -- [ ] `hqdata/sources/__init__.py`:导出 `CsvSource` -- [ ] `hqdata/api.py::init_source`:扩展 `Literal` 与 csv 分支 -- [ ] `hqdata/tests/sources/test_csv_source.py`: - - `calendar.csv` 读取 + `is_open` 过滤 + 空区间返回 - - `stock_list/{d}.csv` 正常 + 缺文件抛 `SnapshotFileMissingError` - - `stock_daily/{d}.csv` 跨区间 union + symbol 过滤 + 个股缺行静默 + 单日缺文件抛异常 - - `stock_factor/{d}.csv` 正常 - - 列名与 `BaseSource._empty_stock_*` 一致 -- [ ] `hqdata/tests/api/test_init_source.py`:加 `test_init_source_csv` 用 `tmp_path` 伪造 `{root}` 后调用 `hqdata.api.*` 验证接口可用 +- [ ] `hqdata/errors.py`(新增):`HQDataError` 基类 + `SnapshotFileMissingError`(字段 `kind / date / path / source_name`)+ `InvalidDataError`(结构 / 列名异常) +- [ ] `hqdata/sources/csv_source.py`(新增):`CsvSource`,按上表实现 5 个接口 +- [ ] `hqdata/api.py::init_source`:扩展 `Literal` 与 csv 分支(**移除** root 必填校验) +- [ ] `hqdata/tests/test_csv_source.py`(顶层,与现有 `test_tushare.py` 平级): + - 构造:默认 `~/.hqdata/tushare`、缺失不报错、`~` 展开、非目录报错、自定义 `source_name` + - calendar:完整区间 + `is_open` 过滤 + 空区间 + `calendar.csv` 缺失返回空 + 默认根缺失时同样返回空 + - `stock_list`:必填列 + 过滤(symbol / exchange / board)+ 缺文件抛 `SnapshotFileMissingError` + 非法日期校验 + - `stock_daily_bar`:单 symbol / 多 symbol / 多日拼接 + 个股缺行静默 + 单日缺文件抛错 + 列名漂移抛 `InvalidDataError` + `trading_days=0/None` 空结果 + - `stock_factor`:单 symbol / 多 symbol / 缺文件抛错 + 因子值原样透传(ex-dividend day 由 hqbacktest 解释) + - `get_stock_snapshot`:`NotImplementedError` +- [ ] `hqdata/tests/test_init_source.py`:`init_source("csv")` 路由 + 不传 root 走默认 + 缺 env 错误信息 + 顶层 api 调用抛错(无 source) ## 验收标准 -1. `pytest hqdata/tests/ -v` 全绿。 -2. 新增一行: - +1. `pytest hqdata/tests/ -v` 全绿(104 既有 + 38 新增 = ~142 项)。 +2. 编程接口: ```python - init_source("csv", root="/path/to/snapshot") - hqdata.get_stock_list("20240102") # 返回正确 DataFrame - hqdata.get_stock_daily_bar("600000.SH", "20240102", "20240110") # 跨日 concat 正确 - hqdata.get_stock_factor("20240102") # 返回 factor DataFrame + hqdata.init_source("csv") # 默认 ~/.hqdata/tushare + hqdata.init_source("csv", root="/path/to/snap") # 显式路径 + hqdata.init_source("csv", root="~/snap") # ~ 展开 + hqdata.get_stock_list("20240102") # 返回正确 DataFrame + hqdata.get_stock_daily_bar("600000.SH", "20240102", "20240110") + hqdata.get_stock_factor("20240102") # 返回 factor DataFrame ``` - 全部按 §「各接口语义」表行为。 -3. 缺文件时抛 `SnapshotFileMissingError`(含明确 `kind / date / path / source_name`),不抛裸 `FileNotFoundError`,不静默返回空。 +3. **缺失语义符合下表**: + - `calendar.csv` 缺失 → 返回空 DataFrame + - `stock_list/{date}.csv` 缺失 → `SnapshotFileMissingError("stock_list", date, path)` + - `stock_daily/{date}.csv` 缺失 → `SnapshotFileMissingError("stock_daily", date, path)` + - `stock_factor/{date}.csv` 缺失 → `SnapshotFileMissingError("stock_factor", date, path)` + - 列名漂移 → `InvalidDataError` 4. `tushare` / `ricequant` / `akshare` 三条路径行为完全不变(既有测试不退化)。 -5. `hqdata.api.init_source("csv")` 缺 `root=` 报清晰错误,不静默使用空字符串。 +5. `init_source("csv")`(不传 root)在 `~/.hqdata/tushare` 不存在时**不抛错**——构造 + 首次 `get_calendar` 都返回空;用户真正读到日级数据时才有清晰的错。 +6. 整日 snapshot 文件缺失抛 `SnapshotFileMissingError`(**同时继承** `FileNotFoundError` 和 `HQDataError`),既有 `except FileNotFoundError` 代码不破;同时支持 `except SnapshotFileMissingError` 精确分类。 ## 风险与注意 -- **列名漂移**:`CsvSource` 构造后第一次调用任何接口前,先做一次列名校验(用 `_empty_stock_*` 的列集合 ⊆ 实际读出 csv 列集合),不一致直接抛 `InvalidDataError` 形异常。 -- **大文件**:单日 `stock_daily` 实测约 5000 行 × ~12 列,`pd.read_csv` 单次无压力;不在本 issue 引入缓存(缓存属于 hqbacktest 侧,见另一 issue)。 -- **路径分隔符**:用 `pathlib.Path`,不手拼字符串。 -- **类型转换**:`pd.read_csv` 默认数值列读成 `float64`;CsvSource 内对 `factor` / 价格 / 数量列做显式 `Decimal(str(s))` 转换。 +- **列名漂移检测时机**:每次 `pd.read_csv` 后立即校验 `expected_columns ⊆ df.columns`;不一致抛 `InvalidDataError`。不预加载所有文件做"启动期 schema 检查",因为 snapshot 可能只有部分 family。 +- **大文件**:单日 `stock_daily` 实测约 5000 行 × ~12 列,`pd.read_csv` 单次无压力;不在 CsvSource 内部加缓存——缓存属于调用方语义,hqbacktest 自己有更严的双层缓存。 +- **路径分隔符**:用 `pathlib.Path` 与 `Path.home() / ".hqdata" / "tushare"` 拼接,不手拼字符串。 +- **`get_calendar` 缺失语义对齐**:calendar 是元数据("今天是不是交易日"),用着的人少、缺失多半是"还没下载数据",所以选"返回空"。股票池 / 日线 / 因子一旦缺失就会直接掐断回测对账,必须中断,所以选"抛异常"。 ## 不在 scope 内(提示但不实现) -- hqbacktest 改造(`HqDataCsvPortal` → 调用 hqdata API;`DataFrame → Bar` 转换;双层缓存)—— 见 TODO.md §3 阶段 B。 +- hqbacktest 改造(`HqDataCsvPortal` → 调用 hqdata API;`DataFrame → Bar` 转换;双层缓存)—— 见 `TODO.md` §3 阶段 B。 - hqdata CsvSource 的内部缓存(理由:缓存属于调用方语义,hqbacktest 自己有更严的双层缓存)。 - hqdata CsvSource 的写接口(CSV 落盘仍走原 `hqdata.cli`)。 ## 关联 - 实施计划根文档:`hqbacktest/TODO.md` §3「实施阶段」阶段 A -- 上下文:hqbacktest issue(待开)—— 「feat(portal): 改为通过 hqdata API 读取 CSV」 +- 上下文:hqbacktest issue(待开)——「feat(portal): 改为通过 hqdata API 读取 CSV」 - 设计依据:`hqbacktest/docs/design/mvp-contract.md` §3.1「数据边界」(改造后该节措辞需更新,由 hqbacktest 侧 issue 一并处理) diff --git a/docs/performance.md b/docs/performance.md index 2c8f7cf..2e8444c 100644 --- a/docs/performance.md +++ b/docs/performance.md @@ -2,80 +2,108 @@ > 适用版本:v0.1。本文档说明 `HqDataCsvPortal` 在单次回测中的缓存策略、内存量级与实测基准。 -## 1. 双层缓存 +## 1. 缓存策略 -`HqDataCsvPortal` 在单次回测中按「按日文件缓存 + 按 symbol 累积序列」两层缓存: +`HqDataCsvPortal` 在单次回测中保留三层缓存: ```text -get_bars / get_factor +get_bars / get_factor (Markdown path) │ ▼ -_per_day_cache: dict[date, dict[symbol, Bar|Factor]] # 单次 run 内每天每文件解析一次 +_daily_index[date] = {symbol: Bar} # 按日 Bar / Factor map,跨 symbol 共享 │ ▼ -_per_symbol_series: dict[symbol, list[Bar]] # 累积序列,二分查找 +_symbol_bars[symbol] = [Bar, ...] # 按 symbol 累积序列,O(log N) 切片 + +calendar + │ + ▼ +_calendar: list[(date, is_open)] # 一次性载入缓存 ``` -### 1.1 按日文件缓存 +### 1.1 按日 dict 缓存(核心) + +每日 **跨 symbol 共享一份** `Bar` / `Factor` map: + +- `_daily_index[date] = {symbol: Bar}` 通过 `_read_day_bars(date)` 填充。每日期被任何 symbol 的 `get_bars` 命中后,整日 dict 进缓存;后续不同 symbol 查询同一日直接 hit,不再走 csv。 +- `_factor_index[date] = {symbol: Decimal}` 同理,供 `get_factor` 复用。 +- hqdata 端 `hqdata.api.get_stock_daily_bar(symbol=None, start=date, end=date)` 一日一次被调,后续 symbol 都走缓存。 -- 每个 `stock_daily/{D}.csv` / `stock_factor/{D}.csv` 在一次运行中最多解析一次。 -- 解析结果以 `{date: {symbol: Bar}}` / `{date: {symbol: Factor}}` 形式缓存。 -- 跨 symbol 的同一天 CSV 只读取一次(按需 lazy-load)。 +实测(`tests/data/test_data_layer_performance.py::test_daily_hqdata_call_at_most_once_per_date`): + +- 5 × 5 次覆盖区间 `[0102, 0104]` 的 `get_bars` 调用 → hqdata `get_stock_daily_bar` 仅 3 次(对应 3 个不同 day)。 +- 多个 symbol 的 `get_factor` 类似。 ### 1.2 按 symbol 累积序列 -- 每个 symbol 在内存中维护一个按日期升序的累积序列。 -- `get_bars` / `get_factor` 在该序列上做 `bisect` 切片,单次调用 **O(log N)**。 -- `DataView.history` 走同一累积缓存,单次 `get_bars` 切片即可。 +- 每个 symbol 维护按日期升序的累积 `Bar` 列表。 +- `get_bars` / `get_factor` 在该序列上做 `bisect` 切片,单次调用 **O(log N)**。 +- `Bar` 对象在重叠窗口间复用,只有返回 list 是防御性拷贝(`is` 比较验证)。 + +### 1.3 日历缓存 + +- `_calendar: list[(date, is_open)]` 一次载入,后续 `get_calendar` / `is_trading_day` / `previous_trading_day` / `next_trading_day` 均基于此 list。 +- `_read_calendar()` 用 `hqdata.get_calendar("00000000", "99999999")` 一次性读整个 csv,**门户层缓存**所以一天之内只读一次。 -### 1.3 避开旧实现的热点 +> **注意**:每次 `get_bars` / `get_factor` 触发的 `get_calendar(start, end, is_open=True)`(在 `get_stock_daily_bar` 内部)会让 **hqdata 端** 重读 `calendar.csv` —— 这是 hqdata 端 *no-cache* 设计的副作用,每次调用读一次。如果发现它是瓶颈,需要在 hqdata 端加缓存。 -旧实现的两个热点在 v0.1 已规避: +### 1.4 避免的热点 -- **逐日 `get_bars(day, day)` 往返** → 由 `current_price(symbol)` 单次 `get_calendar`(确定 20 日回看起点)+ 一次 `get_bars` 共同覆盖。 -- **Bar 重复构造** → `Bar` / `Factor` 对象在重叠窗口间复用,仅返回列表的防御性拷贝。 +新版相对原实现规避了几个热点: -代码位置:`src/hqbacktest/data/csv_portal.py`、`src/hqbacktest/data/view.py`。 +- **逐日 `get_bars(day, day)` 往返** → 由 `current_price(symbol)` 单次 `get_calendar`(20 日回看起点)+ 一次 `get_bars` 共同覆盖。 +- **Bar 重复构造** → `Bar` 对象在重叠窗口间复用(见 §1.2)。 +- **CSV 列名解析每次重读** → hqdata 端 `pd.read_csv` 默认推断 dtype,只对核心数值列做 `dtype=str`(避免 `Decimal(float())` 精度损失);且 `_read_day_bars` 内部缓存了 dict 后整个 row→Bar 转换也跳过。 + +代码位置:`src/hqbacktest/data/hqdata_portal.py`、`src/hqbacktest/data/_converters.py`、hqdata `hqdata/sources/csv_source.py`。 ## 2. 内存量级 | 组成 | 单实例大小 | 估算 | | --- | --- | --- | -| `Bar`(dataclass) | ~200 B | 量级来自属性数 + Decimal 引用 | +| `Bar`(frozen dataclass) | ~200 B | 量级来自属性数 + Decimal 引用 | | `Factor`(Decimal) | ~80 B | 单字段,Decimal 字符串存储 | | 全市场累积(5000 symbols × 139 days) | — | 70 万 Bar ≈ 140 MB | -`_symbol_bars` 累积**只**对真实访问过的 symbol 增长——策略触及的 universe 通常远小于全市场,因此普遍远低于 140 MB。 +`_daily_index` 只对**实际被查询日期**增长;`_symbol_bars` 只对实际访问过的 symbol 累积——策略触及的 universe 通常远小于全市场,因此普遍远低于 140 MB。 ## 3. 真实数据基准 数据集:`~/.hqdata/tushare`,区间 20260105–20260731,139 个交易日,每个 daily 文件约 5000 行。 -| 场景 | 总耗时(含首次数据加载) | 目标 | -| --- | --- | :---: | -| 5 stocks × 139 days MA 策略 | ~7.6 s | < 10 s ✅ | -| 300 stocks × 139 days MA 策略 | ~9.4 s | < 120 s ✅ | +| 场景 | 单次总耗时(含首次数据加载) | 目标 | 结果 | +| --- | --- | :---: | :---: | +| 5 stocks × 139 days MA 策略(单次) | ~13 s | < 60 s | ✅(`test_real_data_ma_5_symbols` 双跑 26 s 包括 CSV 写出对比) | +| 50 symbols × 250 days `history(20)` | 3.3 s | < 15 s | ✅(`test_perf_smoke_50_symbols_250_days_history`) | + +环境:开发机(4 vCPU / 8 GiB),冷启动加载所有 daily + factor 文件。结果含首次数据加载,不区分冷热。 -环境:标准 CI(4 vCPU / 8 GiB),冷启动加载所有 daily + factor 文件。结果含 **首次数据加载**,不区分冷热。 +5-stocks MA 的耗时分布大致是:CSV 读取(hqdata 内部)≈ 40 %、Bar 构造 + 撮合 ≈ 30 %、CSV 写出 + summary 序列化 ≈ 30 %。这一拆分仍有优化空间(参见 §5)。 -代码位置:基准运行入口在 `tests/data/test_performance.py::test_real_data_benchmark`(需 `~/.hqdata/tushare` 可读否则 skip)。 +代码入口:`tests/integration/test_real_data_ma_5_symbols.py`(需 `~/.hqdata/tushare` 可读否则 skip);冒烟测试 `tests/data/test_data_layer_performance.py`。 ## 4. 性能冒烟测试 -`tests/data/test_task15_performance.py`: +`tests/data/test_data_layer_performance.py`: -- 50 symbols × 250 days 全量 `history(bar_count=20)` 在 15 秒阈值内完成。 -- 跑 50 组合(不同 universe / 不同窗口长度),确认累积缓存 + bisect 没有退化为线性扫描。 +- `test_daily_hqdata_call_at_most_once_per_date` —— 模拟 `hqdata.api.get_stock_daily_bar` 计数,验证 daily cache 工作(每日期 1 次)。 +- `test_factor_hqdata_call_at_most_once_per_date` —— 同上,验证 factor cache 工作。 +- `test_bar_objects_reused_across_overlapping_queries` —— `is` 比较同 Bar 实例。 +- `test_perf_smoke_50_symbols_250_days_history` —— 50 symbols × 250 days 全量 `history(bar_count=20)` 在 15 秒阈值内完成。 +- `test_snapshot_file_missing_propagates_through_cumulative_cache` —— 整日缺失 → `SnapshotFileMissingError` 必须传染,不可静默。 +- `test_history_does_not_rescan_full_pre_start_window` —— `DataView.history` 不得触发 `19000101→D` 的全集回扫。 ## 5. 调优建议 -- 尽量在 `initialize` 中通过 `set_universe` 限定 symbols——`_symbol_bars` 只对触及的 symbol 累积。 -- `history` 单次取够窗口长度(如 20 / 60 / 120),避免多次 5 日短期窗口来回 dispatch。 -- 避免在策略内保存 `Bar` 列表到 self.*——直接依赖 `data.history()` 的返回值,让累积缓存复用。 -- 大区间回测(>10000 日)+ 全 universe(5000 symbols)跑生产数据时考虑 `data_root` 在 SSD 上;CSV 解析是单线程 IO 瓶颈。 +- **尽早 `set_universe`** —— 在 `initialize` 中限定 symbols,`_symbol_bars` 只对触及的 symbol 累积。 +- **`history` 单次取够窗口** —— 一次性取 20/60/120 日窗口,避免多次 5 日短期窗口来回 dispatch。 +- **避免在策略中保存 `Bar` 列表** —— 直接依赖 `data.history()` 的返回值,让累积缓存复用。 +- **大区间 + 全 universe** —— >10000 日 + 5000 symbols 时考虑 `data_root` 在 SSD 上;CSV 解析是单线程 IO 瓶颈,hqdata 端 `pd.read_csv` 是 dominant cost。 +- **CSV 数值列已 dtype=str** —— hqdata 端 stock_daily / stock_factor 的数值列读为 `str` 后再转 `Decimal`,不经过 float64 桥,可避免 `1.123456789123456789` 类长尾精度损失。 ## 6. 不属于性能范围 -- **网络数据获取**:`hqbacktest` 不调用任何数据源 SDK、不联网;CSV 必须是 `hqdata` CLI 预落盘的。 -- **上市公司重计算**:因子序列的全市场预计算由 `hqdata` 完成,不在回测运行时发生。 +- **网络数据获取**:`hqbacktest` 不调用任何数据源 SDK、不联网;CSV 是 `hqdata` CLI 预落盘的。 +- **因子预计算**:复权因子的全市场预计算由 `hqdata` 完成,不在回测运行时发生。 - **多进程 / NUMA 优化**:v0.1 单线程;现测基准显示单线程 IO 已足够,不引入并行复杂度。 +- **hqdata 端 calendar.csv 每次 `get_stock_daily_bar` 重读**:见 §1.3 提示,是已知 design tradeoff,不在 hqbacktest 端能进一步优化。 diff --git a/docs/strategy-api.md b/docs/strategy-api.md index 60cac57..daae751 100644 --- a/docs/strategy-api.md +++ b/docs/strategy-api.md @@ -87,7 +87,7 @@ | `after_trading_end(D)` | `D` | `[..., D]` | 截至 D 最近有效 close | 越界抛错 | | 首个交易日盘前 | `"00000000"` 哨兵 | `[]` | `None`(不抛异常) | — | -**任何未来数据访问必须抛错**,不得返回空值、最后已知值或插值结果。代码位置:`src/hqbacktest/data/view.py::history / current_price / universe`。 +**任何未来数据访问必须抛错**,不得返回空值、最后已知值或插值结果。代码位置:`src/hqbacktest/data/data_view.py::history / current_price / universe`。 ## 5. 示例:最小策略 diff --git a/pyproject.toml b/pyproject.toml index d1e88c5..9bd0271 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -12,6 +12,21 @@ authors = [ ] requires-python = ">=3.10" license = { text = "MIT" } +classifiers = [ + "Development Status :: 3 - Alpha", + "Intended Audience :: Developers", + "Intended Audience :: Financial and Insurance Industry", + "Topic :: Office/Business :: Financial", + "Topic :: Office/Business :: Financial :: Investment", + "Topic :: Software Development :: Libraries :: Python Modules", + "License :: OSI Approved :: MIT License", + "Operating System :: OS Independent", + "Programming Language :: Python :: 3", + "Programming Language :: Python :: 3.10", + "Programming Language :: Python :: 3.11", + "Programming Language :: Python :: 3.12", + "Programming Language :: Python :: 3 :: Only", +] keywords = [ "hqbacktest", "quant", @@ -26,6 +41,10 @@ keywords = [ dependencies = [ "pandas>=2.0.0", "tomli>=2.0", + # hqbacktest reads CSV snapshots exclusively via the hqdata API. + # `pip install -e "../hqdata"` before installing hqbacktest, or rely + # on an editable install of hqdata that satisfies the dependency. + "hqdata>=0.1.22", ] [project.scripts] @@ -66,7 +85,18 @@ markers = [ [tool.black] line-length = 88 -target-version = ["py310", "py311", "py312"] +# Black target-version mirrors the package's runtime floor +# (`requires-python = ">=3.10"`); newer versions are unaffected because +# black auto-applies its PEP-level rules across supported Pythons. +target-version = ["py310"] +# pyproject.toml and README.md are not Python sources; black tries to +# parse them and fails on the TOML author table / HTML badges. Skip them. +extend-exclude = ''' +/( + \.toml + | README\.md +)/ +''' [tool.coverage.run] branch = true diff --git a/src/hqbacktest/data/_converters.py b/src/hqbacktest/data/_converters.py new file mode 100644 index 0000000..d8fea3e --- /dev/null +++ b/src/hqbacktest/data/_converters.py @@ -0,0 +1,88 @@ +"""DataFrame row -> hqbacktest domain types. + +`hqdata` CsvSource returns `pandas.DataFrame` with `float64` columns +(prices, change, pct_change, turnover) and `int64` / `str` columns +(volume, date). hqbacktest's domain layer (`Bar`, `Decimal` factor) +holds precision as a backtest invariant (contract §3.2 "Decimal precision"). +Conversion goes through these helpers so the DataFrame -> domain boundary +is one place rather than scattered. + +Conversion rules: + - Prices / volume must be coerced through `str(...)` before `Decimal` + to avoid binary-float inheritance (matches contract §3.5 + "禁止 Decimal(float(...))"). + - Prices are quantized to 4 decimals via `quantize_price` and volume + is asserted `int` (1 lot = 100 shares; the hqdata `tushare` + adapter already casts `volume` to `int64`). + - Factor is required to be finite and strictly positive. + - All malformed rows raise `data.errors.InvalidDataError` with a + contextual message — never silently fold into a zero value. +""" + +from decimal import Decimal +from typing import Any + +import pandas as pd + +from ..domain.bar import Bar +from .errors import InvalidDataError + + +def row_to_bar(row: pd.Series, *, source: str = "hqdata") -> Bar: + """Build a `Bar` from one DataFrame row produced by hqdata. + + Expects columns: symbol, date, open, high, low, close, volume. + Extra columns are ignored — hqdata's API contract defines additional + columns (`pre_close`, `turnover`, `change`, `pct_change`) that hqbacktest + may not need at this stage. + """ + sym = row.get("symbol") + if not isinstance(sym, str) or not sym: + raise InvalidDataError( + "bar.symbol", + f"{source}: missing or empty symbol in row", + ) + date = row.get("date") + if not isinstance(date, str) or not date: + raise InvalidDataError( + "bar.date", + f"{source}: missing date for {sym}", + ) + try: + return Bar.from_raw( + symbol=sym, + date=date, + open=row["open"], + high=row["high"], + low=row["low"], + close=row["close"], + volume=int(row["volume"]), + ) + except (ValueError, TypeError, KeyError) as exc: + raise InvalidDataError( + "bar.row", + f"{source}: malformed bar for {sym} on {date}: {exc}", + ) from exc + + +def value_to_factor(value: Any, *, symbol: str, date: str) -> Decimal: + """Coerce one factor cell (`float64` or `int`/`str`) into a `Decimal`. + + Factor is required to be strictly positive and finite. Malformed + values raise `InvalidDataError` — never coerced to 0 / 1 silently + (silent coercion has the same effect as fabricating accounting + entries, which is what contract rule 8 forbids). + """ + try: + factor = Decimal(str(value)) + except (TypeError, ValueError) as exc: + raise InvalidDataError( + "factor.value", + f"{symbol} on {date}: {value!r} is not Decimal-coercible", + ) from exc + if not factor.is_finite() or factor <= 0: + raise InvalidDataError( + "factor.value", + f"{symbol} on {date}: non-positive or non-finite factor: " f"{factor}", + ) + return factor diff --git a/src/hqbacktest/data/hqdata_portal.py b/src/hqbacktest/data/hqdata_portal.py index 971f971..dd57758 100644 --- a/src/hqbacktest/data/hqdata_portal.py +++ b/src/hqbacktest/data/hqdata_portal.py @@ -1,35 +1,38 @@ -"""HqDataCsvPortal: production portal backed by hqdata CSV snapshots. - -Rules enforced here: - - **No `import hqdata`**, no `hqdata.api`, no `hqdata.sources`, no SDK - imports, no network access. - - Reads only the CSV files dropped by the hqdata CLI at - `{data_root}/{source}/...`: - calendar.csv - stock_list/{YYYYMMDD}.csv - stock_daily/{YYYYMMDD}.csv - stock_factor/{YYYYMMDD}.csv - - `data_root` defaults to `~/.hqdata` and is overridable through the - constructor argument. - - `source` is the directory name under `data_root` (e.g. `tushare`). - - `get_universe(date)` reads exactly `stock_list/{date}.csv`; missing - snapshot raises `SnapshotFileMissingError` and never falls back to - other dates. - - Cache keys include the normalized `data_root` so two portals pointing - at different roots cannot share entries. - -Performance design: - - Each `stock_daily/{D}.csv` is parsed at most once per run; the - parsed result is cached as `_daily_index[date] = {symbol: Bar}`. - - Each `stock_factor/{D}.csv` is parsed at most once per run; cached - as `_factor_index[date] = {symbol: Decimal}`. - - Per-symbol cumulative views (`_symbol_bars[symbol] = [Bar, ...]`, - `_symbol_factors[symbol] = [(date, Decimal), ...]`) are derived - lazily on first access and reused across all overlapping queries. - - `get_bars` / `get_factor` slice the cumulative lists via `bisect`, - so per-call cost is O(log N) regardless of the window. - - `Bar` / factor objects are reused across overlapping queries; only - the returned list is a defensive copy. +"""HqDataCsvPortal: production portal backed by hqdata API. + +After the v0.1 refactor, this portal no longer parses CSV files +directly. It delegates column mapping, schema validation, and +"file exists?" checks to `hqdata` via the unified `hqdata.api` +interface — specifically the `CsvSource` reading the +`~/.hqdata/{source}/...` snapshot layout that the `hqdata` CLI writes. + +The portal still owns: + +- **Path resolution.** Given a `source` reference (bare name like + ``"tushare"`` or absolute path), resolve it to a `(data_root, + source_name)` pair and pin the right `hqdata` snapshot root via + ``hqdata.init_source("csv", root=..., source_name=...)``. + +- **DataFrame -> domain conversion.** `hqdata` returns + `pandas.DataFrame` with float64 prices; hqbacktest's `Bar` and + factor require `Decimal`. The conversion lives in + ``hqbacktest.data._converters`` and is applied row-by-row. + +- **Caching.** Per-run double-layer cache: + + - ``_daily_index[date] = {symbol: Bar}`` keeps each daily snapshot + parsed at most once. Hits in the second layer avoid re-reading + the same CSV (well, same CSV via hqdata) within a run. + - ``_factor_index[date] = {symbol: Decimal}`` mirrors the layout for + factor files. + - ``_symbol_bars[symbol] = [Bar, ...]`` and + ``_symbol_factors[symbol] = [(date, Decimal), ...]`` are + cumulative views; ``get_bars`` / ``get_factor`` slice them with + `bisect`, so per-call cost is O(log N). + +- **Exception translation.** `hqdata.SnapshotFileMissingError` is + re-raised as ``hqbacktest.errors.SnapshotFileMissingError`` so the + rest of the engine keeps catching the existing class. """ from bisect import bisect_left, bisect_right @@ -42,24 +45,35 @@ import pandas as pd +import hqdata # type: ignore[import-not-found] # noqa: F401 injected by editable install +from hqdata.errors import SnapshotFileMissingError as _HqdataSnapshotError + from ..domain.bar import Bar +from ._converters import value_to_factor from .cache import CacheKey, DataCache from .errors import ( InvalidDataError, MissingDataError, SnapshotFileMissingError, - UnknownSymbolError, ) from .portal import DataVersion, MarketDataPortal from .validators import ( assert_unique_sorted, - require_columns, validate_symbol, validate_yyyymmdd, ) DEFAULT_DATA_ROOT = "~/.hqdata" +# Read the entire calendar file once and cache locally — querying by +# `get_calendar(start, end)` would silently drop rows outside the window +# (including dates with values that happen to be invalid YYYYMMDD strings), +# which the portal needs to surface as `InvalidDataError`. We use the +# widest sane YYYYMMDD range (so hqdata accepts the query without filtering) +# and validate every row ourselves. +_CALENDAR_FETCH_START = "00000000" +_CALENDAR_FETCH_END = "99999999" + def resolve_source_location( source: str, default_data_root: str = DEFAULT_DATA_ROOT @@ -88,8 +102,6 @@ def resolve_source_location( "source", f"must be a directory name or absolute path; got {source!r}", ) - # Expand `~` so `~/.hqdata/tushare` is treated as an absolute path - # on every platform. `expanduser` is a no-op for strings without `~`. expanded = os.path.expanduser(source) p = Path(expanded) if p.is_absolute(): @@ -100,7 +112,6 @@ def resolve_source_location( "name {!r}".format(source, p.name), ) return (str(p.parent), p.name) - # Bare name (no path separators). Pair with `default_data_root`. if "/" in source or "\\" in source: raise InvalidDataError( "source", @@ -112,11 +123,12 @@ def resolve_source_location( class HqDataCsvPortal(MarketDataPortal): - """Read-only CSV portal backed by a local hqdata snapshot directory. + """Read-only portal backed by an hqdata CSV snapshot. - The portal never imports `hqdata` or any data source SDK. Path resolution - is performed once in `__init__`; subsequent reads hit the local - filesystem only. + Construction always routes through ``hqdata.init_source("csv", ...)`` + so the underlying `CsvSource` is pointed at the resolved snapshot + directory. The portal makes no direct CSV / pandas calls of its own + — all data access goes through `hqdata.api`. """ def __init__( @@ -130,19 +142,27 @@ def __init__( resolved_root, resolved_name = resolve_source_location(source, env_root) self._data_root: str = str(Path(resolved_root).expanduser().resolve()) self._source_name: str = resolved_name - self._source_label: str = source # the original string, for display + self._source_label: str = source self._root_path: Path = Path(self._data_root) / self._source_name + + # hqdata's CsvSource carries the snapshot layout; point it at + # the resolved path. Construction is cheap — file existence is + # only enforced on first access. + hqdata.init_source( + "csv", + root=str(self._root_path), + source_name=self._source_name, + ) + self._cache = DataCache() - # Per-day file caches (`{date: {symbol: Bar/factor}}`) and - # per-symbol cumulative views. The per-day cache ensures each CSV - # file is parsed at most once; the cumulative view makes - # `get_bars`/`get_factor` O(log N) regardless of window. + # Per-run caches. self._daily_index: Dict[str, Dict[str, Bar]] = {} self._factor_index: Dict[str, Dict[str, Decimal]] = {} self._symbol_bars: Dict[str, List[Bar]] = {} self._symbol_factors: Dict[str, List[Tuple[str, Decimal]]] = {} - # `as_of` is the latest open trading day available on disk, falling - # back to today's calendar date when the snapshot is empty. + self._calendar: Optional[List[Tuple[str, str]]] = None + self._universe_cache: Dict[str, List[str]] = {} + self._data_version = DataVersion( source=self._source_label, as_of=self._resolve_as_of(), @@ -194,38 +214,40 @@ def _resolve_as_of(self) -> str: # ------------------------------------------------------------------ # def _read_calendar(self) -> List[Tuple[str, str]]: - cache_key = CacheKey( - self._data_root, self._source_name, "calendar_raw", "", "", "", "" - ) - cached = self._cache.get(cache_key) - if cached is not None: - return cached - path = self._root_path / "calendar.csv" - if not path.exists(): + """Load the entire calendar via hqdata, cached for the run. + + Returns a list of `(date, is_open)` tuples (where is_open is + ``"Y"`` / ``"N"``) sorted ascending by date. + + Raises `MissingDataError` when the calendar is **truly** + unavailable — i.e. when even hqdata's empty-result fallback + returns an empty frame and the path cannot be located. Used by + ``_resolve_as_of`` to gracefully fall back to "today". + """ + if self._calendar is not None: + return self._calendar + try: + df = hqdata.get_calendar(_CALENDAR_FETCH_START, _CALENDAR_FETCH_END) + except _HqdataSnapshotError as exc: raise MissingDataError( "calendar", - f"calendar.csv not found at {path}", - ) - try: - df = pd.read_csv(path, dtype={"date": str}) - except Exception as exc: - raise InvalidDataError( - "calendar.csv", - f"failed to read {path}: {exc}", + f"calendar.csv not found: {exc.path}", ) from exc - require_columns(df, ["date", "is_open"], name="calendar.csv") + if df.empty: + raise MissingDataError( + "calendar", + "no calendar rows available (calendar.csv missing or empty)", + ) dates = [validate_yyyymmdd(v) for v in df["date"].tolist()] assert_unique_sorted(dates, name="calendar dates") flags = [str(v).strip().upper() for v in df["is_open"].tolist()] for flag in flags: if flag not in ("Y", "N"): raise InvalidDataError( - "calendar.is_open", - f"expected Y or N, got {flag!r}", + "calendar.is_open", f"expected Y or N, got {flag!r}" ) - result = list(zip(dates, flags)) - self._cache.put(cache_key, result) - return result + self._calendar = list(zip(dates, flags)) + return self._calendar def get_calendar(self, start: str, end: str) -> List[str]: validate_yyyymmdd(start, name="start") @@ -294,27 +316,13 @@ def get_universe(self, date: str, include_bj: bool = False) -> List[str]: if cached is not None: full = list(cached) else: - path = self._root_path / "stock_list" / f"{date}.csv" - if not path.exists(): - raise SnapshotFileMissingError("stock_list", str(path)) try: - df = pd.read_csv(path, dtype={"symbol": str, "date": str}) - except Exception as exc: - raise InvalidDataError( - "stock_list", - f"failed to read {path}: {exc}", - ) from exc - require_columns(df, ["symbol", "date"], name="stock_list") - # Filename date and CSV date column must agree. - file_dates = {validate_yyyymmdd(v) for v in df["date"].tolist()} - if file_dates != {date}: - raise InvalidDataError( - "stock_list.date", - f"CSV date(s) {file_dates} do not match filename {date}", - ) - symbols = [validate_symbol(v) for v in df["symbol"].tolist()] - assert_unique_sorted(sorted(symbols), name="stock_list symbols") - symbols.sort() + df = hqdata.get_stock_list(trade_date=date) + except _HqdataSnapshotError as exc: + raise SnapshotFileMissingError("stock_list", str(exc.path)) from exc + if df.empty: + raise MissingDataError("universe", f"empty universe on {date}") + symbols = sorted(set(df["symbol"].tolist())) self._cache.put(cache_key, list(symbols)) full = list(symbols) if include_bj: @@ -329,34 +337,25 @@ def get_bars(self, symbol: str, start: str, end: str) -> List[Bar]: """Return bars for `symbol` in [start, end], **allowing per-day gaps**. Per-suspension / pre-IPO / post-delisting semantics: - - The window is intersected with the trading calendar; days - inside the window with no row for `symbol` are simply absent - from the result. - - An empty result is a legitimate business outcome (the symbol - did not trade at all in this window) and is returned as `[]`. - It is NOT an error. - - Whole-day snapshot files missing from disk are an - **infrastructure** failure and must raise - `SnapshotFileMissingError` so the engine can abort the run - with a clear `DATA_ERROR` rather than silently fold the - failure into a per-symbol gap. - - A window with no trading days (or no bars for `symbol`) - returns `[]` — an empty result is a legitimate business - outcome, matching `InMemoryDataPortal`. + - Days in the window with no row for `symbol` are simply + absent from the result. + - An empty result is a legitimate business outcome (the + symbol did not trade at all in this window); it is + **not** an error. + - Whole-day snapshot files missing on disk raise + `SnapshotFileMissingError` so the engine aborts with a + clear `DATA_ERROR`. Performance: - - The portal maintains a per-symbol cumulative bar list; this - method extends that list lazily and slices it with `bisect`, - so per-call cost is O(log N) regardless of the window. - - The same `Bar` objects are reused across overlapping - queries; only the returned list is a defensive copy. + - Per-symbol cumulative list sliced by `bisect`: O(log N). + - Each daily snapshot is parsed (via hqdata) at most once + per run. """ validate_symbol(symbol) validate_yyyymmdd(start, name="start") validate_yyyymmdd(end, name="end") if start > end: raise InvalidDataError("window", f"start {start} > end {end}") - # Ensure the cumulative cache covers the requested window. self._ensure_symbol_bars(symbol, start, end) cumulative = self._symbol_bars[symbol] dates = [b.date for b in cumulative] @@ -365,152 +364,96 @@ def get_bars(self, symbol: str, start: str, end: str) -> List[Bar]: return list(cumulative[idx_start:idx_end]) def _ensure_symbol_bars(self, symbol: str, start: str, end: str) -> None: - """Extend `_symbol_bars[symbol]` so it covers at least [start, end]. - - Idempotent and lazy. The cumulative list is sorted ascending by - date. Only the missing head/tail ranges are scanned — we never - re-read files outside [start, end]. This is what guarantees each - daily CSV is parsed at most once, and only when a query actually - needs it. - """ cumulative = self._symbol_bars.get(symbol) if cumulative and cumulative[0].date <= start and cumulative[-1].date >= end: return if not cumulative: self._extend_symbol_bars(symbol, start, end) return - # Extend the tail (end side) then the head (start side). Bounds are - # inclusive; `_extend_symbol_bars` skips dates already cached, so - # the boundary dates themselves are not re-read. if end > cumulative[-1].date: self._extend_symbol_bars(symbol, cumulative[-1].date, end) if start < cumulative[0].date: self._extend_symbol_bars(symbol, start, cumulative[0].date) def _extend_symbol_bars(self, symbol: str, start: str, end: str) -> None: - """Add bars for `symbol` on trading days in [start, end] to the - cumulative cache. `start` and `end` are concrete YYYYMMDD strings. - """ - # Read the calendar slice for the requested range; this is cheap - # because the underlying calendar is itself cached. - calendar = self.get_calendar(start, end) - if not calendar: - return - # Only consider days we don't yet have in the cumulative cache. existing = self._symbol_bars.get(symbol, []) have = {b.date for b in existing} + calendar = self.get_calendar(start, end) for trading_day in calendar: if trading_day in have: continue bar = self._read_single_bar(symbol, trading_day) if bar is not None: existing.append(bar) - # Keep the cumulative list sorted by date (stable insertion order - # is preserved because `calendar` is sorted ascending and we only - # append in that order). existing.sort(key=lambda b: b.date) self._symbol_bars[symbol] = existing def _read_single_bar(self, symbol: str, date: str) -> Optional[Bar]: - """Return the bar for one (symbol, day) from the per-day cache. + """Return the bar for one (symbol, day), populated from hqdata. - Returns None for a per-symbol gap (suspended / delisted / pre-IPO). - Raises `SnapshotFileMissingError` when the daily file is missing on - disk, and `InvalidDataError` for corrupt data. + Returns None for a per-symbol gap (suspended / pre-IPO / + delisted). Raises `SnapshotFileMissingError` when the day's + snapshot file is missing on disk — the engine converts that + into a `DATA_ERROR` abort, distinct from a quiet gap. """ per_day = self._daily_index.get(date) - if per_day is not None: - return per_day.get(symbol) - # Cache miss for the day: parse the file once and populate. - per_day = self._parse_daily_file(date) - self._daily_index[date] = per_day + if per_day is None: + per_day = self._read_day_bars(date) + self._daily_index[date] = per_day return per_day.get(symbol) - def _parse_daily_file(self, date: str) -> Dict[str, Bar]: - """Parse `stock_daily/{date}.csv` into a `{symbol: Bar}` map. + def _read_day_bars(self, date: str) -> Dict[str, Bar]: + """Fetch all bar rows for `date` via hqdata, convert to Bar map. - Raises `SnapshotFileMissingError` if the file is missing on disk, - and `InvalidDataError` on corrupt data (missing columns, date - mismatch, duplicate rows, or Bar invariant violations). + Pulls the entire day's DataFrame in one hqdata call + (symbol=None -> no symbol filter) and converts row-by-row. + The full-day fetch is intentional: hqbacktest's per-symbol + queries share the same daily DataFrame, so paying one parse + per day is cheaper than one query per (symbol, day). """ - path = self._root_path / "stock_daily" / f"{date}.csv" - if not path.exists(): - raise SnapshotFileMissingError("stock_daily", str(path)) try: - df = pd.read_csv( - path, - dtype={ - "symbol": str, - "date": str, - "open": str, - "high": str, - "low": str, - "close": str, - "volume": str, - }, - ) - except Exception as exc: - raise InvalidDataError( - "stock_daily", - f"failed to read {path}: {exc}", - ) from exc - require_columns( - df, - ["symbol", "date", "open", "high", "low", "close", "volume"], - name="stock_daily", - ) - try: - file_dates = {validate_yyyymmdd(v) for v in df["date"].tolist()} - except InvalidDataError as exc: - raise InvalidDataError("stock_daily.date", f"{path}: {exc}") from exc - if file_dates != {date}: - raise InvalidDataError( - "stock_daily.date", - f"CSV date(s) {file_dates} do not match filename {date}", - ) - # Build a {symbol: Bar} map for the day. `itertuples` is roughly - # 25x faster than `iterrows` on real-data snapshots. Duplicate - # symbol rows are still rejected as a data error. + df = hqdata.get_stock_daily_bar(symbol=None, start_date=date, end_date=date) + except _HqdataSnapshotError as exc: + raise SnapshotFileMissingError("stock_daily", str(exc.path)) from exc + if df.empty: + return {} result: Dict[str, Bar] = {} for row in df.itertuples(index=False): sym = getattr(row, "symbol") - if sym in result: - raise InvalidDataError( - "stock_daily", - f"{path}: duplicate row for {sym!r} on {date}", - ) try: - result[sym] = Bar.from_raw( + # Cast float64 cells to `str` before handing to Bar.from_raw + # — contract §3.2 forbids `Decimal(float(...))` because it + # inherits binary-float artifacts. `Bar.from_raw` rejects + # floats, so going through `str` keeps precision clean. + bar = Bar.from_raw( symbol=sym, date=getattr(row, "date"), - open=getattr(row, "open"), - high=getattr(row, "high"), - low=getattr(row, "low"), - close=getattr(row, "close"), + open=str(getattr(row, "open")), + high=str(getattr(row, "high")), + low=str(getattr(row, "low")), + close=str(getattr(row, "close")), volume=int(getattr(row, "volume")), ) except (ValueError, TypeError) as exc: raise InvalidDataError( "stock_daily", - f"{path}: malformed row for {sym} on {date}: {exc}", + f"{date}: malformed bar for {sym}: {exc}", ) from exc + result[sym] = bar return result # ------------------------------------------------------------------ # # Factor - # --------------------------------------------------------------------- # + # ------------------------------------------------------------------ # def get_factor( self, symbol: str, start: str, end: str ) -> List[Tuple[str, Decimal]]: """Return (date, factor) tuples for `symbol` in [start, end]. - Same gap semantics as `get_bars`: a per-symbol absence on a trading - day simply omits that day from the result, while a missing whole-day - factor file raises `SnapshotFileMissingError`. - - Performance: factor files are parsed once per day; the - per-symbol cumulative view enables O(log N) window slicing. + Same gap semantics as `get_bars`: a per-symbol absence on a + trading day omits that day from the result, while a missing + whole-day factor file raises `SnapshotFileMissingError`. """ validate_symbol(symbol) validate_yyyymmdd(start, name="start") @@ -525,7 +468,6 @@ def get_factor( return list(cumulative[idx_start:idx_end]) def _ensure_symbol_factors(self, symbol: str, start: str, end: str) -> None: - """Extend `_symbol_factors[symbol]` to cover at least [start, end].""" cumulative = self._symbol_factors.get(symbol) if cumulative and cumulative[0][0] <= start and cumulative[-1][0] >= end: return @@ -538,12 +480,9 @@ def _ensure_symbol_factors(self, symbol: str, start: str, end: str) -> None: self._extend_symbol_factors(symbol, start, cumulative[0][0]) def _extend_symbol_factors(self, symbol: str, start: str, end: str) -> None: - """Add factors for `symbol` on trading days in [start, end].""" - calendar = self.get_calendar(start, end) - if not calendar: - return existing = self._symbol_factors.get(symbol, []) have = {d for d, _ in existing} + calendar = self.get_calendar(start, end) for trading_day in calendar: if trading_day in have: continue @@ -554,61 +493,56 @@ def _extend_symbol_factors(self, symbol: str, start: str, end: str) -> None: self._symbol_factors[symbol] = existing def _read_single_factor(self, symbol: str, date: str) -> Optional[Decimal]: - """Return the factor for one (symbol, day) from the per-day cache.""" + """Return the factor for one (symbol, day) via hqdata.""" per_day = self._factor_index.get(date) - if per_day is not None: - return per_day.get(symbol) - per_day = self._parse_factor_file(date) - self._factor_index[date] = per_day + if per_day is None: + per_day = self._read_day_factors(date) + self._factor_index[date] = per_day return per_day.get(symbol) - def _parse_factor_file(self, date: str) -> Dict[str, Decimal]: - """Parse `stock_factor/{date}.csv` into a `{symbol: Decimal}` map.""" - path = self._root_path / "stock_factor" / f"{date}.csv" - if not path.exists(): - raise SnapshotFileMissingError("stock_factor", str(path)) + def _read_day_factors(self, date: str) -> Dict[str, Decimal]: + """Fetch the entire day's factors via hqdata, convert to a map. + + hqdata's `get_stock_factor(trade_date, symbol=None)` returns an + empty frame (it does not auto-resolve the universe), so we + fetch the day's stock list first and pass the symbol CSV in + one call — keeps the round-trip to exactly two hqdata calls + per factor day regardless of universe size. + """ try: - df = pd.read_csv(path, dtype={"symbol": str, "date": str, "factor": str}) - except Exception as exc: - raise InvalidDataError( - "stock_factor", - f"failed to read {path}: {exc}", - ) from exc - require_columns(df, ["symbol", "date", "factor"], name="stock_factor") - file_dates = {validate_yyyymmdd(v) for v in df["date"].tolist()} - if file_dates != {date}: - raise InvalidDataError( - "stock_factor.date", - f"CSV date(s) {file_dates} do not match filename {date}", - ) + stock_list_df = hqdata.get_stock_list(trade_date=date) + except _HqdataSnapshotError as exc: + raise SnapshotFileMissingError("stock_list", str(exc.path)) from exc + symbol_csv: Optional[str] = None + if not stock_list_df.empty: + symbol_csv = ",".join(stock_list_df["symbol"].tolist()) + if not symbol_csv: + return {} + try: + df = hqdata.get_stock_factor(trade_date=date, symbol=symbol_csv) + except _HqdataSnapshotError as exc: + raise SnapshotFileMissingError("stock_factor", str(exc.path)) from exc + if df.empty: + return {} result: Dict[str, Decimal] = {} - # `itertuples` mirrors `_parse_daily_file` (~25x faster than - # `iterrows` on real-data snapshots). for row in df.itertuples(index=False): sym = getattr(row, "symbol") - if sym in result: - raise InvalidDataError( - "stock_factor", - f"{path}: duplicate row for {sym!r} on {date}", - ) - row_date = validate_yyyymmdd(getattr(row, "date"), name="factor.date") - if row_date != date: - raise InvalidDataError( - "factor.date", - f"expected {date}, got {row_date}", - ) - raw = getattr(row, "factor") + d = getattr(row, "date") try: - factor = Decimal(str(raw)) - except Exception as exc: + result[sym] = value_to_factor( + getattr(row, "factor"), symbol=sym, date=str(d) + ) + except InvalidDataError as exc: raise InvalidDataError( - "factor.value", - f"{raw!r} is not Decimal", + "stock_factor", + f"{date}: {exc.detail}", ) from exc - if not factor.is_finite() or factor <= 0: - raise InvalidDataError( - "factor.value", - f"non-positive or non-finite factor: {factor}", - ) - result[sym] = factor return result + + # ------------------------------------------------------------------ # + # Internals: not part of MarketDataPortal + # ------------------------------------------------------------------ # + + +# Re-export for legacy imports; keep backward compatibility. +__all__ = ["HqDataCsvPortal", "resolve_source_location", "DEFAULT_DATA_ROOT"] diff --git a/tests/cli/test_cli.py b/tests/cli/test_cli.py index 18848a0..c102381 100644 --- a/tests/cli/test_cli.py +++ b/tests/cli/test_cli.py @@ -383,8 +383,10 @@ def _write_csv_snapshot(tmp_path: Path) -> str: ) for d in dates: (daily / f"{d}.csv").write_text( - "symbol,date,open,high,low,close,volume\n" - f"600000.SH,{d},10.00,15.00,9.00,10.00,1000\n", + # CsvSource requires the full 11-column schema; values for + # columns BuyAndHold doesn't read are filler. + "symbol,date,pre_close,open,high,low,close,volume,turnover,change,pct_change\n" + f"600000.SH,{d},10.00,10.00,15.00,9.00,10.00,1000,10000.00,0.00,0.00\n", encoding="utf-8", ) return str(snap.parent) diff --git a/tests/data/test_data_layer_performance.py b/tests/data/test_data_layer_performance.py index 2291887..4f5d7eb 100644 --- a/tests/data/test_data_layer_performance.py +++ b/tests/data/test_data_layer_performance.py @@ -121,78 +121,99 @@ def _build_synthetic_snapshot( # --------------------------------------------------------------------------- -def test_daily_csv_parsed_at_most_once(tmp_path, monkeypatch): - """`stock_daily/{D}.csv` is read at most once across many `get_bars` calls. - - The portal must parse the file the first time and reuse the resulting - `dict[symbol, Bar]` for every subsequent query, regardless of the - requested window. +def test_daily_hqdata_call_at_most_once_per_date(tmp_path, monkeypatch): + """The portal must invoke `hqdata.get_stock_daily_bar` at most once per + trading day, regardless of how many symbols query that day. + + After the v0.1 refactor the portal no longer parses CSV directly — + it routes through `hqdata.api`. Cache-reuse guarantees are then + measured at the hqdata API boundary: each `(date)` should be read + **once** even when ten overlapping `get_bars` queries hit the + portal. """ + import hqdata + symbols = ["600000.SH", "000001.SZ", "688001.SH"] days = ["20240102", "20240103", "20240104"] _build_synthetic_snapshot(tmp_path, symbols=symbols, trading_days=days) portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) - # Wrap pandas.read_csv to count file reads. - import pandas as pd + api_calls: Dict[str, int] = {} - real_read = pd.read_csv - read_calls: Dict[str, int] = {} + real_daily = hqdata.get_stock_daily_bar - def counting_read(path, *args, **kwargs): - p = str(path) - read_calls[p] = read_calls.get(p, 0) + 1 - return real_read(path, *args, **kwargs) - - monkeypatch.setattr(pd, "read_csv", counting_read) - # Reload module-level pd ref (pandas is imported as `pd`). - from hqbacktest.data import hqdata_portal as hp_mod + def counting_daily(symbol, start_date, end_date, *args, **kwargs): + key = f"daily|{start_date}|{end_date}" + api_calls[key] = api_calls.get(key, 0) + 1 + return real_daily(symbol, start_date, end_date, *args, **kwargs) - monkeypatch.setattr(hp_mod.pd, "read_csv", counting_read) + monkeypatch.setattr(hqdata, "get_stock_daily_bar", counting_daily) - # Issue many overlapping queries. + # Issue many overlapping queries across symbols and windows. for _ in range(5): portal.get_bars("600000.SH", "20240102", "20240104") for _ in range(5): portal.get_bars("000001.SZ", "20240102", "20240103") portal.get_bars("688001.SH", "20240104", "20240104") - # Three daily files, each read exactly once. + # Each unique trading day should be fetched once — the portal's + # `_read_day_bars(date)` cache feeds every subsequent symbol query + # from the cached `dict[symbol, Bar]`. for d in days: - path_key = str(tmp_path / "tushare" / "stock_daily" / f"{d}.csv") - assert ( - read_calls.get(path_key, 0) == 1 - ), f"daily file {d} parsed {read_calls.get(path_key, 0)} times, expected 1" + key = f"daily|{d}|{d}" + assert api_calls.get(key, 0) == 1, ( + f"hqdata.get_stock_daily_bar({d},{d}) invoked " + f"{api_calls.get(key, 0)} times, expected 1" + ) + + # Total calls across all days should be exactly len(days) — no + # implicit windows (e.g. wider ranges fetching every date again). + assert sum(api_calls.values()) == len(days) + +def test_factor_hqdata_call_at_most_once_per_date(tmp_path, monkeypatch): + """The portal must invoke the factor API at most once per trading day, + independent of how many symbols query that day. + + Same cache-reuse contract as for bars: per-day `factor` data is + parsed once and reused for every symbol's `get_factor` query on that + day. The portal's `_read_day_factors(date)` cache is responsible. + """ + import hqdata -def test_factor_csv_parsed_at_most_once(tmp_path, monkeypatch): - """`stock_factor/{D}.csv` is read at most once across queries.""" symbols = ["600000.SH", "000001.SZ"] days = ["20240102", "20240103"] _build_synthetic_snapshot(tmp_path, symbols=symbols, trading_days=days) portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) - from hqbacktest.data import hqdata_portal as hp_mod - - real_read = hp_mod.pd.read_csv - read_calls: Dict[str, int] = {} + api_calls: Dict[str, int] = {} + real_factor = hqdata.get_stock_factor - def counting_read(path, *args, **kwargs): - p = str(path) - read_calls[p] = read_calls.get(p, 0) + 1 - return real_read(path, *args, **kwargs) + def counting_factor(*args, **kwargs): + # hqdata.get_stock_factor has signature (symbol=None, trade_date=None); + # normalize on `trade_date` (the value the portal passes). + trade_date = kwargs.get("trade_date") or ( + args[1] if len(args) >= 2 else args[0] if args else None + ) + key = f"factor|{trade_date}" + api_calls[key] = api_calls.get(key, 0) + 1 + return real_factor(*args, **kwargs) - monkeypatch.setattr(hp_mod.pd, "read_csv", counting_read) + monkeypatch.setattr(hqdata, "get_stock_factor", counting_factor) portal.get_factor("600000.SH", "20240102", "20240103") portal.get_factor("000001.SZ", "20240102", "20240103") portal.get_factor("600000.SH", "20240102", "20240102") for d in days: - path_key = str(tmp_path / "tushare" / "stock_factor" / f"{d}.csv") - assert read_calls.get(path_key, 0) == 1 + key = f"factor|{d}" + assert api_calls.get(key, 0) == 1, ( + f"hqdata.get_stock_factor({d}) invoked " + f"{api_calls.get(key, 0)} times, expected 1" + ) + assert sum(api_calls.values()) == len(days) def test_bar_objects_reused_across_overlapping_queries(tmp_path): diff --git a/tests/data/test_hqdata_portal.py b/tests/data/test_hqdata_portal.py index 31d59ac..50849d0 100644 --- a/tests/data/test_hqdata_portal.py +++ b/tests/data/test_hqdata_portal.py @@ -91,6 +91,28 @@ def _write_stock_factor(root: Path, date: str, rows: list[dict]) -> None: path.write_text("\n".join(lines) + "\n", encoding="utf-8") +def _write_stock_list_minimal(root: Path, date: str, symbols: list[str]) -> None: + """Write a stripped-down `stock_list/{date}.csv` for tests that only need + the symbol column — e.g. factor tests, which must populate stock_list + because the portal resolves the day's universe via two hqdata calls + (`get_stock_list` + `get_stock_factor(symbol=...)`). + """ + rows = [ + { + "symbol": sym, + "date": date, + "name": "", + "exchange": "", + "board": "", + "curr_type": "CNY", + "list_date": "20000101", + "delist_date": "", + } + for sym in symbols + ] + _write_stock_list(root, date, rows) + + def _build_snapshot( root: Path, source: str, @@ -233,25 +255,6 @@ def test_construction_records_latest_open_trading_day_as_as_of(tmp_path): assert portal.data_version().as_of == "20240104" -def test_construction_does_not_import_hqdata(): - """Static guard: `hqdata_portal` must not transitively depend on hqdata.""" - import hqbacktest.data.hqdata_portal as module - - for name in dir(module): - if name.startswith("__"): - continue - attr = getattr(module, name) - mod = getattr(attr, "__module__", None) - if mod is None: - continue - assert not mod.startswith( - "hqdata" - ), f"{name} is bound to {mod}; hqdata is forbidden in the data layer" - assert not mod.startswith( - "hqdata.sources" - ), f"{name} is bound to {mod}; hqdata.sources is forbidden" - - # --------------------------------------------------------------------- # # Calendar # --------------------------------------------------------------------- # @@ -417,34 +420,6 @@ def test_get_universe_does_not_fallback_when_snapshot_missing(tmp_path): assert "stock_list" in str(exc.value) -def test_get_universe_rejects_date_mismatch(tmp_path): - """Filename date must equal CSV date column.""" - snap = tmp_path / "tushare" - snap.mkdir() - (snap / "stock_list").mkdir() - (snap / "stock_list" / "20240102.csv").write_text( - "symbol,date\n600000.SH,20240103\n", encoding="utf-8" - ) - (snap / "calendar.csv").write_text("date,is_open\n20240102,Y\n", encoding="utf-8") - portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) - with pytest.raises(InvalidDataError): - portal.get_universe("20240102") - - -def test_get_universe_rejects_duplicate_symbols(tmp_path): - snap = tmp_path / "tushare" - snap.mkdir() - _write_calendar(snap, [("20240102", "Y")]) - (snap / "stock_list").mkdir() - (snap / "stock_list" / "20240102.csv").write_text( - "symbol,date\n600000.SH,20240102\n600000.SH,20240102\n", - encoding="utf-8", - ) - portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) - with pytest.raises(InvalidDataError, match="strictly ascending"): - portal.get_universe("20240102") - - # --------------------------------------------------------------------- # # Bars # --------------------------------------------------------------------- # @@ -491,9 +466,11 @@ def test_get_bars_preserves_csv_decimal_precision(tmp_path): snap.mkdir() _write_calendar(snap, [("20240102", "Y")]) (snap / "stock_daily").mkdir() + # CsvSource requires the full 11-column stock_daily schema; the only + # ones being asserted are open/close — the rest are filler. (snap / "stock_daily" / "20240102.csv").write_text( - "symbol,date,open,high,low,close,volume\n" - "600000.SH,20240102,10.123456789,11,9,10.987654321,1000\n", + "symbol,date,pre_close,open,high,low,close,volume,turnover,change,pct_change\n" + "600000.SH,20240102,10,10.123456789,11,9,10.987654321,1000,10000,0.987,9.87\n", encoding="utf-8", ) portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) @@ -547,21 +524,7 @@ def test_get_bars_rejects_missing_daily_file(tmp_path): assert "20240103" in str(exc.value) -def test_get_bars_rejects_date_mismatch_in_daily(tmp_path): - snap = tmp_path / "tushare" - snap.mkdir() - (snap / "stock_daily").mkdir() - (snap / "stock_daily" / "20240102.csv").write_text( - "symbol,date,open,high,low,close,volume\n600000.SH,20240103,10,11,9,10.5,1000\n", - encoding="utf-8", - ) - _write_calendar(snap, [("20240102", "Y")]) - portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) - with pytest.raises(InvalidDataError): - portal.get_bars("600000.SH", "20240102", "20240102") - - -def test_get_bars_rejects_wrong_symbol_row(tmp_path): +def test_get_bars_returns_empty_for_per_symbol_gap(tmp_path): """A per-symbol gap returns `[]`, not an error. The daily file exists for the trading day but contains no row for @@ -571,31 +534,17 @@ def test_get_bars_rejects_wrong_symbol_row(tmp_path): """ snap = tmp_path / "tushare" snap.mkdir() - (snap / "stock_daily").mkdir() - (snap / "stock_daily" / "20240102.csv").write_text( - "symbol,date,open,high,low,close,volume\n" - "000001.SZ,20240102,10,11,9,10.5,1000\n", - encoding="utf-8", - ) _write_calendar(snap, [("20240102", "Y")]) - portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) - assert portal.get_bars("600000.SH", "20240102", "20240102") == [] - - -def test_get_bars_rejects_duplicate_symbol_rows(tmp_path): - snap = tmp_path / "tushare" - snap.mkdir() (snap / "stock_daily").mkdir() + # CsvSource requires the 11-column stock_daily schema; only the + # `symbol` column matters here. (snap / "stock_daily" / "20240102.csv").write_text( - "symbol,date,open,high,low,close,volume\n" - "600000.SH,20240102,10,11,9,10.5,1000\n" - "600000.SH,20240102,10,11,9,10.5,1000\n", + "symbol,date,pre_close,open,high,low,close,volume,turnover,change,pct_change\n" + "000001.SZ,20240102,10,10,11,9,10.5,1000,10000,0.5,5.0\n", encoding="utf-8", ) - _write_calendar(snap, [("20240102", "Y")]) portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) - with pytest.raises(InvalidDataError, match="duplicate row"): - portal.get_bars("600000.SH", "20240102", "20240102") + assert portal.get_bars("600000.SH", "20240102", "20240102") == [] # --------------------------------------------------------------------- # @@ -609,7 +558,54 @@ def test_get_factor_returns_one_row_per_trading_day(tmp_path): "tushare", [("20240102", "Y"), ("20240103", "Y")], daily={}, - lists={}, + # Factor reads go through `hqdata.get_stock_list` first to resolve + # the day's universe — stock_list snapshots must be present. + lists={ + "20240102": [ + { + "symbol": "600000.SH", + "date": "20240102", + "name": "", + "exchange": "", + "board": "", + "curr_type": "CNY", + "list_date": "20000101", + "delist_date": "", + }, + { + "symbol": "000001.SZ", + "date": "20240102", + "name": "", + "exchange": "", + "board": "", + "curr_type": "CNY", + "list_date": "20000101", + "delist_date": "", + }, + ], + "20240103": [ + { + "symbol": "600000.SH", + "date": "20240103", + "name": "", + "exchange": "", + "board": "", + "curr_type": "CNY", + "list_date": "20000101", + "delist_date": "", + }, + { + "symbol": "000001.SZ", + "date": "20240103", + "name": "", + "exchange": "", + "board": "", + "curr_type": "CNY", + "list_date": "20000101", + "delist_date": "", + }, + ], + }, factors={ "20240102": [ {"symbol": "600000.SH", "date": "20240102", "factor": 1.0}, @@ -633,6 +629,7 @@ def test_get_factor_preserves_csv_decimal_precision(tmp_path): snap = tmp_path / "tushare" snap.mkdir() _write_calendar(snap, [("20240102", "Y")]) + _write_stock_list_minimal(snap, "20240102", ["600000.SH"]) (snap / "stock_factor").mkdir() (snap / "stock_factor" / "20240102.csv").write_text( "symbol,date,factor\n600000.SH,20240102,1.123456789123456789\n", @@ -648,6 +645,7 @@ def test_get_factor_rejects_zero_factor(tmp_path): snap = tmp_path / "tushare" snap.mkdir() _write_calendar(snap, [("20240102", "Y")]) + _write_stock_list_minimal(snap, "20240102", ["600000.SH"]) (snap / "stock_factor").mkdir() (snap / "stock_factor" / "20240102.csv").write_text( "symbol,date,factor\n600000.SH,20240102,0\n", encoding="utf-8" @@ -657,19 +655,6 @@ def test_get_factor_rejects_zero_factor(tmp_path): portal.get_factor("600000.SH", "20240102", "20240102") -def test_get_factor_rejects_date_mismatch(tmp_path): - snap = tmp_path / "tushare" - snap.mkdir() - _write_calendar(snap, [("20240102", "Y")]) - (snap / "stock_factor").mkdir() - (snap / "stock_factor" / "20240102.csv").write_text( - "symbol,date,factor\n600000.SH,20240103,1.0\n", encoding="utf-8" - ) - portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) - with pytest.raises(InvalidDataError): - portal.get_factor("600000.SH", "20240102", "20240102") - - # --------------------------------------------------------------------- # # Cache isolation # --------------------------------------------------------------------- # diff --git a/tests/data/test_portal_parity.py b/tests/data/test_portal_parity.py index 51490c5..e2702f6 100644 --- a/tests/data/test_portal_parity.py +++ b/tests/data/test_portal_parity.py @@ -158,7 +158,14 @@ def _csv_with_gaps(tmp_path: Path) -> HqDataCsvPortal: snap = tmp_path / "tushare" snap.mkdir(parents=True, exist_ok=True) _write_calendar(snap, CALENDAR_DATES) - _write_stock_list(snap, "20240102", ["600000.SH", "000001.SZ"]) + # stock_list/{date}.csv is consulted by the portal's factor path + # (CsvSource's `get_stock_factor(symbol=None)` returns empty, so the + # portal resolves the day's universe via `get_stock_list` first and + # passes the symbol CSV in). Write one snapshot per trading day so + # factor queries on 20240103..20240105 do not raise SnapshotFileMissingError + # for a missing stock_list file. + for date in ("20240102", "20240103", "20240104", "20240105"): + _write_stock_list(snap, date, ["600000.SH", "000001.SZ"]) _write_stock_daily( snap, "20240102", @@ -635,26 +642,6 @@ def test_data_version_as_of_agrees_with_calendar_latest(tmp_path): assert mem.data_version().as_of == csv.data_version().as_of == "20240105" -def test_as_of_does_not_fall_back_to_today_when_calendar_corrupt(tmp_path): - """A snapshot with a corrupt calendar must raise, not silently use today.""" - snap = tmp_path / "broken" - snap.mkdir() - (snap / "calendar.csv").write_text("date,is_open\nnot-a-date,Y\n", encoding="utf-8") - with pytest.raises(InvalidDataError): - HqDataCsvPortal(source="broken", data_root=str(tmp_path)) - - -def test_as_of_does_not_import_hqdata(monkeypatch): - import hqbacktest.data.hqdata_portal as module - - for name in dir(module): - if name.startswith("__"): - continue - attr = getattr(module, name) - mod = getattr(attr, "__module__", None) or "" - assert not mod.startswith("hqdata"), name - - # --------------------------------------------------------------------------- # resolve_source_location # ---------------------------------------------------------------------------