| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115 |
- from __future__ import annotations
- from collections.abc import Iterable
- from sqlalchemy import delete, select, tuple_
- from sqlalchemy.dialects.mysql import insert
- from supply_infra.db.models.demand_belong_pool_rel import DemandBelongPoolRel
- from supply_infra.db.repositories.base import BaseRepository
- _BATCH_SIZE = 1000
- RelPair = tuple[int, int]
- class DemandBelongPoolRelRepository(BaseRepository[DemandBelongPoolRel]):
- """需求词 ↔ 需求池匹配边 — 按 (belong_id, pool_id) 增量写入。"""
- model = DemandBelongPoolRel
- def get_belong_ids_by_pool_ids(self, pool_ids: Iterable[int]) -> dict[int, list[int]]:
- """反查池表行归属的需求归属分类 id:pool_id -> [belong_id, ...]。"""
- id_list = [int(p) for p in pool_ids]
- if not id_list:
- return {}
- result: dict[int, list[int]] = {}
- for i in range(0, len(id_list), _BATCH_SIZE):
- batch = id_list[i : i + _BATCH_SIZE]
- stmt = select(
- DemandBelongPoolRel.multi_demand_pool_di_id,
- DemandBelongPoolRel.demand_belong_category_id,
- ).where(DemandBelongPoolRel.multi_demand_pool_di_id.in_(batch))
- for pool_id, belong_id in self.session.execute(stmt).all():
- result.setdefault(int(pool_id), []).append(int(belong_id))
- return result
- def get_pool_ids_by_belong_ids(
- self,
- belong_ids: Iterable[int],
- *,
- biz_dt: str | None = None,
- ) -> dict[int, list[int]]:
- """正向查询归属词关联的需求池行:belong_id -> [pool_id, ...]。"""
- id_list = [int(belong_id) for belong_id in belong_ids]
- if not id_list:
- return {}
- result: dict[int, list[int]] = {}
- for i in range(0, len(id_list), _BATCH_SIZE):
- batch = id_list[i : i + _BATCH_SIZE]
- stmt = select(
- DemandBelongPoolRel.demand_belong_category_id,
- DemandBelongPoolRel.multi_demand_pool_di_id,
- ).where(
- DemandBelongPoolRel.demand_belong_category_id.in_(batch),
- DemandBelongPoolRel.status == "active",
- )
- if biz_dt is not None:
- stmt = stmt.where(DemandBelongPoolRel.biz_dt == biz_dt)
- for belong_id, pool_id in self.session.execute(stmt).all():
- result.setdefault(int(belong_id), []).append(int(pool_id))
- return result
- def get_existing_pairs(self, pairs: Iterable[RelPair]) -> set[RelPair]:
- """返回 pairs 中已存在的 (belong_id, pool_id)。"""
- pair_list = [(int(b), int(p)) for b, p in pairs]
- if not pair_list:
- return set()
- existing: set[RelPair] = set()
- for i in range(0, len(pair_list), _BATCH_SIZE):
- batch = pair_list[i : i + _BATCH_SIZE]
- stmt = select(
- DemandBelongPoolRel.demand_belong_category_id,
- DemandBelongPoolRel.multi_demand_pool_di_id,
- ).where(
- tuple_(
- DemandBelongPoolRel.demand_belong_category_id,
- DemandBelongPoolRel.multi_demand_pool_di_id,
- ).in_(batch)
- )
- existing.update(
- (int(b), int(p)) for b, p in self.session.execute(stmt).all()
- )
- return existing
- def delete_by_pool_ids(self, pool_ids: Iterable[int]) -> int:
- """删除指定需求池行的关系,用于源行修订/删除后的精确重建。"""
- id_list = sorted({int(pool_id) for pool_id in pool_ids})
- if not id_list:
- return 0
- deleted = 0
- for i in range(0, len(id_list), _BATCH_SIZE):
- batch = id_list[i : i + _BATCH_SIZE]
- stmt = delete(DemandBelongPoolRel).where(
- DemandBelongPoolRel.multi_demand_pool_di_id.in_(batch)
- )
- result = self.session.execute(stmt)
- deleted += result.rowcount or 0
- return deleted
- def bulk_insert_ignore(self, rows: list[dict]) -> int:
- """批量插入,忽略已存在的唯一键。"""
- if not rows:
- return 0
- inserted = 0
- for i in range(0, len(rows), _BATCH_SIZE):
- batch = rows[i : i + _BATCH_SIZE]
- stmt = insert(DemandBelongPoolRel).values(batch).prefix_with("IGNORE")
- result = self.session.execute(stmt)
- inserted += result.rowcount
- return inserted
|