| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123 |
- """当天节点分配状态查询(不含校验重试逻辑)。"""
- from __future__ import annotations
- from typing import Any
- from agents.demand_grade_orchestrator_agent.common.tree_state import has_hung_demand, load_tree_state
- from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
- from supply_infra.db.session import get_session
- MAX_DAILY_BATCHES = 200
- def get_required_hanging_category_ids(biz_dt: str) -> set[int]:
- by_id, _children, weights = load_tree_state(biz_dt)
- return {
- category_id
- for category_id, weight in weights.items()
- if category_id in by_id and has_hung_demand(weight)
- }
- def get_assigned_category_ids(biz_dt: str) -> set[int]:
- with get_session() as session:
- return DemandGradePlanRepository(session).get_assigned_category_ids(biz_dt)
- def get_existing_group_count(biz_dt: str) -> int:
- with get_session() as session:
- snapshot = DemandGradePlanRepository(session).get_execution_snapshot(biz_dt)
- return int(snapshot["planned_groups"])
- def resolve_planning_state(biz_dt: str) -> dict[str, Any]:
- """汇总当天有需求节点、已分批节点与剩余批次额度。"""
- required = get_required_hanging_category_ids(biz_dt)
- assigned = get_assigned_category_ids(biz_dt)
- unassigned = sorted(required - assigned)
- existing_groups = get_existing_group_count(biz_dt)
- remaining_batch_quota = max(0, MAX_DAILY_BATCHES - existing_groups)
- batch_limit_reached = existing_groups >= MAX_DAILY_BATCHES
- has_unassigned_nodes = bool(unassigned)
- can_plan_more = has_unassigned_nodes and not batch_limit_reached
- skip_reason: str | None = None
- if batch_limit_reached:
- skip_reason = f"当天批次已达上限 {MAX_DAILY_BATCHES}"
- elif not has_unassigned_nodes:
- skip_reason = "当天有需求节点均已分批,无待分配节点"
- return {
- "biz_dt": biz_dt,
- "total_hanging_nodes": len(required),
- "required_category_ids": sorted(required),
- "assigned_category_ids": sorted(assigned),
- "unassigned_category_ids": unassigned,
- "existing_groups": existing_groups,
- "remaining_batch_quota": remaining_batch_quota,
- "batch_limit_reached": batch_limit_reached,
- "has_unassigned_nodes": has_unassigned_nodes,
- "can_plan_more": can_plan_more,
- "skip_reason": skip_reason,
- }
- def get_unassigned_hanging_category_ids(biz_dt: str) -> set[int]:
- """返回当天有需求且尚未进入任何批次的分类节点。"""
- return get_required_hanging_category_ids(biz_dt) - get_assigned_category_ids(biz_dt)
- def _sync_group_positions(group: dict[str, Any]) -> dict[str, Any]:
- category_ids = {int(value) for value in group.get("category_ids") or []}
- positions = group.get("category_global_positions")
- if not isinstance(positions, list):
- return group
- return {
- **group,
- "category_global_positions": [
- item for item in positions
- if int(item.get("category_id", -1)) in category_ids
- ],
- }
- def dedupe_cross_group_category_ids(plan: dict[str, Any]) -> list[int]:
- """同计划内后批次若含前批已出现的节点,从后批中移除。"""
- removed: list[int] = []
- seen: set[int] = set()
- cleaned_groups: list[dict[str, Any]] = []
- for group in plan.get("groups") or []:
- kept: list[int] = []
- for category_id in group.get("category_ids") or []:
- try:
- cid = int(category_id)
- except (TypeError, ValueError):
- continue
- if cid in seen:
- removed.append(cid)
- continue
- seen.add(cid)
- kept.append(cid)
- if kept:
- cleaned_groups.append(_sync_group_positions({**group, "category_ids": kept}))
- plan["groups"] = cleaned_groups
- return sorted(set(removed))
- def strip_assigned_category_ids(plan: dict[str, Any], assigned_ids: set[int]) -> list[int]:
- """从计划中移除当天已分批的分类,避免重复落库。"""
- removed: list[int] = []
- cleaned_groups: list[dict[str, Any]] = []
- for group in plan.get("groups") or []:
- kept: list[int] = []
- for category_id in group.get("category_ids") or []:
- try:
- cid = int(category_id)
- except (TypeError, ValueError):
- continue
- if cid in assigned_ids:
- removed.append(cid)
- else:
- kept.append(cid)
- if kept:
- cleaned_groups.append(_sync_group_positions({**group, "category_ids": kept}))
- plan["groups"] = cleaned_groups
- return sorted(set(removed))
|