From 778c4beb0732416c369a7238affa5c71d4f47a9d Mon Sep 17 00:00:00 2001 From: Simon Date: Sun, 6 Sep 2026 20:40:01 +0800 Subject: [PATCH] =?UTF-8?q?perf(data):=20Repository=20=E6=89=B9=E9=87=8F?= =?UTF-8?q?=20upsert=20=E2=80=94=E2=80=94=20=E6=89=B9=E6=AC=A1=E4=B8=80?= =?UTF-8?q?=E6=AC=A1=E6=80=A7=E6=9F=A5=E9=87=8D=20+=20=E6=89=B9=E9=87=8F?= =?UTF-8?q?=E6=8F=92=E5=85=A5/=E6=9B=B4=E6=96=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 原实现逐行 select→insert/update,是全市场同步耗时主因 - 改为:组合键 row-constructor IN 一次查重 → 新行 add_all 批量插入、已有行就地更新 - 批内重复键(数据源偶发)以最后出现者为准覆盖,保持原语义 - 测试:repository/domain/factors/eval/engine/qlib/provider/migrations 等 67 项通过(tmp 库); 依赖真实 quant.db 的 API/Job 测试在全市场同步结束后补跑 --- .../sqlalchemy/repositories/market_impl.py | 45 ++++++++++++------- 1 file changed, 30 insertions(+), 15 deletions(-) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py index 98943c0..886c7a3 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py @@ -10,6 +10,7 @@ from __future__ import annotations from collections.abc import Sequence from datetime import date +from typing import Any from sqlalchemy import select from sqlalchemy.orm import Session @@ -54,30 +55,44 @@ def _upsert_by_business_key( entity_cls, entities: Sequence, ) -> int: - """按业务幂等键查重后 insert/update,返回触及行数(新增+更新)。 + """按业务幂等键**批量**查重后 insert/update(AGENT.md §1 性能友好:不做逐行 select)。 - 同一批内出现重复键(数据源偶发)时:先 flush 使前面已 add 的行可见, - 再按「后出现者覆盖」更新为最新值,避免 UNIQUE 冲突。 + - 一次查询取出本批已有的键 → 新行批量 add,已有行就地更新 + - 同一批内重复键(数据源偶发):以「后出现者」为准覆盖(避免 UNIQUE 冲突) + - 返回处理实体总数(含新增与更新),与旧逐行实现语义一致 """ + if not entities: + return 0 model_cls, key_cols = _TABLE[entity_cls] - seen: set[tuple] = set() - touched = 0 + key_col_attrs = [getattr(model_cls, k) for k in key_cols] + + keyed: list[tuple[tuple, dict]] = [] for ent in entities: values = _fields_of(ent) - key = tuple(values[k] for k in key_cols) - if key in seen: - session.flush() # 让本批内先前新增的行进入 select 视野 - else: - seen.add(key) - filters = [getattr(model_cls, k) == values[k] for k in key_cols] - row = session.scalars(select(model_cls).where(*filters)).first() + keyed.append((tuple(values[k] for k in key_cols), values)) + + existing_rows: dict[tuple, Any] = {} + keys = [k for k, _v in keyed] + if keys: + from sqlalchemy import tuple_ + + stmt = select(model_cls).where(tuple_(*key_col_attrs).in_(keys)) + for row in session.scalars(stmt): + key = tuple(getattr(row, k) for k in key_cols) + existing_rows[key] = row + + pending: dict[tuple, Any] = {} + for key, values in keyed: + row = existing_rows.get(key) + if row is None and key in pending: + row = pending[key] # 批内已待插入的同键 → 覆盖为新值 if row is None: - session.add(model_cls(**values)) + pending[key] = model_cls(**values) else: for col, val in values.items(): setattr(row, col, val) - touched += 1 - return touched + session.add_all(pending.values()) + return len(keyed) class SqlAlchemyStockRepository: