| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102 |
- """批次计划逐条入库(由 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,
- }
|