assignment.py 4.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123
  1. """当天节点分配状态查询(不含校验重试逻辑)。"""
  2. from __future__ import annotations
  3. from typing import Any
  4. from agents.demand_grade_orchestrator_agent.common.tree_state import has_hung_demand, load_tree_state
  5. from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
  6. from supply_infra.db.session import get_session
  7. MAX_DAILY_BATCHES = 200
  8. def get_required_hanging_category_ids(biz_dt: str) -> set[int]:
  9. by_id, _children, weights = load_tree_state(biz_dt)
  10. return {
  11. category_id
  12. for category_id, weight in weights.items()
  13. if category_id in by_id and has_hung_demand(weight)
  14. }
  15. def get_assigned_category_ids(biz_dt: str) -> set[int]:
  16. with get_session() as session:
  17. return DemandGradePlanRepository(session).get_assigned_category_ids(biz_dt)
  18. def get_existing_group_count(biz_dt: str) -> int:
  19. with get_session() as session:
  20. snapshot = DemandGradePlanRepository(session).get_execution_snapshot(biz_dt)
  21. return int(snapshot["planned_groups"])
  22. def resolve_planning_state(biz_dt: str) -> dict[str, Any]:
  23. """汇总当天有需求节点、已分批节点与剩余批次额度。"""
  24. required = get_required_hanging_category_ids(biz_dt)
  25. assigned = get_assigned_category_ids(biz_dt)
  26. unassigned = sorted(required - assigned)
  27. existing_groups = get_existing_group_count(biz_dt)
  28. remaining_batch_quota = max(0, MAX_DAILY_BATCHES - existing_groups)
  29. batch_limit_reached = existing_groups >= MAX_DAILY_BATCHES
  30. has_unassigned_nodes = bool(unassigned)
  31. can_plan_more = has_unassigned_nodes and not batch_limit_reached
  32. skip_reason: str | None = None
  33. if batch_limit_reached:
  34. skip_reason = f"当天批次已达上限 {MAX_DAILY_BATCHES}"
  35. elif not has_unassigned_nodes:
  36. skip_reason = "当天有需求节点均已分批,无待分配节点"
  37. return {
  38. "biz_dt": biz_dt,
  39. "total_hanging_nodes": len(required),
  40. "required_category_ids": sorted(required),
  41. "assigned_category_ids": sorted(assigned),
  42. "unassigned_category_ids": unassigned,
  43. "existing_groups": existing_groups,
  44. "remaining_batch_quota": remaining_batch_quota,
  45. "batch_limit_reached": batch_limit_reached,
  46. "has_unassigned_nodes": has_unassigned_nodes,
  47. "can_plan_more": can_plan_more,
  48. "skip_reason": skip_reason,
  49. }
  50. def get_unassigned_hanging_category_ids(biz_dt: str) -> set[int]:
  51. """返回当天有需求且尚未进入任何批次的分类节点。"""
  52. return get_required_hanging_category_ids(biz_dt) - get_assigned_category_ids(biz_dt)
  53. def _sync_group_positions(group: dict[str, Any]) -> dict[str, Any]:
  54. category_ids = {int(value) for value in group.get("category_ids") or []}
  55. positions = group.get("category_global_positions")
  56. if not isinstance(positions, list):
  57. return group
  58. return {
  59. **group,
  60. "category_global_positions": [
  61. item for item in positions
  62. if int(item.get("category_id", -1)) in category_ids
  63. ],
  64. }
  65. def dedupe_cross_group_category_ids(plan: dict[str, Any]) -> list[int]:
  66. """同计划内后批次若含前批已出现的节点,从后批中移除。"""
  67. removed: list[int] = []
  68. seen: set[int] = set()
  69. cleaned_groups: list[dict[str, Any]] = []
  70. for group in plan.get("groups") or []:
  71. kept: list[int] = []
  72. for category_id in group.get("category_ids") or []:
  73. try:
  74. cid = int(category_id)
  75. except (TypeError, ValueError):
  76. continue
  77. if cid in seen:
  78. removed.append(cid)
  79. continue
  80. seen.add(cid)
  81. kept.append(cid)
  82. if kept:
  83. cleaned_groups.append(_sync_group_positions({**group, "category_ids": kept}))
  84. plan["groups"] = cleaned_groups
  85. return sorted(set(removed))
  86. def strip_assigned_category_ids(plan: dict[str, Any], assigned_ids: set[int]) -> list[int]:
  87. """从计划中移除当天已分批的分类,避免重复落库。"""
  88. removed: list[int] = []
  89. cleaned_groups: list[dict[str, Any]] = []
  90. for group in plan.get("groups") or []:
  91. kept: list[int] = []
  92. for category_id in group.get("category_ids") or []:
  93. try:
  94. cid = int(category_id)
  95. except (TypeError, ValueError):
  96. continue
  97. if cid in assigned_ids:
  98. removed.append(cid)
  99. else:
  100. kept.append(cid)
  101. if kept:
  102. cleaned_groups.append(_sync_group_positions({**group, "category_ids": kept}))
  103. plan["groups"] = cleaned_groups
  104. return sorted(set(removed))