perf(selection): 优化选股执行性能

This commit is contained in:
yuxuanhui
2026-08-12 09:45:16 +08:00
parent dd04933d63
commit 8963c067b3
22 changed files with 1333 additions and 120 deletions
+2
View File
@@ -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
@@ -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 适配器模式,避免引入无必要抽象。"}
@@ -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 | 采用 | 用户确认的初始并发度,后续以基准调整 |
@@ -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 和应用组合模式,避免重复基础设施。"}
@@ -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;两者保持配置化,后续以基准结果调整。
@@ -0,0 +1,26 @@
{
"id": "optimize-selection-execution-performance",
"name": "optimize-selection-execution-performance",
"title": "优化选股执行性能",
"description": "",
"status": "in_progress",
"dev_type": null,
"scope": null,
"package": null,
"priority": "P2",
"creator": "yuxuanhui",
"assignee": "yuxuanhui",
"createdAt": "2026-08-12",
"completedAt": null,
"branch": null,
"base_branch": "main",
"worktree_path": null,
"commit": null,
"pr_url": null,
"subtasks": [],
"children": [],
"parent": null,
"relatedFiles": [],
"notes": "",
"meta": {}
}
+2
View File
@@ -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:
+2
View File
@@ -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,
read_seconds += time.perf_counter() - read_started
history_rows += sum(
len(history.bars) for history in histories if history is not None
)
evaluate_started = time.perf_counter()
items = tuple(
_to_item(
stock.ts_code,
stock.name,
evaluation,
)
evaluation = SelectionEvaluation(
ts_code=stock.ts_code,
target_trade_date=prepared.source.target_trade_date,
status="data_error",
reason=_safe_item_error(exc),
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"]
@@ -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(
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
@@ -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,21 +177,22 @@ 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."""
try:
with self._connection() as connection, connection.transaction():
connection.execute(
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.
"""
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
""",
if not items:
return
item_values = tuple(
(
run_id,
item.ts_code,
@@ -159,26 +200,10 @@ class PostgresSelectionRunRepository(SelectionRunStore):
item.status,
item.signal_count,
item.reason,
),
)
connection.execute(
"DELETE FROM selection_signal WHERE run_id = %s AND ts_code = %s",
(run_id, item.ts_code),
for item in items
)
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
""",
signal_values = tuple(
(
run_id,
signal.ts_code,
@@ -188,11 +213,26 @@ class PostgresSelectionRunRepository(SelectionRunStore):
signal.category.value,
signal.close,
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(
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:
if self.pool is None:
with psycopg.connect(self.database_url) as 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
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",
]
+18 -1
View File
@@ -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