| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350 |
- from __future__ import annotations
- import json
- import uuid
- from typing import Any
- from sqlalchemy import func, select, update
- from supply_infra.db.models.demand_grade_plan import (
- DemandGradePlan,
- DemandGradePlanGroup,
- DemandGradePlanGroupItem,
- )
- from supply_infra.db.repositories.base import BaseRepository
- from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
- from supply_infra.pipeline.dates import china_now
- class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
- model = DemandGradePlan
- def get_latest_plan(self, biz_dt: str) -> DemandGradePlan | None:
- return self.session.scalar(
- select(DemandGradePlan)
- .where(DemandGradePlan.biz_dt == biz_dt)
- .order_by(DemandGradePlan.create_time.desc())
- .limit(1)
- )
- def list_groups_by_biz_dt(self, biz_dt: str) -> list[DemandGradePlanGroup]:
- return list(self.session.scalars(
- select(DemandGradePlanGroup)
- .where(DemandGradePlanGroup.biz_dt == biz_dt)
- .order_by(DemandGradePlanGroup.id)
- ).all())
- def get_assigned_category_ids(self, biz_dt: str) -> set[int]:
- assigned: set[int] = set()
- for group in self.list_groups_by_biz_dt(biz_dt):
- for category_id in json.loads(group.category_ids):
- try:
- assigned.add(int(category_id))
- except (TypeError, ValueError):
- continue
- return assigned
- def get_execution_snapshot(self, biz_dt: str) -> dict[str, Any]:
- """在 session 内汇总当天计划分组执行情况,避免 ORM 脱离会话。"""
- groups = self.list_groups_by_biz_dt(biz_dt)
- group_status = self.summarize(biz_dt)
- assigned_category_ids = sorted(self.get_assigned_category_ids(biz_dt))
- claimable_groups = group_status.get("pending", 0)
- unfinished_groups = claimable_groups + group_status.get("running", 0)
- return {
- "planned_groups": len(groups),
- "assigned_category_ids": assigned_category_ids,
- "group_status": group_status,
- "claimable_groups": claimable_groups,
- "unfinished_groups": unfinished_groups,
- "execution_complete": unfinished_groups == 0,
- }
- def create_plan(self, biz_dt: str, payload: dict[str, Any]) -> None:
- plan_id = str(uuid.uuid4())
- groups = payload.get("groups") or []
- coverage_complete = bool(payload.get("coverage_complete"))
- # 分组明细已写入 demand_grade_plan_group,计划表仅保存可检索的紧凑摘要,避免 TEXT 溢出。
- summary = {
- "biz_dt": biz_dt,
- "grouping_strategy": payload.get("grouping_strategy"),
- "total_hanging_nodes": payload.get("total_hanging_nodes"),
- "group_count": len(groups),
- "heat_level_definition": payload.get("heat_level_definition", {}),
- "batch_heat_level_counts": payload.get("batch_heat_level_counts", {}),
- "coverage_complete": coverage_complete,
- "covered_category_ids": payload.get("covered_category_ids", []),
- "uncovered_category_ids": payload.get("uncovered_category_ids", []),
- }
- self.add(DemandGradePlan(
- plan_id=plan_id, biz_dt=biz_dt, status="planned",
- total_hanging_nodes=int(payload.get("total_hanging_nodes") or 0),
- group_count=len(groups), coverage_complete=1 if coverage_complete else 0,
- plan_json=json.dumps(summary, ensure_ascii=False),
- ))
- for group_no, group in enumerate(groups, start=1):
- shared_traits = json.dumps({
- "description": str(group["shared_traits"]),
- "batch_heat_level": group.get("batch_heat_level"),
- "batch_heat_label": group.get("batch_heat_label"),
- "batch_global_rank_score": group.get("batch_global_rank_score"),
- "batch_total_score_avg": group.get("batch_total_score_avg"),
- "category_global_positions": group.get("category_global_positions", []),
- }, ensure_ascii=False)
- self.add(DemandGradePlanGroup(
- plan_id=plan_id, biz_dt=biz_dt, group_no=group_no, group_key=str(group["group_id"]),
- category_ids=json.dumps(group["category_ids"], ensure_ascii=False),
- planning_reason=str(group["planning_reason"]), shared_traits=shared_traits,
- status="pending",
- ))
- group_row = self.session.scalar(
- select(DemandGradePlanGroup)
- .where(DemandGradePlanGroup.plan_id == plan_id, DemandGradePlanGroup.group_no == group_no)
- .limit(1)
- )
- if group_row is not None:
- graded_names = DemandGradeRepository(self.session).get_existing_demand_names(biz_dt)
- self.materialize_group_items(
- int(group_row.id),
- biz_dt,
- list(group.get("category_ids") or []),
- graded_names=graded_names,
- )
- def list_pending_group_ids(self, biz_dt: str) -> list[int]:
- """返回当天待执行的 pending 任务 id。"""
- stmt = (
- select(DemandGradePlanGroup.id)
- .where(
- DemandGradePlanGroup.biz_dt == biz_dt,
- DemandGradePlanGroup.status == "pending",
- )
- .order_by(DemandGradePlanGroup.id)
- )
- return [int(group_id) for group_id in self.session.scalars(stmt).all()]
- def claim_group(self, biz_dt: str, group_id: int) -> dict[str, Any] | None:
- """按固定 ID 原子领取任务;状态已变化时返回 None。"""
- stmt = (
- select(DemandGradePlanGroup)
- .where(
- DemandGradePlanGroup.id == int(group_id),
- DemandGradePlanGroup.biz_dt == biz_dt,
- DemandGradePlanGroup.status == "pending",
- )
- .with_for_update(skip_locked=True)
- )
- group = self.session.scalar(stmt)
- if group is None:
- return None
- return self._mark_claimed(group)
- @staticmethod
- def _mark_claimed(group: DemandGradePlanGroup) -> dict[str, Any]:
- group.status = "running"
- group.attempts += 1
- group.started_at = china_now()
- group.error_message = None
- group.finished_at = None
- return {
- "id": int(group.id),
- "group_key": group.group_key,
- "category_ids": json.loads(group.category_ids),
- }
- def finish_group(self, group_id: int, *, success: bool, error_message: str | None = None) -> None:
- self.session.execute(
- update(DemandGradePlanGroup)
- .where(DemandGradePlanGroup.id == group_id)
- .values(
- status="finished" if success else "failed",
- error_message=error_message,
- finished_at=china_now(),
- )
- )
- def summarize(self, biz_dt: str) -> dict[str, int]:
- rows = self.session.execute(
- select(DemandGradePlanGroup.status).where(DemandGradePlanGroup.biz_dt == biz_dt)
- ).scalars().all()
- return {status: sum(value == status for value in rows) for status in ("pending", "running", "finished", "failed")}
- def group_item_count(self, group_id: int) -> int:
- return int(self.session.scalar(
- select(func.count())
- .select_from(DemandGradePlanGroupItem)
- .where(DemandGradePlanGroupItem.group_id == int(group_id))
- ) or 0)
- def materialize_group_items(
- self,
- group_id: int,
- biz_dt: str,
- category_ids: list[int],
- *,
- graded_names: set[str] | None = None,
- ) -> int:
- """将 category_ids 展开为组内需求明细;已存在明细时跳过。"""
- if self.group_item_count(group_id) > 0:
- return 0
- graded = graded_names or set()
- from supply_infra.scheduler.plan_group_batch import resolve_demands_for_category_ids
- demands = resolve_demands_for_category_ids(
- biz_dt,
- category_ids,
- session=self.session,
- )
- created = 0
- for sort_order, demand in enumerate(demands, start=1):
- demand_name = str(demand["demand_name"])
- status = "skipped" if demand_name in graded else "pending"
- self.add(DemandGradePlanGroupItem(
- group_id=int(group_id),
- biz_dt=biz_dt,
- pool_id=int(demand["pool_id"]),
- demand_name=demand_name,
- sort_order=sort_order,
- status=status,
- ))
- created += 1
- return created
- def materialize_pending_groups(self, biz_dt: str, *, graded_names: set[str] | None = None) -> int:
- """为当天尚未物化明细的 pending 计划组补写需求列表。"""
- created = 0
- for group in self.list_groups_by_biz_dt(biz_dt):
- if group.status != "pending":
- continue
- if self.group_item_count(int(group.id)) > 0:
- continue
- category_ids = json.loads(group.category_ids)
- created += self.materialize_group_items(
- int(group.id),
- biz_dt,
- category_ids,
- graded_names=graded_names,
- )
- return created
- def list_pending_group_items(
- self,
- group_id: int,
- *,
- limit: int | None = None,
- ) -> list[dict[str, Any]]:
- stmt = (
- select(DemandGradePlanGroupItem)
- .where(
- DemandGradePlanGroupItem.group_id == int(group_id),
- DemandGradePlanGroupItem.status == "pending",
- )
- .order_by(DemandGradePlanGroupItem.sort_order, DemandGradePlanGroupItem.id)
- )
- if limit is not None:
- stmt = stmt.limit(max(1, int(limit)))
- rows = self.session.scalars(stmt).all()
- return [
- {
- "item_id": int(row.id),
- "pool_id": int(row.pool_id),
- "demand_name": str(row.demand_name),
- }
- for row in rows
- ]
- def summarize_group_items(self, group_id: int) -> dict[str, int]:
- rows = self.session.execute(
- select(DemandGradePlanGroupItem.status)
- .where(DemandGradePlanGroupItem.group_id == int(group_id))
- ).scalars().all()
- statuses = ("pending", "finished", "failed", "skipped")
- return {status: sum(value == status for value in rows) for status in statuses}
- def mark_group_items_status(
- self,
- item_ids: list[int],
- *,
- status: str,
- error_message: str | None = None,
- ) -> None:
- if not item_ids:
- return
- values: dict[str, Any] = {"status": status, "error_message": error_message}
- if status in {"finished", "failed", "skipped"}:
- values["finished_at"] = china_now()
- self.session.execute(
- update(DemandGradePlanGroupItem)
- .where(DemandGradePlanGroupItem.id.in_([int(item_id) for item_id in item_ids]))
- .values(**values)
- )
- def list_failed_group_items(
- self,
- biz_dt: str | None = None,
- *,
- group_ids: list[int] | None = None,
- ) -> list[dict[str, Any]]:
- """返回 status=failed 的组内需求明细。"""
- stmt = select(DemandGradePlanGroupItem).where(DemandGradePlanGroupItem.status == "failed")
- if biz_dt:
- stmt = stmt.where(DemandGradePlanGroupItem.biz_dt == biz_dt)
- if group_ids:
- stmt = stmt.where(
- DemandGradePlanGroupItem.group_id.in_([int(group_id) for group_id in group_ids])
- )
- stmt = stmt.order_by(
- DemandGradePlanGroupItem.biz_dt,
- DemandGradePlanGroupItem.group_id,
- DemandGradePlanGroupItem.sort_order,
- DemandGradePlanGroupItem.id,
- )
- rows = self.session.scalars(stmt).all()
- return [
- {
- "item_id": int(row.id),
- "group_id": int(row.group_id),
- "biz_dt": str(row.biz_dt),
- "pool_id": int(row.pool_id),
- "demand_name": str(row.demand_name),
- "error_message": row.error_message,
- }
- for row in rows
- ]
- def reset_failed_items_to_pending(
- self,
- biz_dt: str | None = None,
- *,
- group_ids: list[int] | None = None,
- ) -> dict[str, Any]:
- """将 failed 明细重置为 pending,并将所属计划组重置为 pending 以便重新领取。"""
- failed_rows = self.list_failed_group_items(biz_dt, group_ids=group_ids)
- if not failed_rows:
- return {"reset_items": 0, "reset_groups": 0, "group_ids": [], "items": []}
- item_ids = [int(row["item_id"]) for row in failed_rows]
- affected_group_ids = sorted({int(row["group_id"]) for row in failed_rows})
- self.session.execute(
- update(DemandGradePlanGroupItem)
- .where(DemandGradePlanGroupItem.id.in_(item_ids))
- .values(status="pending", error_message=None, finished_at=None)
- )
- self.session.execute(
- update(DemandGradePlanGroup)
- .where(DemandGradePlanGroup.id.in_(affected_group_ids))
- .values(
- status="pending",
- error_message=None,
- started_at=None,
- finished_at=None,
- )
- )
- return {
- "reset_items": len(item_ids),
- "reset_groups": len(affected_group_ids),
- "group_ids": affected_group_ids,
- "items": failed_rows,
- }
|