demand_grade_repo.py 3.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103
  1. from __future__ import annotations
  2. from collections.abc import Iterable
  3. from sqlalchemy import func, select
  4. from sqlalchemy.dialects.mysql import insert
  5. from supply_infra.db.models.demand_grade import DemandGrade
  6. from supply_infra.db.repositories.base import BaseRepository
  7. _BATCH_SIZE = 500
  8. _UPSERT_COLUMNS = (
  9. "category_ids",
  10. "grade",
  11. "score",
  12. "prior_total_score",
  13. "posterior_rov_avg",
  14. "posterior_rov_count",
  15. "has_posterior",
  16. "related_pool_ids",
  17. "video_list",
  18. "strategies",
  19. "reason",
  20. )
  21. class DemandGradeRepository(BaseRepository[DemandGrade]):
  22. """需求分级结果表 — 按 (biz_dt, demand_name) 增量/更新写入。"""
  23. model = DemandGrade
  24. def get_existing_demand_names(self, biz_dt: str, names: Iterable[str] | None = None) -> set[str]:
  25. """返回指定业务日已分级的需求名集合;传入 names 时只在其中查交集。"""
  26. stmt = select(DemandGrade.demand_name).where(DemandGrade.biz_dt == biz_dt)
  27. if names is not None:
  28. name_list = [n for n in names if n]
  29. if not name_list:
  30. return set()
  31. stmt = stmt.where(DemandGrade.demand_name.in_(name_list))
  32. return {n for n in self.session.scalars(stmt).all() if n}
  33. def count_by_biz_dt(self, biz_dt: str) -> int:
  34. """统计指定业务日已分级的需求数。"""
  35. stmt = select(func.count()).select_from(DemandGrade).where(DemandGrade.biz_dt == biz_dt)
  36. return int(self.session.scalar(stmt) or 0)
  37. def list_by_biz_dt(self, biz_dt: str) -> list[DemandGrade]:
  38. """返回指定业务日的全部分级结果,按等级、需求名排序。"""
  39. stmt = (
  40. select(DemandGrade)
  41. .where(DemandGrade.biz_dt == biz_dt)
  42. .order_by(DemandGrade.grade, DemandGrade.demand_name)
  43. )
  44. return list(self.session.scalars(stmt).all())
  45. def list_by_biz_dt_and_grades(
  46. self,
  47. biz_dt: str,
  48. grades: Iterable[str] = ("S", "A"),
  49. ) -> list[DemandGrade]:
  50. """返回指定业务日、指定等级的分级结果,按等级、需求名排序。"""
  51. grade_list = [g for g in grades if g]
  52. if not grade_list:
  53. return []
  54. stmt = (
  55. select(DemandGrade)
  56. .where(DemandGrade.biz_dt == biz_dt, DemandGrade.grade.in_(grade_list))
  57. .order_by(DemandGrade.grade, DemandGrade.demand_name)
  58. )
  59. return list(self.session.scalars(stmt).all())
  60. def get_latest_biz_dt(self) -> str | None:
  61. """返回 demand_grade 中最新的业务日期。"""
  62. stmt = select(func.max(DemandGrade.biz_dt))
  63. return self.session.scalar(stmt)
  64. def get_ids_by_names(self, biz_dt: str, names: Iterable[str]) -> dict[str, int]:
  65. """按 (biz_dt, demand_name) 反查 id,供 upsert 后写关联表使用。"""
  66. name_list = [n for n in names if n]
  67. if not name_list:
  68. return {}
  69. stmt = select(DemandGrade.demand_name, DemandGrade.id).where(
  70. DemandGrade.biz_dt == biz_dt,
  71. DemandGrade.demand_name.in_(name_list),
  72. )
  73. return {name: int(id_) for name, id_ in self.session.execute(stmt).all()}
  74. def bulk_upsert(self, rows: list[dict]) -> int:
  75. """按 (biz_dt, demand_name) 批量 upsert。"""
  76. if not rows:
  77. return 0
  78. affected = 0
  79. for i in range(0, len(rows), _BATCH_SIZE):
  80. batch = rows[i : i + _BATCH_SIZE]
  81. stmt = insert(DemandGrade).values(batch)
  82. stmt = stmt.on_duplicate_key_update(
  83. **{col: stmt.inserted[col] for col in _UPSERT_COLUMNS}
  84. )
  85. result = self.session.execute(stmt)
  86. affected += result.rowcount or 0
  87. return affected