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 collections.abc import Sequence
|
||||||
from datetime import date
|
from datetime import date
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -54,30 +55,44 @@ def _upsert_by_business_key(
|
|||||||
entity_cls,
|
entity_cls,
|
||||||
entities: Sequence,
|
entities: Sequence,
|
||||||
) -> int:
|
) -> int:
|
||||||
"""按业务幂等键查重后 insert/update,返回触及行数(新增+更新)。
|
"""按业务幂等键**批量**查重后 insert/update(AGENT.md §1 性能友好:不做逐行 select)。
|
||||||
|
|
||||||
同一批内出现重复键(数据源偶发)时:先 flush 使前面已 add 的行可见,
|
- 一次查询取出本批已有的键 → 新行批量 add,已有行就地更新
|
||||||
再按「后出现者覆盖」更新为最新值,避免 UNIQUE 冲突。
|
- 同一批内重复键(数据源偶发):以「后出现者」为准覆盖(避免 UNIQUE 冲突)
|
||||||
|
- 返回处理实体总数(含新增与更新),与旧逐行实现语义一致
|
||||||
"""
|
"""
|
||||||
|
if not entities:
|
||||||
|
return 0
|
||||||
model_cls, key_cols = _TABLE[entity_cls]
|
model_cls, key_cols = _TABLE[entity_cls]
|
||||||
seen: set[tuple] = set()
|
key_col_attrs = [getattr(model_cls, k) for k in key_cols]
|
||||||
touched = 0
|
|
||||||
|
keyed: list[tuple[tuple, dict]] = []
|
||||||
for ent in entities:
|
for ent in entities:
|
||||||
values = _fields_of(ent)
|
values = _fields_of(ent)
|
||||||
key = tuple(values[k] for k in key_cols)
|
keyed.append((tuple(values[k] for k in key_cols), values))
|
||||||
if key in seen:
|
|
||||||
session.flush() # 让本批内先前新增的行进入 select 视野
|
existing_rows: dict[tuple, Any] = {}
|
||||||
else:
|
keys = [k for k, _v in keyed]
|
||||||
seen.add(key)
|
if keys:
|
||||||
filters = [getattr(model_cls, k) == values[k] for k in key_cols]
|
from sqlalchemy import tuple_
|
||||||
row = session.scalars(select(model_cls).where(*filters)).first()
|
|
||||||
|
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:
|
if row is None:
|
||||||
session.add(model_cls(**values))
|
pending[key] = model_cls(**values)
|
||||||
else:
|
else:
|
||||||
for col, val in values.items():
|
for col, val in values.items():
|
||||||
setattr(row, col, val)
|
setattr(row, col, val)
|
||||||
touched += 1
|
session.add_all(pending.values())
|
||||||
return touched
|
return len(keyed)
|
||||||
|
|
||||||
|
|
||||||
class SqlAlchemyStockRepository:
|
class SqlAlchemyStockRepository:
|
||||||
|
|||||||
Reference in New Issue
Block a user