xueyiming před 1 týdnem
rodič
revize
7c37480cb8
29 změnil soubory, kde provedl 1138 přidání a 208 odebrání
  1. 9 1
      agents/demand_grade_agent/prompt/system_prompt.md
  2. 17 1
      agents/demand_grade_agent/run.py
  3. 3 0
      agents/demand_grade_agent/tools/__init__.py
  4. 47 0
      agents/demand_grade_agent/tools/build_grade_plan_context.py
  5. 5 0
      agents/demand_grade_orchestrator_agent/__init__.py
  6. 24 0
      agents/demand_grade_orchestrator_agent/agent.py
  7. 17 0
      agents/demand_grade_orchestrator_agent/common/__init__.py
  8. 64 0
      agents/demand_grade_orchestrator_agent/common/plan_builder.py
  9. 52 0
      agents/demand_grade_orchestrator_agent/common/tree_state.py
  10. 17 0
      agents/demand_grade_orchestrator_agent/prompt/system_prompt.md
  11. 241 0
      agents/demand_grade_orchestrator_agent/run.py
  12. 16 0
      agents/demand_grade_orchestrator_agent/tools/__init__.py
  13. 20 0
      agents/demand_grade_orchestrator_agent/tools/build_full_day_grade_plan.py
  14. 45 0
      agents/demand_grade_orchestrator_agent/tools/query_global_heat_tree.py
  15. 38 0
      agents/demand_grade_orchestrator_agent/tools/query_heat_node_group.py
  16. 25 0
      agents/demand_grade_orchestrator_agent/validation/__init__.py
  17. 85 0
      agents/demand_grade_orchestrator_agent/validation/assignment.py
  18. 87 0
      agents/demand_grade_orchestrator_agent/validation/full_day_grade_plan.py
  19. 1 1
      jobs/backfill_multi_demand_video_list.py
  20. 8 14
      jobs/grade_demand_pool.py
  21. 3 0
      supply_infra/db/models/__init__.py
  22. 53 0
      supply_infra/db/models/demand_grade_plan.py
  23. 2 0
      supply_infra/db/repositories/__init__.py
  24. 17 0
      supply_infra/db/repositories/demand_belong_pool_rel_repo.py
  25. 119 0
      supply_infra/db/repositories/demand_grade_plan_repo.py
  26. 35 0
      supply_infra/scheduler/jobs/backfill_multi_demand_pool_video_list.py
  27. 81 166
      supply_infra/scheduler/jobs/grade_demand_pool.py
  28. 7 1
      supply_infra/scheduler/jobs/run_supply_pipeline.py
  29. 0 24
      supply_infra/scheduler/jobs/sync_multi_demand_pool_odps_to_mysql.py

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

@@ -3,6 +3,9 @@
 和对应的 biz_dt,任务是结合其归属树节点的**先验热度**与**后验真实效果**,逐一划分 S/A/B/C/D
 五档优先级,并调用工具落库到 `demand_grade` 表,供下游选题/投放决策参考。
 
+调用方通常还会提供一份“全局树热度交接件”:其中包含需求节点、父节点、兄弟节点的热度和局部排名。
+它是评分的必需输入,不可只看需求自身数据。
+
 你只做分级判断,不生成新需求词,也不修改需求池原始数据;只处理消息中给定的这些需求词,
 不需要自行查找或列举其他待处理需求。
 
@@ -25,6 +28,10 @@
   - 偏低 → C
   - 先验也很低、几乎无信号(total_score 为 — 或覆盖维度=0/4)→ D
 - 一个需求可能挂在多个树节点上:以最相关/得分最高的节点为主要依据,reason 中说明取用了哪个节点。
+- **局部热度校正**:节点自身热度高、父节点热且在兄弟中靠前时,可上调同等全局分下的等级或 `score`;
+  节点自身偏冷、父节点与兄弟整体也偏冷时,应下调等级或 `score`。节点自身与局部环境冲突时,
+  不直接套规则:结合词级后验判断它是局部突发还是弱信号,并在 reason 中写明冲突。
+- 父节点/兄弟节点只用于校正,不得覆盖明确的低后验:有充分真实后验且表现差时,仍应降级。
 
 ## 同义/相似需求合并
 同一语义的需求可能因措辞不同而在需求池里表现为多条独立记录(例如「减脂期加餐」与
@@ -46,6 +53,7 @@
 - `query_category_path(category_ids)`:查询类目根到叶路径文本,用于写 reason。
 - `query_demand_popularity_by_word(demand_word_names, biz_dt=None)`:按需求词粒度直接查热度统计,**可一次传入多个词** 交叉验证树节点级结论。
 - `query_score_distribution(biz_dt=None)`:查询先验/后验分数分布,制定本批次统一分档标准。
+- `build_grade_plan_context(...)`:按统筹规划选定的树节点生成节点组原因、共同特征和节点/父/兄弟热度;调用方未提供或需要复查时使用。
 - `batch_save_demand_grades(items, biz_dt=None)`:批量落库分级结果,可重复调用按 (biz_dt, demand_name) upsert 覆盖修正。`related_pool_ids` 必填,`video_list`/`strategies` 自动推导。
 
 ## 工作流程
@@ -56,7 +64,7 @@
    - 一次 `query_demand_category_and_weight(demand_names=[...], biz_dt=...)` 批量取归属与权重;
    - 必要时一次 `query_demand_popularity_by_word(demand_word_names=[...])` 做词粒度交叉验证。
    各工具返回结果每段均标注原始查询词(如 `--- demand_name: xxx ---`),便于对应落库。
-4. 对给定列表中的每一个需求词,结合上述批量结果判定 S/A/B/C/D,reason 写清引用的具体数值(如 "total_score=3.2,位于本批p80,无后验数据");`related_pool_ids` 取自 `search_related_pool_demands` 返回的 `[id=...]`。
+4. 对给定列表中的每一个需求词,结合上述批量结果和交接件中的父/兄弟局部热度判定 S/A/B/C/D。reason 必须同时写清全局位置和局部判断(如“节点 total_score=3.2、兄弟前 10%、父节点高热,无后验数据,因此上调至 A”);`related_pool_ids` 取自 `search_related_pool_demands` 返回的 `[id=...]`。
 5. 全部处理完后,调用一次(或分 2~3 次)`batch_save_demand_grades` 落库,覆盖这批给定的所有需求词。
 6. 简要汇报本批次的分级结果后结束本轮任务。
 

+ 17 - 1
agents/demand_grade_agent/run.py

@@ -6,10 +6,17 @@
 """
 from __future__ import annotations
 
+import json
+from typing import Any
+
 from agents.demand_grade_agent import create_demand_grade_agent
 
 
-def main(demand_names: list[str], biz_dt: str | None = None) -> None:
+def main(
+    demand_names: list[str],
+    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()}")
@@ -17,12 +24,21 @@ def main(demand_names: list[str], biz_dt: str | None = None) -> None:
 
     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}
+
+    以下是编排 Agent 提供的全局树热度交接件。它包含节点自身、父节点、兄弟节点热度及局部排名;
+    必须将其纳入评分理由。必要时可调用 build_grade_plan_context 复查,但不得忽略局部冷热信号:
+    {handoff}
     """
     result = agent.run(user_input)
     print(result.content)

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

@@ -10,6 +10,7 @@ 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_path import query_category_path
 from agents.demand_grade_agent.tools.query_demand_category_and_weight import (
     query_demand_category_and_weight,
@@ -31,12 +32,14 @@ ALL_TOOLS: list[Callable[..., Any]] = [
     query_category_path,
     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_path",
     "query_demand_category_and_weight",
     "query_demand_popularity_by_word",

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

@@ -0,0 +1,47 @@
+"""将已落库的节点组任务转换为分级 Agent 的小批次输入。"""
+from __future__ import annotations
+
+import json
+
+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
+
+
+@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_rows = MultiDemandPoolDiRepository(session).get_by_ids(pool_ids)
+    names = list(dict.fromkeys(
+        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
+    ))[: max(1, min(int(max_demands), 100))]
+    if not names:
+        return "所选树节点下没有待分级需求"
+    return json.dumps({
+        "biz_dt": biz_dt,
+        "items": [{"demand_name": name} for name in names],
+        "plan": {
+            "selected_category_ids": selected_ids,
+            "planning_reason": planning_reason.strip(),
+            "shared_traits": shared_traits.strip(),
+            "selection_method": "global_tree_heat_first",
+        },
+    }, ensure_ascii=False)

+ 5 - 0
agents/demand_grade_orchestrator_agent/__init__.py

@@ -0,0 +1,5 @@
+"""demand_grade_orchestrator_agent — 从全局树统筹需求分级规划。"""
+
+from agents.demand_grade_orchestrator_agent.agent import create_demand_grade_orchestrator_agent
+
+__all__ = ["create_demand_grade_orchestrator_agent"]

+ 24 - 0
agents/demand_grade_orchestrator_agent/agent.py

@@ -0,0 +1,24 @@
+"""需求分级统筹规划 Agent 工厂。"""
+from __future__ import annotations
+
+from pathlib import Path
+
+from supply_agent import Agent
+from supply_agent.config import Settings
+from agents.demand_grade_orchestrator_agent.tools import register_all_tools
+
+_PROMPT_PATH = Path(__file__).parent / "prompt" / "system_prompt.md"
+
+
+def create_demand_grade_orchestrator_agent(
+    settings: Settings | None = None, *, model: str | None = None
+) -> Agent:
+    agent = Agent(
+        settings=settings,
+        name="demand_grade_orchestrator_agent",
+        model=model,
+        system_prompt=_PROMPT_PATH.read_text(encoding="utf-8"),
+        max_iterations=24,
+    )
+    register_all_tools(agent.tools)
+    return agent

+ 17 - 0
agents/demand_grade_orchestrator_agent/common/__init__.py

@@ -0,0 +1,17 @@
+"""统筹规划 Agent 的共享数据访问与格式化逻辑。"""
+
+from agents.demand_grade_orchestrator_agent.common.plan_builder import build_grade_plan_for_category_ids
+from agents.demand_grade_orchestrator_agent.common.tree_state import (
+    format_heat_score,
+    has_hung_demand,
+    load_tree_state,
+    path,
+)
+
+__all__ = [
+    "build_grade_plan_for_category_ids",
+    "format_heat_score",
+    "has_hung_demand",
+    "load_tree_state",
+    "path",
+]

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

@@ -0,0 +1,64 @@
+"""按分类节点构建统筹分级计划。"""
+from __future__ import annotations
+
+from collections import defaultdict
+from typing import Any
+
+from agents.demand_grade_orchestrator_agent.common.tree_state import (
+    format_heat_score,
+    has_hung_demand,
+    load_tree_state,
+    path,
+)
+
+
+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]:
+    """为指定分类节点生成分组计划,不重复包含目标集之外的节点。"""
+    targets = {int(category_id) for category_id in category_ids}
+    by_id, _children, weights = load_tree_state(biz_dt)
+    by_parent: dict[int | None, list[int]] = defaultdict(list)
+    for category_id in sorted(targets):
+        if category_id not in by_id:
+            continue
+        weight = weights.get(category_id)
+        if not has_hung_demand(weight):
+            continue
+        parent_id = by_id[category_id].parent_id
+        by_parent[int(parent_id) if parent_id not in (None, 0) else None].append(category_id)
+
+    width = max(1, min(int(max_nodes_per_group), 20))
+    groups: list[dict[str, Any]] = []
+    for parent_id, node_ids in sorted(by_parent.items(), key=lambda item: (item[0] is None, item[0] or 0)):
+        node_ids.sort(
+            key=lambda cid: (
+                weights[cid].total_score is None,
+                -(float(weights[cid].total_score) if weights[cid].total_score is not None else 0),
+                cid,
+            )
+        )
+        parent_path = path(parent_id, by_id) if parent_id is not None else "根层"
+        for start in range(0, len(node_ids), width):
+            chunk = node_ids[start : start + width]
+            heat_scores = "、".join(format_heat_score(weights[cid]) for cid in chunk)
+            groups.append({
+                "group_id": f"parent-{parent_id or 0}-{start // width + 1}",
+                "category_ids": chunk,
+                "planning_reason": f"同属「{parent_path}」分支,按全局热度分从高到低编排({heat_scores})。",
+                "shared_traits": f"共同父节点={parent_path};{grouping_strategy.strip() or '按同分支、全局热度分顺序分组'}",
+            })
+
+    covered = sorted({cid for group in groups for cid in group["category_ids"]})
+    return {
+        "biz_dt": biz_dt,
+        "grouping_strategy": grouping_strategy.strip(),
+        "groups": groups,
+        "covered_category_ids": covered,
+        "uncovered_category_ids": sorted(targets - set(covered)),
+        "coverage_complete": not (targets - set(covered)),
+    }

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

@@ -0,0 +1,52 @@
+"""全局树热度状态的加载与格式化。"""
+from __future__ import annotations
+
+from collections import defaultdict
+from types import SimpleNamespace
+from typing import Any
+
+from supply_infra.db.repositories.category_tree_weight_repo import CategoryTreeWeightRepository
+from supply_infra.db.repositories.global_tree_category_repo import GlobalTreeCategoryRepository
+from supply_infra.db.session import get_session
+
+
+def load_tree_state(biz_dt: str) -> tuple[dict[int, Any], dict[int | None, list[int]], dict[int, Any]]:
+    with get_session() as session:
+        categories = [
+            SimpleNamespace(id=int(row.id), name=row.name, parent_id=row.parent_id, level=row.level)
+            for row in GlobalTreeCategoryRepository(session).list_active_categories()
+        ]
+        weights = [
+            SimpleNamespace(category_id=int(row.category_id), total_score=row.total_score, hung_word_count=row.hung_word_count)
+            for row in CategoryTreeWeightRepository(session).list_by_biz_dt(biz_dt)
+        ]
+    by_id = {row.id: row for row in categories}
+    children: dict[int | None, list[int]] = defaultdict(list)
+    for row in categories:
+        parent_id = int(row.parent_id) if row.parent_id not in (None, 0) else None
+        children[parent_id].append(row.id)
+    for ids in children.values():
+        ids.sort()
+    return by_id, children, {row.category_id: row for row in weights}
+
+
+def format_heat_score(weight: Any | None) -> str:
+    """格式化原始 total_score:只保留两位小数,无数据时明确为 null。"""
+    if weight is None or weight.total_score is None:
+        return "null"
+    return f"{float(weight.total_score):.2f}"
+
+
+def has_hung_demand(weight: Any | None) -> bool:
+    """判断分类自身是否挂有需求。"""
+    return weight is not None and int(weight.hung_word_count or 0) > 0
+
+
+def path(category_id: int | None, by_id: dict[int, Any]) -> str:
+    names: list[str] = []
+    current = by_id.get(category_id) if category_id is not None else None
+    while current is not None:
+        names.append(current.name or str(current.id))
+        parent_id = int(current.parent_id) if current.parent_id not in (None, 0) else None
+        current = by_id.get(parent_id) if parent_id is not None else None
+    return " > ".join(reversed(names))

+ 17 - 0
agents/demand_grade_orchestrator_agent/prompt/system_prompt.md

@@ -0,0 +1,17 @@
+## 角色与任务
+
+你是需求分级统筹规划 Agent。你不直接给需求定级,也不从需求清单中挑词;你必须先从全局分类树的热度和需求分布出发,确定本轮应处理的一个或多个树节点组,再把规划依据交给下游 `demand_grade_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 是当天完整执行计划。
+
+## 规划原则
+
+- 优先规划高热节点簇、同父节点下共同上升的兄弟节点,或“局部突发但大盘偏冷”的对照节点组。
+- 下游会逐组解析节点下的待分级需求;不要反过来根据需求名称拼凑批次。
+- 热度为空或样本数为 0 是数据不足,不是低热;在规划原因中明确说明。
+- 不自行评级、不调用落库工具。只依据工具真实返回的路径、节点和分数。

+ 241 - 0
agents/demand_grade_orchestrator_agent/run.py

@@ -0,0 +1,241 @@
+"""运行统筹规划 Agent 并将计划落库。"""
+from __future__ import annotations
+
+import json
+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 supply_agent.types import Role
+from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
+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":
+            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
+
+
+def _orchestrate_via_agent(
+    biz_dt: str,
+    *,
+    max_nodes_per_group: int,
+    assignment_state: dict[str, Any],
+    target_ids: set[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)}。"
+    )
+
+    for attempt in range(1, _MAX_PLAN_ATTEMPTS + 1):
+        agent = create_demand_grade_orchestrator_agent()
+        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}
+"""
+        )
+        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)
+
+
+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(
+        biz_dt,
+        payload,
+        assigned_ids=set(assignment_state["assigned_category_ids"]),
+    )
+    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",
+            biz_dt,
+            validation["uncovered_category_ids"],
+        )
+    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"]:
+        logger.info(
+            "当天全部有需求节点已分配,跳过统筹: biz_dt=%s assigned=%s",
+            biz_dt,
+            assignment_state["assigned_category_ids"],
+        )
+        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(
+            biz_dt,
+            max_nodes_per_group=max_nodes_per_group,
+            assignment_state=assignment_state,
+            target_ids=set(assignment_state["required_category_ids"]),
+        )
+        return
+
+    logger.info(
+        "当天存在已分配节点,仅处理未分配节点: biz_dt=%s unassigned=%s",
+        biz_dt,
+        unassigned_ids,
+    )
+    _orchestrate_incremental(
+        biz_dt,
+        max_nodes_per_group=max_nodes_per_group,
+        assignment_state=assignment_state,
+        unassigned_ids=unassigned_ids,
+    )
+
+
+def main(biz_dt: str, *, max_nodes_per_group: int = 4) -> 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)
+    with get_session() as session:
+        repo = DemandGradePlanRepository(session)
+        plan = repo.get_latest_plan(biz_dt)
+        groups = repo.list_groups_by_biz_dt(biz_dt)
+        group_status = repo.summarize(biz_dt)
+        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,
+            "plan_count": 1 if plan is not None else 0,
+            "coverage_complete": assignment_after["assignment_complete"],
+            "total_hanging_nodes": assignment_after["total_hanging_nodes"],
+            "group_count": len(groups),
+            "group_status": group_status,
+            "uncovered_category_ids": assignment_after["unassigned_category_ids"],
+            "sample_groups": [
+                {
+                    "group_key": group.group_key,
+                    "category_ids": json.loads(group.category_ids),
+                    "status": group.status,
+                }
+                for group in groups[:3]
+            ],
+            "latest_plan_summary": plan_summary,
+        }
+    print(json.dumps(result, ensure_ascii=False, indent=2))
+    return result
+
+
+if __name__ == "__main__":
+    import sys
+
+    main(sys.argv[1] if len(sys.argv) > 1 else "20260714")

+ 16 - 0
agents/demand_grade_orchestrator_agent/tools/__init__.py

@@ -0,0 +1,16 @@
+"""需求分级统筹规划 Agent 的工具。"""
+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 supply_agent.tools.registry import ToolRegistry
+
+ALL_TOOLS: list[Callable[..., Any]] = [query_global_heat_tree, query_heat_node_group, build_full_day_grade_plan]
+
+
+def register_all_tools(registry: ToolRegistry) -> ToolRegistry:
+    return registry.from_decorated(*ALL_TOOLS)

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

@@ -0,0 +1,20 @@
+"""构建覆盖全部有需求节点的紧凑日计划。"""
+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:
+    """按共同父节点、原始热度分降序生成全天节点组;返回紧凑 JSON,完整任务明细将单独落表。"""
+    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)

+ 45 - 0
agents/demand_grade_orchestrator_agent/tools/query_global_heat_tree.py

@@ -0,0 +1,45 @@
+"""以层级文本展示当天有需求分类所在的全局树及原始热度分。"""
+from __future__ import annotations
+
+from agents.demand_grade_orchestrator_agent.common import format_heat_score, has_hung_demand, load_tree_state
+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)
+    demand_nodes = {cid for cid, weight in weights.items() if has_hung_demand(weight)}
+    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, []))
+        visible_cache[category_id] = value
+        return value
+
+    lines = [
+        f"biz_dt={biz_dt} | 格式:[分类ID]分类名称[全局热度分] + | "
+        "全局热度分为 total_score,保留两位小数;无数据为 null;+ 表示该分类自身有需求"
+    ]
+
+    def render(category_id: int, depth: int) -> None:
+        if not visible(category_id):
+            return
+        category = by_id[category_id]
+        weight = weights.get(category_id)
+        score_text = format_heat_score(weight)
+        suffix = " +" if category_id in demand_nodes else ""
+        lines.append(f"{'  ' * depth}[{category_id}]{category.name or ''}[{score_text}]{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 __name__ == '__main__':
+    res = query_global_heat_tree('20260714')
+    print(res)

+ 38 - 0
agents/demand_grade_orchestrator_agent/tools/query_heat_node_group.py

@@ -0,0 +1,38 @@
+"""下钻节点,查看父子/兄弟间的原始全局热度分。"""
+from __future__ import annotations
+
+from agents.demand_grade_orchestrator_agent.common import format_heat_score, has_hung_demand, load_tree_state, path
+from supply_agent.tools import tool
+
+
+@tool
+def query_heat_node_group(biz_dt: str, category_ids: list[int]) -> str:
+    """返回指定节点、其父节点及直接子节点的原始热度分;“+”表示节点自身有需求。"""
+    by_id, children, weights = load_tree_state(biz_dt)
+    lines = [
+        f"biz_dt={biz_dt} | 格式:[分类ID]分类名称[全局热度分] + | "
+        "全局热度分为 total_score,保留两位小数;无数据为 null;+ 表示该分类自身有需求"
+    ]
+
+    def describe(category_id: int) -> str:
+        category = by_id[category_id]
+        weight = weights.get(category_id)
+        suffix = " +" if has_hung_demand(weight) else ""
+        return f"[{category_id}]{category.name or ''}[{format_heat_score(weight)}]{suffix}"
+
+    for category_id in dict.fromkeys(int(value) for value in category_ids):
+        category = by_id.get(category_id)
+        if category is None:
+            continue
+        parent_id = int(category.parent_id) if category.parent_id not in (None, 0) else None
+        lines.append(f"节点:{describe(category_id)}")
+        lines.append(f"  路径:{path(category_id, by_id)}")
+        if parent_id is not None:
+            lines.append(f"  父节点:{describe(parent_id)}")
+        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)

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

@@ -0,0 +1,25 @@
+"""统筹规划 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",
+]

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

@@ -0,0 +1,85 @@
+"""当天节点分配状态校验。"""
+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

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

@@ -0,0 +1,87 @@
+"""校验单日统筹计划是否覆盖全部挂载需求分类。"""
+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))

+ 1 - 1
jobs/backfill_multi_demand_video_list.py

@@ -12,7 +12,7 @@ import logging
 import sys
 from datetime import datetime
 
-from supply_infra.scheduler.jobs.sync_multi_demand_pool_odps_to_mysql import (
+from supply_infra.scheduler.jobs.backfill_multi_demand_pool_video_list import (
     backfill_video_list,
 )
 

+ 8 - 14
jobs/grade_demand_pool.py

@@ -1,13 +1,13 @@
 #!/usr/bin/env python3
-"""手动执行需求分级:循环取批次调用 demand_grade_agent,直到需求池分级完毕
+"""手动执行树热度驱动的需求分级。
 
-默认每轮 5 个线程并行,各处理一批互不重叠的需求词
+统筹规划 Agent 先落库当天全量节点组计划,再由多个 worker 领取任务并调用分级 Agent
 
 用法:
-    python jobs/grade_demand_pool.py                      # 默认业务日、批次20、最多200批
+    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 批(用于试跑)
+    python jobs/grade_demand_pool.py 20260716 20 5          # 再指定 5 个并发 worker
 """
 from __future__ import annotations
 
@@ -25,18 +25,12 @@ logging.basicConfig(
 def main(
     biz_dt: str | None = None,
     batch_size_arg: str | None = None,
-    max_batches_arg: str | None = None,
+    workers_arg: str | None = None,
 ) -> dict:
     batch_size = int(batch_size_arg) if batch_size_arg else 20
-    if max_batches_arg is None:
-        # 交给下游 job 根据数据库数据自动计算动态上限。
-        max_batches = None
-    elif max_batches_arg.lower() in {"all", "0", "-1"}:
-        max_batches = None
-    else:
-        max_batches = int(max_batches_arg)
-
-    result = grade_demand_pool(biz_dt, batch_size=batch_size, max_batches=max_batches)
+    workers = int(workers_arg) if workers_arg else 5
+
+    result = grade_demand_pool(biz_dt, batch_size=batch_size, workers=workers)
     print(result)
     return result
 

+ 3 - 0
supply_infra/db/models/__init__.py

@@ -5,6 +5,7 @@ from supply_infra.db.models.demand_belong_category import DemandBelongCategory
 from supply_infra.db.models.demand_belong_pool_rel import DemandBelongPoolRel
 from supply_infra.db.models.demand_grade import DemandGrade
 from supply_infra.db.models.demand_grade_category_rel import DemandGradeCategoryRel
+from supply_infra.db.models.demand_grade_plan import DemandGradePlan, DemandGradePlanGroup
 from supply_infra.db.models.demand_popularity_stats import DemandPopularityStats
 from supply_infra.db.models.generated_demand import GeneratedDemand
 from supply_infra.db.models.global_tree_category import GlobalTreeCategory
@@ -20,6 +21,8 @@ __all__ = [
     "DemandBelongPoolRel",
     "DemandGrade",
     "DemandGradeCategoryRel",
+    "DemandGradePlan",
+    "DemandGradePlanGroup",
     "DemandPopularityStats",
     "GeneratedDemand",
     "GlobalTreeCategory",

+ 53 - 0
supply_infra/db/models/demand_grade_plan.py

@@ -0,0 +1,53 @@
+from __future__ import annotations
+
+from datetime import datetime
+
+from sqlalchemy import BigInteger, Index, Integer, String, Text, UniqueConstraint, func
+from sqlalchemy.orm import Mapped, mapped_column
+
+from supply_infra.db.base import Base
+
+
+class DemandGradePlan(Base):
+    """统筹规划 Agent 产出的单日全量分级计划。"""
+
+    __tablename__ = "demand_grade_plan"
+    __table_args__ = (Index("idx_demand_grade_plan_biz_dt", "biz_dt"),)
+
+    id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
+    plan_id: Mapped[str] = mapped_column(String(36), nullable=False, unique=True)
+    biz_dt: Mapped[str] = mapped_column(String(8), nullable=False)
+    status: Mapped[str] = mapped_column(String(32), nullable=False, default="planned")
+    total_hanging_nodes: Mapped[int] = mapped_column(Integer, nullable=False)
+    group_count: Mapped[int] = mapped_column(Integer, nullable=False)
+    coverage_complete: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
+    plan_json: Mapped[str] = mapped_column(Text, nullable=False)
+    create_time: Mapped[datetime] = mapped_column(nullable=False, server_default=func.now())
+    update_time: Mapped[datetime] = mapped_column(nullable=False, server_default=func.now(), onupdate=func.now())
+
+
+class DemandGradePlanGroup(Base):
+    """日计划中的节点组任务,由并发分级 worker 领取执行。"""
+
+    __tablename__ = "demand_grade_plan_group"
+    __table_args__ = (
+        UniqueConstraint("plan_id", "group_no", name="uk_demand_grade_plan_group"),
+        Index("idx_demand_grade_plan_group_claim", "plan_id", "status", "group_no"),
+        Index("idx_demand_grade_plan_group_biz_dt", "biz_dt"),
+    )
+
+    id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
+    plan_id: Mapped[str] = mapped_column(String(36), nullable=False)
+    biz_dt: Mapped[str] = mapped_column(String(8), nullable=False)
+    group_no: Mapped[int] = mapped_column(Integer, nullable=False)
+    group_key: Mapped[str] = mapped_column(String(128), nullable=False)
+    category_ids: Mapped[str] = mapped_column(Text, nullable=False)
+    planning_reason: Mapped[str] = mapped_column(Text, nullable=False)
+    shared_traits: Mapped[str] = mapped_column(Text, nullable=False)
+    status: Mapped[str] = mapped_column(String(32), nullable=False, default="pending")
+    attempts: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
+    error_message: Mapped[str | None] = mapped_column(Text, nullable=True)
+    started_at: Mapped[datetime | None] = mapped_column(nullable=True)
+    finished_at: Mapped[datetime | None] = mapped_column(nullable=True)
+    create_time: Mapped[datetime] = mapped_column(nullable=False, server_default=func.now())
+    update_time: Mapped[datetime] = mapped_column(nullable=False, server_default=func.now(), onupdate=func.now())

+ 2 - 0
supply_infra/db/repositories/__init__.py

@@ -14,6 +14,7 @@ from supply_infra.db.repositories.demand_grade_category_rel_repo import (
     DemandGradeCategoryRelRepository,
 )
 from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
+from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
 from supply_infra.db.repositories.demand_popularity_stats_repo import (
     DemandPopularityStatsRepository,
 )
@@ -36,6 +37,7 @@ __all__ = [
     "DemandBelongPoolRelRepository",
     "DemandGradeCategoryRelRepository",
     "DemandGradeRepository",
+    "DemandGradePlanRepository",
     "DemandPopularityStatsRepository",
     "GeneratedDemandRepository",
     "GlobalTreeCategoryRepository",

+ 17 - 0
supply_infra/db/repositories/demand_belong_pool_rel_repo.py

@@ -35,6 +35,23 @@ class DemandBelongPoolRelRepository(BaseRepository[DemandBelongPoolRel]):
                 result.setdefault(int(pool_id), []).append(int(belong_id))
         return result
 
+    def get_pool_ids_by_belong_ids(self, belong_ids: Iterable[int]) -> dict[int, list[int]]:
+        """正向查询归属词关联的需求池行:belong_id -> [pool_id, ...]。"""
+        id_list = [int(belong_id) for belong_id in belong_ids]
+        if not id_list:
+            return {}
+
+        result: dict[int, list[int]] = {}
+        for i in range(0, len(id_list), _BATCH_SIZE):
+            batch = id_list[i : i + _BATCH_SIZE]
+            stmt = select(
+                DemandBelongPoolRel.demand_belong_category_id,
+                DemandBelongPoolRel.multi_demand_pool_di_id,
+            ).where(DemandBelongPoolRel.demand_belong_category_id.in_(batch))
+            for belong_id, pool_id in self.session.execute(stmt).all():
+                result.setdefault(int(belong_id), []).append(int(pool_id))
+        return result
+
     def get_existing_pairs(self, pairs: Iterable[RelPair]) -> set[RelPair]:
         """返回 pairs 中已存在的 (belong_id, pool_id)。"""
         pair_list = [(int(b), int(p)) for b, p in pairs]

+ 119 - 0
supply_infra/db/repositories/demand_grade_plan_repo.py

@@ -0,0 +1,119 @@
+from __future__ import annotations
+
+import json
+import uuid
+from datetime import datetime
+from typing import Any
+
+from sqlalchemy import select, update
+
+from supply_infra.db.models.demand_grade_plan import DemandGradePlan, DemandGradePlanGroup
+from supply_infra.db.repositories.base import BaseRepository
+
+
+class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
+    model = DemandGradePlan
+
+    def get_latest_plan(self, biz_dt: str) -> DemandGradePlan | None:
+        return self.session.scalar(
+            select(DemandGradePlan)
+            .where(DemandGradePlan.biz_dt == biz_dt)
+            .order_by(DemandGradePlan.create_time.desc())
+            .limit(1)
+        )
+
+    def list_groups_by_biz_dt(self, biz_dt: str) -> list[DemandGradePlanGroup]:
+        return list(self.session.scalars(
+            select(DemandGradePlanGroup)
+            .where(DemandGradePlanGroup.biz_dt == biz_dt)
+            .order_by(DemandGradePlanGroup.id)
+        ).all())
+
+    def get_assigned_category_ids(self, biz_dt: str) -> set[int]:
+        assigned: set[int] = set()
+        for group in self.list_groups_by_biz_dt(biz_dt):
+            for category_id in json.loads(group.category_ids):
+                try:
+                    assigned.add(int(category_id))
+                except (TypeError, ValueError):
+                    continue
+        return assigned
+
+    def get_execution_snapshot(self, biz_dt: str) -> dict[str, Any]:
+        """在 session 内汇总当天计划分组执行情况,避免 ORM 脱离会话。"""
+        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)
+        return {
+            "planned_groups": len(groups),
+            "assigned_category_ids": assigned_category_ids,
+            "group_status": group_status,
+            "claimable_groups": claimable_groups,
+        }
+
+    def create_plan(self, biz_dt: str, payload: dict[str, Any]) -> None:
+        plan_id = str(uuid.uuid4())
+        groups = payload.get("groups") or []
+        coverage_complete = bool(payload.get("coverage_complete"))
+        # 分组明细已写入 demand_grade_plan_group,计划表仅保存可检索的紧凑摘要,避免 TEXT 溢出。
+        summary = {
+            "biz_dt": biz_dt,
+            "grouping_strategy": payload.get("grouping_strategy"),
+            "total_hanging_nodes": payload.get("total_hanging_nodes"),
+            "group_count": len(groups),
+            "coverage_complete": coverage_complete,
+            "covered_category_ids": payload.get("covered_category_ids", []),
+            "uncovered_category_ids": payload.get("uncovered_category_ids", []),
+        }
+        self.add(DemandGradePlan(
+            plan_id=plan_id, biz_dt=biz_dt, status="planned",
+            total_hanging_nodes=int(payload.get("total_hanging_nodes") or 0),
+            group_count=len(groups), coverage_complete=1 if coverage_complete else 0,
+            plan_json=json.dumps(summary, ensure_ascii=False),
+        ))
+        for group_no, group in enumerate(groups, start=1):
+            self.add(DemandGradePlanGroup(
+                plan_id=plan_id, biz_dt=biz_dt, group_no=group_no, group_key=str(group["group_id"]),
+                category_ids=json.dumps(group["category_ids"], ensure_ascii=False),
+                planning_reason=str(group["planning_reason"]), shared_traits=str(group["shared_traits"]),
+                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(DemandGradePlanGroup.id)
+            .limit(1)
+            .with_for_update(skip_locked=True)
+        )
+        group = self.session.scalar(stmt)
+        if group is None:
+            return None
+        group.status = "running"
+        group.attempts += 1
+        group.started_at = datetime.now()
+        return {
+            "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:
+        self.session.execute(
+            update(DemandGradePlanGroup)
+            .where(DemandGradePlanGroup.id == group_id)
+            .values(status="finished" if success else "failed", error_message=error_message, finished_at=datetime.now())
+        )
+
+    def summarize(self, biz_dt: str) -> dict[str, int]:
+        rows = self.session.execute(
+            select(DemandGradePlanGroup.status).where(DemandGradePlanGroup.biz_dt == biz_dt)
+        ).scalars().all()
+        return {status: sum(value == status for value in rows) for status in ("pending", "running", "finished", "failed")}

+ 35 - 0
supply_infra/scheduler/jobs/backfill_multi_demand_pool_video_list.py

@@ -0,0 +1,35 @@
+"""手动维护任务:回填需求池记录的 video_list / video_count。"""
+from __future__ import annotations
+
+import logging
+from typing import Any
+
+from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
+from supply_infra.db.session import get_session
+from supply_infra.odps.client import get_odps_client
+from supply_infra.scheduler.jobs.sync_multi_demand_pool_odps_to_mysql import _to_mysql_rows
+
+logger = logging.getLogger(__name__)
+
+
+def backfill_video_list(partition_date: str) -> dict[str, Any]:
+    """从 ODPS 回填指定分区的 video_list / video_count(每条最多前 10 个 video_id)。"""
+    logger.info("Backfill video_list for partition: %s", partition_date)
+    raw_rows = get_odps_client().fetch_multi_demand_pool(partition_date)
+    mysql_rows = _to_mysql_rows(raw_rows, partition_date)
+
+    with get_session() as session:
+        updated = MultiDemandPoolDiRepository(session).update_video_fields(
+            partition_date,
+            mysql_rows,
+        )
+
+    result = {
+        "partition_date": partition_date,
+        "fetched": len(raw_rows),
+        "unique_rows": len(mysql_rows),
+        "updated": updated,
+        "with_video": sum(1 for row in mysql_rows if row.get("video_list")),
+    }
+    logger.info("Backfill video_list completed: %s", result)
+    return result

+ 81 - 166
supply_infra/scheduler/jobs/grade_demand_pool.py

@@ -1,23 +1,18 @@
-"""
-定时任务:对 multi_demand_pool_di 需求池中的需求循环打分级。
-
-循环控制在本文件(调度任务侧),而不是让 agent 在一次运行内自行分页遍历——
-每一轮并行启动多个 worker(默认 5 个线程),每个 worker 处理一批互不重叠的
-需求词(默认每批 20 个),调用一次 demand_grade_agent 并等待其运行结束
-(agent 内部会调用 batch_save_demand_grades 落库),再重新查询"已分级"集合、
-分配下一批新词,如此循环直到没有更多待分级需求或达到本次运行的批次上限,
-从而避免单次 agent 会话上下文无限增长,同时提升吞吐。
-"""
+"""统筹规划落库后,由多个 worker 领取节点组任务并调用分级 Agent。"""
 from __future__ import annotations
 
+import json
 import logging
-import math
 from concurrent.futures import ThreadPoolExecutor, as_completed
 from datetime import datetime
 from zoneinfo import ZoneInfo
 
 from agents.demand_grade_agent.run import main as grade_demand_words
+from agents.demand_grade_orchestrator_agent.run import orchestrate_daily_grade_plan
+from agents.demand_grade_orchestrator_agent.validation import resolve_assignment_state
+from agents.demand_grade_agent.tools.build_grade_plan_context import build_grade_plan_context
 from supply_infra.config import get_infra_settings
+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.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
 from supply_infra.db.session import get_session
@@ -28,188 +23,108 @@ _DEFAULT_BATCH_SIZE = 20
 _DEFAULT_WORKERS = 5
 
 
-def _resolve_biz_dt(biz_dt: str | None) -> str | None:
+def _resolve_biz_dt(biz_dt: str | None) -> str:
     if biz_dt:
         return biz_dt
-    # 定时任务默认按“当天业务日”执行;使用 scheduler_timezone 以免机器时区不同导致 biz_dt 偏移。
-    settings = get_infra_settings()
-    tz = ZoneInfo(settings.scheduler_timezone)
-    return datetime.now(tz).strftime("%Y%m%d")
+    return datetime.now(ZoneInfo(get_infra_settings().scheduler_timezone)).strftime("%Y%m%d")
 
 
-def _fetch_next_batch(biz_dt: str, batch_size: int, exclude_names: set[str]) -> list[str]:
-    """按 exclude_names 过滤已分配/已分级词,取下一批(不用 offset,避免和落库进度错位)。"""
-    with get_session() as session:
-        summaries = MultiDemandPoolDiRepository(session).list_distinct_demand_name_summaries(
-            biz_dt,
-            limit=batch_size,
-            offset=0,
-            exclude_names=list(exclude_names) if exclude_names else None,
-        )
-    return [item["demand_name"] for item in summaries]
-
-
-def _allocate_parallel_batches(
-    biz_dt: str,
-    batch_size: int,
-    exclude_names: set[str],
-    *,
-    num_workers: int,
-) -> list[list[str]]:
-    """
-    为同一轮并行 worker 分配互不重叠的批次。
-
-    每取出一批后立即加入 reserved,下一批查询时排除,保证单轮内不会重复分配。
-    """
-    batches: list[list[str]] = []
-    reserved = set(exclude_names)
-    for _ in range(num_workers):
-        batch = _fetch_next_batch(biz_dt, batch_size, reserved)
-        if not batch:
-            break
-        batches.append(batch)
-        reserved.update(batch)
-    return batches
-
-
-def _fetch_graded_names(biz_dt: str) -> set[str]:
+def _graded_names(biz_dt: str) -> set[str]:
     with get_session() as session:
         return DemandGradeRepository(session).get_existing_demand_names(biz_dt)
 
 
-def _run_batch_in_thread(batch: list[str], biz_dt: str, batch_label: str) -> None:
-    logger.info("Grade demand pool %s (size=%d): %s", batch_label, len(batch), batch)
-    grade_demand_words(batch, biz_dt=biz_dt)
+def _run_group(biz_dt: str, batch_size: int) -> int:
+    """一个 worker 持续领取任务;同组需求过多时使用同一计划上下文分批处理。"""
+    processed_batches = 0
+    while True:
+        with get_session() as session:
+            group = DemandGradePlanRepository(session).claim_next_group(biz_dt)
+        if group is None:
+            return processed_batches
+        try:
+            while True:
+                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,
+                    excluded_demand_names=sorted(_graded_names(biz_dt)),
+                )
+                try:
+                    context = json.loads(raw_context)
+                except json.JSONDecodeError:
+                    break
+                names = [item["demand_name"] for item in context.get("items", [])]
+                if not names:
+                    break
+                grade_demand_words(names, biz_dt=biz_dt, tree_context=context)
+                processed_batches += 1
+            with get_session() as session:
+                DemandGradePlanRepository(session).finish_group(group["id"], success=True)
+        except Exception as exc:
+            logger.exception("Grade plan group failed: biz_dt=%s group=%s", biz_dt, group["group_key"])
+            with get_session() as session:
+                DemandGradePlanRepository(session).finish_group(group["id"], success=False, error_message=str(exc))
 
 
 def grade_demand_pool(
-    biz_dt: str | None = None,
-    *,
-    batch_size: int = _DEFAULT_BATCH_SIZE,
-    max_batches: int | None = None,
-    workers: int = _DEFAULT_WORKERS,
+    biz_dt: str | None = None, *, batch_size: int = _DEFAULT_BATCH_SIZE, workers: int = _DEFAULT_WORKERS
 ) -> dict:
-    """
-    循环对需求池中尚未分级的需求打分级,每轮并行调用多个 agent 并等待其完成。
-
-    Args:
-        biz_dt: 业务日期 YYYYMMDD;不传则取需求池最新业务日。
-        batch_size: 每批交给 agent 的需求词数量。
-        max_batches: 本次运行最多执行多少批(安全阀,避免单次运行时间过长/无限循环)。
-                     不传/为 None 时,将根据当天数据库待分级需求词数量动态计算上限:
-                     `ceil(total / batch_size) + 10`。
-        workers: 并行线程数;每轮最多同时跑 workers 个互不重叠的批次。
-
-    Returns:
-        运行统计:{"biz_dt", "total", "graded_before", "graded_after", "batches_run", "stopped_reason"}
-    """
+    """先统筹落库(按需增量),再并发消费当天待执行节点组。"""
     resolved_biz_dt = _resolve_biz_dt(biz_dt)
     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)
 
-    graded_names = _fetch_graded_names(resolved_biz_dt)
-    graded_before = len(graded_names)
-
-    # max_batches 默认为动态上限:ceil(total / batch_size) + 10
-    # 其中 total 使用“当天需求词去重数”,与本 job 的分页维度一致。
-    if max_batches is None:
-        batches_needed = math.ceil(total / batch_size) if batch_size > 0 else 0
-        max_batches = batches_needed + 10
-
-    worker_count = max(1, workers)
-
-    logger.info(
-        "Grade demand pool start: biz_dt=%s total=%d graded=%d batch_size=%d "
-        "workers=%d max_batches=%s",
-        resolved_biz_dt,
-        total,
-        graded_before,
-        batch_size,
-        worker_count,
-        max_batches,
-    )
-
-    batches_run = 0
-    stopped_reason = "no_more_pending"
-
-    while True:
-        if max_batches is not None and batches_run >= max_batches:
-            stopped_reason = "max_batches_reached"
-            break
-
-        if max_batches is not None:
-            slots = max_batches - batches_run
-            num_workers = min(worker_count, slots)
-        else:
-            num_workers = worker_count
+    assignment_state = resolve_assignment_state(resolved_biz_dt)
+    orchestrate_daily_grade_plan(biz_dt=resolved_biz_dt)
+    assignment_after = resolve_assignment_state(resolved_biz_dt)
 
-        if num_workers <= 0:
-            stopped_reason = "max_batches_reached"
-            break
+    with get_session() as session:
+        snapshot = DemandGradePlanRepository(session).get_execution_snapshot(resolved_biz_dt)
 
-        batches = _allocate_parallel_batches(
+    if not assignment_after["assignment_complete"]:
+        logger.error(
+            "统筹规划未覆盖当天全部有需求的树节点,将按部分计划继续执行: biz_dt=%s uncovered=%s",
             resolved_biz_dt,
-            batch_size,
-            graded_names,
-            num_workers=num_workers,
+            assignment_after["unassigned_category_ids"],
         )
-        if not batches:
-            stopped_reason = "no_more_pending"
-            break
-
-        round_start_batch = batches_run + 1
-        round_failed = False
-
-        with ThreadPoolExecutor(max_workers=len(batches)) as executor:
-            futures = {
-                executor.submit(
-                    _run_batch_in_thread,
-                    batch,
-                    resolved_biz_dt,
-                    f"batch {round_start_batch + idx}",
-                ): batch
-                for idx, batch in enumerate(batches)
-            }
+
+    claimable_groups = snapshot["claimable_groups"]
+    worker_count = max(1, min(int(workers), claimable_groups)) if claimable_groups else 0
+    batches_run = 0
+    if worker_count > 0:
+        with ThreadPoolExecutor(max_workers=worker_count) as executor:
+            futures = [
+                executor.submit(_run_group, resolved_biz_dt, max(1, batch_size))
+                for _ in range(worker_count)
+            ]
             for future in as_completed(futures):
-                batch = futures[future]
-                try:
-                    future.result()
-                except Exception as e:
-                    logger.error(
-                        "Grade demand pool batch failed: %s (batch=%s)",
-                        e,
-                        batch,
-                        exc_info=True,
-                    )
-                    round_failed = True
-
-        batches_run += len(batches)
-
-        if round_failed:
-            stopped_reason = "batch_failed"
-            break
-
-        new_graded_names = _fetch_graded_names(resolved_biz_dt)
-        if len(new_graded_names) <= len(graded_names):
-            # agent 没有对本轮产生任何新的落库结果,避免死循环重复拿到同一批
-            logger.warning(
-                "Grade demand pool round made no progress (graded still %d), stop",
-                len(graded_names),
-            )
-            stopped_reason = "stalled"
-            graded_names = new_graded_names
-            break
-        graded_names = new_graded_names
+                batches_run += future.result()
+    else:
+        logger.info("当天无待执行节点组,跳过分级 worker: biz_dt=%s", resolved_biz_dt)
+
+    with get_session() as session:
+        repo = DemandGradePlanRepository(session)
+        final_snapshot = repo.get_execution_snapshot(resolved_biz_dt)
+        graded_after = DemandGradeRepository(session).count_by_biz_dt(resolved_biz_dt)
 
     result = {
         "biz_dt": resolved_biz_dt,
         "total": total,
         "graded_before": graded_before,
-        "graded_after": len(graded_names),
-        "batches_run": batches_run,
+        "graded_after": graded_after,
+        "assignment_before": assignment_state,
+        "assignment_after": assignment_after,
+        "planned_category_count": len(final_snapshot["assigned_category_ids"]),
+        "planned_groups": final_snapshot["planned_groups"],
+        "claimable_groups": claimable_groups,
+        "group_status": final_snapshot["group_status"],
         "workers": worker_count,
-        "stopped_reason": stopped_reason,
+        "batches_run": batches_run,
         "run_at": datetime.now().isoformat(),
     }
-    logger.info("Grade demand pool completed: %s", result)
+    logger.info("Tree-first grade completed: %s", result)
     return result

+ 7 - 1
supply_infra/scheduler/jobs/run_supply_pipeline.py

@@ -14,7 +14,9 @@ import logging
 import threading
 from datetime import datetime, timedelta
 from typing import Any
+from zoneinfo import ZoneInfo
 
+from supply_infra.config import get_infra_settings
 from supply_infra.scheduler.constants import (
     SUPPLY_PIPELINE_JOB_ID,
     SUPPLY_PIPELINE_JOB_NAME,
@@ -33,7 +35,11 @@ _pipeline_lock = threading.Lock()
 
 def _resolve_dates(biz_dt: str | None) -> tuple[str, str]:
     """返回 (biz_dt, global_tree_partition_date),global_tree 使用 biz_dt 前一日。"""
-    resolved_biz_dt = biz_dt or datetime.now().strftime("%Y%m%d")
+    if biz_dt:
+        resolved_biz_dt = biz_dt
+    else:
+        timezone = ZoneInfo(get_infra_settings().scheduler_timezone)
+        resolved_biz_dt = datetime.now(timezone).strftime("%Y%m%d")
     tree_partition = (
         datetime.strptime(resolved_biz_dt, "%Y%m%d") - timedelta(days=1)
     ).strftime("%Y%m%d")

+ 0 - 24
supply_infra/scheduler/jobs/sync_multi_demand_pool_odps_to_mysql.py

@@ -238,30 +238,6 @@ def _sync_diff(partition_date: str) -> dict[str, Any]:
     }
 
 
-def backfill_video_list(partition_date: str) -> dict[str, Any]:
-    """从 ODPS 回填指定分区的 video_list / video_count(每条最多前 10 个 video_id)。"""
-    logger.info("Backfill video_list for partition: %s", partition_date)
-    odps = get_odps_client()
-    raw_rows = odps.fetch_multi_demand_pool(partition_date)
-    mysql_rows = _to_mysql_rows(raw_rows, partition_date)
-
-    with get_session() as session:
-        updated = MultiDemandPoolDiRepository(session).update_video_fields(
-            partition_date,
-            mysql_rows,
-        )
-
-    result = {
-        "partition_date": partition_date,
-        "fetched": len(raw_rows),
-        "unique_rows": len(mysql_rows),
-        "updated": updated,
-        "with_video": sum(1 for r in mysql_rows if r.get("video_list")),
-    }
-    logger.info("Backfill video_list completed: %s", result)
-    return result
-
-
 def _to_float(value: Any) -> float | None:
     if value is None:
         return None