perf(data): Repository 批量 upsert —— 批次一次性查重 + 批量插入/更新
- 原实现逐行 select→insert/update,是全市场同步耗时主因 - 改为:组合键 row-constructor IN 一次查重 → 新行 add_all 批量插入、已有行就地更新 - 批内重复键(数据源偶发)以最后出现者为准覆盖,保持原语义 - 测试:repository/domain/factors/eval/engine/qlib/provider/migrations 等 67 项通过(tmp 库); 依赖真实 quant.db 的 API/Job 测试在全市场同步结束后补跑
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user