Przeglądaj źródła

需求分配完整流程修改

xueyiming 1 tydzień temu
rodzic
commit
c729938062
30 zmienionych plików z 1174 dodań i 1320 usunięć
  1. 2 9
      agents/demand_grade_agent/prompt/system_prompt.md
  2. 36 24
      agents/demand_grade_agent/run.py
  3. 0 3
      agents/demand_grade_agent/tools/__init__.py
  4. 46 6
      agents/demand_grade_agent/tools/batch_save_demand_grades.py
  5. 0 169
      agents/demand_grade_agent/tools/build_grade_plan_context.py
  6. 134 0
      agents/demand_grade_orchestrator_agent/_verify_logic.py
  7. 1 1
      agents/demand_grade_orchestrator_agent/agent.py
  8. 4 2
      agents/demand_grade_orchestrator_agent/common/__init__.py
  9. 123 0
      agents/demand_grade_orchestrator_agent/common/assignment.py
  10. 0 211
      agents/demand_grade_orchestrator_agent/common/plan_builder.py
  11. 102 0
      agents/demand_grade_orchestrator_agent/common/plan_persist.py
  12. 198 0
      agents/demand_grade_orchestrator_agent/common/plan_record.py
  13. 7 0
      agents/demand_grade_orchestrator_agent/common/tree_state.py
  14. 24 11
      agents/demand_grade_orchestrator_agent/prompt/system_prompt.md
  15. 74 157
      agents/demand_grade_orchestrator_agent/run.py
  16. 2 2
      agents/demand_grade_orchestrator_agent/tools/__init__.py
  17. 0 20
      agents/demand_grade_orchestrator_agent/tools/build_full_day_grade_plan.py
  18. 20 9
      agents/demand_grade_orchestrator_agent/tools/query_global_heat_tree.py
  19. 6 7
      agents/demand_grade_orchestrator_agent/tools/query_heat_node_group.py
  20. 132 0
      agents/demand_grade_orchestrator_agent/tools/save_grade_plan.py
  21. 0 25
      agents/demand_grade_orchestrator_agent/validation/__init__.py
  22. 0 85
      agents/demand_grade_orchestrator_agent/validation/assignment.py
  23. 0 87
      agents/demand_grade_orchestrator_agent/validation/full_day_grade_plan.py
  24. 8 8
      jobs/grade_demand_pool.py
  25. 85 0
      scripts/run_grade_plan_groups.py
  26. 6 30
      supply_infra/db/repositories/demand_grade_plan_repo.py
  27. 3 3
      supply_infra/scheduler/app.py
  28. 0 143
      supply_infra/scheduler/grade_assignment.py
  29. 102 308
      supply_infra/scheduler/jobs/grade_demand_pool.py
  30. 59 0
      supply_infra/scheduler/plan_group_batch.py

+ 2 - 9
agents/demand_grade_agent/prompt/system_prompt.md

@@ -3,11 +3,8 @@
 和对应的 biz_dt,任务是结合其归属树节点的**全局热度**与**后验真实效果**,逐一划分 S/A/B/C/D
 五档优先级,并调用工具落库到 `demand_grade` 表,供下游选题/投放决策参考。
 
-调用方通常还会提供一份“全局树热度交接件”:其中包含批次热度等级、需求节点、父节点、全部兄弟节点的热度、局部排名和整树排名。
-它是评分的必需输入,不可只看需求自身数据,也不可只看附近节点。
-
 你只做分级判断,不生成新需求词,也不修改需求池原始数据;只处理消息中给定的这些需求词,
-不需要自行查找或列举其他待处理需求。
+不需要自行查找或列举其他待处理需求。所有分类、热度、后验等证据必须通过工具从数据库查询。
 
 ## 全局热度 / 后验含义
 - **全局热度**:`category_tree_weight.total_score`。反映该需求所在树节点在类目树里的历史热度排名,是"没有真实上线数据时"的兜底依据。
@@ -17,7 +14,6 @@
   - `real_rov_7d_count = 0`(或数据缺失):说明效果未知,只能用全局热度兜底判断。**无论全局热度多高,都不建议给到 S 级**(因为没有真实验证支撑),一般封顶在 A。
 
 ## 四类证据必须分开
-- **批次热度等级**:统筹 Agent 根据分类节点在整棵树中的排名分划为 S/A/B/C/D;U 表示分类热度数据不足。它描述整批所在的全局热区,不等于批内每条需求的最终等级。
 - **分类节点全局/局部证据**:`category_tree_weight.total_score`、节点整树名次、父节点和全部兄弟节点。它描述需求所在分类环境。
 - **需求自身来源归一分**:`demand_priority.source_rank_score`,范围 0-100。先在每个 strategy 内独立按原始 `weight` 排名归一化,再对该需求已有来源的归一分取均值。它是具体需求之间可比较的先验信号。
 - **需求词后验**:真实 ROV/VOV 及样本数,优先级高于纯先验。
@@ -65,7 +61,6 @@
 - `query_category_local_heat(biz_dt, category_ids)`:查询节点自身、父节点、**全部兄弟节点**的全局热度 total_score + 后验 real_rov_7d/real_vov_7d,并给出兄弟内排名;判级时必须参考列出的全部相关节点,不得只看部分节点。
 - `query_demand_popularity_by_word(demand_word_names, biz_dt=None)`:按需求词粒度直接查后验热度统计,**可一次传入多个词** 交叉验证树节点级结论。
 - `query_score_distribution(biz_dt=None)`:分别查询分类树全局热度、需求自身来源归一分与后验分布,制定跨批次一致标准。
-- `build_grade_plan_context(...)`:按统筹规划选定的树节点生成待分级需求、节点组原因,并附带 `local_heat` 局部热度快照;需要复查局部环境时可再调用 `query_category_local_heat`。
 - `batch_save_demand_grades(items, biz_dt=None)`:批量落库分级结果,可重复调用按 (biz_dt, demand_name) upsert 覆盖修正。`related_pool_ids` 必填,`video_list`/`strategies` 自动推导。
 
 ## 工作流程
@@ -78,10 +73,9 @@
    - 必要时一次 `query_demand_popularity_by_word(demand_word_names=[...])` 做词粒度交叉验证。
    各工具返回结果每段均标注原始查询词(如 `--- demand_name: xxx ---`),便于对应落库。
 4. 对给定列表中的每一个需求词,结合以下数据判定 S/A/B/C/D:
-   - 需求词自身:交接件 `items[].demand_priority` / `search_related_pool_demands` / `query_demand_popularity_by_word`
+   - 需求词自身:`search_related_pool_demands` / `query_demand_popularity_by_word`
    - 归属分类节点:`query_demand_category_and_weight`
    - 全局与局部环境(整树名次 + 父节点 + 全部兄弟节点完整权重):`query_category_local_heat`
-   - 批次热度:交接件 `plan.batch_heat_level`,只作为环境证据,不直接复制成需求等级
    reason 必须同时写清需求自身来源归一分/有效来源数、分类节点整树位置、局部判断和词级后验;`related_pool_ids` 取自 `search_related_pool_demands` 返回的 `[id=...]`。
 5. 全部处理完后,调用一次(或分 2~3 次)`batch_save_demand_grades` 落库,覆盖这批给定的所有需求词。
    `score` 不需要自行计算或传入,保存工具会用当日全量需求池确定性重算 `source_rank_score` 并落库;禁止另造一套模型分。
@@ -92,4 +86,3 @@
 - reason 必须具体:写清引用的全局热度/后验数值、是否合并了同义词、依据哪个树节点。
 - 找不到归属树节点或权重数据的需求:如实说明"无法评级/数据缺失",不要强行给出等级去凑数。
 - 有后验数据始终优先于纯全局热度判断;无后验数据时保持谨慎,不给最高档。
-- 批次等级和需求等级是两个不同结论;禁止把 S 热批次里的所有需求直接评为 S。

+ 36 - 24
agents/demand_grade_agent/run.py

@@ -6,46 +6,58 @@
 """
 from __future__ import annotations
 
-import json
 from typing import Any
 
 from agents.demand_grade_agent import create_demand_grade_agent
 
 
+def build_grade_user_input(demands: list[dict[str, Any]], biz_dt: str) -> str:
+    """构建传给分级 Agent 的用户消息。"""
+    demand_lines = []
+    for item in demands:
+        pool_id = item.get("pool_id", item.get("id"))
+        if pool_id is None:
+            raise ValueError(f"demand 缺少 pool_id: {item!r}")
+        demand_lines.append(f"[{int(pool_id)}] {str(item['demand_name']).strip()}")
+
+    lines_text = "\n".join(demand_lines)
+    return f"""请对以下 {len(demand_lines)} 个需求词逐一评级(S/A/B/C/D)。
+只处理这一批,不要尝试查找或列举更多需求词。判级完成后调用 batch_save_demand_grades 落库。
+落库时 related_pool_ids 使用列表中对应行的 pool_id;分类、热度、后验等证据请通过工具从数据库查询。
+
+biz_dt={biz_dt}
+需求词列表 [pool_id] demand_name
+{lines_text}
+"""
+
+
 def main(
-    demand_names: list[str],
+    demands: list[dict[str, Any]],
     biz_dt: str | None = None,
-    tree_context: dict[str, Any] | None = None,
 ) -> None:
     agent = create_demand_grade_agent()
     print(f"demand_grade_agent ready | model={agent.model}")
     print(f"tools: {agent.tools.list_tools()}")
     print()
 
-    names_str = "\n".join(f"- {name}" for name in demand_names)
-    biz_dt_note = f"biz_dt={biz_dt}" if biz_dt else "未指定 biz_dt,请先调用 query_latest_biz_dt() 确定"
-    handoff = (
-        json.dumps(tree_context, ensure_ascii=False)
-        if tree_context is not None
-        else "未提供;请通过 build_grade_plan_context 查询。"
-    )
-    user_input = f"""
-    请对以下 {len(demand_names)} 个需求词逐一评级(S/A/B/C/D),{biz_dt_note}。
-    只处理这一批,不要尝试查找或列举更多需求词。判级完成后调用 batch_save_demand_grades 落库。
-
-    需求词列表:
-    {names_str}
-
-    以下是统筹规划提供的节点组上下文(含批次热度等级、需求自身来源归一分、
-    `local_heat` 局部热度与整树名次快照)。批次等级不等于需求最终等级。
-    必须结合整树位置、父节点、全部兄弟节点和需求自身分做校正;需要更细信息时可调用
-    `query_category_local_heat`。不同来源原始 weight 不得直接相加:
-    {handoff}
-    """
+    if not demands:
+        raise ValueError("demands 不能为空")
+
+    batch_biz_dt = (biz_dt or "").strip()
+    if not batch_biz_dt:
+        raise ValueError("biz_dt 不能为空")
+
+    user_input = build_grade_user_input(demands, batch_biz_dt)
     result = agent.run(user_input)
     print(result.content)
     print(f"\n[iterations={result.iterations}, tool_calls={result.tool_calls_made}]")
 
 
 if __name__ == "__main__":
-    main(["因果报应", "降半旗"], biz_dt="20260714")
+    main(
+        [
+            {"pool_id": 101, "demand_name": "减脂期加餐"},
+            {"pool_id": 102, "demand_name": "减脂期"},
+        ],
+        biz_dt="20260721",
+    )

+ 0 - 3
agents/demand_grade_agent/tools/__init__.py

@@ -10,7 +10,6 @@ from collections.abc import Callable
 from typing import Any
 
 from agents.demand_grade_agent.tools.batch_save_demand_grades import batch_save_demand_grades
-from agents.demand_grade_agent.tools.build_grade_plan_context import build_grade_plan_context
 from agents.demand_grade_agent.tools.query_category_local_heat import query_category_local_heat
 from agents.demand_grade_agent.tools.query_category_path import query_category_path
 from agents.demand_grade_agent.tools.query_demand_category_and_weight import (
@@ -34,14 +33,12 @@ ALL_TOOLS: list[Callable[..., Any]] = [
     query_category_local_heat,
     query_demand_popularity_by_word,
     query_score_distribution,
-    build_grade_plan_context,
     batch_save_demand_grades,
 ]
 
 __all__ = [
     "ALL_TOOLS",
     "batch_save_demand_grades",
-    "build_grade_plan_context",
     "query_category_local_heat",
     "query_category_path",
     "query_demand_category_and_weight",

+ 46 - 6
agents/demand_grade_agent/tools/batch_save_demand_grades.py

@@ -3,6 +3,7 @@
 """
 from __future__ import annotations
 
+import json
 import logging
 from decimal import Decimal
 from typing import Any, Optional
@@ -24,6 +25,35 @@ from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPo
 from supply_infra.db.session import get_session
 
 logger = logging.getLogger(__name__)
+_MAX_ERROR_DETAILS = 10
+
+
+def _format_errors(errors: list[str]) -> str:
+    if not errors:
+        return ""
+    if len(errors) <= _MAX_ERROR_DETAILS:
+        return ";".join(errors)
+    hidden = len(errors) - _MAX_ERROR_DETAILS
+    return ";".join(errors[:_MAX_ERROR_DETAILS]) + f";...另有 {hidden} 条类似错误"
+
+
+def _coerce_items(raw: Any) -> tuple[list[Any], str | None]:
+    """将 Agent 传入的 items 规范为列表,兼容误传 JSON 字符串。"""
+    if raw is None:
+        return [], "items 不能为空"
+    if isinstance(raw, str):
+        text = raw.strip()
+        if not text:
+            return [], "items 不能为空"
+        try:
+            raw = json.loads(text)
+        except json.JSONDecodeError:
+            return [], "items 必须是对象数组,不能把未解析的 JSON 字符串直接传入"
+    if isinstance(raw, dict):
+        return [raw], None
+    if not isinstance(raw, list):
+        return [], f"items 必须是数组,当前类型: {type(raw).__name__}"
+    return raw, None
 
 
 def _optional_decimal(value: Any, field: str, idx: int, errors: list[str]) -> Decimal | None:
@@ -99,6 +129,12 @@ def _normalize_items(
             continue
 
         related_pool_ids = _optional_int_list(item.get("related_pool_ids"), "related_pool_ids", idx, errors)
+        if not related_pool_ids and item.get("pool_id") is not None:
+            try:
+                related_pool_ids = [int(item["pool_id"])]
+            except (TypeError, ValueError):
+                errors.append(f"第 {idx} 项 pool_id 无效: {item.get('pool_id')!r}")
+                continue
         if not related_pool_ids:
             errors.append(
                 f"第 {idx} 项缺少 related_pool_ids(必填,需先用 search_related_pool_demands 找到对应的 "
@@ -164,8 +200,8 @@ def batch_save_demand_grades(items: list[dict[str, Any]], biz_dt: Optional[str]
             - demand_name (必填): 需求名称
             - grade (必填): S/A/B/C/D 之一
             - reason (必填): 判断依据,需引用具体的先验/后验数值
-            - related_pool_ids (必填): 该需求对应的 multi_demand_pool_di.id 列表,需先调用
-              search_related_pool_demands 找到;用于关联原始需求,并自动推导 video_list/strategies
+            - related_pool_ids (必填): 该需求对应的 multi_demand_pool_di.id 列表;若输入里给了
+              pool_id,也可直接写 "pool_id": 123 代替 related_pool_ids
             - score: 无需传入;保存时按当日全量需求池自动计算需求自身来源归一分(0-100),
               即各 strategy 内独立排名归一化后,对该需求已有来源取均值
             - category_ids (可选): 归属的树节点 id 列表,会写入 demand_grade_category_rel 映射表
@@ -185,11 +221,15 @@ def batch_save_demand_grades(items: list[dict[str, Any]], biz_dt: Optional[str]
     if err:
         return err
 
+    coerced_items, coerce_err = _coerce_items(items)
+    if coerce_err:
+        return coerce_err
+
     rows, related_pool_id_lists, category_id_lists, errors = _normalize_items(
-        items, default_biz_dt=default_biz_dt
+        coerced_items, default_biz_dt=default_biz_dt
     )
     if not rows:
-        detail = ";".join(errors) if errors else "无有效数据"
+        detail = _format_errors(errors) if errors else "无有效数据"
         return f"没有可保存的数据: {detail}"
 
     try:
@@ -229,7 +269,7 @@ def batch_save_demand_grades(items: list[dict[str, Any]], biz_dt: Optional[str]
                 saved_indices.append(i)
 
             if not final_rows:
-                detail = ";".join(errors) if errors else "无有效数据"
+                detail = _format_errors(errors) if errors else "无有效数据"
                 return f"没有可保存的数据: {detail}"
 
             grade_repo = DemandGradeRepository(session)
@@ -253,7 +293,7 @@ def batch_save_demand_grades(items: list[dict[str, Any]], biz_dt: Optional[str]
 
         parts = [f"提交 {len(rows)} 条,成功写入/更新 {affected} 条({len(final_rows)} 条通过校验)"]
         if errors:
-            parts.append(f"校验失败/警告 {len(errors)} 条: " + ";".join(errors))
+            parts.append(f"校验失败/警告 {len(errors)} 条: {_format_errors(errors)}")
 
         message = "。".join(parts)
         logger.info("batch_save_demand_grades completed: %s", message)

+ 0 - 169
agents/demand_grade_agent/tools/build_grade_plan_context.py

@@ -1,169 +0,0 @@
-"""将已落库的节点组任务转换为分级 Agent 的小批次输入。"""
-from __future__ import annotations
-
-import json
-
-from agents.demand_grade_agent.tools.demand_priority import (
-    DEMAND_PRIORITY_SCORE_METHOD,
-    build_demand_priority_index,
-)
-from agents.demand_grade_agent.tools.tree_local import build_local_heat_snapshot
-from supply_agent.tools import tool
-from supply_infra.db.repositories.demand_belong_category_repo import DemandBelongCategoryRepository
-from supply_infra.db.repositories.demand_belong_pool_rel_repo import DemandBelongPoolRelRepository
-from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
-from supply_infra.db.session import get_session
-
-
-def _parse_group_metadata(shared_traits: str) -> dict:
-    try:
-        parsed = json.loads(shared_traits)
-    except (json.JSONDecodeError, TypeError):
-        return {"description": shared_traits.strip()}
-    if not isinstance(parsed, dict):
-        return {"description": shared_traits.strip()}
-    return parsed
-
-
-def build_supplement_grade_context(biz_dt: str, demand_names: list[str]) -> dict:
-    """为计划任务执行后仍未分级的需求构建补偿上下文。"""
-    requested_names = list(dict.fromkeys(
-        str(name).strip()
-        for name in demand_names
-        if name is not None and str(name).strip()
-    ))
-    if not requested_names:
-        return {
-            "biz_dt": biz_dt,
-            "items": [],
-            "plan": {"selection_method": "supplement_missing_grades"},
-            "local_heat": [],
-        }
-
-    requested_set = set(requested_names)
-    with get_session() as session:
-        pool_repo = MultiDemandPoolDiRepository(session)
-        all_pool_rows = pool_repo.list_by_biz_dt(biz_dt)
-        matched_pool_rows = [
-            row for row in all_pool_rows if row.demand_name in requested_set
-        ]
-        pool_name_by_id = {
-            int(row.id): str(row.demand_name)
-            for row in matched_pool_rows
-            if row.demand_name
-        }
-        priority_index = build_demand_priority_index(all_pool_rows)
-        belong_ids_by_pool = DemandBelongPoolRelRepository(session).get_belong_ids_by_pool_ids(
-            list(pool_name_by_id)
-        )
-        all_belong_ids = {
-            belong_id for values in belong_ids_by_pool.values() for belong_id in values
-        }
-        belongs = DemandBelongCategoryRepository(session).get_by_ids(all_belong_ids)
-        category_by_belong = {
-            int(row.id): int(row.category_id)
-            for row in belongs
-            if row.category_id is not None
-        }
-
-    category_ids_by_name: dict[str, set[int]] = {
-        name: set() for name in requested_names
-    }
-    for pool_id, demand_name in pool_name_by_id.items():
-        for belong_id in belong_ids_by_pool.get(pool_id, []):
-            category_id = category_by_belong.get(int(belong_id))
-            if category_id is not None:
-                category_ids_by_name[demand_name].add(category_id)
-
-    ordered_names = sorted(requested_names, key=lambda name: (
-        priority_index.get(name, {}).get("source_rank_score") is None,
-        -float(priority_index.get(name, {}).get("source_rank_score") or 0),
-        float(priority_index.get(name, {}).get("global_demand_rank") or float("inf")),
-        name,
-    ))
-    selected_category_ids = sorted({
-        category_id
-        for name in ordered_names
-        for category_id in category_ids_by_name.get(name, set())
-    })
-    return {
-        "biz_dt": biz_dt,
-        "items": [
-            {
-                "demand_name": name,
-                "category_ids": sorted(category_ids_by_name.get(name, set())),
-                "demand_priority": priority_index.get(name),
-            }
-            for name in ordered_names
-        ],
-        "demand_priority_score_method": DEMAND_PRIORITY_SCORE_METHOD,
-        "plan": {
-            "selected_category_ids": selected_category_ids,
-            "planning_reason": "计划任务执行后发现树上需求仍缺少等级,进入补充分级",
-            "shared_traits": "补偿任务按缺失需求直接分批,每批最多30个",
-            "batch_heat_level": None,
-            "selection_method": "supplement_missing_grades",
-            "is_supplement": True,
-        },
-        "local_heat": build_local_heat_snapshot(biz_dt, selected_category_ids),
-    }
-
-
-@tool
-def build_grade_plan_context(
-    biz_dt: str,
-    category_ids: list[int],
-    planning_reason: str,
-    shared_traits: str,
-    max_demands: int = 20,
-    excluded_demand_names: list[str] | None = None,
-) -> str:
-    """按节点组取一小批待分级需求,并保留统筹原因与共同特征。"""
-    selected_ids = list(dict.fromkeys(int(value) for value in category_ids))
-    excluded = set(excluded_demand_names or [])
-    with get_session() as session:
-        belongs = DemandBelongCategoryRepository(session).list_by_category_ids(selected_ids)
-        pool_ids_by_belong = DemandBelongPoolRelRepository(session).get_pool_ids_by_belong_ids(
-            [int(row.id) for row in belongs]
-        )
-        pool_ids = sorted({pool_id for values in pool_ids_by_belong.values() for pool_id in values})
-        pool_repo = MultiDemandPoolDiRepository(session)
-        pool_rows = pool_repo.get_by_ids(pool_ids)
-        priority_index = build_demand_priority_index(pool_repo.list_by_biz_dt(biz_dt))
-        candidate_names = {
-            str(row.demand_name)
-            for row in pool_rows
-            if row.biz_dt == biz_dt and row.demand_name and row.demand_name not in excluded
-        }
-    names = sorted(candidate_names, key=lambda name: (
-        priority_index.get(name, {}).get("source_rank_score") is None,
-        -float(priority_index.get(name, {}).get("source_rank_score") or 0),
-        float(priority_index.get(name, {}).get("global_demand_rank") or float("inf")),
-        name,
-    ))[: max(1, min(int(max_demands), 100))]
-    if not names:
-        return "所选树节点下没有待分级需求"
-    group_metadata = _parse_group_metadata(shared_traits)
-    return json.dumps({
-        "biz_dt": biz_dt,
-        "items": [
-            {
-                "demand_name": name,
-                "demand_priority": priority_index.get(name),
-            }
-            for name in names
-        ],
-        "demand_priority_score_method": DEMAND_PRIORITY_SCORE_METHOD,
-        "plan": {
-            "selected_category_ids": selected_ids,
-            "planning_reason": planning_reason.strip(),
-            "shared_traits": group_metadata.get("description", shared_traits.strip()),
-            "batch_heat_level": group_metadata.get("batch_heat_level"),
-            "batch_heat_label": group_metadata.get("batch_heat_label"),
-            "batch_global_rank_score": group_metadata.get("batch_global_rank_score"),
-            "batch_total_score_avg": group_metadata.get("batch_total_score_avg"),
-            "category_global_positions": group_metadata.get("category_global_positions", []),
-            "selection_method": "batch_heat_level_and_tree_adjacency",
-        },
-        "local_heat": build_local_heat_snapshot(biz_dt, selected_ids),
-    }, ensure_ascii=False)

+ 134 - 0
agents/demand_grade_orchestrator_agent/_verify_logic.py

@@ -0,0 +1,134 @@
+"""统筹 Agent 逻辑与工具自检。"""
+from __future__ import annotations
+
+import json
+import sys
+from unittest.mock import patch
+
+from agents.demand_grade_orchestrator_agent.common.assignment import (
+    MAX_DAILY_BATCHES,
+    dedupe_cross_group_category_ids,
+    strip_assigned_category_ids,
+)
+from agents.demand_grade_orchestrator_agent.common.plan_record import prepare_grade_groups
+from agents.demand_grade_orchestrator_agent.run import _summarize_agent_saves
+from supply_agent.types import Message, Role
+
+
+def _ok(name: str) -> None:
+    print(f"  ✓ {name}")
+
+
+def test_prepare_grade_groups_permissive() -> None:
+    fake_by_id = {
+        101: type("C", (), {"id": 101, "name": "A", "parent_id": None, "level": 1})(),
+        102: type("C", (), {"id": 102, "name": "B", "parent_id": None, "level": 1})(),
+    }
+    fake_weights = {
+        101: type("W", (), {"category_id": 101, "total_score": 0.9, "hung_word_count": 5})(),
+        102: type("W", (), {"category_id": 102, "total_score": 0.8, "hung_word_count": 0})(),
+    }
+    groups = [
+        {"category_ids": [101, 102, 999], "batch_heat_level": "X", "planning_reason": "", "shared_traits": ""},
+        {"category_ids": [101], "batch_heat_level": "A", "planning_reason": "dup", "shared_traits": "dup"},
+    ]
+    with patch(
+        "agents.demand_grade_orchestrator_agent.common.plan_record.load_tree_state",
+        return_value=(fake_by_id, {None: [101, 102]}, fake_weights),
+    ), patch(
+        "agents.demand_grade_orchestrator_agent.common.plan_record.global_heat_positions",
+        return_value={101: {"rank": 1, "total": 1, "normalized_score": 0.95}},
+    ), patch(
+        "agents.demand_grade_orchestrator_agent.common.plan_record.has_hung_demand",
+        side_effect=lambda w: w is not None and int(w.hung_word_count or 0) > 0,
+    ), patch(
+        "agents.demand_grade_orchestrator_agent.common.plan_record.path",
+        side_effect=lambda cid, _by: f"path-{cid}",
+    ), patch(
+        "agents.demand_grade_orchestrator_agent.common.plan_record.heat_level",
+        return_value="A",
+    ):
+        prepared = prepare_grade_groups(
+            "20260714",
+            "策略",
+            groups,
+            assigned_category_ids=set(),
+        )
+    assert len(prepared["groups"]) == 1
+    assert prepared["groups"][0]["category_ids"] == [101]
+    assert prepared["groups"][0]["batch_heat_level"] == "A"
+    assert prepared["groups"][0]["planning_reason"]
+    _ok("无需求/非法字段不报错,仅过滤后入库")
+
+
+def test_dedupe_cross_group() -> None:
+    plan = {"groups": [{"category_ids": [1, 2]}, {"category_ids": [2, 3]}]}
+    removed = dedupe_cross_group_category_ids(plan)
+    assert removed == [2] and plan["groups"][1]["category_ids"] == [3]
+    _ok("后批重复节点过滤")
+
+
+def test_summarize_agent_saves() -> None:
+    class Result:
+        messages = [
+            Message(
+                role=Role.TOOL,
+                name="save_grade_plan",
+                content=json.dumps({"ok": True, "persisted": True, "persisted_group_count": 2}),
+            ),
+        ]
+
+    summary = _summarize_agent_saves(Result())
+    assert summary["save_count"] == 1 and summary["persisted_groups"] == 2
+    _ok("统计入库结果")
+
+
+def test_save_grade_plan_db(biz_dt: str) -> None:
+    from agents.demand_grade_orchestrator_agent.tools.save_grade_plan import save_grade_plan
+    from agents.demand_grade_orchestrator_agent.common.assignment import resolve_planning_state
+
+    state = resolve_planning_state(biz_dt)
+    if not state["unassigned_category_ids"] or state["remaining_batch_quota"] <= 0:
+        print("  · save_grade_plan 跳过(无待分配或额度已满)")
+        return
+
+    cid = state["unassigned_category_ids"][0]
+    with patch(
+        "agents.demand_grade_orchestrator_agent.tools.save_grade_plan.persist_groups_one_by_one",
+        return_value={
+            "persisted_group_count": 1,
+            "persisted_groups": [{"category_ids": [cid]}],
+            "skipped_quota": 0,
+            "skipped_empty": 0,
+            "failed_groups": [],
+            "existing_groups": 1,
+            "remaining_batch_quota": MAX_DAILY_BATCHES - 1,
+            "unassigned_category_ids": state["unassigned_category_ids"][1:],
+            "coverage_complete": False,
+            "total_hanging_nodes": state["total_hanging_nodes"],
+        },
+    ):
+        result = json.loads(
+            save_grade_plan(
+                biz_dt,
+                "自检",
+                [{"category_ids": [cid, cid, 999999], "batch_heat_level": "Z"}],
+            )
+        )
+    assert result["ok"] is True and result["persisted"] is True
+    _ok("save_grade_plan 宽松校验 + 逐批入库路径")
+
+
+def main() -> None:
+    biz_dt = sys.argv[1] if len(sys.argv) > 1 else "20260714"
+    print("=== 统筹 Agent 逻辑自检 ===\n[单元测试]")
+    test_prepare_grade_groups_permissive()
+    test_dedupe_cross_group()
+    test_summarize_agent_saves()
+    print("\n[DB 集成检测]")
+    test_save_grade_plan_db(biz_dt)
+    print("\n全部通过。")
+
+
+if __name__ == "__main__":
+    main()

+ 1 - 1
agents/demand_grade_orchestrator_agent/agent.py

@@ -18,7 +18,7 @@ def create_demand_grade_orchestrator_agent(
         name="demand_grade_orchestrator_agent",
         model=model,
         system_prompt=_PROMPT_PATH.read_text(encoding="utf-8"),
-        max_iterations=24,
+        max_iterations=48,
     )
     register_all_tools(agent.tools)
     return agent

+ 4 - 2
agents/demand_grade_orchestrator_agent/common/__init__.py

@@ -1,8 +1,9 @@
 """统筹规划 Agent 的共享数据访问与格式化逻辑。"""
 
-from agents.demand_grade_orchestrator_agent.common.plan_builder import build_grade_plan_for_category_ids
+from agents.demand_grade_orchestrator_agent.common.plan_record import prepare_grade_groups
 from agents.demand_grade_orchestrator_agent.common.tree_state import (
     HEAT_LEVEL_DEFINITION,
+    format_demand_count,
     format_heat_score,
     format_rank,
     global_heat_positions,
@@ -13,8 +14,9 @@ from agents.demand_grade_orchestrator_agent.common.tree_state import (
 )
 
 __all__ = [
-    "build_grade_plan_for_category_ids",
+    "prepare_grade_groups",
     "HEAT_LEVEL_DEFINITION",
+    "format_demand_count",
     "format_heat_score",
     "format_rank",
     "global_heat_positions",

+ 123 - 0
agents/demand_grade_orchestrator_agent/common/assignment.py

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

+ 0 - 211
agents/demand_grade_orchestrator_agent/common/plan_builder.py

@@ -1,211 +0,0 @@
-"""按分类节点构建带批次热度等级的统筹分级计划。"""
-from __future__ import annotations
-
-from collections import Counter, defaultdict
-from typing import Any
-
-from agents.demand_grade_orchestrator_agent.common.tree_state import (
-    HEAT_LEVEL_DEFINITION,
-    format_heat_score,
-    format_rank,
-    global_heat_positions,
-    has_hung_demand,
-    heat_level,
-    load_tree_state,
-    path,
-)
-
-
-_LEVEL_ORDER = {level: index for index, level in enumerate(("S", "A", "B", "C", "D", "U"))}
-
-
-def _parent_id(category: Any) -> int | None:
-    return int(category.parent_id) if category.parent_id not in (None, 0) else None
-
-
-def _position_sort_key(
-    category_id: int,
-    positions: dict[int, dict[str, float | int]],
-) -> tuple[bool, float, int]:
-    position = positions.get(category_id)
-    return (
-        position is None,
-        float(position["rank"]) if position is not None else float("inf"),
-        category_id,
-    )
-
-
-def _connected_components(
-    target_ids: set[int],
-    by_id: dict[int, Any],
-    levels: dict[int, str],
-) -> list[list[int]]:
-    """把同等级且树上相邻/同父的目标节点聚成语义邻接分量。"""
-    edges: dict[int, set[int]] = defaultdict(set)
-    sibling_buckets: dict[tuple[int | None, str], list[int]] = defaultdict(list)
-    for category_id in target_ids:
-        parent_id = _parent_id(by_id[category_id])
-        sibling_buckets[(parent_id, levels[category_id])].append(category_id)
-        if parent_id in target_ids and levels[parent_id] == levels[category_id]:
-            edges[category_id].add(parent_id)
-            edges[parent_id].add(category_id)
-
-    for sibling_ids in sibling_buckets.values():
-        if len(sibling_ids) < 2:
-            continue
-        anchor = sibling_ids[0]
-        for sibling_id in sibling_ids[1:]:
-            edges[anchor].add(sibling_id)
-            edges[sibling_id].add(anchor)
-
-    components: list[list[int]] = []
-    unseen = set(target_ids)
-    while unseen:
-        start = min(unseen)
-        stack = [start]
-        unseen.remove(start)
-        component: list[int] = []
-        while stack:
-            current = stack.pop()
-            component.append(current)
-            for neighbor in edges.get(current, set()):
-                if neighbor in unseen:
-                    unseen.remove(neighbor)
-                    stack.append(neighbor)
-        components.append(component)
-    return components
-
-
-def _relation_summary(category_ids: list[int], by_id: dict[int, Any]) -> str:
-    parents = {_parent_id(by_id[category_id]) for category_id in category_ids}
-    has_parent_child = any(_parent_id(by_id[category_id]) in category_ids for category_id in category_ids)
-    if len(category_ids) == 1:
-        return "单节点批次(无同等级相邻目标节点)"
-    if len(parents) == 1:
-        parent_id = next(iter(parents))
-        return f"同父兄弟节点;共同父节点={path(parent_id, by_id) or '根层'}"
-    if has_parent_child:
-        return "树上直接相邻的父子/近邻节点"
-    return "同一邻接分类簇"
-
-
-def _node_position_payload(
-    category_id: int,
-    by_id: dict[int, Any],
-    weights: dict[int, Any],
-    positions: dict[int, dict[str, float | int]],
-) -> dict[str, Any]:
-    position = positions.get(category_id)
-    return {
-        "category_id": category_id,
-        "path": path(category_id, by_id),
-        "total_score": (
-            float(weights[category_id].total_score)
-            if weights.get(category_id) is not None and weights[category_id].total_score is not None
-            else None
-        ),
-        "global_rank": float(position["rank"]) if position is not None else None,
-        "global_scored_node_count": int(position["total"]) if position is not None else 0,
-        "global_rank_score": float(position["normalized_score"]) if position is not None else None,
-    }
-
-
-def build_grade_plan_for_category_ids(
-    biz_dt: str,
-    category_ids: list[int],
-    grouping_strategy: str,
-    *,
-    max_nodes_per_group: int = 4,
-) -> dict[str, Any]:
-    """生成批次计划:先分热度等级,再把同等级相邻节点合为一批。"""
-    requested = {int(category_id) for category_id in category_ids}
-    by_id, _children, weights = load_tree_state(biz_dt)
-    targets = {
-        category_id
-        for category_id in requested
-        if category_id in by_id and has_hung_demand(weights.get(category_id))
-    }
-    positions = global_heat_positions(weights)
-    levels = {category_id: heat_level(positions.get(category_id)) for category_id in targets}
-    width = max(1, min(int(max_nodes_per_group), 20))
-
-    drafts: list[dict[str, Any]] = []
-    for component in _connected_components(targets, by_id, levels):
-        ordered = sorted(component, key=lambda cid: _position_sort_key(cid, positions))
-        for start in range(0, len(ordered), width):
-            chunk = ordered[start : start + width]
-            level = levels[chunk[0]]
-            node_positions = [
-                _node_position_payload(category_id, by_id, weights, positions)
-                for category_id in chunk
-            ]
-            rank_scores = [
-                float(item["global_rank_score"])
-                for item in node_positions
-                if item["global_rank_score"] is not None
-            ]
-            raw_scores = [
-                float(item["total_score"])
-                for item in node_positions
-                if item["total_score"] is not None
-            ]
-            relation = _relation_summary(chunk, by_id)
-            score_text = "、".join(format_heat_score(weights.get(cid)) for cid in chunk)
-            rank_text = "、".join(format_rank(positions.get(cid)) for cid in chunk)
-            label = str(HEAT_LEVEL_DEFINITION[level]["label"])
-            drafts.append({
-                "category_ids": chunk,
-                "batch_heat_level": level,
-                "batch_heat_label": label,
-                "batch_global_rank_score": (
-                    sum(rank_scores) / len(rank_scores) if rank_scores else None
-                ),
-                "batch_total_score_avg": sum(raw_scores) / len(raw_scores) if raw_scores else None,
-                "category_global_positions": node_positions,
-                "planning_reason": (
-                    f"批次热度等级={level}({label});节点在整棵树有分节点中的名次="
-                    f"{rank_text},total_score={score_text}。{relation},因此统一成批。"
-                ),
-                "shared_traits": (
-                    f"{relation};同属批次热度等级 {level};"
-                    f"{grouping_strategy.strip() or '按整树热度等级与树邻接关系分批'}"
-                ),
-                "_sort_rank": min(
-                    (
-                        float(positions[cid]["rank"])
-                        for cid in chunk
-                        if cid in positions
-                    ),
-                    default=float("inf"),
-                ),
-            })
-
-    drafts.sort(key=lambda group: (
-        _LEVEL_ORDER[group["batch_heat_level"]],
-        group["_sort_rank"],
-        group["category_ids"][0],
-    ))
-    groups: list[dict[str, Any]] = []
-    for sequence_no, draft in enumerate(drafts, start=1):
-        level = draft["batch_heat_level"]
-        draft.pop("_sort_rank", None)
-        groups.append({
-            "group_id": f"{level}-batch-{sequence_no:03d}",
-            "sequence_no": sequence_no,
-            **draft,
-        })
-
-    covered = sorted({cid for group in groups for cid in group["category_ids"]})
-    uncovered = sorted(requested - set(covered))
-    level_counts = Counter(group["batch_heat_level"] for group in groups)
-    return {
-        "biz_dt": biz_dt,
-        "grouping_strategy": grouping_strategy.strip(),
-        "heat_level_definition": HEAT_LEVEL_DEFINITION,
-        "global_scored_node_count": len(positions),
-        "batch_heat_level_counts": dict(sorted(level_counts.items(), key=lambda item: _LEVEL_ORDER[item[0]])),
-        "groups": groups,
-        "covered_category_ids": covered,
-        "uncovered_category_ids": uncovered,
-        "coverage_complete": not uncovered,
-    }

+ 102 - 0
agents/demand_grade_orchestrator_agent/common/plan_persist.py

@@ -0,0 +1,102 @@
+"""批次计划逐条入库(由 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,
+    }

+ 198 - 0
agents/demand_grade_orchestrator_agent/common/plan_record.py

@@ -0,0 +1,198 @@
+"""将 Agent 提交的批次转为可入库结构(仅过滤无需求/已分配节点,不报错)。"""
+from __future__ import annotations
+
+import logging
+from collections import Counter
+from typing import Any
+
+from agents.demand_grade_orchestrator_agent.common.tree_state import (
+    HEAT_LEVEL_DEFINITION,
+    global_heat_positions,
+    has_hung_demand,
+    heat_level,
+    load_tree_state,
+    path,
+)
+
+logger = logging.getLogger(__name__)
+
+_LEVEL_ORDER = {level: index for index, level in enumerate(("S", "A", "B", "C", "D", "U"))}
+_VALID_LEVELS = frozenset(_LEVEL_ORDER)
+_DEFAULT_REASON = "Agent 统筹批次"
+_DEFAULT_TRAITS = "按全局树热度与结构分批"
+
+
+def _parse_category_ids(raw: Any) -> list[int]:
+    if not isinstance(raw, list):
+        return []
+    category_ids: list[int] = []
+    for value in raw:
+        try:
+            category_ids.append(int(value))
+        except (TypeError, ValueError):
+            continue
+    return list(dict.fromkeys(category_ids))
+
+
+def _normalize_batch_heat_level(raw: Any, category_ids: list[int], positions: dict[int, dict]) -> str:
+    level = str(raw or "").strip().upper()
+    if level in _VALID_LEVELS:
+        return level
+    for category_id in category_ids:
+        computed = heat_level(positions.get(category_id))
+        if computed in _VALID_LEVELS:
+            return computed
+    return "U"
+
+
+def _node_position_payload(
+    category_id: int,
+    by_id: dict[int, Any],
+    weights: dict[int, Any],
+    positions: dict[int, dict[str, float | int]],
+) -> dict[str, Any]:
+    position = positions.get(category_id)
+    return {
+        "category_id": category_id,
+        "path": path(category_id, by_id),
+        "total_score": (
+            float(weights[category_id].total_score)
+            if weights.get(category_id) is not None and weights[category_id].total_score is not None
+            else None
+        ),
+        "global_rank": float(position["rank"]) if position is not None else None,
+        "global_scored_node_count": int(position["total"]) if position is not None else 0,
+        "global_rank_score": float(position["normalized_score"]) if position is not None else None,
+    }
+
+
+def _enrich_group(
+    group: dict[str, Any],
+    *,
+    sequence_no: int,
+    by_id: dict[int, Any],
+    weights: dict[int, Any],
+    positions: dict[int, dict[str, float | int]],
+    grouping_strategy: str,
+) -> dict[str, Any]:
+    category_ids = group["category_ids"]
+    batch_heat_level = _normalize_batch_heat_level(group.get("batch_heat_level"), category_ids, positions)
+    node_positions = [
+        _node_position_payload(category_id, by_id, weights, positions)
+        for category_id in category_ids
+    ]
+    rank_scores = [
+        float(item["global_rank_score"])
+        for item in node_positions
+        if item.get("global_rank_score") is not None
+    ]
+    raw_scores = [
+        float(item["total_score"])
+        for item in node_positions
+        if item.get("total_score") is not None
+    ]
+    label = str(HEAT_LEVEL_DEFINITION.get(batch_heat_level, HEAT_LEVEL_DEFINITION["U"])["label"])
+    planning_reason = str(group.get("planning_reason") or "").strip() or _DEFAULT_REASON
+    shared_traits = str(group.get("shared_traits") or "").strip() or grouping_strategy.strip() or _DEFAULT_TRAITS
+    return {
+        "group_id": f"{batch_heat_level}-batch-{sequence_no:03d}",
+        "sequence_no": sequence_no,
+        "category_ids": category_ids,
+        "batch_heat_level": batch_heat_level,
+        "batch_heat_label": label,
+        "batch_global_rank_score": sum(rank_scores) / len(rank_scores) if rank_scores else None,
+        "batch_total_score_avg": sum(raw_scores) / len(raw_scores) if raw_scores else None,
+        "category_global_positions": node_positions,
+        "planning_reason": planning_reason,
+        "shared_traits": shared_traits,
+    }
+
+
+def prepare_grade_groups(
+    biz_dt: str,
+    grouping_strategy: str,
+    groups: list[dict[str, Any]],
+    *,
+    assigned_category_ids: set[int],
+) -> dict[str, Any]:
+    """过滤并补全批次;仅剔除无需求、不存在、已分配、重复节点。"""
+    by_id, _children, weights = load_tree_state(biz_dt)
+    positions = global_heat_positions(weights)
+    seen: set[int] = set()
+    filtered_category_ids: list[int] = []
+    filtered_duplicates: list[int] = []
+    failed_prepare_groups: list[dict[str, Any]] = []
+    draft_groups: list[dict[str, Any]] = []
+
+    for index, raw in enumerate(groups or [], start=1):
+        try:
+            kept_ids: list[int] = []
+            for category_id in _parse_category_ids(raw.get("category_ids")):
+                if category_id in seen:
+                    filtered_duplicates.append(category_id)
+                    continue
+                if category_id in assigned_category_ids:
+                    filtered_category_ids.append(category_id)
+                    continue
+                if category_id not in by_id:
+                    filtered_category_ids.append(category_id)
+                    continue
+                if not has_hung_demand(weights.get(category_id)):
+                    filtered_category_ids.append(category_id)
+                    continue
+                seen.add(category_id)
+                kept_ids.append(category_id)
+            if not kept_ids:
+                continue
+            draft_groups.append({
+                "category_ids": kept_ids,
+                "batch_heat_level": raw.get("batch_heat_level"),
+                "planning_reason": raw.get("planning_reason"),
+                "shared_traits": raw.get("shared_traits"),
+                "_sequence_no": index,
+            })
+        except Exception as exc:
+            logger.exception("准备第 %s 批失败,已跳过: biz_dt=%s", index, biz_dt)
+            failed_prepare_groups.append({"index": index, "error": str(exc)})
+
+    enriched_groups: list[dict[str, Any]] = []
+    for group in draft_groups:
+        sequence_no = int(group.get("_sequence_no", len(enriched_groups) + 1))
+        try:
+            enriched_groups.append(
+                _enrich_group(
+                    group,
+                    sequence_no=sequence_no,
+                    by_id=by_id,
+                    weights=weights,
+                    positions=positions,
+                    grouping_strategy=grouping_strategy,
+                )
+            )
+        except Exception as exc:
+            logger.exception(
+                "补全第 %s 批失败,已跳过: biz_dt=%s category_ids=%s",
+                sequence_no,
+                biz_dt,
+                group.get("category_ids"),
+            )
+            failed_prepare_groups.append({
+                "index": sequence_no,
+                "category_ids": group.get("category_ids"),
+                "error": str(exc),
+            })
+    level_counts = Counter(group["batch_heat_level"] for group in enriched_groups)
+    return {
+        "biz_dt": biz_dt,
+        "grouping_strategy": grouping_strategy.strip(),
+        "heat_level_definition": HEAT_LEVEL_DEFINITION,
+        "global_scored_node_count": len(positions),
+        "batch_heat_level_counts": dict(
+            sorted(level_counts.items(), key=lambda item: _LEVEL_ORDER[item[0]])
+        ),
+        "groups": enriched_groups,
+        "covered_category_ids": sorted(seen),
+        "filtered_category_ids": sorted(set(filtered_category_ids)),
+        "filtered_duplicate_category_ids": sorted(set(filtered_duplicates)),
+        "failed_prepare_groups": failed_prepare_groups,
+    }

+ 7 - 0
agents/demand_grade_orchestrator_agent/common/tree_state.py

@@ -43,6 +43,13 @@ def has_hung_demand(weight: Any | None) -> bool:
     return weight is not None and int(weight.hung_word_count or 0) > 0
 
 
+def format_demand_count(weight: Any | None) -> str:
+    """格式化节点挂载需求数量;无需求时返回空字符串。"""
+    if not has_hung_demand(weight):
+        return ""
+    return str(int(weight.hung_word_count or 0))
+
+
 HEAT_LEVEL_DEFINITION: dict[str, dict[str, str | float | None]] = {
     "S": {"label": "高热", "min_global_rank_score": 0.90},
     "A": {"label": "较高热", "min_global_rank_score": 0.75},

+ 24 - 11
agents/demand_grade_orchestrator_agent/prompt/system_prompt.md

@@ -1,19 +1,32 @@
 ## 角色与任务
 
-你是需求分级统筹规划 Agent。你不直接给需求定级,也不从需求清单中挑词;你必须先从全局分类树的热度和需求分布出发,确定本轮应处理的一个或多个树节点组,再把规划依据交给下游 `demand_grade_agent`
+你是需求分级统筹规划 Agent。你不直接给需求定级,也不从需求清单中挑词;你必须先从全局分类树的热度和需求分布出发,**自行判断**如何将有需求的分类节点划分成批次,再调用工具记录计划供下游执行
 
 ## 固定工作流
 
-1. 先调用 `query_global_heat_tree(biz_dt)`,查看有需求节点所在的全局树及其原始 `total_score` 热度分。不得跳过这一步。树中 `+` 仅表示该节点自身有需求,`null` 表示没有热度数据。
-2. 选择一个热点节点,或多个具有共同路径、相近热度等级或父节点背景的节点组成节点组;调用 `query_heat_node_group(biz_dt, category_ids)` 下钻验证,可多轮下钻。
-3. 总结该节点组的 `planning_reason` 与 `shared_traits`:必须说明其整树名次、批次热度等级、父/兄弟关系,以及为什么这些节点适合被同一批处理。
-4. 统筹的目标是覆盖当天**全部**有挂载需求的分类节点,不是只规划下一批。最后调用 `build_full_day_grade_plan(biz_dt, grouping_strategy, max_nodes_per_group)`,确认 `coverage_complete=true` 且 `uncovered_category_ids=[]`。其 JSON 是当天完整执行计划。
+1. 先调用 `query_global_heat_tree(biz_dt)`,查看**当天未分批**的有需求节点及其祖先路径、热度分与需求数量。不得跳过这一步。树中 `+` 与需求数仅对未分批节点展示;祖先节点仅提供路径上下文。
+2. 选择热点节点或节点组,调用 `query_heat_node_group(biz_dt, category_ids)` 下钻验证,可多轮下钻,直到理解各簇的父/兄弟关系与热度等级。`+` 同样仅标记未分批节点。
+3. **由你决定**如何分批:为每个批次确定 `category_ids`、`batch_heat_level`、`planning_reason`、`shared_traits`。必须说明整树名次、批次热度等级、父/兄弟关系,以及为何这些节点适合同批处理。
+4. 调用 `save_grade_plan(biz_dt, grouping_strategy, groups)` 提交并入库。**可多次调用**,每轮可提交多条 group;工具会逐批入库,仅过滤无需求、已分配节点,其余问题不阻断。根据返回的 `unassigned_category_ids` 与 `remaining_batch_quota` 继续提交,直到覆盖完成或达到每日上限。
+5. 工具返回 `persisted_group_count`、`filtered_category_ids`;即使部分批次跳过,只要 `persisted=true` 即表示有成功入库。
 
 ## 规划原则
 
-- 优先规划高热节点簇、同父节点下共同上升的兄弟节点,或“局部突发但大盘偏冷”的对照节点组。
-- 每个批次必须有明确的 `batch_heat_level`:S/A/B/C/D 来自节点在整棵有分分类树中的排名分,U 表示热度数据不足。批次等级用于区分热度,不要求下游严格按等级串行执行。
-- 相邻/相似节点只有在同一热度等级时才合批,避免一个极热节点把冷节点所在批次整体抬高。
-- 下游会逐组解析节点下的待分级需求;不要反过来根据需求名称拼凑批次。
-- 热度为空或样本数为 0 是数据不足,不是低热;在规划原因中明确说明。
-- 不自行评级、不调用落库工具。只依据工具真实返回的路径、节点和分数。
+- 统筹目标是尽量覆盖**全部未分批**分类节点。
+- 每天最多 200 个批次,优先保证效果好的批次保留。
+- 每个批次建议不超过 20 个节点。
+- `batch_heat_level`、`planning_reason`、`shared_traits` 缺失时工具会使用默认值补全。
+- 相邻/相似节点优先在同一热度等级时合批,避免一个极热节点把冷节点所在批次整体抬高。
+- 下游会逐组从数据库读取节点下的待分级需求;不要反过来根据需求名称拼凑批次。
+- 热度为空或样本数为 0 是数据不足,不是低热;在 `planning_reason` 中明确说明。
+- 不自行评级。只依据工具真实返回的路径、节点和分数。
+
+## groups 提交格式
+
+每个 group 为对象,包含:
+- `category_ids`: 整数列表
+- `batch_heat_level`: `S` / `A` / `B` / `C` / `D` / `U`
+- `planning_reason`: 本批规划原因(必填)
+- `shared_traits`: 本批节点共同特征(必填)
+
+`grouping_strategy` 写本轮提交的策略摘要(多轮提交时可各写各轮)。

+ 74 - 157
agents/demand_grade_orchestrator_agent/run.py

@@ -1,4 +1,4 @@
-"""运行统筹规划 Agent 并将计划落库。"""
+"""运行统筹规划 Agent(save_grade_plan 工具内直接入库)。"""
 from __future__ import annotations
 
 import json
@@ -6,15 +6,7 @@ import logging
 from typing import Any
 
 from agents.demand_grade_orchestrator_agent import create_demand_grade_orchestrator_agent
-from agents.demand_grade_orchestrator_agent.common import build_grade_plan_for_category_ids
-from agents.demand_grade_orchestrator_agent.validation import (
-    enrich_plan_assignment_summary,
-    resolve_assignment_state,
-    strip_assigned_category_ids,
-    strip_invalid_category_ids,
-    validate_full_day_grade_plan,
-    validate_plan_covers_targets,
-)
+from agents.demand_grade_orchestrator_agent.common.assignment import MAX_DAILY_BATCHES, resolve_planning_state
 from supply_agent.types import Role
 from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
 from supply_infra.db.session import get_session
@@ -22,189 +14,115 @@ from supply_infra.db.session import get_session
 logger = logging.getLogger(__name__)
 
 
-_MAX_PLAN_ATTEMPTS = 5
-
-
-def _extract_plan(result: Any) -> dict[str, Any] | None:
-    """从 Agent 的工具调用记录中提取最后一份完整计划。"""
-    for message in reversed(result.messages):
-        if message.role != Role.TOOL or message.name != "build_full_day_grade_plan":
+def _summarize_agent_saves(result: Any) -> dict[str, Any]:
+    """统计 Agent 会话中 save_grade_plan 成功入库的次数与批次数。"""
+    save_count = 0
+    persisted_groups = 0
+    last_payload: dict[str, Any] | None = None
+    for message in result.messages:
+        if message.role != Role.TOOL or message.name != "save_grade_plan":
             continue
         try:
             payload = json.loads(message.content or "")
         except json.JSONDecodeError:
             continue
-        if isinstance(payload, dict) and payload.get("groups") is not None:
-            return payload
-    return None
-
-
-def _persist_plan(biz_dt: str, payload: dict[str, Any]) -> None:
-    if not payload.get("groups"):
-        logger.info("无新增节点组需要落库: biz_dt=%s", biz_dt)
-        return
-    with get_session() as session:
-        DemandGradePlanRepository(session).create_plan(biz_dt, payload)
-
-
-def _sanitize_plan(
-    biz_dt: str,
-    plan: dict[str, Any] | None,
-    *,
-    assigned_ids: set[int],
-) -> dict[str, Any] | None:
-    if plan is None:
-        return None
-    removed_invalid = strip_invalid_category_ids(biz_dt, plan)
-    if removed_invalid:
-        logger.error(
-            "统筹计划包含无效分类 ID,已自动移除: biz_dt=%s removed=%s",
-            biz_dt,
-            removed_invalid,
-        )
-    removed_duplicates = strip_assigned_category_ids(plan, assigned_ids)
-    if removed_duplicates:
-        logger.warning(
-            "统筹计划包含当天已分配分类,已自动去重: biz_dt=%s removed=%s",
-            biz_dt,
-            removed_duplicates,
-        )
-    return plan if plan.get("groups") else None
+        if not isinstance(payload, dict):
+            continue
+        if payload.get("ok") is True and payload.get("persisted") is True:
+            save_count += 1
+            persisted_groups += int(payload.get("persisted_group_count") or 0)
+            last_payload = payload
+    return {
+        "save_count": save_count,
+        "persisted_groups": persisted_groups,
+        "last_payload": last_payload,
+    }
 
 
 def _orchestrate_via_agent(
     biz_dt: str,
     *,
-    max_nodes_per_group: int,
-    assignment_state: dict[str, Any],
-    target_ids: set[int],
+    planning_state: dict[str, Any],
+    unassigned_count: int,
 ) -> None:
-    assigned_ids = set(assignment_state["assigned_category_ids"])
-    feedback = ""
-    last_validation: dict[str, Any] | None = None
-    last_payload: dict[str, Any] | None = None
-    target_hint = (
-        "请覆盖当天全部有需求分类节点。"
-        if target_ids == set(assignment_state["required_category_ids"])
-        else f"只需覆盖以下未分配分类 ID:{sorted(target_ids)}。"
-    )
+    remaining_quota = int(planning_state["remaining_batch_quota"])
 
-    for attempt in range(1, _MAX_PLAN_ATTEMPTS + 1):
-        agent = create_demand_grade_orchestrator_agent()
+    agent = create_demand_grade_orchestrator_agent()
+    try:
         result = agent.run(
             f"""请为业务日 {biz_dt} 制定全局树需求分级计划。
 必须从 query_global_heat_tree 开始,经过至少一次 query_heat_node_group 下钻,
-最后调用 build_full_day_grade_plan,并将每组最多节点数设为 {max_nodes_per_group}。
-{target_hint}
-{feedback}
+由你自行划分批次后调用 save_grade_plan 提交(可多次调用,每轮提交一部分批次)。
+
+待分批节点数:{unassigned_count}(请通过 query_global_heat_tree 查看,勿依赖本消息枚举 ID)。
+剩余可新增批次数:{remaining_quota}(每日总上限 {MAX_DAILY_BATCHES},当前已存在 {planning_state["existing_groups"]} 个批次)。
+每批最多 20 个节点;节点较多时分多轮 save_grade_plan,每轮关注工具返回的 unassigned_category_ids 与 remaining_batch_quota。
 """
         )
-        payload = _sanitize_plan(biz_dt, _extract_plan(result), assigned_ids=assigned_ids)
-        validation = (
-            validate_full_day_grade_plan(biz_dt, payload)
-            if target_ids == set(assignment_state["required_category_ids"])
-            else validate_plan_covers_targets(payload, target_ids)
-        )
-        last_validation = validation
-        if payload is not None:
-            last_payload = payload
-        if payload is not None and validation["coverage_complete"]:
-            enrich_plan_assignment_summary(biz_dt, payload, assignment_state=assignment_state)
-            _persist_plan(biz_dt, payload)
-            return
-
-        missing = validation["uncovered_category_ids"]
-        feedback = (
-            "上一次计划未通过独立校验。请重新查询并重新调用 "
-            "build_full_day_grade_plan;"
-            f"必须补充分配未覆盖分类 ID:{missing or '无'}。"
-        )
-
-    uncovered = (last_validation or {}).get("uncovered_category_ids", [])
-    logger.error(
-        "统筹规划 Agent 在最多 %d 次尝试后仍未覆盖目标分类: biz_dt=%s uncovered=%s",
-        _MAX_PLAN_ATTEMPTS,
-        biz_dt,
-        uncovered,
-    )
-    if last_payload is not None:
-        enrich_plan_assignment_summary(biz_dt, last_payload, assignment_state=assignment_state)
-        _persist_plan(biz_dt, last_payload)
-
+    except Exception:
+        logger.exception("统筹 Agent 运行异常: biz_dt=%s", biz_dt)
+        return
 
-def _orchestrate_incremental(
-    biz_dt: str,
-    *,
-    max_nodes_per_group: int,
-    assignment_state: dict[str, Any],
-    unassigned_ids: list[int],
-) -> None:
-    payload = build_grade_plan_for_category_ids(
-        biz_dt,
-        unassigned_ids,
-        "补充分配当天未覆盖节点",
-        max_nodes_per_group=max_nodes_per_group,
-    )
-    payload = _sanitize_plan(
+    summary = _summarize_agent_saves(result)
+    if summary["save_count"] == 0:
+        logger.warning("统筹 Agent 未成功入库任何批次: biz_dt=%s", biz_dt)
+        return
+    logger.info(
+        "统筹 Agent 完成: biz_dt=%s save_count=%s persisted_groups=%s",
         biz_dt,
-        payload,
-        assigned_ids=set(assignment_state["assigned_category_ids"]),
+        summary["save_count"],
+        summary["persisted_groups"],
     )
-    if payload is None:
-        logger.error("未能为未分配节点生成分组: biz_dt=%s unassigned=%s", biz_dt, unassigned_ids)
-        return
-
-    validation = validate_plan_covers_targets(payload, set(unassigned_ids))
-    if not validation["coverage_complete"]:
-        logger.error(
-            "增量统筹计划未覆盖全部未分配节点: biz_dt=%s uncovered=%s",
+    last = summary["last_payload"] or {}
+    if last.get("unassigned_category_ids"):
+        logger.info(
+            "统筹后仍有未分批节点: biz_dt=%s remaining=%s quota=%s",
             biz_dt,
-            validation["uncovered_category_ids"],
+            len(last["unassigned_category_ids"]),
+            last.get("remaining_batch_quota"),
         )
-    enrich_plan_assignment_summary(biz_dt, payload, assignment_state=assignment_state)
-    _persist_plan(biz_dt, payload)
 
 
-def orchestrate_daily_grade_plan(*, biz_dt: str, max_nodes_per_group: int = 4) -> None:
-    """运行统筹规划并将新增节点组写入数据库。"""
-    assignment_state = resolve_assignment_state(biz_dt)
-    if assignment_state["assignment_complete"]:
+def orchestrate_daily_grade_plan(*, biz_dt: str) -> None:
+    """运行统筹规划(批次由 Agent 划分,save_grade_plan 工具入库)。"""
+    planning_state = resolve_planning_state(biz_dt)
+
+    if planning_state["batch_limit_reached"]:
         logger.info(
-            "当天全部有需求节点已分配,跳过统筹: biz_dt=%s assigned=%s",
+            "跳过统筹 Agent:当天批次已达上限 biz_dt=%s existing_groups=%s limit=%s",
             biz_dt,
-            assignment_state["assigned_category_ids"],
+            planning_state["existing_groups"],
+            MAX_DAILY_BATCHES,
         )
         return
 
-    unassigned_ids = assignment_state["unassigned_category_ids"]
-    if not assignment_state["assigned_category_ids"]:
-        logger.info("当天尚无已分配节点,执行全量统筹: biz_dt=%s", biz_dt)
-        _orchestrate_via_agent(
+    unassigned_ids = planning_state["unassigned_category_ids"]
+    if not unassigned_ids:
+        logger.info(
+            "跳过统筹 Agent:当天有需求节点均已分批 biz_dt=%s total_hanging_nodes=%s",
             biz_dt,
-            max_nodes_per_group=max_nodes_per_group,
-            assignment_state=assignment_state,
-            target_ids=set(assignment_state["required_category_ids"]),
+            planning_state["total_hanging_nodes"],
         )
         return
 
     logger.info(
-        "当天存在已分配节点,仅处理未分配节点: biz_dt=%s unassigned=%s",
+        "执行统筹 Agent: biz_dt=%s unassigned=%s remaining_quota=%s",
         biz_dt,
-        unassigned_ids,
+        len(unassigned_ids),
+        planning_state["remaining_batch_quota"],
     )
-    _orchestrate_incremental(
+    _orchestrate_via_agent(
         biz_dt,
-        max_nodes_per_group=max_nodes_per_group,
-        assignment_state=assignment_state,
-        unassigned_ids=unassigned_ids,
+        planning_state=planning_state,
+        unassigned_count=len(unassigned_ids),
     )
 
 
-def main(biz_dt: str, *, max_nodes_per_group: int = 4) -> dict[str, Any]:
+def main(biz_dt: str) -> dict[str, Any]:
     """手动测试:运行统筹规划并打印当天分配摘要。"""
-    assignment_before = resolve_assignment_state(biz_dt)
-    orchestrate_daily_grade_plan(biz_dt=biz_dt, max_nodes_per_group=max_nodes_per_group)
-    assignment_after = resolve_assignment_state(biz_dt)
+    planning_before = resolve_planning_state(biz_dt)
+    orchestrate_daily_grade_plan(biz_dt=biz_dt)
+    planning_after = resolve_planning_state(biz_dt)
     with get_session() as session:
         repo = DemandGradePlanRepository(session)
         plan = repo.get_latest_plan(biz_dt)
@@ -213,14 +131,13 @@ def main(biz_dt: str, *, max_nodes_per_group: int = 4) -> dict[str, Any]:
         plan_summary = json.loads(plan.plan_json) if plan is not None else {}
         result = {
             "biz_dt": biz_dt,
-            "assignment_before": assignment_before,
-            "assignment_after": assignment_after,
+            "planning_before": planning_before,
+            "planning_after": planning_after,
             "plan_count": 1 if plan is not None else 0,
-            "coverage_complete": assignment_after["assignment_complete"],
-            "total_hanging_nodes": assignment_after["total_hanging_nodes"],
+            "total_hanging_nodes": planning_after["total_hanging_nodes"],
             "group_count": len(groups),
             "group_status": group_status,
-            "uncovered_category_ids": assignment_after["unassigned_category_ids"],
+            "unassigned_category_ids": planning_after["unassigned_category_ids"],
             "sample_groups": [
                 {
                     "group_key": group.group_key,

+ 2 - 2
agents/demand_grade_orchestrator_agent/tools/__init__.py

@@ -4,12 +4,12 @@ from __future__ import annotations
 from collections.abc import Callable
 from typing import Any
 
-from agents.demand_grade_orchestrator_agent.tools.build_full_day_grade_plan import build_full_day_grade_plan
 from agents.demand_grade_orchestrator_agent.tools.query_global_heat_tree import query_global_heat_tree
 from agents.demand_grade_orchestrator_agent.tools.query_heat_node_group import query_heat_node_group
+from agents.demand_grade_orchestrator_agent.tools.save_grade_plan import save_grade_plan
 from supply_agent.tools.registry import ToolRegistry
 
-ALL_TOOLS: list[Callable[..., Any]] = [query_global_heat_tree, query_heat_node_group, build_full_day_grade_plan]
+ALL_TOOLS: list[Callable[..., Any]] = [query_global_heat_tree, query_heat_node_group, save_grade_plan]
 
 
 def register_all_tools(registry: ToolRegistry) -> ToolRegistry:

+ 0 - 20
agents/demand_grade_orchestrator_agent/tools/build_full_day_grade_plan.py

@@ -1,20 +0,0 @@
-"""构建覆盖全部有需求节点的紧凑日计划。"""
-from __future__ import annotations
-
-import json
-
-from agents.demand_grade_orchestrator_agent.common import build_grade_plan_for_category_ids
-from agents.demand_grade_orchestrator_agent.validation import get_required_hanging_category_ids
-from supply_agent.tools import tool
-
-
-@tool
-def build_full_day_grade_plan(biz_dt: str, grouping_strategy: str, max_nodes_per_group: int = 4) -> str:
-    """按整树热度等级和树邻接关系生成全天节点组;每组明确标注 S/A/B/C/D/U。"""
-    payload = build_grade_plan_for_category_ids(
-        biz_dt,
-        sorted(get_required_hanging_category_ids(biz_dt)),
-        grouping_strategy,
-        max_nodes_per_group=max_nodes_per_group,
-    )
-    return json.dumps(payload, ensure_ascii=False)

+ 20 - 9
agents/demand_grade_orchestrator_agent/tools/query_global_heat_tree.py

@@ -1,7 +1,9 @@
-"""以层级文本展示当天有需求分类所在的全局树及原始热度分。"""
+"""以层级文本展示当天未分批分类节点所在的全局树及原始热度分。"""
 from __future__ import annotations
 
+from agents.demand_grade_orchestrator_agent.common.assignment import get_unassigned_hanging_category_ids
 from agents.demand_grade_orchestrator_agent.common import (
+    format_demand_count,
     format_heat_score,
     format_rank,
     global_heat_positions,
@@ -14,22 +16,25 @@ from supply_agent.tools import tool
 
 @tool
 def query_global_heat_tree(biz_dt: str) -> str:
-    """返回有需求节点及其祖先路径;“+”仅表示当前分类自身挂有需求。"""
+    """返回未分批节点及其祖先路径;+ 与需求数仅展示未分批节点自身。"""
     by_id, children, weights = load_tree_state(biz_dt)
     positions = global_heat_positions(weights)
-    demand_nodes = {cid for cid, weight in weights.items() if has_hung_demand(weight)}
+    unassigned_nodes = get_unassigned_hanging_category_ids(biz_dt)
     visible_cache: dict[int, bool] = {}
 
     def visible(category_id: int) -> bool:
         if category_id in visible_cache:
             return visible_cache[category_id]
-        value = category_id in demand_nodes or any(visible(child) for child in children.get(category_id, []))
+        value = category_id in unassigned_nodes or any(
+            visible(child) for child in children.get(category_id, [])
+        )
         visible_cache[category_id] = value
         return value
 
     lines = [
-        f"biz_dt={biz_dt} | 格式:[分类ID]分类名称[total_score|整树名次|热度等级] + | "
-        "名次在整棵有分节点中计算;无数据为 null/U;+ 表示该分类自身有需求"
+        f"biz_dt={biz_dt} | 格式:[分类ID]分类名称[total_score|整树名次|热度等级|需求数] + | "
+        "仅展示未分批节点及其祖先路径;+ 与需求数仅对未分批且有挂载需求的节点展示;"
+        "名次在整棵有分节点中计算;无数据为 null/U"
     ]
 
     def render(category_id: int, depth: int) -> None:
@@ -39,17 +44,23 @@ def query_global_heat_tree(biz_dt: str) -> str:
         weight = weights.get(category_id)
         score_text = format_heat_score(weight)
         position = positions.get(category_id)
-        suffix = " +" if category_id in demand_nodes else ""
+        is_unassigned = category_id in unassigned_nodes
+        demand_count = format_demand_count(weight) if is_unassigned and has_hung_demand(weight) else ""
+        demand_suffix = f"|{demand_count}" if demand_count else ""
+        suffix = " +" if is_unassigned and has_hung_demand(weight) else ""
         lines.append(
             f"{'  ' * depth}[{category_id}]{category.name or ''}"
-            f"[{score_text}|{format_rank(position)}|{heat_level(position)}]{suffix}"
+            f"[{score_text}|{format_rank(position)}|{heat_level(position)}{demand_suffix}]{suffix}"
         )
         for child in children.get(category_id, []):
             render(child, depth + 1)
 
     for root_id in children.get(None, []):
         render(root_id, 0)
-    return "\n".join(lines) if len(lines) > 1 else f"biz_dt={biz_dt} 无有需求的分类节点"
+    if len(lines) == 1:
+        return f"biz_dt={biz_dt} 无未分批的有需求分类节点"
+    lines.append(f"未分批节点数={len(unassigned_nodes)}")
+    return "\n".join(lines)
 
 
 if __name__ == '__main__':

+ 6 - 7
agents/demand_grade_orchestrator_agent/tools/query_heat_node_group.py

@@ -1,6 +1,7 @@
-"""下钻节点,查看整树位置及父子/兄弟间的全局热度。"""
+"""下钻节点,查看整树位置及父子/兄弟间的全局热度(未分批节点标 +)。"""
 from __future__ import annotations
 
+from agents.demand_grade_orchestrator_agent.common.assignment import get_unassigned_hanging_category_ids
 from agents.demand_grade_orchestrator_agent.common import (
     format_heat_score,
     format_rank,
@@ -18,16 +19,18 @@ def query_heat_node_group(biz_dt: str, category_ids: list[int]) -> str:
     """返回指定节点、父节点、全部兄弟和直接子节点的整树名次及热度等级。"""
     by_id, children, weights = load_tree_state(biz_dt)
     positions = global_heat_positions(weights)
+    unassigned_nodes = get_unassigned_hanging_category_ids(biz_dt)
     lines = [
         f"biz_dt={biz_dt} | 格式:[分类ID]分类名称[total_score|整树名次|热度等级] + | "
-        "名次在整棵有分节点中计算;无数据为 null/U;+ 表示该分类自身有需求"
+        "名次在整棵有分节点中计算;无数据为 null/U;+ 仅表示未分批且有挂载需求"
     ]
 
     def describe(category_id: int) -> str:
         category = by_id[category_id]
         weight = weights.get(category_id)
         position = positions.get(category_id)
-        suffix = " +" if has_hung_demand(weight) else ""
+        is_unassigned = category_id in unassigned_nodes
+        suffix = " +" if is_unassigned and has_hung_demand(weight) else ""
         return (
             f"[{category_id}]{category.name or ''}"
             f"[{format_heat_score(weight)}|{format_rank(position)}|{heat_level(position)}]{suffix}"
@@ -47,7 +50,3 @@ def query_heat_node_group(biz_dt: str, category_ids: list[int]) -> str:
         child_text = "、".join(describe(child) for child in children.get(category_id, []))
         lines.append(f"  直接子节点:{child_text or '无'}")
     return "\n".join(lines)
-
-if __name__ == '__main__':
-    res = query_heat_node_group('20260714',[313,686])
-    print(res)

+ 132 - 0
agents/demand_grade_orchestrator_agent/tools/save_grade_plan.py

@@ -0,0 +1,132 @@
+"""记录 Agent 制定的批次计划并入库。"""
+from __future__ import annotations
+
+import json
+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.plan_persist import persist_groups_one_by_one
+from agents.demand_grade_orchestrator_agent.common.plan_record import prepare_grade_groups
+from supply_agent.tools import tool
+
+logger = logging.getLogger(__name__)
+
+
+@tool
+def save_grade_plan(
+    biz_dt: str,
+    grouping_strategy: str,
+    groups: list[dict[str, Any]],
+) -> str:
+    """提交你自行划分的批次计划并入库。可多次调用,每轮可提交多条 group。
+
+    每个 group 需包含:
+    - category_ids: 本批分类 ID 列表
+    - batch_heat_level: S/A/B/C/D/U
+    - planning_reason: 本批规划原因
+    - shared_traits: 本批节点共同特征
+    """
+    existing_groups = get_existing_group_count(biz_dt)
+    remaining_quota = max(0, MAX_DAILY_BATCHES - existing_groups)
+    required = get_required_hanging_category_ids(biz_dt)
+    assigned = get_assigned_category_ids(biz_dt)
+    unassigned = get_unassigned_hanging_category_ids(biz_dt)
+
+    base_response: dict[str, Any] = {
+        "biz_dt": biz_dt,
+        "existing_groups": existing_groups,
+        "remaining_batch_quota": remaining_quota,
+        "unassigned_category_ids": sorted(unassigned),
+        "total_hanging_nodes": len(required),
+    }
+
+    if remaining_quota <= 0:
+        return json.dumps({
+            **base_response,
+            "ok": True,
+            "persisted": False,
+            "message": f"当天批次已达上限 {MAX_DAILY_BATCHES},本批未入库。",
+            "groups": [],
+        }, ensure_ascii=False)
+
+    if not groups:
+        return json.dumps({
+            **base_response,
+            "ok": True,
+            "persisted": False,
+            "message": "groups 为空,未入库。",
+            "groups": [],
+        }, ensure_ascii=False)
+
+    try:
+        prepared = prepare_grade_groups(
+            biz_dt,
+            grouping_strategy,
+            groups,
+            assigned_category_ids=assigned,
+        )
+    except Exception:
+        logger.exception("准备批次计划失败: biz_dt=%s", biz_dt)
+        return json.dumps({
+            **base_response,
+            "ok": True,
+            "persisted": False,
+            "message": "准备批次时发生异常,未入库。",
+            "groups": [],
+        }, ensure_ascii=False)
+
+    if not prepared.get("groups"):
+        return json.dumps({
+            **base_response,
+            "ok": True,
+            "persisted": False,
+            "message": "过滤后无有效节点可入库。",
+            "groups": [],
+            "filtered_category_ids": prepared.get("filtered_category_ids", []),
+            "filtered_duplicate_category_ids": prepared.get("filtered_duplicate_category_ids", []),
+            "failed_prepare_groups": prepared.get("failed_prepare_groups", []),
+        }, ensure_ascii=False)
+
+    persist_result = persist_groups_one_by_one(biz_dt, prepared)
+
+    persisted = persist_result["persisted_group_count"] > 0
+    failed_count = len(persist_result["failed_groups"])
+    if persisted:
+        message = f"已入库 {persist_result['persisted_group_count']} 批"
+        if failed_count:
+            message += f",失败 {failed_count} 批"
+        message += "。"
+    else:
+        message = "本批无成功入库记录(可能均已分配、额度已满或全部失败)。"
+    return json.dumps({
+        "ok": True,
+        "persisted": persisted,
+        "biz_dt": biz_dt,
+        "grouping_strategy": prepared.get("grouping_strategy"),
+        "groups": persist_result.get("persisted_groups", []),
+        "persisted_group_count": persist_result["persisted_group_count"],
+        "skipped_quota": persist_result["skipped_quota"],
+        "skipped_empty": persist_result["skipped_empty"],
+        "failed_groups": persist_result["failed_groups"],
+        "failed_prepare_groups": prepared.get("failed_prepare_groups", []),
+        "filtered_category_ids": prepared.get("filtered_category_ids", []),
+        "filtered_duplicate_category_ids": prepared.get("filtered_duplicate_category_ids", []),
+        "covered_category_ids": [
+            cid
+            for group in persist_result.get("persisted_groups", [])
+            for cid in group.get("category_ids", [])
+        ],
+        "existing_groups": persist_result["existing_groups"],
+        "remaining_batch_quota": persist_result["remaining_batch_quota"],
+        "unassigned_category_ids": persist_result["unassigned_category_ids"],
+        "coverage_complete": persist_result["coverage_complete"],
+        "total_hanging_nodes": persist_result["total_hanging_nodes"],
+        "message": message,
+    }, ensure_ascii=False)

+ 0 - 25
agents/demand_grade_orchestrator_agent/validation/__init__.py

@@ -1,25 +0,0 @@
-"""统筹规划 Agent 的校验逻辑。"""
-
-from agents.demand_grade_orchestrator_agent.validation.assignment import (
-    enrich_plan_assignment_summary,
-    get_assigned_category_ids,
-    get_required_hanging_category_ids,
-    resolve_assignment_state,
-    strip_assigned_category_ids,
-)
-from agents.demand_grade_orchestrator_agent.validation.full_day_grade_plan import (
-    strip_invalid_category_ids,
-    validate_full_day_grade_plan,
-    validate_plan_covers_targets,
-)
-
-__all__ = [
-    "enrich_plan_assignment_summary",
-    "get_assigned_category_ids",
-    "get_required_hanging_category_ids",
-    "resolve_assignment_state",
-    "strip_assigned_category_ids",
-    "strip_invalid_category_ids",
-    "validate_full_day_grade_plan",
-    "validate_plan_covers_targets",
-]

+ 0 - 85
agents/demand_grade_orchestrator_agent/validation/assignment.py

@@ -1,85 +0,0 @@
-"""当天节点分配状态校验。"""
-from __future__ import annotations
-
-from typing import Any
-
-from agents.demand_grade_orchestrator_agent.common 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
-
-
-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 resolve_assignment_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)
-    return {
-        "biz_dt": biz_dt,
-        "total_hanging_nodes": len(required),
-        "required_category_ids": sorted(required),
-        "assigned_category_ids": sorted(assigned),
-        "unassigned_category_ids": unassigned,
-        "assignment_complete": not unassigned,
-    }
-
-
-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({**group, "category_ids": kept})
-    plan["groups"] = cleaned_groups
-    return sorted(set(removed))
-
-
-def enrich_plan_assignment_summary(
-    biz_dt: str,
-    plan: dict[str, Any],
-    *,
-    assignment_state: dict[str, Any] | None = None,
-) -> dict[str, Any]:
-    """根据当天全量分配情况补充计划的覆盖摘要。"""
-    state = assignment_state or resolve_assignment_state(biz_dt)
-    required = set(state["required_category_ids"])
-    already_assigned = set(state["assigned_category_ids"])
-    new_assigned = {
-        int(category_id)
-        for group in plan.get("groups") or []
-        for category_id in group.get("category_ids") or []
-    }
-    all_assigned = already_assigned | new_assigned
-    uncovered = sorted(required - all_assigned)
-    plan.update({
-        "biz_dt": biz_dt,
-        "total_hanging_nodes": len(required),
-        "covered_category_ids": sorted(required & all_assigned),
-        "uncovered_category_ids": uncovered,
-        "coverage_complete": not uncovered,
-    })
-    return plan

+ 0 - 87
agents/demand_grade_orchestrator_agent/validation/full_day_grade_plan.py

@@ -1,87 +0,0 @@
-"""校验单日统筹计划是否覆盖全部挂载需求分类。"""
-from __future__ import annotations
-
-from typing import Any
-
-from agents.demand_grade_orchestrator_agent.common import load_tree_state
-from agents.demand_grade_orchestrator_agent.validation.assignment import get_required_hanging_category_ids
-
-
-def strip_invalid_category_ids(biz_dt: str, plan: dict[str, Any]) -> list[int]:
-    """从计划中移除不存在于全局树的分类 ID,并返回被移除的 ID 列表。"""
-    by_id, _, _ = load_tree_state(biz_dt)
-    valid_ids = set(by_id)
-    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 valid_ids:
-                kept.append(cid)
-            else:
-                removed.append(cid)
-        if kept:
-            cleaned_groups.append({**group, "category_ids": kept})
-    plan["groups"] = cleaned_groups
-    return sorted(set(removed))
-
-
-def validate_full_day_grade_plan(
-    biz_dt: str, plan: dict[str, Any] | None
-) -> dict[str, Any]:
-    """重新查询当天需求节点,并校验计划中的节点分配是否完整。
-
-    校验以数据库中的 ``hung_word_count`` 为准,而不是信任 Agent 返回的
-    ``coverage_complete`` 字段。计划组中的重复节点会被去重。
-
-    调用方应先用 ``strip_invalid_category_ids`` 清理无效分类,本函数只判断
-    是否覆盖全部有需求节点。
-    """
-    required_ids = get_required_hanging_category_ids(biz_dt)
-
-    assigned_ids: set[int] = set()
-    for group in (plan or {}).get("groups") or []:
-        for category_id in group.get("category_ids") or []:
-            try:
-                assigned_ids.add(int(category_id))
-            except (TypeError, ValueError):
-                continue
-
-    covered_ids = sorted(required_ids & assigned_ids)
-    uncovered_ids = sorted(required_ids - assigned_ids)
-    return {
-        "biz_dt": biz_dt,
-        "total_hanging_nodes": len(required_ids),
-        "covered_category_ids": covered_ids,
-        "uncovered_category_ids": uncovered_ids,
-        "coverage_complete": not uncovered_ids,
-    }
-
-
-def validate_plan_covers_targets(plan: dict[str, Any] | None, target_ids: set[int]) -> dict[str, Any]:
-    """校验计划是否覆盖指定目标分类集合。"""
-    assigned_ids: set[int] = set()
-    for group in (plan or {}).get("groups") or []:
-        for category_id in group.get("category_ids") or []:
-            try:
-                assigned_ids.add(int(category_id))
-            except (TypeError, ValueError):
-                continue
-
-    covered_ids = sorted(target_ids & assigned_ids)
-    uncovered_ids = sorted(target_ids - assigned_ids)
-    return {
-        "covered_category_ids": covered_ids,
-        "uncovered_category_ids": uncovered_ids,
-        "coverage_complete": not uncovered_ids,
-    }
-
-
-if __name__ == "__main__":
-    import json
-
-    print(json.dumps(validate_full_day_grade_plan("20260714", None), ensure_ascii=False))

+ 8 - 8
jobs/grade_demand_pool.py

@@ -4,10 +4,9 @@
 统筹规划 Agent 先落库当天全量节点组计划,再由多个 worker 领取任务并调用分级 Agent。
 
 用法:
-    python jobs/grade_demand_pool.py                      # 默认业务日、批次20、5 个 worker
-    python jobs/grade_demand_pool.py 20260716              # 指定业务日
-    python jobs/grade_demand_pool.py 20260716 20           # 指定业务日 + 批次大小
-    python jobs/grade_demand_pool.py 20260716 20 5          # 再指定 5 个并发 worker
+    python jobs/grade_demand_pool.py                  # 默认业务日、5 个 worker
+    python jobs/grade_demand_pool.py 20260716         # 指定业务日
+    python jobs/grade_demand_pool.py 20260716 5       # 指定业务日 + 5 个并发 worker
 """
 from __future__ import annotations
 
@@ -24,13 +23,15 @@ logging.basicConfig(
 
 def main(
     biz_dt: str | None = None,
-    batch_size_arg: str | None = None,
     workers_arg: str | None = None,
 ) -> dict:
-    batch_size = int(batch_size_arg) if batch_size_arg else 20
     workers = int(workers_arg) if workers_arg else 5
 
-    result = grade_demand_pool(biz_dt, batch_size=batch_size, workers=workers)
+    result = grade_demand_pool(
+        biz_dt,
+        workers=workers,
+        with_orchestrate=True,
+    )
     print(result)
     return result
 
@@ -39,5 +40,4 @@ if __name__ == "__main__":
     main(
         sys.argv[1] if len(sys.argv) > 1 else None,
         sys.argv[2] if len(sys.argv) > 2 else None,
-        sys.argv[3] if len(sys.argv) > 3 else None,
     )

+ 85 - 0
scripts/run_grade_plan_groups.py

@@ -0,0 +1,85 @@
+#!/usr/bin/env python3
+"""批量执行 demand_grade_plan_group 分级任务。
+
+与定时任务共用 supply_infra.scheduler.jobs.grade_demand_pool.grade_demand_pool。
+
+Usage:
+  .venv/bin/python scripts/run_grade_plan_groups.py
+  .venv/bin/python scripts/run_grade_plan_groups.py --biz-dt 20260721
+  .venv/bin/python scripts/run_grade_plan_groups.py --biz-dt 20260721 --workers 5
+  .venv/bin/python scripts/run_grade_plan_groups.py --biz-dt 20260721 --with-orchestrate
+"""
+from __future__ import annotations
+
+import argparse
+import json
+import logging
+import sys
+from pathlib import Path
+
+_ROOT = Path(__file__).resolve().parents[1]
+if str(_ROOT) not in sys.path:
+    sys.path.insert(0, str(_ROOT))
+
+from supply_infra.scheduler.plan_group_batch import MAX_DEMANDS_PER_BATCH
+from supply_infra.scheduler.jobs.grade_demand_pool import grade_demand_pool
+
+logger = logging.getLogger(__name__)
+
+
+def main(argv: list[str] | None = None) -> int:
+    parser = argparse.ArgumentParser(description="批量执行 demand_grade_plan_group 分级任务")
+    parser.add_argument("--biz-dt", default="20260721", help="业务日期 YYYYMMDD,默认 20260721")
+    parser.add_argument(
+        "--max-demands-per-batch",
+        type=int,
+        default=MAX_DEMANDS_PER_BATCH,
+        help=f"每个 Agent 子批次最多处理的需求条数,默认 {MAX_DEMANDS_PER_BATCH}",
+    )
+    parser.add_argument("--workers", type=int, default=5, help="并发执行的 plan_group 数")
+    parser.add_argument(
+        "--max-rounds",
+        type=int,
+        default=0,
+        help="最多执行轮数,0 表示直到没有 pending 任务",
+    )
+    parser.add_argument(
+        "--with-orchestrate",
+        action="store_true",
+        help="执行前先跑统筹 Agent 生成/补充计划",
+    )
+    parser.add_argument(
+        "--json",
+        action="store_true",
+        help="最终以 JSON 打印摘要",
+    )
+    args = parser.parse_args(argv)
+
+    logging.basicConfig(
+        level=logging.INFO,
+        format="%(asctime)s %(levelname)s %(name)s: %(message)s",
+    )
+
+    result = grade_demand_pool(
+        str(args.biz_dt).strip(),
+        workers=max(1, int(args.workers)),
+        max_demands_per_batch=max(1, min(int(args.max_demands_per_batch), MAX_DEMANDS_PER_BATCH)),
+        with_orchestrate=bool(args.with_orchestrate),
+        max_rounds=max(0, int(args.max_rounds)),
+    )
+
+    if args.json:
+        print(json.dumps(result, ensure_ascii=False, indent=2, default=str))
+    else:
+        print("\n=== 批量分级完成 ===")
+        print(f"biz_dt={result.get('biz_dt')}")
+        print(f"完成任务组={result.get('groups_run')}")
+        print(f"已分级: {result.get('graded_before')} -> {result.get('graded_after')}")
+        print(f"任务状态: {(result.get('group_status') or result.get('plan_execution', {}).get('final_snapshot', {}).get('group_status'))}")
+        print(f"是否全部完成: {result.get('success')}")
+
+    return 0 if result.get("success") else 1
+
+
+if __name__ == "__main__":
+    raise SystemExit(main())

+ 6 - 30
supply_infra/db/repositories/demand_grade_plan_repo.py

@@ -5,7 +5,7 @@ import uuid
 from datetime import datetime
 from typing import Any
 
-from sqlalchemy import case, select, update
+from sqlalchemy import select, update
 
 from supply_infra.db.models.demand_grade_plan import DemandGradePlan, DemandGradePlanGroup
 from supply_infra.db.repositories.base import BaseRepository
@@ -44,7 +44,7 @@ class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
         groups = self.list_groups_by_biz_dt(biz_dt)
         group_status = self.summarize(biz_dt)
         assigned_category_ids = sorted(self.get_assigned_category_ids(biz_dt))
-        claimable_groups = group_status.get("pending", 0) + group_status.get("failed", 0)
+        claimable_groups = group_status.get("pending", 0)
         unfinished_groups = claimable_groups + group_status.get("running", 0)
         return {
             "planned_groups": len(groups),
@@ -93,35 +93,13 @@ class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
                 status="pending",
             ))
 
-    def claim_next_group(self, biz_dt: str) -> dict[str, Any] | None:
-        """兼容入口:失败任务优先,其次领取未执行任务。"""
-        stmt = (
-            select(DemandGradePlanGroup)
-            .where(
-                DemandGradePlanGroup.biz_dt == biz_dt,
-                DemandGradePlanGroup.status.in_(("pending", "failed")),
-            )
-            .order_by(
-                case((DemandGradePlanGroup.status == "failed", 0), else_=1),
-                DemandGradePlanGroup.id,
-            )
-            .limit(1)
-            .with_for_update(skip_locked=True)
-        )
-        group = self.session.scalar(stmt)
-        if group is None:
-            return None
-        return self._mark_claimed(group)
-
-    def list_group_ids_by_status(self, biz_dt: str, status: str) -> list[int]:
-        """按 id 返回某状态任务;调度器用它冻结单轮任务集合,避免轮内无限重试。"""
-        if status not in {"pending", "failed"}:
-            return []
+    def list_pending_group_ids(self, biz_dt: str) -> list[int]:
+        """返回当天待执行的 pending 任务 id。"""
         stmt = (
             select(DemandGradePlanGroup.id)
             .where(
                 DemandGradePlanGroup.biz_dt == biz_dt,
-                DemandGradePlanGroup.status == status,
+                DemandGradePlanGroup.status == "pending",
             )
             .order_by(DemandGradePlanGroup.id)
         )
@@ -134,7 +112,7 @@ class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
             .where(
                 DemandGradePlanGroup.id == int(group_id),
                 DemandGradePlanGroup.biz_dt == biz_dt,
-                DemandGradePlanGroup.status.in_(("pending", "failed")),
+                DemandGradePlanGroup.status == "pending",
             )
             .with_for_update(skip_locked=True)
         )
@@ -154,8 +132,6 @@ class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
             "id": int(group.id),
             "group_key": group.group_key,
             "category_ids": json.loads(group.category_ids),
-            "planning_reason": group.planning_reason,
-            "shared_traits": group.shared_traits,
         }
 
     def finish_group(self, group_id: int, *, success: bool, error_message: str | None = None) -> None:

+ 3 - 3
supply_infra/scheduler/app.py

@@ -19,7 +19,7 @@ logger = logging.getLogger(__name__)
 
 _scheduler: BackgroundScheduler | None = None
 
-_PIPELINE_CRON_HOURS = "9,15,21"
+_PIPELINE_CRON_HOUR = 12
 
 
 def create_scheduler() -> BackgroundScheduler:
@@ -27,10 +27,10 @@ def create_scheduler() -> BackgroundScheduler:
     settings = get_infra_settings()
     scheduler = BackgroundScheduler(timezone=settings.scheduler_timezone)
 
-    # 每天 9:00 / 15:00 / 21:00 串行执行:全局树 → 需求池 → 分级
+    # 每天 12:00 串行执行:全局树 → 需求池 → 分级
     scheduler.add_job(
         run_supply_pipeline,
-        trigger=CronTrigger(hour=_PIPELINE_CRON_HOURS, minute=0),
+        trigger=CronTrigger(hour=_PIPELINE_CRON_HOUR, minute=0),
         id=SUPPLY_PIPELINE_JOB_ID,
         name=SUPPLY_PIPELINE_JOB_NAME,
         replace_existing=True,

+ 0 - 143
supply_infra/scheduler/grade_assignment.py

@@ -1,143 +0,0 @@
-"""调度执行层:把已落库计划任务转换为最小化的需求分类分配。"""
-from __future__ import annotations
-
-from agents.demand_grade_agent.tools.demand_priority import build_demand_priority_index
-from supply_infra.db.repositories.demand_belong_category_repo import DemandBelongCategoryRepository
-from supply_infra.db.repositories.demand_belong_pool_rel_repo import DemandBelongPoolRelRepository
-from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
-from supply_infra.db.session import get_session
-
-
-def _category_ids_by_name(
-    pool_rows: list,
-    belong_ids_by_pool: dict[int, list[int]],
-    category_by_belong: dict[int, int],
-) -> dict[str, set[int]]:
-    """把需求池行解析为需求词到全部有效分类 ID 的映射。"""
-    result: dict[str, set[int]] = {}
-    for row in pool_rows:
-        if not row.demand_name:
-            continue
-        category_ids = result.setdefault(str(row.demand_name), set())
-        for belong_id in belong_ids_by_pool.get(int(row.id), []):
-            category_id = category_by_belong.get(int(belong_id))
-            if category_id is not None:
-                category_ids.add(category_id)
-    return result
-
-
-def _load_category_ids_by_name(pool_rows: list) -> dict[str, set[int]]:
-    with get_session() as session:
-        belong_ids_by_pool = DemandBelongPoolRelRepository(session).get_belong_ids_by_pool_ids(
-            [int(row.id) for row in pool_rows]
-        )
-        all_belong_ids = {
-            belong_id
-            for values in belong_ids_by_pool.values()
-            for belong_id in values
-        }
-        category_by_belong = {
-            int(row.id): int(row.category_id)
-            for row in DemandBelongCategoryRepository(session).get_by_ids(all_belong_ids)
-            if row.category_id is not None
-        }
-    return _category_ids_by_name(pool_rows, belong_ids_by_pool, category_by_belong)
-
-
-def build_supplement_grade_assignment(biz_dt: str, demand_names: list[str]) -> dict:
-    """为仍未分级的需求生成需求词与所属分类的最小交接件。"""
-    requested_names = list(
-        dict.fromkeys(
-            str(name).strip()
-            for name in demand_names
-            if name is not None and str(name).strip()
-        )
-    )
-    if not requested_names:
-        return {"biz_dt": biz_dt, "items": []}
-
-    requested_set = set(requested_names)
-    with get_session() as session:
-        pool_repo = MultiDemandPoolDiRepository(session)
-        all_pool_rows = pool_repo.list_by_biz_dt(biz_dt)
-        matched_pool_rows = [
-            row for row in all_pool_rows if row.demand_name in requested_set
-        ]
-        priority_index = build_demand_priority_index(all_pool_rows)
-
-    category_ids_by_name = _load_category_ids_by_name(matched_pool_rows)
-    ordered_names = sorted(
-        requested_names,
-        key=lambda name: (
-            priority_index.get(name, {}).get("source_rank_score") is None,
-            -float(priority_index.get(name, {}).get("source_rank_score") or 0),
-            float(priority_index.get(name, {}).get("global_demand_rank") or float("inf")),
-            name,
-        ),
-    )
-    return {
-        "biz_dt": biz_dt,
-        "items": [
-            {
-                "demand_name": name,
-                "category_ids": sorted(category_ids_by_name.get(name, set())),
-            }
-            for name in ordered_names
-        ],
-    }
-
-
-def build_grade_batch_assignment(
-    biz_dt: str,
-    category_ids: list[int],
-    max_demands: int = 20,
-    excluded_demand_names: list[str] | None = None,
-) -> dict:
-    """按计划节点取一小批需求,只返回需求词与其全部所属分类。"""
-    selected_ids = list(dict.fromkeys(int(value) for value in category_ids))
-    excluded = set(excluded_demand_names or [])
-    with get_session() as session:
-        belongs = DemandBelongCategoryRepository(session).list_by_category_ids(selected_ids)
-        pool_ids_by_belong = DemandBelongPoolRelRepository(session).get_pool_ids_by_belong_ids(
-            [int(row.id) for row in belongs]
-        )
-        pool_ids = sorted({pool_id for values in pool_ids_by_belong.values() for pool_id in values})
-        pool_repo = MultiDemandPoolDiRepository(session)
-        pool_rows = pool_repo.get_by_ids(pool_ids)
-        all_day_pool_rows = pool_repo.list_by_biz_dt(biz_dt)
-        priority_index = build_demand_priority_index(all_day_pool_rows)
-        candidate_names = {
-            str(row.demand_name)
-            for row in pool_rows
-            if row.biz_dt == biz_dt and row.demand_name and row.demand_name not in excluded
-        }
-
-    names = sorted(
-        candidate_names,
-        key=lambda name: (
-            priority_index.get(name, {}).get("source_rank_score") is None,
-            -float(priority_index.get(name, {}).get("source_rank_score") or 0),
-            float(priority_index.get(name, {}).get("global_demand_rank") or float("inf")),
-            name,
-        ),
-    )[: max(1, min(int(max_demands), 100))]
-    if not names:
-        return {"biz_dt": biz_dt, "items": []}
-
-    selected_name_set = set(names)
-    selected_pool_rows = [
-        row
-        for row in all_day_pool_rows
-        if row.demand_name in selected_name_set
-    ]
-    category_ids_by_name = _load_category_ids_by_name(selected_pool_rows)
-    return {
-        "biz_dt": biz_dt,
-        "items": [
-            {
-                "demand_name": name,
-                "category_ids": sorted(category_ids_by_name.get(name, set())),
-            }
-            for name in names
-        ],
-    }

+ 102 - 308
supply_infra/scheduler/jobs/grade_demand_pool.py

@@ -1,7 +1,6 @@
-"""统筹落库后执行分级任务,并对任务状态与树上需求等级做补偿闭环。"""
+"""统筹落库后执行分级任务。"""
 from __future__ import annotations
 
-import json
 import logging
 from concurrent.futures import ThreadPoolExecutor, as_completed
 from datetime import datetime
@@ -9,32 +8,20 @@ from typing import Any
 from zoneinfo import ZoneInfo
 
 from agents.demand_grade_agent.run import main as grade_demand_words
-from agents.demand_grade_agent.tools.build_grade_plan_context import (
-    build_grade_plan_context,
-    build_supplement_grade_context,
-)
 from agents.demand_grade_orchestrator_agent.run import orchestrate_daily_grade_plan
-from agents.demand_grade_orchestrator_agent.validation import resolve_assignment_state
 from supply_infra.config import get_infra_settings
-from supply_infra.db.repositories.demand_belong_category_repo import (
-    DemandBelongCategoryRepository,
-)
-from supply_infra.db.repositories.demand_belong_pool_rel_repo import (
-    DemandBelongPoolRelRepository,
-)
 from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
 from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
-from supply_infra.db.repositories.global_tree_category_repo import GlobalTreeCategoryRepository
 from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
 from supply_infra.db.session import get_session
+from supply_infra.scheduler.plan_group_batch import (
+    MAX_DEMANDS_PER_BATCH,
+    list_pending_demands_by_category_ids,
+)
 
 logger = logging.getLogger(__name__)
 
-_DEFAULT_BATCH_SIZE = 20
 _DEFAULT_WORKERS = 5
-_MAX_PLAN_EXECUTION_ROUNDS = 5
-_MAX_SUPPLEMENT_ROUNDS = 5
-_SUPPLEMENT_BATCH_SIZE = 30
 
 
 def _resolve_biz_dt(biz_dt: str | None) -> str:
@@ -48,8 +35,8 @@ def _graded_names(biz_dt: str) -> set[str]:
         return DemandGradeRepository(session).get_existing_demand_names(biz_dt)
 
 
-def _run_group(biz_dt: str, batch_size: int, group_id: int) -> int:
-    """领取并执行一个固定任务;失败留到下一轮,不在当前轮内再次领取。"""
+def _run_group(biz_dt: str, group_id: int, *, max_demands_per_batch: int) -> int:
+    """领取并执行一个固定任务,组内按批处理,失败子批跳过并继续。"""
     with get_session() as session:
         group = DemandGradePlanRepository(session).claim_group(biz_dt, group_id)
     if group is None:
@@ -57,44 +44,42 @@ def _run_group(biz_dt: str, batch_size: int, group_id: int) -> int:
         return 0
 
     processed_batches = 0
+    batch_errors: list[str] = []
+    skipped_names: set[str] = set()
     try:
         while True:
-            graded_before = _graded_names(biz_dt)
-            raw_context = build_grade_plan_context(
-                biz_dt=biz_dt,
-                category_ids=group["category_ids"],
-                planning_reason=group["planning_reason"],
-                shared_traits=group["shared_traits"],
-                max_demands=batch_size,
+            graded_before = _graded_names(biz_dt) | skipped_names
+            demands = list_pending_demands_by_category_ids(
+                biz_dt,
+                group["category_ids"],
+                max_demands=max_demands_per_batch,
                 excluded_demand_names=sorted(graded_before),
             )
-            try:
-                context = json.loads(raw_context)
-            except json.JSONDecodeError as exc:
-                if raw_context == "所选树节点下没有待分级需求":
-                    break
-                raise RuntimeError(f"构建分级上下文失败: {raw_context}") from exc
-
-            names = [
-                str(item["demand_name"])
-                for item in context.get("items", [])
-                if item.get("demand_name")
-            ]
-            if not names:
+            if not demands:
                 break
 
-            grade_demand_words(names, biz_dt=biz_dt, tree_context=context)
-            graded_after = _graded_names(biz_dt)
-            not_saved = sorted(set(names) - graded_after)
-            if not_saved:
-                raise RuntimeError(f"分级 Agent 未落库本批全部需求: {not_saved}")
-            processed_batches += 1
+            try:
+                grade_demand_words(demands, biz_dt=biz_dt)
+                processed_batches += 1
+            except Exception as exc:
+                logger.exception(
+                    "分级子批次失败,跳过并继续本组其余需求: biz_dt=%s group=%s count=%s",
+                    biz_dt,
+                    group["group_key"],
+                    len(demands),
+                )
+                batch_errors.append(str(exc))
+                skipped_names.update(item["demand_name"] for item in demands)
 
         with get_session() as session:
-            DemandGradePlanRepository(session).finish_group(group["id"], success=True)
+            DemandGradePlanRepository(session).finish_group(
+                group["id"],
+                success=processed_batches > 0,
+                error_message="; ".join(batch_errors) if batch_errors else None,
+            )
     except Exception as exc:
         logger.exception(
-            "分级计划任务失败,留待下一轮重试: biz_dt=%s group=%s",
+            "分级计划任务失败: biz_dt=%s group=%s",
             biz_dt,
             group["group_key"],
         )
@@ -111,340 +96,147 @@ def _run_group(biz_dt: str, batch_size: int, group_id: int) -> int:
                 biz_dt,
                 group["group_key"],
             )
+        return 0
     return processed_batches
 
 
-def _execute_group_phase(
+def _execute_plan_tasks(
     biz_dt: str,
     *,
-    status: str,
-    batch_size: int,
     workers: int,
-) -> dict[str, int]:
-    """执行一个状态阶段;调用方先执行 failed,再执行 pending。"""
+    max_demands_per_batch: int,
+) -> dict[str, Any]:
+    """并发执行当天全部 pending 计划任务,仅执行一轮。"""
     with get_session() as session:
-        group_ids = DemandGradePlanRepository(session).list_group_ids_by_status(biz_dt, status)
+        group_ids = DemandGradePlanRepository(session).list_pending_group_ids(biz_dt)
     if not group_ids:
-        return {"attempted_groups": 0, "batches_run": 0, "workers": 0}
+        with get_session() as session:
+            final_snapshot = DemandGradePlanRepository(session).get_execution_snapshot(biz_dt)
+        return {
+            "attempted_groups": 0,
+            "groups_run": 0,
+            "workers": 0,
+            "final_snapshot": final_snapshot,
+            "execution_complete": final_snapshot["execution_complete"],
+        }
 
     worker_count = max(1, min(int(workers), len(group_ids)))
-    batches_run = 0
+    groups_run = 0
     with ThreadPoolExecutor(max_workers=worker_count) as executor:
         futures = [
-            executor.submit(_run_group, biz_dt, batch_size, group_id)
+            executor.submit(_run_group, biz_dt, group_id, max_demands_per_batch=max_demands_per_batch)
             for group_id in group_ids
         ]
         for future in as_completed(futures):
             try:
-                batches_run += future.result()
+                groups_run += future.result()
             except Exception:
-                # _run_group 已自行兜底;这里再兜一次,禁止 worker 异常中断定时任务。
-                logger.exception("分级任务 worker 出现未捕获错误: biz_dt=%s status=%s", biz_dt, status)
+                logger.exception("分级任务 worker 出现未捕获错误: biz_dt=%s", biz_dt)
+
+    with get_session() as session:
+        final_snapshot = DemandGradePlanRepository(session).get_execution_snapshot(biz_dt)
     return {
         "attempted_groups": len(group_ids),
-        "batches_run": batches_run,
+        "groups_run": groups_run,
         "workers": worker_count,
+        "final_snapshot": final_snapshot,
+        "execution_complete": final_snapshot["execution_complete"],
     }
 
 
-def _execute_plan_tasks_with_retries(
+def execute_plan_tasks_until_complete(
     biz_dt: str,
     *,
-    batch_size: int,
     workers: int,
-    max_rounds: int = _MAX_PLAN_EXECUTION_ROUNDS,
+    max_demands_per_batch: int = MAX_DEMANDS_PER_BATCH,
+    max_rounds: int = 0,
 ) -> dict[str, Any]:
-    """检查并执行计划任务,失败优先,任务状态检查最多循环五轮。"""
-    max_rounds = max(1, min(int(max_rounds), _MAX_PLAN_EXECUTION_ROUNDS))
-    history: list[dict[str, Any]] = []
-    total_batches = 0
-    max_workers_used = 0
-
-    for round_no in range(1, max_rounds + 1):
-        with get_session() as session:
-            before = DemandGradePlanRepository(session).get_execution_snapshot(biz_dt)
-        if before["execution_complete"]:
-            logger.info("全部分级计划任务已执行完成,跳过分级 worker: biz_dt=%s", biz_dt)
+    """循环执行 pending 计划任务,直到全部完成或达到 max_rounds。"""
+    rounds: list[dict[str, Any]] = []
+    groups_run = 0
+    round_no = 0
+    while True:
+        if max_rounds > 0 and round_no >= max_rounds:
             break
-
-        failed_result = _execute_group_phase(
-            biz_dt,
-            status="failed",
-            batch_size=batch_size,
-            workers=workers,
-        )
-        # 严格等失败阶段结束后再处理从未执行的任务。
-        pending_result = _execute_group_phase(
+        round_no += 1
+        result = _execute_plan_tasks(
             biz_dt,
-            status="pending",
-            batch_size=batch_size,
             workers=workers,
+            max_demands_per_batch=max_demands_per_batch,
         )
-        total_batches += failed_result["batches_run"] + pending_result["batches_run"]
-        max_workers_used = max(
-            max_workers_used,
-            failed_result["workers"],
-            pending_result["workers"],
-        )
+        rounds.append(result)
+        groups_run += int(result.get("groups_run") or 0)
+        if int(result.get("attempted_groups") or 0) == 0:
+            break
 
+    final_snapshot = rounds[-1]["final_snapshot"] if rounds else None
+    if final_snapshot is None:
         with get_session() as session:
-            after = DemandGradePlanRepository(session).get_execution_snapshot(biz_dt)
-        history.append({
-            "round": round_no,
-            "before": before["group_status"],
-            "failed_attempted": failed_result["attempted_groups"],
-            "pending_attempted": pending_result["attempted_groups"],
-            "after": after["group_status"],
-            "execution_complete": after["execution_complete"],
-        })
-        if after["execution_complete"]:
-            break
-        if failed_result["attempted_groups"] + pending_result["attempted_groups"] == 0:
-            logger.error(
-                "计划任务尚未全部完成,但当前没有可执行的 failed/pending 任务: "
-                "biz_dt=%s round=%s status=%s",
-                biz_dt,
-                round_no,
-                after["group_status"],
-            )
+            final_snapshot = DemandGradePlanRepository(session).get_execution_snapshot(biz_dt)
 
-    with get_session() as session:
-        final_snapshot = DemandGradePlanRepository(session).get_execution_snapshot(biz_dt)
-    if not final_snapshot["execution_complete"]:
-        logger.error(
-            "计划任务最多执行 %s 轮后仍未全部完成,继续执行需求覆盖检查: biz_dt=%s status=%s",
-            max_rounds,
-            biz_dt,
-            final_snapshot["group_status"],
-        )
     return {
-        "rounds": len(history),
-        "history": history,
-        "batches_run": total_batches,
-        "workers": max_workers_used,
+        "rounds": round_no,
+        "groups_run": groups_run,
+        "workers": max(1, int(workers)),
+        "round_details": rounds,
         "final_snapshot": final_snapshot,
-        "execution_complete": final_snapshot["execution_complete"],
-    }
-
-
-def _tree_attached_demand_names(biz_dt: str) -> set[str]:
-    """返回业务日内通过有效挂载关系连接到当前全局树的需求池需求名。"""
-    with get_session() as session:
-        active_category_ids = {
-            int(row.id)
-            for row in GlobalTreeCategoryRepository(session).list_active_categories()
-        }
-        belongs = [
-            row
-            for row in DemandBelongCategoryRepository(session).list_active()
-            if row.category_id is not None and int(row.category_id) in active_category_ids
-        ]
-        pool_ids_by_belong = DemandBelongPoolRelRepository(session).get_pool_ids_by_belong_ids(
-            [int(row.id) for row in belongs]
-        )
-        pool_ids = {
-            int(pool_id)
-            for values in pool_ids_by_belong.values()
-            for pool_id in values
-        }
-        pool_rows = MultiDemandPoolDiRepository(session).get_by_ids(sorted(pool_ids))
-        return {
-            str(row.demand_name)
-            for row in pool_rows
-            if row.biz_dt == biz_dt and row.demand_name
-        }
-
-
-def _resolve_tree_grade_coverage(biz_dt: str) -> dict[str, Any]:
-    tree_demand_names = _tree_attached_demand_names(biz_dt)
-    graded = _graded_names(biz_dt)
-    missing = sorted(tree_demand_names - graded)
-    return {
-        "biz_dt": biz_dt,
-        "tree_demand_count": len(tree_demand_names),
-        "graded_tree_demand_count": len(tree_demand_names & graded),
-        "missing_count": len(missing),
-        "missing_demand_names": missing,
-        "coverage_complete": not missing,
-    }
-
-
-def _coverage_summary(coverage: dict[str, Any]) -> dict[str, Any]:
-    """压缩返回与日志内容,避免大量缺失需求撑大定时任务结果。"""
-    return {
-        "biz_dt": coverage["biz_dt"],
-        "tree_demand_count": coverage["tree_demand_count"],
-        "graded_tree_demand_count": coverage["graded_tree_demand_count"],
-        "missing_count": coverage["missing_count"],
-        "missing_demand_sample": coverage["missing_demand_names"][:30],
-        "coverage_complete": coverage["coverage_complete"],
-    }
-
-
-def _chunks(values: list[str], size: int) -> list[list[str]]:
-    return [values[start : start + size] for start in range(0, len(values), size)]
-
-
-def _supplement_missing_grades(
-    biz_dt: str,
-    *,
-    max_rounds: int = _MAX_SUPPLEMENT_ROUNDS,
-    batch_size: int = _SUPPLEMENT_BATCH_SIZE,
-) -> dict[str, Any]:
-    """补充分级树上缺失需求,每批最多30个,检查与补偿最多循环五轮。"""
-    max_rounds = max(1, min(int(max_rounds), _MAX_SUPPLEMENT_ROUNDS))
-    batch_size = max(1, min(int(batch_size), _SUPPLEMENT_BATCH_SIZE))
-    history: list[dict[str, Any]] = []
-    batches_run = 0
-    initial_coverage = _resolve_tree_grade_coverage(biz_dt)
-
-    for round_no in range(1, max_rounds + 1):
-        before = _resolve_tree_grade_coverage(biz_dt)
-        if before["coverage_complete"]:
-            logger.info("树上全部需求均已分级,跳过补充任务: biz_dt=%s", biz_dt)
-            break
-
-        missing_names = list(before["missing_demand_names"])
-        attempted_batches = 0
-        for batch_no, demand_batch in enumerate(_chunks(missing_names, batch_size), start=1):
-            attempted_batches += 1
-            try:
-                context = build_supplement_grade_context(biz_dt, demand_batch)
-                context_names = [
-                    str(item["demand_name"])
-                    for item in context.get("items", [])
-                    if item.get("demand_name")
-                ]
-                if not context_names:
-                    logger.error(
-                        "补充分级上下文没有有效需求: biz_dt=%s round=%s batch=%s input=%s",
-                        biz_dt,
-                        round_no,
-                        batch_no,
-                        demand_batch,
-                    )
-                    continue
-                grade_demand_words(context_names, biz_dt=biz_dt, tree_context=context)
-                batches_run += 1
-            except Exception:
-                logger.exception(
-                    "补充分级批次失败,继续处理其他批次: biz_dt=%s round=%s batch=%s names=%s",
-                    biz_dt,
-                    round_no,
-                    batch_no,
-                    demand_batch,
-                )
-
-        after = _resolve_tree_grade_coverage(biz_dt)
-        history.append({
-            "round": round_no,
-            "missing_before": before["missing_count"],
-            "attempted_batches": attempted_batches,
-            "missing_after": after["missing_count"],
-            "coverage_complete": after["coverage_complete"],
-        })
-        if after["coverage_complete"]:
-            break
-        if after["missing_count"] >= before["missing_count"]:
-            logger.error(
-                "补充分级本轮没有减少缺失需求,将进入下一轮重试: "
-                "biz_dt=%s round=%s missing=%s",
-                biz_dt,
-                round_no,
-                after["missing_count"],
-            )
-
-    final_coverage = _resolve_tree_grade_coverage(biz_dt)
-    if not final_coverage["coverage_complete"]:
-        logger.error(
-            "补充分级最多执行 %s 轮后仍有树上需求未分级: biz_dt=%s missing_count=%s sample=%s",
-            max_rounds,
-            biz_dt,
-            final_coverage["missing_count"],
-            final_coverage["missing_demand_names"][:30],
-        )
-    return {
-        "rounds": len(history),
-        "history": history,
-        "batches_run": batches_run,
-        "initial_coverage": _coverage_summary(initial_coverage),
-        "final_coverage": _coverage_summary(final_coverage),
-        "coverage_complete": final_coverage["coverage_complete"],
+        "execution_complete": bool(final_snapshot["execution_complete"]),
     }
 
 
 def _grade_demand_pool_impl(
     resolved_biz_dt: str,
     *,
-    batch_size: int,
     workers: int,
+    max_demands_per_batch: int = MAX_DEMANDS_PER_BATCH,
+    with_orchestrate: bool = True,
+    max_rounds: int = 0,
 ) -> dict[str, Any]:
     with get_session() as session:
         total = MultiDemandPoolDiRepository(session).count_distinct_demand_names(resolved_biz_dt)
         graded_before = DemandGradeRepository(session).count_by_biz_dt(resolved_biz_dt)
 
-    try:
-        assignment_before = resolve_assignment_state(resolved_biz_dt)
-    except Exception as exc:
-        logger.exception("查询统筹分配前状态失败,继续尝试统筹: biz_dt=%s", resolved_biz_dt)
-        assignment_before = {"biz_dt": resolved_biz_dt, "error": str(exc)}
-
-    try:
-        orchestrate_daily_grade_plan(biz_dt=resolved_biz_dt)
-    except Exception:
-        logger.exception("统筹 Agent 执行失败,继续处理数据库中已有任务: biz_dt=%s", resolved_biz_dt)
-
-    try:
-        assignment_after = resolve_assignment_state(resolved_biz_dt)
-    except Exception as exc:
-        logger.exception("查询统筹分配后状态失败,继续执行已有任务: biz_dt=%s", resolved_biz_dt)
-        assignment_after = {"biz_dt": resolved_biz_dt, "error": str(exc)}
-
-    if assignment_after.get("assignment_complete") is False:
-        logger.error(
-            "统筹规划未覆盖当天全部有需求的树节点,将按部分计划继续执行: biz_dt=%s uncovered=%s",
-            resolved_biz_dt,
-            assignment_after.get("unassigned_category_ids"),
-        )
+    if with_orchestrate:
+        try:
+            orchestrate_daily_grade_plan(biz_dt=resolved_biz_dt)
+        except Exception:
+            logger.exception("统筹 Agent 执行失败,继续处理数据库中已有任务: biz_dt=%s", resolved_biz_dt)
 
-    plan_execution = _execute_plan_tasks_with_retries(
+    plan_execution = execute_plan_tasks_until_complete(
         resolved_biz_dt,
-        batch_size=max(1, int(batch_size)),
         workers=max(1, int(workers)),
+        max_demands_per_batch=max(1, min(int(max_demands_per_batch), MAX_DEMANDS_PER_BATCH)),
+        max_rounds=max_rounds,
     )
-    # 只在统筹结束并完成计划任务状态检查后,验证树上需求等级覆盖并补偿。
-    supplement = _supplement_missing_grades(resolved_biz_dt)
 
     with get_session() as session:
         graded_after = DemandGradeRepository(session).count_by_biz_dt(resolved_biz_dt)
     final_snapshot = plan_execution["final_snapshot"]
     result = {
-        "success": bool(plan_execution["execution_complete"] and supplement["coverage_complete"]),
+        "success": bool(plan_execution["execution_complete"]),
         "biz_dt": resolved_biz_dt,
         "total": total,
         "graded_before": graded_before,
         "graded_after": graded_after,
-        "assignment_before": assignment_before,
-        "assignment_after": assignment_after,
         "planned_category_count": len(final_snapshot["assigned_category_ids"]),
         "planned_groups": final_snapshot["planned_groups"],
         "group_status": final_snapshot["group_status"],
         "plan_execution": plan_execution,
-        "supplement": supplement,
         "workers": plan_execution["workers"],
-        "batches_run": plan_execution["batches_run"],
-        "supplement_batches_run": supplement["batches_run"],
+        "groups_run": plan_execution["groups_run"],
         "run_at": datetime.now().isoformat(),
     }
-    logger.info("Tree-first grade completed: %s", result)
+    logger.info("Grade demand pool completed: %s", result)
     return result
 
 
 def grade_demand_pool(
     biz_dt: str | None = None,
     *,
-    batch_size: int = _DEFAULT_BATCH_SIZE,
     workers: int = _DEFAULT_WORKERS,
+    max_demands_per_batch: int = MAX_DEMANDS_PER_BATCH,
+    with_orchestrate: bool = True,
+    max_rounds: int = 0,
 ) -> dict[str, Any]:
     """执行完整分级闭环;任何错误只记录日志并返回,不向上中断定时任务。"""
     resolved_biz_dt = str(biz_dt or "")
@@ -452,8 +244,10 @@ def grade_demand_pool(
         resolved_biz_dt = _resolve_biz_dt(biz_dt)
         return _grade_demand_pool_impl(
             resolved_biz_dt,
-            batch_size=batch_size,
             workers=workers,
+            max_demands_per_batch=max_demands_per_batch,
+            with_orchestrate=with_orchestrate,
+            max_rounds=max_rounds,
         )
     except Exception as exc:
         logger.exception("需求分级任务发生未捕获错误,已阻止异常中断定时任务: biz_dt=%s", resolved_biz_dt)

+ 59 - 0
supply_infra/scheduler/plan_group_batch.py

@@ -0,0 +1,59 @@
+"""从 plan_group 的 category_ids 查询待分级需求词。"""
+from __future__ import annotations
+
+from typing import Any
+
+from agents.demand_grade_agent.tools.demand_priority import build_demand_priority_index
+from supply_infra.db.repositories.demand_belong_category_repo import DemandBelongCategoryRepository
+from supply_infra.db.repositories.demand_belong_pool_rel_repo import DemandBelongPoolRelRepository
+from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
+from supply_infra.db.session import get_session
+
+MAX_DEMANDS_PER_BATCH = 30
+
+
+def _priority_sort_key(name: str, priority_index: dict) -> tuple:
+    return (
+        priority_index.get(name, {}).get("source_rank_score") is None,
+        -float(priority_index.get(name, {}).get("source_rank_score") or 0),
+        float(priority_index.get(name, {}).get("global_demand_rank") or float("inf")),
+        name,
+    )
+
+
+def list_pending_demands_by_category_ids(
+    biz_dt: str,
+    category_ids: list[int],
+    *,
+    max_demands: int = MAX_DEMANDS_PER_BATCH,
+    excluded_demand_names: list[str] | None = None,
+) -> list[dict[str, Any]]:
+    """按分类节点取待分级需求池记录(pool_id + demand_name),每批最多 max_demands 条。"""
+    selected_ids = list(dict.fromkeys(int(value) for value in category_ids))
+    excluded = set(excluded_demand_names or [])
+    with get_session() as session:
+        belongs = DemandBelongCategoryRepository(session).list_by_category_ids(selected_ids)
+        pool_ids_by_belong = DemandBelongPoolRelRepository(session).get_pool_ids_by_belong_ids(
+            [int(row.id) for row in belongs]
+        )
+        pool_ids = sorted({pool_id for values in pool_ids_by_belong.values() for pool_id in values})
+        pool_repo = MultiDemandPoolDiRepository(session)
+        pool_rows = pool_repo.get_by_ids(pool_ids)
+        priority_index = build_demand_priority_index(pool_repo.list_by_biz_dt(biz_dt))
+        candidates: list[dict[str, Any]] = []
+        for row in pool_rows:
+            if row.biz_dt != biz_dt or not row.demand_name:
+                continue
+            demand_name = str(row.demand_name)
+            if demand_name in excluded:
+                continue
+            candidates.append({"pool_id": int(row.id), "demand_name": demand_name})
+
+    candidates.sort(
+        key=lambda item: (
+            *_priority_sort_key(item["demand_name"], priority_index),
+            item["pool_id"],
+        )
+    )
+    limit = max(1, min(int(max_demands), MAX_DEMANDS_PER_BATCH))
+    return candidates[:limit]