demand_belong_pool_rel_repo.py 4.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115
  1. from __future__ import annotations
  2. from collections.abc import Iterable
  3. from sqlalchemy import delete, select, tuple_
  4. from sqlalchemy.dialects.mysql import insert
  5. from supply_infra.db.models.demand_belong_pool_rel import DemandBelongPoolRel
  6. from supply_infra.db.repositories.base import BaseRepository
  7. _BATCH_SIZE = 1000
  8. RelPair = tuple[int, int]
  9. class DemandBelongPoolRelRepository(BaseRepository[DemandBelongPoolRel]):
  10. """需求词 ↔ 需求池匹配边 — 按 (belong_id, pool_id) 增量写入。"""
  11. model = DemandBelongPoolRel
  12. def get_belong_ids_by_pool_ids(self, pool_ids: Iterable[int]) -> dict[int, list[int]]:
  13. """反查池表行归属的需求归属分类 id:pool_id -> [belong_id, ...]。"""
  14. id_list = [int(p) for p in pool_ids]
  15. if not id_list:
  16. return {}
  17. result: dict[int, list[int]] = {}
  18. for i in range(0, len(id_list), _BATCH_SIZE):
  19. batch = id_list[i : i + _BATCH_SIZE]
  20. stmt = select(
  21. DemandBelongPoolRel.multi_demand_pool_di_id,
  22. DemandBelongPoolRel.demand_belong_category_id,
  23. ).where(DemandBelongPoolRel.multi_demand_pool_di_id.in_(batch))
  24. for pool_id, belong_id in self.session.execute(stmt).all():
  25. result.setdefault(int(pool_id), []).append(int(belong_id))
  26. return result
  27. def get_pool_ids_by_belong_ids(
  28. self,
  29. belong_ids: Iterable[int],
  30. *,
  31. biz_dt: str | None = None,
  32. ) -> dict[int, list[int]]:
  33. """正向查询归属词关联的需求池行:belong_id -> [pool_id, ...]。"""
  34. id_list = [int(belong_id) for belong_id in belong_ids]
  35. if not id_list:
  36. return {}
  37. result: dict[int, list[int]] = {}
  38. for i in range(0, len(id_list), _BATCH_SIZE):
  39. batch = id_list[i : i + _BATCH_SIZE]
  40. stmt = select(
  41. DemandBelongPoolRel.demand_belong_category_id,
  42. DemandBelongPoolRel.multi_demand_pool_di_id,
  43. ).where(
  44. DemandBelongPoolRel.demand_belong_category_id.in_(batch),
  45. DemandBelongPoolRel.status == "active",
  46. )
  47. if biz_dt is not None:
  48. stmt = stmt.where(DemandBelongPoolRel.biz_dt == biz_dt)
  49. for belong_id, pool_id in self.session.execute(stmt).all():
  50. result.setdefault(int(belong_id), []).append(int(pool_id))
  51. return result
  52. def get_existing_pairs(self, pairs: Iterable[RelPair]) -> set[RelPair]:
  53. """返回 pairs 中已存在的 (belong_id, pool_id)。"""
  54. pair_list = [(int(b), int(p)) for b, p in pairs]
  55. if not pair_list:
  56. return set()
  57. existing: set[RelPair] = set()
  58. for i in range(0, len(pair_list), _BATCH_SIZE):
  59. batch = pair_list[i : i + _BATCH_SIZE]
  60. stmt = select(
  61. DemandBelongPoolRel.demand_belong_category_id,
  62. DemandBelongPoolRel.multi_demand_pool_di_id,
  63. ).where(
  64. tuple_(
  65. DemandBelongPoolRel.demand_belong_category_id,
  66. DemandBelongPoolRel.multi_demand_pool_di_id,
  67. ).in_(batch)
  68. )
  69. existing.update(
  70. (int(b), int(p)) for b, p in self.session.execute(stmt).all()
  71. )
  72. return existing
  73. def delete_by_pool_ids(self, pool_ids: Iterable[int]) -> int:
  74. """删除指定需求池行的关系,用于源行修订/删除后的精确重建。"""
  75. id_list = sorted({int(pool_id) for pool_id in pool_ids})
  76. if not id_list:
  77. return 0
  78. deleted = 0
  79. for i in range(0, len(id_list), _BATCH_SIZE):
  80. batch = id_list[i : i + _BATCH_SIZE]
  81. stmt = delete(DemandBelongPoolRel).where(
  82. DemandBelongPoolRel.multi_demand_pool_di_id.in_(batch)
  83. )
  84. result = self.session.execute(stmt)
  85. deleted += result.rowcount or 0
  86. return deleted
  87. def bulk_insert_ignore(self, rows: list[dict]) -> int:
  88. """批量插入,忽略已存在的唯一键。"""
  89. if not rows:
  90. return 0
  91. inserted = 0
  92. for i in range(0, len(rows), _BATCH_SIZE):
  93. batch = rows[i : i + _BATCH_SIZE]
  94. stmt = insert(DemandBelongPoolRel).values(batch).prefix_with("IGNORE")
  95. result = self.session.execute(stmt)
  96. inserted += result.rowcount
  97. return inserted