Merge pull request 'Develop' (#12) from develop into main
Deploy Production / deploy (push) Successful in 20s
Deploy Production / deploy (push) Successful in 20s
Reviewed-on: sakibcc/zhixing-system#12
This commit was merged in pull request #12.
This commit is contained in:
@@ -23,4 +23,6 @@ ZHIXING_MARKET_DATA_MAX_RETRIES=3
|
|||||||
ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS=1.0
|
ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS=1.0
|
||||||
ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS=0.2
|
ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS=0.2
|
||||||
ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY=7380521
|
ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY=7380521
|
||||||
|
ZHIXING_SELECTION_MAX_WORKERS=4
|
||||||
|
ZHIXING_SELECTION_BATCH_SIZE=200
|
||||||
API_UPSTREAM=http://server:8000
|
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 -->
|
<!-- @@@auto:current-status -->
|
||||||
- **Active File**: `journal-1.md`
|
- **Active File**: `journal-1.md`
|
||||||
- **Total Sessions**: 8
|
- **Total Sessions**: 9
|
||||||
- **Last Active**: 2026-08-11
|
- **Last Active**: 2026-08-12
|
||||||
<!-- @@@/auto:current-status -->
|
<!-- @@@/auto:current-status -->
|
||||||
|
|
||||||
---
|
---
|
||||||
@@ -19,7 +19,7 @@
|
|||||||
<!-- @@@auto:active-documents -->
|
<!-- @@@auto:active-documents -->
|
||||||
| File | Lines | Status |
|
| File | Lines | Status |
|
||||||
|------|-------|--------|
|
|------|-------|--------|
|
||||||
| `journal-1.md` | ~222 | Active |
|
| `journal-1.md` | ~243 | Active |
|
||||||
<!-- @@@/auto:active-documents -->
|
<!-- @@@/auto:active-documents -->
|
||||||
|
|
||||||
---
|
---
|
||||||
@@ -29,6 +29,7 @@
|
|||||||
<!-- @@@auto:session-history -->
|
<!-- @@@auto:session-history -->
|
||||||
| # | Date | Title | Commits | Branch |
|
| # | Date | Title | Commits | Branch |
|
||||||
|---|------|-------|---------|--------|
|
|---|------|-------|---------|--------|
|
||||||
|
| 9 | 2026-08-12 | 完成选股执行性能优化 | `8963c06` | `develop` |
|
||||||
| 8 | 2026-08-11 | 完成市场数据同步与完整性检查 | `7ce1154`, `8f5f504` | `develop` |
|
| 8 | 2026-08-11 | 完成市场数据同步与完整性检查 | `7ce1154`, `8f5f504` | `develop` |
|
||||||
| 7 | 2026-08-10 | 完成选股执行状态抽屉与紧凑布局 | `17237e0` | `develop` |
|
| 7 | 2026-08-10 | 完成选股执行状态抽屉与紧凑布局 | `17237e0` | `develop` |
|
||||||
| 6 | 2026-08-10 | 按原型完善选股结果分页接口 | `ed7bdda`, `3af97bf` | `develop` |
|
| 6 | 2026-08-10 | 按原型完善选股结果分页接口 | `ed7bdda`, `3af97bf` | `develop` |
|
||||||
|
|||||||
@@ -220,3 +220,24 @@
|
|||||||
### Next Steps
|
### Next Steps
|
||||||
|
|
||||||
- 配置 ZHIXING_TEST_DATABASE_URL 后执行真实 PostgreSQL 集成验证
|
- 配置 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_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_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_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:-}
|
ZHIXING_TUSHARE_TOKEN: ${ZHIXING_TUSHARE_TOKEN:-}
|
||||||
init: true
|
init: true
|
||||||
ports:
|
ports:
|
||||||
|
|||||||
@@ -18,6 +18,8 @@ services:
|
|||||||
ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS: ${ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS:-1.0}
|
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_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_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:-}
|
ZHIXING_TUSHARE_TOKEN: ${ZHIXING_TUSHARE_TOKEN:-}
|
||||||
init: true
|
init: true
|
||||||
expose:
|
expose:
|
||||||
|
|||||||
@@ -24,6 +24,8 @@ class Settings(BaseSettings):
|
|||||||
market_data_max_retries: int = 3
|
market_data_max_retries: int = 3
|
||||||
market_data_retry_backoff_seconds: float = 1.0
|
market_data_retry_backoff_seconds: float = 1.0
|
||||||
market_data_advisory_lock_key: int = 7_380_521
|
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(
|
model_config = SettingsConfigDict(
|
||||||
env_file=".env",
|
env_file=".env",
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
from datetime import date
|
from datetime import date
|
||||||
|
|
||||||
from ..domain.models import SelectionEvaluation, StockHistory
|
from ..domain.models import SelectionEvaluation, StockHistory
|
||||||
@@ -44,3 +45,12 @@ class EvaluateZhixingB1:
|
|||||||
"""Evaluate an already loaded history for deterministic unit tests."""
|
"""Evaluate an already loaded history for deterministic unit tests."""
|
||||||
|
|
||||||
return self.strategy.evaluate(history, target_trade_date)
|
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
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import time
|
||||||
|
from collections.abc import Callable, Sequence
|
||||||
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import date
|
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 (
|
from ..domain.runs import (
|
||||||
|
BatchSelectionRunStore,
|
||||||
|
BatchSelectionUniverseReader,
|
||||||
SelectionExecutionSource,
|
SelectionExecutionSource,
|
||||||
SelectionRerunRequired,
|
SelectionRerunRequired,
|
||||||
SelectionResultQuery,
|
SelectionResultQuery,
|
||||||
@@ -17,6 +22,7 @@ from ..domain.runs import (
|
|||||||
SelectionRunItem,
|
SelectionRunItem,
|
||||||
SelectionRunStatus,
|
SelectionRunStatus,
|
||||||
SelectionRunStore,
|
SelectionRunStore,
|
||||||
|
SelectionStock,
|
||||||
SelectionUniverseReader,
|
SelectionUniverseReader,
|
||||||
)
|
)
|
||||||
from .evaluate import EvaluateZhixingB1
|
from .evaluate import EvaluateZhixingB1
|
||||||
@@ -48,12 +54,21 @@ class RunZhixingB1:
|
|||||||
reader: SelectionUniverseReader,
|
reader: SelectionUniverseReader,
|
||||||
store: SelectionRunStore,
|
store: SelectionRunStore,
|
||||||
evaluator: SelectionEvaluator | None = None,
|
evaluator: SelectionEvaluator | None = None,
|
||||||
|
*,
|
||||||
|
max_workers: int = 4,
|
||||||
|
batch_size: int = 200,
|
||||||
) -> None:
|
) -> 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.reader = reader
|
||||||
self.store = store
|
self.store = store
|
||||||
self.evaluator = evaluator or EvaluateZhixingB1(reader)
|
self.evaluator = evaluator or EvaluateZhixingB1(reader)
|
||||||
|
self.max_workers = max_workers
|
||||||
|
self.batch_size = batch_size
|
||||||
|
|
||||||
def prepare(
|
def prepare(
|
||||||
self,
|
self,
|
||||||
@@ -81,35 +96,57 @@ class RunZhixingB1:
|
|||||||
returns so the UI never mistakes a lost worker exception for success.
|
returns so the UI never mistakes a lost worker exception for success.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
stocks = _unique_stocks(prepared.source.stocks)
|
||||||
evaluated_count = 0
|
evaluated_count = 0
|
||||||
selected_stock_count = 0
|
selected_stock_count = 0
|
||||||
signal_count = 0
|
signal_count = 0
|
||||||
failed_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:
|
try:
|
||||||
for stock in prepared.source.stocks:
|
with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
|
||||||
try:
|
for batch_stocks in _chunks(stocks, self.batch_size):
|
||||||
evaluation = self.evaluator.execute(
|
read_started = time.perf_counter()
|
||||||
stock.ts_code,
|
histories = self._load_histories(
|
||||||
|
batch_stocks,
|
||||||
prepared.source.target_trade_date,
|
prepared.source.target_trade_date,
|
||||||
)
|
)
|
||||||
except Exception as exc: # noqa: BLE001 - isolate one stock from the batch
|
read_seconds += time.perf_counter() - read_started
|
||||||
logger.exception(
|
history_rows += sum(
|
||||||
"selection_item_failed run_id=%s ts_code=%s",
|
len(history.bars) for history in histories if history is not None
|
||||||
prepared.run.id,
|
)
|
||||||
|
|
||||||
|
evaluate_started = time.perf_counter()
|
||||||
|
items = tuple(
|
||||||
|
_to_item(
|
||||||
stock.ts_code,
|
stock.ts_code,
|
||||||
|
stock.name,
|
||||||
|
evaluation,
|
||||||
)
|
)
|
||||||
evaluation = SelectionEvaluation(
|
for stock, evaluation in zip(
|
||||||
ts_code=stock.ts_code,
|
batch_stocks,
|
||||||
target_trade_date=prepared.source.target_trade_date,
|
executor.map(
|
||||||
status="data_error",
|
self._evaluate_stock,
|
||||||
reason=_safe_item_error(exc),
|
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)
|
evaluate_seconds += time.perf_counter() - evaluate_started
|
||||||
evaluated_count += 1
|
|
||||||
selected_stock_count += evaluation.status == "selected"
|
evaluated_count += len(items)
|
||||||
signal_count += len(evaluation.signals)
|
selected_stock_count += sum(item.status == "selected" for item in items)
|
||||||
failed_count += evaluation.status in _FAILURE_STATUSES
|
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)
|
status = _run_status(evaluated_count, failed_count)
|
||||||
self.store.finish_run(
|
self.store.finish_run(
|
||||||
@@ -121,7 +158,12 @@ class RunZhixingB1:
|
|||||||
failed_count=failed_count,
|
failed_count=failed_count,
|
||||||
)
|
)
|
||||||
except Exception as exc: # noqa: BLE001 - worker boundary must persist failure state
|
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:
|
try:
|
||||||
self.store.finish_run(
|
self.store.finish_run(
|
||||||
prepared.run.id,
|
prepared.run.id,
|
||||||
@@ -134,7 +176,94 @@ class RunZhixingB1:
|
|||||||
error_message=str(exc),
|
error_message=str(exc),
|
||||||
)
|
)
|
||||||
except Exception: # noqa: BLE001 - preserve the original worker failure
|
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(
|
def get_run(
|
||||||
self,
|
self,
|
||||||
@@ -187,6 +316,35 @@ def _safe_item_error(error: Exception) -> str:
|
|||||||
return " ".join(str(error).split())[:500] or error.__class__.__name__
|
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__ = [
|
__all__ = [
|
||||||
"PreparedSelectionRun",
|
"PreparedSelectionRun",
|
||||||
"RunZhixingB1",
|
"RunZhixingB1",
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import date, datetime
|
from datetime import date, datetime
|
||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
@@ -150,3 +151,19 @@ class SelectionUniverseReader(Protocol):
|
|||||||
) -> SelectionExecutionSource: ...
|
) -> SelectionExecutionSource: ...
|
||||||
|
|
||||||
def load_history(self, ts_code: str, target_trade_date: date) -> StockHistory: ...
|
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"]
|
||||||
+163
-31
@@ -2,9 +2,11 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Generator, Sequence
|
||||||
|
from contextlib import contextmanager
|
||||||
from datetime import date, datetime
|
from datetime import date, datetime
|
||||||
from decimal import Decimal, InvalidOperation
|
from decimal import Decimal, InvalidOperation
|
||||||
from typing import cast
|
from typing import Any, cast
|
||||||
|
|
||||||
import psycopg
|
import psycopg
|
||||||
|
|
||||||
@@ -12,6 +14,7 @@ from ....bootstrap.config import Settings
|
|||||||
from ..domain.models import SelectionBar, SelectionDailyBasic, StockHistory
|
from ..domain.models import SelectionBar, SelectionDailyBasic, StockHistory
|
||||||
from ..domain.ports import MarketDataReaderError
|
from ..domain.ports import MarketDataReaderError
|
||||||
from ..domain.runs import SelectionExecutionSource, SelectionStock
|
from ..domain.runs import SelectionExecutionSource, SelectionStock
|
||||||
|
from .postgres_pool import SelectionConnectionPool, SelectionPostgresPool
|
||||||
|
|
||||||
|
|
||||||
class SelectionReaderError(MarketDataReaderError):
|
class SelectionReaderError(MarketDataReaderError):
|
||||||
@@ -23,6 +26,25 @@ class SelectionMarketDataNotReady(MarketDataReaderError):
|
|||||||
|
|
||||||
|
|
||||||
_HISTORY_QUERY = """
|
_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
|
SELECT
|
||||||
bar.ts_code,
|
bar.ts_code,
|
||||||
stock.name,
|
stock.name,
|
||||||
@@ -111,10 +133,30 @@ def _as_float(value: object) -> float | None:
|
|||||||
class PostgresMarketDataReader:
|
class PostgresMarketDataReader:
|
||||||
"""Load qfq bars and same-day basic facts without writing market data."""
|
"""Load qfq bars and same-day basic facts without writing market data."""
|
||||||
|
|
||||||
def __init__(self, settings: Settings | str) -> None:
|
def __init__(
|
||||||
"""Create a reader from injected settings or a compatible URL string."""
|
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
|
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:
|
def load_history(self, ts_code: str, target_trade_date: date) -> StockHistory:
|
||||||
"""Read all retained qfq rows through the explicit target date.
|
"""Read all retained qfq rows through the explicit target date.
|
||||||
@@ -133,34 +175,57 @@ class PostgresMarketDataReader:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
try:
|
try:
|
||||||
with psycopg.connect(self.database_url) as connection:
|
with self._connection() as connection:
|
||||||
rows = connection.execute(
|
rows = connection.execute(
|
||||||
_HISTORY_QUERY,
|
_SINGLE_HISTORY_QUERY,
|
||||||
(ts_code, target_trade_date),
|
(ts_code, target_trade_date),
|
||||||
).fetchall()
|
).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(
|
raise SelectionReaderError(
|
||||||
f"failed to load market history for {ts_code} at {target_trade_date.isoformat()}"
|
f"failed to load market history for {ts_code} at {target_trade_date.isoformat()}"
|
||||||
) from exc
|
) from exc
|
||||||
|
|
||||||
bars: dict[date, SelectionBar] = {}
|
return self._histories_from_rows(
|
||||||
daily_basic: dict[date, SelectionDailyBasic] = {}
|
rows,
|
||||||
name = ""
|
(SelectionStock(ts_code, ""),),
|
||||||
for raw_row in rows:
|
target_trade_date,
|
||||||
row = cast(tuple[object, ...], raw_row)
|
)[0]
|
||||||
row_code, row_name, bar, basic = self._map_row(row, ts_code)
|
|
||||||
if row_code != ts_code:
|
def load_histories(
|
||||||
raise ValueError(f"reader returned unexpected stock code: {row_code}")
|
self,
|
||||||
name = row_name or name
|
stocks: Sequence[SelectionStock] | Sequence[str],
|
||||||
if bar.trade_date <= target_trade_date:
|
target_trade_date: date,
|
||||||
bars[bar.trade_date] = bar
|
) -> tuple[StockHistory, ...]:
|
||||||
daily_basic[bar.trade_date] = basic
|
"""Read one bounded stock chunk with one parameterized qfq query.
|
||||||
return StockHistory(
|
|
||||||
ts_code=ts_code,
|
Historical daily-basic values are deliberately not joined here: B1
|
||||||
name=name,
|
only needs OHLCV for its historical formula. The execution-source
|
||||||
bars=tuple(bars[trade_date] for trade_date in sorted(bars)),
|
query still requires a complete target-day basic row before a stock is
|
||||||
daily_basic={trade_date: daily_basic[trade_date] for trade_date in sorted(daily_basic)},
|
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(
|
def load_execution_source(
|
||||||
self,
|
self,
|
||||||
@@ -187,7 +252,7 @@ class PostgresMarketDataReader:
|
|||||||
if strategy != "zhixing_b1":
|
if strategy != "zhixing_b1":
|
||||||
raise SelectionMarketDataNotReady(f"unsupported selection strategy: {strategy}")
|
raise SelectionMarketDataNotReady(f"unsupported selection strategy: {strategy}")
|
||||||
try:
|
try:
|
||||||
with psycopg.connect(self.database_url) as connection:
|
with self._connection() as connection:
|
||||||
source_row = connection.execute(_SOURCE_QUERY, (target_trade_date,)).fetchone()
|
source_row = connection.execute(_SOURCE_QUERY, (target_trade_date,)).fetchone()
|
||||||
if source_row is None:
|
if source_row is None:
|
||||||
raise SelectionMarketDataNotReady(
|
raise SelectionMarketDataNotReady(
|
||||||
@@ -197,9 +262,9 @@ class PostgresMarketDataReader:
|
|||||||
_ELIGIBLE_STOCKS_QUERY,
|
_ELIGIBLE_STOCKS_QUERY,
|
||||||
(target_trade_date, target_trade_date),
|
(target_trade_date, target_trade_date),
|
||||||
).fetchall()
|
).fetchall()
|
||||||
except SelectionMarketDataNotReady:
|
except (SelectionMarketDataNotReady, SelectionReaderError):
|
||||||
raise
|
raise
|
||||||
except psycopg.Error as exc:
|
except Exception as exc: # noqa: BLE001 - redact driver/pool details at the port boundary
|
||||||
raise SelectionReaderError(
|
raise SelectionReaderError(
|
||||||
f"failed to load selection source at {target_trade_date.isoformat()}"
|
f"failed to load selection source at {target_trade_date.isoformat()}"
|
||||||
) from exc
|
) from exc
|
||||||
@@ -224,8 +289,8 @@ class PostgresMarketDataReader:
|
|||||||
def _map_row(
|
def _map_row(
|
||||||
row: tuple[object, ...],
|
row: tuple[object, ...],
|
||||||
expected_code: str,
|
expected_code: str,
|
||||||
) -> tuple[str, str, SelectionBar, SelectionDailyBasic]:
|
) -> tuple[str, str, SelectionBar, SelectionDailyBasic | None]:
|
||||||
"""Map the current query row, tolerating a legacy test row without name."""
|
"""Map qfq OHLCV rows and tolerate the legacy basic-join test shape."""
|
||||||
|
|
||||||
if len(row) >= 10:
|
if len(row) >= 10:
|
||||||
code, raw_name, raw_date = row[0], row[1], row[2]
|
code, raw_name, raw_date = row[0], row[1], row[2]
|
||||||
@@ -233,13 +298,19 @@ class PostgresMarketDataReader:
|
|||||||
elif len(row) >= 9:
|
elif len(row) >= 9:
|
||||||
code, raw_name, raw_date = row[0], "", row[1]
|
code, raw_name, raw_date = row[0], "", row[1]
|
||||||
values = row[2:]
|
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:
|
else:
|
||||||
raise ValueError("market history row has too few columns")
|
raise ValueError("market history row has too few columns")
|
||||||
row_code = str(code or expected_code)
|
row_code = str(code or expected_code)
|
||||||
name = str(raw_name or "")
|
name = str(raw_name or "")
|
||||||
trade_date = _as_date(raw_date)
|
trade_date = _as_date(raw_date)
|
||||||
if len(values) < 7:
|
if len(values) < 5:
|
||||||
raise ValueError("market history row is missing OHLCV/basic columns")
|
raise ValueError("market history row is missing OHLCV columns")
|
||||||
bar = SelectionBar(
|
bar = SelectionBar(
|
||||||
trade_date=trade_date,
|
trade_date=trade_date,
|
||||||
open=_as_float(values[0]),
|
open=_as_float(values[0]),
|
||||||
@@ -248,9 +319,70 @@ class PostgresMarketDataReader:
|
|||||||
close=_as_float(values[3]),
|
close=_as_float(values[3]),
|
||||||
volume=_as_float(values[4]),
|
volume=_as_float(values[4]),
|
||||||
)
|
)
|
||||||
basic = SelectionDailyBasic(
|
basic = (
|
||||||
|
SelectionDailyBasic(
|
||||||
trade_date=trade_date,
|
trade_date=trade_date,
|
||||||
turnover_rate=_as_float(values[5]),
|
turnover_rate=_as_float(values[5]),
|
||||||
total_mv=_as_float(values[6]),
|
total_mv=_as_float(values[6]),
|
||||||
)
|
)
|
||||||
|
if len(values) >= 7
|
||||||
|
else None
|
||||||
|
)
|
||||||
return row_code, name, bar, basic
|
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
|
||||||
|
|||||||
+101
-38
@@ -4,7 +4,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import json
|
import json
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from collections.abc import Generator, Mapping
|
from collections.abc import Generator, Mapping, Sequence
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from datetime import date, datetime
|
from datetime import date, datetime
|
||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
@@ -28,6 +28,7 @@ from ..domain.runs import (
|
|||||||
SelectionRunStoreError,
|
SelectionRunStoreError,
|
||||||
)
|
)
|
||||||
from ..domain.zhixing_b1 import ZHIXING_B1_SIGNAL_ORDER
|
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)}
|
_SIGNAL_PRIORITY = {category: index for index, category in enumerate(ZHIXING_B1_SIGNAL_ORDER)}
|
||||||
_CATEGORY_PREFIXES = {
|
_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):
|
class PostgresSelectionRunRepository(SelectionRunStore):
|
||||||
"""Persist one current result attempt per strategy and target date."""
|
"""Persist one current result attempt per strategy and target date."""
|
||||||
|
|
||||||
def __init__(self, database_url: str) -> None:
|
def __init__(
|
||||||
"""Create the adapter with an injected PostgreSQL URL."""
|
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
|
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(
|
def prepare_run(
|
||||||
self,
|
self,
|
||||||
@@ -137,21 +177,22 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def record_item(self, run_id: str, item: SelectionRunItem) -> None:
|
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."""
|
||||||
|
|
||||||
try:
|
self.record_items(run_id, (item,))
|
||||||
with self._connection() as connection, connection.transaction():
|
|
||||||
connection.execute(
|
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.
|
||||||
"""
|
"""
|
||||||
INSERT INTO selection_run_item
|
|
||||||
(run_id, ts_code, name, status, signal_count, reason)
|
if not items:
|
||||||
VALUES (%s, %s, %s, %s, %s, %s)
|
return
|
||||||
ON CONFLICT (run_id, ts_code) DO UPDATE SET
|
item_values = tuple(
|
||||||
name = EXCLUDED.name,
|
|
||||||
status = EXCLUDED.status,
|
|
||||||
signal_count = EXCLUDED.signal_count,
|
|
||||||
reason = EXCLUDED.reason
|
|
||||||
""",
|
|
||||||
(
|
(
|
||||||
run_id,
|
run_id,
|
||||||
item.ts_code,
|
item.ts_code,
|
||||||
@@ -159,26 +200,10 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
|||||||
item.status,
|
item.status,
|
||||||
item.signal_count,
|
item.signal_count,
|
||||||
item.reason,
|
item.reason,
|
||||||
),
|
|
||||||
)
|
)
|
||||||
connection.execute(
|
for item in items
|
||||||
"DELETE FROM selection_signal WHERE run_id = %s AND ts_code = %s",
|
|
||||||
(run_id, item.ts_code),
|
|
||||||
)
|
)
|
||||||
for signal in item.signals:
|
signal_values = tuple(
|
||||||
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,
|
run_id,
|
||||||
signal.ts_code,
|
signal.ts_code,
|
||||||
@@ -188,11 +213,26 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
|||||||
signal.category.value,
|
signal.category.value,
|
||||||
signal.close,
|
signal.close,
|
||||||
Jsonb(dict(signal.details)),
|
Jsonb(dict(signal.details)),
|
||||||
),
|
|
||||||
)
|
)
|
||||||
except psycopg.Error as exc:
|
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(
|
||||||
|
"DELETE FROM selection_signal WHERE run_id = %s AND ts_code = ANY(%s)",
|
||||||
|
(run_id, codes),
|
||||||
|
)
|
||||||
|
_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(
|
raise SelectionRunStoreError(
|
||||||
f"failed to persist selection item {item.ts_code}"
|
f"failed to persist selection item {code_context}"
|
||||||
) from exc
|
) from exc
|
||||||
|
|
||||||
def finish_run(
|
def finish_run(
|
||||||
@@ -398,12 +438,35 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
|||||||
"""Translate psycopg failures without exposing driver details."""
|
"""Translate psycopg failures without exposing driver details."""
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
if self.pool is None:
|
||||||
with psycopg.connect(self.database_url) as connection:
|
with psycopg.connect(self.database_url) as connection:
|
||||||
yield connection
|
yield connection
|
||||||
except psycopg.Error as exc:
|
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
|
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:
|
def _signal_from_row(row: tuple[object, ...]) -> SelectionSignal:
|
||||||
"""Map a persisted signal row back to the domain signal model."""
|
"""Map a persisted signal row back to the domain signal model."""
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
"""HTTP presentation for persisted strategy execution results."""
|
"""HTTP presentation for persisted strategy execution results."""
|
||||||
|
|
||||||
|
import atexit
|
||||||
|
import threading
|
||||||
from datetime import date, datetime
|
from datetime import date, datetime
|
||||||
from typing import Annotated, Literal
|
from typing import Annotated, Literal
|
||||||
|
|
||||||
@@ -17,6 +19,7 @@ from zhixing_server.modules.selection.domain.runs import (
|
|||||||
SelectionRunInProgress,
|
SelectionRunInProgress,
|
||||||
SelectionRunStoreError,
|
SelectionRunStoreError,
|
||||||
)
|
)
|
||||||
|
from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool
|
||||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||||
PostgresMarketDataReader,
|
PostgresMarketDataReader,
|
||||||
SelectionMarketDataNotReady,
|
SelectionMarketDataNotReady,
|
||||||
@@ -27,6 +30,8 @@ from zhixing_server.modules.selection.infrastructure.postgres_runs import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
selection_router = APIRouter()
|
selection_router = APIRouter()
|
||||||
|
_SELECTION_POOL_CACHE_LOCK = threading.Lock()
|
||||||
|
_SELECTION_POOL_CACHE: dict[tuple[str, int], SelectionPostgresPool] = {}
|
||||||
|
|
||||||
StrategyValue = Literal["zhixing_b1"]
|
StrategyValue = Literal["zhixing_b1"]
|
||||||
SelectionStatusValue = Literal[
|
SelectionStatusValue = Literal[
|
||||||
@@ -117,11 +122,45 @@ class SelectionResultsResponse(BaseModel):
|
|||||||
def get_selection_service(
|
def get_selection_service(
|
||||||
settings: Annotated[Settings, Depends(get_settings)],
|
settings: Annotated[Settings, Depends(get_settings)],
|
||||||
) -> RunZhixingB1:
|
) -> RunZhixingB1:
|
||||||
"""Build one request-scoped selection application service."""
|
"""Build the selection service on top of process-scoped shared resources."""
|
||||||
|
|
||||||
reader = PostgresMarketDataReader(settings)
|
pool = get_selection_postgres_pool(settings)
|
||||||
store = PostgresSelectionRunRepository(settings.database_url)
|
reader = PostgresMarketDataReader(settings, pool=pool)
|
||||||
return RunZhixingB1(reader, store)
|
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(
|
@selection_router.post(
|
||||||
@@ -292,5 +331,6 @@ __all__ = [
|
|||||||
"SelectionRunAcceptedResponse",
|
"SelectionRunAcceptedResponse",
|
||||||
"SelectionRunRequest",
|
"SelectionRunRequest",
|
||||||
"get_selection_service",
|
"get_selection_service",
|
||||||
|
"get_selection_postgres_pool",
|
||||||
"selection_router",
|
"selection_router",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -3,9 +3,12 @@
|
|||||||
from datetime import date
|
from datetime import date
|
||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
|
|
||||||
|
import pytest
|
||||||
from fastapi.testclient import TestClient
|
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.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.application.run import PreparedSelectionRun
|
||||||
from zhixing_server.modules.selection.domain.models import SelectionSignal, ZhixingB1Category
|
from zhixing_server.modules.selection.domain.models import SelectionSignal, ZhixingB1Category
|
||||||
from zhixing_server.modules.selection.domain.runs import (
|
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 (
|
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||||
SelectionMarketDataNotReady,
|
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)
|
TARGET = date(2026, 8, 8)
|
||||||
|
|
||||||
@@ -128,6 +134,17 @@ def _client(service: FakeSelectionService) -> TestClient:
|
|||||||
return TestClient(app)
|
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:
|
def test_trigger_returns_accepted_run_and_schedules_execution() -> None:
|
||||||
service = FakeSelectionService()
|
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."""
|
"""PostgreSQL reader contract tests using a fake connection."""
|
||||||
|
|
||||||
|
from collections.abc import Generator
|
||||||
|
from contextlib import contextmanager
|
||||||
from datetime import date
|
from datetime import date
|
||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
from typing import cast
|
from typing import cast
|
||||||
@@ -7,6 +9,8 @@ from typing import cast
|
|||||||
import psycopg
|
import psycopg
|
||||||
import pytest
|
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 (
|
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||||
PostgresMarketDataReader,
|
PostgresMarketDataReader,
|
||||||
SelectionMarketDataNotReady,
|
SelectionMarketDataNotReady,
|
||||||
@@ -39,6 +43,23 @@ class FakeResult:
|
|||||||
return self.rows
|
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:
|
def test_reader_parameterizes_target_and_maps_left_join(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
connection = FakeConnection(
|
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)
|
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:
|
class SourceConnection:
|
||||||
def __init__(self, source_row: tuple[object, ...] | None) -> None:
|
def __init__(self, source_row: tuple[object, ...] | None) -> None:
|
||||||
self.source_row = source_row
|
self.source_row = source_row
|
||||||
|
|||||||
@@ -61,6 +61,33 @@ class FakeConnection:
|
|||||||
return FakeResult()
|
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:
|
def _source() -> SelectionExecutionSource:
|
||||||
return SelectionExecutionSource(
|
return SelectionExecutionSource(
|
||||||
market_sync_batch_id="market-run-1",
|
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)
|
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:
|
class LoadConnection:
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self.statements: list[tuple[str, tuple[object, ...]]] = []
|
self.statements: list[tuple[str, tuple[object, ...]]] = []
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
"""Application tests for persisted whole-universe selection runs."""
|
"""Application tests for persisted whole-universe selection runs."""
|
||||||
|
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
from datetime import date
|
from datetime import date
|
||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
@@ -42,6 +44,30 @@ class FakeReader:
|
|||||||
raise AssertionError("the fake evaluator should be used")
|
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:
|
class FakeStore:
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self.items: list[SelectionRunItem] = []
|
self.items: list[SelectionRunItem] = []
|
||||||
@@ -114,6 +140,26 @@ class FakeStore:
|
|||||||
return None
|
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:
|
class FakeEvaluator:
|
||||||
def __init__(self, results: dict[str, SelectionEvaluation]) -> None:
|
def __init__(self, results: dict[str, SelectionEvaluation]) -> None:
|
||||||
self.results = results
|
self.results = results
|
||||||
@@ -129,6 +175,29 @@ class RaisingEvaluator:
|
|||||||
return SelectionEvaluation(ts_code, target_trade_date, "no_signal")
|
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:
|
def _source() -> SelectionExecutionSource:
|
||||||
return SelectionExecutionSource(
|
return SelectionExecutionSource(
|
||||||
market_sync_batch_id="market-run-1",
|
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[0:2] == ("run-1", "partial_success")
|
||||||
assert store.finished[2]["evaluated_count"] == 2
|
assert store.finished[2]["evaluated_count"] == 2
|
||||||
assert store.finished[2]["failed_count"] == 1
|
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