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: