From 8963c067b3afae559f0d4c57ff0c472b0651a142 Mon Sep 17 00:00:00 2001 From: yuxuanhui Date: Wed, 12 Aug 2026 09:45:16 +0800 Subject: [PATCH] =?UTF-8?q?perf(selection):=20=E4=BC=98=E5=8C=96=E9=80=89?= =?UTF-8?q?=E8=82=A1=E6=89=A7=E8=A1=8C=E6=80=A7=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .env.example | 2 + .../check.jsonl | 4 + .../design.md | 151 +++++++++++++ .../implement.jsonl | 4 + .../implement.md | 60 +++++ .../prd.md | 87 ++++++++ .../task.json | 26 +++ docker-compose.dev.yml | 2 + docker-compose.prod.yml | 2 + .../src/zhixing_server/bootstrap/config.py | 2 + .../modules/selection/application/evaluate.py | 10 + .../modules/selection/application/run.py | 208 +++++++++++++++--- .../modules/selection/domain/runs.py | 17 ++ .../selection/infrastructure/postgres_pool.py | 96 ++++++++ .../infrastructure/postgres_reader.py | 200 ++++++++++++++--- .../selection/infrastructure/postgres_runs.py | 175 ++++++++++----- .../modules/selection/presentation/http.py | 48 +++- zhixing-server/tests/test_selection_http.py | 19 +- .../unit/selection/test_postgres_pool.py | 45 ++++ .../unit/selection/test_postgres_reader.py | 79 +++++++ .../unit/selection/test_postgres_runs.py | 85 +++++++ .../tests/unit/selection/test_run.py | 131 +++++++++++ 22 files changed, 1333 insertions(+), 120 deletions(-) create mode 100644 .trellis/tasks/08-12-optimize-selection-execution-performance/check.jsonl create mode 100644 .trellis/tasks/08-12-optimize-selection-execution-performance/design.md create mode 100644 .trellis/tasks/08-12-optimize-selection-execution-performance/implement.jsonl create mode 100644 .trellis/tasks/08-12-optimize-selection-execution-performance/implement.md create mode 100644 .trellis/tasks/08-12-optimize-selection-execution-performance/prd.md create mode 100644 .trellis/tasks/08-12-optimize-selection-execution-performance/task.json create mode 100644 zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_pool.py create mode 100644 zhixing-server/tests/unit/selection/test_postgres_pool.py diff --git a/.env.example b/.env.example index 6b201f6..3d1199c 100644 --- a/.env.example +++ b/.env.example @@ -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 diff --git a/.trellis/tasks/08-12-optimize-selection-execution-performance/check.jsonl b/.trellis/tasks/08-12-optimize-selection-execution-performance/check.jsonl new file mode 100644 index 0000000..ffae1c3 --- /dev/null +++ b/.trellis/tasks/08-12-optimize-selection-execution-performance/check.jsonl @@ -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 适配器模式,避免引入无必要抽象。"} diff --git a/.trellis/tasks/08-12-optimize-selection-execution-performance/design.md b/.trellis/tasks/08-12-optimize-selection-execution-performance/design.md new file mode 100644 index 0000000..55d7453 --- /dev/null +++ b/.trellis/tasks/08-12-optimize-selection-execution-performance/design.md @@ -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 | 采用 | 用户确认的初始并发度,后续以基准调整 | diff --git a/.trellis/tasks/08-12-optimize-selection-execution-performance/implement.jsonl b/.trellis/tasks/08-12-optimize-selection-execution-performance/implement.jsonl new file mode 100644 index 0000000..3e0442c --- /dev/null +++ b/.trellis/tasks/08-12-optimize-selection-execution-performance/implement.jsonl @@ -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 和应用组合模式,避免重复基础设施。"} diff --git a/.trellis/tasks/08-12-optimize-selection-execution-performance/implement.md b/.trellis/tasks/08-12-optimize-selection-execution-performance/implement.md new file mode 100644 index 0000000..eca047c --- /dev/null +++ b/.trellis/tasks/08-12-optimize-selection-execution-performance/implement.md @@ -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。 diff --git a/.trellis/tasks/08-12-optimize-selection-execution-performance/prd.md b/.trellis/tasks/08-12-optimize-selection-execution-performance/prd.md new file mode 100644 index 0000000..1fa24ff --- /dev/null +++ b/.trellis/tasks/08-12-optimize-selection-execution-performance/prd.md @@ -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;两者保持配置化,后续以基准结果调整。 diff --git a/.trellis/tasks/08-12-optimize-selection-execution-performance/task.json b/.trellis/tasks/08-12-optimize-selection-execution-performance/task.json new file mode 100644 index 0000000..d940002 --- /dev/null +++ b/.trellis/tasks/08-12-optimize-selection-execution-performance/task.json @@ -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": {} +} \ No newline at end of file diff --git a/docker-compose.dev.yml b/docker-compose.dev.yml index 6386ad0..605f5a8 100644 --- a/docker-compose.dev.yml +++ b/docker-compose.dev.yml @@ -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: diff --git a/docker-compose.prod.yml b/docker-compose.prod.yml index b61c934..01d77a4 100644 --- a/docker-compose.prod.yml +++ b/docker-compose.prod.yml @@ -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: diff --git a/zhixing-server/src/zhixing_server/bootstrap/config.py b/zhixing-server/src/zhixing_server/bootstrap/config.py index 0455cc3..afa7572 100644 --- a/zhixing-server/src/zhixing_server/bootstrap/config.py +++ b/zhixing-server/src/zhixing_server/bootstrap/config.py @@ -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", diff --git a/zhixing-server/src/zhixing_server/modules/selection/application/evaluate.py b/zhixing-server/src/zhixing_server/modules/selection/application/evaluate.py index 429d7ef..17ef3fd 100644 --- a/zhixing-server/src/zhixing_server/modules/selection/application/evaluate.py +++ b/zhixing-server/src/zhixing_server/modules/selection/application/evaluate.py @@ -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) diff --git a/zhixing-server/src/zhixing_server/modules/selection/application/run.py b/zhixing-server/src/zhixing_server/modules/selection/application/run.py index 6d19ecc..8f87e02 100644 --- a/zhixing-server/src/zhixing_server/modules/selection/application/run.py +++ b/zhixing-server/src/zhixing_server/modules/selection/application/run.py @@ -3,12 +3,17 @@ from __future__ import annotations import logging +import time +from collections.abc import Callable, Sequence +from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from datetime import date -from typing import Literal, Protocol +from typing import Literal, Protocol, cast -from ..domain.models import SelectionEvaluation +from ..domain.models import SelectionEvaluation, StockHistory from ..domain.runs import ( + BatchSelectionRunStore, + BatchSelectionUniverseReader, SelectionExecutionSource, SelectionRerunRequired, SelectionResultQuery, @@ -17,6 +22,7 @@ from ..domain.runs import ( SelectionRunItem, SelectionRunStatus, SelectionRunStore, + SelectionStock, SelectionUniverseReader, ) from .evaluate import EvaluateZhixingB1 @@ -48,12 +54,21 @@ class RunZhixingB1: reader: SelectionUniverseReader, store: SelectionRunStore, evaluator: SelectionEvaluator | None = None, + *, + max_workers: int = 4, + batch_size: int = 200, ) -> None: - """Inject storage ports and optionally a test evaluator.""" + """Inject storage ports and configure bounded chunk execution.""" + if max_workers < 1: + raise ValueError("max_workers must be at least 1") + if batch_size < 1: + raise ValueError("batch_size must be at least 1") self.reader = reader self.store = store self.evaluator = evaluator or EvaluateZhixingB1(reader) + self.max_workers = max_workers + self.batch_size = batch_size def prepare( self, @@ -81,35 +96,57 @@ class RunZhixingB1: returns so the UI never mistakes a lost worker exception for success. """ + stocks = _unique_stocks(prepared.source.stocks) evaluated_count = 0 selected_stock_count = 0 signal_count = 0 failed_count = 0 + history_rows = 0 + batch_count = _chunk_count(len(stocks), self.batch_size) + read_seconds = 0.0 + evaluate_seconds = 0.0 + persist_seconds = 0.0 try: - for stock in prepared.source.stocks: - try: - evaluation = self.evaluator.execute( - stock.ts_code, + with ThreadPoolExecutor(max_workers=self.max_workers) as executor: + for batch_stocks in _chunks(stocks, self.batch_size): + read_started = time.perf_counter() + histories = self._load_histories( + batch_stocks, prepared.source.target_trade_date, ) - except Exception as exc: # noqa: BLE001 - isolate one stock from the batch - logger.exception( - "selection_item_failed run_id=%s ts_code=%s", - prepared.run.id, - stock.ts_code, + read_seconds += time.perf_counter() - read_started + history_rows += sum( + len(history.bars) for history in histories if history is not None ) - evaluation = SelectionEvaluation( - ts_code=stock.ts_code, - target_trade_date=prepared.source.target_trade_date, - status="data_error", - reason=_safe_item_error(exc), + + evaluate_started = time.perf_counter() + items = tuple( + _to_item( + stock.ts_code, + stock.name, + evaluation, + ) + for stock, evaluation in zip( + batch_stocks, + executor.map( + self._evaluate_stock, + batch_stocks, + histories, + [prepared.source.target_trade_date] * len(batch_stocks), + ), + strict=True, + ) ) - item = _to_item(stock.ts_code, stock.name, evaluation) - self.store.record_item(prepared.run.id, item) - evaluated_count += 1 - selected_stock_count += evaluation.status == "selected" - signal_count += len(evaluation.signals) - failed_count += evaluation.status in _FAILURE_STATUSES + evaluate_seconds += time.perf_counter() - evaluate_started + + evaluated_count += len(items) + selected_stock_count += sum(item.status == "selected" for item in items) + signal_count += sum(item.signal_count for item in items) + failed_count += sum(item.status in _FAILURE_STATUSES for item in items) + + persist_started = time.perf_counter() + self._record_items(prepared.run.id, items) + persist_seconds += time.perf_counter() - persist_started status = _run_status(evaluated_count, failed_count) self.store.finish_run( @@ -121,7 +158,12 @@ class RunZhixingB1: failed_count=failed_count, ) except Exception as exc: # noqa: BLE001 - worker boundary must persist failure state - logger.exception("selection_run_failed run_id=%s", prepared.run.id) + logger.error( + "selection_run_failed run_id=%s error_type=%s reason=%s", + prepared.run.id, + exc.__class__.__name__, + _safe_item_error(exc), + ) try: self.store.finish_run( prepared.run.id, @@ -134,7 +176,94 @@ class RunZhixingB1: error_message=str(exc), ) except Exception: # noqa: BLE001 - preserve the original worker failure - logger.exception("selection_run_failure_persist_failed run_id=%s", prepared.run.id) + logger.error( + "selection_run_failure_persist_failed run_id=%s", + prepared.run.id, + ) + finally: + logger.info( + "selection_run_summary run_id=%s stock_count=%d history_rows=%d " + "batch_count=%d worker_count=%d read_seconds=%.3f " + "evaluate_seconds=%.3f persist_seconds=%.3f", + prepared.run.id, + len(stocks), + history_rows, + batch_count, + self.max_workers, + read_seconds, + evaluate_seconds, + persist_seconds, + ) + + def _load_histories( + self, + stocks: Sequence[SelectionStock], + target_trade_date: date, + ) -> tuple[StockHistory | None, ...]: + """Load one chunk when the reader supports it, with old-path fallback.""" + + typed_stocks = tuple(stocks) + loader = getattr(self.reader, "load_histories", None) + if callable(loader): + batch_reader = cast(BatchSelectionUniverseReader, self.reader) + loaded = batch_reader.load_histories(typed_stocks, target_trade_date) + histories_by_code = {history.ts_code: history for history in loaded} + return tuple( + histories_by_code.get( + stock.ts_code, + StockHistory(ts_code=stock.ts_code, name=stock.name), + ) + for stock in typed_stocks + ) + + if isinstance(self.evaluator, EvaluateZhixingB1): + return tuple( + self.reader.load_history(stock.ts_code, target_trade_date) for stock in typed_stocks + ) + return (None,) * len(typed_stocks) + + def _evaluate_stock( + self, + stock: SelectionStock, + history: StockHistory | None, + target_trade_date: date, + ) -> SelectionEvaluation: + """Evaluate one stock inside a worker and isolate its exception.""" + + ts_code = stock.ts_code + try: + execute_history: Callable[[StockHistory, date], SelectionEvaluation] | None = getattr( + self.evaluator, + "execute_history", + None, + ) + if history is not None and execute_history is not None: + return execute_history(history, target_trade_date) + return self.evaluator.execute(ts_code, target_trade_date) + except Exception as exc: # noqa: BLE001 - isolate one stock from the batch + logger.warning( + "selection_item_failed ts_code=%s error_type=%s reason=%s", + ts_code, + exc.__class__.__name__, + _safe_item_error(exc), + ) + return SelectionEvaluation( + ts_code=ts_code, + target_trade_date=target_trade_date, + status="data_error", + reason=_safe_item_error(exc), + ) + + def _record_items(self, run_id: str, items: Sequence[SelectionRunItem]) -> None: + """Use batch persistence while retaining the old single-item seam.""" + + record_items = getattr(self.store, "record_items", None) + if callable(record_items): + batch_store = cast(BatchSelectionRunStore, self.store) + batch_store.record_items(run_id, tuple(items)) + return + for item in items: + self.store.record_item(run_id, item) def get_run( self, @@ -187,6 +316,35 @@ def _safe_item_error(error: Exception) -> str: return " ".join(str(error).split())[:500] or error.__class__.__name__ +def _chunks( + values: Sequence[SelectionStock], + size: int, +) -> tuple[tuple[SelectionStock, ...], ...]: + """Split a stable stock sequence into bounded immutable chunks.""" + + return tuple(tuple(values[index : index + size]) for index in range(0, len(values), size)) + + +def _chunk_count(value_count: int, size: int) -> int: + """Return the number of chunks without materializing empty chunks.""" + + return (value_count + size - 1) // size + + +def _unique_stocks(stocks: Sequence[SelectionStock]) -> tuple[SelectionStock, ...]: + """Keep the first source row for each stock so it is evaluated once.""" + + seen: set[str] = set() + unique: list[SelectionStock] = [] + for stock in stocks: + ts_code = stock.ts_code + if ts_code in seen: + continue + seen.add(ts_code) + unique.append(stock) + return tuple(unique) + + __all__ = [ "PreparedSelectionRun", "RunZhixingB1", diff --git a/zhixing-server/src/zhixing_server/modules/selection/domain/runs.py b/zhixing-server/src/zhixing_server/modules/selection/domain/runs.py index 9f7826d..6452d94 100644 --- a/zhixing-server/src/zhixing_server/modules/selection/domain/runs.py +++ b/zhixing-server/src/zhixing_server/modules/selection/domain/runs.py @@ -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, ...]: ... diff --git a/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_pool.py b/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_pool.py new file mode 100644 index 0000000..4db71e2 --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_pool.py @@ -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"] diff --git a/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_reader.py b/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_reader.py index c46e4ba..c7123c1 100644 --- a/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_reader.py +++ b/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_reader.py @@ -2,9 +2,11 @@ from __future__ import annotations +from collections.abc import Generator, Sequence +from contextlib import contextmanager from datetime import date, datetime from decimal import Decimal, InvalidOperation -from typing import cast +from typing import Any, cast import psycopg @@ -12,6 +14,7 @@ from ....bootstrap.config import Settings from ..domain.models import SelectionBar, SelectionDailyBasic, StockHistory from ..domain.ports import MarketDataReaderError from ..domain.runs import SelectionExecutionSource, SelectionStock +from .postgres_pool import SelectionConnectionPool, SelectionPostgresPool class SelectionReaderError(MarketDataReaderError): @@ -23,6 +26,25 @@ class SelectionMarketDataNotReady(MarketDataReaderError): _HISTORY_QUERY = """ +SELECT + bar.ts_code, + stock.name, + bar.trade_date, + bar.open, + bar.high, + bar.low, + bar.close, + bar.vol +FROM market_daily_bar AS bar +LEFT JOIN market_stock AS stock + ON stock.ts_code = bar.ts_code +WHERE bar.ts_code = ANY(%s) + AND bar.source_adj = 'qfq' + AND bar.trade_date <= %s +ORDER BY bar.ts_code ASC, bar.trade_date ASC +""" + +_SINGLE_HISTORY_QUERY = """ SELECT bar.ts_code, stock.name, @@ -111,10 +133,30 @@ def _as_float(value: object) -> float | None: class PostgresMarketDataReader: """Load qfq bars and same-day basic facts without writing market data.""" - def __init__(self, settings: Settings | str) -> None: - """Create a reader from injected settings or a compatible URL string.""" + def __init__( + self, + settings: Settings | str, + *, + pool: SelectionPostgresPool | SelectionConnectionPool | None = None, + ) -> None: + """Create a reader from settings/URL and an optional shared pool. + + Omitting ``pool`` intentionally retains the direct ``psycopg.connect`` + path used by one-shot callers and existing adapter tests. The HTTP + composition layer always supplies the process-scoped selection pool. + """ self.database_url = settings.database_url if isinstance(settings, Settings) else settings + if isinstance(pool, SelectionPostgresPool): + self.pool: SelectionPostgresPool | None = pool + elif pool is not None: + self.pool = SelectionPostgresPool( + self.database_url, + max_connections=1, + pool=pool, + ) + else: + self.pool = None def load_history(self, ts_code: str, target_trade_date: date) -> StockHistory: """Read all retained qfq rows through the explicit target date. @@ -133,34 +175,57 @@ class PostgresMarketDataReader: """ try: - with psycopg.connect(self.database_url) as connection: + with self._connection() as connection: rows = connection.execute( - _HISTORY_QUERY, + _SINGLE_HISTORY_QUERY, (ts_code, target_trade_date), ).fetchall() - except psycopg.Error as exc: + except SelectionReaderError: + raise + except Exception as exc: # noqa: BLE001 - redact driver/pool details at the port boundary raise SelectionReaderError( f"failed to load market history for {ts_code} at {target_trade_date.isoformat()}" ) from exc - bars: dict[date, SelectionBar] = {} - daily_basic: dict[date, SelectionDailyBasic] = {} - name = "" - for raw_row in rows: - row = cast(tuple[object, ...], raw_row) - row_code, row_name, bar, basic = self._map_row(row, ts_code) - if row_code != ts_code: - raise ValueError(f"reader returned unexpected stock code: {row_code}") - name = row_name or name - if bar.trade_date <= target_trade_date: - bars[bar.trade_date] = bar - daily_basic[bar.trade_date] = basic - return StockHistory( - ts_code=ts_code, - name=name, - bars=tuple(bars[trade_date] for trade_date in sorted(bars)), - daily_basic={trade_date: daily_basic[trade_date] for trade_date in sorted(daily_basic)}, + return self._histories_from_rows( + rows, + (SelectionStock(ts_code, ""),), + target_trade_date, + )[0] + + def load_histories( + self, + stocks: Sequence[SelectionStock] | Sequence[str], + target_trade_date: date, + ) -> tuple[StockHistory, ...]: + """Read one bounded stock chunk with one parameterized qfq query. + + Historical daily-basic values are deliberately not joined here: B1 + only needs OHLCV for its historical formula. The execution-source + query still requires a complete target-day basic row before a stock is + admitted to a run. + """ + + normalized = tuple( + stock if isinstance(stock, SelectionStock) else SelectionStock(stock, "") + for stock in stocks ) + if not normalized: + return () + codes = [stock.ts_code for stock in normalized] + try: + with self._connection() as connection: + rows = connection.execute( + _HISTORY_QUERY, + (codes, target_trade_date), + ).fetchall() + except SelectionReaderError: + raise + except Exception as exc: # noqa: BLE001 - redact driver/pool details at the port boundary + raise SelectionReaderError( + f"failed to load market history batch at {target_trade_date.isoformat()}" + ) from exc + return self._histories_from_rows(rows, normalized, target_trade_date) def load_execution_source( self, @@ -187,7 +252,7 @@ class PostgresMarketDataReader: if strategy != "zhixing_b1": raise SelectionMarketDataNotReady(f"unsupported selection strategy: {strategy}") try: - with psycopg.connect(self.database_url) as connection: + with self._connection() as connection: source_row = connection.execute(_SOURCE_QUERY, (target_trade_date,)).fetchone() if source_row is None: raise SelectionMarketDataNotReady( @@ -197,9 +262,9 @@ class PostgresMarketDataReader: _ELIGIBLE_STOCKS_QUERY, (target_trade_date, target_trade_date), ).fetchall() - except SelectionMarketDataNotReady: + except (SelectionMarketDataNotReady, SelectionReaderError): raise - except psycopg.Error as exc: + except Exception as exc: # noqa: BLE001 - redact driver/pool details at the port boundary raise SelectionReaderError( f"failed to load selection source at {target_trade_date.isoformat()}" ) from exc @@ -224,8 +289,8 @@ class PostgresMarketDataReader: def _map_row( row: tuple[object, ...], expected_code: str, - ) -> tuple[str, str, SelectionBar, SelectionDailyBasic]: - """Map the current query row, tolerating a legacy test row without name.""" + ) -> tuple[str, str, SelectionBar, SelectionDailyBasic | None]: + """Map qfq OHLCV rows and tolerate the legacy basic-join test shape.""" if len(row) >= 10: code, raw_name, raw_date = row[0], row[1], row[2] @@ -233,13 +298,19 @@ class PostgresMarketDataReader: elif len(row) >= 9: code, raw_name, raw_date = row[0], "", row[1] values = row[2:] + elif len(row) >= 8: + code, raw_name, raw_date = row[0], row[1], row[2] + values = row[3:] + elif len(row) >= 7: + code, raw_name, raw_date = row[0], "", row[1] + values = row[2:] else: raise ValueError("market history row has too few columns") row_code = str(code or expected_code) name = str(raw_name or "") trade_date = _as_date(raw_date) - if len(values) < 7: - raise ValueError("market history row is missing OHLCV/basic columns") + if len(values) < 5: + raise ValueError("market history row is missing OHLCV columns") bar = SelectionBar( trade_date=trade_date, open=_as_float(values[0]), @@ -248,9 +319,70 @@ class PostgresMarketDataReader: close=_as_float(values[3]), volume=_as_float(values[4]), ) - basic = SelectionDailyBasic( - trade_date=trade_date, - turnover_rate=_as_float(values[5]), - total_mv=_as_float(values[6]), + basic = ( + SelectionDailyBasic( + trade_date=trade_date, + turnover_rate=_as_float(values[5]), + total_mv=_as_float(values[6]), + ) + if len(values) >= 7 + else None ) return row_code, name, bar, basic + + @classmethod + def _histories_from_rows( + cls, + rows: Sequence[object], + stocks: Sequence[SelectionStock], + target_trade_date: date, + ) -> tuple[StockHistory, ...]: + """Group sorted/possibly duplicated database rows by requested stock.""" + + requested = {stock.ts_code: stock for stock in stocks} + bars_by_code: dict[str, dict[date, SelectionBar]] = {code: {} for code in requested} + basics_by_code: dict[str, dict[date, SelectionDailyBasic]] = { + code: {} for code in requested + } + names = {stock.ts_code: stock.name for stock in stocks} + for raw_row in rows: + row = cast(tuple[object, ...], raw_row) + row_hint = str(row[0]) if row and row[0] is not None else "" + row_code, row_name, bar, basic = cls._map_row(row, row_hint) + if row_code not in requested: + raise ValueError(f"reader returned unexpected stock code: {row_code}") + names[row_code] = row_name or names[row_code] + if bar.trade_date <= target_trade_date: + bars_by_code[row_code][bar.trade_date] = bar + if basic is not None: + basics_by_code[row_code][bar.trade_date] = basic + histories: list[StockHistory] = [] + for stock in stocks: + code = stock.ts_code + bars = bars_by_code[code] + basics = basics_by_code[code] + histories.append( + StockHistory( + ts_code=code, + name=names[code], + bars=tuple(bars[trade_date] for trade_date in sorted(bars)), + daily_basic={trade_date: basics[trade_date] for trade_date in sorted(basics)}, + ) + ) + return tuple(histories) + + @contextmanager + def _connection(self) -> Generator[Any, None, None]: + """Borrow from the shared pool or use the legacy direct connection.""" + + try: + if self.pool is None: + with psycopg.connect(self.database_url) as connection: + yield connection + else: + with self.pool.connection() as connection: + yield connection + except (SelectionMarketDataNotReady, SelectionReaderError, ValueError): + raise + except Exception as exc: # noqa: BLE001 - normalize pool/driver failures + raise SelectionReaderError("selection database operation failed") from exc diff --git a/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_runs.py b/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_runs.py index 2a9fcb9..c694dc2 100644 --- a/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_runs.py +++ b/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_runs.py @@ -4,7 +4,7 @@ from __future__ import annotations import json from collections import defaultdict -from collections.abc import Generator, Mapping +from collections.abc import Generator, Mapping, Sequence from contextlib import contextmanager from datetime import date, datetime from decimal import Decimal @@ -28,6 +28,7 @@ from ..domain.runs import ( SelectionRunStoreError, ) from ..domain.zhixing_b1 import ZHIXING_B1_SIGNAL_ORDER +from .postgres_pool import SelectionConnectionPool, SelectionPostgresPool _SIGNAL_PRIORITY = {category: index for index, category in enumerate(ZHIXING_B1_SIGNAL_ORDER)} _CATEGORY_PREFIXES = { @@ -45,13 +46,52 @@ _SIGNAL_ORDER_SQL = ( ) +_ITEM_UPSERT = """ +INSERT INTO selection_run_item + (run_id, ts_code, name, status, signal_count, reason) +VALUES (%s, %s, %s, %s, %s, %s) +ON CONFLICT (run_id, ts_code) DO UPDATE SET + name = EXCLUDED.name, + status = EXCLUDED.status, + signal_count = EXCLUDED.signal_count, + reason = EXCLUDED.reason +""" +_SIGNAL_UPSERT = """ +INSERT INTO selection_signal + ( + run_id, ts_code, name, target_trade_date, strategy, + category, close, details + ) +VALUES (%s, %s, %s, %s, %s, %s, %s, %s) +ON CONFLICT (run_id, ts_code, category) DO UPDATE SET + name = EXCLUDED.name, + close = EXCLUDED.close, + details = EXCLUDED.details +""" + + class PostgresSelectionRunRepository(SelectionRunStore): """Persist one current result attempt per strategy and target date.""" - def __init__(self, database_url: str) -> None: - """Create the adapter with an injected PostgreSQL URL.""" + def __init__( + self, + database_url: str, + *, + pool: SelectionPostgresPool | SelectionConnectionPool | None = None, + ) -> None: + """Create the adapter with a URL and optional shared PostgreSQL pool.""" self.database_url = database_url + if isinstance(pool, SelectionPostgresPool): + self.pool: SelectionPostgresPool | None = pool + elif pool is not None: + self.pool = SelectionPostgresPool( + database_url, + max_connections=1, + pool=pool, + ) + else: + self.pool = None def prepare_run( self, @@ -137,62 +177,62 @@ class PostgresSelectionRunRepository(SelectionRunStore): ) def record_item(self, run_id: str, item: SelectionRunItem) -> None: - """Upsert one stock outcome and all of its independent signal rows.""" + """Persist one item through the batch path for compatibility.""" + self.record_items(run_id, (item,)) + + def record_items(self, run_id: str, items: Sequence[SelectionRunItem]) -> None: + """Persist one chunk in one transaction with set-based driver calls. + + Existing signal rows are removed before the upserts so retrying a + chunk cannot retain a category that disappeared from a recalculation. + ``executemany`` is used for both materialized tables; the small + fallback keeps the direct fake connections used by older tests usable. + """ + + if not items: + return + item_values = tuple( + ( + run_id, + item.ts_code, + item.name, + item.status, + item.signal_count, + item.reason, + ) + for item in items + ) + signal_values = tuple( + ( + run_id, + signal.ts_code, + signal.name, + signal.target_trade_date, + signal.strategy, + signal.category.value, + signal.close, + Jsonb(dict(signal.details)), + ) + for item in items + for signal in item.signals + ) + codes = [item.ts_code for item in items] try: with self._connection() as connection, connection.transaction(): connection.execute( - """ - INSERT INTO selection_run_item - (run_id, ts_code, name, status, signal_count, reason) - VALUES (%s, %s, %s, %s, %s, %s) - ON CONFLICT (run_id, ts_code) DO UPDATE SET - name = EXCLUDED.name, - status = EXCLUDED.status, - signal_count = EXCLUDED.signal_count, - reason = EXCLUDED.reason - """, - ( - run_id, - item.ts_code, - item.name, - item.status, - item.signal_count, - item.reason, - ), + "DELETE FROM selection_signal WHERE run_id = %s AND ts_code = ANY(%s)", + (run_id, codes), ) - connection.execute( - "DELETE FROM selection_signal WHERE run_id = %s AND ts_code = %s", - (run_id, item.ts_code), - ) - for signal in item.signals: - connection.execute( - """ - INSERT INTO selection_signal - ( - run_id, ts_code, name, target_trade_date, strategy, - category, close, details - ) - VALUES (%s, %s, %s, %s, %s, %s, %s, %s) - ON CONFLICT (run_id, ts_code, category) DO UPDATE SET - name = EXCLUDED.name, - close = EXCLUDED.close, - details = EXCLUDED.details - """, - ( - run_id, - signal.ts_code, - signal.name, - signal.target_trade_date, - signal.strategy, - signal.category.value, - signal.close, - Jsonb(dict(signal.details)), - ), - ) - except psycopg.Error as exc: + _executemany(connection, _ITEM_UPSERT, item_values) + if signal_values: + _executemany(connection, _SIGNAL_UPSERT, signal_values) + except SelectionRunError: + raise + except Exception as exc: # noqa: BLE001 - redact driver/pool details + code_context = items[0].ts_code if len(items) == 1 else f"{len(items)} items" raise SelectionRunStoreError( - f"failed to persist selection item {item.ts_code}" + f"failed to persist selection item {code_context}" ) from exc def finish_run( @@ -398,12 +438,35 @@ class PostgresSelectionRunRepository(SelectionRunStore): """Translate psycopg failures without exposing driver details.""" try: - with psycopg.connect(self.database_url) as connection: - yield connection - except psycopg.Error as exc: + if self.pool is None: + with psycopg.connect(self.database_url) as connection: + yield connection + else: + with self.pool.connection() as connection: + yield connection + except SelectionRunError: + raise + except Exception as exc: # noqa: BLE001 - normalize driver/pool errors raise SelectionRunStoreError("selection database operation failed") from exc +def _executemany(connection: Any, query: str, parameters: Sequence[tuple[object, ...]]) -> None: + """Use psycopg's batch API while retaining a minimal fake connection seam.""" + + executemany = getattr(connection, "executemany", None) + if callable(executemany): + executemany(query, parameters) + return + cursor_factory = getattr(connection, "cursor", None) + if callable(cursor_factory): + cursor_context = cast(Any, cursor_factory()) + with cursor_context as cursor: + cursor.executemany(query, parameters) + return + for values in parameters: + connection.execute(query, values) + + def _signal_from_row(row: tuple[object, ...]) -> SelectionSignal: """Map a persisted signal row back to the domain signal model.""" diff --git a/zhixing-server/src/zhixing_server/modules/selection/presentation/http.py b/zhixing-server/src/zhixing_server/modules/selection/presentation/http.py index 3356164..bed0c8d 100644 --- a/zhixing-server/src/zhixing_server/modules/selection/presentation/http.py +++ b/zhixing-server/src/zhixing_server/modules/selection/presentation/http.py @@ -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", ] diff --git a/zhixing-server/tests/test_selection_http.py b/zhixing-server/tests/test_selection_http.py index b8ff818..c113776 100644 --- a/zhixing-server/tests/test_selection_http.py +++ b/zhixing-server/tests/test_selection_http.py @@ -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() diff --git a/zhixing-server/tests/unit/selection/test_postgres_pool.py b/zhixing-server/tests/unit/selection/test_postgres_pool.py new file mode 100644 index 0000000..ca77181 --- /dev/null +++ b/zhixing-server/tests/unit/selection/test_postgres_pool.py @@ -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 diff --git a/zhixing-server/tests/unit/selection/test_postgres_reader.py b/zhixing-server/tests/unit/selection/test_postgres_reader.py index ee7d509..7f5d748 100644 --- a/zhixing-server/tests/unit/selection/test_postgres_reader.py +++ b/zhixing-server/tests/unit/selection/test_postgres_reader.py @@ -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 diff --git a/zhixing-server/tests/unit/selection/test_postgres_runs.py b/zhixing-server/tests/unit/selection/test_postgres_runs.py index 1029c8d..b862b74 100644 --- a/zhixing-server/tests/unit/selection/test_postgres_runs.py +++ b/zhixing-server/tests/unit/selection/test_postgres_runs.py @@ -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, ...]]] = [] diff --git a/zhixing-server/tests/unit/selection/test_run.py b/zhixing-server/tests/unit/selection/test_run.py index 2fe51d9..bb3ea0a 100644 --- a/zhixing-server/tests/unit/selection/test_run.py +++ b/zhixing-server/tests/unit/selection/test_run.py @@ -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