Develop #12
@@ -23,4 +23,6 @@ ZHIXING_MARKET_DATA_MAX_RETRIES=3
|
||||
ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS=1.0
|
||||
ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS=0.2
|
||||
ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY=7380521
|
||||
ZHIXING_SELECTION_MAX_WORKERS=4
|
||||
ZHIXING_SELECTION_BATCH_SIZE=200
|
||||
API_UPSTREAM=http://server:8000
|
||||
|
||||
+4
@@ -0,0 +1,4 @@
|
||||
{"file":".trellis/spec/backend/selection.md","reason":"检查批量执行后公式结果、状态矩阵、独立 signals、分页排序和重跑契约没有漂移。"}
|
||||
{"file":".trellis/spec/backend/quality-guidelines.md","reason":"检查类型、异常隔离、测试形状和全量质量命令是否满足后端规格。"}
|
||||
{"file":".trellis/spec/backend/configuration-and-runtime.md","reason":"检查 pool/worker 生命周期、Settings 环境变量和进程关闭行为。"}
|
||||
{"file":".trellis/spec/guides/code-reuse-thinking-guide.md","reason":"检查是否复用了已有池化 PostgreSQL 适配器模式,避免引入无必要抽象。"}
|
||||
+151
@@ -0,0 +1,151 @@
|
||||
# 技术设计:选股执行性能优化
|
||||
|
||||
## 1. 设计目标与边界
|
||||
|
||||
本次只优化 `selection` bounded context 的执行数据流,不改变 `ZhixingB1Strategy`
|
||||
的公式实现、`StockHistory` 的业务语义和现有 HTTP 结果契约。
|
||||
|
||||
保留以下外部事实:
|
||||
|
||||
- PostgreSQL 是行情和选股最终结果的事实源;
|
||||
- 选股运行仍由 `POST /api/v1/selection/runs` 创建并返回 `run_id`;
|
||||
- 每个股票仍有一个 `SelectionRunItem`,每类命中仍有一个独立
|
||||
`SelectionSignal`;
|
||||
- 单股评估异常继续隔离,不中断整个运行;
|
||||
- `(strategy, target_trade_date)` 的运行 claim、重跑保护和 advisory lock 保持不变。
|
||||
|
||||
不在本次设计中加入 Redis、外部任务队列、进程级任务恢复或新的前端进度协议。
|
||||
|
||||
## 2. 模块与接缝
|
||||
|
||||
### 2.1 应用模块
|
||||
|
||||
继续由 `RunZhixingB1` 作为深模块,对 presentation 暴露 prepare/execute/query
|
||||
能力。它内部增加三个私有阶段:
|
||||
|
||||
1. `load_histories`:按股票分块从 reader 读取历史;
|
||||
2. `evaluate_histories`:使用固定数量的线程 worker 调用现有
|
||||
`EvaluateZhixingB1.execute_history`;
|
||||
3. `record_items`:按结果分块提交 PostgreSQL。
|
||||
|
||||
公式策略本身仍只接收一个 `StockHistory`,避免把数据库、连接池和并发细节带入
|
||||
domain。
|
||||
|
||||
### 2.2 批量读取接口
|
||||
|
||||
为兼容已有单股调用者和测试 fake,在 `SelectionUniverseReader` 旁增加可选的批量
|
||||
历史读取扩展协议;生产适配器实现该扩展,执行器运行时优先探测批量方法,缺失时
|
||||
保留原单股 fallback。批量历史读取能力为:
|
||||
|
||||
```python
|
||||
load_histories(
|
||||
stocks: Sequence[SelectionStock],
|
||||
target_trade_date: date,
|
||||
) -> tuple[StockHistory, ...]
|
||||
```
|
||||
|
||||
`load_history(ts_code, target_trade_date)` 保留给单股调用者和兼容测试;批量执行路径
|
||||
不再调用它。
|
||||
|
||||
`PostgresMarketDataReader.load_histories` 使用一条参数化 SQL 读取一个分块:
|
||||
|
||||
- `bar.ts_code = ANY(%s)`;
|
||||
- `bar.source_adj = 'qfq'`;
|
||||
- `bar.trade_date <= %s`;
|
||||
- 按 `bar.ts_code, bar.trade_date` 升序返回;
|
||||
- 当前 B1 路径只读取股票名和 OHLCV,不再为历史每一行 LEFT JOIN
|
||||
`market_daily_basic`;目标日 basic 的完整性仍由执行源查询检查。
|
||||
|
||||
读取结果按股票代码分组,缺失代码返回空 `StockHistory`,由现有策略状态矩阵映射为
|
||||
`missing_target_bar` 或 `insufficient_history`。不会改变目标日截断、qfq 和排序契约。
|
||||
|
||||
分块大小由 `selection_batch_size` 控制,默认 200;每个分块读取完成、评估完成并写入
|
||||
后才释放历史对象,避免全市场历史同时驻留内存。
|
||||
|
||||
### 2.3 有界 PostgreSQL 资源
|
||||
|
||||
新增 selection infrastructure 的轻量资源 owner,内部持有一个
|
||||
`psycopg_pool.ConnectionPool`:
|
||||
|
||||
- `max_size = selection_max_workers + 2`,默认 6;
|
||||
- reader 和 run repository 共享同一个 pool;
|
||||
- pool 在进程内按 `(database_url, max_workers)` 缓存;
|
||||
- 首次使用时 open,进程退出时 close;
|
||||
- 测试通过构造函数注入 fake pool/connection。
|
||||
|
||||
这沿用市场数据模块现有的 pool 生命周期模式,不让每个 HTTP 请求或每只股票拥有
|
||||
独立 pool。读写 adapter 只借用短生命周期连接,业务事务仍由 adapter 控制。
|
||||
|
||||
### 2.4 批量写入接口
|
||||
|
||||
为兼容已有单项调用者和测试 fake,在 `SelectionRunStore` 旁增加可选的批量写入
|
||||
扩展协议:
|
||||
|
||||
```python
|
||||
record_items(run_id: str, items: Sequence[SelectionRunItem]) -> None
|
||||
```
|
||||
|
||||
`PostgresSelectionRunRepository.record_items` 在一个事务内:
|
||||
|
||||
1. 按分块股票代码删除该 run/股票已有 signal,保证重试幂等;
|
||||
2. 使用 psycopg connection cursor 的 `executemany` upsert 全部
|
||||
`selection_run_item`;
|
||||
3. 使用同一批量 API 插入全部独立 `selection_signal`;旧 fake connection 没有
|
||||
cursor 时保留逐条 execute fallback;
|
||||
4. 事务成功后返回。
|
||||
|
||||
原 `record_item` 保留为单项兼容 wrapper,并委托给 `record_items([item])`;应用批量
|
||||
路径不再调用它。新 run 的每个写入分块只提交一次事务,单股公式异常仍在应用层被
|
||||
转成 `data_error` 后进入该批次。
|
||||
|
||||
### 2.5 四 worker 执行模型
|
||||
|
||||
`RunZhixingB1` 构造时接收 `max_workers=4`,也可由 Settings 注入。每个 history 分块
|
||||
使用一个长期存在的 `ThreadPoolExecutor(max_workers=4)` 评估;worker 不直接写库。
|
||||
|
||||
选择线程而不是立即引入进程池的原因:
|
||||
|
||||
- 历史读取已经在应用线程按分块完成,无需跨进程复制数据库连接;
|
||||
- pandas/numpy 的一部分计算可以释放 GIL;
|
||||
- 线程共享只读 `StockHistory`,实现和回滚成本较低;
|
||||
- 若基准证明公式 CPU/GIL 成为主瓶颈,后续可以在同一 evaluator 接缝替换为进程
|
||||
worker,不影响 reader/store 契约。
|
||||
|
||||
每个分块保持如下状态流:
|
||||
|
||||
```text
|
||||
读取分块 → 4 worker 评估 → 聚合计数 → 一次批量写入 → 处理下一分块
|
||||
```
|
||||
|
||||
评估完成顺序不作为业务契约;查询端仍按股票代码和公式优先级稳定排序。
|
||||
|
||||
## 3. 配置与可观测性
|
||||
|
||||
新增 Settings 字段:
|
||||
|
||||
- `selection_max_workers: int = 4`,环境变量 `ZHIXING_SELECTION_MAX_WORKERS`;
|
||||
- `selection_batch_size: int = 200`,环境变量 `ZHIXING_SELECTION_BATCH_SIZE`。
|
||||
|
||||
执行结束时写一条安全的汇总日志,包含:run id、股票数、历史行数、分块数、worker
|
||||
数以及 read/evaluate/persist 的 wall-clock 秒数。不得写入连接串、token、行情详情
|
||||
或完整异常堆栈。
|
||||
|
||||
## 4. 兼容性与回滚
|
||||
|
||||
- 不新增或修改数据库表、索引和 HTTP 字段;无需 migration。
|
||||
- 如果批量读取/写入出现问题,可暂时将 `selection_batch_size=1`,保留相同接口并
|
||||
回到单股票批次;`max_workers=1` 可关闭并发以定位问题。
|
||||
- 通过 golden、应用层多分类测试和结果存储测试保证公式结果不漂移。
|
||||
- 任何 pool 获取失败都按现有 storage error 语义处理,不把半批次伪装成成功。
|
||||
|
||||
## 5. 关键取舍
|
||||
|
||||
| 选择 | 本次决定 | 原因 |
|
||||
| --- | --- | --- |
|
||||
| Redis | 不引入 | 当前问题是 PostgreSQL 往返和串行执行,已有持久化事实源足够 |
|
||||
| 全市场一次读取 | 不采用 | 六年历史乘以全股票池会增加内存峰值 |
|
||||
| 分块读取 | 采用,默认 200 | 控制内存并保留中间进度/失败隔离 |
|
||||
| 每股事务 | 不采用 | 事务数量随股票数线性增长 |
|
||||
| 分块事务 | 采用 | 减少提交次数,同时保留可控的部分进度 |
|
||||
| 无界并发 | 不采用 | 可能耗尽 PostgreSQL 连接和内存 |
|
||||
| 4 个线程 worker | 采用 | 用户确认的初始并发度,后续以基准调整 |
|
||||
+4
@@ -0,0 +1,4 @@
|
||||
{"file":".trellis/spec/backend/selection.md","reason":"保留 qfq、目标交易日、七类独立信号、run 重跑保护和批次状态契约。"}
|
||||
{"file":".trellis/spec/backend/configuration-and-runtime.md","reason":"新增 worker/批次配置并复用应用级资源生命周期、Settings 注入和进程关闭约定。"}
|
||||
{"file":".trellis/spec/backend/quality-guidelines.md","reason":"实现连接池、批量写入和并发后执行 Ruff、Pyright、pytest 质量门禁。"}
|
||||
{"file":".trellis/spec/guides/code-reuse-thinking-guide.md","reason":"优先复用现有市场数据 ConnectionPool 和应用组合模式,避免重复基础设施。"}
|
||||
+60
@@ -0,0 +1,60 @@
|
||||
# 实现计划:选股执行性能优化
|
||||
|
||||
## Phase 1:基线与配置
|
||||
|
||||
1. 增加 `selection_max_workers=4` 和 `selection_batch_size=200` 配置,并同步
|
||||
`.env.example`、开发/生产 compose 的可配置环境变量。
|
||||
2. 为运行执行器增加 read/evaluate/persist 的安全汇总计时;不要输出单股行情或
|
||||
凭据。
|
||||
3. 先运行现有 selection 测试,记录基线;确认工作区中没有用户并行改动。
|
||||
|
||||
## Phase 2:连接池与批量存储
|
||||
|
||||
4. 新增 selection PostgreSQL pool resource owner,支持注入 fake pool、open/close
|
||||
和借用连接;在 presentation 组合层按数据库配置缓存并注册退出清理。
|
||||
5. 改造 `PostgresMarketDataReader` 和 `PostgresSelectionRunRepository` 使用共享
|
||||
pool,同时保留直接构造/测试兼容路径。
|
||||
6. 在 `SelectionRunStore` 增加 `record_items`;实现 chunk 内一次事务、item
|
||||
`executemany`、signal `executemany` 和重试幂等删除。
|
||||
7. 更新应用层 FakeStore、PostgreSQL adapter tests,锁定每批写入的 SQL 数量和
|
||||
独立 signal 不丢失。
|
||||
|
||||
## Phase 3:批量读取与四 worker
|
||||
|
||||
8. 在 reader 的可选批量扩展和 `EvaluateZhixingB1` 测试 seam 中加入批量 history
|
||||
读取/已有 history 评估能力;保持单股 `execute` 兼容。
|
||||
9. 实现按 `ts_code` 分组的 qfq 批量 SQL,移除当前 B1 不使用的历史 daily-basic
|
||||
join/字段,保留执行源的目标日 basic 完整性校验。
|
||||
10. 将 `RunZhixingB1.execute` 改为 200 股票分块:批量读、4 worker 评估、聚合、批量
|
||||
写入;批量扩展缺失时 fallback 到旧单股/单项接口;保留单股异常隔离和最终状态统计。
|
||||
11. 增加批量读取、缺失 history、worker 异常、重复批次写入和结果顺序稳定性测试。
|
||||
|
||||
## Phase 4:质量门禁与实测
|
||||
|
||||
12. 运行 selection 单元/golden/HTTP 测试,确认七类独立信号和重跑契约不变。
|
||||
13. 如 PostgreSQL 环境可用,使用固定目标日运行一次真实 harness,记录 worker=1
|
||||
与 worker=4 的 read/evaluate/persist 以及总耗时;检查连接数没有超过 pool 上限。
|
||||
14. 运行完整后端格式、lint、type-check、pytest;必要时运行根目录 check/test。
|
||||
15. 检查 diff 只包含本任务文件,确认不包含 Redis 或无关前端改动。
|
||||
|
||||
## 主要验证命令
|
||||
|
||||
```bash
|
||||
cd zhixing-server
|
||||
uv run pytest tests/unit/selection tests/integration/test_zhixing_b1_golden.py -q
|
||||
uv run pytest tests/test_selection_http.py -q
|
||||
uv run ruff format --check .
|
||||
uv run ruff check .
|
||||
uv run pyright
|
||||
uv run pytest
|
||||
```
|
||||
|
||||
## 风险与回滚点
|
||||
|
||||
- pool 生命周期:先用 fake pool 测试 open/close,再接入 HTTP dependency;发生资源
|
||||
泄漏时回滚组合层缓存,不动公式。
|
||||
- 批量 SQL:先保持旧单股 reader/record wrapper,批量路径验证通过后再切换应用调用。
|
||||
- 线程 worker:先以 `max_workers=1` 验证结果等价,再使用默认 4;任何状态/信号差异
|
||||
都回滚并发切换,保留批量读取/写入的独立改动。
|
||||
- 内存:固定分块 200 并在每批完成后释放 histories;如果真实数据峰值过高,先调低
|
||||
`selection_batch_size`,不增加 Redis。
|
||||
@@ -0,0 +1,87 @@
|
||||
# 优化选股执行性能
|
||||
|
||||
## Goal
|
||||
|
||||
在不改变知行 B1 公式语义和结果持久化契约的前提下,降低全量选股执行的
|
||||
数据库连接、查询和事务开销,并默认使用 4 个受限 worker 并发评估股票。
|
||||
|
||||
用户价值:执行同一目标交易日的选股策略时,系统更快完成且仍能保留完整的
|
||||
逐股状态、失败原因和七类独立子信号。
|
||||
|
||||
## Background and Confirmed Facts
|
||||
|
||||
- `RunZhixingB1.execute` 当前按 `prepared.source.stocks` 串行逐股执行:
|
||||
[application/run.py:76-122](../../../zhixing-server/src/zhixing_server/modules/selection/application/run.py:76)。
|
||||
- `PostgresMarketDataReader.load_history` 每只股票建立一次直接 PostgreSQL
|
||||
连接并读取目标日前的全部 qfq 行情:
|
||||
[infrastructure/postgres_reader.py:25-163](../../../zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_reader.py:25)。
|
||||
- `PostgresSelectionRunRepository.record_item` 每只股票建立独立事务,并逐条
|
||||
插入其信号:
|
||||
[infrastructure/postgres_runs.py:139-197](../../../zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_runs.py:139)。
|
||||
- 市场数据模块已经有可复用的有上限 `psycopg_pool.ConnectionPool` 模式:
|
||||
[market_data/infrastructure/postgres.py:34-53](../../../zhixing-server/src/zhixing_server/modules/market_data/infrastructure/postgres.py:34)。
|
||||
- 知行 B1 需要保留七类独立子信号;历史 golden 和现有测试是结果兼容基线。
|
||||
- 本任务不引入 Redis,不改选股公式规则,不实现独立任务队列或新的进度 API。
|
||||
|
||||
## Requirements
|
||||
|
||||
### R1. 可观测的性能基线
|
||||
|
||||
- 为一次运行记录可复现的阶段指标:有效股票数、历史行数、读取耗时、公式
|
||||
评估耗时、持久化耗时、worker 数和批次数。
|
||||
- 指标不能输出数据库 URL、密码、Tushare token 或单股完整行情。
|
||||
- 运行结果、信号分类和失败状态仍以 PostgreSQL 持久化结果为准。
|
||||
|
||||
### R2. 连接池化
|
||||
|
||||
- 选股读取适配器和结果存储适配器使用应用生命周期内的有上限 PostgreSQL
|
||||
连接池,不再为每只股票创建和销毁连接。
|
||||
- 连接池大小必须覆盖 4 个 worker、批量写入和必要的查询余量,不能无限增长。
|
||||
- 应用关闭时可靠关闭连接池;单元测试可以注入 fake pool/connection。
|
||||
|
||||
### R3. 批量历史读取
|
||||
|
||||
- 保持 qfq、目标交易日截断、升序日期和六年数据保留语义不变。
|
||||
- 将逐股历史读取改为按股票批量/分块读取;分块大小可配置但必须有默认上限,
|
||||
防止一次性把全市场历史全部载入内存。
|
||||
- 当前 B1 公式不使用历史 `turnover_rate` 和 `total_mv`;本任务可以移除历史
|
||||
查询中不必要的 daily-basic 字段/连接,但目标日数据完整性校验必须保留。
|
||||
|
||||
### R4. 批量结果写入
|
||||
|
||||
- 将逐股 `record_item` 改为按批次写入 `selection_run_item` 和
|
||||
`selection_signal`,默认批次大小为 200,且保留逐股评估异常隔离。
|
||||
- 批次提交失败时不能标记为成功;运行最终状态必须正确收敛为
|
||||
`success`、`partial_success` 或 `failed`。
|
||||
- 结果查询、重跑唯一性和七类独立信号身份不变。
|
||||
|
||||
### R5. 四 worker 有界并发
|
||||
|
||||
- 默认使用 4 个 worker;并发度必须可配置且至少为 1。
|
||||
- worker 不得无限创建连接、线程或进程;数据库连接数受池上限约束。
|
||||
- 同一股票只评估一次;结果顺序不作为业务契约,API 查询仍按既有稳定规则排序。
|
||||
- 公式计算失败仍只影响该股票,不能丢失其他股票结果。
|
||||
|
||||
## Out of Scope
|
||||
|
||||
- Redis、Celery/RQ/Arq 等外部任务队列或独立 worker 服务。
|
||||
- 前端轮询协议和执行进度接口重构。
|
||||
- 公式阈值、指标定义、历史窗口、股票池范围和信号分类调整。
|
||||
- 结果表结构的大规模迁移或删除历史结果。
|
||||
|
||||
## Acceptance Criteria
|
||||
|
||||
- [ ] 在固定历史 fixture 上,现有 golden、七类独立信号和选股状态全部保持一致。
|
||||
- [ ] 单元测试覆盖连接池生命周期、批量历史按股票分组、批量写入、4 worker
|
||||
并发上限、单股异常隔离和批次失败状态收敛。
|
||||
- [ ] 一次选股运行不再产生每股一次的数据库连接;读取和写入均通过连接池。
|
||||
- [ ] 一次选股运行不再为每只股票单独提交结果事务;结果按批次提交。
|
||||
- [ ] 运行日志/基准输出包含 R1 指标,并能区分 read/evaluate/persist 三段耗时。
|
||||
- [ ] 使用实际 PostgreSQL 数据或等价可复现 harness 验证 4 worker 下运行成功,
|
||||
且没有出现超出连接池上限的连接创建。
|
||||
- [ ] 选股 HTTP 契约、重跑保护、最终分页结果和失败列表相关测试通过。
|
||||
- [ ] 不引入 Redis,工作区只包含本任务相关改动。
|
||||
|
||||
## Open Questions
|
||||
|
||||
无阻塞问题。批量大小默认 200,worker 默认 4;两者保持配置化,后续以基准结果调整。
|
||||
+26
@@ -0,0 +1,26 @@
|
||||
{
|
||||
"id": "optimize-selection-execution-performance",
|
||||
"name": "optimize-selection-execution-performance",
|
||||
"title": "优化选股执行性能",
|
||||
"description": "",
|
||||
"status": "completed",
|
||||
"dev_type": null,
|
||||
"scope": null,
|
||||
"package": null,
|
||||
"priority": "P2",
|
||||
"creator": "yuxuanhui",
|
||||
"assignee": "yuxuanhui",
|
||||
"createdAt": "2026-08-12",
|
||||
"completedAt": "2026-08-12",
|
||||
"branch": null,
|
||||
"base_branch": "main",
|
||||
"worktree_path": null,
|
||||
"commit": null,
|
||||
"pr_url": null,
|
||||
"subtasks": [],
|
||||
"children": [],
|
||||
"parent": null,
|
||||
"relatedFiles": [],
|
||||
"notes": "",
|
||||
"meta": {}
|
||||
}
|
||||
@@ -8,8 +8,8 @@
|
||||
|
||||
<!-- @@@auto:current-status -->
|
||||
- **Active File**: `journal-1.md`
|
||||
- **Total Sessions**: 8
|
||||
- **Last Active**: 2026-08-11
|
||||
- **Total Sessions**: 9
|
||||
- **Last Active**: 2026-08-12
|
||||
<!-- @@@/auto:current-status -->
|
||||
|
||||
---
|
||||
@@ -19,7 +19,7 @@
|
||||
<!-- @@@auto:active-documents -->
|
||||
| File | Lines | Status |
|
||||
|------|-------|--------|
|
||||
| `journal-1.md` | ~222 | Active |
|
||||
| `journal-1.md` | ~243 | Active |
|
||||
<!-- @@@/auto:active-documents -->
|
||||
|
||||
---
|
||||
@@ -29,6 +29,7 @@
|
||||
<!-- @@@auto:session-history -->
|
||||
| # | Date | Title | Commits | Branch |
|
||||
|---|------|-------|---------|--------|
|
||||
| 9 | 2026-08-12 | 完成选股执行性能优化 | `8963c06` | `develop` |
|
||||
| 8 | 2026-08-11 | 完成市场数据同步与完整性检查 | `7ce1154`, `8f5f504` | `develop` |
|
||||
| 7 | 2026-08-10 | 完成选股执行状态抽屉与紧凑布局 | `17237e0` | `develop` |
|
||||
| 6 | 2026-08-10 | 按原型完善选股结果分页接口 | `ed7bdda`, `3af97bf` | `develop` |
|
||||
|
||||
@@ -220,3 +220,24 @@
|
||||
### Next Steps
|
||||
|
||||
- 配置 ZHIXING_TEST_DATABASE_URL 后执行真实 PostgreSQL 集成验证
|
||||
|
||||
|
||||
## Session 9: 完成选股执行性能优化
|
||||
|
||||
**Date**: 2026-08-12
|
||||
**Task**: 完成选股执行性能优化
|
||||
**Branch**: `develop`
|
||||
|
||||
### Summary
|
||||
|
||||
完成共享 PostgreSQL 连接池、批量历史读取、4 worker 评估、批量结果写入和阶段耗时日志;全量测试 82 passed、2 skipped,Ruff 与 Pyright 通过。
|
||||
|
||||
### Git Commits
|
||||
|
||||
| Hash | Message |
|
||||
|------|---------|
|
||||
| `8963c06` | (see git log) |
|
||||
|
||||
### Status
|
||||
|
||||
[OK] **Completed**
|
||||
|
||||
@@ -36,6 +36,8 @@ services:
|
||||
ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS: ${ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS:-1.0}
|
||||
ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS: ${ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS:-0.2}
|
||||
ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY: ${ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY:-7380521}
|
||||
ZHIXING_SELECTION_MAX_WORKERS: ${ZHIXING_SELECTION_MAX_WORKERS:-4}
|
||||
ZHIXING_SELECTION_BATCH_SIZE: ${ZHIXING_SELECTION_BATCH_SIZE:-200}
|
||||
ZHIXING_TUSHARE_TOKEN: ${ZHIXING_TUSHARE_TOKEN:-}
|
||||
init: true
|
||||
ports:
|
||||
|
||||
@@ -18,6 +18,8 @@ services:
|
||||
ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS: ${ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS:-1.0}
|
||||
ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS: ${ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS:-0.2}
|
||||
ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY: ${ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY:-7380521}
|
||||
ZHIXING_SELECTION_MAX_WORKERS: ${ZHIXING_SELECTION_MAX_WORKERS:-4}
|
||||
ZHIXING_SELECTION_BATCH_SIZE: ${ZHIXING_SELECTION_BATCH_SIZE:-200}
|
||||
ZHIXING_TUSHARE_TOKEN: ${ZHIXING_TUSHARE_TOKEN:-}
|
||||
init: true
|
||||
expose:
|
||||
|
||||
@@ -24,6 +24,8 @@ class Settings(BaseSettings):
|
||||
market_data_max_retries: int = 3
|
||||
market_data_retry_backoff_seconds: float = 1.0
|
||||
market_data_advisory_lock_key: int = 7_380_521
|
||||
selection_max_workers: int = Field(default=4, ge=1)
|
||||
selection_batch_size: int = Field(default=200, ge=1)
|
||||
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from datetime import date
|
||||
|
||||
from ..domain.models import SelectionEvaluation, StockHistory
|
||||
@@ -44,3 +45,12 @@ class EvaluateZhixingB1:
|
||||
"""Evaluate an already loaded history for deterministic unit tests."""
|
||||
|
||||
return self.strategy.evaluate(history, target_trade_date)
|
||||
|
||||
def execute_histories(
|
||||
self,
|
||||
histories: Sequence[StockHistory],
|
||||
target_trade_date: date,
|
||||
) -> tuple[SelectionEvaluation, ...]:
|
||||
"""Evaluate loaded histories without issuing one read per stock."""
|
||||
|
||||
return tuple(self.execute_history(history, target_trade_date) for history in histories)
|
||||
|
||||
@@ -3,12 +3,17 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Callable, Sequence
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from datetime import date
|
||||
from typing import Literal, Protocol
|
||||
from typing import Literal, Protocol, cast
|
||||
|
||||
from ..domain.models import SelectionEvaluation
|
||||
from ..domain.models import SelectionEvaluation, StockHistory
|
||||
from ..domain.runs import (
|
||||
BatchSelectionRunStore,
|
||||
BatchSelectionUniverseReader,
|
||||
SelectionExecutionSource,
|
||||
SelectionRerunRequired,
|
||||
SelectionResultQuery,
|
||||
@@ -17,6 +22,7 @@ from ..domain.runs import (
|
||||
SelectionRunItem,
|
||||
SelectionRunStatus,
|
||||
SelectionRunStore,
|
||||
SelectionStock,
|
||||
SelectionUniverseReader,
|
||||
)
|
||||
from .evaluate import EvaluateZhixingB1
|
||||
@@ -48,12 +54,21 @@ class RunZhixingB1:
|
||||
reader: SelectionUniverseReader,
|
||||
store: SelectionRunStore,
|
||||
evaluator: SelectionEvaluator | None = None,
|
||||
*,
|
||||
max_workers: int = 4,
|
||||
batch_size: int = 200,
|
||||
) -> None:
|
||||
"""Inject storage ports and optionally a test evaluator."""
|
||||
"""Inject storage ports and configure bounded chunk execution."""
|
||||
|
||||
if max_workers < 1:
|
||||
raise ValueError("max_workers must be at least 1")
|
||||
if batch_size < 1:
|
||||
raise ValueError("batch_size must be at least 1")
|
||||
self.reader = reader
|
||||
self.store = store
|
||||
self.evaluator = evaluator or EvaluateZhixingB1(reader)
|
||||
self.max_workers = max_workers
|
||||
self.batch_size = batch_size
|
||||
|
||||
def prepare(
|
||||
self,
|
||||
@@ -81,35 +96,57 @@ class RunZhixingB1:
|
||||
returns so the UI never mistakes a lost worker exception for success.
|
||||
"""
|
||||
|
||||
stocks = _unique_stocks(prepared.source.stocks)
|
||||
evaluated_count = 0
|
||||
selected_stock_count = 0
|
||||
signal_count = 0
|
||||
failed_count = 0
|
||||
history_rows = 0
|
||||
batch_count = _chunk_count(len(stocks), self.batch_size)
|
||||
read_seconds = 0.0
|
||||
evaluate_seconds = 0.0
|
||||
persist_seconds = 0.0
|
||||
try:
|
||||
for stock in prepared.source.stocks:
|
||||
try:
|
||||
evaluation = self.evaluator.execute(
|
||||
stock.ts_code,
|
||||
with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
|
||||
for batch_stocks in _chunks(stocks, self.batch_size):
|
||||
read_started = time.perf_counter()
|
||||
histories = self._load_histories(
|
||||
batch_stocks,
|
||||
prepared.source.target_trade_date,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 - isolate one stock from the batch
|
||||
logger.exception(
|
||||
"selection_item_failed run_id=%s ts_code=%s",
|
||||
prepared.run.id,
|
||||
stock.ts_code,
|
||||
read_seconds += time.perf_counter() - read_started
|
||||
history_rows += sum(
|
||||
len(history.bars) for history in histories if history is not None
|
||||
)
|
||||
evaluation = SelectionEvaluation(
|
||||
ts_code=stock.ts_code,
|
||||
target_trade_date=prepared.source.target_trade_date,
|
||||
status="data_error",
|
||||
reason=_safe_item_error(exc),
|
||||
|
||||
evaluate_started = time.perf_counter()
|
||||
items = tuple(
|
||||
_to_item(
|
||||
stock.ts_code,
|
||||
stock.name,
|
||||
evaluation,
|
||||
)
|
||||
for stock, evaluation in zip(
|
||||
batch_stocks,
|
||||
executor.map(
|
||||
self._evaluate_stock,
|
||||
batch_stocks,
|
||||
histories,
|
||||
[prepared.source.target_trade_date] * len(batch_stocks),
|
||||
),
|
||||
strict=True,
|
||||
)
|
||||
)
|
||||
item = _to_item(stock.ts_code, stock.name, evaluation)
|
||||
self.store.record_item(prepared.run.id, item)
|
||||
evaluated_count += 1
|
||||
selected_stock_count += evaluation.status == "selected"
|
||||
signal_count += len(evaluation.signals)
|
||||
failed_count += evaluation.status in _FAILURE_STATUSES
|
||||
evaluate_seconds += time.perf_counter() - evaluate_started
|
||||
|
||||
evaluated_count += len(items)
|
||||
selected_stock_count += sum(item.status == "selected" for item in items)
|
||||
signal_count += sum(item.signal_count for item in items)
|
||||
failed_count += sum(item.status in _FAILURE_STATUSES for item in items)
|
||||
|
||||
persist_started = time.perf_counter()
|
||||
self._record_items(prepared.run.id, items)
|
||||
persist_seconds += time.perf_counter() - persist_started
|
||||
|
||||
status = _run_status(evaluated_count, failed_count)
|
||||
self.store.finish_run(
|
||||
@@ -121,7 +158,12 @@ class RunZhixingB1:
|
||||
failed_count=failed_count,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 - worker boundary must persist failure state
|
||||
logger.exception("selection_run_failed run_id=%s", prepared.run.id)
|
||||
logger.error(
|
||||
"selection_run_failed run_id=%s error_type=%s reason=%s",
|
||||
prepared.run.id,
|
||||
exc.__class__.__name__,
|
||||
_safe_item_error(exc),
|
||||
)
|
||||
try:
|
||||
self.store.finish_run(
|
||||
prepared.run.id,
|
||||
@@ -134,7 +176,94 @@ class RunZhixingB1:
|
||||
error_message=str(exc),
|
||||
)
|
||||
except Exception: # noqa: BLE001 - preserve the original worker failure
|
||||
logger.exception("selection_run_failure_persist_failed run_id=%s", prepared.run.id)
|
||||
logger.error(
|
||||
"selection_run_failure_persist_failed run_id=%s",
|
||||
prepared.run.id,
|
||||
)
|
||||
finally:
|
||||
logger.info(
|
||||
"selection_run_summary run_id=%s stock_count=%d history_rows=%d "
|
||||
"batch_count=%d worker_count=%d read_seconds=%.3f "
|
||||
"evaluate_seconds=%.3f persist_seconds=%.3f",
|
||||
prepared.run.id,
|
||||
len(stocks),
|
||||
history_rows,
|
||||
batch_count,
|
||||
self.max_workers,
|
||||
read_seconds,
|
||||
evaluate_seconds,
|
||||
persist_seconds,
|
||||
)
|
||||
|
||||
def _load_histories(
|
||||
self,
|
||||
stocks: Sequence[SelectionStock],
|
||||
target_trade_date: date,
|
||||
) -> tuple[StockHistory | None, ...]:
|
||||
"""Load one chunk when the reader supports it, with old-path fallback."""
|
||||
|
||||
typed_stocks = tuple(stocks)
|
||||
loader = getattr(self.reader, "load_histories", None)
|
||||
if callable(loader):
|
||||
batch_reader = cast(BatchSelectionUniverseReader, self.reader)
|
||||
loaded = batch_reader.load_histories(typed_stocks, target_trade_date)
|
||||
histories_by_code = {history.ts_code: history for history in loaded}
|
||||
return tuple(
|
||||
histories_by_code.get(
|
||||
stock.ts_code,
|
||||
StockHistory(ts_code=stock.ts_code, name=stock.name),
|
||||
)
|
||||
for stock in typed_stocks
|
||||
)
|
||||
|
||||
if isinstance(self.evaluator, EvaluateZhixingB1):
|
||||
return tuple(
|
||||
self.reader.load_history(stock.ts_code, target_trade_date) for stock in typed_stocks
|
||||
)
|
||||
return (None,) * len(typed_stocks)
|
||||
|
||||
def _evaluate_stock(
|
||||
self,
|
||||
stock: SelectionStock,
|
||||
history: StockHistory | None,
|
||||
target_trade_date: date,
|
||||
) -> SelectionEvaluation:
|
||||
"""Evaluate one stock inside a worker and isolate its exception."""
|
||||
|
||||
ts_code = stock.ts_code
|
||||
try:
|
||||
execute_history: Callable[[StockHistory, date], SelectionEvaluation] | None = getattr(
|
||||
self.evaluator,
|
||||
"execute_history",
|
||||
None,
|
||||
)
|
||||
if history is not None and execute_history is not None:
|
||||
return execute_history(history, target_trade_date)
|
||||
return self.evaluator.execute(ts_code, target_trade_date)
|
||||
except Exception as exc: # noqa: BLE001 - isolate one stock from the batch
|
||||
logger.warning(
|
||||
"selection_item_failed ts_code=%s error_type=%s reason=%s",
|
||||
ts_code,
|
||||
exc.__class__.__name__,
|
||||
_safe_item_error(exc),
|
||||
)
|
||||
return SelectionEvaluation(
|
||||
ts_code=ts_code,
|
||||
target_trade_date=target_trade_date,
|
||||
status="data_error",
|
||||
reason=_safe_item_error(exc),
|
||||
)
|
||||
|
||||
def _record_items(self, run_id: str, items: Sequence[SelectionRunItem]) -> None:
|
||||
"""Use batch persistence while retaining the old single-item seam."""
|
||||
|
||||
record_items = getattr(self.store, "record_items", None)
|
||||
if callable(record_items):
|
||||
batch_store = cast(BatchSelectionRunStore, self.store)
|
||||
batch_store.record_items(run_id, tuple(items))
|
||||
return
|
||||
for item in items:
|
||||
self.store.record_item(run_id, item)
|
||||
|
||||
def get_run(
|
||||
self,
|
||||
@@ -187,6 +316,35 @@ def _safe_item_error(error: Exception) -> str:
|
||||
return " ".join(str(error).split())[:500] or error.__class__.__name__
|
||||
|
||||
|
||||
def _chunks(
|
||||
values: Sequence[SelectionStock],
|
||||
size: int,
|
||||
) -> tuple[tuple[SelectionStock, ...], ...]:
|
||||
"""Split a stable stock sequence into bounded immutable chunks."""
|
||||
|
||||
return tuple(tuple(values[index : index + size]) for index in range(0, len(values), size))
|
||||
|
||||
|
||||
def _chunk_count(value_count: int, size: int) -> int:
|
||||
"""Return the number of chunks without materializing empty chunks."""
|
||||
|
||||
return (value_count + size - 1) // size
|
||||
|
||||
|
||||
def _unique_stocks(stocks: Sequence[SelectionStock]) -> tuple[SelectionStock, ...]:
|
||||
"""Keep the first source row for each stock so it is evaluated once."""
|
||||
|
||||
seen: set[str] = set()
|
||||
unique: list[SelectionStock] = []
|
||||
for stock in stocks:
|
||||
ts_code = stock.ts_code
|
||||
if ts_code in seen:
|
||||
continue
|
||||
seen.add(ts_code)
|
||||
unique.append(stock)
|
||||
return tuple(unique)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"PreparedSelectionRun",
|
||||
"RunZhixingB1",
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
@@ -150,3 +151,19 @@ class SelectionUniverseReader(Protocol):
|
||||
) -> SelectionExecutionSource: ...
|
||||
|
||||
def load_history(self, ts_code: str, target_trade_date: date) -> StockHistory: ...
|
||||
|
||||
|
||||
class BatchSelectionRunStore(Protocol):
|
||||
"""Optional batch-write extension for selection stores."""
|
||||
|
||||
def record_items(self, run_id: str, items: Sequence[SelectionRunItem]) -> None: ...
|
||||
|
||||
|
||||
class BatchSelectionUniverseReader(Protocol):
|
||||
"""Optional bounded batch-history extension for selection readers."""
|
||||
|
||||
def load_histories(
|
||||
self,
|
||||
stocks: Sequence[SelectionStock],
|
||||
target_trade_date: date,
|
||||
) -> tuple[StockHistory, ...]: ...
|
||||
|
||||
@@ -0,0 +1,96 @@
|
||||
"""Bounded PostgreSQL pool ownership for selection adapters."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Protocol, cast
|
||||
|
||||
from psycopg_pool import ConnectionPool
|
||||
|
||||
|
||||
class SelectionConnectionPool(Protocol):
|
||||
"""Small pool surface shared by psycopg and unit-test fakes."""
|
||||
|
||||
def open(self, *, wait: bool = True) -> None: ...
|
||||
|
||||
def close(self) -> None: ...
|
||||
|
||||
def connection(self) -> Any: ...
|
||||
|
||||
|
||||
class SelectionPostgresPool:
|
||||
"""Own one bounded PostgreSQL pool for the selection read/write adapters.
|
||||
|
||||
The owner opens lazily on the first borrowed connection, which keeps app
|
||||
construction cheap while still making the pool lifetime process-scoped in
|
||||
the HTTP composition layer. A fake pool can be injected for unit tests.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
database_url: str,
|
||||
*,
|
||||
max_connections: int,
|
||||
pool: SelectionConnectionPool | None = None,
|
||||
) -> None:
|
||||
"""Create a bounded pool owner with an optional injected pool."""
|
||||
|
||||
if max_connections < 1:
|
||||
raise ValueError("max_connections must be at least 1")
|
||||
self.database_url = database_url
|
||||
self.max_connections = max_connections
|
||||
self.pool: SelectionConnectionPool = pool or cast(
|
||||
SelectionConnectionPool,
|
||||
ConnectionPool(
|
||||
conninfo=database_url,
|
||||
min_size=1,
|
||||
max_size=max_connections,
|
||||
open=False,
|
||||
),
|
||||
)
|
||||
self._pool_open = False
|
||||
self._pool_state_lock = threading.Lock()
|
||||
|
||||
def open(self) -> None:
|
||||
"""Open the underlying pool once and wait for its minimum connection."""
|
||||
|
||||
with self._pool_state_lock:
|
||||
if self._pool_open:
|
||||
return
|
||||
if bool(getattr(self.pool, "_opened", False)):
|
||||
self._pool_open = True
|
||||
return
|
||||
self.pool.open(wait=True)
|
||||
self._pool_open = True
|
||||
|
||||
def close(self) -> None:
|
||||
"""Close the pool after borrowed connections have been returned."""
|
||||
|
||||
with self._pool_state_lock:
|
||||
if self._pool_open or bool(getattr(self.pool, "_opened", False)):
|
||||
self.pool.close()
|
||||
self._pool_open = False
|
||||
|
||||
@contextmanager
|
||||
def connection(self) -> Generator[Any, None, None]:
|
||||
"""Borrow one connection and return it to the bounded pool."""
|
||||
|
||||
self.open()
|
||||
with self.pool.connection() as connection:
|
||||
yield connection
|
||||
|
||||
def __enter__(self) -> SelectionPostgresPool:
|
||||
"""Open and return this resource owner."""
|
||||
|
||||
self.open()
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type: object, exc_value: object, traceback: object) -> None:
|
||||
"""Release the owned pool at the end of a context."""
|
||||
|
||||
self.close()
|
||||
|
||||
|
||||
__all__ = ["SelectionConnectionPool", "SelectionPostgresPool"]
|
||||
+166
-34
@@ -2,9 +2,11 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Generator, Sequence
|
||||
from contextlib import contextmanager
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal, InvalidOperation
|
||||
from typing import cast
|
||||
from typing import Any, cast
|
||||
|
||||
import psycopg
|
||||
|
||||
@@ -12,6 +14,7 @@ from ....bootstrap.config import Settings
|
||||
from ..domain.models import SelectionBar, SelectionDailyBasic, StockHistory
|
||||
from ..domain.ports import MarketDataReaderError
|
||||
from ..domain.runs import SelectionExecutionSource, SelectionStock
|
||||
from .postgres_pool import SelectionConnectionPool, SelectionPostgresPool
|
||||
|
||||
|
||||
class SelectionReaderError(MarketDataReaderError):
|
||||
@@ -23,6 +26,25 @@ class SelectionMarketDataNotReady(MarketDataReaderError):
|
||||
|
||||
|
||||
_HISTORY_QUERY = """
|
||||
SELECT
|
||||
bar.ts_code,
|
||||
stock.name,
|
||||
bar.trade_date,
|
||||
bar.open,
|
||||
bar.high,
|
||||
bar.low,
|
||||
bar.close,
|
||||
bar.vol
|
||||
FROM market_daily_bar AS bar
|
||||
LEFT JOIN market_stock AS stock
|
||||
ON stock.ts_code = bar.ts_code
|
||||
WHERE bar.ts_code = ANY(%s)
|
||||
AND bar.source_adj = 'qfq'
|
||||
AND bar.trade_date <= %s
|
||||
ORDER BY bar.ts_code ASC, bar.trade_date ASC
|
||||
"""
|
||||
|
||||
_SINGLE_HISTORY_QUERY = """
|
||||
SELECT
|
||||
bar.ts_code,
|
||||
stock.name,
|
||||
@@ -111,10 +133,30 @@ def _as_float(value: object) -> float | None:
|
||||
class PostgresMarketDataReader:
|
||||
"""Load qfq bars and same-day basic facts without writing market data."""
|
||||
|
||||
def __init__(self, settings: Settings | str) -> None:
|
||||
"""Create a reader from injected settings or a compatible URL string."""
|
||||
def __init__(
|
||||
self,
|
||||
settings: Settings | str,
|
||||
*,
|
||||
pool: SelectionPostgresPool | SelectionConnectionPool | None = None,
|
||||
) -> None:
|
||||
"""Create a reader from settings/URL and an optional shared pool.
|
||||
|
||||
Omitting ``pool`` intentionally retains the direct ``psycopg.connect``
|
||||
path used by one-shot callers and existing adapter tests. The HTTP
|
||||
composition layer always supplies the process-scoped selection pool.
|
||||
"""
|
||||
|
||||
self.database_url = settings.database_url if isinstance(settings, Settings) else settings
|
||||
if isinstance(pool, SelectionPostgresPool):
|
||||
self.pool: SelectionPostgresPool | None = pool
|
||||
elif pool is not None:
|
||||
self.pool = SelectionPostgresPool(
|
||||
self.database_url,
|
||||
max_connections=1,
|
||||
pool=pool,
|
||||
)
|
||||
else:
|
||||
self.pool = None
|
||||
|
||||
def load_history(self, ts_code: str, target_trade_date: date) -> StockHistory:
|
||||
"""Read all retained qfq rows through the explicit target date.
|
||||
@@ -133,34 +175,57 @@ class PostgresMarketDataReader:
|
||||
"""
|
||||
|
||||
try:
|
||||
with psycopg.connect(self.database_url) as connection:
|
||||
with self._connection() as connection:
|
||||
rows = connection.execute(
|
||||
_HISTORY_QUERY,
|
||||
_SINGLE_HISTORY_QUERY,
|
||||
(ts_code, target_trade_date),
|
||||
).fetchall()
|
||||
except psycopg.Error as exc:
|
||||
except SelectionReaderError:
|
||||
raise
|
||||
except Exception as exc: # noqa: BLE001 - redact driver/pool details at the port boundary
|
||||
raise SelectionReaderError(
|
||||
f"failed to load market history for {ts_code} at {target_trade_date.isoformat()}"
|
||||
) from exc
|
||||
|
||||
bars: dict[date, SelectionBar] = {}
|
||||
daily_basic: dict[date, SelectionDailyBasic] = {}
|
||||
name = ""
|
||||
for raw_row in rows:
|
||||
row = cast(tuple[object, ...], raw_row)
|
||||
row_code, row_name, bar, basic = self._map_row(row, ts_code)
|
||||
if row_code != ts_code:
|
||||
raise ValueError(f"reader returned unexpected stock code: {row_code}")
|
||||
name = row_name or name
|
||||
if bar.trade_date <= target_trade_date:
|
||||
bars[bar.trade_date] = bar
|
||||
daily_basic[bar.trade_date] = basic
|
||||
return StockHistory(
|
||||
ts_code=ts_code,
|
||||
name=name,
|
||||
bars=tuple(bars[trade_date] for trade_date in sorted(bars)),
|
||||
daily_basic={trade_date: daily_basic[trade_date] for trade_date in sorted(daily_basic)},
|
||||
return self._histories_from_rows(
|
||||
rows,
|
||||
(SelectionStock(ts_code, ""),),
|
||||
target_trade_date,
|
||||
)[0]
|
||||
|
||||
def load_histories(
|
||||
self,
|
||||
stocks: Sequence[SelectionStock] | Sequence[str],
|
||||
target_trade_date: date,
|
||||
) -> tuple[StockHistory, ...]:
|
||||
"""Read one bounded stock chunk with one parameterized qfq query.
|
||||
|
||||
Historical daily-basic values are deliberately not joined here: B1
|
||||
only needs OHLCV for its historical formula. The execution-source
|
||||
query still requires a complete target-day basic row before a stock is
|
||||
admitted to a run.
|
||||
"""
|
||||
|
||||
normalized = tuple(
|
||||
stock if isinstance(stock, SelectionStock) else SelectionStock(stock, "")
|
||||
for stock in stocks
|
||||
)
|
||||
if not normalized:
|
||||
return ()
|
||||
codes = [stock.ts_code for stock in normalized]
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
rows = connection.execute(
|
||||
_HISTORY_QUERY,
|
||||
(codes, target_trade_date),
|
||||
).fetchall()
|
||||
except SelectionReaderError:
|
||||
raise
|
||||
except Exception as exc: # noqa: BLE001 - redact driver/pool details at the port boundary
|
||||
raise SelectionReaderError(
|
||||
f"failed to load market history batch at {target_trade_date.isoformat()}"
|
||||
) from exc
|
||||
return self._histories_from_rows(rows, normalized, target_trade_date)
|
||||
|
||||
def load_execution_source(
|
||||
self,
|
||||
@@ -187,7 +252,7 @@ class PostgresMarketDataReader:
|
||||
if strategy != "zhixing_b1":
|
||||
raise SelectionMarketDataNotReady(f"unsupported selection strategy: {strategy}")
|
||||
try:
|
||||
with psycopg.connect(self.database_url) as connection:
|
||||
with self._connection() as connection:
|
||||
source_row = connection.execute(_SOURCE_QUERY, (target_trade_date,)).fetchone()
|
||||
if source_row is None:
|
||||
raise SelectionMarketDataNotReady(
|
||||
@@ -197,9 +262,9 @@ class PostgresMarketDataReader:
|
||||
_ELIGIBLE_STOCKS_QUERY,
|
||||
(target_trade_date, target_trade_date),
|
||||
).fetchall()
|
||||
except SelectionMarketDataNotReady:
|
||||
except (SelectionMarketDataNotReady, SelectionReaderError):
|
||||
raise
|
||||
except psycopg.Error as exc:
|
||||
except Exception as exc: # noqa: BLE001 - redact driver/pool details at the port boundary
|
||||
raise SelectionReaderError(
|
||||
f"failed to load selection source at {target_trade_date.isoformat()}"
|
||||
) from exc
|
||||
@@ -224,8 +289,8 @@ class PostgresMarketDataReader:
|
||||
def _map_row(
|
||||
row: tuple[object, ...],
|
||||
expected_code: str,
|
||||
) -> tuple[str, str, SelectionBar, SelectionDailyBasic]:
|
||||
"""Map the current query row, tolerating a legacy test row without name."""
|
||||
) -> tuple[str, str, SelectionBar, SelectionDailyBasic | None]:
|
||||
"""Map qfq OHLCV rows and tolerate the legacy basic-join test shape."""
|
||||
|
||||
if len(row) >= 10:
|
||||
code, raw_name, raw_date = row[0], row[1], row[2]
|
||||
@@ -233,13 +298,19 @@ class PostgresMarketDataReader:
|
||||
elif len(row) >= 9:
|
||||
code, raw_name, raw_date = row[0], "", row[1]
|
||||
values = row[2:]
|
||||
elif len(row) >= 8:
|
||||
code, raw_name, raw_date = row[0], row[1], row[2]
|
||||
values = row[3:]
|
||||
elif len(row) >= 7:
|
||||
code, raw_name, raw_date = row[0], "", row[1]
|
||||
values = row[2:]
|
||||
else:
|
||||
raise ValueError("market history row has too few columns")
|
||||
row_code = str(code or expected_code)
|
||||
name = str(raw_name or "")
|
||||
trade_date = _as_date(raw_date)
|
||||
if len(values) < 7:
|
||||
raise ValueError("market history row is missing OHLCV/basic columns")
|
||||
if len(values) < 5:
|
||||
raise ValueError("market history row is missing OHLCV columns")
|
||||
bar = SelectionBar(
|
||||
trade_date=trade_date,
|
||||
open=_as_float(values[0]),
|
||||
@@ -248,9 +319,70 @@ class PostgresMarketDataReader:
|
||||
close=_as_float(values[3]),
|
||||
volume=_as_float(values[4]),
|
||||
)
|
||||
basic = SelectionDailyBasic(
|
||||
trade_date=trade_date,
|
||||
turnover_rate=_as_float(values[5]),
|
||||
total_mv=_as_float(values[6]),
|
||||
basic = (
|
||||
SelectionDailyBasic(
|
||||
trade_date=trade_date,
|
||||
turnover_rate=_as_float(values[5]),
|
||||
total_mv=_as_float(values[6]),
|
||||
)
|
||||
if len(values) >= 7
|
||||
else None
|
||||
)
|
||||
return row_code, name, bar, basic
|
||||
|
||||
@classmethod
|
||||
def _histories_from_rows(
|
||||
cls,
|
||||
rows: Sequence[object],
|
||||
stocks: Sequence[SelectionStock],
|
||||
target_trade_date: date,
|
||||
) -> tuple[StockHistory, ...]:
|
||||
"""Group sorted/possibly duplicated database rows by requested stock."""
|
||||
|
||||
requested = {stock.ts_code: stock for stock in stocks}
|
||||
bars_by_code: dict[str, dict[date, SelectionBar]] = {code: {} for code in requested}
|
||||
basics_by_code: dict[str, dict[date, SelectionDailyBasic]] = {
|
||||
code: {} for code in requested
|
||||
}
|
||||
names = {stock.ts_code: stock.name for stock in stocks}
|
||||
for raw_row in rows:
|
||||
row = cast(tuple[object, ...], raw_row)
|
||||
row_hint = str(row[0]) if row and row[0] is not None else ""
|
||||
row_code, row_name, bar, basic = cls._map_row(row, row_hint)
|
||||
if row_code not in requested:
|
||||
raise ValueError(f"reader returned unexpected stock code: {row_code}")
|
||||
names[row_code] = row_name or names[row_code]
|
||||
if bar.trade_date <= target_trade_date:
|
||||
bars_by_code[row_code][bar.trade_date] = bar
|
||||
if basic is not None:
|
||||
basics_by_code[row_code][bar.trade_date] = basic
|
||||
histories: list[StockHistory] = []
|
||||
for stock in stocks:
|
||||
code = stock.ts_code
|
||||
bars = bars_by_code[code]
|
||||
basics = basics_by_code[code]
|
||||
histories.append(
|
||||
StockHistory(
|
||||
ts_code=code,
|
||||
name=names[code],
|
||||
bars=tuple(bars[trade_date] for trade_date in sorted(bars)),
|
||||
daily_basic={trade_date: basics[trade_date] for trade_date in sorted(basics)},
|
||||
)
|
||||
)
|
||||
return tuple(histories)
|
||||
|
||||
@contextmanager
|
||||
def _connection(self) -> Generator[Any, None, None]:
|
||||
"""Borrow from the shared pool or use the legacy direct connection."""
|
||||
|
||||
try:
|
||||
if self.pool is None:
|
||||
with psycopg.connect(self.database_url) as connection:
|
||||
yield connection
|
||||
else:
|
||||
with self.pool.connection() as connection:
|
||||
yield connection
|
||||
except (SelectionMarketDataNotReady, SelectionReaderError, ValueError):
|
||||
raise
|
||||
except Exception as exc: # noqa: BLE001 - normalize pool/driver failures
|
||||
raise SelectionReaderError("selection database operation failed") from exc
|
||||
|
||||
+119
-56
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from collections.abc import Generator, Mapping
|
||||
from collections.abc import Generator, Mapping, Sequence
|
||||
from contextlib import contextmanager
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
@@ -28,6 +28,7 @@ from ..domain.runs import (
|
||||
SelectionRunStoreError,
|
||||
)
|
||||
from ..domain.zhixing_b1 import ZHIXING_B1_SIGNAL_ORDER
|
||||
from .postgres_pool import SelectionConnectionPool, SelectionPostgresPool
|
||||
|
||||
_SIGNAL_PRIORITY = {category: index for index, category in enumerate(ZHIXING_B1_SIGNAL_ORDER)}
|
||||
_CATEGORY_PREFIXES = {
|
||||
@@ -45,13 +46,52 @@ _SIGNAL_ORDER_SQL = (
|
||||
)
|
||||
|
||||
|
||||
_ITEM_UPSERT = """
|
||||
INSERT INTO selection_run_item
|
||||
(run_id, ts_code, name, status, signal_count, reason)
|
||||
VALUES (%s, %s, %s, %s, %s, %s)
|
||||
ON CONFLICT (run_id, ts_code) DO UPDATE SET
|
||||
name = EXCLUDED.name,
|
||||
status = EXCLUDED.status,
|
||||
signal_count = EXCLUDED.signal_count,
|
||||
reason = EXCLUDED.reason
|
||||
"""
|
||||
_SIGNAL_UPSERT = """
|
||||
INSERT INTO selection_signal
|
||||
(
|
||||
run_id, ts_code, name, target_trade_date, strategy,
|
||||
category, close, details
|
||||
)
|
||||
VALUES (%s, %s, %s, %s, %s, %s, %s, %s)
|
||||
ON CONFLICT (run_id, ts_code, category) DO UPDATE SET
|
||||
name = EXCLUDED.name,
|
||||
close = EXCLUDED.close,
|
||||
details = EXCLUDED.details
|
||||
"""
|
||||
|
||||
|
||||
class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
"""Persist one current result attempt per strategy and target date."""
|
||||
|
||||
def __init__(self, database_url: str) -> None:
|
||||
"""Create the adapter with an injected PostgreSQL URL."""
|
||||
def __init__(
|
||||
self,
|
||||
database_url: str,
|
||||
*,
|
||||
pool: SelectionPostgresPool | SelectionConnectionPool | None = None,
|
||||
) -> None:
|
||||
"""Create the adapter with a URL and optional shared PostgreSQL pool."""
|
||||
|
||||
self.database_url = database_url
|
||||
if isinstance(pool, SelectionPostgresPool):
|
||||
self.pool: SelectionPostgresPool | None = pool
|
||||
elif pool is not None:
|
||||
self.pool = SelectionPostgresPool(
|
||||
database_url,
|
||||
max_connections=1,
|
||||
pool=pool,
|
||||
)
|
||||
else:
|
||||
self.pool = None
|
||||
|
||||
def prepare_run(
|
||||
self,
|
||||
@@ -137,62 +177,62 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
)
|
||||
|
||||
def record_item(self, run_id: str, item: SelectionRunItem) -> None:
|
||||
"""Upsert one stock outcome and all of its independent signal rows."""
|
||||
"""Persist one item through the batch path for compatibility."""
|
||||
|
||||
self.record_items(run_id, (item,))
|
||||
|
||||
def record_items(self, run_id: str, items: Sequence[SelectionRunItem]) -> None:
|
||||
"""Persist one chunk in one transaction with set-based driver calls.
|
||||
|
||||
Existing signal rows are removed before the upserts so retrying a
|
||||
chunk cannot retain a category that disappeared from a recalculation.
|
||||
``executemany`` is used for both materialized tables; the small
|
||||
fallback keeps the direct fake connections used by older tests usable.
|
||||
"""
|
||||
|
||||
if not items:
|
||||
return
|
||||
item_values = tuple(
|
||||
(
|
||||
run_id,
|
||||
item.ts_code,
|
||||
item.name,
|
||||
item.status,
|
||||
item.signal_count,
|
||||
item.reason,
|
||||
)
|
||||
for item in items
|
||||
)
|
||||
signal_values = tuple(
|
||||
(
|
||||
run_id,
|
||||
signal.ts_code,
|
||||
signal.name,
|
||||
signal.target_trade_date,
|
||||
signal.strategy,
|
||||
signal.category.value,
|
||||
signal.close,
|
||||
Jsonb(dict(signal.details)),
|
||||
)
|
||||
for item in items
|
||||
for signal in item.signals
|
||||
)
|
||||
codes = [item.ts_code for item in items]
|
||||
try:
|
||||
with self._connection() as connection, connection.transaction():
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO selection_run_item
|
||||
(run_id, ts_code, name, status, signal_count, reason)
|
||||
VALUES (%s, %s, %s, %s, %s, %s)
|
||||
ON CONFLICT (run_id, ts_code) DO UPDATE SET
|
||||
name = EXCLUDED.name,
|
||||
status = EXCLUDED.status,
|
||||
signal_count = EXCLUDED.signal_count,
|
||||
reason = EXCLUDED.reason
|
||||
""",
|
||||
(
|
||||
run_id,
|
||||
item.ts_code,
|
||||
item.name,
|
||||
item.status,
|
||||
item.signal_count,
|
||||
item.reason,
|
||||
),
|
||||
"DELETE FROM selection_signal WHERE run_id = %s AND ts_code = ANY(%s)",
|
||||
(run_id, codes),
|
||||
)
|
||||
connection.execute(
|
||||
"DELETE FROM selection_signal WHERE run_id = %s AND ts_code = %s",
|
||||
(run_id, item.ts_code),
|
||||
)
|
||||
for signal in item.signals:
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO selection_signal
|
||||
(
|
||||
run_id, ts_code, name, target_trade_date, strategy,
|
||||
category, close, details
|
||||
)
|
||||
VALUES (%s, %s, %s, %s, %s, %s, %s, %s)
|
||||
ON CONFLICT (run_id, ts_code, category) DO UPDATE SET
|
||||
name = EXCLUDED.name,
|
||||
close = EXCLUDED.close,
|
||||
details = EXCLUDED.details
|
||||
""",
|
||||
(
|
||||
run_id,
|
||||
signal.ts_code,
|
||||
signal.name,
|
||||
signal.target_trade_date,
|
||||
signal.strategy,
|
||||
signal.category.value,
|
||||
signal.close,
|
||||
Jsonb(dict(signal.details)),
|
||||
),
|
||||
)
|
||||
except psycopg.Error as exc:
|
||||
_executemany(connection, _ITEM_UPSERT, item_values)
|
||||
if signal_values:
|
||||
_executemany(connection, _SIGNAL_UPSERT, signal_values)
|
||||
except SelectionRunError:
|
||||
raise
|
||||
except Exception as exc: # noqa: BLE001 - redact driver/pool details
|
||||
code_context = items[0].ts_code if len(items) == 1 else f"{len(items)} items"
|
||||
raise SelectionRunStoreError(
|
||||
f"failed to persist selection item {item.ts_code}"
|
||||
f"failed to persist selection item {code_context}"
|
||||
) from exc
|
||||
|
||||
def finish_run(
|
||||
@@ -398,12 +438,35 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
"""Translate psycopg failures without exposing driver details."""
|
||||
|
||||
try:
|
||||
with psycopg.connect(self.database_url) as connection:
|
||||
yield connection
|
||||
except psycopg.Error as exc:
|
||||
if self.pool is None:
|
||||
with psycopg.connect(self.database_url) as connection:
|
||||
yield connection
|
||||
else:
|
||||
with self.pool.connection() as connection:
|
||||
yield connection
|
||||
except SelectionRunError:
|
||||
raise
|
||||
except Exception as exc: # noqa: BLE001 - normalize driver/pool errors
|
||||
raise SelectionRunStoreError("selection database operation failed") from exc
|
||||
|
||||
|
||||
def _executemany(connection: Any, query: str, parameters: Sequence[tuple[object, ...]]) -> None:
|
||||
"""Use psycopg's batch API while retaining a minimal fake connection seam."""
|
||||
|
||||
executemany = getattr(connection, "executemany", None)
|
||||
if callable(executemany):
|
||||
executemany(query, parameters)
|
||||
return
|
||||
cursor_factory = getattr(connection, "cursor", None)
|
||||
if callable(cursor_factory):
|
||||
cursor_context = cast(Any, cursor_factory())
|
||||
with cursor_context as cursor:
|
||||
cursor.executemany(query, parameters)
|
||||
return
|
||||
for values in parameters:
|
||||
connection.execute(query, values)
|
||||
|
||||
|
||||
def _signal_from_row(row: tuple[object, ...]) -> SelectionSignal:
|
||||
"""Map a persisted signal row back to the domain signal model."""
|
||||
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""HTTP presentation for persisted strategy execution results."""
|
||||
|
||||
import atexit
|
||||
import threading
|
||||
from datetime import date, datetime
|
||||
from typing import Annotated, Literal
|
||||
|
||||
@@ -17,6 +19,7 @@ from zhixing_server.modules.selection.domain.runs import (
|
||||
SelectionRunInProgress,
|
||||
SelectionRunStoreError,
|
||||
)
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||
PostgresMarketDataReader,
|
||||
SelectionMarketDataNotReady,
|
||||
@@ -27,6 +30,8 @@ from zhixing_server.modules.selection.infrastructure.postgres_runs import (
|
||||
)
|
||||
|
||||
selection_router = APIRouter()
|
||||
_SELECTION_POOL_CACHE_LOCK = threading.Lock()
|
||||
_SELECTION_POOL_CACHE: dict[tuple[str, int], SelectionPostgresPool] = {}
|
||||
|
||||
StrategyValue = Literal["zhixing_b1"]
|
||||
SelectionStatusValue = Literal[
|
||||
@@ -117,11 +122,45 @@ class SelectionResultsResponse(BaseModel):
|
||||
def get_selection_service(
|
||||
settings: Annotated[Settings, Depends(get_settings)],
|
||||
) -> RunZhixingB1:
|
||||
"""Build one request-scoped selection application service."""
|
||||
"""Build the selection service on top of process-scoped shared resources."""
|
||||
|
||||
reader = PostgresMarketDataReader(settings)
|
||||
store = PostgresSelectionRunRepository(settings.database_url)
|
||||
return RunZhixingB1(reader, store)
|
||||
pool = get_selection_postgres_pool(settings)
|
||||
reader = PostgresMarketDataReader(settings, pool=pool)
|
||||
store = PostgresSelectionRunRepository(settings.database_url, pool=pool)
|
||||
return RunZhixingB1(
|
||||
reader,
|
||||
store,
|
||||
max_workers=settings.selection_max_workers,
|
||||
batch_size=settings.selection_batch_size,
|
||||
)
|
||||
|
||||
|
||||
def get_selection_postgres_pool(settings: Settings) -> SelectionPostgresPool:
|
||||
"""Return the cached bounded pool shared by selection adapters."""
|
||||
|
||||
key = (settings.database_url, settings.selection_max_workers + 2)
|
||||
with _SELECTION_POOL_CACHE_LOCK:
|
||||
pool = _SELECTION_POOL_CACHE.get(key)
|
||||
if pool is None:
|
||||
pool = SelectionPostgresPool(
|
||||
settings.database_url,
|
||||
max_connections=key[1],
|
||||
)
|
||||
_SELECTION_POOL_CACHE[key] = pool
|
||||
return pool
|
||||
|
||||
|
||||
def _close_cached_selection_pools() -> None:
|
||||
"""Close all process-cached selection pools during interpreter shutdown."""
|
||||
|
||||
with _SELECTION_POOL_CACHE_LOCK:
|
||||
pools = tuple(_SELECTION_POOL_CACHE.values())
|
||||
_SELECTION_POOL_CACHE.clear()
|
||||
for pool in pools:
|
||||
pool.close()
|
||||
|
||||
|
||||
atexit.register(_close_cached_selection_pools)
|
||||
|
||||
|
||||
@selection_router.post(
|
||||
@@ -292,5 +331,6 @@ __all__ = [
|
||||
"SelectionRunAcceptedResponse",
|
||||
"SelectionRunRequest",
|
||||
"get_selection_service",
|
||||
"get_selection_postgres_pool",
|
||||
"selection_router",
|
||||
]
|
||||
|
||||
@@ -3,9 +3,12 @@
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
import zhixing_server.modules.selection.presentation.http as selection_http
|
||||
from zhixing_server.bootstrap.app import create_app
|
||||
from zhixing_server.bootstrap.config import Settings
|
||||
from zhixing_server.modules.selection.application.run import PreparedSelectionRun
|
||||
from zhixing_server.modules.selection.domain.models import SelectionSignal, ZhixingB1Category
|
||||
from zhixing_server.modules.selection.domain.runs import (
|
||||
@@ -19,7 +22,10 @@ from zhixing_server.modules.selection.domain.runs import (
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||
SelectionMarketDataNotReady,
|
||||
)
|
||||
from zhixing_server.modules.selection.presentation.http import get_selection_service
|
||||
from zhixing_server.modules.selection.presentation.http import (
|
||||
get_selection_postgres_pool,
|
||||
get_selection_service,
|
||||
)
|
||||
|
||||
TARGET = date(2026, 8, 8)
|
||||
|
||||
@@ -128,6 +134,17 @@ def _client(service: FakeSelectionService) -> TestClient:
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_selection_dependency_reuses_one_bounded_pool(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(selection_http, "_SELECTION_POOL_CACHE", {})
|
||||
settings = Settings(database_url="postgresql://test", selection_max_workers=4)
|
||||
|
||||
first = get_selection_postgres_pool(settings)
|
||||
second = get_selection_postgres_pool(settings)
|
||||
|
||||
assert first is second
|
||||
assert first.max_connections == 6
|
||||
|
||||
|
||||
def test_trigger_returns_accepted_run_and_schedules_execution() -> None:
|
||||
service = FakeSelectionService()
|
||||
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
"""Lifecycle tests for the shared selection PostgreSQL pool owner."""
|
||||
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool
|
||||
|
||||
|
||||
class FakePool:
|
||||
def __init__(self) -> None:
|
||||
self.open_calls = 0
|
||||
self.close_calls = 0
|
||||
self.connection_calls = 0
|
||||
|
||||
def open(self, *, wait: bool = True) -> None:
|
||||
assert wait is True
|
||||
self.open_calls += 1
|
||||
|
||||
def close(self) -> None:
|
||||
self.close_calls += 1
|
||||
|
||||
@contextmanager
|
||||
def connection(self) -> Generator[str, None, None]:
|
||||
self.connection_calls += 1
|
||||
yield "connection"
|
||||
|
||||
|
||||
def test_selection_pool_opens_once_borrows_and_closes_injected_pool() -> None:
|
||||
fake = FakePool()
|
||||
owner = SelectionPostgresPool(
|
||||
"postgresql://test",
|
||||
max_connections=6,
|
||||
pool=fake,
|
||||
)
|
||||
|
||||
with owner.connection() as connection:
|
||||
assert connection == "connection"
|
||||
with owner.connection() as connection:
|
||||
assert connection == "connection"
|
||||
|
||||
assert owner.max_connections == 6
|
||||
assert fake.open_calls == 1
|
||||
assert fake.connection_calls == 2
|
||||
owner.close()
|
||||
assert fake.close_calls == 1
|
||||
@@ -1,5 +1,7 @@
|
||||
"""PostgreSQL reader contract tests using a fake connection."""
|
||||
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from typing import cast
|
||||
@@ -7,6 +9,8 @@ from typing import cast
|
||||
import psycopg
|
||||
import pytest
|
||||
|
||||
from zhixing_server.modules.selection.domain.runs import SelectionStock
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||
PostgresMarketDataReader,
|
||||
SelectionMarketDataNotReady,
|
||||
@@ -39,6 +43,23 @@ class FakeResult:
|
||||
return self.rows
|
||||
|
||||
|
||||
class Pool:
|
||||
def __init__(self, connection: FakeConnection) -> None:
|
||||
self.connection_value = connection
|
||||
self.opened = 0
|
||||
|
||||
def open(self, *, wait: bool = True) -> None:
|
||||
assert wait is True
|
||||
self.opened += 1
|
||||
|
||||
def close(self) -> None:
|
||||
return None
|
||||
|
||||
@contextmanager
|
||||
def connection(self) -> Generator[FakeConnection, None, None]:
|
||||
yield self.connection_value
|
||||
|
||||
|
||||
def test_reader_parameterizes_target_and_maps_left_join(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
connection = FakeConnection(
|
||||
[
|
||||
@@ -86,6 +107,64 @@ def test_reader_parameterizes_target_and_maps_left_join(monkeypatch: pytest.Monk
|
||||
assert "trade_date <= %s" in cast(str, connection.query)
|
||||
|
||||
|
||||
def test_reader_batches_qfq_rows_by_stock_without_historical_basic_join() -> None:
|
||||
connection = FakeConnection(
|
||||
[
|
||||
(
|
||||
"600000.SH",
|
||||
"浦发银行",
|
||||
date(2024, 1, 2),
|
||||
"8",
|
||||
"8.5",
|
||||
"7.8",
|
||||
"8.2",
|
||||
"900",
|
||||
),
|
||||
(
|
||||
"000001.SZ",
|
||||
"平安银行",
|
||||
date(2024, 1, 3),
|
||||
"10.5",
|
||||
"11",
|
||||
"10",
|
||||
"10.8",
|
||||
"1200",
|
||||
),
|
||||
(
|
||||
"000001.SZ",
|
||||
"平安银行",
|
||||
date(2024, 1, 2),
|
||||
"10",
|
||||
"11",
|
||||
"9",
|
||||
"10.5",
|
||||
"1000",
|
||||
),
|
||||
]
|
||||
)
|
||||
pool = Pool(connection)
|
||||
owner = SelectionPostgresPool("postgresql://test", max_connections=6, pool=pool)
|
||||
|
||||
histories = PostgresMarketDataReader("postgresql://test", pool=owner).load_histories(
|
||||
(
|
||||
SelectionStock("000001.SZ", "平安银行"),
|
||||
SelectionStock("600000.SH", "浦发银行"),
|
||||
),
|
||||
date(2024, 1, 3),
|
||||
)
|
||||
|
||||
assert [history.ts_code for history in histories] == ["000001.SZ", "600000.SH"]
|
||||
assert [bar.trade_date for bar in histories[0].bars] == [
|
||||
date(2024, 1, 2),
|
||||
date(2024, 1, 3),
|
||||
]
|
||||
assert histories[0].daily_basic == {}
|
||||
assert connection.parameters == (["000001.SZ", "600000.SH"], date(2024, 1, 3))
|
||||
assert "bar.ts_code = ANY(%s)" in cast(str, connection.query)
|
||||
assert "market_daily_basic" not in cast(str, connection.query)
|
||||
assert pool.opened == 1
|
||||
|
||||
|
||||
class SourceConnection:
|
||||
def __init__(self, source_row: tuple[object, ...] | None) -> None:
|
||||
self.source_row = source_row
|
||||
|
||||
@@ -61,6 +61,33 @@ class FakeConnection:
|
||||
return FakeResult()
|
||||
|
||||
|
||||
class BatchConnection(FakeConnection):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(None)
|
||||
self.executemany_calls: list[tuple[str, tuple[tuple[object, ...], ...]]] = []
|
||||
|
||||
def cursor(self) -> "BatchCursor":
|
||||
return BatchCursor(self.executemany_calls)
|
||||
|
||||
|
||||
class BatchCursor:
|
||||
def __init__(self, calls: list[tuple[str, tuple[tuple[object, ...], ...]]]) -> None:
|
||||
self.calls = calls
|
||||
|
||||
def __enter__(self) -> "BatchCursor":
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
def executemany(
|
||||
self,
|
||||
query: str,
|
||||
parameters: tuple[tuple[object, ...], ...],
|
||||
) -> None:
|
||||
self.calls.append((query, parameters))
|
||||
|
||||
|
||||
def _source() -> SelectionExecutionSource:
|
||||
return SelectionExecutionSource(
|
||||
market_sync_batch_id="market-run-1",
|
||||
@@ -163,6 +190,64 @@ def test_record_item_persists_independent_signals_as_jsonb(monkeypatch: pytest.M
|
||||
assert isinstance(signal_insert[-1], Jsonb)
|
||||
|
||||
|
||||
def test_record_items_uses_one_delete_and_two_batch_upserts(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
connection = BatchConnection()
|
||||
repository = _repository(monkeypatch, connection)
|
||||
first = SelectionSignal(
|
||||
ts_code="000001.SZ",
|
||||
name="平安银行",
|
||||
target_trade_date=TARGET,
|
||||
strategy="zhixing_b1",
|
||||
category=ZHIXING_B1_SIGNAL_ORDER[0],
|
||||
close=10.5,
|
||||
details={"j": 12.0},
|
||||
)
|
||||
second = SelectionSignal(
|
||||
ts_code="000001.SZ",
|
||||
name="平安银行",
|
||||
target_trade_date=TARGET,
|
||||
strategy="zhixing_b1",
|
||||
category=ZHIXING_B1_SIGNAL_ORDER[-1],
|
||||
close=10.5,
|
||||
details={"j": 13.0},
|
||||
)
|
||||
|
||||
repository.record_items(
|
||||
"run-1",
|
||||
(
|
||||
SelectionRunItem(
|
||||
ts_code="000001.SZ",
|
||||
name="平安银行",
|
||||
status="selected",
|
||||
signal_count=2,
|
||||
signals=(first, second),
|
||||
),
|
||||
SelectionRunItem(
|
||||
ts_code="600000.SH",
|
||||
name="浦发银行",
|
||||
status="no_signal",
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
delete_query, delete_parameters = connection.statements[0]
|
||||
assert "DELETE FROM selection_signal" in delete_query
|
||||
assert "ANY(%s)" in delete_query
|
||||
assert delete_parameters == ("run-1", ["000001.SZ", "600000.SH"])
|
||||
assert len(connection.executemany_calls) == 2
|
||||
assert "INSERT INTO selection_run_item" in connection.executemany_calls[0][0]
|
||||
assert "INSERT INTO selection_signal" in connection.executemany_calls[1][0]
|
||||
signal_parameters = connection.executemany_calls[1][1]
|
||||
assert len(signal_parameters) == 2
|
||||
assert {values[5] for values in signal_parameters} == {
|
||||
ZHIXING_B1_SIGNAL_ORDER[0].value,
|
||||
ZHIXING_B1_SIGNAL_ORDER[-1].value,
|
||||
}
|
||||
assert all(isinstance(values[-1], Jsonb) for values in signal_parameters)
|
||||
|
||||
|
||||
class LoadConnection:
|
||||
def __init__(self) -> None:
|
||||
self.statements: list[tuple[str, tuple[object, ...]]] = []
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Application tests for persisted whole-universe selection runs."""
|
||||
|
||||
import threading
|
||||
import time
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from typing import Literal
|
||||
@@ -42,6 +44,30 @@ class FakeReader:
|
||||
raise AssertionError("the fake evaluator should be used")
|
||||
|
||||
|
||||
class BatchReader(FakeReader):
|
||||
def __init__(self, source: SelectionExecutionSource) -> None:
|
||||
super().__init__(source)
|
||||
self.batch_calls: list[tuple[str, ...]] = []
|
||||
|
||||
def load_histories(
|
||||
self,
|
||||
stocks: tuple[SelectionStock, ...],
|
||||
target_trade_date: date,
|
||||
) -> tuple[StockHistory, ...]:
|
||||
self.batch_calls.append(tuple(stock.ts_code for stock in stocks))
|
||||
return tuple(StockHistory(ts_code=stock.ts_code, name=stock.name) for stock in stocks)
|
||||
|
||||
|
||||
class PartialBatchReader(BatchReader):
|
||||
def load_histories(
|
||||
self,
|
||||
stocks: tuple[SelectionStock, ...],
|
||||
target_trade_date: date,
|
||||
) -> tuple[StockHistory, ...]:
|
||||
self.batch_calls.append(tuple(stock.ts_code for stock in stocks))
|
||||
return ()
|
||||
|
||||
|
||||
class FakeStore:
|
||||
def __init__(self) -> None:
|
||||
self.items: list[SelectionRunItem] = []
|
||||
@@ -114,6 +140,26 @@ class FakeStore:
|
||||
return None
|
||||
|
||||
|
||||
class BatchStore(FakeStore):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.batches: list[tuple[SelectionRunItem, ...]] = []
|
||||
|
||||
def record_item(self, run_id: str, item: SelectionRunItem) -> None:
|
||||
raise AssertionError("the batch path should use record_items")
|
||||
|
||||
def record_items(self, run_id: str, items: tuple[SelectionRunItem, ...]) -> None:
|
||||
assert run_id == "run-1"
|
||||
batch = tuple(items)
|
||||
self.batches.append(batch)
|
||||
self.items.extend(batch)
|
||||
|
||||
|
||||
class FailingBatchStore(BatchStore):
|
||||
def record_items(self, run_id: str, items: tuple[SelectionRunItem, ...]) -> None:
|
||||
raise RuntimeError("batch write unavailable")
|
||||
|
||||
|
||||
class FakeEvaluator:
|
||||
def __init__(self, results: dict[str, SelectionEvaluation]) -> None:
|
||||
self.results = results
|
||||
@@ -129,6 +175,29 @@ class RaisingEvaluator:
|
||||
return SelectionEvaluation(ts_code, target_trade_date, "no_signal")
|
||||
|
||||
|
||||
class ConcurrentHistoryEvaluator:
|
||||
def __init__(self) -> None:
|
||||
self.active = 0
|
||||
self.peak = 0
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def execute(self, ts_code: str, target_trade_date: date) -> SelectionEvaluation:
|
||||
raise AssertionError("the batch evaluator path should be used")
|
||||
|
||||
def execute_history(
|
||||
self,
|
||||
history: StockHistory,
|
||||
target_trade_date: date,
|
||||
) -> SelectionEvaluation:
|
||||
with self._lock:
|
||||
self.active += 1
|
||||
self.peak = max(self.peak, self.active)
|
||||
time.sleep(0.02)
|
||||
with self._lock:
|
||||
self.active -= 1
|
||||
return SelectionEvaluation(history.ts_code, target_trade_date, "no_signal")
|
||||
|
||||
|
||||
def _source() -> SelectionExecutionSource:
|
||||
return SelectionExecutionSource(
|
||||
market_sync_batch_id="market-run-1",
|
||||
@@ -233,3 +302,65 @@ def test_execute_isolates_unexpected_single_stock_failure() -> None:
|
||||
assert store.finished[0:2] == ("run-1", "partial_success")
|
||||
assert store.finished[2]["evaluated_count"] == 2
|
||||
assert store.finished[2]["failed_count"] == 1
|
||||
|
||||
|
||||
def test_execute_batches_history_reads_writes_and_limits_evaluation_workers() -> None:
|
||||
source = SelectionExecutionSource(
|
||||
market_sync_batch_id="market-run-1",
|
||||
target_trade_date=TARGET,
|
||||
target_count=8,
|
||||
valid_count=8,
|
||||
coverage=Decimal("1"),
|
||||
stocks=tuple(SelectionStock(f"{index:06d}.SZ", f"stock-{index}") for index in range(8)),
|
||||
)
|
||||
reader = BatchReader(source)
|
||||
store = BatchStore()
|
||||
evaluator = ConcurrentHistoryEvaluator()
|
||||
|
||||
service = RunZhixingB1(
|
||||
reader,
|
||||
store,
|
||||
evaluator,
|
||||
max_workers=4,
|
||||
batch_size=4,
|
||||
)
|
||||
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
|
||||
|
||||
assert reader.batch_calls == [
|
||||
("000000.SZ", "000001.SZ", "000002.SZ", "000003.SZ"),
|
||||
("000004.SZ", "000005.SZ", "000006.SZ", "000007.SZ"),
|
||||
]
|
||||
assert [len(batch) for batch in store.batches] == [4, 4]
|
||||
assert evaluator.peak <= 4
|
||||
assert evaluator.peak >= 2
|
||||
assert store.finished is not None
|
||||
assert store.finished[0:2] == ("run-1", "success")
|
||||
|
||||
|
||||
def test_execute_maps_missing_batch_history_without_a_single_stock_read() -> None:
|
||||
source = _source()
|
||||
reader = PartialBatchReader(source)
|
||||
store = FakeStore()
|
||||
|
||||
service = RunZhixingB1(reader, store)
|
||||
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
|
||||
|
||||
assert reader.batch_calls == [("000001.SZ", "600000.SH")]
|
||||
assert [item.status for item in store.items] == ["missing_target_bar", "missing_target_bar"]
|
||||
assert store.finished is not None
|
||||
assert store.finished[0:2] == ("run-1", "failed")
|
||||
|
||||
|
||||
def test_execute_marks_batch_write_failure_as_failed() -> None:
|
||||
source = _source()
|
||||
reader = BatchReader(source)
|
||||
store = FailingBatchStore()
|
||||
evaluator = ConcurrentHistoryEvaluator()
|
||||
|
||||
service = RunZhixingB1(reader, store, evaluator)
|
||||
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
|
||||
|
||||
assert store.finished is not None
|
||||
assert store.finished[0:2] == ("run-1", "failed")
|
||||
assert store.finished[2]["error_type"] == "batch_error"
|
||||
assert store.finished[2]["failed_count"] == 1
|
||||
|
||||
Reference in New Issue
Block a user