plan_group_batch.py 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145
  1. """计划组需求明细:物化与读取。"""
  2. from __future__ import annotations
  3. from typing import Any
  4. from sqlalchemy.orm import Session
  5. from supply_infra.db.repositories.demand_belong_category_repo import DemandBelongCategoryRepository
  6. from supply_infra.db.repositories.demand_belong_pool_rel_repo import DemandBelongPoolRelRepository
  7. from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
  8. from supply_infra.db.session import get_session
  9. MAX_DEMANDS_PER_BATCH = 30
  10. def _sort_category_units(
  11. units: list[tuple[int, int | None, list[Any]]],
  12. ) -> list[tuple[int, int | None, list[Any]]]:
  13. """同父分类相邻,同分类作为不可分割单元。"""
  14. return sorted(
  15. units,
  16. key=lambda item: (
  17. item[1] is None,
  18. item[1] if item[1] is not None else -1,
  19. item[0],
  20. ),
  21. )
  22. def pack_category_units(
  23. units: list[tuple[int, int | None, list[Any]]],
  24. *,
  25. max_demands_per_group: int = MAX_DEMANDS_PER_BATCH,
  26. ) -> list[list[int]]:
  27. """将分类需求单元打包为计划组;每组尽量不超过 max_demands_per_group 条需求。"""
  28. cap = max(1, int(max_demands_per_group))
  29. groups: list[list[int]] = []
  30. current: list[int] = []
  31. current_size = 0
  32. for category_id, _parent_id, demands in _sort_category_units(units):
  33. size = len(demands)
  34. if size <= 0:
  35. continue
  36. if size > cap:
  37. if current:
  38. groups.append(current)
  39. current = []
  40. current_size = 0
  41. groups.append([category_id])
  42. continue
  43. if current and current_size + size > cap:
  44. groups.append(current)
  45. current = []
  46. current_size = 0
  47. current.append(category_id)
  48. current_size += size
  49. if current:
  50. groups.append(current)
  51. return groups
  52. def split_even_batches(items: list[Any], *, max_per_batch: int = MAX_DEMANDS_PER_BATCH) -> list[list[Any]]:
  53. """按总数均分子批次:总数不超过上限则一批;否则递增组数直到每组不超过上限,余数从前组分配。"""
  54. n = len(items)
  55. if n == 0:
  56. return []
  57. cap = max(1, int(max_per_batch))
  58. if n <= cap:
  59. return [items]
  60. k = 2
  61. while True:
  62. base = n // k
  63. rem = n % k
  64. max_size = base + 1 if rem > 0 else base
  65. if max_size <= cap:
  66. break
  67. k += 1
  68. batches: list[list[Any]] = []
  69. idx = 0
  70. for i in range(k):
  71. size = base + (1 if i < rem else 0)
  72. batches.append(items[idx : idx + size])
  73. idx += size
  74. return batches
  75. def _priority_sort_key(name: str, priority_index: dict) -> tuple:
  76. return (
  77. priority_index.get(name, {}).get("source_rank_score") is None,
  78. -float(priority_index.get(name, {}).get("source_rank_score") or 0),
  79. float(priority_index.get(name, {}).get("global_demand_rank") or float("inf")),
  80. name,
  81. )
  82. def resolve_demands_for_category_ids(
  83. biz_dt: str,
  84. category_ids: list[int],
  85. *,
  86. session: Session | None = None,
  87. ) -> list[dict[str, Any]]:
  88. """按分类节点解析全部需求池记录(pool_id + demand_name)。"""
  89. selected_ids = list(dict.fromkeys(int(value) for value in category_ids))
  90. if session is None:
  91. with get_session() as owned_session:
  92. return resolve_demands_for_category_ids(
  93. biz_dt,
  94. selected_ids,
  95. session=owned_session,
  96. )
  97. belongs = DemandBelongCategoryRepository(session).list_by_category_ids(selected_ids)
  98. pool_ids_by_belong = DemandBelongPoolRelRepository(session).get_pool_ids_by_belong_ids(
  99. [int(row.id) for row in belongs]
  100. )
  101. pool_ids = sorted({pool_id for values in pool_ids_by_belong.values() for pool_id in values})
  102. pool_repo = MultiDemandPoolDiRepository(session)
  103. pool_rows = pool_repo.get_by_ids(pool_ids)
  104. from agents.demand_grade_agent.tools.demand_priority import build_demand_priority_index
  105. priority_index = build_demand_priority_index(pool_repo.list_by_biz_dt(biz_dt))
  106. candidates: list[dict[str, Any]] = []
  107. for row in pool_rows:
  108. if row.biz_dt != biz_dt or not row.demand_name:
  109. continue
  110. candidates.append({
  111. "pool_id": int(row.id),
  112. "demand_name": str(row.demand_name),
  113. })
  114. candidates.sort(
  115. key=lambda item: (
  116. *_priority_sort_key(item["demand_name"], priority_index),
  117. item["pool_id"],
  118. )
  119. )
  120. return candidates