demand_grade_plan_repo.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350
  1. from __future__ import annotations
  2. import json
  3. import uuid
  4. from typing import Any
  5. from sqlalchemy import func, select, update
  6. from supply_infra.db.models.demand_grade_plan import (
  7. DemandGradePlan,
  8. DemandGradePlanGroup,
  9. DemandGradePlanGroupItem,
  10. )
  11. from supply_infra.db.repositories.base import BaseRepository
  12. from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
  13. from supply_infra.pipeline.dates import china_now
  14. class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
  15. model = DemandGradePlan
  16. def get_latest_plan(self, biz_dt: str) -> DemandGradePlan | None:
  17. return self.session.scalar(
  18. select(DemandGradePlan)
  19. .where(DemandGradePlan.biz_dt == biz_dt)
  20. .order_by(DemandGradePlan.create_time.desc())
  21. .limit(1)
  22. )
  23. def list_groups_by_biz_dt(self, biz_dt: str) -> list[DemandGradePlanGroup]:
  24. return list(self.session.scalars(
  25. select(DemandGradePlanGroup)
  26. .where(DemandGradePlanGroup.biz_dt == biz_dt)
  27. .order_by(DemandGradePlanGroup.id)
  28. ).all())
  29. def get_assigned_category_ids(self, biz_dt: str) -> set[int]:
  30. assigned: set[int] = set()
  31. for group in self.list_groups_by_biz_dt(biz_dt):
  32. for category_id in json.loads(group.category_ids):
  33. try:
  34. assigned.add(int(category_id))
  35. except (TypeError, ValueError):
  36. continue
  37. return assigned
  38. def get_execution_snapshot(self, biz_dt: str) -> dict[str, Any]:
  39. """在 session 内汇总当天计划分组执行情况,避免 ORM 脱离会话。"""
  40. groups = self.list_groups_by_biz_dt(biz_dt)
  41. group_status = self.summarize(biz_dt)
  42. assigned_category_ids = sorted(self.get_assigned_category_ids(biz_dt))
  43. claimable_groups = group_status.get("pending", 0)
  44. unfinished_groups = claimable_groups + group_status.get("running", 0)
  45. return {
  46. "planned_groups": len(groups),
  47. "assigned_category_ids": assigned_category_ids,
  48. "group_status": group_status,
  49. "claimable_groups": claimable_groups,
  50. "unfinished_groups": unfinished_groups,
  51. "execution_complete": unfinished_groups == 0,
  52. }
  53. def create_plan(self, biz_dt: str, payload: dict[str, Any]) -> None:
  54. plan_id = str(uuid.uuid4())
  55. groups = payload.get("groups") or []
  56. coverage_complete = bool(payload.get("coverage_complete"))
  57. # 分组明细已写入 demand_grade_plan_group,计划表仅保存可检索的紧凑摘要,避免 TEXT 溢出。
  58. summary = {
  59. "biz_dt": biz_dt,
  60. "grouping_strategy": payload.get("grouping_strategy"),
  61. "total_hanging_nodes": payload.get("total_hanging_nodes"),
  62. "group_count": len(groups),
  63. "heat_level_definition": payload.get("heat_level_definition", {}),
  64. "batch_heat_level_counts": payload.get("batch_heat_level_counts", {}),
  65. "coverage_complete": coverage_complete,
  66. "covered_category_ids": payload.get("covered_category_ids", []),
  67. "uncovered_category_ids": payload.get("uncovered_category_ids", []),
  68. }
  69. self.add(DemandGradePlan(
  70. plan_id=plan_id, biz_dt=biz_dt, status="planned",
  71. total_hanging_nodes=int(payload.get("total_hanging_nodes") or 0),
  72. group_count=len(groups), coverage_complete=1 if coverage_complete else 0,
  73. plan_json=json.dumps(summary, ensure_ascii=False),
  74. ))
  75. for group_no, group in enumerate(groups, start=1):
  76. shared_traits = json.dumps({
  77. "description": str(group["shared_traits"]),
  78. "batch_heat_level": group.get("batch_heat_level"),
  79. "batch_heat_label": group.get("batch_heat_label"),
  80. "batch_global_rank_score": group.get("batch_global_rank_score"),
  81. "batch_total_score_avg": group.get("batch_total_score_avg"),
  82. "category_global_positions": group.get("category_global_positions", []),
  83. }, ensure_ascii=False)
  84. self.add(DemandGradePlanGroup(
  85. plan_id=plan_id, biz_dt=biz_dt, group_no=group_no, group_key=str(group["group_id"]),
  86. category_ids=json.dumps(group["category_ids"], ensure_ascii=False),
  87. planning_reason=str(group["planning_reason"]), shared_traits=shared_traits,
  88. status="pending",
  89. ))
  90. group_row = self.session.scalar(
  91. select(DemandGradePlanGroup)
  92. .where(DemandGradePlanGroup.plan_id == plan_id, DemandGradePlanGroup.group_no == group_no)
  93. .limit(1)
  94. )
  95. if group_row is not None:
  96. graded_names = DemandGradeRepository(self.session).get_existing_demand_names(biz_dt)
  97. self.materialize_group_items(
  98. int(group_row.id),
  99. biz_dt,
  100. list(group.get("category_ids") or []),
  101. graded_names=graded_names,
  102. )
  103. def list_pending_group_ids(self, biz_dt: str) -> list[int]:
  104. """返回当天待执行的 pending 任务 id。"""
  105. stmt = (
  106. select(DemandGradePlanGroup.id)
  107. .where(
  108. DemandGradePlanGroup.biz_dt == biz_dt,
  109. DemandGradePlanGroup.status == "pending",
  110. )
  111. .order_by(DemandGradePlanGroup.id)
  112. )
  113. return [int(group_id) for group_id in self.session.scalars(stmt).all()]
  114. def claim_group(self, biz_dt: str, group_id: int) -> dict[str, Any] | None:
  115. """按固定 ID 原子领取任务;状态已变化时返回 None。"""
  116. stmt = (
  117. select(DemandGradePlanGroup)
  118. .where(
  119. DemandGradePlanGroup.id == int(group_id),
  120. DemandGradePlanGroup.biz_dt == biz_dt,
  121. DemandGradePlanGroup.status == "pending",
  122. )
  123. .with_for_update(skip_locked=True)
  124. )
  125. group = self.session.scalar(stmt)
  126. if group is None:
  127. return None
  128. return self._mark_claimed(group)
  129. @staticmethod
  130. def _mark_claimed(group: DemandGradePlanGroup) -> dict[str, Any]:
  131. group.status = "running"
  132. group.attempts += 1
  133. group.started_at = china_now()
  134. group.error_message = None
  135. group.finished_at = None
  136. return {
  137. "id": int(group.id),
  138. "group_key": group.group_key,
  139. "category_ids": json.loads(group.category_ids),
  140. }
  141. def finish_group(self, group_id: int, *, success: bool, error_message: str | None = None) -> None:
  142. self.session.execute(
  143. update(DemandGradePlanGroup)
  144. .where(DemandGradePlanGroup.id == group_id)
  145. .values(
  146. status="finished" if success else "failed",
  147. error_message=error_message,
  148. finished_at=china_now(),
  149. )
  150. )
  151. def summarize(self, biz_dt: str) -> dict[str, int]:
  152. rows = self.session.execute(
  153. select(DemandGradePlanGroup.status).where(DemandGradePlanGroup.biz_dt == biz_dt)
  154. ).scalars().all()
  155. return {status: sum(value == status for value in rows) for status in ("pending", "running", "finished", "failed")}
  156. def group_item_count(self, group_id: int) -> int:
  157. return int(self.session.scalar(
  158. select(func.count())
  159. .select_from(DemandGradePlanGroupItem)
  160. .where(DemandGradePlanGroupItem.group_id == int(group_id))
  161. ) or 0)
  162. def materialize_group_items(
  163. self,
  164. group_id: int,
  165. biz_dt: str,
  166. category_ids: list[int],
  167. *,
  168. graded_names: set[str] | None = None,
  169. ) -> int:
  170. """将 category_ids 展开为组内需求明细;已存在明细时跳过。"""
  171. if self.group_item_count(group_id) > 0:
  172. return 0
  173. graded = graded_names or set()
  174. from supply_infra.scheduler.plan_group_batch import resolve_demands_for_category_ids
  175. demands = resolve_demands_for_category_ids(
  176. biz_dt,
  177. category_ids,
  178. session=self.session,
  179. )
  180. created = 0
  181. for sort_order, demand in enumerate(demands, start=1):
  182. demand_name = str(demand["demand_name"])
  183. status = "skipped" if demand_name in graded else "pending"
  184. self.add(DemandGradePlanGroupItem(
  185. group_id=int(group_id),
  186. biz_dt=biz_dt,
  187. pool_id=int(demand["pool_id"]),
  188. demand_name=demand_name,
  189. sort_order=sort_order,
  190. status=status,
  191. ))
  192. created += 1
  193. return created
  194. def materialize_pending_groups(self, biz_dt: str, *, graded_names: set[str] | None = None) -> int:
  195. """为当天尚未物化明细的 pending 计划组补写需求列表。"""
  196. created = 0
  197. for group in self.list_groups_by_biz_dt(biz_dt):
  198. if group.status != "pending":
  199. continue
  200. if self.group_item_count(int(group.id)) > 0:
  201. continue
  202. category_ids = json.loads(group.category_ids)
  203. created += self.materialize_group_items(
  204. int(group.id),
  205. biz_dt,
  206. category_ids,
  207. graded_names=graded_names,
  208. )
  209. return created
  210. def list_pending_group_items(
  211. self,
  212. group_id: int,
  213. *,
  214. limit: int | None = None,
  215. ) -> list[dict[str, Any]]:
  216. stmt = (
  217. select(DemandGradePlanGroupItem)
  218. .where(
  219. DemandGradePlanGroupItem.group_id == int(group_id),
  220. DemandGradePlanGroupItem.status == "pending",
  221. )
  222. .order_by(DemandGradePlanGroupItem.sort_order, DemandGradePlanGroupItem.id)
  223. )
  224. if limit is not None:
  225. stmt = stmt.limit(max(1, int(limit)))
  226. rows = self.session.scalars(stmt).all()
  227. return [
  228. {
  229. "item_id": int(row.id),
  230. "pool_id": int(row.pool_id),
  231. "demand_name": str(row.demand_name),
  232. }
  233. for row in rows
  234. ]
  235. def summarize_group_items(self, group_id: int) -> dict[str, int]:
  236. rows = self.session.execute(
  237. select(DemandGradePlanGroupItem.status)
  238. .where(DemandGradePlanGroupItem.group_id == int(group_id))
  239. ).scalars().all()
  240. statuses = ("pending", "finished", "failed", "skipped")
  241. return {status: sum(value == status for value in rows) for status in statuses}
  242. def mark_group_items_status(
  243. self,
  244. item_ids: list[int],
  245. *,
  246. status: str,
  247. error_message: str | None = None,
  248. ) -> None:
  249. if not item_ids:
  250. return
  251. values: dict[str, Any] = {"status": status, "error_message": error_message}
  252. if status in {"finished", "failed", "skipped"}:
  253. values["finished_at"] = china_now()
  254. self.session.execute(
  255. update(DemandGradePlanGroupItem)
  256. .where(DemandGradePlanGroupItem.id.in_([int(item_id) for item_id in item_ids]))
  257. .values(**values)
  258. )
  259. def list_failed_group_items(
  260. self,
  261. biz_dt: str | None = None,
  262. *,
  263. group_ids: list[int] | None = None,
  264. ) -> list[dict[str, Any]]:
  265. """返回 status=failed 的组内需求明细。"""
  266. stmt = select(DemandGradePlanGroupItem).where(DemandGradePlanGroupItem.status == "failed")
  267. if biz_dt:
  268. stmt = stmt.where(DemandGradePlanGroupItem.biz_dt == biz_dt)
  269. if group_ids:
  270. stmt = stmt.where(
  271. DemandGradePlanGroupItem.group_id.in_([int(group_id) for group_id in group_ids])
  272. )
  273. stmt = stmt.order_by(
  274. DemandGradePlanGroupItem.biz_dt,
  275. DemandGradePlanGroupItem.group_id,
  276. DemandGradePlanGroupItem.sort_order,
  277. DemandGradePlanGroupItem.id,
  278. )
  279. rows = self.session.scalars(stmt).all()
  280. return [
  281. {
  282. "item_id": int(row.id),
  283. "group_id": int(row.group_id),
  284. "biz_dt": str(row.biz_dt),
  285. "pool_id": int(row.pool_id),
  286. "demand_name": str(row.demand_name),
  287. "error_message": row.error_message,
  288. }
  289. for row in rows
  290. ]
  291. def reset_failed_items_to_pending(
  292. self,
  293. biz_dt: str | None = None,
  294. *,
  295. group_ids: list[int] | None = None,
  296. ) -> dict[str, Any]:
  297. """将 failed 明细重置为 pending,并将所属计划组重置为 pending 以便重新领取。"""
  298. failed_rows = self.list_failed_group_items(biz_dt, group_ids=group_ids)
  299. if not failed_rows:
  300. return {"reset_items": 0, "reset_groups": 0, "group_ids": [], "items": []}
  301. item_ids = [int(row["item_id"]) for row in failed_rows]
  302. affected_group_ids = sorted({int(row["group_id"]) for row in failed_rows})
  303. self.session.execute(
  304. update(DemandGradePlanGroupItem)
  305. .where(DemandGradePlanGroupItem.id.in_(item_ids))
  306. .values(status="pending", error_message=None, finished_at=None)
  307. )
  308. self.session.execute(
  309. update(DemandGradePlanGroup)
  310. .where(DemandGradePlanGroup.id.in_(affected_group_ids))
  311. .values(
  312. status="pending",
  313. error_message=None,
  314. started_at=None,
  315. finished_at=None,
  316. )
  317. )
  318. return {
  319. "reset_items": len(item_ids),
  320. "reset_groups": len(affected_group_ids),
  321. "group_ids": affected_group_ids,
  322. "items": failed_rows,
  323. }