diff --git a/CHANGELOG.md b/CHANGELOG.md index 26b0dc0..09c0cc2 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,110 @@ All notable changes to `hqbacktest` are documented in this file. The format follows [Keep a Changelog](https://keepachangelog.com/en/1.1.0/), and the project adheres to [Semantic Versioning](https://semver.org/). +## [0.1.1] - 2026-08-25 + +Patch release that hardens `hqbacktest` against the v0.1 real-data +review (see [TODO.md](./TODO.md)「v0.1 评审结论」). All changes are +backward-compatible unless called out below. + +### Added +- **Data layer hardening (task 14):** `get_bars` / `get_factor` allow + per-day gaps in the window; `SnapshotFileMissingError` (subclass of + `MissingDataError`) distinguishes a missing whole-day snapshot from + a per-symbol gap. `current_price` walks back up to 20 trading days + for the most recent valid close. The first-trading-day sentinel + `visible_through="00000000"` no longer raises; `history` returns + `[]` and `current_price` returns `None`. `InMemoryDataPortal` + drops its forward-walk universe fallback to match the CSV portal's + per-date semantics. `.BJ` symbols are excluded by default + (`include_bj=True` opt-in). `Bar.volume` is documented as **手**. + Defensive copies returned from cached lists. +- **Data layer performance (task 15):** per-day file cache + (`{date: {symbol: Bar}}`) plus per-symbol cumulative sequences + (`_symbol_bars[symbol]`) with `bisect` slicing. The 5-symbol + moving-average strategy over the v0.1.1-calibrated 139-day window + finishes well under the 60-second budget. +- **Match & ledger semantics (task 16):** `SimulatedBroker.match` + matches all SELL orders first, then BUYs, with rolling cash so + same-day "卖旧买新" rotations are not falsely rejected for cash. + SELL orders are no longer lot-rounded; `order_target(symbol, 0)` + can flatten a position that contains odd-lot shares. BUY-only lot + rounding is preserved. 7 contract-level invariants are pinned + in `docs/design/mvp-contract.md` §3.4 (T+1 whole-order rejection, + `realized_pnl` excludes fees, `ROUND_HALF_EVEN`, etc.). +- **Equity curve & metrics baseline (task 17):** first-day P&L now + flows into `daily_return` and `drawdown` (anchored to + `initial_cash`), so the chained-product identity `∏(1+daily_return) + == 1+total_return` holds for any run length. `daily_volatility` + returns `None` for runs with fewer than 2 daily returns (no more + misleading 0 / `nan`). `metrics.py` rebuilds `Decimal` via + `Decimal(str(...))` to avoid `Decimal(float)` artifacts. +- **Strategy isolation & audit trail (task 18):** `Order` is now + `@dataclass(frozen=True)` with `fill_ids: tuple[str, ...]`; strategies + cannot mutate Order objects returned from `Context.pending_orders()`. + `DataView.portal` is now a private `_portal` field; strategies cannot + bypass `visible_through`. `set_universe(...)` enforces trading scope; + orders outside the universe are rejected with + `RejectReason.OUT_OF_UNIVERSE`. New `Context.historical_universe()` + returns the historical stock list through the guarded data view. +- **Factor diagnostics on holdings (task 19):** the engine runs + `analyze_factor_series` against holdings-period factor series + with a 0.1% relative jump band. Any holding-period factor jump emits + a `DATA_WARNING` event and a `FactorDiagnostic` entry; the + diagnostics are observability-only — cash, position and equity are + byte-identical with the no-diagnostics baseline. CLI prints a one-line + summary at run end when diagnostics fired. **`adjustment_policy=none` + still excludes dividends from the NAV** (contract task 9 invariant); + the diagnostics surface this bias; long-window NAV remains + unsuitable for return estimation. +- **CLI first-mile + documentation honesty (task 20):** + `hqbacktest run` (the console script) prepends the config file's + directory and the current working directory to `sys.path` so the + strategy module can be resolved by name alone (matching + `python -m hqbacktest run`). Config validation rejects `nan` / + `inf` / float `initial_cash`, impossible calendar dates, and + empty trading-day windows with single-line `ConfigError` (CLI exit 2). + Output directories that already contain prior-run files are rejected + with exit 3; `--force` overrides. `Context.order_value` accepts + `int` / `str` cash amounts. `run_metadata.json`'s `git_commit` now + records the hqbacktest package's own commit (not the user's cwd). + README "项目状态" / "命令行" / "错误信息" / "包布局" sections + brought into line with the implementation. `BaseStrategy.__init__` + accepts and stores `**kwargs` so `[strategy].kwargs` round-trips. + +### Added (test infrastructure) +- **`tests/integration/`** (task 21): four real-data smoke scenarios + against `~/.hqdata/tushare`, auto-skipped when the snapshot is + missing (no credentials, no network): + 1. buy_and_hold across 600000.SH's 2026-07-16 dividend ex-date — the + task-14/19 contract (`factor jump 16.5935 → 17.3774` produces a + `DATA_WARNING`) is enforced end-to-end. + 2. 5-symbol moving-average strategy over the full 139-day window — + byte-deterministic across two runs and below the 60-second budget. + 3. universe containing the known suspended symbol `000008.SZ` + (suspended 2026-07-07..2026-07-13) — no crash, the fallback-close + valuation `DATA_WARNING` is recorded. + 4. first-trading-day `before_trading_start` reads `current_price` — + returns `None` against the sentinel without crashing. + +### Constraints (unchanged from v0.1) +- `adjustment_policy` MUST be `"none"`. +- Market orders only; limit / stop / partial fills raise + `UnsupportedOrderTypeError` / are rejected. +- Only Chinese A-share common stocks (沪深); no ST / 涨跌停 / 新股 / + 北交所 / 融资融券 / 期权 support. +- Default A-share cost model only. +- CSV-only data ingestion via hqdata CLI; no network calls; no tokens + read or written. + +### Known limitations carried forward +- `adjustment_policy=none` means the NAV systematically underestimates + cross-ex-date windows (no dividend accounting). Task-19 factor + diagnostics surface the jumps but do not fabricate dividends. +- `Position.update_buy` uses simple-average cost (not FIFO). +- `BacktestResult.metrics` reconstruction on `load()` is best-effort. +- README's "路线图" section lists capabilities deferred to v0.2+. + ## [0.1.0] - 2026-08-23 First public release of `hqbacktest`. The project implements tasks 1-13 of @@ -66,4 +170,5 @@ First public release of `hqbacktest`. The project implements tasks 1-13 of byte-stable across runs; the live engine retains the full set. - README's §路线图 section lists capabilities deferred to v0.2+. +[0.1.1]: https://github.com/HonestQuantTech/hqbacktest/releases/tag/v0.1.1 [0.1.0]: https://github.com/HonestQuantTech/hqbacktest/releases/tag/v0.1.0 diff --git a/README.md b/README.md index f870416..e8e9413 100644 --- a/README.md +++ b/README.md @@ -1,17 +1,20 @@ # hqbacktest - A股量化策略回测与交易模拟引擎

- +

## 项目状态 -`hqbacktest` 当前处于**v0.1 发布候选**: +`hqbacktest` 当前发布 **`v0.1.1`**: -- **已实现(任务 1–13 完成):** 产品契约、可安装的 Python 包、领域模型(订单、成交、持仓、账本、快照)、订单状态机、`AdjustmentPolicy` 枚举、`CorporateAction` 数据结构草案、`Decimal` 精度与 JSON 序列化;`MarketDataPortal`、`HqDataCsvPortal`(CSV 快照门户)、`InMemoryDataPortal`、`DataView`、内存缓存和无未来函数校验;日频事件时钟、`BacktestEngine`、五阶段调度(`SESSION_START → BEFORE_TRADING_START → OPEN_MATCH → BAR_CLOSE → AFTER_TRADING_END`)、按阶段的数据可见性切换和可追溯事件日志;`BaseStrategy` 生命周期与受控 `Context` API、下单意图;`SimulatedBroker`(`OPEN_MATCH` 阶段按当日 `bar.open` 全额成交市价单);`TradingRuleSet`(`LongOnly` / `LotSize` / `NonTradingDay` / `InvalidPrice` / `InsufficientCash` / `T1Sellable` 六条默认规则)和 `CostModel`(默认 A 股费率:0.025% 佣金 + 5 元保底 + 0.1% 卖出印花税,0 过户费;所有费率在 README 与代码中显式声明,无隐藏常量);账本拒绝原因(`INSUFFICIENT_CASH` / `INSUFFICIENT_SHARES`)和 T+1 日终结算;末交易日 `BACKTEST_ENDED` 自动撤销。`BacktestConfig.adjustment_policy` 严格只接受 `"none"`;`CorporateActionProvider` 为设计草案,因子诊断接口与 `analyze_factor_series` 分析器已就位。`BacktestResult` 含 `equity_curve` / `orders_table` / `fills_table` / `positions_table` / `costs_table` / `PerformanceMetrics` / `events.jsonl` / `data_version` / `factor_diagnostics`;`save(dir)` / `load(dir)` 导出 CSV+JSON 并可重建。`examples/buy_and_hold.py` 与 `examples/moving_average.py` 用公共 API 跑通端到端流程并有 7 天确定性 `InMemoryDataPortal` 数据 fixture,10 项端到端回归测试覆盖完整生命周期。`hqbacktest run --config FILE --output DIR` 命令行(`hqbacktest/cli/` 包,TOML 配置 + 校验 + 策略导入 + 元数据 + 独立输出目录,绝不泄露凭证);26 项 CLI 测试覆盖端到端、配置验证、可复现性与错误信息。`.github/workflows/ci.yml` 覆盖 Python 3.10 / 3.11 / 3.12、`black`、`pytest`、`pytest-cov`、示例 smoke 与 CLI smoke;`python -m build` 产出 sdist + wheel;`CHANGELOG.md` 记录 v0.1。 +- **v0.1.1 新增(任务 14–21):** 数据层缺行/停牌/首日语义修复(任务 14);按日文件缓存 + 按 symbol 累积序列 + bisect 切片,性能从「小时级」降到「秒级」(任务 15);同批撮合 SELL→BUY 滚动现金 + SELL 零股可卖 + 钉死 T+1/realized_pnl/ROUND_HALF_EVEN/撮合顺序(任务 16);首日 P&L 进入收益曲线 + running peak 含 `initial_cash` + 波动率样本不足返回 `None` + 恒等式 ∏(1+r)=1+total_return(任务 17);`Order` 不可变 + `DataView.portal` 私有 + universe 生效(任务 18);持仓期间因子跳变自动诊断 + CLI 汇总警告 + 账本零影响(任务 19);console script 策略模块导入 + nan/inf/空窗口校验 + 输出目录防护 + `order_value` 接受 int/str + 文档一致性(任务 20);`tests/integration/` 真实数据冒烟基线(任务 21)。 +- **v0.1 已实现(任务 1–13):** 产品契约、可安装的 Python 包、领域模型(订单、成交、持仓、账本、快照)、订单状态机、`AdjustmentPolicy` 枚举、`CorporateAction` 数据结构草案、`Decimal` 精度与 JSON 序列化;`MarketDataPortal`、`HqDataCsvPortal`(CSV 快照门户)、`InMemoryDataPortal`、`DataView`、内存缓存和无未来函数校验;日频事件时钟、`BacktestEngine`、五阶段调度(`SESSION_START → BEFORE_TRADING_START → OPEN_MATCH → BAR_CLOSE → AFTER_TRADING_END`)、按阶段的数据可见性切换和可追溯事件日志;`BaseStrategy` 生命周期与受控 `Context` API、下单意图;`SimulatedBroker`(`OPEN_MATCH` 阶段按当日 `bar.open` 全额成交市价单);`TradingRuleSet`(`LongOnly` / `LotSize` / `NonTradingDay` / `InvalidPrice` / `InsufficientCash` / `T1Sellable` 六条默认规则)和 `CostModel`(默认 A 股费率:0.025% 佣金 + 5 元保底 + 0.1% 卖出印花税,0 过户费);账本拒绝原因(`INSUFFICIENT_CASH` / `INSUFFICIENT_SHARES`)和 T+1 日终结算;末交易日 `BACKTEST_ENDED` 自动撤销。`BacktestConfig.adjustment_policy` 严格只接受 `"none"`。`BacktestResult` 含 `equity_curve` / `orders_table` / `fills_table` / `positions_table` / `costs_table` / `PerformanceMetrics` / `events.jsonl` / `data_version` / `factor_diagnostics`;`save(dir)` / `load(dir)` 导出 CSV+JSON 并可重建。`examples/buy_and_hold.py` 与 `examples/moving_average.py` 用公共 API 跑通端到端流程并有 7 天确定性 `InMemoryDataPortal` 数据 fixture。`hqbacktest run --config FILE --output DIR` 命令行(`hqbacktest/cli/` 包,TOML 配置 + 校验 + 策略导入 + 元数据 + 独立输出目录,绝不泄露凭证)。`.github/workflows/ci.yml` 覆盖 Python 3.10 / 3.11 / 3.12、`black`、`pytest`、`pytest-cov`、示例 smoke 与 CLI smoke;`python -m build` 产出 sdist + wheel;`CHANGELOG.md` 记录 v0.1 与 v0.1.1。 -- **不在 v0.1 内(路线图):** 限价 / 止损单、成交量参与率、部分成交;ST / 涨跌停 / 新股 / 北交所规则;融资融券 / 期货 / 期权 / 多账户;指数基准与归因;Notebook 与远程策略入口;JSON Schema 校验以外的策略注册中心;交互图表 / HTML 报告。 +> **⚠️ v0.1.1 仍未做:** 分红会计 / 涨跌停 / 新股 / 北交所 / 限价单 / 指数基准 / 多账户 / 分钟线 / 实盘对接 —— 见 [CHANGELOG.md](CHANGELOG.md) 与 [TODO.md](TODO.md) 「发布后再排期的增强项」。`adjustment_policy=none` 下跨除权日的净值仍**系统性低估**(少分红现金),任务 19 因子诊断会显式记录此类跳变,但长区间结果不可直接用于收益评估。 + +- **不在 v0.1.1 内(路线图):** 限价 / 止损单、成交量参与率、部分成交;ST / 涨跌停 / 新股 / 北交所规则;融资融券 / 期货 / 期权 / 多账户;指数基准与归因;Notebook 与远程策略入口;JSON Schema 校验以外的策略注册中心;交互图表 / HTML 报告。 本文描述的用法、命令行和功能表与 [`docs/design/mvp-contract.md`](docs/design/mvp-contract.md) 一致;其中的示例已经可以按 §示例 章节运行。功能表区分「已实现」与「路线图 / 计划中」;任何契约变更必须先更新契约文档。开发顺序与 AI 协作提示见 [TODO.md](TODO.md)。 @@ -34,8 +37,9 @@ | 功能 | 目标接口/产物 | 首版语义 | 当前状态 | | --- | --- | --- | :---: | -| 交易日与历史股票池 | `MarketDataPortal` | 按回测日获取交易日和股票池,避免以今日股票列表产生幸存者偏差 | 已实现 | -| 日线数据可见性 | `DataView.history()` | 盘前最多看到前一交易日;当天收盘后才可读取当天日线 | 已实现 | +| 交易日与历史股票池 | `MarketDataPortal` | 按回测日获取交易日和股票池,避免以今日股票列表产生幸存者偏差;`.BJ` 股票默认过滤,`include_bj=True` 可保留 | 已实现 | +| 日线数据可见性 | `DataView.history()` | 盘前最多看到前一交易日;当天收盘后才可读取当天日线;首日盘前哨兵 `visible_through="00000000"` 不抛异常 | 已实现 | +| 缺行/停牌/估值口径 | `get_bars`/`DataView.current_price`/日终估值 | `get_bars` 允许逐日间隙;停牌持仓按 20 日回看最近收盘估值并写 `DATA_WARNING`;整日快照缺失抛 `SnapshotFileMissingError`;`Bar.volume` 单位「手」 | 已实现 | | 日频事件时钟 | `BacktestEngine`、`EventLog` | 五阶段固定顺序;盘前 D-1、收盘 D 的可见性切换;事件日志记录日期、阶段与错误原因;策略异常带日期和阶段 | 已实现 | | 策略生命周期 | `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_*()` | 首版只支持市价委托,策略只能提交意图,不能直接改账户;仅盘前与收盘回调可下单,订单创建/撤销写入事件日志 | 已实现 | @@ -52,10 +56,12 @@ ### 已规划的支持范围 -- **市场与频率:** 沪深普通股票的日线回测;标的使用 `600000.SH`、`000001.SZ` 这类统一代码。 +- **市场与频率:** 沪深普通股票的日线回测;标的使用 `600000.SH`、`000001.SZ` 这类统一代码。`.BJ`(北交所)股票默认从 `get_universe` 中过滤,需要时通过 `include_bj=True` 显式启用。 - **账户:** 单个人民币现金账户、股票现货多头;不使用杠杆或保证金。 - **数据:** 每次回测固定使用一个 `hqdata` 数据源。需要日线时,首选 Tushare 或 RiceQuant;当前 `hqdata` 的 AkShare 适配器不提供稳定的日线能力,不能作为首版回测数据源。 - **时间:** 日期一律使用 `YYYYMMDD`。`before_trading_start(D)` 只能访问 D-1 及以前的数据,可在 D 开盘参与撮合;`on_bar(D)` 在 D 收盘后才看到 D 日线,所提交订单最早在 D+1 开盘处理。 +- **缺行与停牌(任务 14):** `get_bars` 允许逐日间隙(窗口内无任何行返回 `[]`);个股当日缺行(停牌 / 未上市 / 已退市)属正常业务结果;**整日快照文件缺失** 是基础设施错误,必须中止运行。停牌持仓估值采用「最近 20 个交易日内最近一个有效收盘价」并写入 `DATA_WARNING` 事件;首日盘前哨兵 `visible_through="00000000"` 不抛异常。 +- **成交量单位:** `Bar.volume` 单位为「**手**」(1 手 = 100 股;与 Tushare `hqdata` 适配器口径一致);需要股数时应乘以 `LOT_SIZE`。 - **成交:** 首版市价单按符合规则的开盘价全额成交;订单、拒绝、费用和成交都要保留可追溯记录。 - **复权:** 成交、现金账本和 v0.1 净值均使用未复权价格,且 `adjustment_policy` 固定为 `none`。同源复权因子可用于数据质量诊断,但不用于伪造现金分红、送配、配股或税务会计。 - **结果:** 每次运行应导出净值曲线、订单、成交、每日持仓、成本、配置和运行元数据。 @@ -103,9 +109,11 @@ hqbacktest/ ├── pyproject.toml # 构建、依赖、pytest 与 black 配置 ├── src/hqbacktest/ # 引擎源码(src 布局) │ ├── __init__.py # 版本及稳定的公开 API 导出 +│ ├── __main__.py # `hqbacktest run` CLI 入口(任务 12) │ ├── domain/ # 任务 3 的模型、状态机、精度与序列化 │ ├── data/ # 任务 4 的数据门户、DataView、缓存与校验 -│ └── engine/ # 任务 5–6 的事件时钟、调度器、BacktestEngine 与受控策略接口 +│ ├── engine/ # 任务 5–6 的事件时钟、调度器、BacktestEngine 与受控策略接口 +│ └── cli/ # 任务 12 的 TOML 配置解析、CLI runner、退出码 ├── tests/ # 单元测试,必须不依赖网络或本地行情文件 ├── examples/ # 端到端示例(任务 11:buy_and_hold / moving_average) └── docs/design/ # 设计文档(如 mvp-contract.md) @@ -124,9 +132,75 @@ hqbacktest/ 数据集必须包含 `calendar.csv`,以及按交易日组织的 `stock_list/{YYYYMMDD}.csv`、`stock_daily/{YYYYMMDD}.csv` 与 `stock_factor/{YYYYMMDD}.csv`。这些文件由 `hqdata` CLI 在回测前写入;hqbacktest 既不下载数据,也不保存凭证。任何真实 token、账户号或私密配置都**不应**提交到仓库,也不应出现在回测结果目录中。 +### 性能与内存(任务 15) + +`HqDataCsvPortal` 在单次回测中按「按日文件缓存 + 按 symbol 累积序列」两层缓存: + +- 每个 `stock_daily/{D}.csv` / `stock_factor/{D}.csv` 在一次运行中最多解析一次;解析结果以 `{date: {symbol: Bar}}` 形式缓存。 +- 每个 symbol 在内存中维护一个按日期升序的累积序列;`get_bars` / `get_factor` 在该序列上做 `bisect` 切片,单次调用 O(log N)。 +- `DataView.history` 走累积缓存,单次 `get_bars` 切片即可;`DataView.current_price` 仅需一次 `get_calendar`(有缓存)确定 20 个交易日回看起点 + 一次 `get_bars`,二者都避免了旧实现的逐日 `get_bars(day, day)` 往返。 +- `Bar` / 因子对象在重叠窗口间复用,仅返回列表的防御性拷贝。 + +真实数据基准(`~/.hqdata/tushare`,20260105–20260731,139 个交易日,每个 daily 文件约 5000 行): + +| 场景 | 总耗时(含首次数据加载) | 任务 15 目标 | +| --- | --- | --- | +| 5 stocks × 139 days MA 策略 | ~7.6 s | < 10 s ✅ | +| 300 stocks × 139 days MA 策略 | ~9.4 s | < 120 s ✅ | + +内存量级:每个 `Bar` 约 200 字节。覆盖完整窗口(5000 symbols × 139 days ≈ 70 万 Bar)约 140 MB;策略触及的 universe 通常远小于全市场。`_symbol_bars` 累积只对真实访问过的 symbol 增长。 + +性能冒烟测试在 `tests/data/test_task15_performance.py`,50 symbols × 250 days 全量 `history(bar_count=20)` 在 15 秒阈值内完成。 + +### 撮合与账本口径(任务 16) + +- **同日撮合顺序:** 单个 `OPEN_MATCH(today)` batch 内**所有 SELL 先撮合、再撮合 BUY**(A 股「卖出资金当日可用」),同侧内保持策略提交顺序。资金检查用**滚动现金**而非撮合前快照。 +- **整手取整:** 仅 BUY 按 100 股整手向下取整;SELL 允许任意正整数股(含零股),静默截到整手违反契约。`order_target(symbol, 0)` 必须能清仓含零股的持仓。 +- **T+1 可卖不足:** 整单拒绝(`INSUFFICIENT_SHARES`),不支持部分截断。 +- **`realized_pnl` 不含费用**:仅 `(sell_price - avg_cost) × quantity`;费用(commission / stamp_tax / other_fee)只走现金账。 +- **金额量化:** 全部 `ROUND_HALF_EVEN`(`quantize_cash` 0.01 元 / `quantize_price` 0.0001 元),与券商「四舍五入到分」存在 1 分级差异。 +- **`Fill.BUY` 不携带 `stamp_tax`:** 印花税仅在 SELL 收取;`Fill.__post_init__` 拒绝 BUY 携带非零 stamp_tax 以保证账本与 costs 表一致。 +- **`intents.target_quantity_for_value(0)` 返回 0**(flatten),与 docstring 一致。 +- **CLI `initial_cash` 拒绝 float**(TOML 字面 `100000.0` 报错),与 `BacktestConfig` 严格度对齐(contract rule 5)。 + +完整口径见 `docs/design/mvp-contract.md` §3.4。手算回归测试在 `tests/engine/test_task16_matching.py`。 + +### 净值与绩效指标口径(任务 17) + +- **首日 P&L 进入曲线:** `daily_return[0] = total_equity[0] / initial_cash - 1`、`drawdown[0] = (initial_cash - total_equity[0]) / initial_cash`,不再硬编码 0;**后续日回撤 running peak = `max(initial_cash, 历史 total_equity)`**(峰值序列以 `initial_cash` 为初始峰值);满足恒等式 `∏(1 + daily_return) = 1 + total_return`(Decimal 精度内)。 +- **波动率样本不足返回 `None`:** `< 2` 个日收益时 `daily_volatility` / `annualized_volatility` / `sharpe_ratio` 返回 `None` 并附 note,禁止错报 0。真正 0 波动率才返回 `Decimal('0')`。 +- **幂运算桥接:** `(1 + total_return) ** (n / 252)` 通过 `Decimal(str(float(...)))` 重建为 Decimal,避免 `Decimal(float(...))` 直接继承二进制浮点。 +- **Decimal 量化:** 所有 `float` 桥接的 Decimal 输出统一 quantize 到 `Decimal('0.000000000001')`,保证 `summary.json` 干净。 +- **`positions.sellable_quantity` 口径:** **结转后**(D 行快照显示 D+1 起始时可卖数)。engine 在 `_snapshot_equity` 前调用 `settle_t1`,D 行的 `sellable_quantity` 已包含当日成交的滚动。 + +完整口径见 `docs/design/mvp-contract.md` §3.5。手算回归测试在 `tests/engine/test_task17_metrics.py`。 + +### 策略隔离与审计完整性(任务 18) + +- **`Order` 不可变:** `@dataclass(frozen=True)`,策略通过 `Context.pending_orders()` 拿到 Order 后无法修改任何字段(`quantity` / `avg_fill_price` / `fill_ids` 等)。`transition` / `record_fill` 用 `object.__setattr__` 绕过冻结,仅 engine / broker 可调用。 +- **`DataView.portal` 私有:** 字段名 `_portal`,策略无法通过 `view.portal.get_bars(sym, future_date)` 绕过 `visible_through`;所有数据访问走 `view.history` / `view.current_price` / `view.universe`。 +- **Universe 生效:** `set_universe([...])` 后对未声明符号下单立即拒绝(`RejectReason.OUT_OF_UNIVERSE`,含 ORDER_CREATED + ORDER_REJECTED 事件,Order 不经过 broker、停在 `_out_of_universe_orders` 并在 result 构造时折入 `orders_table`);未设 universe 时不限制。 +- **历史股票池:** `Context.historical_universe()` 返回 `visible_through` 当日的 portal 股票池(默认排除 `.BJ`),受可见性约束,不暴露 raw portal。 +- **返回值防御性:** `pending_orders()` / `universe()` / `historical_universe()` 均返回 list 副本;Bar / Factor 跨查询复用(任务 15)。 + +完整口径见 `docs/design/mvp-contract.md` §3.6。回归测试在 `tests/engine/test_task18_isolation.py`。 + +### 因子诊断与分红偏差显性化(任务 19) + +> ⚠️ **`adjustment_policy=none` 下,跨除权日的净值系统性低估(少分红现金)。** + +`v0.1` 严格不实现分红会计(契约任务 9)。任务 19 让偏差**可见**: + +- 引擎对**当前持仓**标的的因子跳变(|Δfactor/factor| > **0.1%**)自动生成 `DATA_WARNING` 事件、`FactorDiagnostic` 记录;写入 `result.factor_diagnostics`、`summary.json`、`events.jsonl`。清仓后(持仓归零)停止告警,持有期结束。 +- **账本零影响**:诊断是只读观测,cash / position / equity 与无诊断时 byte-identical(`test_diagnostics_do_not_change_ledger` 锁定)。 +- CLI 末尾打印一行汇总:`warning: N corporate-action factor jumps detected during holding periods; NAV excludes dividends (adjustment_policy=none), see summary.json`。 +- 跨除权日长区间结果不可用于收益评估,必须先评估 factor 跳变并改用 `factor_total_return`(v0.1.1+ 后续任务);README 显著位置保留此声明并链接诊断输出。 + +完整口径见 `docs/design/mvp-contract.md` §3.7。复刻 600000.SH 2026-07-16 除权案例的回归测试在 `tests/engine/test_task19_factor_diagnostics.py`。 + ## Python 用法(已实现) -> 下面的代码就是 `examples/moving_average.py` 的简化版,可直接 `python -m hqbacktest run` 跑通(见 §命令行)。完整示例见 `examples/`。 +> 下面的代码就是 `examples/moving_average.py` 的简化版,演示了 `Context` 与 `DataView` 的基本交互。**端到端运行请用 CLI(见 §命令行)或直接 `python -m hqbacktest run`**;完整示例见 `examples/`。 ```python from decimal import Decimal @@ -202,6 +276,8 @@ hqbacktest run --config configs/moving_average.toml --output results/moving-aver python -m hqbacktest run --config configs/moving_average.toml --output results/moving-average ``` +**策略模块解析(任务 20):** `hqbacktest run` 会把 config 文件所在目录和当前工作目录加入 `sys.path`,让 `[strategy].module = "my_strategy"` 这类不带点号的写法能直接 import 成功(与 `python -m hqbacktest run` 行为一致)。 + `--output` 可选:省略时使用配置中 `[output].directory`;提供时覆盖该值(例如 CI 里把结果重定向到临时目录)。 ### 配置 schema @@ -264,13 +340,19 @@ results/run-1/ | 必填字段缺失 | 2 | `[start] missing required key 'start_date'` | | 未知 section / key | 2 | `unknown config sections: ['extra']; allowed: [...]` | | 日期格式错 | 2 | `[start].start_date: must be 8 digits` | +| `initial_cash = nan` / `inf` | 2 | `[capital].initial_cash=NaN must be a finite number ...` | +| `initial_cash = float` | 2 | `[capital].initial_cash must be int/str/Decimal; float is forbidden ...` | +| 空交易窗口 | 2 | `no trading days in [...] for source 'memory'; ...` | | 策略模块无法导入 | 2 | `could not import strategy module 'examples.foo': ...` | | 策略类非 BaseStrategy 子类 | 2 | `MyStrategy is not a BaseStrategy subclass` | | 策略无 class_name 且模块无 BaseStrategy | 2 | `no BaseStrategy subclass found in ...` | | 输出目录不可创建 / 不可写 | 3 | `cannot create output directory ...: ...` | +| 输出目录已含旧结果文件 | 3 | `output directory ... already contains prior-run files; pass force=True ...` | | 引擎异常(非 RunFailed) | 4 | `backtest run failed: ...` | | 成功 | 0 | stdout: `hqbacktest: wrote results to results/run-1` | +CLI 同时支持 `--force` 覆盖已有输出目录:`hqbacktest run --config FILE --output DIR --force`。 + ### 复现性 两次相同输入 + 相同数据 + 相同 `data_root` 的运行: diff --git a/docs/design/mvp-contract.md b/docs/design/mvp-contract.md index 66d700c..fa2de63 100644 --- a/docs/design/mvp-contract.md +++ b/docs/design/mvp-contract.md @@ -138,6 +138,81 @@ - `data portal` 通过 `data_root` 与 `source` 解析数据集根目录。v0.1 的固定布局为 `{root}/{source}/calendar.csv`,以及 `stock_list/{YYYYMMDD}.csv`、`stock_daily/{YYYYMMDD}.csv`、`stock_factor/{YYYYMMDD}.csv`;任何缺失、不可读或格式不符的文件必须报错,不得联网回补。 - hqdata CSV 快照是叶子数据边界;更新数据只能在回测运行前通过 hqdata CLI 完成。 +### 3.3 数据可见性与缺行语义(任务 14 固化) + +| 维度 | v0.1 默认决定 | +| --- | --- | +| `get_bars(symbol, start, end)` | 返回窗口内实际存在的行,**允许逐日间隙**;窗口内无任何行返回 `[]` 而非报错。 | +| `get_bar(symbol, date)` 失败分类 | 个股当日缺行(停牌 / 未上市 / 已退市) → 返回「无当日行情」的空结果;**整日快照文件缺失** → `SnapshotFileMissingError`(`MissingDataError` 子类),引擎不得当作「该股无价」处理。 | +| `current_price(symbol)` | 返回截至 `visible_through` 的最近一个有效收盘价,**回看上限 20 个交易日**,超出返回 `None`;停牌持仓按最近收盘估值并记录 `DATA_WARNING` 事件,禁止静默按 0 计入。 | +| 首个交易日盘前(`visible_through="00000000"`) | `history` 返回 `[]`、`current_price` 返回 `None`,**不抛异常**。 | +| `get_universe(date)` | **按精确日期查询**,不做向前回退;`.BJ`(北交所)股票默认过滤,可通过 `include_bj=True` 保留。 | +| `Bar.volume` 单位 | **手**(1 手 = 100 股;与 Tushare `hqdata` 适配器口径一致);调用方需要股数时应乘以 `LOT_SIZE`。 | +| 双门户一致性 | `InMemoryDataPortal` 与 `HqDataCsvPortal` 行为完全一致(parity 测试覆盖)。 | + +### 3.4 撮合与账本口径(任务 16 固化) + +| 维度 | v0.1 默认决定 | +| --- | --- | +| 同日撮合顺序 | 单个 `OPEN_MATCH(today)` batch 内**所有 SELL 先撮合、再撮合 BUY**(A 股「卖出资金当日可用」),同侧内保持策略提交顺序;`broker.match` 按此顺序返回结果以保证 `running_cash` 滚动检查生效。 | +| 资金检查 | `InsufficientCashRule` 检查**滚动现金**(running_cash):SELL 净回款增加、BUY 成本扣除,拒绝后不滚动。 | +| 整手取整 | **仅 BUY** 按 100 股整手向下取整;SELL 允许任意正整数股(含零股),静默截到整手违反契约。 | +| `order_target(symbol, 0)` | 必须能清仓含零股的持仓;`order(sym, -N)` 对任意正整数 N 提交原数量。 | +| T+1 可卖不足 | **整单拒绝**(`RejectReason.INSUFFICIENT_SHARES`),不支持部分截断。 | +| `realized_pnl` | **不含任何费用**(commission / stamp_tax / other_fee),仅 `(sell_price - avg_cost) × quantity`;费用只走现金账。 | +| 金额量化 | 全部使用 `ROUND_HALF_EVEN`(`quantize_cash` 0.01 元 / `quantize_price` 0.0001 元),与券商「四舍五入到分」存在 1 分级差异。 | +| 同日同价「先买后卖」与「先卖后买」 | 提交顺序决定成交时点;`realized_pnl` 在费用外对相同 `(price, avg_cost, quantity)` 相同,**不含费用的现金额**依提交顺序而不同。 | +| BUY 成交 `stamp_tax` | 必须为 0(印花税仅在 SELL 收取),`Fill.__post_init__` 校验拒绝非零。 | +| `target_quantity_for_value(0)` | 返回 `0`(flatten),与 docstring 一致。 | +| CLI `initial_cash` | 拒绝 `float`(TOML 字面 `100000.0` 报错),与引擎层 `BacktestConfig` 严格度对齐。 | + +### 3.5 净值与绩效指标口径(任务 17 固化) + +| 维度 | v0.1 默认决定 | +| --- | --- | +| 首日 `daily_return` | `total_equity[0] / initial_cash - 1`(**不再硬编码 0**),首日 P&L 进入收益序列。 | +| 首日 `drawdown` | `(initial_cash - total_equity[0]) / initial_cash`(首日下跌时为正;不再硬编码 0),后续日 running peak = `max(initial_cash, 历史 total_equity)`,首日跌幅进入回撤峰值序列。 | +| 恒等式 | `∏(1 + daily_return) == 1 + total_return`(Decimal 精度内)。 | +| 波动率样本不足 | `< 2` 个日收益时 `daily_volatility` / `annualized_volatility` / `sharpe_ratio` 返回 `None` + note(**禁止错报 0**);真正 0 波动率才返回 `Decimal('0')`。 | +| `annualized_return` 幂运算 | `float(growth) ** float(exponent)` 通过 `Decimal(str(...))` 重建为 Decimal,**禁止 `Decimal(float(...))`**。 | +| Decimal 量化 | 所有 `float` 桥接的 Decimal 输出统一 quantize 到 `Decimal('0.000000000001')`,保证 `summary.json` 干净。 | +| `positions.sellable_quantity` | **结转后**(D 行快照为 D+1 起始时可卖数);engine 在 `_snapshot_equity` 前调用 `settle_t1`,D 行的 `sellable_quantity` 即已包含当日成交的滚动。 | + +### 3.6 策略隔离与审计完整性(任务 18 固化) + +| 维度 | v0.1 默认决定 | +| --- | --- | +| `Order` 不可变 | `@dataclass(frozen=True)`,`fill_ids: tuple[str, ...]`;策略收到 `pending_orders()` 后无法修改任何字段(quantity / avg_fill_price / fill_ids 等);`transition` / `record_fill` 用 `object.__setattr__` 绕过冻结,仅 engine / broker 可调用。 | +| `DataView.portal` 私有 | 字段名 `_portal`(私有),**策略无法**通过 `view.portal.get_bars(sym, future_date)` 绕过 `visible_through`;所有数据访问走 `view.history` / `view.current_price` / `view.universe`。 | +| Universe 生效 | `set_universe(...)` 后,对未声明的符号下单立即拒绝(`RejectReason.OUT_OF_UNIVERSE`,含 ORDER_CREATED + ORDER_REJECTED 事件,Order 不经过 broker、停留在 `_out_of_universe_orders` 并在 result 构造时折入 `orders_table`);**未设 universe 时不限制**。 | +| 历史股票池 | `Context.historical_universe()` 返回 `visible_through` 当日的 portal 股票池(默认排除 `.BJ`),受可见性约束;不暴露 raw portal。 | +| 返回值防御性 | `pending_orders()` / `universe()` / `historical_universe()` 均返回 list 副本;Bar / Factor 跨查询复用(任务 15)。 | + +### 3.7 因子诊断接入与分红偏差显性化(任务 19 固化) + +| 维度 | v0.1 默认决定 | +| --- | --- | +| 启用条件 | `adjustment_policy=none`(v0.1 唯一接受值);引擎对**当前持仓**(quantity > 0)标的的因子跳变自动诊断,清仓后停止(持有期结束)。 | +| 跳变阈值 | **0.1%**(holdings-period 阈值,`jump_band=(0.999, 1.001)`);`analyze_factor_series` 默认 `(0.5, 2.0)` 用于一般因子质量诊断,holdings-period 用更严阈值。 | +| 数据来源 | `portal.get_factor(symbol, today, today)`;快照缺失 / 无因子静默跳过(不影响 run),仅在**当前持仓**且超出阈值时产出警告。 | +| 写入位置 | `EngineEvent` 阶段 `DATA_WARNING`(detail 包含前后因子值与 ratio);`FactorDiagnosticCollector` → `BacktestResult.factor_diagnostics` → `summary.json` 的 `factor_diagnostics` 字段。 | +| 账本影响 | **零**:诊断是只读观测,绝不修改 cash / position / equity;baseline 与诊断版的逐字节相同(byte-identical ledger)。 | +| CLI 警告 | `run_from_config` 末尾若 `result.factor_diagnostics` 非空,打印一行 `warning: N corporate-action factor jumps detected during holding periods; NAV excludes dividends (adjustment_policy=none), see summary.json`。 | +| 文档承诺 | README 显著位置明示:`adjustment_policy=none` 下跨除权日的净值**系统性低估**(少分红现金),长区间结果不可用于收益评估,并链接因子诊断输出(`summary.json` / `events.jsonl`)。 | + +### 3.8 CLI 易用性与契约承诺一致性(任务 20 固化) + +| 维度 | v0.1 默认决定 | +| --- | --- | +| 策略模块解析 | `hqbacktest run` 把 config 文件所在目录和当前工作目录加入 `sys.path`(与 `python -m hqbacktest run` 行为一致),`[strategy].module = "my_strategy"` 这类不带点号的写法直接可 import。 | +| `initial_cash` 校验 | 拒绝 `nan` / `+inf` / `-inf` / `float`;接受 `int` / `str` / `Decimal`;错误为单行 `ConfigError`(CLI exit 2)。 | +| 日期校验 | `YYYYMMDD` 格式 + 真实日历;`20241399` 等非法日期立即拒。 | +| 空交易窗口 | `[start, end]` 区间在 portal 日历上无交易日 → `ConfigurationError`,不得静默成功写出空结果。 | +| 输出目录复用 | 已存在且非空的输出目录默认拒绝(CLI exit 3);`--force` 可覆盖。 | +| `order_value` 等下单函数 | 接受 `int` / `str` 金额(继续拒绝 `float`),降低策略样板代码。 | +| `git_commit` 语义 | `run_metadata.json` 记录 **hqbacktest 自身** 的 git commit(engine 来源),不再记录用户 cwd 仓库 commit。 | +| 文档一致性 | README「命令行」「错误信息」章节逐条与实现对齐;包布局含 `cli/` 子包;CLI 错误示例覆盖 exit 2/3/4 各档。 | + ## 6. 不可变规则 以下 13 条规则在 v0.1 期间**不允许任何代码绕过**;新增能力时若必须突破某条,必须先在本文档登记例外并同步 README。 @@ -186,4 +261,10 @@ | 2026-08-23 | v0.1 | 任务 9:公司行为扩展设计门槛落地——`BacktestConfig.adjustment_policy` 严格只接受 `"none"`;`CorporateActionProvider` 列为设计草案并锁定 10 个权威字段;`BacktestResult.adjustment_policy` 与 `factor_diagnostics` 字段已就位;因子诊断接口存在但 v0.1 不启用 | hqbacktest 维护者 | | 2026-08-17 | v0.1(已被后续修订取代) | 曾将 `source` 交给 `hqdata` 解析;该 API 驱动的数据边界已在 2026-08-23 被 CSV 快照契约取代 | hqbacktest 维护者 | | 2026-08-23 | v0.1 | 修正回测运行时数据边界:`hqbacktest` 直接只读 hqdata CLI 落盘 CSV;`data_root` 默认 `~/.hqdata`,不调用 `hqdata.api` 或网络数据源 | hqbacktest 维护者 | -| 2026-08-23 | v0.1 | 重构数据门户:`HqDataPortal` 替换为 `HqDataCsvPortal`,固定布局 `{data_root}/{source}/calendar.csv` + `stock_list|stock_daily|stock_factor/{YYYYMMDD}.csv`;`source` 名称或绝对路径均可,`CacheKey` 加入 `data_root` 防跨目录污染 | hqbacktest 维护者 | \ No newline at end of file +| 2026-08-23 | v0.1 | 重构数据门户:`HqDataPortal` 替换为 `HqDataCsvPortal`,固定布局 `{data_root}/{source}/calendar.csv` + `stock_list|stock_daily|stock_factor/{YYYYMMDD}.csv`;`source` 名称或绝对路径均可,`CacheKey` 加入 `data_root` 防跨目录污染 | hqbacktest 维护者 | +| 2026-08-24 | v0.1 | 任务 14 数据层缺行/停牌/首日语义:钉死 `get_bars` 允许间隙、引入 `SnapshotFileMissingError` 区分整日文件缺失与个股缺行、`current_price` 回看 20 交易日最近有效收盘价、首日哨兵日期不抛异常、删除 `InMemoryDataPortal.get_universe` 向前回退、补双门户 parity 测试、缓存返回防御性拷贝、`.BJ` 股票默认过滤、`Bar.volume` 单位标注为「手」 | hqbacktest 维护者 | +| 2026-08-24 | v0.1 | 任务 16 撮合与账本语义:同批撮合 SELL 先于 BUY(滚动现金)、SELL 不整手取整(`order_target(0)` 可清零股)、`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 维护者 | +| 2026-08-24 | v0.1 | 任务 17 净值与指标基准:首日 `daily_return` / `drawdown` 以 `initial_cash` 为基准(不再硬编码 0)、后续日 running peak = `max(initial_cash, 历史 total_equity)`、波动率样本不足返回 `None` 而非 0、`Decimal(str(float(...)))` 替代 `Decimal(float(...))` 幂运算桥接、`positions.sellable_quantity` 口径登记为「结转后」;恒等式 `∏(1 + daily_return) = 1 + total_return` 成立;登记 §3.5 | hqbacktest 维护者 | +| 2026-08-24 | v0.1 | 任务 18 策略隔离与审计完整性:`Order` 改为 `frozen=True`(策略无法篡改 `pending_orders()` 返回的 Order)、`DataView.portal` 改为私有 `_portal`、universe 生效(`RejectReason.OUT_OF_UNIVERSE`)、`Context.historical_universe()` 转发 `DataView.universe()` 受可见性约束;登记 §3.6 | hqbacktest 维护者 | +| 2026-08-24 | v0.1 | 任务 19 因子诊断接入与分红偏差显性化:engine 在持仓/成交标的的因子跳变(阈值 0.1%)自动生成 DATA_WARNING + `FactorDiagnostic`,结果写入 `summary.json` / `events.jsonl`;CLI 末尾打印汇总警告;账本与净值完全不变;登记 §3.7 | hqbacktest 维护者 | +| 2026-08-24 | v0.1 | 任务 20 CLI 易用性与文档真实性:console script 把 config dir + cwd 加入 sys.path(策略模块解析与 `python -m` 对齐);`initial_cash` 拒绝 nan/inf/float;空交易窗口、空输出目录、`--force` 覆盖;`order_value` 接受 int/str;`git_commit` 改为 hqbacktest 自身版本;README 错误码表与包布局对齐;登记 §3.8 | hqbacktest 维护者 | \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index b8ad48f..266d0d7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "hqbacktest" -version = "0.1.0" +version = "0.1.1" description = "A股量化策略回测与交易模拟引擎" readme = "README.md" authors = [ @@ -58,6 +58,10 @@ include = ["hqbacktest*"] [tool.pytest.ini_options] testpaths = ["tests"] pythonpath = ["."] +# Integration tests live under tests/integration/ and are skipped by +# default (no snapshot on the test machine). Run with +# `pytest tests/integration/ -v` on machines that have `~/.hqdata/tushare`. +addopts = "--ignore=tests/integration" markers = [ "integration: marks tests as integration tests against real hqdata sources (skipped by default)", ] diff --git a/src/hqbacktest/__init__.py b/src/hqbacktest/__init__.py index 153f179..4886b54 100644 --- a/src/hqbacktest/__init__.py +++ b/src/hqbacktest/__init__.py @@ -51,7 +51,7 @@ TradingRuleSet, ) -__version__ = "0.1.0" +__version__ = "0.1.1" __all__ = [ "__version__", diff --git a/src/hqbacktest/cli/__main__.py b/src/hqbacktest/cli/__main__.py index 2fd0dc8..e6b19e3 100644 --- a/src/hqbacktest/cli/__main__.py +++ b/src/hqbacktest/cli/__main__.py @@ -7,11 +7,13 @@ from __future__ import annotations import argparse +import importlib +import os import sys -from typing import List, Optional, Sequence +from typing import Optional, Sequence from .config import ConfigError -from .runner import RunResult, run_from_file +from .runner import RunResult, _prepare_sys_path, run_from_file def build_parser() -> argparse.ArgumentParser: @@ -42,6 +44,14 @@ def build_parser() -> argparse.ArgumentParser: "fills.csv, positions.csv, costs.csv and summary.json." ), ) + run.add_argument( + "--force", + action="store_true", + help=( + "Overwrite an output directory that already contains " + "prior-run files (task 20)." + ), + ) return parser @@ -59,8 +69,21 @@ def main(argv: Optional[Sequence[str]] = None) -> int: def _run(args: argparse.Namespace) -> int: + # Task 20: prepend the config file's directory and the current + # working directory to `sys.path` so the strategy module can be + # resolved by name alone, matching the documented "first-mile" + # workflow. This mirrors what `python -m` would do for an + # in-tree import and makes the console script behave the same + # way as `python -m hqbacktest run`. + _prepare_sys_path(args.config) + # Honor a test-only env hook so the CLI can be driven against an + # in-memory portal in subprocess tests without touching the real + # `~/.hqdata` snapshot. Task 20. + _maybe_load_test_bootstrap() try: - result: RunResult = run_from_file(args.config, output_dir=args.output) + result: RunResult = run_from_file( + args.config, output_dir=args.output, force=args.force + ) except ConfigError as exc: print(f"hqbacktest: {exc}", file=sys.stderr) return 2 @@ -76,4 +99,17 @@ def _run(args: argparse.Namespace) -> int: return 0 +def _maybe_load_test_bootstrap() -> None: + """If `HQBACKTEST_CLI_BOOTSTRAP` is set, import that module by name. + + Test-only hook used by `tests/cli/test_task20_cli.py` to swap the + portal builder in a subprocess without writing to the real + `~/.hqdata` snapshot. Production users never set this. + """ + name = os.environ.get("HQBACKTEST_CLI_BOOTSTRAP") + if not name: + return + importlib.import_module(name) + + __all__ = ["build_parser", "main"] diff --git a/src/hqbacktest/cli/config.py b/src/hqbacktest/cli/config.py index 37cadcf..8f3e19d 100644 --- a/src/hqbacktest/cli/config.py +++ b/src/hqbacktest/cli/config.py @@ -237,7 +237,16 @@ def _require_decimal( value = section[key] if isinstance(value, bool): raise ConfigError(f"[{section_name}].{key} must be a number, not bool") - if isinstance(value, (int, float, str)): + # Task 16: float is forbidden at the CLI layer too, matching the + # engine's contract rule 5. Without this check a TOML like + # `initial_cash = 100000.0` would silently convert to a Decimal via + # `Decimal(str(float))`, masking the precision concern. + if isinstance(value, float): + raise ConfigError( + f"[{section_name}].{key} must be int/str/Decimal; float is " + "forbidden (contract rule 5)" + ) + if isinstance(value, (int, str)): try: d = Decimal(str(value)) except Exception as exc: @@ -250,6 +259,16 @@ def _require_decimal( raise ConfigError( f"[{section_name}].{key} must be a number, got {type(value).__name__}" ) + # Task 20: NaN / +Inf / -Inf are technically parseable by + # `Decimal(str('nan'))` but break every downstream comparison + # (e.g. `nan < 0` raises InvalidOperation). Reject them here + # so the user gets a clean single-line ConfigError instead of + # a traceback. + if not d.is_finite(): + raise ConfigError( + f"[{section_name}].{key}={d} must be a finite number " + "(NaN / +Inf / -Inf are not allowed)" + ) if min_value is not None and d < min_value: raise ConfigError(f"[{section_name}].{key}={d} must be >= {min_value}") return d diff --git a/src/hqbacktest/cli/runner.py b/src/hqbacktest/cli/runner.py index 9951377..52d6735 100644 --- a/src/hqbacktest/cli/runner.py +++ b/src/hqbacktest/cli/runner.py @@ -48,17 +48,25 @@ def run_from_file( - config_path: str, output_dir: Optional[str] = None + config_path: str, output_dir: Optional[str] = None, force: bool = False ) -> "RunResult": # type: ignore[name-defined] """Load the config, build the engine, run, and write the output dir. `output_dir` (from the CLI `--output` flag) overrides the config's `[output].directory` when given. Returns a `RunResult` with the exit code (0 on success) and the path of the output directory. The CLI - converts this to a process exit status. + converts this to a process exit status. `force=True` lets the + runner overwrite a non-empty output directory (task 20). + + Task 20: prepend the config's directory and cwd to `sys.path` so + the user-supplied strategy module can be imported by name alone + (the documented first-mile workflow). """ + _prepare_sys_path(config_path) config_file = load_config_file(config_path) - return run_from_config(config_file, source_path=config_path, output_dir=output_dir) + return run_from_config( + config_file, source_path=config_path, output_dir=output_dir, force=force + ) def run_from_config( @@ -66,8 +74,16 @@ def run_from_config( *, source_path: Optional[str] = None, output_dir: Optional[str] = None, + force: bool = False, ) -> "RunResult": # type: ignore[name-defined] - """Build the engine from a validated `ConfigFile` and run the backtest.""" + """Build the engine from a validated `ConfigFile` and run the backtest. + + `force=True` lets the run overwrite an output directory that + already contains prior-run files (task 20). Without `force`, + mixing a fresh run with stale CSVs / summary.json from a + previous run is rejected with exit code 3 to keep the audit + trail honest. + """ effective_output = output_dir or config_file.output_directory try: strategy = resolve_strategy(config_file) @@ -76,6 +92,30 @@ def run_from_config( backtest_config = build_backtest_config(config_file) portal = _resolve_portal(backtest_config.source, backtest_config.data_root) output_path = Path(effective_output) + if output_path.exists() and not output_path.is_dir(): + # The configured output path is an existing FILE, not a + # directory. Refuse rather than silently failing later. + return RunResult( + exit_code=3, + output_dir=None, + message=( + f"output path {output_path!s} is not a directory; " + f"remove the file or change [output].directory" + ), + ) + if output_path.exists() and not force: + # Reject if the directory already holds files from a prior run; + # an empty directory is allowed (first run). + if any(output_path.iterdir()): + return RunResult( + exit_code=3, + output_dir=output_path, + message=( + f"output directory {output_path!s} already contains " + f"prior-run files; pass force=True to overwrite or " + f"choose a fresh directory" + ), + ) try: output_path.mkdir(parents=True, exist_ok=True) except OSError as exc: @@ -92,10 +132,15 @@ def run_from_config( ) from ..engine.engine import BacktestEngine # local: avoid circular + from ..engine.errors import ConfigurationError try: engine = BacktestEngine(backtest_config, strategy=strategy, portal=portal) result = engine.run() + except ConfigurationError as exc: + # A configuration error (e.g. an empty trading-day window) is a + # user-input problem, not a run failure: exit 2, single line. + return RunResult(exit_code=2, output_dir=output_path, message=str(exc)) except Exception as exc: # We do NOT swallow this as a normal RunFailed; the CLI surfaces # the message verbatim. @@ -112,6 +157,19 @@ def run_from_config( ) result.save(str(output_path)) + # Task 19: warn the operator when factor diagnostics surfaced a + # holding-period jump. The NAV excludes dividends (policy="none"), + # so this is the only place a human sees the bias. The CLI is the + # operator's terminal; we print one summary line. + holdings_diag = result.factor_diagnostics or [] + if holdings_diag: + print( + f"warning: {len(holdings_diag)} corporate-action factor " + f"jumps detected during holding periods; NAV excludes " + f"dividends (adjustment_policy=none), see summary.json", + flush=True, + ) + return RunResult(exit_code=0, output_dir=output_path, message="") @@ -194,22 +252,62 @@ def _write_run_metadata( def _git_commit() -> Optional[str]: - """Return the current short git commit, or `None` if unavailable. - - We never raise from here; failure to read git is a no-op. + """Return the hqbacktest package's own short git commit, or `None` + if unavailable. + + Task 20: the commit recorded here is the commit of the + hqbacktest repo, NOT the user's cwd repository. A user + running the CLI from inside their own strategy repo gets + hqbacktest's commit (the engine that produced the result); + if they want their strategy's commit too they can record it + themselves in the config. We never raise from here; failure + to read git is a no-op. """ try: - out = subprocess.check_output( - ["git", "rev-parse", "--short", "HEAD"], - stderr=subprocess.DEVNULL, - timeout=2, - ) - text = out.decode("utf-8").strip() - return text or None + import hqbacktest as _hq_pkg + + pkg_dir = Path(_hq_pkg.__file__).resolve().parent + # Walk up to find the directory that contains `.git`. + for parent in [pkg_dir, *pkg_dir.parents]: + if (parent / ".git").exists(): + out = subprocess.check_output( + ["git", "-C", str(parent), "rev-parse", "--short", "HEAD"], + stderr=subprocess.DEVNULL, + timeout=2, + ) + text = out.decode("utf-8").strip() + return text or None + return None except Exception: return None +def _prepare_sys_path(config_path: str) -> None: + """Add the config file's directory and cwd to `sys.path`. + + Task 20: lets the user-supplied strategy module be imported by + name alone (e.g. `module = 'strategy'`) without having to add + `sys.path` boilerplate in the strategy file. Mirrors what + `python -m hqbacktest run` would do for in-tree imports. + """ + candidates: List[str] = [] + try: + cfg_dir = str(Path(config_path).resolve().parent) + if cfg_dir and cfg_dir not in candidates: + candidates.append(cfg_dir) + except OSError: + pass + try: + cwd = str(Path.cwd()) + if cwd and cwd not in candidates: + candidates.append(cwd) + except OSError: + pass + for entry in candidates: + if entry and entry not in sys.path: + sys.path.insert(0, entry) + + __all__ = [ "RunResult", "run_from_config", diff --git a/src/hqbacktest/data/__init__.py b/src/hqbacktest/data/__init__.py index 11327ec..de1d044 100644 --- a/src/hqbacktest/data/__init__.py +++ b/src/hqbacktest/data/__init__.py @@ -12,6 +12,7 @@ FutureDataAccessError, InvalidDataError, MissingDataError, + SnapshotFileMissingError, UnknownSymbolError, ) from .hqdata_portal import ( @@ -42,6 +43,7 @@ "InvalidDataError", "MarketDataPortal", "MissingDataError", + "SnapshotFileMissingError", "UnknownSymbolError", "assert_unique_sorted", "require_columns", diff --git a/src/hqbacktest/data/data_view.py b/src/hqbacktest/data/data_view.py index a327732..b44c01e 100644 --- a/src/hqbacktest/data/data_view.py +++ b/src/hqbacktest/data/data_view.py @@ -5,21 +5,46 @@ what the strategy can see) and exposes a small ergonomic API: view.history(symbol, field, bar_count) -> list of values - view.current_price(symbol) -> latest close <= visible_through + view.current_price(symbol) -> latest close within the lookback window view.universe() -> list[str] (as of visible_through) + +Task 14 semantics (per `docs/design/mvp-contract.md` and `TODO.md`): + - `history(bar_count=N)` only queries the relevant slice (the lookback + window); it does NOT scan the full pre-start window of every symbol. + - `current_price(symbol)` returns the most recent valid close within the + last `CURRENT_PRICE_LOOKBACK` (20) trading days. Suspended / delisted / + pre-IPO symbols therefore return a usable price instead of an empty + series. + - The sentinel `visible_through="00000000"` used on the very first + trading day must NOT raise: `history` returns `[]` and + `current_price` returns `None`. """ from dataclasses import dataclass from decimal import Decimal from typing import List, Optional -from .errors import FutureDataAccessError, MissingDataError +from .errors import ( + FutureDataAccessError, + MissingDataError, + SnapshotFileMissingError, +) from .portal import MarketDataPortal from .validators import validate_symbol, validate_yyyymmdd VALID_FIELDS = ("open", "high", "low", "close", "volume") DEFAULT_HISTORY_START = "19000101" +# Cap on how far back `current_price` walks to find the most recent close. +# Calibrated for A股: with a normal trading week of ~5 trading days, 20 +# covers roughly a month of holidays plus one typical multi-day suspension. +CURRENT_PRICE_LOOKBACK = 20 + +# Sentinel value used by the scheduler when no prior trading day exists yet. +# Any explicit "00000000" request must be treated as "no data visible" +# rather than as a literal date lookup. +NO_HISTORY_SENTINEL = "00000000" + @dataclass class DataView: @@ -28,46 +53,108 @@ class DataView: `universe_start` is the earliest date the strategy is allowed to query. Together with `visible_through`, it forms the half-open window `[universe_start, visible_through]` that all reads are restricted to. + + `visible_through="00000000"` is a legal sentinel that exposes no data: + `history(...)` returns `[]` and `current_price(...)` returns `None`. + + Task 18: the `portal` attribute is **private** (name-mangled to + `_portal`). Strategies cannot reach the raw `MarketDataPortal` and + bypass `visible_through` via `view.portal.get_bars(...)`. All + data-layer access goes through the guarded methods on this view. + The constructor still accepts `portal=...` (kwarg) so existing + call sites don't break, but the value is stored only on the + private field and is never re-exposed. """ - portal: MarketDataPortal + _portal: MarketDataPortal visible_through: str universe_start: Optional[str] = None + def __init__( + self, + portal: MarketDataPortal, + visible_through: str, + universe_start: Optional[str] = None, + ) -> None: + # Accept `portal=` for backward compatibility, but store it on + # the private `_portal` field. Strategies that try to read + # `view.portal` get `AttributeError` (task 18 isolation). + self._portal = portal + self.visible_through = visible_through + self.universe_start = universe_start + self.__post_init__() + def __post_init__(self) -> None: + # The sentinel "00000000" is allowed as a special value. + if self.visible_through == NO_HISTORY_SENTINEL: + if ( + self.universe_start is not None + and self.universe_start != NO_HISTORY_SENTINEL + ): + raise FutureDataAccessError(self.universe_start, self.visible_through) + return validate_yyyymmdd(self.visible_through, name="visible_through") if self.universe_start is not None: + if self.universe_start == NO_HISTORY_SENTINEL: + raise FutureDataAccessError(self.universe_start, self.visible_through) validate_yyyymmdd(self.universe_start, name="universe_start") if self.universe_start > self.visible_through: raise FutureDataAccessError(self.universe_start, self.visible_through) def _guard(self, requested: str) -> None: - if requested > self.visible_through: + """Reject queries that escape the visibility window. + + `00000000` is always considered to lie outside the visible window + because the portal never indexes dates earlier than its first + snapshot day; passing it through would force every underlying call + to scan the full pre-start history. + """ + if requested == NO_HISTORY_SENTINEL: raise FutureDataAccessError(requested, self.visible_through) - if self.universe_start is not None and requested < self.universe_start: + if requested > self.visible_through: raise FutureDataAccessError(requested, self.visible_through) + if ( + self.universe_start is not None + and self.universe_start != NO_HISTORY_SENTINEL + and requested < self.universe_start + ): + # Reading before the data start is NOT future-data access; the + # window simply predates what the strategy is allowed to see. + raise MissingDataError( + "requested window starts before universe_start", + f"{requested} < {self.universe_start}", + ) # ------------------------------------------------------------------ # # Pass-through (still guarded) # ------------------------------------------------------------------ # - def get_bars(self, symbol: str, start: str, end: str): + def get_bars(self, symbol: str, start: str, end: str) -> List: + # Sentinel view: empty by construction. + if self.visible_through == NO_HISTORY_SENTINEL: + return [] self._guard(start) self._guard(end) - return self.portal.get_bars(symbol, start, end) + return self._portal.get_bars(symbol, start, end) def get_factor(self, symbol: str, start: str, end: str): + if self.visible_through == NO_HISTORY_SENTINEL: + return [] self._guard(start) self._guard(end) - return self.portal.get_factor(symbol, start, end) + return self._portal.get_factor(symbol, start, end) - def get_universe(self, date: Optional[str] = None) -> List[str]: + def get_universe( + self, date: Optional[str] = None, include_bj: bool = False + ) -> List[str]: target = self.visible_through if date is None else date + if self.visible_through == NO_HISTORY_SENTINEL: + return [] self._guard(target) - return self.portal.get_universe(target) + return self._portal.get_universe(target, include_bj=include_bj) def universe(self) -> List[str]: - """Historical universe as of `visible_through`.""" + """Historical universe as of `visible_through` (excluding .BJ).""" return self.get_universe() # ------------------------------------------------------------------ # @@ -84,34 +171,126 @@ def history( The result is ordered ascending by date. If `bar_count` exceeds the number of available bars, the shorter list is returned. + + Per task 14: when `universe_start` is set, only bars within the + `[universe_start, visible_through]` window are queried — never the + full pre-start history. The first-trading-day sentinel returns + `[]` rather than raising. """ validate_symbol(symbol) if field not in VALID_FIELDS: raise ValueError(f"field must be one of {VALID_FIELDS}, got {field!r}") if bar_count <= 0: raise ValueError(f"bar_count must be positive, got {bar_count}") - start = ( - self.universe_start - if self.universe_start is not None - else DEFAULT_HISTORY_START - ) - bars = self.portal.get_bars(symbol, start, self.visible_through) + if self.visible_through == NO_HISTORY_SENTINEL: + return [] + # Cap the lookback to bar_count trading days so we never ask the + # portal for the full pre-start history of every symbol. The cap + # is a small constant; bar_count itself is what determines the + # requested number of values. + if ( + self.universe_start is not None + and self.universe_start != NO_HISTORY_SENTINEL + ): + start = self.universe_start + else: + start = self._resolve_history_start(bar_count) + try: + bars = self._portal.get_bars(symbol, start, self.visible_through) + except SnapshotFileMissingError: + # A missing whole-day snapshot is an infrastructure failure, not + # a per-symbol gap. It must propagate so the engine aborts the + # run instead of silently producing an empty history. + raise + except FutureDataAccessError: + # Defensive: a portal that re-checks visible_through would + # raise here. Empty is the truthful answer. + return [] + except MissingDataError: + return [] if len(bars) > bar_count: bars = bars[-bar_count:] return [getattr(b, field) for b in bars] def current_price(self, symbol: str) -> Optional[Decimal]: - """Return the close price on or before `visible_through`, or None.""" + """Return the most recent valid close on or before `visible_through`. + + Task 14 semantics: walk back up to `CURRENT_PRICE_LOOKBACK` (20) + trading days from `visible_through` and return the latest close in + that window. Returns `None` when no bar exists in the lookback + window (e.g. pre-IPO or first-trading-day sentinel). Suspended + symbols therefore keep their last traded price for valuation. + + `InvalidDataError` propagates: a corrupt row is an infrastructure + failure that must not be silently folded into "no price". + `SnapshotFileMissingError` (a whole-day snapshot file missing on + disk) is an infrastructure failure and propagates so the run aborts; + `MissingDataError` (suspended / delisted / pre-IPO) is absorbed. + + Task 15 performance: one `get_calendar` (cached) resolves the + 20-trading-day cutoff, then a single `get_bars(cutoff, end)` is + dispatched. This avoids the old N `get_bars(day, day)` round-trips + while preserving the 20-**trading-day** lookback bound. + """ validate_symbol(symbol) - start = ( - self.universe_start - if self.universe_start is not None - else DEFAULT_HISTORY_START + if self.visible_through == NO_HISTORY_SENTINEL: + return None + # Resolve the 20-trading-day cutoff. The lookback is bounded by + # trading days, NOT by bar count: a symbol suspended for longer + # than the lookback must exhaust the window (return None), not + # reach further back to its pre-suspension close. + trading_days = self._portal.get_calendar( + _trading_day_lookback_start(self.visible_through), + self.visible_through, ) + if not trading_days: + return None + lookback = trading_days[-CURRENT_PRICE_LOOKBACK:] + cutoff = lookback[0] try: - bars = self.portal.get_bars(symbol, start, self.visible_through) + bars = self._portal.get_bars(symbol, cutoff, self.visible_through) + except SnapshotFileMissingError: + # Whole-day file gone: infrastructure failure, propagate so + # the engine aborts the run with DATA_ERROR. + raise except MissingDataError: return None - if not bars: - return None - return bars[-1].close + for bar in reversed(bars): + close = bar.close + if close is not None and close > 0: + return close + return None + + # ------------------------------------------------------------------ # + # Internal helpers + # ------------------------------------------------------------------ # + + def _resolve_history_start(self, bar_count: int) -> str: + """Compute a bounded start date for `history(bar_count=N)`. + + Task 14: do NOT scan the full pre-start history (19000101). A + 5-year lookback window is comfortably larger than any realistic + `bar_count` (most strategies look back < 250 trading days, roughly + one calendar year) and bounds the worst-case file scan. The portal + intersects the window with the trading calendar and filters + per-symbol gaps, so an over-wide window is cheap and never changes + the returned bar sequence. + + `bar_count` is accepted for API symmetry; a precise + calendar-position bound is deferred to task 15's cache design. + """ + return _trading_day_lookback_start(self.visible_through) + + +def _trading_day_lookback_start(visible_through: str) -> str: + """Return a YYYYMMDD string that comfortably covers the lookback window. + + The portal intersects with the calendar, so an overly wide window is + cheap. We pick a 5-year backstop (never earlier than + `DEFAULT_HISTORY_START`) to avoid generating dates that fall before the + snapshot started. + """ + year = int(visible_through) // 10000 + floor_year = int(DEFAULT_HISTORY_START[:4]) + start_year = max(year - 5, floor_year) + return f"{start_year}0101" diff --git a/src/hqbacktest/data/errors.py b/src/hqbacktest/data/errors.py index cd7d430..813e43d 100644 --- a/src/hqbacktest/data/errors.py +++ b/src/hqbacktest/data/errors.py @@ -24,7 +24,18 @@ def __init__(self, requested: str, visible_through: str) -> None: class MissingDataError(DataError): - """Requested data is not available (no rows, missing dates, etc.).""" + """Requested data is not available (no rows, missing dates, etc.). + + Use this for ordinary business-level absences: a symbol suspended on a + given trading day, an IPO that has not started yet, a stock that has + already delisted, etc. These are recoverable per-symbol outcomes and the + engine must NOT abort the run because of them. + + For infrastructure-level failures (the whole daily snapshot file is + missing on disk), raise `SnapshotFileMissingError` instead so the engine + can distinguish "this stock has no row today" from "we cannot read any + row at all today". + """ def __init__(self, what: str, detail: str = "") -> None: message = f"missing data: {what}" @@ -35,6 +46,24 @@ def __init__(self, what: str, detail: str = "") -> None: self.detail = detail +class SnapshotFileMissingError(MissingDataError): + """A whole-daily snapshot file is missing on disk (data infrastructure). + + Distinct from `MissingDataError` (a per-symbol gap such as a suspended or + delisted stock) so the engine and broker can refuse to silently fold + an infrastructure failure into a business rejection. The engine must + abort the run with a clear `DATA_ERROR` rather than treating the + missing file as "no quote available". + + `path` is the filesystem path that was expected; `what` describes the + snapshot family (e.g. `stock_daily`). + """ + + def __init__(self, what: str, path: str) -> None: + super().__init__(what, f"snapshot file missing: {path}") + self.path = path + + class InvalidDataError(DataError): """Data returned by the source violates hqbacktest invariants.""" diff --git a/src/hqbacktest/data/hqdata_portal.py b/src/hqbacktest/data/hqdata_portal.py index bf2b5ae..e2af446 100644 --- a/src/hqbacktest/data/hqdata_portal.py +++ b/src/hqbacktest/data/hqdata_portal.py @@ -1,6 +1,6 @@ """HqDataCsvPortal: production portal backed by hqdata CSV snapshots. -Rules enforced here (TODO task 4 contract): +Rules enforced here (TODO task 4 + task 15): - **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 @@ -13,15 +13,30 @@ 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 `MissingDataError` and never falls back to other dates. + snapshot raises `SnapshotFileMissingError` and never falls back to + other dates (task 14). - Cache keys include the normalized `data_root` so two portals pointing at different roots cannot share entries. + +Task 15 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. """ +from bisect import bisect_left, bisect_right from datetime import date as _date from decimal import Decimal from pathlib import Path -from typing import List, Optional, Tuple +from typing import Dict, List, Optional, Tuple import pandas as pd @@ -30,6 +45,7 @@ from .errors import ( InvalidDataError, MissingDataError, + SnapshotFileMissingError, UnknownSymbolError, ) from .portal import DataVersion, MarketDataPortal @@ -79,6 +95,14 @@ def __init__( self._source_label: str = source # the original string, for display self._root_path: Path = Path(self._data_root) / self._source_name self._cache = DataCache() + # Task 15: 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. + 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._data_version = DataVersion( @@ -109,9 +133,18 @@ def cache(self) -> DataCache: return self._cache def _resolve_as_of(self) -> str: + """Pick the snapshot's `as_of` from the calendar. + + `MissingDataError` (no calendar.csv at all) is treated as "snapshot + empty": fall back to today's date so the engine can still publish a + `data_version`. `InvalidDataError` (calendar exists but is corrupt) + is a real data infrastructure failure and MUST propagate; silently + masking it with `_date.today()` would let the engine believe the + snapshot is healthy when it is not. + """ try: calendar = self._read_calendar() - except (MissingDataError, InvalidDataError): + except MissingDataError: return _date.today().strftime("%Y%m%d") opens = [d for d, is_open in calendar if is_open == "Y"] if opens: @@ -166,11 +199,11 @@ def get_calendar(self, start: str, end: str) -> List[str]: ) cached = self._cache.get(cache_key) if cached is not None: - return cached + return list(cached) calendar = self._read_calendar() result = [d for d, flag in calendar if flag == "Y" and start <= d <= end] - self._cache.put(cache_key, result) - return result + self._cache.put(cache_key, list(result)) + return list(result) def is_trading_day(self, d: str) -> bool: validate_yyyymmdd(d) @@ -207,93 +240,164 @@ def next_trading_day(self, d: str) -> str: # Universe # ------------------------------------------------------------------ # - def get_universe(self, date: str) -> List[str]: + def get_universe(self, date: str, include_bj: bool = False) -> List[str]: + """Return the historical stock list as of `date`. + + Per task 14: `.BJ` (Beijing Stock Exchange) symbols are excluded by + default since v0.1 does not yet support BSE-specific trading rules + (no first-day limit-up/down, distinct trading calendar, etc.). + Pass `include_bj=True` to opt in. + """ validate_yyyymmdd(date) cache_key = CacheKey( self._data_root, self._source_name, "universe", "", "", date, date ) cached = self._cache.get(cache_key) if cached is not None: - return cached - path = self._root_path / "stock_list" / f"{date}.csv" - if not path.exists(): - raise MissingDataError( - "stock_list snapshot", - f"no snapshot for {date} at {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() - self._cache.put(cache_key, symbols) - return symbols + 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() + self._cache.put(cache_key, list(symbols)) + full = list(symbols) + if include_bj: + return full + return [sym for sym in full if not sym.endswith(".BJ")] # ------------------------------------------------------------------ # # Bars # ------------------------------------------------------------------ # 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 (task 14): + - 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`. + + Task 15 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. + """ validate_symbol(symbol) validate_yyyymmdd(start, name="start") validate_yyyymmdd(end, name="end") if start > end: raise InvalidDataError("window", f"start {start} > end {end}") - cache_key = CacheKey( - self._data_root, self._source_name, "bars", symbol, "", start, end - ) - cached = self._cache.get(cache_key) - if cached is not None: - return cached + # 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] + idx_start = bisect_left(dates, start) + idx_end = bisect_right(dates, end) + 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 (task 15). + """ + 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: - raise MissingDataError("calendar", f"no trading days in [{start}, {end}]") - bars: List[Bar] = [] + 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} for trading_day in calendar: - try: - day_bars = self._read_bars_for_day(symbol, trading_day) - except MissingDataError: - # No daily file for this symbol on this trading day: - # contract §3.1 "交易日覆盖范围" + task 4 validation. - raise MissingDataError( - "bars", - f"no bars for {symbol} on trading day {trading_day}", - ) - bars.extend(day_bars) - if not bars: - raise MissingDataError("bars", f"no bars for {symbol} in [{start}, {end}]") - # Universe membership + per-row symbol sanity. - for bar in bars: - if bar.symbol != symbol: - raise UnknownSymbolError( - f"CSV returned symbol {bar.symbol!r} when {symbol!r} was requested" - ) - # All bars must land on calendar trading days (already enforced above - # by iterating the calendar) and the requested symbol must be present - # for every trading day. - self._cache.put(cache_key, bars) - return bars - - def _read_bars_for_day(self, symbol: str, date: str) -> List[Bar]: + 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. + + 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. + """ + 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 + 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. + + Raises `SnapshotFileMissingError` if the file is missing on disk, + and `InvalidDataError` on corrupt data (missing columns, date + mismatch, duplicate rows, or Bar invariant violations). + """ path = self._root_path / "stock_daily" / f"{date}.csv" if not path.exists(): - raise MissingDataError( - "stock_daily snapshot", - f"no snapshot for {date} at {path}", - ) + raise SnapshotFileMissingError("stock_daily", str(path)) try: df = pd.read_csv( path, @@ -317,40 +421,43 @@ def _read_bars_for_day(self, symbol: str, date: str) -> List[Bar]: ["symbol", "date", "open", "high", "low", "close", "volume"], name="stock_daily", ) - # Filename date and CSV date column must agree. - file_dates = {validate_yyyymmdd(v) for v in df["date"].tolist()} + 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}", ) - # Filter to the requested symbol; all remaining rows must match. - rows = df[df["symbol"] == symbol] - if len(rows) > 1: - raise InvalidDataError( - "stock_daily", - f"expected at most one row for {symbol!r} on {date}, got {len(rows)}", - ) - out: List[Bar] = [] - for _, row in rows.iterrows(): - # Read price text directly from CSV before creating Decimal values. - out.append( - Bar.from_raw( - symbol=str(row["symbol"]), - date=str(row["date"]), - open=str(row["open"]), - high=str(row["high"]), - low=str(row["low"]), - close=str(row["close"]), - volume=int(row["volume"]), + # Build a {symbol: Bar} map for the day. `itertuples` is roughly + # 25x faster than `iterrows` on real-data snapshots (task 15 + # benchmark). Duplicate symbol rows are still rejected as a + # data error. + 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}", ) - ) - if not out: - raise MissingDataError( - "bars", - f"symbol {symbol!r} missing from {path}", - ) - return out + try: + result[sym] = 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"), + volume=int(getattr(row, "volume")), + ) + except (ValueError, TypeError) as exc: + raise InvalidDataError( + "stock_daily", + f"{path}: malformed row for {sym} on {date}: {exc}", + ) from exc + return result # ------------------------------------------------------------------ # # Factor @@ -359,74 +466,112 @@ def _read_bars_for_day(self, symbol: str, date: str) -> List[Bar]: 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`. + + Task 15 performance: factor files are parsed once per day; the + per-symbol cumulative view enables O(log N) window slicing. + """ validate_symbol(symbol) validate_yyyymmdd(start, name="start") validate_yyyymmdd(end, name="end") if start > end: raise InvalidDataError("window", f"start {start} > end {end}") - cache_key = CacheKey( - self._data_root, self._source_name, "factor", symbol, "", start, end - ) - cached = self._cache.get(cache_key) - if cached is not None: - return cached + self._ensure_symbol_factors(symbol, start, end) + cumulative = self._symbol_factors[symbol] + dates = [d for d, _ in cumulative] + idx_start = bisect_left(dates, start) + idx_end = bisect_right(dates, end) + 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 + if not cumulative: + self._extend_symbol_factors(symbol, start, end) + return + if end > cumulative[-1][0]: + self._extend_symbol_factors(symbol, cumulative[-1][0], end) + if start < cumulative[0][0]: + 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: - raise MissingDataError("factor", f"no trading days in [{start}, {end}]") - rows: List[Tuple[str, Decimal]] = [] + return + existing = self._symbol_factors.get(symbol, []) + have = {d for d, _ in existing} for trading_day in calendar: - path = self._root_path / "stock_factor" / f"{trading_day}.csv" - if not path.exists(): - raise MissingDataError( - "stock_factor snapshot", - f"no snapshot for {trading_day} at {path}", - ) - try: - df = pd.read_csv( - path, dtype={"symbol": str, "date": str, "factor": str} - ) - except Exception as exc: + if trading_day in have: + continue + factor = self._read_single_factor(symbol, trading_day) + if factor is not None: + existing.append((trading_day, factor)) + existing.sort(key=lambda x: x[0]) + 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.""" + 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 + 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)) + 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}", + ) + result: Dict[str, Decimal] = {} + # `itertuples` mirrors `_parse_daily_file` (task 15: ~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"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 != {trading_day}: - raise InvalidDataError( - "stock_factor.date", - f"CSV date(s) {file_dates} do not match filename {trading_day}", - ) - symbol_rows = df[df["symbol"] == symbol] - if symbol_rows.empty: - raise MissingDataError( - "factor", - f"symbol {symbol!r} missing from {path}", + f"{path}: duplicate row for {sym!r} on {date}", ) - if len(symbol_rows) != 1: - raise InvalidDataError( - "factor", - f"expected one row for {symbol!r} on {trading_day}, got {len(symbol_rows)}", - ) - row = symbol_rows.iloc[0] - row_date = validate_yyyymmdd(row["date"], name="factor.date") - if row_date != trading_day: + row_date = validate_yyyymmdd(getattr(row, "date"), name="factor.date") + if row_date != date: raise InvalidDataError( "factor.date", - f"expected {trading_day}, got {row_date}", + f"expected {date}, got {row_date}", ) + raw = getattr(row, "factor") try: - factor = Decimal(str(row["factor"])) + factor = Decimal(str(raw)) except Exception as exc: raise InvalidDataError( "factor.value", - f"{row['factor']!r} is not Decimal", + f"{raw!r} is not Decimal", ) from exc if not factor.is_finite() or factor <= 0: raise InvalidDataError( "factor.value", f"non-positive or non-finite factor: {factor}", ) - rows.append((row_date, factor)) - self._cache.put(cache_key, rows) - return rows + result[sym] = factor + return result diff --git a/src/hqbacktest/data/memory_portal.py b/src/hqbacktest/data/memory_portal.py index 64d65b8..cfcee87 100644 --- a/src/hqbacktest/data/memory_portal.py +++ b/src/hqbacktest/data/memory_portal.py @@ -12,8 +12,7 @@ from typing import Dict, Iterable, List, Mapping, Sequence, Tuple from ..domain.bar import Bar -from ..domain.enums import EventType -from .errors import InvalidDataError, MissingDataError +from .errors import InvalidDataError, MissingDataError, SnapshotFileMissingError from .portal import DataVersion, MarketDataPortal from .validators import ( assert_unique_sorted, @@ -149,6 +148,8 @@ def get_calendar(self, start: str, end: str) -> List[str]: raise InvalidDataError("window", f"start {start} > end {end}") idx_start = bisect_left(self.calendar, start) idx_end = bisect_right(self.calendar, end) + # Return a defensive copy; mutating it must never corrupt the + # portal's internal state. return list(self.calendar[idx_start:idx_end]) def is_trading_day(self, date: str) -> bool: @@ -169,30 +170,40 @@ def next_trading_day(self, date: str) -> str: raise MissingDataError("next trading day", f"no day after {date}") return self.calendar[idx] - def get_universe(self, date: str) -> List[str]: + def get_universe(self, date: str, include_bj: bool = False) -> List[str]: validate_yyyymmdd(date) if date not in self.universe_by_date: - # Convention: walk back to the most recent universe snapshot <= date. - for snapshot_date in sorted(self.universe_by_date.keys(), reverse=True): - if snapshot_date <= date: - return sorted(self.universe_by_date[snapshot_date]) - raise MissingDataError("universe", f"no universe on or before {date}") - return sorted(self.universe_by_date[date]) + # Match the production CSV portal: a missing whole-day stock-list + # snapshot is an infrastructure failure (SnapshotFileMissingError), + # not a per-symbol gap. This keeps the two portals' exception + # types identical for the same fixture (task 14 parity). + raise SnapshotFileMissingError( + "stock_list", f"/stock_list/{date}.csv" + ) + snapshot = sorted(self.universe_by_date[date]) + if include_bj: + return snapshot + return [sym for sym in snapshot if not sym.endswith(".BJ")] def get_bars(self, symbol: str, start: str, end: str) -> List[Bar]: + """Return bars in [start, end], allowing per-day gaps. + + An empty result (the symbol did not trade at all in the window, or + the window has no trading days) is returned as `[]` rather than + raising. This matches the production `HqDataCsvPortal` semantics + introduced in task 14. + """ bars = self._bars_in_window(symbol, start, end) - if not bars: - raise MissingDataError("bars", f"no bars for {symbol} in [{start}, {end}]") return list(bars) def get_factor( self, symbol: str, start: str, end: str ) -> List[Tuple[str, Decimal]]: + """Return (date, factor) tuples in [start, end], allowing gaps. + + Empty result is `[]`, not an error, matching the CSV portal. + """ rows = self._factors_in_window(symbol, start, end) - if not rows: - raise MissingDataError( - "factor", f"no factor for {symbol} in [{start}, {end}]" - ) return list(rows) def calendar_set(self) -> set: diff --git a/src/hqbacktest/data/portal.py b/src/hqbacktest/data/portal.py index 552c181..181e6b5 100644 --- a/src/hqbacktest/data/portal.py +++ b/src/hqbacktest/data/portal.py @@ -47,8 +47,16 @@ def next_trading_day(self, date: str) -> str: """Return the trading day strictly after `date`, or raise.""" ... - def get_universe(self, date: str) -> List[str]: - """Historical stock list as of `date` (per contract §3.1).""" + def get_universe(self, date: str, include_bj: bool = False) -> List[str]: + """Historical stock list as of `date` (per contract §3.1). + + `include_bj=False` (the default) excludes Beijing Stock Exchange + (`.BJ`) symbols, which are not supported in v0.1 (no first-day + limit-up/down rule, distinct trading calendar, etc.). Pass + `include_bj=True` to opt in; the engine still treats `.BJ` symbols + as ordinary stocks from a data-layer perspective, but the broker + and rule set do not yet enforce BSE-specific rules. + """ ... def get_bars(self, symbol: str, start: str, end: str) -> List["Bar"]: diff --git a/src/hqbacktest/domain/bar.py b/src/hqbacktest/domain/bar.py index 9e75fe1..d2f6a5c 100644 --- a/src/hqbacktest/domain/bar.py +++ b/src/hqbacktest/domain/bar.py @@ -10,7 +10,14 @@ class Bar: """A single trading day's bar for a single symbol. - All prices are unadjusted (contract §3.1). OHLC and volume are immutable. + All prices are unadjusted (contract §3.1). OHLC and volume are + immutable. + + `volume` is the day's traded volume expressed in **lots (手)** — the + unit used by the upstream Tushare `hqdata` adapter. 1 lot = 100 + shares for both Shanghai and Shenzhen. Callers that need shares must + multiply by `LOT_SIZE` (see `domain.money`). Strategies that compare + volumes across symbols MUST do so on this same lot basis. """ symbol: str @@ -19,7 +26,7 @@ class Bar: high: Decimal low: Decimal close: Decimal - volume: int + volume: int # 单位:手 (1 手 = 100 股) def __post_init__(self) -> None: if not self.symbol or not isinstance(self.symbol, str): diff --git a/src/hqbacktest/domain/enums.py b/src/hqbacktest/domain/enums.py index 5e2ac39..8a51f13 100644 --- a/src/hqbacktest/domain/enums.py +++ b/src/hqbacktest/domain/enums.py @@ -57,6 +57,9 @@ class RejectReason(Enum): MISSING_DATA = "MISSING_DATA" DUPLICATE_ORDER = "DUPLICATE_ORDER" BACKTEST_ENDED = "BACKTEST_ENDED" + # Task 18: order targets a symbol outside the strategy's declared + # universe. Only emitted when `set_universe` has been called. + OUT_OF_UNIVERSE = "OUT_OF_UNIVERSE" OTHER = "OTHER" @@ -81,6 +84,7 @@ class EventType(Enum): ORDER_FILLED = "ORDER_FILLED" DATA_ERROR = "DATA_ERROR" + DATA_WARNING = "DATA_WARNING" RUN_FAILED = "RUN_FAILED" diff --git a/src/hqbacktest/domain/fill.py b/src/hqbacktest/domain/fill.py index 0afb5db..a042ed0 100644 --- a/src/hqbacktest/domain/fill.py +++ b/src/hqbacktest/domain/fill.py @@ -65,6 +65,13 @@ def __post_init__(self) -> None: raise ValueError(f"{field_name} must be quantized to 2 decimal places") if len(self.filled_at) != 8 or not self.filled_at.isdigit(): raise ValueError(f"filled_at must be YYYYMMDD, got {self.filled_at!r}") + # Task 16: 印花税 is charged only on SELL. A BUY fill carrying a + # non-zero stamp_tax would let the cash ledger drift from the + # costs table, so we reject it at construction time. + if self.side is Side.BUY and self.stamp_tax != 0: + raise ValueError( + f"stamp_tax on BUY must be 0 (印花税 only on SELL), got {self.stamp_tax}" + ) gross = cash_for_trade(self.quantity, self.price) expected_amount = gross if self.side is Side.BUY else -gross if self.amount != expected_amount: diff --git a/src/hqbacktest/domain/order.py b/src/hqbacktest/domain/order.py index fa8214f..ba7fbda 100644 --- a/src/hqbacktest/domain/order.py +++ b/src/hqbacktest/domain/order.py @@ -8,19 +8,26 @@ from dataclasses import dataclass, field from decimal import Decimal -from typing import List, Optional +from typing import Optional, Tuple from .enums import EventType, OrderStatus, OrderType, RejectReason, Side from .money import PRICE_QUANT from .state_machine import validate_transition -@dataclass +@dataclass(frozen=True) class Order: """A single strategy-issued order. - `filled_quantity` accumulates across fills; `avg_fill_price` is recomputed - on every fill using a running weighted average. + `filled_quantity` accumulates across fills; `avg_fill_price` is + recomputed on every fill using a running weighted average. + + Task 18: the dataclass is **frozen**. Strategies that receive an + Order via `Context.pending_orders()` cannot mutate any field — + attempted assignment raises `FrozenInstanceError`. Lifecycle + mutations (`transition`, `record_fill`) use `object.__setattr__` + to bypass the freeze; that path is engine-internal and never + exposed to strategy code. """ order_id: str @@ -45,7 +52,7 @@ class Order: reject_reason: Optional[RejectReason] = None reject_detail: Optional[str] = None - fill_ids: List[str] = field(default_factory=list) + fill_ids: Tuple[str, ...] = field(default_factory=tuple) notes: str = "" def __post_init__(self) -> None: @@ -108,28 +115,32 @@ def transition( reason: Optional[RejectReason] = None, detail: Optional[str] = None, ) -> None: - """Move the order to `target`, stamping the matching timestamp.""" + """Move the order to `target`, stamping the matching timestamp. + + Task 18: this mutator uses `object.__setattr__` to bypass the + frozen-dataclass guard. Only the engine / broker may call it. + """ self._validate_date(at, "at") validate_transition(self.status, target) - self.status = target + object.__setattr__(self, "status", target) if target is OrderStatus.ACCEPTED: - self.accepted_at = at + object.__setattr__(self, "accepted_at", at) elif target is OrderStatus.PENDING: - self.pending_at = at + object.__setattr__(self, "pending_at", at) elif target is OrderStatus.CANCELLED: - self.cancelled_at = at + object.__setattr__(self, "cancelled_at", at) if reason is not None: - self.reject_reason = reason + object.__setattr__(self, "reject_reason", reason) if detail is not None: - self.reject_detail = detail + object.__setattr__(self, "reject_detail", detail) elif target is OrderStatus.REJECTED: - self.rejected_at = at - self.reject_reason = reason or RejectReason.OTHER - self.reject_detail = detail or "" + object.__setattr__(self, "rejected_at", at) + object.__setattr__(self, "reject_reason", reason or RejectReason.OTHER) + object.__setattr__(self, "reject_detail", detail or "") elif target is OrderStatus.FILLED: - self.filled_at = at + object.__setattr__(self, "filled_at", at) elif target is OrderStatus.PARTIALLY_FILLED: - self.partially_filled_at = at + object.__setattr__(self, "partially_filled_at", at) def record_fill( self, @@ -138,9 +149,16 @@ def record_fill( price: Decimal, at: str, ) -> None: - """Record a single fill and update running aggregates.""" + """Record a single fill and update running aggregates. + + v0.1 always moves orders through ACCEPTED → PENDING before any + fill arrives, so the ACCEPTED branch in the status guard was + unreachable in practice (task 16). + + Task 18: this mutator uses `object.__setattr__` to bypass the + frozen-dataclass guard. Only the engine / broker may call it. + """ if self.status not in ( - OrderStatus.ACCEPTED, OrderStatus.PENDING, OrderStatus.PARTIALLY_FILLED, ): @@ -164,15 +182,19 @@ def record_fill( f"fill overflow: {self.filled_quantity}+{quantity} > {self.quantity}" ) if self.avg_fill_price is None: - self.avg_fill_price = price.quantize(PRICE_QUANT) + new_avg = price.quantize(PRICE_QUANT) else: total = self.avg_fill_price * Decimal( self.filled_quantity ) + price * Decimal(quantity) - self.avg_fill_price = (total / Decimal(new_filled)).quantize(PRICE_QUANT) - self.filled_quantity = new_filled - self.fill_ids.append(fill_id) - if self.filled_quantity == self.quantity: + new_avg = (total / Decimal(new_filled)).quantize(PRICE_QUANT) + object.__setattr__(self, "avg_fill_price", new_avg) + object.__setattr__(self, "filled_quantity", new_filled) + # fill_ids is an immutable tuple so a strategy holding the frozen + # Order cannot append/clear it in place (task 18). Rebuild the + # tuple here via object.__setattr__ to bypass the frozen guard. + object.__setattr__(self, "fill_ids", self.fill_ids + (fill_id,)) + if new_filled == self.quantity: self.transition(OrderStatus.FILLED, at=at) elif self.status is not OrderStatus.PARTIALLY_FILLED: self.transition(OrderStatus.PARTIALLY_FILLED, at=at) diff --git a/src/hqbacktest/engine/broker.py b/src/hqbacktest/engine/broker.py index 34c4d92..9ec2fd1 100644 --- a/src/hqbacktest/engine/broker.py +++ b/src/hqbacktest/engine/broker.py @@ -1,11 +1,11 @@ """SimulatedBroker: market-on-open matching with rule set and cost model. -Rules (task 7 + task 8): +Rules (task 7 + task 8 + task 16): * Only `OrderType.MARKET` orders are supported; any other type is rejected by `Context` before reaching the broker. * The bar for `today` is read first: `MissingDataError` (suspended / no bar) becomes `bar_available=False` for the rule set, while - `InvalidDataError` and I/O errors propagate and abort the run. + `InvalidDataError` / I/O errors propagate and abort the run. * Each order then goes through `TradingRuleSet.first_denial`; the first denial short-circuits with a typed `RejectReason`. * Fees come from `CostModel.compute(order, price, quantity)`; the @@ -14,17 +14,31 @@ * Ledger-side rejections (`InsufficientCashError`, `InsufficientSharesError`) are produced by `Portfolio.apply_fill` and converted by the engine into typed rejections. + +Task 16 batch-matching order (A-share convention): + * Within a single `OPEN_MATCH(today)` batch, SELL orders match + **before** BUY orders. This mirrors "卖出资金当日可用": the + proceeds of a SELL can fund a same-batch BUY, so a rotation + ("卖旧买新") is not falsely rejected for INSUFFICIENT_CASH. + * The rule set's `InsufficientCashRule` checks against a rolling + `running_cash` that starts at `portfolio_cash`, increases by each + successful SELL's net proceeds, and decreases by each successful + BUY's cost. The ordering within each side (SELL-only or BUY-only) + preserves the strategy's submission order. """ from decimal import Decimal from typing import Callable, List, Optional, Tuple -from ..data.errors import MissingDataError +from ..data.errors import ( + MissingDataError, + SnapshotFileMissingError, +) from ..data.portal import MarketDataPortal from ..domain.bar import Bar from ..domain.enums import EventType, OrderStatus, RejectReason, Side from ..domain.fill import Fill -from ..domain.money import PRICE_QUANT +from ..domain.money import PRICE_QUANT, quantize_cash from ..domain.order import Order from .cost_model import CostModel, DefaultCostModel from .rule_set import RuleCheckContext, TradingRuleSet @@ -55,25 +69,50 @@ def match( `sellable_quantity_for(symbol)` is a callable returning the current T+1 sellable shares for the symbol (0 if no position). It mirrors - the engine's view of the portfolio so the rule set stays pure. - - `portfolio_cash` is a snapshot taken before this batch: for later - orders in the same batch the rule-level cash check may be stale. - That is safe because `Portfolio.apply_fill` re-checks every fill - against the live ledger and is the final arbiter. + the engine's view of the portfolio so the rule set stays pure; + the value is captured **before** the batch runs (T+1 state from + the previous day's settlement), and is not re-checked after each + SELL because a SELL only reduces `sellable_quantity`, so stale + values only over-estimate availability (fail-safe for the rule). + + Task 16: orders are partitioned into SELLs and BUYs. All SELLs + match first (in submission order) so that their proceeds are + available to fund the subsequent BUYs (rolling cash). The + partition is a stable sort: the relative order within each side + is preserved, but cross-side order is normalized to + [SELLs..., BUYs...] as required by A-share convention. + + Results are returned in **matching order** ([SELLs..., BUYs...]), + NOT submission order: the engine applies fills in the returned + order, so a BUY must be applied only after its funding SELL has + already credited the portfolio's cash. Re-assembling into + submission order would break the rolling-cash guarantee (a BUY + submitted before its funding SELL would be applied first and + falsely rejected for INSUFFICIENT_CASH). """ + sell_orders = [o for o in orders if o.side is Side.SELL] + buy_orders = [o for o in orders if o.side is Side.BUY] + # Stable partition: preserve relative submission order within each + # side (the original `orders` list is already insertion-ordered). + ordered = sell_orders + buy_orders + running_cash = quantize_cash(portfolio_cash) results: List[MatchResult] = [] - for order in orders: - results.append( - self._match_one( - order, - portal, - today, - rule_set, - portfolio_cash, - sellable_quantity_for(order.symbol), - ) + for order in ordered: + match_result = self._match_one( + order, + portal, + today, + rule_set, + running_cash, + sellable_quantity_for(order.symbol), ) + results.append(match_result) + # Roll cash: only successful fills update the running balance. + # `net_amount()` is already signed (SELL proceeds positive, + # BUY costs negative), so a single addition covers both sides. + _, fill, _, _ = match_result + if fill is not None: + running_cash = quantize_cash(running_cash + fill.net_amount()) return results # ------------------------------------------------------------------ # @@ -97,13 +136,20 @@ def _match_one( f"order not PENDING (status={order.status.name})", ) - # Read the bar for today. `MissingDataError` (suspended symbol / - # no bar for this trading day) is a BUSINESS outcome: the order is - # rejected via the rule set (contract §4: 停牌标的不可成交, but the - # run continues). `InvalidDataError` / I/O errors are infrastructure - # failures and propagate so the engine aborts with `RunFailed`. + # Read the bar for today. + # * `MissingDataError` (per-symbol gap: suspended / delisted / + # pre-IPO): the bar is simply unavailable. The order is rejected + # by the rule set (`bar_available=False`); the run continues. + # * `SnapshotFileMissingError` (the whole daily file is missing + # on disk): a data infrastructure failure. It MUST propagate so + # the engine aborts the run with `RunFailed` rather than + # silently treating the snapshot as empty. + # * `InvalidDataError` / I/O errors: infrastructure failures and + # propagate. try: bars = portal.get_bars(order.symbol, today, today) + except SnapshotFileMissingError: + raise except MissingDataError: bars = [] bar_available = bool(bars) diff --git a/src/hqbacktest/engine/context.py b/src/hqbacktest/engine/context.py index 92bf291..ab97540 100644 --- a/src/hqbacktest/engine/context.py +++ b/src/hqbacktest/engine/context.py @@ -30,7 +30,7 @@ from ..data.data_view import DataView from ..data.validators import validate_symbol -from ..domain.enums import EventType, OrderType, OrderStatus, Side +from ..domain.enums import EventType, OrderType, OrderStatus, RejectReason, Side from ..domain.money import LOT_SIZE, is_positive, quantize_cash, round_lot from ..domain.order import Order from ..domain.portfolio import Portfolio @@ -78,6 +78,7 @@ def __init__( self._data_view = data_view self._universe: List[str] = [] self._pending_orders: List[Order] = [] + self._out_of_universe_orders: List[Order] = [] self._initialized: bool = False self._universe_locked: bool = False self._run_finished: bool = False @@ -114,6 +115,21 @@ def _consume_pending_orders(self) -> List[Order]: self._pending_orders = [] return orders + def _consume_out_of_universe_orders(self) -> List[Order]: + """Return and clear out-of-universe orders for audit-trail merge.""" + orders = list(self._out_of_universe_orders) + self._out_of_universe_orders = [] + return orders + + def _has_out_of_universe_orders(self) -> bool: + """True when out-of-universe rejections are waiting to be drained. + + Peek-only (does not clear): the scheduler uses this to decide + whether to invoke the matcher, while the engine drains the list + via `_consume_out_of_universe_orders` (task 18). + """ + return bool(self._out_of_universe_orders) + # ------------------------------------------------------------------ # # Guards # ------------------------------------------------------------------ # @@ -195,11 +211,30 @@ def position(self, symbol: str) -> Optional[Position]: return replace(pos) if pos is not None else None def universe(self) -> List[str]: + """Snapshot of the strategy's declared universe (defensive copy).""" self._require_active("universe") return list(self._universe) + def historical_universe(self) -> List[str]: + """The historical stock list as of the current `visible_through`. + + Task 18: this is the only universe accessor that reads through + the data portal. It is constrained by `visible_through` and + must not be used to bypass the data view. When no `DataView` + is published (e.g. in `initialize`), an empty list is returned. + """ + self._require_active("historical_universe") + if self._data_view is None: + return [] + return self._data_view.universe() + def pending_orders(self) -> List[Order]: - """Snapshot of in-flight orders (engine clears them at matching).""" + """Snapshot of in-flight orders (engine clears them at matching). + + Task 18: the returned list and its `Order` elements are + defensive copies / frozen instances — strategies cannot + mutate the engine's view of the order. + """ self._require_active("pending_orders") return list(self._pending_orders) @@ -259,6 +294,35 @@ def _next_order_id(self) -> str: self._order_counter += 1 return f"O{self._current_date}-{self._order_counter:06d}" + def _coerce_amount(self, value, name: str) -> Decimal: + """Coerce a monetary value to `Decimal`, accepting int/str/Decimal. + + `float` and `bool` are rejected (contract rule 5: no binary + float enters the ledger), and NaN / Inf are rejected. Task 20: + lets strategies use literal cash values (`15000`, `'15000'`) + without wrapping them in `Decimal(...)`. + """ + if isinstance(value, bool): + raise StrategyLifecycleError(f"{name} must be a number, got bool") + if isinstance(value, float): + raise StrategyLifecycleError( + f"{name} must not be float (contract rule 5); got {value!r}" + ) + if not isinstance(value, (Decimal, int, str)): + raise StrategyLifecycleError( + f"{name} must be Decimal/int/str, got {type(value).__name__}" + ) + if isinstance(value, (int, str)): + try: + value = Decimal(str(value)) + except Exception as exc: + raise StrategyLifecycleError( + f"{name}={value!r} is not a valid Decimal: {exc}" + ) from exc + if not value.is_finite(): + raise StrategyLifecycleError(f"{name} must be finite, got {value}") + return value + # ------------------------------------------------------------------ # # Order intents (contract: never mutates the ledger) # ------------------------------------------------------------------ # @@ -287,15 +351,18 @@ def order( def order_value( self, symbol: str, - value: Decimal, + value, order_type: OrderType = OrderType.MARKET, ) -> Optional[Order]: - """Place an order for `value` worth of `symbol` (positive => BUY).""" + """Place an order for `value` worth of `symbol` (positive => BUY). + + `value` may be a `Decimal`, `int`, or numeric string; `float` + remains forbidden (contract rule 5). Task 20: widening the + accepted types here lets strategies use literal cash values + without first wrapping them in `Decimal(...)`. + """ self._require_orderable("order_value") - if not isinstance(value, Decimal): - raise StrategyLifecycleError( - f"value must be Decimal, got {type(value).__name__}" - ) + value = self._coerce_amount(value, "value") if value == 0: return None price = self.current_price(symbol) @@ -342,15 +409,16 @@ def order_target( def order_target_value( self, symbol: str, - target_value: Decimal, + target_value, order_type: OrderType = OrderType.MARKET, ) -> Optional[Order]: - """Reconcile holdings towards `target_value` worth of `symbol`.""" + """Reconcile holdings towards `target_value` worth of `symbol`. + + `target_value` may be a `Decimal`, `int`, or numeric string; + `float` remains forbidden (contract rule 5). Task 20. + """ self._require_orderable("order_target_value") - if not isinstance(target_value, Decimal): - raise StrategyLifecycleError( - f"target_value must be Decimal, got {type(target_value).__name__}" - ) + target_value = self._coerce_amount(target_value, "target_value") if target_value < 0: raise StrategyLifecycleError( f"target_value must be non-negative, got {target_value}" @@ -425,7 +493,7 @@ def _create_order( side: Side, quantity: int, order_type: OrderType = OrderType.MARKET, - ) -> Order: + ) -> Optional[Order]: # Contract rule 7: every order path funnels through here, so the # order-type allow-list and symbol validation cannot be bypassed by # the convenience helpers. @@ -440,10 +508,27 @@ def _create_order( validate_symbol(symbol) if not is_positive(Decimal(quantity)): raise StrategyLifecycleError(f"quantity must be positive, got {quantity}") - lot_aligned = round_lot(quantity, lot_size=LOT_SIZE) - if lot_aligned == 0: - raise StrategyLifecycleError( - f"quantity {quantity} is below one lot of {LOT_SIZE} shares" + # Task 16: lot-alignment applies to BUY only. A-share rules allow + # odd-lot SELLs so positions holding non-lot quantities can be + # fully closed; round_lot() on a SELL would silently shrink + # the order (e.g. 150 -> 100), violating the contract that the + # broker sees the exact share count the strategy submitted. + if side is Side.BUY: + lot_aligned = round_lot(quantity, lot_size=LOT_SIZE) + if lot_aligned == 0: + raise StrategyLifecycleError( + f"quantity {quantity} is below one lot of {LOT_SIZE} shares" + ) + final_quantity = lot_aligned + else: + final_quantity = quantity + # Task 18: when a universe has been declared, orders for symbols + # outside it must be rejected (typed reason + audit-trail event). + # When the strategy has not called `set_universe`, the universe + # is empty and trading is unrestricted. + if self._universe and symbol not in self._universe: + return self._reject_out_of_universe( + symbol=symbol, side=side, quantity=final_quantity ) # Guaranteed by _require_orderable in every public order method. created_session = self._phase @@ -451,7 +536,7 @@ def _create_order( order_id=self._next_order_id(), symbol=symbol, side=side, - quantity=lot_aligned, + quantity=final_quantity, order_type=order_type, created_at=self._current_date, created_session=created_session, @@ -468,9 +553,70 @@ def _create_order( phase=EventType.ORDER_CREATED, order_id=order.order_id, detail=( - f"{side.name} {lot_aligned} {symbol} " + f"{side.name} {final_quantity} {symbol} " f"(session={created_session.name})" ), ) ) return order + + def _reject_out_of_universe( + self, + *, + symbol: str, + side: Side, + quantity: int, + ) -> Optional[Order]: + """Build a REJECTED order for a symbol outside the declared + universe. Records ORDER_REJECTED + ORDER_CREATED events so the + audit trail is complete (task 18). The order is NOT appended + to `_pending_orders` so the broker never sees it (REJECTED is + a terminal status and would corrupt the broker's state + machine); instead it lands in `_out_of_universe_orders` and + is folded into the engine's `orders_table` at result build. + """ + created_session = self._phase + order = Order( + order_id=self._next_order_id(), + symbol=symbol, + side=side, + quantity=quantity, + order_type=OrderType.MARKET, + created_at=self._current_date, + created_session=created_session, + ) + order.transition(OrderStatus.ACCEPTED, at=self._current_date) + # Stamp the reason onto the Order itself (not only the event log) + # so `orders_table.reject_reason` and the ORDER_REJECTED event + # agree (task 18 audit integrity). + order.transition( + OrderStatus.REJECTED, + at=self._current_date, + reason=RejectReason.OUT_OF_UNIVERSE, + detail=( + f"{symbol} not in declared universe " + f"({len(self._universe)} symbols); order rejected" + ), + ) + self._out_of_universe_orders.append(order) + self._event_log.record( + EngineEvent( + date=self._current_date, + phase=EventType.ORDER_CREATED, + order_id=order.order_id, + detail=f"{side.name} {quantity} {symbol} (session={created_session.name})", + ) + ) + self._event_log.record( + EngineEvent( + date=self._current_date, + phase=EventType.ORDER_REJECTED, + order_id=order.order_id, + error=RejectReason.OUT_OF_UNIVERSE.name, + detail=( + f"{symbol} not in declared universe " + f"({len(self._universe)} symbols); order rejected" + ), + ) + ) + return order diff --git a/src/hqbacktest/engine/engine.py b/src/hqbacktest/engine/engine.py index b4c0210..9fcd052 100644 --- a/src/hqbacktest/engine/engine.py +++ b/src/hqbacktest/engine/engine.py @@ -24,10 +24,10 @@ from dataclasses import asdict from decimal import Decimal -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Optional, Tuple from ..data.data_view import DataView -from ..data.errors import MissingDataError +from ..data.errors import DataError, MissingDataError, SnapshotFileMissingError from ..data.hqdata_portal import HqDataCsvPortal, resolve_source_location from ..data.portal import MarketDataPortal from ..domain.enums import EventType, OrderStatus, RejectReason @@ -39,8 +39,11 @@ from .config import BacktestConfig from .context import Context from .corporate_actions import ( + DEFAULT_JUMP_BAND, + FactorDiagnostic, FactorDiagnosticCollector, V01_ADJUSTMENT_POLICY, + analyze_factor_series, ) from .errors import ( ConfigurationError, @@ -48,6 +51,7 @@ RunFailed, StrategyLifecycleError, ) +from ..data.data_view import CURRENT_PRICE_LOOKBACK from .events import EngineEvent, EventLog from .iterator import TradingDayIterator from .metrics import EquityPoint, compute_metrics @@ -56,6 +60,19 @@ from .strategy import NullStrategy, Strategy +def _lookback_start_date(today: str) -> str: + """Compute the lookback-window start date for valuation fallbacks. + + Mirrors `DataView._trading_day_lookback_start`: a generous 5-year + window that comfortably covers `CURRENT_PRICE_LOOKBACK` trading days + without forcing the portal to scan the full pre-start history. + """ + yyyymmdd = int(today) + year = yyyymmdd // 10000 + start_year = max(year - 5, 1900) + return f"{start_year}0101" + + class BacktestEngine: """Drive the daily event loop for a backtest run.""" @@ -81,6 +98,19 @@ def __init__( self._event_log = EventLog() self._portfolio = Portfolio(initial_cash=config.initial_cash) self._factor_diagnostics = FactorDiagnosticCollector() + # Task 19: per-symbol cumulative factor history (sorted by + # date) so the engine can run holdings-period factor + # diagnostics incrementally without re-reading the portal. + self._factor_history: Dict[str, List[Tuple[str, Decimal]]] = {} + # Holding-period jump band (relative). A factor ratio outside + # this band while a symbol is held is a strong dividend / split + # signal. 0.1% matches task 19's "cannot be ignored" threshold; + # the default `DEFAULT_JUMP_BAND` (0.5, 2.0) is reserved for + # general factor-quality diagnostics. + self._holding_jump_band: Tuple[Decimal, Decimal] = ( + Decimal("0.999"), + Decimal("1.001"), + ) self._fills: List[Fill] = [] self._equity_curve: List[EquityPoint] = [] # Every order the engine has consumed (insertion-ordered), so the @@ -259,6 +289,11 @@ def _run_day_safely( # End-of-day settlement: roll today's buys into sellable (T+1). self._portfolio.settle_t1(today=today, previous_date=None) self._snapshot_equity(today, portal) + # Task 19: run holdings-period factor diagnostics for + # symbols held or traded today. Diagnostics are pure + # observations: they never mutate cash, positions or + # equity, and they cannot abort the run. + self._run_factor_diagnostics(today, portal) except RunFailed: raise except Exception as exc: @@ -339,17 +374,32 @@ def _order_row(self, order: Order) -> Dict[str, Any]: def _snapshot_equity(self, today: str, portal: MarketDataPortal) -> None: """Record one EquityPoint using today's close for market value. - Contract §4: day-end valuation uses D's valid unadjusted close. If a - HELD symbol has no valid close (missing bar, or close <= 0), the run - FAILS with a DATA_ERROR event — v0.1 never silently skips valuation, - never uses previous closes, and never values holdings at zero. + Contract §4 + task 14 valuation semantics: + * Preferred source: today's valid unadjusted close (D's bar). + * Fallback: when the held symbol has no bar on `today` + (suspended / delisted / pre-IPO), use the most recent valid + close from the same `current_price` lookback window and + record a `DATA_WARNING` event so the audit trail reflects + the deviation. The fallback is bounded by + `DataView.CURRENT_PRICE_LOOKBACK` (20) trading days. + * If neither today's close nor any lookback close is + available, the run FAILS with `DATA_ERROR` — v0.1 never + values holdings at zero and never silently drops a holding. """ prices: Dict[str, Decimal] = {} for symbol, position in self._portfolio.positions.items(): if position.quantity == 0: continue - close = self._close_price_or_none(portal, symbol, today) - if close is None: + today_price = self._close_price_or_none(portal, symbol, today) + if today_price is not None: + prices[symbol] = today_price + continue + # No bar for `symbol` today (suspended / delisted / pre-IPO). + # Fall back to the most recent valid close within the same + # lookback window used by `DataView.current_price` and emit a + # `DATA_WARNING` so the audit trail reflects the deviation. + fallback = self._lookback_price_or_none(portal, symbol, today) + if fallback is None: self._event_log.record( EngineEvent( date=today, @@ -357,7 +407,7 @@ def _snapshot_equity(self, today: str, portal: MarketDataPortal) -> None: error="MissingDataError", detail=( f"no valid close for held symbol {symbol} on " - f"{today}; valuation aborted" + f"{today} (lookback exhausted); valuation aborted" ), ) ) @@ -365,26 +415,53 @@ def _snapshot_equity(self, today: str, portal: MarketDataPortal) -> None: today, "AFTER_TRADING_END", MissingDataError( - "close", f"no valid close for held symbol {symbol} on {today}" + "close", + f"no valid close for held symbol {symbol} on {today}", ), ) - prices[symbol] = close + self._event_log.record( + EngineEvent( + date=today, + phase=EventType.DATA_WARNING, + detail=( + f"held symbol {symbol} has no bar on {today}; " + f"valued at fallback close {fallback}" + ), + ) + ) + prices[symbol] = fallback market_value = self._portfolio.market_value(prices) total_equity = self._portfolio.cash + market_value + # Task 17: anchor the first day's `daily_return` and the + # `drawdown` series to `initial_cash` (not a zero seed). A first- + # day P&L must flow into the equity curve so `∏(1 + r) == 1 + + # total_return` and `max_drawdown` can see day-1 drawdowns. prev_total = self._equity_curve[-1].total_equity if self._equity_curve else None - if prev_total is None or prev_total == 0: + if prev_total is None: + # First trading day: benchmark the return against initial_cash + # (task 17) so a first-day P&L flows into the return series. + daily_return = ( + total_equity / self._config.initial_cash - Decimal("1") + if self._config.initial_cash > 0 + else Decimal("0") + ) + elif prev_total == 0: + # Defensive: a zero prior equity cannot produce a return ratio. daily_return = Decimal("0") else: daily_return = total_equity / prev_total - Decimal("1") + # Drawdown: the running peak must include `initial_cash`, otherwise + # a first-day loss is silently lost once a later day's equity stays + # below initial_cash but above the prior day's equity (task 17: + # "回撤峰值序列以 initial_cash 为初始峰值"). + peak = self._config.initial_cash if self._equity_curve: - peak = max(pt.total_equity for pt in self._equity_curve) - drawdown = ( - max(Decimal("0"), (peak - total_equity) / peak) - if peak > 0 - else Decimal("0") - ) - else: - drawdown = Decimal("0") + peak = max(peak, *(pt.total_equity for pt in self._equity_curve)) + drawdown = ( + max(Decimal("0"), (peak - total_equity) / peak) + if peak > 0 + else Decimal("0") + ) self._equity_curve.append( EquityPoint( date=today, @@ -395,7 +472,8 @@ def _snapshot_equity(self, today: str, portal: MarketDataPortal) -> None: drawdown=drawdown, ) ) - # Per-day position snapshot with today's actual close prices. + # Per-day position snapshot with the valuation price actually used + # (today's close or a lookback close for suspended symbols). for symbol, position in self._portfolio.positions.items(): if position.quantity == 0: continue @@ -420,9 +498,13 @@ def _close_price_or_none( Only `MissingDataError` maps to None; corrupt data or I/O errors propagate and abort the run via the caller's RunFailed wrapping. + `SnapshotFileMissingError` (whole-day file gone) is an + infrastructure failure and propagates as well. """ try: bars = portal.get_bars(symbol, today, today) + except SnapshotFileMissingError: + raise except MissingDataError: return None if not bars: @@ -432,11 +514,44 @@ def _close_price_or_none( return None return close + @staticmethod + def _lookback_price_or_none( + portal: MarketDataPortal, symbol: str, today: str + ) -> Optional[Decimal]: + """Most recent valid close for `symbol` within the lookback window. + + Mirrors `DataView.current_price` exactly so the engine's valuation + fallback agrees with what strategies see through `DataView`. Used + only when `_close_price_or_none` returns None for the same day. + """ + try: + trading_days = portal.get_calendar(_lookback_start_date(today), today) + except MissingDataError: + trading_days = [] + if not trading_days: + return None + lookback = trading_days[-CURRENT_PRICE_LOOKBACK:] + for day in reversed(lookback): + try: + bars = portal.get_bars(symbol, day, day) + except SnapshotFileMissingError: + raise + except MissingDataError: + continue + if not bars: + continue + close = bars[0].close + if close is not None and close > 0: + return close + return None + def _sellable_for(self, symbol: str) -> int: pos = self._portfolio.positions.get(symbol) return pos.sellable_quantity if pos else 0 - def _on_open_match(self, today: str, pending: List[Order]) -> None: + def _on_open_match( + self, today: str, pending: List[Order], context: Context + ) -> None: """Match pending orders at `OPEN_MATCH(today)` and apply fills. Each order goes through `TradingRuleSet.evaluate` first; the first @@ -449,6 +564,11 @@ def _on_open_match(self, today: str, pending: List[Order]) -> None: """ for order in pending: self._orders[order.order_id] = order + # Task 18: fold out-of-universe rejections into the orders + # table so the audit trail sees them. These orders never + # reached the broker. + for order in context._consume_out_of_universe_orders(): + self._orders[order.order_id] = order results = self._broker.match( pending, self.portal, @@ -553,3 +673,118 @@ def _cancel_leftover_orders( def data_view(self, today: str) -> DataView: """Build a `DataView` snapshot as the engine would at `BAR_CLOSE(today)`.""" return DataView(portal=self.portal, visible_through=today) + + # ------------------------------------------------------------------ # + # Task 19: holdings-period factor diagnostics + # ------------------------------------------------------------------ # + + def _run_factor_diagnostics(self, today: str, portal: MarketDataPortal) -> None: + """For each currently-held symbol, run factor diagnostics. + + Reads the cumulative factor history for `today`, calls + `analyze_factor_series` with the tighter holdings-period + jump band (0.1% relative), and records observations on + the `FactorDiagnosticCollector` and a `DATA_WARNING` + event for each new anomaly. + + Only symbols with a non-zero position at day-end are scanned + (task 19: "持仓涉及的标的"). A symbol that was fully sold + (position back to zero) must NOT keep emitting holdings-period + warnings — its holding period has ended, and the accumulated + factor history is therefore reset. + + Diagnostics are pure observations: they MUST NOT mutate + cash, positions, or equity (contract task 9 invariant). + Missing factors or whole-day snapshot failures are silently + skipped so the run continues (the existing + `_snapshot_equity` / broker paths raise DATA_ERROR on + infrastructure failures, and we do not want to raise a + second error here). + """ + relevant = { + sym for sym, pos in self._portfolio.positions.items() if pos.quantity > 0 + } + # Drop factor history for symbols no longer held so a future + # re-entry does not compare against a pre-gap, stale factor. + for sym in list(self._factor_history): + if sym not in relevant: + del self._factor_history[sym] + for sym in sorted(relevant): + history = self._factor_history.get(sym, []) + new_factor_rows = self._load_factor_rows(sym, today) + for d, f in new_factor_rows: + if not history or history[-1][0] < d: + history.append((d, f)) + self._factor_history[sym] = history + if len(history) < 2: + continue + diagnostics = analyze_factor_series( + symbol=sym, + expected_dates=[d for d, _ in history], + factors=history, + jump_band=self._holding_jump_band, + ) + existing = { + (d.symbol, d.date, d.kind, d.detail) + for d in self._factor_diagnostics.all() + } + for diag in diagnostics: + key = (diag.symbol, diag.date, diag.kind, diag.detail) + if key in existing: + continue + self._factor_diagnostics.record(diag) + # Include the actual factor values in the audit event so + # the human can verify the ex-date dividend event from + # the event log alone (without re-reading the snapshot). + prev_factor = self._prev_factor_before(history, diag.date) + new_factor = self._factor_on(history, diag.date) + detail = ( + f"factor diagnostic: {diag.symbol} {diag.kind} on " + f"{diag.date}: factor {prev_factor} -> {new_factor}; " + f"{diag.detail}" + ) + self._event_log.record( + EngineEvent( + date=today, + phase=EventType.DATA_WARNING, + detail=detail, + ) + ) + + def _load_factor_rows(self, symbol: str, today: str) -> List[Tuple[str, Decimal]]: + """Read today's factor for `symbol`, returning a one-row list. + + Tolerates data-layer absences (MissingDataError / + SnapshotFileMissingError / InvalidDataError) by returning an + empty list (the analyzer will simply not see a row for today). + This is intentional: factor-data absences are diagnostic + observations, not run-aborting failures. Programming errors + (anything that is NOT a `DataError`) still propagate. + """ + try: + rows = self.portal.get_factor(symbol, today, today) + except DataError: + return [] + return [(today, f) for _, f in rows] + + @staticmethod + def _prev_factor_before( + history: List[Tuple[str, Decimal]], date: str + ) -> Optional[Decimal]: + """Return the most recent factor in `history` strictly + before `date`, or `None` if no earlier row exists. + """ + prev: Optional[Decimal] = None + for d, f in history: + if d < date: + prev = f + else: + break + return prev + + @staticmethod + def _factor_on(history: List[Tuple[str, Decimal]], date: str) -> Optional[Decimal]: + for d, f in history: + if d == date: + return f + return None diff --git a/src/hqbacktest/engine/intents.py b/src/hqbacktest/engine/intents.py index 77e38b5..4609b78 100644 --- a/src/hqbacktest/engine/intents.py +++ b/src/hqbacktest/engine/intents.py @@ -51,18 +51,24 @@ def side_from_quantity(quantity: int) -> Optional[Side]: def signed_diff_to_lots(target: int, current: int, lot_size: int = 100) -> int: - """Compute the signed delta between `target` and `current`, lot-aligned. + """Compute the signed delta between `target` and `current`. - Used by `order_target*` helpers: the resulting value is the share count + BUY deltas (target > current) are floored to the nearest lot (100 + shares). SELL deltas (target < current) preserve the requested + share count exactly — A-share rules allow odd-lot SELLs so a + position holding a non-lot multiple can still be flattened. Used + by `order_target*` helpers: the resulting value is the share count to send through the broker (positive => BUY, negative => SELL). """ diff = target - current if diff == 0: return 0 - sign = 1 if diff > 0 else -1 - abs_diff = abs(diff) - lots = abs_diff // lot_size # floor; contract: lot-aligned BUY/SELL - return sign * lots * lot_size + if diff < 0: + # SELL: any positive integer is legal; do NOT lot-round. + return diff + # BUY: floor to nearest lot to avoid requesting fractional shares. + lots = diff // lot_size + return lots * lot_size def target_quantity_for_value( @@ -70,7 +76,10 @@ def target_quantity_for_value( ) -> int: """Compute the desired `quantity` for a `target_value` position. - `target_value` may be `Decimal("0")` to flatten the position. + `target_value` may be `Decimal("0")` to flatten the position; in + that case the function returns `0` (the caller, e.g. + `Context.order_target_value`, will translate this into a full + flatten via `order_target(symbol, 0)`). """ if price <= 0: raise StrategyLifecycleError(f"price must be positive, got {price}") @@ -79,7 +88,7 @@ def target_quantity_for_value( f"target_value must be non-negative, got {target_value}" ) if target_value == 0: - return current_quantity # flat → no change + return 0 # flatten (per docstring) target_qty_signed = quantity_from_value(target_value, price, lot_size) # quantity_from_value already returns a lot-aligned signed count. # For a target_value (positive), the signed count is positive. diff --git a/src/hqbacktest/engine/iterator.py b/src/hqbacktest/engine/iterator.py index 6fa1705..c38b32c 100644 --- a/src/hqbacktest/engine/iterator.py +++ b/src/hqbacktest/engine/iterator.py @@ -12,7 +12,13 @@ class TradingDayIterator: - """Yields trading days in `[start, end]` (inclusive), in ascending order.""" + """Yields trading days in `[start, end]` (inclusive), in ascending order. + + Task 20: an empty trading-day window is a hard error. Silent + success on a zero-day run produces misleading "no signals" + reports and breaks reproducibility, so the iterator raises + `ConfigurationError` instead of yielding an empty sequence. + """ def __init__( self, @@ -28,6 +34,19 @@ def __init__( self._start = start self._end = end self._trading_days: List[str] = list(portal.get_calendar(start, end)) + if not self._trading_days: + # Use `source_name()` (CSV portal) or fall back to the + # `source` attribute (memory portal). Both portals expose a + # human-readable name so the error is informative. + source_name = ( + portal.source_name() + if hasattr(portal, "source_name") + else getattr(portal, "source", "") + ) + raise ConfigurationError( + f"no trading days in [{start}, {end}] for source " + f"{source_name!r}; the backtest window has no data" + ) self._index = 0 def __iter__(self) -> Iterator[str]: diff --git a/src/hqbacktest/engine/metrics.py b/src/hqbacktest/engine/metrics.py index ead9095..799684f 100644 --- a/src/hqbacktest/engine/metrics.py +++ b/src/hqbacktest/engine/metrics.py @@ -1,16 +1,36 @@ -"""Performance metrics (task 10). +"""Performance metrics (task 10 + task 17). Formulas (all documented in README and contract doc): * `total_return` = (final_equity / initial_equity) - 1 - * `daily_return` = equity[t] / equity[t-1] - 1 (simple) - * `annualized_return` = (1 + total_return) ** (trading_days / + * `daily_return` = equity[t] / equity[t-1] - 1 (t >= 1) + The engine anchors t=0 to `initial_cash` + so a first-day P&L flows into the + return series (task 17). The chained- + product identity + `∏(1 + daily_return) == 1 + total_return` + therefore holds for any run. + * `annualized_return` = (1 + total_return) ** (n / annual_trading_days) - 1 - * `daily_volatility` = stdev(daily_returns) (sample stddev; ddof=1) + The exponent is computed as `float` + then re-encoded as `Decimal(str(...))` + so the ledger never sees a Decimal + built directly from `float`. + * `daily_volatility` = stdev(daily_returns[1:]) (sample, ddof=1) + `None` when fewer than 2 daily returns + are available (single-day run, or two + trading days with only one observed + return). Reports `0` only when the + series is genuinely flat. * `annualized_volatility` = daily_volatility * sqrt(annual_trading_days) + `None` iff `daily_volatility is None`. * `sharpe_ratio` = (annualized_return - risk_free_rate) / annualized_volatility + `None` whenever `annualized_volatility` + is `None` or zero (zero-volatility note). * `max_drawdown` = max(peak - current) / peak over the - whole equity curve + whole equity curve. The peak sequence + starts at `initial_cash` so first-day + drawdowns contribute. * `turnover` = (sum(buy value) + sum(sell value)) / 2 / initial_equity (one-sided average) * `trade_count` = number of `Fill` records @@ -18,12 +38,14 @@ cost / total SELL fills; `None` if no SELL fills -Edge cases (per task 10 verification "空回测 / 单日 / 零波动 / 全亏损 / 无交易"): +Edge cases (per task 10/17 verification "空回测 / 单日 / 零波动 / 全亏损 / +无交易 / 样本不足"): * `len(equity_curve) == 0` (empty run): all metrics 0 or `None`; notes record "no trading days". * `len(equity_curve) == 1` (single day): `total_return` is the only computable ratio; annualised / vol / sharpe are `None` and a note is added. + * `daily_volatility` requires >= 2 daily returns; otherwise `None`. * `daily_volatility == 0` (flat equity): `sharpe_ratio` is `None`; a note records the reason. * `no SELL fills`: `win_rate` is `None`. @@ -43,6 +65,11 @@ from ..domain.fill import Fill from ..domain.enums import Side +# Quantization used when re-encoding `float` results as `Decimal`. 12 +# decimal places comfortably exceed the precision any real-data +# metric reaches and keeps `summary.json` clean. +_METRIC_QUANT = Decimal("0.000000000001") + @dataclass(frozen=True) class MetricsConfig: @@ -135,7 +162,24 @@ def compute_metrics( initial_cash: Decimal, config: MetricsConfig, ) -> PerformanceMetrics: - """Compute all v0.1 metrics from an equity curve and the fill list.""" + """Compute all v0.1 metrics from an equity curve and the fill list. + + Task 17 invariants: + * `daily_return` is recomputed from `total_equity` via + `_daily_returns`, which anchors the first day's return to + `initial_cash` (engine writes the same value to the + `EquityPoint`). The chained-product identity therefore + holds regardless of how the engine seeded day 0. + * `daily_volatility` is `None` whenever fewer than 2 daily + returns are available (single-day run, or a two-day run that + has only one observed return). It is `0` only when the series + is genuinely flat — task 17 forbids returning `0` for + undefined statistics. + * All Decimal metrics that involve `float` arithmetic go + through `Decimal(str(...))` so the ledger never holds a + `Decimal` constructed directly from a binary float (contract + rule 5). + """ notes: List[str] = [] n_days = len(equity_curve) final_equity = equity_curve[-1].total_equity if n_days else initial_cash @@ -159,23 +203,34 @@ def compute_metrics( annualized_return = None notes.append("annualized_return: total return <= -100%") else: - # Use float for the power operation; cast back to Decimal. - exponent = Decimal(n_days) / Decimal(config.annual_trading_days) - annualized_return = Decimal(float(growth) ** float(exponent)) - Decimal("1") + # Compute the power via float (Decimal has no built-in + # exponentiation), then re-encode through str() so the + # resulting Decimal never directly inherits binary-float + # bits. Quantize to a fixed precision so summary.json stays + # clean (task 17). + exponent = float(n_days) / float(config.annual_trading_days) + annualized_return = Decimal(str(float(growth) ** exponent)).quantize( + _METRIC_QUANT + ) - Decimal("1") # Daily volatility (per-day stdev) and its annualisation; Sharpe uses - # the annualised pair so the units match. + # the annualised pair so the units match. Task 17: insufficient + # samples (< 2 daily returns) returns `None`, not 0. returns = _daily_returns([pt.total_equity for pt in equity_curve]) annualized_volatility: Optional[Decimal] - if n_days < 2: - daily_volatility = None + if len(returns) - 1 < 2: + # returns[0] is the seed (0); subsequent entries are the actual + # daily returns. < 2 means stdev cannot be computed. + daily_volatility: Optional[Decimal] = None annualized_volatility = None - sharpe_ratio = None + sharpe_ratio: Optional[Decimal] = None notes.append("daily_volatility: requires >= 2 daily returns") else: try: vol_per_day = Decimal(str(stdev([float(r) for r in returns[1:]]))) except StatisticsError: + # All-zero series: stdev is undefined in `statistics` for + # 0-variance; treat as zero-volatility (a defined value). vol_per_day = Decimal("0") daily_volatility = vol_per_day if vol_per_day == 0: diff --git a/src/hqbacktest/engine/scheduler.py b/src/hqbacktest/engine/scheduler.py index 11a5065..c4180c1 100644 --- a/src/hqbacktest/engine/scheduler.py +++ b/src/hqbacktest/engine/scheduler.py @@ -26,7 +26,9 @@ # Callback the engine plugs in to consume pending orders at OPEN_MATCH. -OpenMatchCallback = Callable[[str, List[Order]], None] +# `context` is passed so the engine can drain out-of-universe rejections +# alongside regular pending orders (task 18). +OpenMatchCallback = Callable[[str, List[Order], "Context"], None] # Each phase's `visible_through` rule (contract §4). A value of `None` means @@ -77,20 +79,27 @@ def build_view( schedule: PhaseSchedule, today: str, ) -> DataView: - """Construct a `DataView` whose `visible_through` matches the phase rule.""" + """Construct a `DataView` whose `visible_through` matches the phase rule. + + `universe_start` is deliberately left unset here: it denotes the earliest + date the strategy may query (a run-level bound derived from the backtest + window), not the phase visibility. Bounding `history(bar_count=N)`'s + lookback window is `DataView`'s own responsibility (task 14), independent + of per-phase visibility. + """ if schedule.visible_through_mode == SAME_DAY_VISIBLE_THROUGH: - visible_through = today - elif schedule.visible_through_mode == PRE_BAR_VISIBLE_THROUGH: + return DataView(portal=portal, visible_through=today) + if schedule.visible_through_mode == PRE_BAR_VISIBLE_THROUGH: prev = previous_trading_day(portal, today) - # On the first trading day no history exists yet; the sentinel keeps - # the view legal but restricts reads to dates < today, so the - # strategy simply sees no bars. - visible_through = prev if prev is not None else NO_HISTORY_VISIBLE_THROUGH - else: # pragma: no cover - defensive - raise ValueError( - f"unknown visible_through mode: {schedule.visible_through_mode}" - ) - return DataView(portal=portal, visible_through=visible_through) + if prev is None: + # First trading day: no history exists yet. The sentinel keeps + # the view legal but exposes no data. + return DataView( + portal=portal, + visible_through=NO_HISTORY_VISIBLE_THROUGH, + ) + return DataView(portal=portal, visible_through=prev) + raise ValueError(f"unknown visible_through mode: {schedule.visible_through_mode}") def run_day( @@ -137,8 +146,14 @@ def run_day( # Only consume when a matcher is wired in; otherwise pending # orders would silently vanish without any event. pending = context._consume_pending_orders() - if pending: - on_open_match(today, pending) + # Task 18: invoke the matcher whenever there is anything to + # process — pending orders OR out-of-universe rejections to + # fold into the audit table. The engine drains both in a + # single call; do NOT pre-consume the out-of-universe list + # here (the engine must consume it, or those rejections + # would be silently dropped from the orders table). + if pending or context._has_out_of_universe_orders(): + on_open_match(today, pending, context) continue elif entry.phase is EventType.BAR_CLOSE: strategy.on_bar(context, view) diff --git a/src/hqbacktest/engine/strategy.py b/src/hqbacktest/engine/strategy.py index c47027d..d30ca8f 100644 --- a/src/hqbacktest/engine/strategy.py +++ b/src/hqbacktest/engine/strategy.py @@ -14,7 +14,7 @@ for explicit `super().initialize(context)` patterns. """ -from typing import Any, Protocol, runtime_checkable +from typing import Any, Dict, Protocol, runtime_checkable from ..domain.enums import EventType from .events import EngineEvent @@ -58,8 +58,22 @@ class BaseStrategy: only after `initialize`; * raises `StrategyLifecycleError` if the strategy mis-uses the API (e.g. calling `context.order(...)` outside a callback). + + Task 21: the constructor accepts and stores `**kwargs` so the CLI + can pass user-supplied parameters via `[strategy].kwargs`. + Subclasses that need parameters should accept them in their own + `__init__` (the base ``**kwargs`` is swallowed so no TypeError + on stray keyword args). """ + def __init__(self, **kwargs: Any) -> None: + # Store kwargs so tests / introspection can find them; the + # subclass's own __init__ is responsible for consuming the + # parameters it actually needs. Subclasses that override + # __init__ should still call ``super().__init__(**kwargs)`` (or + # accept `**kwargs` themselves) to remain CLI-compatible. + self.kwargs: Dict[str, Any] = dict(kwargs) + # ------------------------------------------------------------------ # # Optional user-facing universe declaration # ------------------------------------------------------------------ # diff --git a/tests/cli/test_task20_cli.py b/tests/cli/test_task20_cli.py new file mode 100644 index 0000000..cb898c7 --- /dev/null +++ b/tests/cli/test_task20_cli.py @@ -0,0 +1,455 @@ +"""Task 20: CLI first-mile + README honesty regression tests. + +Covers: + * Console-script (`hqbacktest run`) works from a fresh working + directory using only the console-script entry point, NOT + `python -m hqbacktest run`. + * `initial_cash = nan` / `inf` / negative yields a ConfigError + (CLI exit code 2, single-line stderr). + * `start_date = 20241399` (impossible date) yields a ConfigError. + * A backtest window with zero trading days raises a ConfigError + (exit 2), not a silent empty result. + * Output directory that already contains prior-run files fails + with exit 3 unless `--force` is given. + * `Context.order_value(symbol, 15000)` accepts `int` and `str` + cash amounts (not only Decimal) without raising. + * `run_metadata.git_commit` reflects the hqbacktest package's own + commit (not the user's cwd repository). +""" + +from __future__ import annotations + +import json +import os +import subprocess +import sys +from decimal import Decimal +from pathlib import Path + +import pytest + +from hqbacktest import BacktestConfig, BacktestEngine, BaseStrategy +from hqbacktest.cli.config import ConfigError, load_config_file +from hqbacktest.cli.runner import _git_commit, run_from_file +from hqbacktest.data import InMemoryDataPortal +from hqbacktest.domain.bar import Bar + + +# --------------------------------------------------------------------------- +# Fixture helpers +# --------------------------------------------------------------------------- + + +def _bar(date: str, sym: str = "600000.SH") -> Bar: + return Bar.from_raw( + symbol=sym, + date=date, + open="10.0000", + high="30.0000", + low="5.0000", + close="10.0000", + volume=1000, + ) + + +def _memory_portal() -> InMemoryDataPortal: + p = InMemoryDataPortal( + calendar=["20240102", "20240103", "20240104"], + universe_by_date={"20240102": ["600000.SH"]}, + as_of="20240104", + ) + for d in ("20240102", "20240103", "20240104"): + p.add_bar(_bar(d)) + return p + + +def _minimal_config(output_dir: Path, strategy_module: str = "strategy") -> str: + return f"""[start] +start_date = '20240102' +end_date = '20240104' +[capital] +initial_cash = '100000' +[data] +source = 'memory' +[strategy] +module = '{strategy_module}' +[output] +directory = '{output_dir}' +""" + + +def _write_strategy_module(tmp: Path) -> Path: + p = tmp / "strategy.py" + p.write_text( + "from hqbacktest import BaseStrategy\n" + "class BuyHold(BaseStrategy):\n" + " def initialize(self, context):\n" + " context.set_universe(['600000.SH'])\n" + " def on_bar(self, context, data):\n" + " if context.now == '20240102':\n" + " context.order('600000.SH', 100)\n" + ) + return p + + +# --------------------------------------------------------------------------- +# Console script first-mile +# --------------------------------------------------------------------------- + + +def test_console_script_resolves_strategy_from_cwd(tmp_path): + """`hqbacktest run` (the console script) must be able to find a + strategy module colocated with the config file in a fresh + working directory. + + We exercise the actual `hqbacktest` console entry point (not + `python -m hqbacktest run`) and verify that the run succeeds. + The portal is monkey-patched inside the subprocess via a + `sitecustomize`-style bootstrap script written to a temp + directory that the subprocess adds to PYTHONPATH. + """ + import shutil + import textwrap + + workdir = tmp_path / "workdir" + workdir.mkdir() + _write_strategy_module(workdir) + config_file = workdir / "config.toml" + config_file.write_text(_minimal_config(workdir / "out")) + + # Bootstrap that swaps `_resolve_portal` for the memory portal. + # The CLI imports this module by name via `HQBACKTEST_CLI_BOOTSTRAP`. + bootstrap = workdir / "_cli_bootstrap.py" + bootstrap.write_text( + textwrap.dedent( + """ + import hqbacktest.cli.runner as _r + from hqbacktest.data import InMemoryDataPortal + from hqbacktest.domain.bar import Bar + _P = InMemoryDataPortal( + calendar=['20240102', '20240103', '20240104'], + universe_by_date={'20240102': ['600000.SH']}, + as_of='20240104', + ) + for _d in ('20240102', '20240103', '20240104'): + _P.add_bar(Bar.from_raw( + symbol='600000.SH', date=_d, open='10', + high='30', low='5', close='10', volume=1000, + )) + _r._resolve_portal = lambda source, data_root: _P + """ + ).strip() + + "\n" + ) + + hqbacktest_bin = shutil.which("hqbacktest") + if hqbacktest_bin is None: + argv = [sys.executable, "-m", "hqbacktest", "run"] + else: + argv = [hqbacktest_bin, "run"] + env = { + **os.environ, + "PYTHONPATH": str(workdir), # so the bootstrap is importable + "HQBACKTEST_CLI_BOOTSTRAP": "_cli_bootstrap", + } + result = subprocess.run( + argv + ["--config", str(config_file)], + capture_output=True, + text=True, + cwd=str(workdir), + check=False, + env=env, + ) + assert result.returncode == 0, ( + f"console script failed (exit={result.returncode}):\n" + f"stdout={result.stdout}\nstderr={result.stderr}" + ) + out = workdir / "out" + assert out.exists() + assert (out / "summary.json").exists() + + +# --------------------------------------------------------------------------- +# Config validation: nan / inf / negative / bad date / empty window +# --------------------------------------------------------------------------- + + +def test_initial_cash_nan_rejected(tmp_path): + cfg = tmp_path / "c.toml" + cfg.write_text( + "[start]\n" + "start_date = '20240102'\n" + "end_date = '20240104'\n" + "[capital]\n" + "initial_cash = 'nan'\n" + "[data]\n" + "source = 'memory'\n" + "[strategy]\n" + "module = 'strategy'\n" + "[output]\n" + f"directory = '{tmp_path / 'out'}'\n" + ) + with pytest.raises(ConfigError, match="nan|valid number|number"): + load_config_file(str(cfg)) + + +def test_initial_cash_inf_rejected(tmp_path): + cfg = tmp_path / "c.toml" + cfg.write_text( + "[start]\n" + "start_date = '20240102'\n" + "end_date = '20240104'\n" + "[capital]\n" + "initial_cash = 'inf'\n" + "[data]\n" + "source = 'memory'\n" + "[strategy]\n" + "module = 'strategy'\n" + "[output]\n" + f"directory = '{tmp_path / 'out'}'\n" + ) + with pytest.raises(ConfigError): + load_config_file(str(cfg)) + + +def test_initial_cash_negative_rejected(tmp_path): + """`_require_decimal` already enforces `min_value=0`, so a + negative literal must surface as ConfigError.""" + cfg = tmp_path / "c.toml" + cfg.write_text( + "[start]\n" + "start_date = '20240102'\n" + "end_date = '20240104'\n" + "[capital]\n" + "initial_cash = '-100'\n" + "[data]\n" + "source = 'memory'\n" + "[strategy]\n" + "module = 'strategy'\n" + "[output]\n" + f"directory = '{tmp_path / 'out'}'\n" + ) + with pytest.raises(ConfigError): + load_config_file(str(cfg)) + + +def test_start_date_impossible_rejected(tmp_path): + """An impossible calendar date (`20241399`) must be rejected.""" + cfg = tmp_path / "c.toml" + cfg.write_text( + "[start]\n" + "start_date = '20241399'\n" + "end_date = '20240104'\n" + "[capital]\n" + "initial_cash = '100000'\n" + "[data]\n" + "source = 'memory'\n" + "[strategy]\n" + "module = 'strategy'\n" + "[output]\n" + f"directory = '{tmp_path / 'out'}'\n" + ) + with pytest.raises(ConfigError): + load_config_file(str(cfg)) + + +def test_engine_window_with_zero_trading_days_raises(tmp_path): + """An engine window with no trading days must raise rather than + silently write an empty result. + """ + cfg = BacktestConfig( + start_date="20240102", + end_date="20240104", + initial_cash=Decimal("100000"), + source="tushare", + ) + + class S(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + p = InMemoryDataPortal( + calendar=[], # empty -> no trading days + universe_by_date={}, + as_of="20991231", + ) + from hqbacktest.engine.errors import ConfigurationError + + # The empty-window check fires at iterator construction time, so + # we get a typed ConfigurationError (mapped to CLI exit 2 by + # `__main__`) rather than a RunFailed mid-run. + with pytest.raises(ConfigurationError, match="no trading days"): + BacktestEngine(cfg, strategy=S(), portal=p).run() + + +# --------------------------------------------------------------------------- +# Output directory reuse: exit 3 unless --force +# --------------------------------------------------------------------------- + + +def test_output_dir_with_prior_files_rejected(tmp_path): + """When the output directory already contains prior-run files, + the runner must NOT silently overwrite them (exit 3). + """ + portal = _memory_portal() + from hqbacktest.cli import runner + + strategy = tmp_path / "strategy.py" + strategy.write_text( + "from hqbacktest import BaseStrategy\n" + "class S(BaseStrategy):\n" + " def initialize(self, context):\n" + " context.set_universe(['600000.SH'])\n" + ) + out_dir = tmp_path / "out" + out_dir.mkdir() + (out_dir / "summary.json").write_text('{"old": true}') + cfg = tmp_path / "c.toml" + cfg.write_text(_minimal_config(out_dir)) + original = runner._resolve_portal + runner._resolve_portal = lambda source, data_root: portal + try: + result = run_from_file(str(cfg)) + finally: + runner._resolve_portal = original + assert result.exit_code == 3 + assert "prior" in result.message.lower() or "exists" in result.message.lower() + # The prior summary.json must be untouched. + assert (out_dir / "summary.json").read_text() == '{"old": true}' + + +# --------------------------------------------------------------------------- +# Context.order_value: int / str accepted +# --------------------------------------------------------------------------- + + +def test_order_value_accepts_int_cash(): + """`order_value(symbol, 15000)` (int) must not raise — int is a + valid monetary literal at the strategy-API layer.""" + + class Spend(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20240102": + context.order_value("600000.SH", 15000) + + cfg = BacktestConfig( + start_date="20240102", + end_date="20240103", + initial_cash=Decimal("100000"), + source="tushare", + ) + engine = BacktestEngine(cfg, strategy=Spend(), portal=_memory_portal()) + engine.run() # must not raise + + +def test_order_value_accepts_str_cash(): + """`order_value(symbol, '15000')` (str) must not raise either.""" + + class Spend(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20240102": + context.order_value("600000.SH", "15000") + + cfg = BacktestConfig( + start_date="20240102", + end_date="20240103", + initial_cash=Decimal("100000"), + source="tushare", + ) + engine = BacktestEngine(cfg, strategy=Spend(), portal=_memory_portal()) + engine.run() # must not raise + + +def test_order_target_value_accepts_int_cash(): + """`order_target_value(symbol, 15000)` (int) must also accept an + int/str monetary literal (task 20), not only `Decimal`.""" + + class TargetSpend(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20240102": + context.order_target_value("600000.SH", 15000) + + cfg = BacktestConfig( + start_date="20240102", + end_date="20240103", + initial_cash=Decimal("100000"), + source="tushare", + ) + engine = BacktestEngine(cfg, strategy=TargetSpend(), portal=_memory_portal()) + engine.run() # must not raise + + +def test_empty_window_cli_returns_exit_2(tmp_path): + """An empty trading-day window must surface as exit code 2 + (configuration error), not exit 4 (run failure).""" + from hqbacktest.cli import runner + + empty = InMemoryDataPortal(calendar=[], universe_by_date={}, as_of="20991231") + strategy = tmp_path / "strategy.py" + strategy.write_text( + "from hqbacktest import BaseStrategy\n" + "class S(BaseStrategy):\n" + " def initialize(self, context):\n" + " context.set_universe(['600000.SH'])\n" + ) + out_dir = tmp_path / "out" + cfg = tmp_path / "c.toml" + cfg.write_text(_minimal_config(out_dir)) + original = runner._resolve_portal + runner._resolve_portal = lambda source, data_root: empty + try: + result = run_from_file(str(cfg)) + finally: + runner._resolve_portal = original + assert result.exit_code == 2, result.message + assert "no trading days" in result.message + + +# --------------------------------------------------------------------------- +# git_commit semantics +# --------------------------------------------------------------------------- + + +def test_git_commit_reports_package_commit_or_none(): + """`_git_commit()` must not raise and returns either a short hex + string or `None`. It is documented as a best-effort lookup + (task 20). + """ + result = _git_commit() + assert result is None or (isinstance(result, str) and len(result) >= 4) + + +def test_run_metadata_records_package_version(tmp_path): + """`run_metadata.json` must record the hqbacktest package + version and the configured start/end dates (task 20: rename + `git_commit` semantics).""" + from hqbacktest.cli import runner + from hqbacktest.cli.config import load_config_file + + portal = _memory_portal() + out = tmp_path / "out" + cfg_file = tmp_path / "c.toml" + _write_strategy_module(tmp_path) + cfg_file.write_text(_minimal_config(out)) + original = runner._resolve_portal + runner._resolve_portal = lambda source, data_root: portal + try: + result = run_from_file(str(cfg_file)) + finally: + runner._resolve_portal = original + assert result.exit_code == 0, result.message + meta = json.loads((result.output_dir / "run_metadata.json").read_text()) + from hqbacktest import __version__ as HQ_VER + + assert meta["hqbacktest_version"] == HQ_VER + assert meta["config_start_date"] == "20240102" + assert meta["config_end_date"] == "20240104" diff --git a/tests/data/test_data_view.py b/tests/data/test_data_view.py index 06fe1e9..245c106 100644 --- a/tests/data/test_data_view.py +++ b/tests/data/test_data_view.py @@ -110,7 +110,13 @@ class BrokenPortal(InMemoryDataPortal): def get_bars(self, symbol, start, end): raise InvalidDataError("bars", "malformed source row") - view = DataView(portal=BrokenPortal(), visible_through="20240102") + # The portal must have at least one calendar day for the lookback + # window to actually invoke get_bars (else current_price returns + # None without touching the data layer). + view = DataView( + portal=BrokenPortal(calendar=["20240102"]), + visible_through="20240102", + ) with pytest.raises(InvalidDataError, match="malformed"): view.current_price("600000.SH") diff --git a/tests/data/test_hqdata_portal.py b/tests/data/test_hqdata_portal.py index 9928648..826e812 100644 --- a/tests/data/test_hqdata_portal.py +++ b/tests/data/test_hqdata_portal.py @@ -278,9 +278,11 @@ def test_calendar_rejects_invalid_is_open(tmp_path): snap = tmp_path / "tushare" snap.mkdir() (snap / "calendar.csv").write_text("date,is_open\n20240102,X\n", encoding="utf-8") - portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + # Per task 14, a corrupt calendar is an infrastructure failure: the + # portal surfaces it at construction time (via `_resolve_as_of`) so + # the engine can never publish a misleading `data_version`. with pytest.raises(InvalidDataError): - portal.get_calendar("20240101", "20240110") + HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) # --------------------------------------------------------------------- # @@ -463,11 +465,22 @@ def test_get_bars_caches(tmp_path): portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) a = portal.get_bars("600000.SH", "20240102", "20240102") b = portal.get_bars("600000.SH", "20240102", "20240102") - assert a is b # served from cache + # Task 14: cached lists are returned as defensive copies, never the + # internal reference. Strategies must not be able to mutate the + # cache by mutating a returned list. + assert a == b + assert a is not b def test_get_bars_rejects_missing_daily_file(tmp_path): - """A trading day without a daily file raises MissingDataError.""" + """A trading day without a daily file raises SnapshotFileMissingError. + + Per task 14, a missing whole-day file is an infrastructure failure + (distinct from a per-symbol gap) and must propagate so the engine can + abort the run with `DATA_ERROR`. + """ + from hqbacktest.data import SnapshotFileMissingError + _build_snapshot( tmp_path, "tushare", @@ -480,9 +493,9 @@ def test_get_bars_rejects_missing_daily_file(tmp_path): factors={}, ) portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) - with pytest.raises(MissingDataError) as exc: + with pytest.raises(SnapshotFileMissingError) as exc: portal.get_bars("600000.SH", "20240102", "20240103") - assert "no bars" in str(exc.value).lower() + assert "20240103" in str(exc.value) def test_get_bars_rejects_date_mismatch_in_daily(tmp_path): @@ -500,6 +513,13 @@ def test_get_bars_rejects_date_mismatch_in_daily(tmp_path): def test_get_bars_rejects_wrong_symbol_row(tmp_path): + """A per-symbol gap returns `[]`, not an error (task 14). + + The daily file exists for the trading day but contains no row for + the requested symbol (suspended / delisted / pre-IPO). This is a + per-symbol gap and is a normal business outcome, so the portal + returns an empty list rather than raising. + """ snap = tmp_path / "tushare" snap.mkdir() (snap / "stock_daily").mkdir() @@ -510,9 +530,7 @@ def test_get_bars_rejects_wrong_symbol_row(tmp_path): ) _write_calendar(snap, [("20240102", "Y")]) portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) - with pytest.raises(MissingDataError): - # The file exists but has no row for 600000.SH. - portal.get_bars("600000.SH", "20240102", "20240102") + assert portal.get_bars("600000.SH", "20240102", "20240102") == [] def test_get_bars_rejects_duplicate_symbol_rows(tmp_path): @@ -527,7 +545,7 @@ def test_get_bars_rejects_duplicate_symbol_rows(tmp_path): ) _write_calendar(snap, [("20240102", "Y")]) portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) - with pytest.raises(InvalidDataError, match="at most one row"): + with pytest.raises(InvalidDataError, match="duplicate row"): portal.get_bars("600000.SH", "20240102", "20240102") diff --git a/tests/data/test_memory_portal.py b/tests/data/test_memory_portal.py index 2af9f14..0347c70 100644 --- a/tests/data/test_memory_portal.py +++ b/tests/data/test_memory_portal.py @@ -8,6 +8,7 @@ InMemoryDataPortal, InvalidDataError, MissingDataError, + SnapshotFileMissingError, ) from hqbacktest.domain.bar import Bar @@ -85,11 +86,11 @@ def test_get_bars_rejects_window_start_after_end(): p.get_bars("600000.SH", "20240105", "20240102") -def test_get_bars_raises_when_window_empty(): +def test_get_bars_returns_empty_when_window_empty(): + """Task 14: a window with no bars returns [] (per-symbol gap semantics).""" p = InMemoryDataPortal() p.add_bar(_bar("600000.SH", "20240102")) - with pytest.raises(MissingDataError): - p.get_bars("600000.SH", "20240201", "20240205") + assert p.get_bars("600000.SH", "20240201", "20240205") == [] def test_get_bars_rejects_invalid_symbol(): @@ -98,21 +99,33 @@ def test_get_bars_rejects_invalid_symbol(): p.get_bars("not-a-symbol", "20240102", "20240105") -def test_universe_walks_back_to_latest_snapshot(): +def test_universe_exact_date_only_no_walk_back(): + """Task 14: per-date snapshot semantics; no implicit walk-back. + + The in-memory portal must match the production CSV portal which only + looks up the snapshot for the exact requested date. This avoids the + silent forward-fallback that violated contract §4 (stock list must be + queried per backtest day, not by walking back to the most recent + snapshot). + """ p = InMemoryDataPortal( universe_by_date={ "20240102": ["600000.SH", "000001.SZ"], "20240105": ["600000.SH", "000002.SZ", "688001.SH"], } ) - assert p.get_universe("20240103") == ["000001.SZ", "600000.SH"] + assert p.get_universe("20240102") == ["000001.SZ", "600000.SH"] assert p.get_universe("20240105") == ["000002.SZ", "600000.SH", "688001.SH"] - assert p.get_universe("20240110") == ["000002.SZ", "600000.SH", "688001.SH"] + # Walk-back is no longer supported. + with pytest.raises(SnapshotFileMissingError): + p.get_universe("20240103") + with pytest.raises(SnapshotFileMissingError): + p.get_universe("20240110") def test_universe_raises_when_no_snapshot_exists(): p = InMemoryDataPortal(universe_by_date={"20240102": ["600000.SH"]}) - with pytest.raises(MissingDataError): + with pytest.raises(SnapshotFileMissingError): p.get_universe("20200101") @@ -189,21 +202,22 @@ def test_constructor_rejects_initial_bar_outside_known_calendar(): ) -def test_missing_bar_on_trading_day_raises_missing_data(): - """停牌/缺行: a window that is entirely empty (no bar for any of its - trading days) must surface as MissingDataError.""" +def test_missing_bar_on_trading_day_returns_empty_list(): + """Task 14: a fully empty window returns [], not MissingDataError. + + "停牌/缺行" is a per-symbol gap, a normal business outcome. The + portal surfaces it as an empty list so the caller (engine, DataView) + can decide the policy. + """ p = InMemoryDataPortal(calendar=["20240102", "20240103", "20240104"]) - # No bars at all -> MissingDataError on any query in this window. - with pytest.raises(MissingDataError) as exc: - p.get_bars("600000.SH", "20240102", "20240103") - assert "no bars" in str(exc.value).lower() + # No bars at all -> empty list, not an error. + assert p.get_bars("600000.SH", "20240102", "20240103") == [] -def test_window_with_only_suspended_days_raises(): +def test_window_with_only_suspended_days_returns_empty(): p = InMemoryDataPortal(calendar=["20240102", "20240103"]) # No bars at all on a valid trading-day window. - with pytest.raises(MissingDataError): - p.get_bars("600000.SH", "20240102", "20240103") + assert p.get_bars("600000.SH", "20240102", "20240103") == [] def test_partial_window_returns_available_bars(): diff --git a/tests/data/test_portal_parity.py b/tests/data/test_portal_parity.py new file mode 100644 index 0000000..a09bd98 --- /dev/null +++ b/tests/data/test_portal_parity.py @@ -0,0 +1,536 @@ +"""Parity tests between InMemoryDataPortal and HqDataCsvPortal. + +The two portals must agree on every observable behavior for the same fixture +data. Per task 14 of TODO.md, any divergence means tests can pass on memory +data while production silently misbehaves on CSV snapshots. Each test in this +file constructs equivalent fixtures for both portals and asserts identical +return values and identical exception types. +""" + +from __future__ import annotations + +from datetime import date as _date +from decimal import Decimal +from pathlib import Path + +import pytest + +from hqbacktest.data import ( + HqDataCsvPortal, + InMemoryDataPortal, + InvalidDataError, + MissingDataError, + SnapshotFileMissingError, + UnknownSymbolError, +) +from hqbacktest.data.hqdata_portal import resolve_source_location +from hqbacktest.domain.bar import Bar + + +# --------------------------------------------------------------------------- +# Fixture builders +# --------------------------------------------------------------------------- + + +# Trading calendar with a weekend-style gap: 20240102-20240105 are trading +# days, 20240106 (Saturday) is excluded. +CALENDAR_DATES: list[tuple[str, str]] = [ + ("20240102", "Y"), + ("20240103", "Y"), + ("20240104", "Y"), + ("20240105", "Y"), +] +LATE_CALENDAR_DATES: list[tuple[str, str]] = [ + ("20240102", "Y"), + ("20240103", "Y"), + ("20240104", "Y"), + ("20240105", "Y"), + ("20240108", "Y"), + ("20240109", "Y"), + ("20240110", "Y"), + ("20240111", "Y"), + ("20240112", "Y"), +] + + +def _make_bar(symbol: str, date: str, close: str = "10.00") -> Bar: + # Wide OHLC envelope so any close in [9, 30] is valid. + return Bar.from_raw( + symbol=symbol, + date=date, + open="10.00", + high="30.00", + low="9.00", + close=close, + volume=1000, + ) + + +def _memory_with_gaps() -> InMemoryDataPortal: + """SUSPENDED on 20240103; never-listed symbol 999999.SH.""" + p = InMemoryDataPortal( + calendar=[d for d, f in CALENDAR_DATES], + universe_by_date={"20240102": ["600000.SH", "000001.SZ"]}, + as_of="20240105", + ) + # 600000.SH: traded on 20240102, suspended 20240103-04, traded 20240105. + p.add_bar(_make_bar("600000.SH", "20240102", "10.00")) + p.add_bar(_make_bar("600000.SH", "20240105", "11.00")) + # 000001.SZ: traded every day. + p.add_bar(_make_bar("000001.SZ", "20240102", "20.00")) + p.add_bar(_make_bar("000001.SZ", "20240103", "20.50")) + p.add_bar(_make_bar("000001.SZ", "20240104", "20.25")) + p.add_bar(_make_bar("000001.SZ", "20240105", "21.00")) + return p + + +def _write_calendar(root: Path, rows: list[tuple[str, str]]) -> None: + lines = ["date,is_open"] + for d, f in rows: + lines.append(f"{d},{f}") + (root / "calendar.csv").write_text("\n".join(lines) + "\n", encoding="utf-8") + + +def _write_stock_list(root: Path, date: str, symbols: list[str]) -> None: + target = root / "stock_list" + target.mkdir(parents=True, exist_ok=True) + lines = ["symbol,date,name,exchange,board,curr_type,list_date,delist_date"] + for sym in symbols: + lines.append(f"{sym},{date},name,SSE,MB,CNY,19990101,") + (target / f"{date}.csv").write_text("\n".join(lines) + "\n", encoding="utf-8") + + +def _write_stock_daily(root: Path, date: str, rows: list[dict]) -> None: + target = root / "stock_daily" + target.mkdir(parents=True, exist_ok=True) + fields = [ + "symbol", + "date", + "pre_close", + "open", + "high", + "low", + "close", + "volume", + "turnover", + "change", + "pct_change", + ] + lines = [",".join(fields)] + for r in rows: + lines.append(",".join(str(r[f]) for f in fields)) + (target / f"{date}.csv").write_text("\n".join(lines) + "\n", encoding="utf-8") + + +def _csv_with_gaps(tmp_path: Path) -> HqDataCsvPortal: + """Same fixture as `_memory_with_gaps` but on disk. + + Layout: + - stock_daily/20240102.csv has both 600000.SH and 000001.SZ. + - stock_daily/20240103.csv only has 000001.SZ (600000.SH suspended). + - stock_daily/20240104.csv only has 000001.SZ (600000.SH still suspended). + - stock_daily/20240105.csv has both 600000.SH and 000001.SZ. + """ + 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"]) + _write_stock_daily( + snap, + "20240102", + [ + { + "symbol": "600000.SH", + "date": "20240102", + "pre_close": 10, + "open": 10, + "high": 11, + "low": 9, + "close": 10, + "volume": 1000, + "turnover": 10000, + "change": 0, + "pct_change": 0, + }, + { + "symbol": "000001.SZ", + "date": "20240102", + "pre_close": 20, + "open": 20, + "high": 21, + "low": 19, + "close": 20, + "volume": 1000, + "turnover": 20000, + "change": 0, + "pct_change": 0, + }, + ], + ) + _write_stock_daily( + snap, + "20240103", + [ + { + "symbol": "000001.SZ", + "date": "20240103", + "pre_close": 20, + "open": 20, + "high": 21, + "low": 19, + "close": 20.5, + "volume": 1000, + "turnover": 20500, + "change": 0.5, + "pct_change": 2.5, + } + ], + ) + _write_stock_daily( + snap, + "20240104", + [ + { + "symbol": "000001.SZ", + "date": "20240104", + "pre_close": 20.5, + "open": 20, + "high": 21, + "low": 19, + "close": 20.25, + "volume": 1000, + "turnover": 20250, + "change": -0.25, + "pct_change": -1.22, + } + ], + ) + _write_stock_daily( + snap, + "20240105", + [ + { + "symbol": "600000.SH", + "date": "20240105", + "pre_close": 10, + "open": 11, + "high": 11.5, + "low": 10.5, + "close": 11, + "volume": 1000, + "turnover": 11000, + "change": 1, + "pct_change": 10, + }, + { + "symbol": "000001.SZ", + "date": "20240105", + "pre_close": 20.25, + "open": 21, + "high": 21.5, + "low": 20.5, + "close": 21, + "volume": 1000, + "turnover": 21000, + "change": 0.75, + "pct_change": 3.7, + }, + ], + ) + return HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + + +# --------------------------------------------------------------------------- +# Calendar parity +# --------------------------------------------------------------------------- + + +def test_calendar_window_returns_same_open_dates(tmp_path): + mem = InMemoryDataPortal(calendar=[d for d, _ in CALENDAR_DATES], as_of="20240105") + csv = _csv_with_gaps(tmp_path) + assert mem.get_calendar("20240102", "20240105") == csv.get_calendar( + "20240102", "20240105" + ) + + +def test_is_trading_day_agrees(tmp_path): + mem = InMemoryDataPortal(calendar=[d for d, _ in CALENDAR_DATES], as_of="20240105") + csv = _csv_with_gaps(tmp_path) + for d, flag in CALENDAR_DATES: + assert mem.is_trading_day(d) == csv.is_trading_day(d) + assert csv.is_trading_day(d) is (flag == "Y") + + +def test_previous_and_next_trading_day_agrees(tmp_path): + mem = InMemoryDataPortal( + calendar=[d for d, _ in LATE_CALENDAR_DATES], as_of="20240112" + ) + snap = tmp_path / "tushare" + snap.mkdir() + _write_calendar(snap, LATE_CALENDAR_DATES) + csv = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + for d in ("20240105", "20240108", "20240111"): + assert mem.previous_trading_day(d) == csv.previous_trading_day(d) + assert mem.next_trading_day(d) == csv.next_trading_day(d) + + +def test_calendar_rejects_start_after_end(tmp_path): + mem = InMemoryDataPortal(calendar=[d for d, _ in CALENDAR_DATES], as_of="20240105") + csv = _csv_with_gaps(tmp_path) + with pytest.raises(InvalidDataError): + mem.get_calendar("20240105", "20240102") + with pytest.raises(InvalidDataError): + csv.get_calendar("20240105", "20240102") + + +# --------------------------------------------------------------------------- +# Universe parity +# --------------------------------------------------------------------------- + + +def test_universe_exact_date_agrees(tmp_path): + mem = _memory_with_gaps() + csv = _csv_with_gaps(tmp_path) + assert mem.get_universe("20240102") == csv.get_universe("20240102") + + +def test_universe_raises_on_missing_snapshot_for_both(tmp_path): + """Neither portal silently walks back to a prior snapshot (task 14). + + A missing whole-day stock-list snapshot is an infrastructure failure in + both portals, so the exception type must be identical + (`SnapshotFileMissingError`), not merely a shared base class. + """ + mem = _memory_with_gaps() + csv = _csv_with_gaps(tmp_path) + with pytest.raises(SnapshotFileMissingError): + mem.get_universe("20240106") + with pytest.raises(SnapshotFileMissingError): + csv.get_universe("20240106") + + +def test_universe_rejects_future_date(tmp_path): + """Both portals validate the date format.""" + mem = _memory_with_gaps() + csv = _csv_with_gaps(tmp_path) + with pytest.raises(InvalidDataError): + mem.get_universe("not-a-date") + with pytest.raises(InvalidDataError): + csv.get_universe("not-a-date") + + +def test_universe_excludes_bj_by_default(tmp_path): + """`.BJ` (Beijing Stock Exchange) symbols are excluded by default.""" + from hqbacktest.data import InMemoryDataPortal + + mem = InMemoryDataPortal( + calendar=["20240102"], + universe_by_date={"20240102": ["600000.SH", "830001.BJ", "000001.SZ"]}, + as_of="20240102", + ) + snap = tmp_path / "tushare" + snap.mkdir() + _write_calendar(snap, [("20240102", "Y")]) + _write_stock_list(snap, "20240102", ["600000.SH", "830001.BJ", "000001.SZ"]) + csv = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + assert mem.get_universe("20240102") == ["000001.SZ", "600000.SH"] + assert csv.get_universe("20240102") == ["000001.SZ", "600000.SH"] + + +def test_universe_includes_bj_when_requested(tmp_path): + """`include_bj=True` keeps `.BJ` symbols in the result.""" + from hqbacktest.data import InMemoryDataPortal + + mem = InMemoryDataPortal( + calendar=["20240102"], + universe_by_date={"20240102": ["600000.SH", "830001.BJ", "000001.SZ"]}, + as_of="20240102", + ) + snap = tmp_path / "tushare" + snap.mkdir() + _write_calendar(snap, [("20240102", "Y")]) + _write_stock_list(snap, "20240102", ["600000.SH", "830001.BJ", "000001.SZ"]) + csv = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + assert mem.get_universe("20240102", include_bj=True) == [ + "000001.SZ", + "600000.SH", + "830001.BJ", + ] + assert csv.get_universe("20240102", include_bj=True) == [ + "000001.SZ", + "600000.SH", + "830001.BJ", + ] + + +# --------------------------------------------------------------------------- +# Bars parity +# --------------------------------------------------------------------------- + + +def test_bars_returns_window_subset_allowing_gaps(tmp_path): + """A suspended symbol must return its actual bars, not raise.""" + mem = _memory_with_gaps() + csv = _csv_with_gaps(tmp_path) + mem_bars = mem.get_bars("600000.SH", "20240102", "20240105") + csv_bars = csv.get_bars("600000.SH", "20240102", "20240105") + assert [b.date for b in mem_bars] == [b.date for b in csv_bars] + assert [b.close for b in mem_bars] == [b.close for b in csv_bars] + assert [b.date for b in mem_bars] == ["20240102", "20240105"] + + +def test_bars_full_coverage_symbol_matches(tmp_path): + mem = _memory_with_gaps() + csv = _csv_with_gaps(tmp_path) + mem_bars = mem.get_bars("000001.SZ", "20240102", "20240105") + csv_bars = csv.get_bars("000001.SZ", "20240102", "20240105") + assert [b.date for b in mem_bars] == [b.date for b in csv_bars] + assert [str(b.close) for b in mem_bars] == [str(b.close) for b in csv_bars] + + +def test_bars_window_empty_when_never_listed(tmp_path): + """A symbol that never traded in the window returns empty, not raise.""" + mem = InMemoryDataPortal(calendar=[d for d, _ in CALENDAR_DATES], as_of="20240105") + csv = _csv_with_gaps(tmp_path) + assert mem.get_bars("999999.SH", "20240102", "20240105") == [] + assert csv.get_bars("999999.SH", "20240102", "20240105") == [] + + +def test_bars_rejects_window_start_after_end_for_both(tmp_path): + mem = _memory_with_gaps() + csv = _csv_with_gaps(tmp_path) + with pytest.raises(InvalidDataError): + mem.get_bars("000001.SZ", "20240105", "20240102") + with pytest.raises(InvalidDataError): + csv.get_bars("000001.SZ", "20240105", "20240102") + + +def test_bars_rejects_bad_symbol_for_both(tmp_path): + mem = _memory_with_gaps() + csv = _csv_with_gaps(tmp_path) + with pytest.raises(InvalidDataError): + mem.get_bars("not-a-symbol", "20240102", "20240105") + with pytest.raises(InvalidDataError): + csv.get_bars("not-a-symbol", "20240102", "20240105") + + +def test_bars_distinguishes_snapshot_missing_from_per_symbol_gap(tmp_path): + """整日快照缺失 must raise SnapshotFileMissingError, not MissingDataError. + + An individual symbol missing from an existing daily file returns []. + """ + snap = tmp_path / "tushare" + snap.mkdir() + _write_calendar( + snap, + [("20240102", "Y"), ("20240103", "Y")], + ) + _write_stock_list(snap, "20240102", ["600000.SH"]) + _write_stock_daily( + snap, + "20240102", + [ + { + "symbol": "600000.SH", + "date": "20240102", + "pre_close": 10, + "open": 10, + "high": 11, + "low": 9, + "close": 10, + "volume": 1000, + "turnover": 10000, + "change": 0, + "pct_change": 0, + } + ], + ) + # NOTE: no 20240103.csv at all + csv = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + with pytest.raises(SnapshotFileMissingError): + csv.get_bars("600000.SH", "20240102", "20240103") + + +def test_bars_snapshot_missing_vs_per_symbol_missing_classification(): + """get_bars raises SnapshotFileMissingError (subclass of MissingDataError). + + The two failure modes must remain distinguishable for the engine: a per- + symbol gap is a normal business outcome (suspended / delisted / IPO'd), + while a missing whole-day file is a data infrastructure failure that must + abort the run. + """ + assert issubclass(SnapshotFileMissingError, MissingDataError) + + +def test_factor_rejects_zero_in_both_portals(tmp_path): + """Both portals reject non-positive factor values.""" + from hqbacktest.data.errors import InvalidDataError as _I + + mem = InMemoryDataPortal(calendar=["20240102"], as_of="20240102") + with pytest.raises(_I): + mem.add_factor("600000.SH", "20240102", Decimal("0")) + # CSV-side invalid factor is exercised in test_hqdata_portal.py. + + +# --------------------------------------------------------------------------- +# As-of parity +# --------------------------------------------------------------------------- + + +def test_data_version_as_of_agrees_with_calendar_latest(tmp_path): + mem = InMemoryDataPortal(calendar=[d for d, _ in CALENDAR_DATES], as_of="20240105") + csv = _csv_with_gaps(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 +# --------------------------------------------------------------------------- + + +def test_resolve_source_location_rejects_dot_dot(tmp_path): + with pytest.raises(InvalidDataError): + resolve_source_location("..", default_data_root=str(tmp_path)) + + +def test_resolve_source_location_rejects_dot(tmp_path): + with pytest.raises(InvalidDataError): + resolve_source_location(".", default_data_root=str(tmp_path)) + + +# --------------------------------------------------------------------------- +# Immutability of cached lists +# --------------------------------------------------------------------------- + + +def test_cached_bar_lists_are_isolated_from_strategy_mutation(tmp_path): + """The portal must never return its internal cached list reference.""" + mem = _memory_with_gaps() + csv = _csv_with_gaps(tmp_path) + mem_bars = mem.get_bars("000001.SZ", "20240102", "20240105") + csv_bars = csv.get_bars("000001.SZ", "20240102", "20240105") + # Caller mutating the returned list must not corrupt later queries. + mem_bars.clear() + csv_bars.clear() + assert mem.get_bars("000001.SZ", "20240102", "20240105") + assert csv.get_bars("000001.SZ", "20240102", "20240105") diff --git a/tests/data/test_task14_semantics.py b/tests/data/test_task14_semantics.py new file mode 100644 index 0000000..085dddd --- /dev/null +++ b/tests/data/test_task14_semantics.py @@ -0,0 +1,338 @@ +"""Task 14 tests for DataView and snapshot-missing error handling. + +Covers: + * Sentinel `visible_through="00000000"` on the first trading day must not + raise; `history` returns [] and `current_price` returns None. + * `current_price` walks back up to 20 trading days for the most recent + valid close (suspended-symbol semantics). + * `history(bar_count=N)` no longer scans the full pre-start window. + * `SnapshotFileMissingError` is a sibling/child of `MissingDataError` so + engine / broker can distinguish business gap from infrastructure failure. +""" + +from __future__ import annotations + +from decimal import Decimal + +import pytest + +from hqbacktest.data import ( + DataView, + FutureDataAccessError, + InMemoryDataPortal, + MissingDataError, + SnapshotFileMissingError, +) +from hqbacktest.data.data_view import DEFAULT_HISTORY_START +from hqbacktest.domain.bar import Bar +from hqbacktest.engine.engine import BacktestEngine + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +def _bar(date: str, close: str) -> Bar: + # Wide OHLC envelope so any close in [9, 30] is valid. + return Bar.from_raw( + symbol="600000.SH", + date=date, + open="10.00", + high="30.00", + low="9.00", + close=close, + volume=1000, + ) + + +def _long_calendar_portal() -> InMemoryDataPortal: + """30 trading days; bars only on every other day from 20240102.""" + dates = [f"2024{(4 + i // 30):04d}{(2 + i % 28):02d}" for i in range(30)] + # Normalize: take 30 unique YYYYMMDD dates by truncating month overflow + base = ["20240102"] + for i in range(1, 30): + d = int(base[i - 1]) + 1 + # skip weekends (very rough) + base.append(f"{d:08d}") + p = InMemoryDataPortal(calendar=base, as_of=base[-1]) + for d in base[::2]: + p.add_bar(_bar(d, "10.00")) + return p + + +# --------------------------------------------------------------------------- +# Sentinel visible_through +# --------------------------------------------------------------------------- + + +def test_first_trading_day_sentinel_does_not_raise_for_history(): + """Before any trading day, the scheduler uses `00000000` as a sentinel. + + `history` must return an empty list, not raise, so that strategies calling + `data.history(...)` on the first day do not crash. + """ + p = _long_calendar_portal() + view = DataView(portal=p, visible_through="00000000") + assert view.history("600000.SH", field="close", bar_count=20) == [] + + +def test_first_trading_day_sentinel_does_not_raise_for_current_price(): + p = _long_calendar_portal() + view = DataView(portal=p, visible_through="00000000") + assert view.current_price("600000.SH") is None + + +def test_sentinel_get_bars_returns_empty(): + p = _long_calendar_portal() + view = DataView(portal=p, visible_through="00000000") + # Even an explicit get_bars call must return [] (not raise) for the + # sentinel, so the strategy can detect "no data yet" gracefully. + assert view.get_bars("600000.SH", "00000000", "00000000") == [] + + +# --------------------------------------------------------------------------- +# current_price: lookback window +# --------------------------------------------------------------------------- + + +def _gap_portal() -> InMemoryDataPortal: + """21 trading days; bar on every 5th day only.""" + dates = [f"{d:08d}" for d in range(20240102, 20240102 + 21)] + p = InMemoryDataPortal(calendar=dates, as_of=dates[-1]) + for i, d in enumerate(dates): + if i % 5 == 0: + p.add_bar(_bar(d, str(10 + i / 10))) + return p + + +def test_current_price_returns_most_recent_close_within_lookback(): + p = _gap_portal() + last_bar_date = "20240102" # the 0th day, the only bar in the first 5 + view = DataView(portal=p, visible_through="20240105") + # Even though no bar exists on 20240105 itself, the most recent bar + # within 20 trading days is on 20240102. + assert view.current_price("600000.SH") == Decimal("10.0000") + + +def test_current_price_returns_none_when_lookback_exhausted(): + """If no valid close exists within 20 trading days, return None.""" + p = InMemoryDataPortal(calendar=[], as_of="20991231") + view = DataView(portal=p, visible_through="20240102") + assert view.current_price("600000.SH") is None + + +def test_current_price_returns_none_for_symbol_with_no_history(): + p = _gap_portal() + view = DataView(portal=p, visible_through="20240121") + # No bar for 000001.SZ at all. + assert view.current_price("000001.SZ") is None + + +def test_current_price_does_not_scan_full_history(): + """`current_price` must NOT call get_bars with the full pre-start window. + + Regression for the task-14 finding that `DataView.history` always used + `19000101→visible_through`, blowing up the data layer cache. + """ + p = _gap_portal() + view = DataView(portal=p, visible_through="20240105") + + called: list[tuple[str, str, str]] = [] + original = p.get_bars + + def spy(symbol: str, start: str, end: str): + called.append((symbol, start, end)) + return original(symbol, start, end) + + p.get_bars = spy # type: ignore[assignment] + view.current_price("600000.SH") + assert called, "expected at least one get_bars call" + # No call should start at DEFAULT_HISTORY_START (19000101) when an + # explicit visible_through is provided. + assert all(start != DEFAULT_HISTORY_START for _, start, _ in called) + + +# --------------------------------------------------------------------------- +# history semantics +# --------------------------------------------------------------------------- + + +def test_history_returns_empty_when_window_empty(): + p = _gap_portal() + view = DataView(portal=p, visible_through="20240105") + closes = view.history("000001.SZ", field="close", bar_count=5) + assert closes == [] + + +def test_history_does_not_use_legacy_start_when_universe_start_set(): + p = _gap_portal() + view = DataView(portal=p, visible_through="20240105", universe_start="20240101") + called: list[tuple[str, str, str]] = [] + original = p.get_bars + + def spy(symbol: str, start: str, end: str): + called.append((symbol, start, end)) + return original(symbol, start, end) + + p.get_bars = spy # type: ignore[assignment] + view.history("600000.SH", bar_count=5) + assert all(start == "20240101" for _, start, _ in called) + + +def test_history_universe_start_before_data_returns_visible_bars(): + """`universe_start` earlier than the actual data must not raise.""" + p = _gap_portal() + view = DataView(portal=p, visible_through="20240105", universe_start="19800101") + # universe_start is well before any data; history returns whatever + # bars are visible in the window, never raising. + closes = view.history("600000.SH", bar_count=5) + assert closes, "the window's visible bars must still be returned" + assert all(Decimal("10") <= c <= Decimal("10") for c in closes) + + +# --------------------------------------------------------------------------- +# SnapshotFileMissingError classification +# --------------------------------------------------------------------------- + + +def test_snapshot_file_missing_error_is_subclass_of_missing_data(): + assert issubclass(SnapshotFileMissingError, MissingDataError) + + +def test_snapshot_file_missing_error_carries_path_info(): + err = SnapshotFileMissingError("stock_daily", "/tmp/foo/stock_daily/20240102.csv") + msg = str(err) + assert "stock_daily" in msg + assert "20240102" in msg + + +# --------------------------------------------------------------------------- +# Existing behaviour preserved +# --------------------------------------------------------------------------- + + +def test_history_with_universe_start_returns_only_in_window(): + p = _long_calendar_portal() + view = DataView( + portal=p, + visible_through=p.calendar[-1], + universe_start=p.calendar[-5], + ) + closes = view.history("600000.SH", bar_count=10) + assert 0 < len(closes) <= 5 + + +def test_current_price_returns_latest_close_on_full_coverage(): + p = _long_calendar_portal() + view = DataView(portal=p, visible_through=p.calendar[-1]) + price = view.current_price("600000.SH") + assert price is not None + assert Decimal("10.0000") <= price <= Decimal("10.0000") + + +# --------------------------------------------------------------------------- +# SnapshotFileMissingError propagation (task 14: infrastructure failure must +# not be silently folded into a per-symbol gap) +# --------------------------------------------------------------------------- + + +class _SnapshotMissingPortal: + """A portal whose whole-day snapshot is missing for every query.""" + + def get_calendar(self, start: str, end: str) -> list[str]: + return ["20240102", "20240103", "20240104"] + + def get_bars(self, symbol: str, start: str, end: str) -> list: + raise SnapshotFileMissingError("stock_daily", f"/tmp/stock_daily/{end}.csv") + + def get_universe(self, date: str, include_bj: bool = False) -> list: + return [] + + def get_factor(self, symbol: str, start: str, end: str) -> list: + return [] + + +def test_history_propagates_snapshot_missing(): + view = DataView(portal=_SnapshotMissingPortal(), visible_through="20240104") + with pytest.raises(SnapshotFileMissingError): + view.history("600000.SH", field="close", bar_count=5) + + +def test_current_price_propagates_snapshot_missing(): + view = DataView(portal=_SnapshotMissingPortal(), visible_through="20240104") + with pytest.raises(SnapshotFileMissingError): + view.current_price("600000.SH") + + +def test_engine_close_price_propagates_snapshot_missing(): + with pytest.raises(SnapshotFileMissingError): + BacktestEngine._close_price_or_none( + _SnapshotMissingPortal(), "600000.SH", "20240104" + ) + + +def test_engine_lookback_price_propagates_snapshot_missing(): + with pytest.raises(SnapshotFileMissingError): + BacktestEngine._lookback_price_or_none( + _SnapshotMissingPortal(), "600000.SH", "20240104" + ) + + +# --------------------------------------------------------------------------- +# history start bound + early-read error type (task 14) +# --------------------------------------------------------------------------- + + +def test_history_does_not_query_legacy_19000101_start(): + """`history` must not scan the full pre-start window (task 14).""" + p = _gap_portal() + view = DataView(portal=p, visible_through="20240105") + + called: list[tuple[str, str, str]] = [] + original = p.get_bars + + def spy(symbol: str, start: str, end: str): + called.append((symbol, start, end)) + return original(symbol, start, end) + + p.get_bars = spy # type: ignore[assignment] + view.history("600000.SH", field="close", bar_count=5) + assert called, "expected at least one get_bars call" + assert all(start != DEFAULT_HISTORY_START for _, start, _ in called) + + +def test_get_bars_before_universe_start_raises_missing_data_not_future(): + """Reading before the data start is missing data, NOT future-data access.""" + p = _gap_portal() + view = DataView(portal=p, visible_through="20240105", universe_start="20240103") + with pytest.raises(MissingDataError) as exc: + view.get_bars("600000.SH", "20240101", "20240102") + # Must not be (and must not be reported as) future-data access. + assert not isinstance(exc.value, FutureDataAccessError) + + +def test_current_price_none_when_suspended_beyond_lookback(): + """The lookback is bounded by 20 *trading days*, not 20 *bars*. + + A symbol suspended for longer than the lookback must return None (the + window is exhausted), not reach back to its pre-suspension close. This + is a regression guard for a task-15 rewrite that briefly trimmed the + result to the last 20 bars instead of the last 20 trading days. + """ + days = [f"{20240102 + i:08d}" for i in range(30)] + p = InMemoryDataPortal(calendar=days, as_of=days[-1]) + for d in days[:3]: + p.add_bar(_bar(d, "10.00")) + view = DataView(portal=p, visible_through=days[-1]) + assert view.current_price("600000.SH") is None + + +def test_current_price_uses_latest_within_20_trading_days(): + """A bar within the 20-trading-day window is still returned.""" + days = [f"{20240102 + i:08d}" for i in range(30)] + p = InMemoryDataPortal(calendar=days, as_of=days[-1]) + # Bar on day index 15 (well within the last 20 trading days) at 10.5. + p.add_bar(_bar(days[15], "10.50")) + view = DataView(portal=p, visible_through=days[-1]) + assert view.current_price("600000.SH") == Decimal("10.5000") diff --git a/tests/data/test_task15_performance.py b/tests/data/test_task15_performance.py new file mode 100644 index 0000000..f1358a1 --- /dev/null +++ b/tests/data/test_task15_performance.py @@ -0,0 +1,433 @@ +"""Performance and cache-reuse smoke tests for task 15. + +Covers: + * The per-day CSV files are parsed **at most once** per run; calling + `get_bars` / `get_factor` repeatedly never re-reads the file. + * `get_bars` on overlapping windows reuses the cached `Bar` objects + (no per-call object reconstruction). + * A 50-symbol × 250-day backtest that calls `data.history(bar_count=20)` + on every (symbol, day) finishes within a CI-friendly time budget. + +Thresholds are deliberately generous so a busy CI runner does not flake. +""" + +from __future__ import annotations + +import time +from bisect import bisect_left, bisect_right +from decimal import Decimal +from pathlib import Path +from typing import Dict, List + +import pytest + +from hqbacktest.data import DataView, HqDataCsvPortal +from hqbacktest.data.hqdata_portal import HqDataCsvPortal as _PortalCls + + +# --------------------------------------------------------------------------- +# Fixture helpers (reused from test_hqdata_portal style) +# --------------------------------------------------------------------------- + + +def _write_calendar(root: Path, rows: list) -> None: + lines = ["date,is_open"] + for d, f in rows: + lines.append(f"{d},{f}") + (root / "calendar.csv").write_text("\n".join(lines) + "\n", encoding="utf-8") + + +def _write_stock_list(root: Path, date: str, symbols: list) -> None: + target = root / "stock_list" + target.mkdir(parents=True, exist_ok=True) + lines = ["symbol,date,name,exchange,board,curr_type,list_date,delist_date"] + for sym in symbols: + lines.append(f"{sym},{date},name,SSE,MB,CNY,19990101,") + (target / f"{date}.csv").write_text("\n".join(lines) + "\n", encoding="utf-8") + + +def _write_stock_daily(root: Path, date: str, rows: list) -> None: + target = root / "stock_daily" + target.mkdir(parents=True, exist_ok=True) + fields = [ + "symbol", + "date", + "pre_close", + "open", + "high", + "low", + "close", + "volume", + "turnover", + "change", + "pct_change", + ] + lines = [",".join(fields)] + for r in rows: + lines.append(",".join(str(r[f]) for f in fields)) + (target / f"{date}.csv").write_text("\n".join(lines) + "\n", encoding="utf-8") + + +def _write_stock_factor(root: Path, date: str, rows: list) -> None: + target = root / "stock_factor" + target.mkdir(parents=True, exist_ok=True) + lines = ["symbol,date,factor"] + for r in rows: + lines.append(f"{r['symbol']},{date},{r['factor']}") + (target / f"{date}.csv").write_text("\n".join(lines) + "\n", encoding="utf-8") + + +def _build_synthetic_snapshot( + root: Path, + *, + symbols: List[str], + trading_days: List[str], + include_factor: bool = True, +) -> Path: + snap = root / "tushare" + snap.mkdir(parents=True, exist_ok=True) + _write_calendar(snap, [(d, "Y") for d in trading_days]) + for d in trading_days: + _write_stock_list(snap, d, symbols) + rows = [] + for i, sym in enumerate(symbols): + close = 10 + (i % 5) + (int(d) % 7) * 0.1 + rows.append( + { + "symbol": sym, + "date": d, + "pre_close": close - 0.1, + "open": close, + "high": close + 0.5, + "low": close - 0.5, + "close": close, + "volume": 10000, + "turnover": 100000, + "change": 0.1, + "pct_change": 1.0, + } + ) + _write_stock_daily(snap, d, rows) + if include_factor: + f_rows = [ + {"symbol": sym, "factor": 1.0 + (int(d) % 3) * 0.01} for sym in symbols + ] + _write_stock_factor(snap, d, f_rows) + return snap + + +# --------------------------------------------------------------------------- +# Cache reuse tests +# --------------------------------------------------------------------------- + + +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. + """ + 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 + + real_read = pd.read_csv + read_calls: Dict[str, int] = {} + + 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 + + monkeypatch.setattr(hp_mod.pd, "read_csv", counting_read) + + # Issue many overlapping queries. + 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. + 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" + + +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] = {} + + 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(hp_mod.pd, "read_csv", counting_read) + + 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 + + +def test_bar_objects_reused_across_overlapping_queries(tmp_path): + """`get_bars` overlapping windows must return the same `Bar` instances. + + Memory control: a fresh object per call would multiply allocation by + the number of overlapping queries. Task 15 says bar objects are + cached and reused. + """ + symbols = ["600000.SH"] + days = ["20240102", "20240103", "20240104"] + _build_synthetic_snapshot(tmp_path, symbols=symbols, trading_days=days) + + portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + + wide = portal.get_bars("600000.SH", "20240102", "20240104") + narrow = portal.get_bars("600000.SH", "20240103", "20240104") + # The bar at 20240104 is the same object across both queries. + assert wide[-1] is narrow[-1] + assert wide[-2] is narrow[-2] + + +def test_history_does_not_rescan_full_pre_start_window(tmp_path): + """`DataView.history` must not hit the portal with `19000101→D` windows. + + With the task-15 cache the underlying `get_bars` call should use a + bounded window (not the legacy 19000101 start). This is a regression + guard for the original 2026-08 finding. + """ + symbols = ["600000.SH"] + days = ["20240102", "20240103", "20240104", "20240105", "20240108"] + _build_synthetic_snapshot(tmp_path, symbols=symbols, trading_days=days) + portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + + captured: List = [] + + real_get_bars = portal.get_bars + + def spy(symbol, start, end): + captured.append((symbol, start, end)) + return real_get_bars(symbol, start, end) + + portal.get_bars = spy # type: ignore[assignment] + + view = DataView(portal=portal, visible_through="20240108") + closes = view.history("600000.SH", field="close", bar_count=5) + assert len(closes) == 5 + # No call should start before the snapshot began (19000101 sentinel). + for _, start, _ in captured: + assert not start.startswith( + "19" + ), f"history called get_bars with legacy {start} start" + + +# --------------------------------------------------------------------------- +# Performance smoke test (CI-friendly) +# --------------------------------------------------------------------------- + + +def test_perf_smoke_50_symbols_250_days_history(tmp_path): + """50 symbols × 250 days, history(bar_count=20) every (symbol, day). + + The legacy portal (pre-task-15) would re-parse each daily file for + every call. With the cache the total wall time must stay under a + generous CI threshold. We pick 15 s — far above the expected + sub-second runtime but well within typical GitHub Actions timeouts. + """ + symbols = [f"{600000 + i:06d}.SH" for i in range(50)] + days = [] + # 250 sequential YYYYMMDD strings starting at 20240102, skipping weekends. + d = 20240102 + while len(days) < 250: + mmdd = d % 10000 + weekday = mmdd % 7 # rough placeholder; we don't actually skip here + days.append(f"{d:08d}") + d += 1 + if d % 100 == 32: + d += 70 # jump a month to keep within 250 entries + # Truncate to exactly 250 just in case the loop overshot. + days = days[:250] + + _build_synthetic_snapshot(tmp_path, symbols=symbols, trading_days=days) + portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + + start = time.monotonic() + for d in days: + view = DataView(portal=portal, visible_through=d) + for sym in symbols: + view.history(sym, field="close", bar_count=20) + elapsed = time.monotonic() - start + + # CI-friendly threshold. The actual runtime on a modern machine is + # well under 1 s for this fixture size. + assert elapsed < 15.0, f"history perf smoke took {elapsed:.2f}s (>15s)" + + +# --------------------------------------------------------------------------- +# Bisect correctness +# --------------------------------------------------------------------------- + + +def test_get_bars_window_returns_correct_slice(tmp_path): + """Window slicing from the cumulative cache must match per-day reads.""" + symbols = ["600000.SH", "000001.SZ"] + days = ["20240102", "20240103", "20240104", "20240105", "20240108"] + _build_synthetic_snapshot(tmp_path, symbols=symbols, trading_days=days) + portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + + full = portal.get_bars("600000.SH", "20240102", "20240108") + mid = portal.get_bars("600000.SH", "20240103", "20240105") + late = portal.get_bars("600000.SH", "20240108", "20240108") + assert [b.date for b in full] == days + assert [b.date for b in mid] == ["20240103", "20240104", "20240105"] + assert [b.date for b in late] == ["20240108"] + + +def test_get_bars_handles_per_symbol_gaps_in_cumulative_cache(tmp_path): + """A suspended day simply yields no bar in the cumulative cache.""" + symbols = ["600000.SH", "000001.SZ"] + days = ["20240102", "20240103", "20240104"] + snap = _build_synthetic_snapshot(tmp_path, symbols=symbols, trading_days=days) + # Suspend 600000.SH on 20240103 by omitting its row. + target = snap / "stock_daily" / "20240103.csv" + lines = target.read_text(encoding="utf-8").splitlines() + header, rest = lines[0], lines[1:] + keep = [ln for ln in rest if not ln.startswith("600000.SH,")] + target.write_text("\n".join([header] + keep) + "\n", encoding="utf-8") + + portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + bars = portal.get_bars("600000.SH", "20240102", "20240104") + assert [b.date for b in bars] == ["20240102", "20240104"] + # 000001.SZ is unaffected. + full = portal.get_bars("000001.SZ", "20240102", "20240104") + assert [b.date for b in full] == days + + +def test_snapshot_file_missing_propagates_through_cumulative_cache( + tmp_path, monkeypatch +): + """If a daily file is missing on disk, the cumulative cache path must + still raise `SnapshotFileMissingError` rather than silently fold the + failure into an empty list (task 14 invariant preserved). + """ + symbols = ["600000.SH"] + days = ["20240102", "20240103"] + _build_synthetic_snapshot(tmp_path, symbols=symbols, trading_days=days) + portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + + # First query populates the cache for 20240102. + portal.get_bars("600000.SH", "20240102", "20240102") + # Now physically remove the 20240103 file. + from hqbacktest.data import SnapshotFileMissingError + + (tmp_path / "tushare" / "stock_daily" / "20240103.csv").unlink() + # Bypass the cached "bars" entry: query a new (start, end) window that + # would force the portal to consult the daily index for 20240103. + from hqbacktest.data import hqdata_portal as hp_mod + + portal._daily_index.pop("20240103", None) + with pytest.raises(SnapshotFileMissingError): + portal.get_bars("600000.SH", "20240103", "20240103") + + +def test_get_factor_uses_cumulative_cache(tmp_path): + """`get_factor` windows return the same factor objects across calls.""" + symbols = ["600000.SH"] + days = ["20240102", "20240103", "20240104"] + _build_synthetic_snapshot(tmp_path, symbols=symbols, trading_days=days) + portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + + full = portal.get_factor("600000.SH", "20240102", "20240104") + late = portal.get_factor("600000.SH", "20240104", "20240104") + assert [d for d, _ in full] == days + assert [d for d, _ in late] == ["20240104"] + # Same Decimal instance (cached tuple element). + assert full[-1][1] is late[-1][1] + + +def test_forward_extend_does_not_parse_unqueried_files(tmp_path): + """Forward extension only reads (cend, end]; never files before the + already-covered range. + + Regression for a bug where the forward-extend path used lo="00000000", + parsing every daily file from the epoch up to `end` — including files + the strategy never queried. A missing/corrupt file outside the query + window must not abort the run (task 15 + task 14 "distinguishable + failure" invariant). + """ + symbols = ["600000.SH"] + snap = tmp_path / "tushare" + snap.mkdir(parents=True, exist_ok=True) + _write_calendar( + snap, + [ + ("20240101", "Y"), + ("20240102", "Y"), + ("20240103", "Y"), + ("20240104", "Y"), + ("20240105", "Y"), + ], + ) + _write_stock_list(snap, "20240102", symbols) + for d in ["20240102", "20240103", "20240104", "20240105"]: + _write_stock_daily( + snap, + d, + [ + { + "symbol": "600000.SH", + "date": d, + "pre_close": 10, + "open": 10, + "high": 11, + "low": 9, + "close": 10, + "volume": 1000, + "turnover": 10000, + "change": 0, + "pct_change": 0, + } + ], + ) + # 20240101.csv is intentionally absent. + portal = HqDataCsvPortal(source="tushare", data_root=str(tmp_path)) + + # First query covers 0103-0104 (never touches 0101). + assert [b.date for b in portal.get_bars("600000.SH", "20240103", "20240104")] == [ + "20240103", + "20240104", + ] + # Forward extension to 0105 must NOT touch the missing 0101 file. + assert [b.date for b in portal.get_bars("600000.SH", "20240105", "20240105")] == [ + "20240105" + ] + # Backward extension to 0102 must not touch 0101 either. + assert [b.date for b in portal.get_bars("600000.SH", "20240102", "20240102")] == [ + "20240102" + ] diff --git a/tests/domain/test_enums.py b/tests/domain/test_enums.py index 7e82e00..2bba755 100644 --- a/tests/domain/test_enums.py +++ b/tests/domain/test_enums.py @@ -38,6 +38,7 @@ def test_reject_reason_is_closed_and_known(): "MISSING_DATA", "DUPLICATE_ORDER", "BACKTEST_ENDED", + "OUT_OF_UNIVERSE", "OTHER", } assert {member.name for member in RejectReason} == expected diff --git a/tests/domain/test_order.py b/tests/domain/test_order.py index 844252d..dbd5d47 100644 --- a/tests/domain/test_order.py +++ b/tests/domain/test_order.py @@ -83,7 +83,7 @@ def test_record_full_fill_moves_to_filled(): o.record_fill("F001", quantity=100, price=Decimal("10.50"), at="20240103") assert o.filled_quantity == 100 assert o.avg_fill_price == Decimal("10.5000") - assert o.fill_ids == ["F001"] + assert o.fill_ids == ("F001",) assert o.status is OrderStatus.FILLED @@ -122,7 +122,7 @@ def test_record_fill_for_new_order_does_not_mutate_order(): assert o.status is OrderStatus.NEW assert o.filled_quantity == 0 assert o.avg_fill_price is None - assert o.fill_ids == [] + assert o.fill_ids == () def test_multiple_partial_fills_keep_partial_status_until_complete(): diff --git a/tests/engine/test_broker.py b/tests/engine/test_broker.py index e739c37..942fc74 100644 --- a/tests/engine/test_broker.py +++ b/tests/engine/test_broker.py @@ -418,7 +418,10 @@ def on_bar(self, context, data): assert portfolio.cash + portfolio.market_value( {"600000.SH": Decimal("10.0000")} ) == Decimal("99984.00") - # Fees live in `realized_pnl` via the broker + portfolio.apply_fill path. + # Fees do NOT live in `realized_pnl`: per contract §3.1 / rule 8, + # realized_pnl is the gross (sell_price - avg_cost) * quantity and + # stays independent of commission / stamp_tax / other_fee. Fees flow + # through `cash` only. fills = [e for e in engine.event_log.all() if e.phase is EventType.ORDER_FILLED] assert len(fills) == 3 diff --git a/tests/engine/test_engine.py b/tests/engine/test_engine.py index 0ed0584..f5c5d66 100644 --- a/tests/engine/test_engine.py +++ b/tests/engine/test_engine.py @@ -126,28 +126,46 @@ def after_trading_end(self, context): ] -def test_engine_before_trading_start_cannot_read_today_close(): - from hqbacktest.data import MissingDataError +def test_engine_before_trading_start_cannot_read_future(): + """The strategy must never see past `visible_through`. + + Per task 14: on the first trading day the sentinel `00000000` is + used and `history(...)` returns `[]` rather than raising (the + strategy is allowed to call it; the data layer simply has nothing). + On every other day, a direct future-data access still must raise. + """ + from hqbacktest.data import FutureDataAccessError + + captured_first_day: list[list] = [] + future_attempts: list[int] = [] class Reader: def initialize(self, context): return None def before_trading_start(self, context, data): - data.history("600000.SH", field="close", bar_count=1) + captured_first_day.append( + data.history("600000.SH", field="close", bar_count=1) + ) def on_bar(self, context, data): - return None + # Try to read past `visible_through` directly via get_bars. + try: + data.get_bars("600000.SH", "20991231", "20991231") + future_attempts.append(0) + except FutureDataAccessError: + future_attempts.append(1) def after_trading_end(self, context): return None engine = BacktestEngine(_config(), strategy=Reader(), portal=_portal()) - with pytest.raises(RunFailed) as exc: - engine.run() - # The first day (no history) and the second day (D=20240103 not yet visible) - # both raise; we just need at least one to fail and be reported. - assert exc.value.phase == "BEFORE_TRADING_START" + engine.run() + assert captured_first_day[0] == [] + assert all(future_attempts), ( + f"Expected every BAR_CLOSE get_bars call past visible_through to raise; " + f"got {future_attempts}" + ) def test_engine_bar_close_can_read_today_close(): @@ -300,7 +318,12 @@ def test_engine_event_log_preserves_chronological_order(): assert list(by_date.keys()) == sorted(by_date.keys()) -def test_engine_emits_zero_events_for_empty_calendar(): +def test_engine_rejects_empty_calendar_window(): + """Task 20: an empty trading-day window must raise rather than + silently produce an empty result. + """ + from hqbacktest.engine.errors import ConfigurationError + p = InMemoryDataPortal(calendar=[]) cfg = BacktestConfig( start_date="20240102", @@ -308,9 +331,8 @@ def test_engine_emits_zero_events_for_empty_calendar(): initial_cash=Decimal("100000"), source="tushare", ) - result = BacktestEngine(cfg, portal=p).run() - assert result.trading_days == [] - assert len(result.event_log) == 0 + with pytest.raises(ConfigurationError, match="no trading days"): + BacktestEngine(cfg, portal=p).run() def test_engine_rejects_invalid_config(): diff --git a/tests/engine/test_intents.py b/tests/engine/test_intents.py index 51a9787..36433b9 100644 --- a/tests/engine/test_intents.py +++ b/tests/engine/test_intents.py @@ -52,8 +52,10 @@ def test_signed_diff_to_lots_buy(): def test_signed_diff_to_lots_sell(): - # Need to go from 300 to 50; diff = -250, lot-rounded = -200. - assert signed_diff_to_lots(50, 300) == -200 + # Need to go from 300 to 50; diff = -250. Task 16: SELL preserves + # the requested share count (odd-lot SELLs allowed per A-share rules), + # so the result is -250, not lot-rounded -200. + assert signed_diff_to_lots(50, 300) == -250 def test_signed_diff_to_lots_no_change(): @@ -70,10 +72,12 @@ def test_target_quantity_for_value_positive(): ) -def test_target_quantity_for_value_zero_returns_current(): +def test_target_quantity_for_value_zero_returns_zero_per_docstring(): + # Task 16: per the function's docstring ("may be 0 to flatten"), a + # zero target returns 0 — the caller flattens via order_target(). assert ( target_quantity_for_value(Decimal("0"), Decimal("12.50"), current_quantity=300) - == 300 + == 0 ) diff --git a/tests/engine/test_iterator.py b/tests/engine/test_iterator.py index dc4b7e6..aa77dd7 100644 --- a/tests/engine/test_iterator.py +++ b/tests/engine/test_iterator.py @@ -24,10 +24,13 @@ def test_iterator_respects_window(): assert list(it) == ["20240103", "20240104"] -def test_iterator_is_empty_when_window_has_no_open_days(): +def test_iterator_raises_when_window_has_no_open_days(): + """Task 20: an empty trading-day window is a hard error, not a + silent success. This avoids the "no signals" misreport bug. + """ p = _portal_with(["20240102"]) - it = TradingDayIterator(p, "20240105", "20240110") - assert list(it) == [] + with pytest.raises(ConfigurationError, match="no trading days"): + TradingDayIterator(p, "20240105", "20240110") def test_iterator_rejects_invalid_dates(): @@ -56,10 +59,6 @@ def test_iterator_len_and_is_empty(): assert len(it) == 2 assert not it.is_empty() - empty = TradingDayIterator(p, "20240110", "20240115") - assert len(empty) == 0 - assert empty.is_empty() - def test_iterator_does_not_invent_natural_days(): """If the calendar lacks a date, the iterator must not yield it.""" diff --git a/tests/engine/test_result_export.py b/tests/engine/test_result_export.py index 5e7a1b8..99cd285 100644 --- a/tests/engine/test_result_export.py +++ b/tests/engine/test_result_export.py @@ -89,7 +89,12 @@ def on_bar(self, context, data): assert isinstance(pt.drawdown, Decimal) -def test_engine_empty_calendar_yields_empty_equity_curve(): +def test_engine_empty_calendar_raises(): + """Task 20: an empty trading-day window must raise rather than + silently produce an empty result. + """ + from hqbacktest.engine.errors import ConfigurationError + class Null(BaseStrategy): def initialize(self, context): pass @@ -99,8 +104,8 @@ def initialize(self, context): strategy=Null(), portal=_portal([]), ) - result = engine.run() - assert result.equity_curve == [] + with pytest.raises(ConfigurationError, match="no trading days"): + engine.run() # --------------------------------------------------------------------- # @@ -299,34 +304,55 @@ def initialize(self, context): # --------------------------------------------------------------------- # -def test_engine_valuation_fails_when_held_symbol_has_no_close(): - """Contract §4: 持仓标的没有有效收盘价 → 运行失败 + DATA_ERROR 事件; - v0.1 never silently values holdings at zero.""" +def test_engine_valuation_uses_lookback_for_suspended_symbol(): + """Contract §4 + task 14: a suspended holding is valued at the most + recent valid close within the lookback window and a DATA_WARNING is + recorded. The run continues normally. + """ class BuyHold(BaseStrategy): def initialize(self, context): context.set_universe(["600000.SH"]) - def on_bar(self, context, data): - if context.now == "20240102": - context.order("600000.SH", 100) + def before_trading_start(self, context, data): + # Buy on the only trading day that has a bar; the subsequent + # valuation days must fall back to the same close. + context.order("600000.SH", 100) p = InMemoryDataPortal(calendar=["20240102", "20240103", "20240104"]) p.add_bar(_bar("20240102")) - p.add_bar(_bar("20240103")) - # No bar on 20240104: the buy filled at 0103 open, and day-end - # valuation on 0104 cannot price the 100-share holding. + # 20240103 / 20240104 have NO bar for 600000.SH; the lookback fallback + # should use the 20240102 close for valuation. engine = BacktestEngine( _config("20240102", "20240104"), strategy=BuyHold(), portal=p ) - from hqbacktest.engine.errors import RunFailed - - with pytest.raises(RunFailed): - engine.run() + result = engine.run() + warnings = [e for e in engine.event_log.all() if e.phase is EventType.DATA_WARNING] + # One DATA_WARNING per suspended day (20240103, 20240104). + assert len(warnings) == 2 + assert all("600000.SH" in e.detail for e in warnings) + # Equity curve populated, no DATA_ERROR (lookback succeeded). data_errors = [e for e in engine.event_log.all() if e.phase is EventType.DATA_ERROR] - assert len(data_errors) == 1 - assert "600000.SH" in data_errors[0].detail - assert engine.result is None # no half-populated result on failure + assert data_errors == [] + assert len(result.equity_curve) == 3 + + +def test_engine_valuation_aborts_when_lookback_exhausted(): + """Direct test of the engine valuation fallbacks: when even the + 20-day lookback cannot find a valid close for a held symbol, the run + aborts with DATA_ERROR. + + Constructed by directly invoking the private `_lookback_price_or_none` + helper, since the engine only ever holds a position after a + successful fill (which requires at least one bar somewhere in the + calendar). This test pins the task-14 contract: lookback or fail, + never silently zero. + """ + from hqbacktest.engine.engine import BacktestEngine + + p = InMemoryDataPortal(calendar=["20240102", "20240103", "20240104"]) + # No bars at all: lookback is empty. + assert BacktestEngine._lookback_price_or_none(p, "600000.SH", "20240104") is None # --------------------------------------------------------------------- # diff --git a/tests/engine/test_task16_matching.py b/tests/engine/test_task16_matching.py new file mode 100644 index 0000000..5854d69 --- /dev/null +++ b/tests/engine/test_task16_matching.py @@ -0,0 +1,429 @@ +"""Task 16 hand-calculated regression tests for matching + lot rounding. + +Covers: + * Same-day SELL proceeds are available to fund the same batch's BUY + orders ("卖旧买新" rotation). + * SELL orders are not lot-rounded: odd-lot SELLs (含零股) succeed and + can fully flatten a position via `order_target(symbol, 0)`. + * SELL 150 must NOT be silently shrunk to 100 (静默篡改策略意图 + violates the contract). + * Same-day BUY-then-SELL vs SELL-then-BUY at the same price produce + the same realized_pnl for the SELL position (fee differences are + not realized_pnl). + * T+1 violation -> whole-order rejection (no partial fills in v0.1). + * `intents.target_quantity_for_value(0)` returns 0 (flatten), per its + docstring. + * CLI `_require_decimal` rejects float (consistent with engine). + * `Fill` with non-zero stamp_tax on a BUY raises (cost-table consistency). +""" + +from decimal import Decimal + +import pytest + +from hqbacktest import BacktestConfig, BacktestEngine, BaseStrategy +from hqbacktest.cli.config import _require_decimal, ConfigError +from hqbacktest.data import InMemoryDataPortal +from hqbacktest.domain.bar import Bar +from hqbacktest.domain.enums import EventType, OrderStatus, Side +from hqbacktest.domain.fill import Fill +from hqbacktest.engine.intents import signed_diff_to_lots, target_quantity_for_value +from hqbacktest.domain.portfolio import Portfolio + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _bar(date: str, open_: str, close: str) -> Bar: + return Bar.from_raw( + symbol="600000.SH", + date=date, + open=open_, + high="20.0000", + low="9.0000", + close=close, + volume=1000, + ) + + +def _bar2(date: str, open_: str, close: str) -> Bar: + return Bar.from_raw( + symbol="000001.SZ", + date=date, + open=open_, + high="20.0000", + low="9.0000", + close=close, + volume=1000, + ) + + +def _portal_two_symbols() -> InMemoryDataPortal: + p = InMemoryDataPortal( + calendar=["20240102", "20240103", "20240104"], + universe_by_date={"20240102": ["600000.SH", "000001.SZ"]}, + as_of="20240104", + ) + for d in ("20240102", "20240103", "20240104"): + p.add_bar(_bar(d, "10.0000", "10.0000")) + p.add_bar(_bar2(d, "10.0000", "10.0000")) + return p + + +def _cfg(start="20240102", end="20240104") -> BacktestConfig: + return BacktestConfig( + start_date=start, + end_date=end, + initial_cash=Decimal("100000"), + source="tushare", + ) + + +# --------------------------------------------------------------------------- +# Same-day SELL-then-BUY: rotation orders must not be rejected +# --------------------------------------------------------------------------- + + +def test_same_day_sell_proceeds_fund_buy_in_same_batch(): + """Hand-calculated regression: hold 900 @ 10 (cost basis 9000), cash + 45; same batch SELL 900 + BUY 900 @ 10. + + Before task 16 the BUY was rejected with INSUFFICIENT_CASH because + the broker matched in submission order using a pre-batch cash + snapshot. After task 16 SELL runs first and its net proceeds + (900 * 10 - 5 commission - 9 stamp = 8986) push cash from 45 to + 9031, funding the 9005-cost BUY in the same batch. + """ + + class Rotate(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH", "000001.SZ"]) + + def before_trading_start(self, context, data): + # Day 1: build the 600000.SH position from initial cash. + if context.now == "20240102": + context.order("600000.SH", 900) + # Day 2: rotation — sell 600000.SH, buy 000001.SZ, same batch. + elif context.now == "20240103": + context.order_target("600000.SH", 0) + context.order("000001.SZ", 900) + + # 9050 = 9000 (cost) + 5 (commission) + 40 spare so day 1 buy fills. + cfg = _cfg() + cfg = BacktestConfig( + start_date=cfg.start_date, + end_date=cfg.end_date, + initial_cash=Decimal("9050"), + source=cfg.source, + ) + engine = BacktestEngine(cfg, strategy=Rotate(), portal=_portal_two_symbols()) + engine.run() + + fills = [e for e in engine.event_log.all() if e.phase is EventType.ORDER_FILLED] + day3_fills = [e for e in fills if e.date == "20240103"] + # 20240103: both legs fill (SELL 600000.SH + BUY 000001.SZ). + assert len(day3_fills) == 2 + # The BUY must not have been rejected for cash. + rejected = [ + e + for e in engine.event_log.all() + if e.phase is EventType.ORDER_REJECTED and e.date == "20240103" + ] + assert all("INSUFFICIENT_CASH" not in (e.error or "") for e in rejected) + # Sell fill exists with the full 900 quantity (not lot-rounded). + sell_fill = next( + ( + e + for e in day3_fills + if e.detail and "qty=900" in e.detail and "stamp=" in e.detail + ), + None, + ) + assert sell_fill is not None, "expected a SELL 900 fill on 20240103" + + +# --------------------------------------------------------------------------- +# Odd-lot SELL: not lot-rounded, can flatten +# --------------------------------------------------------------------------- + + +def test_sell_150_is_not_silently_rounded_to_100(): + """`order_target(symbol, 0)` on a 150-share position must flatten all + 150 shares, not silently shrink to 100. The sell fill must show + `qty=150`, not `qty=100`. + """ + portfolio = Portfolio(initial_cash=Decimal("1000")) + position = portfolio.get_position("600000.SH") + position.update_buy(150, Decimal("10.0000")) + portfolio.settle_t1(today="20240101", previous_date=None) + assert position.sellable_quantity == 150 + + class Flatten(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def before_trading_start(self, context, data): + context.order_target("600000.SH", 0) + + p = _portal_two_symbols() + engine = BacktestEngine(_cfg(), strategy=Flatten(), portal=p) + # Inject the 150-share pre-position directly into the engine ledger. + pre_pos = engine.portfolio.get_position("600000.SH") + pre_pos.update_buy(150, Decimal("10.0000")) + engine.portfolio.settle_t1(today="20240101", previous_date=None) + + engine.run() + + # The SELL must fill for the full 150 shares — qty=150, never qty=100. + sell_fills = [ + e + for e in engine.event_log.all() + if e.phase is EventType.ORDER_FILLED and e.detail and "qty=150" in e.detail + ] + assert sell_fills, "expected a 150-share SELL fill" + # And there must NOT be a 100-share fill for 600000.SH (which would + # be the silent shrink). + bad_fills = [ + e + for e in engine.event_log.all() + if e.phase is EventType.ORDER_FILLED and e.detail and "qty=100" in e.detail + ] + assert not bad_fills, f"unexpected 100-share SELL fill: {bad_fills}" + + +def test_sell_50_odd_lot_is_accepted_by_engine(): + """`order(sym, -50)` on a position with ≥ 50 sellable must succeed. + + Sanity check: when the position holds 100 (1 lot) and the strategy + submits order(sym, -50), the SELL fills for exactly 50 shares. (The + engine must NOT truncate to 100 nor reject for lot-size.) + """ + + class SellHalf(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def before_trading_start(self, context, data): + if context.now == "20240102": + context.order("600000.SH", 100) + elif context.now == "20240103": + context.order("600000.SH", -50) + + p = _portal_two_symbols() + engine = BacktestEngine(_cfg(), strategy=SellHalf(), portal=p) + engine.run() + # Expect a fill on 20240103 with qty=50 (the odd-lot SELL). + sell_fills = [ + e + for e in engine.event_log.all() + if e.phase is EventType.ORDER_FILLED + and e.date == "20240103" + and e.detail + and "qty=50" in e.detail + ] + assert sell_fills, "expected a SELL 50 fill on 20240103" + + +# --------------------------------------------------------------------------- +# Lot rounding helper: SELL no longer lot-aligned +# --------------------------------------------------------------------------- + + +def test_signed_diff_to_lots_sell_does_not_round_to_lot(): + """`signed_diff_to_lots(50, 300)` returns -250, not -200. + + SELL orders preserve the requested odd-lot count; the lot floor + applies only to BUY orders. + """ + assert signed_diff_to_lots(50, 300) == -250 + + +def test_signed_diff_to_lots_buy_still_floors_to_lot(): + assert signed_diff_to_lots(250, 0) == 200 + + +# --------------------------------------------------------------------------- +# target_quantity_for_value(0) flattens per docstring +# --------------------------------------------------------------------------- + + +def test_target_quantity_for_value_zero_returns_zero_per_docstring(): + """`target_quantity_for_value(Decimal('0'), ...)` returns 0 + (flatten), matching the docstring's "may be 0 to flatten". + """ + assert ( + target_quantity_for_value(Decimal("0"), Decimal("12.50"), current_quantity=300) + == 0 + ) + + +# --------------------------------------------------------------------------- +# T+1 whole-order rejection (not partial fill) +# --------------------------------------------------------------------------- + + +def test_t1_violation_rejects_whole_order_not_partial(): + """Selling more than the sellable_quantity rejects the entire order, + not a partial fill. v0.1 does not implement partial fills. + """ + + class SellMoreThanHeld(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def before_trading_start(self, context, data): + if context.now == "20240102": + context.order("600000.SH", 100) # 100 shares bought day 1 + elif context.now == "20240103": + # Try to sell 200 — only 100 sellable (T+1). + context.order("600000.SH", -200) + + p = _portal_two_symbols() + engine = BacktestEngine(_cfg(), strategy=SellMoreThanHeld(), portal=p) + engine.run() + rejected = [ + e + for e in engine.event_log.all() + if e.phase is EventType.ORDER_REJECTED and e.date == "20240103" + ] + assert rejected, "expected a rejection on 20240103 for T+1 violation" + # No partial fill: the position must remain at 100 shares. + fills = [ + e + for e in engine.event_log.all() + if e.phase is EventType.ORDER_FILLED and e.date == "20240103" + ] + assert not fills, f"expected no fills on 20240103, got {fills}" + + +# --------------------------------------------------------------------------- +# Same-day BUY-then-SELL vs SELL-then-BUY: identical realized_pnl for SELL +# --------------------------------------------------------------------------- + + +def test_realized_pnl_independent_of_intraday_match_order(): + """Realized PnL on the SELL equals price - cost, regardless of + whether the BUY in the same batch happened before or after the SELL. + + With fees handled outside realized_pnl (contract rule 8), the + intraday match order does not affect realized_pnl. (Both BUY and + SELL execute at the same price, the SELL's avg_cost is the BUY's + price, so realized = price - price = 0.) + """ + + class TwoStepSameBatch(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH", "000001.SZ"]) + + def before_trading_start(self, context, data): + if context.now == "20240102": + # Pre-position with 100 of 600000.SH for day-3 sell. + context.order("600000.SH", 100) + elif context.now == "20240103": + # Both legs in the same batch. + context.order("000001.SZ", 100) # BUY first + context.order_target("600000.SH", 0) # then SELL + + p = _portal_two_symbols() + engine = BacktestEngine(_cfg(), strategy=TwoStepSameBatch(), portal=p) + engine.run() + # realized_pnl on the 600000.SH position should be 0 + # (sell @ 10 against cost @ 10). Fees are NOT part of realized_pnl. + pos = engine.portfolio.positions.get("600000.SH") + if pos is not None: + assert pos.realized_pnl == Decimal("0.00") + + +# --------------------------------------------------------------------------- +# Fill invariants +# --------------------------------------------------------------------------- + + +def test_fill_buy_with_nonzero_stamp_tax_raises(): + """A BUY fill must not carry stamp_tax (印花税只在 SELL 收取). + Stamp_tax on a BUY would let the cash ledger drift from the costs + table. + """ + with pytest.raises(ValueError, match="stamp_tax"): + Fill.from_trade( + fill_id="F1", + order_id="O1", + symbol="600000.SH", + side=Side.BUY, + quantity=100, + price=Decimal("10.0000"), + commission=Decimal("5.00"), + stamp_tax=Decimal("1.00"), # invalid for BUY + other_fee=Decimal("0"), + filled_at="20240102", + session=EventType.OPEN_MATCH, + ) + + +# --------------------------------------------------------------------------- +# CLI initial_cash rejects float (consistent with engine) +# --------------------------------------------------------------------------- + + +def test_cli_require_decimal_rejects_float(): + """The CLI validator must reject float inputs to match the engine's + contract rule 5 ("Decimal/str/int; float forbidden"). Otherwise a + TOML like `initial_cash = 100000.0` would silently convert to + Decimal via `Decimal(str(100000.0))` at the CLI layer. + """ + with pytest.raises(ConfigError, match="float"): + _require_decimal({"initial_cash": 100000.0}, "capital", "initial_cash") + + +def test_same_day_buy_first_rotation_funds_from_sell(): + """A BUY submitted BEFORE its funding SELL must still fill. + + The broker normalizes the batch to [SELLs..., BUYs...] regardless of + submission order, and returns results in that matching order so the + engine credits the SELL's cash before debiting the BUY. A BUY-first + rotation must NOT be falsely rejected for INSUFFICIENT_CASH. + """ + + class RotateBuyFirst(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH", "000001.SZ"]) + + def before_trading_start(self, context, data): + if context.now == "20240102": + context.order("600000.SH", 900) + elif context.now == "20240103": + # BUY submitted before its funding SELL. + context.order("000001.SZ", 900) + context.order_target("600000.SH", 0) + + cfg = BacktestConfig( + start_date="20240102", + end_date="20240104", + # Enough for day-1 buy (9005), but not enough to fund day-3's + # 9005 BUY without the same-batch SELL proceeds. + initial_cash=Decimal("9050"), + source="tushare", + ) + engine = BacktestEngine( + cfg, strategy=RotateBuyFirst(), portal=_portal_two_symbols() + ) + engine.run() + + day3_fills = [ + e + for e in engine.event_log.all() + if e.phase is EventType.ORDER_FILLED and e.date == "20240103" + ] + assert len(day3_fills) == 2 + rejected = [ + e + for e in engine.event_log.all() + if e.phase is EventType.ORDER_REJECTED and e.date == "20240103" + ] + assert all("INSUFFICIENT_CASH" not in (e.error or "") for e in rejected) + # Rotation completed: 600000.SH flattened, 000001.SZ held. + positions = {s: p.quantity for s, p in engine.portfolio.positions.items()} + assert positions.get("600000.SH", 0) == 0 + assert positions.get("000001.SZ", 0) == 900 diff --git a/tests/engine/test_task17_metrics.py b/tests/engine/test_task17_metrics.py new file mode 100644 index 0000000..0f424a6 --- /dev/null +++ b/tests/engine/test_task17_metrics.py @@ -0,0 +1,376 @@ +"""Task 17: equity curve / metrics baseline regression tests. + +Covers: + * First-day P&L flows into `daily_return` (no longer hard-coded 0). + * First-day P&L flows into `drawdown` (no longer hard-coded 0). + * A 1-day hold that drops 18% reports `max_drawdown` ~= 18%. + * `∏(1 + daily_return) == 1 + total_return` (chained-product identity). + * Single-day and two-day runs return `None` for `daily_volatility` / + `annualized_volatility` / `sharpe_ratio` instead of misleading 0. + * `positions_table.sellable_quantity` records the post-settlement + value (D-row shows the shares sellable on D+1, per contract). + * `metrics.py` never builds Decimal from `float` directly. +""" + +from decimal import Decimal + +import pytest + +from hqbacktest import BacktestConfig, BacktestEngine, BaseStrategy +from hqbacktest.data import InMemoryDataPortal +from hqbacktest.domain.bar import Bar +from hqbacktest.domain.enums import EventType, OrderStatus +from hqbacktest.domain.portfolio import Portfolio +from hqbacktest.domain.position import Position +from hqbacktest.engine.metrics import ( + EquityPoint, + MetricsConfig, + compute_metrics, +) + + +# --------------------------------------------------------------------------- +# Engine-driven: first-day P&L flows into the equity curve +# --------------------------------------------------------------------------- + + +def _bar(date: str, open_: str, close: str, sym: str = "600000.SH") -> Bar: + # Wide OHLC envelope so any close in [5, 30] is valid. + return Bar.from_raw( + symbol=sym, + date=date, + open=open_, + high="30.0000", + low="5.0000", + close=close, + volume=1000, + ) + + +def _two_day_portal() -> InMemoryDataPortal: + p = InMemoryDataPortal( + calendar=["20240102", "20240103", "20240104"], + universe_by_date={"20240102": ["600000.SH"]}, + as_of="20240104", + ) + # Day 1: open 10.00 -> close 8.20 (down 18%). Day 2: 8.20 -> 9.60. + p.add_bar(_bar("20240102", "10.0000", "8.2000")) + p.add_bar(_bar("20240103", "8.2000", "9.6000")) + p.add_bar(_bar("20240104", "9.6000", "10.0000")) + return p + + +def test_first_day_loss_appears_in_daily_return(): + """A 18% drop on the first trading day must register as a negative + `daily_return` on day 1 (was previously hard-coded to 0). + + Setup: BUY @ 10 on day 1's BAR_CLOSE (matches at day-2 open @ 10), + then day-2 close = 8.2, so the first observed P&L on day 2 is ~-18% + relative to initial cash. + """ + + class BuyAll(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20240102": + context.order_target_percent("600000.SH", Decimal("0.95")) + + cfg = BacktestConfig( + start_date="20240102", + end_date="20240103", + initial_cash=Decimal("100000"), + source="tushare", + ) + # Two-day window: day 1 flat, day 2 drops to 8.2. + p = InMemoryDataPortal( + calendar=["20240102", "20240103"], + universe_by_date={"20240102": ["600000.SH"]}, + as_of="20240103", + ) + p.add_bar(_bar("20240102", "10.0000", "10.0000")) + p.add_bar(_bar("20240103", "10.0000", "8.2000")) + + engine = BacktestEngine(cfg, strategy=BuyAll(), portal=p) + engine.run() + eq = engine.result.equity_curve + # Day-2 daily_return is the FIRST observed P&L (the buy filled at + # day-2 open=10 against the day-2 close=8.2 -> -18%). + # Find the first non-zero daily_return. + nonzero = next((pt for pt in eq if pt.daily_return != 0), None) + assert nonzero is not None + assert nonzero.daily_return < Decimal("0") + # -0.17 (with commission drag): 95000 / 10 = 11500 shares @ 10, but + # cash drops to 4976.25 after fees; equity at close = 82876.25 vs + # initial 100000 -> -0.171. We accept either -0.17 or -0.18 within + # 0.005 tolerance for commission / lot rounding noise. + assert abs(nonzero.daily_return - Decimal("-0.17")) < Decimal("0.01") + + +def test_first_day_loss_appears_in_drawdown(): + """A 18% drop on day 2 must drive `max_drawdown` above 0.""" + cfg = BacktestConfig( + start_date="20240102", + end_date="20240103", + initial_cash=Decimal("100000"), + source="tushare", + ) + + class BuyAll(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20240102": + context.order_target_percent("600000.SH", Decimal("0.95")) + + p = InMemoryDataPortal( + calendar=["20240102", "20240103"], + universe_by_date={"20240102": ["600000.SH"]}, + as_of="20240103", + ) + p.add_bar(_bar("20240102", "10.0000", "10.0000")) + p.add_bar(_bar("20240103", "10.0000", "8.2000")) + + engine = BacktestEngine(cfg, strategy=BuyAll(), portal=p) + result = engine.run() + # 95% of 100k = 95000; buy 11500 shares @ 10 = 95023.75; market_value + # at day-2 close = 11500 * 8.2 = 77900. drawdown ~= (100000 - 72876) / + # 100000 ≈ 0.27 once fees are subtracted; but max_drawdown uses the + # equity curve which is initial_cash (no holding) -> -0.18 the day + # the position marks. We just assert drawdown > 0 (was 0 before + # task 17). + assert result.metrics.max_drawdown > Decimal("0.17") + assert result.metrics.max_drawdown < Decimal("0.20") + + +def test_chained_product_identity_first_day_loss_then_recovery(): + """∏(1 + daily_return) must equal 1 + total_return within Decimal + precision. With a first-day loss flowing into the curve, the + identity used to silently fail because day-1's return was 0. + """ + cfg = BacktestConfig( + start_date="20240102", + end_date="20240103", + initial_cash=Decimal("100000"), + source="tushare", + ) + + class BuyAll(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20240102": + context.order_target_percent("600000.SH", Decimal("0.95")) + + p = InMemoryDataPortal( + calendar=["20240102", "20240103"], + universe_by_date={"20240102": ["600000.SH"]}, + as_of="20240103", + ) + p.add_bar(_bar("20240102", "10.0000", "10.0000")) + p.add_bar(_bar("20240103", "10.0000", "8.2000")) + + engine = BacktestEngine(cfg, strategy=BuyAll(), portal=p) + result = engine.run() + daily_returns = [pt.daily_return for pt in result.equity_curve] + growth = Decimal("1") + for r in daily_returns: + growth *= Decimal("1") + r + expected = Decimal("1") + result.metrics.total_return + diff = abs(growth - expected) + assert diff < Decimal("0.0001"), ( + f"chained product identity broken: product={growth}, " + f"1+total_return={expected}, diff={diff}" + ) + + +# --------------------------------------------------------------------------- +# metrics.py: insufficient samples return None +# --------------------------------------------------------------------------- + + +def test_single_day_volatility_is_none(): + """A 1-day equity curve has < 2 daily returns -> daily_volatility is + `None`, not 0. + """ + eq = [ + EquityPoint( + date="20240102", + cash=Decimal("100000"), + market_value=Decimal("0"), + total_equity=Decimal("100000"), + daily_return=Decimal("0"), + drawdown=Decimal("0"), + ) + ] + m = compute_metrics( + equity_curve=eq, + fills=[], + initial_cash=Decimal("100000"), + config=MetricsConfig(), + ) + assert m.daily_volatility is None + assert m.annualized_volatility is None + assert m.sharpe_ratio is None + assert any("requires >= 2" in n for n in m.notes) + + +def test_two_day_volatility_is_none_when_only_one_return(): + """A 2-day equity curve has exactly one daily return -> stdev on a + single value raises StatisticsError; the metric must be `None`, + not 0, and `sharpe_ratio` must therefore also be `None`. + """ + eq = [ + EquityPoint( + "20240102", + Decimal("100000"), + Decimal("0"), + Decimal("100000"), + Decimal("0"), + Decimal("0"), + ), + EquityPoint( + "20240103", + Decimal("90000"), + Decimal("0"), + Decimal("110000"), + Decimal("0.10"), + Decimal("0"), + ), + ] + m = compute_metrics( + equity_curve=eq, + fills=[], + initial_cash=Decimal("100000"), + config=MetricsConfig(), + ) + assert m.daily_volatility is None + assert m.sharpe_ratio is None + assert any("requires >= 2" in n for n in m.notes) + + +# --------------------------------------------------------------------------- +# positions_table.sellable_quantity semantic: post-settlement snapshot +# --------------------------------------------------------------------------- + + +def test_positions_table_records_post_settlement_sellable_quantity(): + """The D-row in `positions_table` records `sellable_quantity` AFTER + the end-of-day T+1 settlement: shares bought on D become sellable + on D+1, so the D-row shows that sellable count. + + Scenario: BUY 100 @ 10 on day 1, hold to day 2. After day 1's + settlement, the position holds 100 sellable shares; day-1 row + must show that count, not 0. + """ + + class HoldOverNight(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20240102": + context.order("600000.SH", 100) + + cfg = BacktestConfig( + start_date="20240102", + end_date="20240103", + initial_cash=Decimal("100000"), + source="tushare", + ) + p = InMemoryDataPortal( + calendar=["20240102", "20240103"], + universe_by_date={"20240102": ["600000.SH"]}, + as_of="20240103", + ) + p.add_bar(_bar("20240102", "10.0000", "10.0000")) + p.add_bar(_bar("20240103", "10.0000", "10.0000")) + + engine = BacktestEngine(cfg, strategy=HoldOverNight(), portal=p) + result = engine.run() + rows = [r for r in result.positions_table if r["symbol"] == "600000.SH"] + # Day 1 row shows 100 sellable (post-settlement snapshot). + assert rows[0]["quantity"] == "100" + assert rows[0]["sellable_quantity"] == "100" + + +# --------------------------------------------------------------------------- +# metrics.py: no direct Decimal(float) +# --------------------------------------------------------------------------- + + +def test_metrics_output_is_clean_decimal_strings(): + """`annualized_return` must not leak float artifacts (e.g. 1.1**0.039...). + The metric is computed via Decimal(str(float_pow_result)). + """ + eq = [] + base = Decimal("100000") + growth_per_day = Decimal("1.005") + for i in range(20): + equity = base * growth_per_day**i + eq.append( + EquityPoint( + date=f"2024{(i // 30) + 1:04d}{(i % 28) + 2:02d}", + cash=equity, + market_value=Decimal("0"), + total_equity=equity, + daily_return=growth_per_day - Decimal("1") if i > 0 else Decimal("0"), + drawdown=Decimal("0"), + ) + ) + m = compute_metrics( + equity_curve=eq, + fills=[], + initial_cash=Decimal("100000"), + config=MetricsConfig(), + ) + assert m.annualized_return is not None + # Quantize to a reasonable number of decimals; the value must not + # contain the float "inf" / "nan" sentinel and must be a proper Decimal. + s = str(m.annualized_return) + assert "inf" not in s.lower() + assert "nan" not in s.lower() + + +def test_drawdown_peak_includes_initial_cash_on_continued_drop(): + """The running drawdown peak must include `initial_cash`, so a first-day + loss followed by a continued drop does NOT under-report drawdown. + + Scenario: BUY 9500 @ 10 at day-1 open, close 9.1 -> equity 91426.25 + (drawdown ~8.57%). Day-2 close 8.9 -> equity 89526.25. The correct + day-2 drawdown is (100000 - 89526.25) / 100000 = 10.47%, NOT + (91426.25 - 89526.25) / 91426.25 = 2.08% — the peak must be + `initial_cash` (task 17: "回撤峰值序列以 initial_cash 为初始峰值"). + """ + + class BuyAtOpen(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def before_trading_start(self, context, data): + if context.now == "20240102": + context.order("600000.SH", 9500) + + p = InMemoryDataPortal( + calendar=["20240102", "20240103"], + universe_by_date={"20240102": ["600000.SH"]}, + as_of="20240103", + ) + p.add_bar(_bar("20240102", "10.0000", "9.1000")) + p.add_bar(_bar("20240103", "9.1000", "8.9000")) + + cfg = BacktestConfig( + start_date="20240102", + end_date="20240103", + initial_cash=Decimal("100000"), + source="tushare", + ) + engine = BacktestEngine(cfg, strategy=BuyAtOpen(), portal=p) + result = engine.run() + second_day = result.equity_curve[1] + expected = (Decimal("100000") - Decimal("89526.25")) / Decimal("100000") + assert abs(second_day.drawdown - expected) < Decimal("0.01") + assert result.metrics.max_drawdown >= second_day.drawdown diff --git a/tests/engine/test_task18_isolation.py b/tests/engine/test_task18_isolation.py new file mode 100644 index 0000000..4bb2ff1 --- /dev/null +++ b/tests/engine/test_task18_isolation.py @@ -0,0 +1,291 @@ +"""Task 18: strategy isolation + audit-trail integrity tests. + +Covers: + * `Order` is frozen after creation: a strategy that mutates the + Order returned from `Context.pending_orders()` does not affect + the engine's view of the order (broker still matches at the + original quantity, audit log still records the original + `avg_fill_price` and `fill_ids`). + * `DataView.portal` is no longer publicly accessible (the + strategy cannot bypass `visible_through` by calling + `view.portal.get_bars(sym, start, future_date)`). + * `set_universe(...)` actually constrains trading: orders against + a symbol outside the declared universe are rejected with a + typed `OUT_OF_UNIVERSE` reason and an audit-trail event. + * When no universe has been declared, behaviour is unchanged (no + false rejection). + * `Context.historical_universe()` exposes the historical stock + list (per `visible_through`) as a read-only view through the + engine-owned data view, not the raw portal. +""" + +from decimal import Decimal + +import pytest + +from hqbacktest import BacktestConfig, BacktestEngine, BaseStrategy +from hqbacktest.data import DataView, InMemoryDataPortal +from hqbacktest.domain.bar import Bar +from hqbacktest.domain.enums import ( + EventType, + OrderStatus, + OrderType, + RejectReason, + Side, +) +from hqbacktest.engine.config import BacktestConfig as _Cfg # noqa: F401 + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +def _bar(date: str, sym: str = "600000.SH") -> Bar: + return Bar.from_raw( + symbol=sym, + date=date, + open="10.0000", + high="30.0000", + low="5.0000", + close="10.0000", + volume=1000, + ) + + +def _two_symbol_portal() -> InMemoryDataPortal: + p = InMemoryDataPortal( + calendar=["20240102", "20240103", "20240104"], + universe_by_date={ + "20240102": ["600000.SH", "000001.SZ"], + "20240103": ["600000.SH", "000001.SZ"], + "20240104": ["600000.SH", "000001.SZ"], + }, + as_of="20240104", + ) + for d in ("20240102", "20240103", "20240104"): + p.add_bar(_bar(d, sym="600000.SH")) + p.add_bar(_bar(d, sym="000001.SZ")) + return p + + +def _cfg(start="20240102", end="20240103") -> BacktestConfig: + return BacktestConfig( + start_date=start, + end_date=end, + initial_cash=Decimal("100000"), + source="tushare", + ) + + +# --------------------------------------------------------------------------- +# Order immutability for the strategy +# --------------------------------------------------------------------------- + + +def test_pending_orders_returns_frozen_order_copies(): + """Order objects handed to the strategy must be immutable. + + Mutating `quantity` on a returned Order must NOT change the order + the broker eventually matches against. The audit trail's + `avg_fill_price` and `fill_ids` must likewise be unaffected. + """ + + captured: list = [] + + class TryToTamper(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now != "20240102": + return + context.order("600000.SH", 100) + orders = context.pending_orders() + # Strategy tries to inflate the SELL order's quantity so it + # could exceed the lot rule. Frozen Order objects must + # raise on attribute assignment. + with pytest.raises((AttributeError, dataclasses.FrozenInstanceError)): + orders[0].quantity = 9999 # type: ignore[misc] + with pytest.raises((AttributeError, dataclasses.FrozenInstanceError)): + orders[0].avg_fill_price = Decimal("99.99") # type: ignore[misc] + captured.append(len(orders)) + + import dataclasses + + engine = BacktestEngine(_cfg(), strategy=TryToTamper(), portal=_two_symbol_portal()) + engine.run() + # The order for 100 shares of 600000.SH on day 1 must have filled at + # exactly 100, NOT the tampered 9999. + fills = [e for e in engine.event_log.all() if e.phase is EventType.ORDER_FILLED] + assert any( + "qty=100" in (e.detail or "") for e in fills + ), f"expected a 100-share BUY fill, got {fills}" + assert not any("qty=9999" in (e.detail or "") for e in fills) + + +# --------------------------------------------------------------------------- +# DataView.portal is private +# --------------------------------------------------------------------------- + + +def test_data_view_portal_is_not_publicly_accessible(): + """The portal attribute must not be reachable from outside the + data layer. Strategies must not bypass `visible_through` by + reading `view.portal.get_bars(sym, start, future_date)`. + """ + p = InMemoryDataPortal( + calendar=["20240102", "20240103"], + as_of="20240103", + ) + p.add_bar(_bar("20240102")) + p.add_bar(_bar("20240103")) + view = DataView(portal=p, visible_through="20240102") + # The public `portal` attribute must not exist; accessing it must + # raise AttributeError so a strategy cannot reach the raw portal. + with pytest.raises(AttributeError): + view.portal # type: ignore[attr-defined] + + +# --------------------------------------------------------------------------- +# Universe enforcement +# --------------------------------------------------------------------------- + + +def test_orders_outside_universe_are_rejected(): + """`set_universe([...])` must enforce trading scope. + + Submitting an order for a symbol outside the declared universe + must: + - be rejected (REJECTED status + audit-trail event) + - never reach the broker + - not be silently re-routed or filled + """ + + class TradeOutsideUniverse(BaseStrategy): + def initialize(self, context): + # Declare only 600000.SH in the universe. + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20240102": + # Attempt to trade 000001.SZ (outside the universe). + context.order("000001.SZ", 100) + + engine = BacktestEngine( + _cfg(), strategy=TradeOutsideUniverse(), portal=_two_symbol_portal() + ) + engine.run() + rejected = [ + e for e in engine.event_log.all() if e.phase is EventType.ORDER_REJECTED + ] + assert any( + "OUT_OF_UNIVERSE" in (e.error or "") or "000001.SZ" in (e.detail or "") + for e in rejected + ), f"expected an OUT_OF_UNIVERSE rejection, got {rejected}" + # No fill must exist for 000001.SZ. + fills = [e for e in engine.event_log.all() if e.phase is EventType.ORDER_FILLED] + assert not any("000001.SZ" in (e.detail or "") for e in fills) + + +def test_unset_universe_does_not_constrain(): + """When `set_universe` has not been called, trading is unrestricted. + + The order for an arbitrary symbol must reach the broker and fill + normally. + """ + + class NoUniverse(BaseStrategy): + def initialize(self, context): + # NOTE: no set_universe call. + return None + + def on_bar(self, context, data): + if context.now == "20240102": + context.order("600000.SH", 100) + + engine = BacktestEngine(_cfg(), strategy=NoUniverse(), portal=_two_symbol_portal()) + engine.run() + fills = [e for e in engine.event_log.all() if e.phase is EventType.ORDER_FILLED] + assert any("qty=100" in (e.detail or "") for e in fills) + + +# --------------------------------------------------------------------------- +# Context.historical_universe() +# --------------------------------------------------------------------------- + + +def test_historical_universe_returns_portal_universe_for_visible_date(): + """`Context.historical_universe()` returns the historical stock + list for `visible_through` (via `DataView.universe()`), respecting + visibility. It does NOT expose the raw portal. + """ + + class Inspect(BaseStrategy): + def __init__(self): + self.snapshot: list = [] + + def on_bar(self, context, data): + if context.now == "20240102": + self.snapshot.append(list(context.historical_universe())) + + strategy = Inspect() + engine = BacktestEngine(_cfg(), strategy=strategy, portal=_two_symbol_portal()) + engine.run() + assert strategy.snapshot[0] == ["000001.SZ", "600000.SH"] + + +def test_out_of_universe_order_appears_in_orders_table(): + """A rejected out-of-universe order must reach `orders_table` with the + `OUT_OF_UNIVERSE` reason, not just the event log. + + Regression for a bug where the scheduler drained the out-of-universe + list but discarded the orders (never folding them into the engine's + `_orders` dict), leaving `orders_table` empty while the event log + recorded the rejection. + """ + + class TradeOutsideUniverse(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20240102": + context.order("000001.SZ", 100) + + engine = BacktestEngine( + _cfg(), strategy=TradeOutsideUniverse(), portal=_two_symbol_portal() + ) + result = engine.run() + rows = [ + r + for r in result.orders_table + if r["symbol"] == "000001.SZ" and r["status"] == "REJECTED" + ] + assert rows, "expected the out-of-universe order in orders_table" + assert rows[0]["reject_reason"] == "OUT_OF_UNIVERSE" + + +def test_order_fill_ids_is_immutable(): + """`Order.fill_ids` must be an immutable tuple so a strategy holding + a frozen Order cannot append/clear it in place (task 18). + + A frozen dataclass only blocks attribute *reassignment*; a `list` + field would still be mutable in place. Switching to `tuple` closes + that last escape hatch. + """ + from hqbacktest.domain.order import Order + + o = Order( + order_id="O001", + symbol="600000.SH", + side=Side.BUY, + quantity=100, + order_type=OrderType.MARKET, + created_at="20240102", + created_session=EventType.BEFORE_TRADING_START, + ) + assert isinstance(o.fill_ids, tuple) + # Attribute reassignment is already blocked by frozen=True, but an + # in-place list append would NOT be — the tuple type prevents it. + assert not hasattr(o.fill_ids, "append") diff --git a/tests/engine/test_task19_factor_diagnostics.py b/tests/engine/test_task19_factor_diagnostics.py new file mode 100644 index 0000000..d78431d --- /dev/null +++ b/tests/engine/test_task19_factor_diagnostics.py @@ -0,0 +1,368 @@ +"""Task 19: factor-diagnostics-on-holding + CLI summary tests. + +Covers: + * Holding through an ex-date with a > 0.1% factor jump emits a + DATA_WARNING event AND records a FactorDiagnostic on the + collector, surfaced in `result.factor_diagnostics`. + * The 600000.SH 2026-07-16 ex-date case (factor 16.59 -> 17.38, + ≈ 4.7% dividend) reproduces as a clear, traceable warning. + * Diagnostics do NOT change cash, position, or equity: the + balance is byte-identical with the no-diagnostics baseline. + * Symbols that are NOT held or traded never generate holdings- + period diagnostics (the engine still records generic + diagnostics but the holding-period summary must exclude them). + * CLI runner prints a one-line summary when any such diagnostics + were recorded. +""" + +from decimal import Decimal +from io import StringIO + +import pytest + +from hqbacktest import BacktestConfig, BacktestEngine, BaseStrategy +from hqbacktest.cli.runner import run_from_config +from hqbacktest.data import InMemoryDataPortal +from hqbacktest.domain.bar import Bar +from hqbacktest.domain.enums import EventType +from hqbacktest.engine.corporate_actions import ( + DEFAULT_JUMP_BAND, + FactorDiagnostic, +) + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +def _bar(date: str, sym: str = "600000.SH") -> Bar: + return Bar.from_raw( + symbol=sym, + date=date, + open="10.0000", + high="30.0000", + low="5.0000", + close="10.0000", + volume=1000, + ) + + +def _hold_then_dividend_portal() -> InMemoryDataPortal: + """Reproduce 600000.SH 2026-07-16 dividend ex-date. + + Calendar: + 2026-07-14, 15, 16 (ex-date), 17, 18 + Per-day factors (decimal strings): + 14: 16.59 + 15: 16.59 + 16: 17.38 (≈ +4.76% jump — dividend) + 17: 17.38 + 18: 17.38 + """ + p = InMemoryDataPortal( + calendar=["20260714", "20260715", "20260716", "20260717", "20260718"], + universe_by_date={"20260714": ["600000.SH"]}, + as_of="20260718", + ) + factors = { + "20260714": [("600000.SH", "16.59")], + "20260715": [("600000.SH", "16.59")], + "20260716": [("600000.SH", "17.38")], + "20260717": [("600000.SH", "17.38")], + "20260718": [("600000.SH", "17.38")], + } + for d in ("20260714", "20260715", "20260716", "20260717", "20260718"): + p.add_bar(_bar(d)) + p.add_factor("600000.SH", d, Decimal(factors[d][0][1])) + return p + + +def _hold_then_dividend_engine(strategy): + cfg = BacktestConfig( + start_date="20260714", + end_date="20260718", + initial_cash=Decimal("100000"), + source="tushare", + ) + return BacktestEngine(cfg, strategy=strategy, portal=_hold_then_dividend_portal()) + + +# --------------------------------------------------------------------------- +# Holding-period factor diagnostics +# --------------------------------------------------------------------------- + + +def test_holding_through_dividend_emits_warning_event(): + """A 4.76% factor jump while holding must produce a DATA_WARNING + event with the symbol + dates + factor values in the detail. + """ + + class Hold(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20260714": + context.order("600000.SH", 100) + + engine = _hold_then_dividend_engine(Hold()) + engine.run() + warnings = [e for e in engine.event_log.all() if e.phase is EventType.DATA_WARNING] + assert warnings, "expected at least one DATA_WARNING for dividend jump" + detail = warnings[0].detail or "" + assert "600000.SH" in detail + assert "16.59" in detail + assert "17.38" in detail + # The ex-date is 2026-07-16 (the day the factor jumped). + assert "20260716" in detail + + +def test_holding_through_dividend_records_factor_diagnostic(): + """The same holding scenario must also surface in + `result.factor_diagnostics` with kind='abnormal_jump' (the + available kind for cross-day factor changes). + """ + + class Hold(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20260714": + context.order("600000.SH", 100) + + engine = _hold_then_dividend_engine(Hold()) + result = engine.run() + diagnostics = result.factor_diagnostics + assert diagnostics, "expected FactorDiagnostic records" + sym_dates = {(d.symbol, d.date) for d in diagnostics} + assert ("600000.SH", "20260716") in sym_dates + # The diagnostic detail surfaces the factor ratio produced by + # `analyze_factor_series`. The audit-trail event carries the + # before/after factor values (see `test_holding_through_dividend_ + # emits_warning_event`). + detail = next( + d.detail + for d in diagnostics + if d.symbol == "600000.SH" and d.date == "20260716" + ) + assert "abnormal_jump" == next( + d.kind for d in diagnostics if d.symbol == "600000.SH" and d.date == "20260716" + ) + assert "factor ratio" in detail or "ratio" in detail + + +def test_diagnostics_do_not_change_ledger(): + """Diagnostics must be observability-only: cash, position and + equity are byte-identical with a baseline that does not run + diagnostics at all. + + The baseline strategy uses identical inputs and trades but the + portal never triggers the diagnostics path because no ex-date + factor jump exists. + """ + from hqbacktest.data import InMemoryDataPortal + + class Hold(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20260714": + context.order("600000.SH", 100) + + flat = InMemoryDataPortal( + calendar=["20260714", "20260715", "20260716", "20260717", "20260718"], + universe_by_date={"20260714": ["600000.SH"]}, + as_of="20260718", + ) + for d in ("20260714", "20260715", "20260716", "20260717", "20260718"): + flat.add_bar(_bar(d)) + flat.add_factor("600000.SH", d, Decimal("1.0000")) + cfg = BacktestConfig( + start_date="20260714", + end_date="20260718", + initial_cash=Decimal("100000"), + source="tushare", + ) + baseline_engine = BacktestEngine(cfg, strategy=Hold(), portal=flat) + baseline_engine.run() + baseline_curve = list(baseline_engine.result.equity_curve) + + jump_engine = _hold_then_dividend_engine(Hold()) + jump_engine.run() + jump_curve = list(jump_engine.result.equity_curve) + + assert len(baseline_curve) == len(jump_curve) + for b, j in zip(baseline_curve, jump_curve): + assert b.cash == j.cash + assert b.market_value == j.market_value + assert b.total_equity == j.total_equity + + +def test_unheld_symbols_produce_no_holding_diagnostic(): + """A factor jump on a symbol that was never held or traded does + not generate a holdings-period diagnostic for that symbol. + """ + p = InMemoryDataPortal( + calendar=["20260714", "20260715", "20260716"], + universe_by_date={"20260714": ["600000.SH", "999999.SH"]}, + as_of="20260716", + ) + # 600000.SH flat factors; 999999.SH has a big jump on the last day. + factor_for_999 = { + "20260714": Decimal("1.00"), + "20260715": Decimal("1.00"), + "20260716": Decimal("2.00"), + } + for d in ("20260714", "20260715", "20260716"): + p.add_bar(_bar(d)) + p.add_bar(_bar(d, sym="999999.SH")) + p.add_factor("600000.SH", d, Decimal("1.00")) + p.add_factor("999999.SH", d, factor_for_999[d]) + + class OnlyHold(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH", "999999.SH"]) + + def on_bar(self, context, data): + if context.now == "20260714": + context.order("600000.SH", 100) + + cfg = BacktestConfig( + start_date="20260714", + end_date="20260716", + initial_cash=Decimal("100000"), + source="tushare", + ) + engine = BacktestEngine(cfg, strategy=OnlyHold(), portal=p) + result = engine.run() + syms = {d.symbol for d in result.factor_diagnostics} + # 999999.SH was never traded; no holding-period diagnostic for it. + assert "999999.SH" not in syms + + +def test_jump_threshold_default_band_is_wider_than_holding_threshold(): + """`analyze_factor_series` keeps its default `jump_band` of + (0.5, 2.0) for general diagnostics, but the engine applies a + tighter 0.1% threshold for the holdings-period holding summary. + The default band is wider so general diagnostics still + surface truly wild factor swings. + """ + assert DEFAULT_JUMP_BAND == (Decimal("0.5"), Decimal("2.0")) + + +def test_sold_symbol_stops_emitting_holding_diagnostics(): + """A symbol that has been fully sold must not keep emitting + holdings-period factor-jump warnings after its holding period ends. + + Regression for a bug where the engine tracked every symbol ever + traded (a `_traded_symbols` set that never shrank), so a factor + jump AFTER the position was flattened still produced a spurious + DATA_WARNING and FactorDiagnostic. + """ + + class BuyThenSell(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == "20260714": + context.order("600000.SH", 100) + elif context.now == "20260715": + context.order_target("600000.SH", 0) # flatten + + p = InMemoryDataPortal( + calendar=["20260714", "20260715", "20260716", "20260717"], + universe_by_date={"20260714": ["600000.SH"]}, + as_of="20260717", + ) + factors = { + "20260714": "16.59", + "20260715": "16.59", + "20260716": "17.38", # ex-date jump AFTER the position is flat + "20260717": "17.38", + } + for d in ("20260714", "20260715", "20260716", "20260717"): + p.add_bar(_bar(d)) + p.add_factor("600000.SH", d, Decimal(factors[d])) + + cfg = BacktestConfig( + start_date="20260714", + end_date="20260717", + initial_cash=Decimal("100000"), + source="tushare", + ) + engine = BacktestEngine(cfg, strategy=BuyThenSell(), portal=p) + result = engine.run() + assert result.factor_diagnostics == [] + warnings = [e for e in engine.event_log.all() if e.phase is EventType.DATA_WARNING] + assert warnings == [] + + +# --------------------------------------------------------------------------- +# CLI stdout summary +# --------------------------------------------------------------------------- + + +def test_cli_runner_prints_summary_when_diagnostics_present(tmp_path): + """`run_from_config` writes a one-line warning to stdout when the + engine recorded any holding-period factor diagnostics. + """ + import sys + + strategy_src = ( + "from hqbacktest import BaseStrategy\n" + "class Hold(BaseStrategy):\n" + " def initialize(self, context):\n" + " context.set_universe(['600000.SH'])\n" + " def on_bar(self, context, data):\n" + " if context.now == '20260714':\n" + " context.order('600000.SH', 100)\n" + ) + strategy_file = tmp_path / "strategy.py" + strategy_file.write_text(strategy_src) + config_file = tmp_path / "config.toml" + config_file.write_text( + "[start]\n" + "start_date = '20260714'\n" + "end_date = '20260718'\n" + "[capital]\n" + "initial_cash = '100000'\n" + "[data]\n" + "source = 'memory'\n" + "[strategy]\n" + "module = 'strategy'\n" + "class_name = 'Hold'\n" + "[output]\n" + f"directory = '{tmp_path / 'out'}'\n" + ) + portal = _hold_then_dividend_portal() + from hqbacktest.cli import runner + from hqbacktest.cli.config import load_config_file + + original = runner._resolve_portal + runner._resolve_portal = lambda source, data_root: portal + sys.path.insert(0, str(tmp_path)) + sys.modules.pop("strategy", None) # avoid stale cache across tests + cfg_file_obj = load_config_file(str(config_file)) + buf = StringIO() + old_stdout = sys.stdout + sys.stdout = buf + try: + run_result = run_from_config(cfg_file_obj, source_path=str(config_file)) + finally: + sys.stdout = old_stdout + sys.path.remove(str(tmp_path)) + sys.modules.pop("strategy", None) + runner._resolve_portal = original + output = buf.getvalue() + assert ( + "factor" in output.lower() + or "diagnostic" in output.lower() + or "warning" in output.lower() + ), ( + f"expected a factor/diagnostic warning in stdout; got: {output!r}; " + f"exit_code={run_result.exit_code}, message={run_result.message}" + ) diff --git a/tests/integration/__init__.py b/tests/integration/__init__.py new file mode 100644 index 0000000..4bb626d --- /dev/null +++ b/tests/integration/__init__.py @@ -0,0 +1,26 @@ +"""Task 21: real-data integration smoke tests for v0.1.1. + +These tests run against the local `~/.hqdata/tushare` snapshot only when +that directory exists AND contains the calendar.csv + a non-empty +stock_daily subfolder. On machines without the snapshot the entire +group is skipped (no credentials, no network). They live under +`tests/integration/` so the default `pytest tests/` run can opt out +via the `tests` testpath. + +Calibration snapshot: `~/.hqdata/tushare`, 2026-01-05 .. 2026-07-31 +(139 trading days, ~5200 symbols). Re-calibrate the assertions if the +local snapshot changes. + +The four scenarios cover the v0.1 findings that previously broke +real-data runs (per TODO.md 21): + + 1. buy_and_hold across 600000.SH's 2026-07-16 dividend ex-date: + factor jumps 16.5935 -> 17.3774 (~4.7%), factor diagnostics + warn on the holding-period jump. + 2. 5-symbol moving-average strategy over the full window: + must complete and be byte-deterministic across two runs. + 3. Universe that contains a known suspended symbol (000008.SZ, + suspended 2026-07-07..2026-07-13): no crash, warnings traceable. + 4. First-trading-day `before_trading_start` reads `current_price`: + returns None (sentinel visible_through) without crashing. +""" diff --git a/tests/integration/_harness.py b/tests/integration/_harness.py new file mode 100644 index 0000000..9d11e88 --- /dev/null +++ b/tests/integration/_harness.py @@ -0,0 +1,69 @@ +"""Shared helpers + auto-skip for the integration smoke tests.""" + +from __future__ import annotations + +import os +from pathlib import Path + +import pytest + +DEFAULT_DATA_ROOT = "~/.hqdata" +DEFAULT_SOURCE = "tushare" + +# Calibration snapshot window (139 trading days for the +# `~/.hqdata/tushare` snapshot on the v0.1.1 release machine). +CALIBRATION_START = "20260105" +CALIBRATION_END = "20260731" + +# Known dividend ex-date (600000.SH, factor 16.5935 -> 17.3774). +CALIBRATION_EX_DATE = "20260716" + +# Known suspended window (000008.SZ). +CALIBRATION_SUSPENDED_START = "20260707" +CALIBRATION_SUSPENDED_END = "20260713" + + +def _data_root() -> Path: + """Return the resolved hqdata root (overridable via HQDATA_ROOT).""" + raw = os.environ.get("HQDATA_ROOT", DEFAULT_DATA_ROOT) + return Path(raw).expanduser() + + +def _source_available(name: str = DEFAULT_SOURCE) -> bool: + """True iff the source's CSV directory exists and looks non-empty.""" + root = _data_root() / name + if not root.exists() or not root.is_dir(): + return False + if not (root / "calendar.csv").exists(): + return False + try: + next(root.iterdir()) + except StopIteration: + return False # empty directory → snapshot is incomplete + return True + + +def skip_if_no_snapshot() -> pytest.MarkDecorator: + """Skip the test if the local `~/.hqdata/tushare` snapshot is missing. + + Use on every real-data test: + @pytest.mark.integration + @skip_if_no_snapshot() + def test_xxx(): ... + """ + reason = "hqdata snapshot not available at ~/.hqdata/tushare" + return pytest.mark.skipif(not _source_available(), reason=reason) + + +__all__ = [ + "DEFAULT_DATA_ROOT", + "DEFAULT_SOURCE", + "CALIBRATION_START", + "CALIBRATION_END", + "CALIBRATION_EX_DATE", + "CALIBRATION_SUSPENDED_START", + "CALIBRATION_SUSPENDED_END", + "_data_root", + "_source_available", + "skip_if_no_snapshot", +] diff --git a/tests/integration/test_real_data_buy_and_hold.py b/tests/integration/test_real_data_buy_and_hold.py new file mode 100644 index 0000000..900e7ad --- /dev/null +++ b/tests/integration/test_real_data_buy_and_hold.py @@ -0,0 +1,63 @@ +"""Integration scenario 1: buy_and_hold across the 2026-07-16 dividend. + +Calibrated against the `~/.hqdata/tushare` snapshot for v0.1.1: + * 600000.SH factor jumps 16.5935 -> 17.3774 (~4.7%) on 20260716. + * Buy-and-hold from 20260105 to 20260731 with 95% sizing must + produce at least one DATA_WARNING factor diagnostic. +""" + +from __future__ import annotations + +import pytest + +from hqbacktest import BacktestConfig, BacktestEngine, BaseStrategy +from hqbacktest.data import HqDataCsvPortal +from hqbacktest.domain.enums import EventType + +from ._harness import ( + CALIBRATION_END, + CALIBRATION_EX_DATE, + CALIBRATION_START, + skip_if_no_snapshot, +) + + +@pytest.mark.integration +@skip_if_no_snapshot() +def test_buy_and_hold_across_dividend_ex_date(): + """A 95%-of-cash buy on day 1 held through the 20260716 ex-date + must produce at least one DATA_WARNING factor diagnostic. + """ + from decimal import Decimal + + class Hold(BaseStrategy): + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def on_bar(self, context, data): + if context.now == CALIBRATION_START: + context.order_target_percent("600000.SH", Decimal("0.95")) + + cfg = BacktestConfig( + start_date=CALIBRATION_START, + end_date=CALIBRATION_END, + initial_cash=Decimal("100000"), + source="tushare", + ) + portal = HqDataCsvPortal(source="tushare") + engine = BacktestEngine(cfg, strategy=Hold(), portal=portal) + result = engine.run() + # Factor-diagnostics collector must contain the 20260716 jump. + diag_dates = {d.date for d in result.factor_diagnostics if d.symbol == "600000.SH"} + assert ( + CALIBRATION_EX_DATE in diag_dates + ), f"expected a 20260716 factor jump for 600000.SH; got {diag_dates}" + # And the warning must appear in the event log. + warnings = [ + e + for e in engine.event_log.all() + if e.phase is EventType.DATA_WARNING and "600000.SH" in (e.detail or "") + ] + assert any( + CALIBRATION_EX_DATE in (e.detail or "") for e in warnings + ), f"no DATA_WARNING mentions {CALIBRATION_EX_DATE}: {warnings}" diff --git a/tests/integration/test_real_data_first_day.py b/tests/integration/test_real_data_first_day.py new file mode 100644 index 0000000..0e4cb27 --- /dev/null +++ b/tests/integration/test_real_data_first_day.py @@ -0,0 +1,63 @@ +"""Integration scenario 4: first-trading-day `before_trading_start` +reading `current_price` against the real snapshot. + +Per task 14, the first trading day uses the sentinel `visible_through` +of `"00000000"`, so `current_price` must return `None` rather than +crashing. This guards the documented sentinel contract end-to-end +against a real CSV portal (where the empty/invalid date would +otherwise trip the validator). +""" + +from __future__ import annotations + +from decimal import Decimal +from pathlib import Path + +import pytest + +from hqbacktest import BacktestConfig, BacktestEngine, BaseStrategy +from hqbacktest.data import HqDataCsvPortal + +from ._harness import ( + CALIBRATION_END, + CALIBRATION_START, + skip_if_no_snapshot, +) + + +class ReadPriceFirstDay(BaseStrategy): + """Read `current_price` on the first trading day's + `before_trading_start`. Must return `None`, not raise. + """ + + seen: list = [] + + def initialize(self, context): + context.set_universe(["600000.SH"]) + + def before_trading_start(self, context, data): + if context.now == CALIBRATION_START: + price = context.current_price("600000.SH") + self.seen.append(price) + + +@pytest.mark.integration +@skip_if_no_snapshot() +def test_before_trading_start_current_price_first_day(): + """First-day `current_price` returns None (sentinel visible_through) + without raising against the real CSV portal. + """ + strategy = ReadPriceFirstDay() + cfg = BacktestConfig( + start_date=CALIBRATION_START, + end_date=CALIBRATION_END, + initial_cash=Decimal("100000"), + source="tushare", + ) + portal = HqDataCsvPortal(source="tushare") + engine = BacktestEngine(cfg, strategy=strategy, portal=portal) + engine.run() # must not raise + # Strategy ran at least once on the first day. + assert len(strategy.seen) == 1 + # The sentinel view returned None, not 0 or a real price. + assert strategy.seen[0] is None diff --git a/tests/integration/test_real_data_ma_5_symbols.py b/tests/integration/test_real_data_ma_5_symbols.py new file mode 100644 index 0000000..faa0b29 --- /dev/null +++ b/tests/integration/test_real_data_ma_5_symbols.py @@ -0,0 +1,140 @@ +"""Integration scenario 2: 5-symbol moving-average strategy over the +full window, with a wall-clock budget and byte-determinism check. + +Calibrated against the `~/.hqdata/tushare` snapshot for v0.1.1: + * 5 picked symbols (well-known large caps). + * Run budget: < 60 s for the full 139-day window. + * Two runs with identical inputs must produce byte-identical output + (excludes the `timestamp_utc` field in `run_metadata.json`, + which is non-deterministic by design). + +The strategy is defined inline (a local `strategy.py` written into +`tmp_path`) so the universe is configurable from the test instead +of being hard-coded inside `examples.moving_average`. +""" + +from __future__ import annotations + +import json +import time +from pathlib import Path + +import pytest + +from hqbacktest.cli.runner import run_from_file + +from ._harness import ( + CALIBRATION_END, + CALIBRATION_START, + skip_if_no_snapshot, +) + + +UNIVERSE = [ + "600000.SH", + "000001.SZ", + "601318.SH", + "600519.SH", + "000333.SZ", +] + + +_STRATEGY_SRC = ''' +from decimal import Decimal +from hqbacktest import BaseStrategy + + +class FiveSymbolMA(BaseStrategy): + """5-symbol moving average, configurable universe via __init__.""" + + def __init__(self, universe): + super().__init__() + self.universe = list(universe) + + def initialize(self, context): + context.set_universe(self.universe) + + def on_bar(self, context, data): + for sym in self.universe: + closes = data.history(sym, field="close", bar_count=5) + if len(closes) < 5: + continue + avg = sum(closes) / Decimal(len(closes)) + if closes[-1] > avg: + context.order_target_percent(sym, Decimal("0.20")) + else: + context.order_target(sym, 0) +''' + + +@pytest.mark.integration +@skip_if_no_snapshot() +def test_5_symbol_moving_average_full_window_deterministic(tmp_path: Path): + """Two consecutive runs of a 5-symbol MA strategy produce the + same equity_curve.csv / summary.json / fills.csv bytes (modulo + the timestamp field). + + Calibrated to < 60 s on the v0.1.1 release machine; the + threshold is intentionally generous to survive CI jitter. + """ + (tmp_path / "strategy.py").write_text(_STRATEGY_SRC) + cfg_template = ( + "[start]\n" + f"start_date = '{CALIBRATION_START}'\n" + f"end_date = '{CALIBRATION_END}'\n" + "[capital]\n" + "initial_cash = '1000000'\n" + "[data]\n" + "source = 'tushare'\n" + "[strategy]\n" + "module = 'strategy'\n" + f"kwargs = {{ universe = {json.dumps(UNIVERSE)} }}\n" + "[output]\n" + "directory = '__OUT__'\n" + ) + + out_a = tmp_path / "a" + out_b = tmp_path / "b" + cfg_a = tmp_path / "c_a.toml" + cfg_a.write_text(cfg_template.replace("__OUT__", str(out_a))) + cfg_b = tmp_path / "c_b.toml" + cfg_b.write_text(cfg_template.replace("__OUT__", str(out_b))) + + t0 = time.monotonic() + result_a = run_from_file(str(cfg_a), force=True) + elapsed = time.monotonic() - t0 + assert result_a.exit_code == 0, result_a.message + assert elapsed < 60.0, f"5-symbol MA took {elapsed:.2f}s (>60s)" + + result_b = run_from_file(str(cfg_b), force=True) + assert result_b.exit_code == 0, result_b.message + + # The deterministic comparison excludes `timestamp_utc` (wall + # clock); every other output file must be byte-identical. + for name in ( + "equity_curve.csv", + "orders.csv", + "fills.csv", + "positions.csv", + "costs.csv", + "summary.json", + "events.jsonl", + ): + a = (out_a / name).read_bytes() + b = (out_b / name).read_bytes() + assert a == b, f"{name} differs between runs" + meta_a = json.loads((out_a / "run_metadata.json").read_text()) + meta_b = json.loads((out_b / "run_metadata.json").read_text()) + # `timestamp_utc` is wall-clock; `config_output_directory` and + # `output_directory` reflect the per-run config path. Both are + # expected to differ; everything else must match. + nondeterministic = { + "timestamp_utc", + "config_output_directory", + "output_directory", + "config_path", + } + for k in meta_a: + if k in nondeterministic: + continue + assert meta_a[k] == meta_b[k], f"run_metadata {k} differs" diff --git a/tests/integration/test_real_data_suspended.py b/tests/integration/test_real_data_suspended.py new file mode 100644 index 0000000..ddd5a60 --- /dev/null +++ b/tests/integration/test_real_data_suspended.py @@ -0,0 +1,83 @@ +"""Integration scenario 3: universe with a suspended stock must not +crash; warnings must be traceable. + +Calibrated against the `~/.hqdata/tushare` snapshot for v0.1.1: + * 000008.SZ is suspended over 20260707..20260713 (7 trading days). + * Universe containing 000008.SZ must let the engine run, with the + suspension recorded via task-14's `DATA_WARNING` (fallback close + valuation) when the holding is in the suspended window. +""" + +from __future__ import annotations + +from decimal import Decimal +from pathlib import Path + +import pytest + +from hqbacktest import BacktestConfig, BacktestEngine, BaseStrategy +from hqbacktest.data import HqDataCsvPortal +from hqbacktest.domain.enums import EventType + +from ._harness import ( + CALIBRATION_END, + CALIBRATION_START, + skip_if_no_snapshot, +) + + +UNIVERSE = [ + "600000.SH", + "000008.SZ", # suspended 20260707..20260713 + "601318.SH", +] + + +@pytest.mark.integration +@skip_if_no_snapshot() +def test_universe_with_suspended_symbol_runs_and_warns(tmp_path: Path): + """Hold 000008.SZ through its suspended window and verify: + - the run completes (no crash); + - a DATA_WARNING is recorded for the suspended valuation; + - the audit trail pinpoints the symbol + window. + + The strategy buys 000008.SZ well before its suspension window + (20260701) and holds through 20260707..20260713. Task-14's + fallback-close valuation must then log a `DATA_WARNING` for the + suspended days. + """ + + class HoldSuspended(BaseStrategy): + def initialize(self, context): + context.set_universe(UNIVERSE) + + def on_bar(self, context, data): + if context.now == "20260701": + context.order_target_percent("000008.SZ", Decimal("0.30")) + + cfg = BacktestConfig( + start_date=CALIBRATION_START, + end_date=CALIBRATION_END, + initial_cash=Decimal("100000"), + source="tushare", + ) + portal = HqDataCsvPortal(source="tushare") + engine = BacktestEngine(cfg, strategy=HoldSuspended(), portal=portal) + result = engine.run() + # Run completes. + assert len(result.equity_curve) > 0 + # Audit-trail warning that mentions the suspended symbol. + warnings = [ + e + for e in engine.event_log.all() + if e.phase is EventType.DATA_WARNING and "000008.SZ" in (e.detail or "") + ] + assert warnings, ( + f"expected at least one DATA_WARNING mentioning 000008.SZ; got " + f"{[e.detail for e in warnings]}" + ) + # The fallback close valuation fires whenever a held symbol has + # no bar. The message mentions "fallback close". + assert any( + "fallback close" in (e.detail or "").lower() for e in warnings + ), "expected a 'fallback close' message in the warnings"