plan_persist.py 4.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102
  1. """批次计划逐条入库(由 save_grade_plan 调用)。"""
  2. from __future__ import annotations
  3. import logging
  4. from typing import Any
  5. from agents.demand_grade_orchestrator_agent.common.assignment import (
  6. MAX_DAILY_BATCHES,
  7. get_assigned_category_ids,
  8. get_existing_group_count,
  9. get_required_hanging_category_ids,
  10. get_unassigned_hanging_category_ids,
  11. )
  12. from agents.demand_grade_orchestrator_agent.common.tree_state import has_hung_demand, load_tree_state
  13. from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
  14. from supply_infra.db.session import get_session
  15. logger = logging.getLogger(__name__)
  16. def _filter_group_category_ids(biz_dt: str, category_ids: list[int]) -> list[int]:
  17. """入库前再次剔除已分配与无需求节点。"""
  18. assigned = get_assigned_category_ids(biz_dt)
  19. _by_id, _children, weights = load_tree_state(biz_dt)
  20. kept: list[int] = []
  21. for category_id in category_ids:
  22. if category_id in assigned:
  23. continue
  24. if not has_hung_demand(weights.get(category_id)):
  25. continue
  26. kept.append(category_id)
  27. return kept
  28. def persist_groups_one_by_one(biz_dt: str, base_payload: dict[str, Any]) -> dict[str, Any]:
  29. """逐批入库,单批失败不影响其余批次;额度用尽时停止。本函数不向外抛异常。"""
  30. total_hanging_nodes = len(get_required_hanging_category_ids(biz_dt))
  31. persisted_group_count = 0
  32. persisted_groups: list[dict[str, Any]] = []
  33. skipped_quota = 0
  34. skipped_empty = 0
  35. failed_groups: list[dict[str, Any]] = []
  36. for group in base_payload.get("groups") or []:
  37. try:
  38. if get_existing_group_count(biz_dt) >= MAX_DAILY_BATCHES:
  39. skipped_quota += 1
  40. continue
  41. category_ids = _filter_group_category_ids(biz_dt, list(group.get("category_ids") or []))
  42. if not category_ids:
  43. skipped_empty += 1
  44. continue
  45. unassigned_before = get_unassigned_hanging_category_ids(biz_dt)
  46. single_group = {**group, "category_ids": category_ids}
  47. single_payload = {
  48. **base_payload,
  49. "total_hanging_nodes": total_hanging_nodes,
  50. "groups": [single_group],
  51. "covered_category_ids": category_ids,
  52. "uncovered_category_ids": sorted(unassigned_before - set(category_ids)),
  53. "coverage_complete": not (unassigned_before - set(category_ids)),
  54. }
  55. with get_session() as session:
  56. DemandGradePlanRepository(session).create_plan(biz_dt, single_payload)
  57. persisted_group_count += 1
  58. persisted_groups.append(single_group)
  59. except Exception as exc:
  60. logger.exception(
  61. "单批入库失败,已跳过并继续: biz_dt=%s group_id=%s category_ids=%s",
  62. biz_dt,
  63. group.get("group_id"),
  64. group.get("category_ids"),
  65. )
  66. failed_groups.append({
  67. "group_id": group.get("group_id"),
  68. "category_ids": group.get("category_ids"),
  69. "error": str(exc),
  70. })
  71. try:
  72. unassigned_after = sorted(get_unassigned_hanging_category_ids(biz_dt))
  73. existing_groups = get_existing_group_count(biz_dt)
  74. except Exception as exc:
  75. logger.exception("读取入库后状态失败: biz_dt=%s", biz_dt)
  76. unassigned_after = []
  77. existing_groups = persisted_group_count
  78. return {
  79. "persisted_group_count": persisted_group_count,
  80. "persisted_groups": persisted_groups,
  81. "skipped_quota": skipped_quota,
  82. "skipped_empty": skipped_empty,
  83. "failed_groups": failed_groups,
  84. "existing_groups": existing_groups,
  85. "remaining_batch_quota": max(0, MAX_DAILY_BATCHES - existing_groups),
  86. "unassigned_category_ids": unassigned_after,
  87. "coverage_complete": not unassigned_after,
  88. "total_hanging_nodes": total_hanging_nodes,
  89. }