"""当天节点分配状态查询(不含校验重试逻辑)。""" 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))