"""批次计划逐条入库(由 save_grade_plan 调用)。""" from __future__ import annotations import logging from typing import Any from agents.demand_grade_orchestrator_agent.common.assignment import ( MAX_DAILY_BATCHES, get_assigned_category_ids, get_existing_group_count, get_required_hanging_category_ids, get_unassigned_hanging_category_ids, ) 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 logger = logging.getLogger(__name__) def _filter_group_category_ids(biz_dt: str, category_ids: list[int]) -> list[int]: """入库前再次剔除已分配与无需求节点。""" assigned = get_assigned_category_ids(biz_dt) _by_id, _children, weights = load_tree_state(biz_dt) kept: list[int] = [] for category_id in category_ids: if category_id in assigned: continue if not has_hung_demand(weights.get(category_id)): continue kept.append(category_id) return kept def persist_groups_one_by_one(biz_dt: str, base_payload: dict[str, Any]) -> dict[str, Any]: """逐批入库,单批失败不影响其余批次;额度用尽时停止。本函数不向外抛异常。""" total_hanging_nodes = len(get_required_hanging_category_ids(biz_dt)) persisted_group_count = 0 persisted_groups: list[dict[str, Any]] = [] skipped_quota = 0 skipped_empty = 0 failed_groups: list[dict[str, Any]] = [] for group in base_payload.get("groups") or []: try: if get_existing_group_count(biz_dt) >= MAX_DAILY_BATCHES: skipped_quota += 1 continue category_ids = _filter_group_category_ids(biz_dt, list(group.get("category_ids") or [])) if not category_ids: skipped_empty += 1 continue unassigned_before = get_unassigned_hanging_category_ids(biz_dt) single_group = {**group, "category_ids": category_ids} single_payload = { **base_payload, "total_hanging_nodes": total_hanging_nodes, "groups": [single_group], "covered_category_ids": category_ids, "uncovered_category_ids": sorted(unassigned_before - set(category_ids)), "coverage_complete": not (unassigned_before - set(category_ids)), } with get_session() as session: DemandGradePlanRepository(session).create_plan(biz_dt, single_payload) persisted_group_count += 1 persisted_groups.append(single_group) except Exception as exc: logger.exception( "单批入库失败,已跳过并继续: biz_dt=%s group_id=%s category_ids=%s", biz_dt, group.get("group_id"), group.get("category_ids"), ) failed_groups.append({ "group_id": group.get("group_id"), "category_ids": group.get("category_ids"), "error": str(exc), }) try: unassigned_after = sorted(get_unassigned_hanging_category_ids(biz_dt)) existing_groups = get_existing_group_count(biz_dt) except Exception as exc: logger.exception("读取入库后状态失败: biz_dt=%s", biz_dt) unassigned_after = [] existing_groups = persisted_group_count return { "persisted_group_count": persisted_group_count, "persisted_groups": persisted_groups, "skipped_quota": skipped_quota, "skipped_empty": skipped_empty, "failed_groups": failed_groups, "existing_groups": existing_groups, "remaining_batch_quota": max(0, MAX_DAILY_BATCHES - existing_groups), "unassigned_category_ids": unassigned_after, "coverage_complete": not unassigned_after, "total_hanging_nodes": total_hanging_nodes, }