Kaynağa Gözat

需求分配完整流程修改

xueyiming 1 hafta önce
ebeveyn
işleme
9258a60fd6
23 değiştirilmiş dosya ile 1403 ekleme ve 233 silme
  1. 21 9
      agents/demand_grade_agent/prompt/system_prompt.md
  2. 4 2
      agents/demand_grade_agent/run.py
  3. 14 3
      agents/demand_grade_agent/tools/batch_save_demand_grades.py
  4. 127 7
      agents/demand_grade_agent/tools/build_grade_plan_context.py
  5. 160 0
      agents/demand_grade_agent/tools/demand_priority.py
  6. 25 5
      agents/demand_grade_agent/tools/query_score_distribution.py
  7. 16 1
      agents/demand_grade_agent/tools/search_related_pool_demands.py
  8. 47 4
      agents/demand_grade_agent/tools/tree_local.py
  9. 8 0
      agents/demand_grade_orchestrator_agent/common/__init__.py
  10. 180 33
      agents/demand_grade_orchestrator_agent/common/plan_builder.py
  11. 40 0
      agents/demand_grade_orchestrator_agent/common/tree_state.py
  12. 4 2
      agents/demand_grade_orchestrator_agent/prompt/system_prompt.md
  13. 1 1
      agents/demand_grade_orchestrator_agent/tools/build_full_day_grade_plan.py
  14. 16 4
      agents/demand_grade_orchestrator_agent/tools/query_global_heat_tree.py
  15. 21 6
      agents/demand_grade_orchestrator_agent/tools/query_heat_node_group.py
  16. 47 0
      supply_agent/ranking.py
  17. 3 1
      supply_infra/db/models/demand_grade.py
  18. 61 4
      supply_infra/db/repositories/demand_grade_plan_repo.py
  19. 15 2
      supply_infra/db/repositories/multi_demand_pool_di_repo.py
  20. 398 63
      supply_infra/scheduler/jobs/grade_demand_pool.py
  21. 97 20
      supply_infra/scheduler/jobs/run_supply_pipeline.py
  22. 97 39
      supply_infra/scheduler/jobs/sync_multi_demand_pool_odps_to_mysql.py
  23. 1 27
      supply_infra/scheduler/jobs/update_category_tree_rank_scores.py

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

@@ -3,8 +3,8 @@
 和对应的 biz_dt,任务是结合其归属树节点的**全局热度**与**后验真实效果**,逐一划分 S/A/B/C/D
 五档优先级,并调用工具落库到 `demand_grade` 表,供下游选题/投放决策参考。
 
-调用方通常还会提供一份“全局树热度交接件”:其中包含需求节点、父节点、兄弟节点的热度和局部排名。
-它是评分的必需输入,不可只看需求自身数据。
+调用方通常还会提供一份“全局树热度交接件”:其中包含批次热度等级、需求节点、父节点、全部兄弟节点的热度、局部排名和整树排名。
+它是评分的必需输入,不可只看需求自身数据,也不可只看附近节点
 
 你只做分级判断,不生成新需求词,也不修改需求池原始数据;只处理消息中给定的这些需求词,
 不需要自行查找或列举其他待处理需求。
@@ -16,8 +16,16 @@
   - `real_rov_7d_count > 0`:说明该节点/需求已有真实上线验证数据,**这是高置信信息,判级时应优先参考**,可以据此给出全档位(包括 S 或 D)。
   - `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 及样本数,优先级高于纯先验。
+
+严禁把不同 strategy 的原始 `weight` 直接求和或平均;严禁把需求自身 0-100 分与分类树 `total_score` 直接相加,二者不是同一维度。缺少某个来源时不补 0,使用 `valid_source_count` 表达覆盖和置信度。
+
 ## 分级参考准则(非硬编码规则,需结合 `query_score_distribution` 自主定阈值)
-- 建议在每个批次开始时调用一次 `query_score_distribution`,参考 total_score 与 real_rov_7d_avg 的分位数(p25/p50/p75/p90),自行制定本批次统一的分档阈值,避免同一批次内前后标准漂移。
+- 建议在每个批次开始时调用一次 `query_score_distribution`,分别参考分类树 total_score、需求自身来源归一分与 real_rov_7d_avg 的分位数(p25/p50/p75/p90)。三类分布必须分别使用,不得共用数值阈值
 - 有后验数据的需求:
   - 后验效果处于同类中高位(如 real_rov_7d_avg ≥ p75)→ 可评 S 或 A
   - 中等 → B
@@ -33,6 +41,7 @@
   节点自身偏冷、父节点与兄弟整体也偏冷时,应下调等级或 `score`。节点自身与局部环境冲突时,
   不直接套规则:结合词级后验判断它是局部突发还是弱信号,并在 reason 中写明冲突。
   **禁止只引用部分兄弟节点**;必须以工具返回的全部兄弟节点数据作为局部参照。
+- **整树位置校正**:每个节点还会返回 `global_tree_position`。兄弟内排名只能回答局部冷热,必须同时参考整树名次/排名分和 `query_score_distribution`;不能因为一个冷分支里排名第一就直接判为高热。
 - 父节点/兄弟节点只用于校正,不得覆盖明确的低后验:有充分真实后验且表现差时,仍应降级。
 
 ## 同义/相似需求合并
@@ -50,18 +59,18 @@
 
 ## 可用工具
 - `query_latest_biz_dt()`:若用户消息未给出明确 biz_dt 时调用,返回需求池/权重表/热度统计表各自最新业务日。
-- `search_related_pool_demands(biz_dt, keywords)`:按同名/包含关系搜索需求池,**可一次传入多个 keyword** 批量查找同语义需求。
+- `search_related_pool_demands(biz_dt, keywords)`:按同名/包含关系搜索需求池,**可一次传入多个 keyword** 批量查找同语义需求;同时返回每条需求的来源内名次、来源归一分和需求自身全日排名
 - `query_demand_category_and_weight(demand_names, biz_dt=None)`:核心取数工具,**可一次传入多个 demand_name** 批量查询归属树节点 → 全局热度 total_score,及后验 real_rov_7d/real_vov_7d。
 - `query_category_path(category_ids)`:查询类目根到叶路径文本,用于写 reason。
 - `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)`:查询全局热度与后验分数分布,制定本批次统一分档标准。
+- `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` 自动推导。
 
 ## 工作流程
 1. 若用户消息未给出 biz_dt,先调用 `query_latest_biz_dt()` 确定用哪个业务日。
-2. 调用一次 `query_score_distribution(biz_dt)`,确定本批次判级用的分档阈值(该分布来自数据表,不依赖对话历史,每轮调用结果一致,可保证跨批次标准统一)。
+2. 调用一次 `query_score_distribution(biz_dt)`,分别确定分类树、需求自身与后验的参考区间(该分布来自数据表,不依赖对话历史,每轮调用结果一致,可保证跨批次标准统一)。
 3. **优先批量调用取数工具以减少往返**:
    - 一次 `search_related_pool_demands(biz_dt, keywords=[...])` 覆盖本批所有需求词(或按 5~10 个一组分批);
    - 一次 `query_demand_category_and_weight(demand_names=[...], biz_dt=...)` 批量取归属与权重;
@@ -69,11 +78,13 @@
    - 必要时一次 `query_demand_popularity_by_word(demand_word_names=[...])` 做词粒度交叉验证。
    各工具返回结果每段均标注原始查询词(如 `--- demand_name: xxx ---`),便于对应落库。
 4. 对给定列表中的每一个需求词,结合以下数据判定 S/A/B/C/D:
-   - 需求词自身:`search_related_pool_demands` / `query_demand_popularity_by_word`
+   - 需求词自身:交接件 `items[].demand_priority` / `search_related_pool_demands` / `query_demand_popularity_by_word`
    - 归属分类节点:`query_demand_category_and_weight`
-   - 局部环境(父节点 + 全部兄弟节点完整权重):`query_category_local_heat`
-   reason 必须同时写清全局位置、局部判断和词级后验(如“节点 total_score=3.2、兄弟 4/4 中排名第 1、父节点高热,无后验数据,因此上调至 A”);`related_pool_ids` 取自 `search_related_pool_demands` 返回的 `[id=...]`。
+   - 全局与局部环境(整树名次 + 父节点 + 全部兄弟节点完整权重):`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` 并落库;禁止另造一套模型分。
 6. 简要汇报本批次的分级结果后结束本轮任务。
 
 ## 原则
@@ -81,3 +92,4 @@
 - reason 必须具体:写清引用的全局热度/后验数值、是否合并了同义词、依据哪个树节点。
 - 找不到归属树节点或权重数据的需求:如实说明"无法评级/数据缺失",不要强行给出等级去凑数。
 - 有后验数据始终优先于纯全局热度判断;无后验数据时保持谨慎,不给最高档。
+- 批次等级和需求等级是两个不同结论;禁止把 S 热批次里的所有需求直接评为 S。

+ 4 - 2
agents/demand_grade_agent/run.py

@@ -36,8 +36,10 @@ def main(
     需求词列表:
     {names_str}
 
-    以下是统筹规划提供的节点组上下文(含 `local_heat` 局部热度快照)。
-    必须结合父节点、兄弟节点排名做局部校正;需要更细信息时可调用 `query_category_local_heat`:
+    以下是统筹规划提供的节点组上下文(含批次热度等级、需求自身来源归一分、
+    `local_heat` 局部热度与整树名次快照)。批次等级不等于需求最终等级。
+    必须结合整树位置、父节点、全部兄弟节点和需求自身分做校正;需要更细信息时可调用
+    `query_category_local_heat`。不同来源原始 weight 不得直接相加:
     {handoff}
     """
     result = agent.run(user_input)

+ 14 - 3
agents/demand_grade_agent/tools/batch_save_demand_grades.py

@@ -7,6 +7,7 @@ import logging
 from decimal import Decimal
 from typing import Any, Optional
 
+from agents.demand_grade_agent.tools.demand_priority import build_demand_priority_index
 from agents.demand_grade_agent.tools.shared import (
     VALID_GRADES,
     collect_strategies,
@@ -107,7 +108,6 @@ def _normalize_items(
 
         seen_keys.add(dedupe_key)
 
-        score = _optional_decimal(item.get("score"), "score", idx, errors)
         prior_raw = item.get("prior_total_score")
         if prior_raw is None or prior_raw == "" or prior_raw == "—":
             prior_total_score = None
@@ -134,7 +134,8 @@ def _normalize_items(
                 "demand_name": demand_name,
                 "category_ids": dump_int_list(category_ids),
                 "grade": grade,
-                "score": score,
+                # 保存阶段会基于当日全量需求池确定性重算,禁止由模型自由填写。
+                "score": None,
                 "prior_total_score": prior_total_score,
                 "posterior_rov_avg": posterior_rov_avg,
                 "posterior_rov_count": posterior_rov_count,
@@ -165,7 +166,8 @@ def batch_save_demand_grades(items: list[dict[str, Any]], biz_dt: Optional[str]
             - reason (必填): 判断依据,需引用具体的先验/后验数值
             - related_pool_ids (必填): 该需求对应的 multi_demand_pool_di.id 列表,需先调用
               search_related_pool_demands 找到;用于关联原始需求,并自动推导 video_list/strategies
-            - score (可选): 0-100 数值分,辅助同级排序
+            - score: 无需传入;保存时按当日全量需求池自动计算需求自身来源归一分(0-100),
+              即各 strategy 内独立排名归一化后,对该需求已有来源取均值
             - category_ids (可选): 归属的树节点 id 列表,会写入 demand_grade_category_rel 映射表
             - prior_total_score (可选): 落库时的先验 total_score 快照
             - posterior_rov_avg / posterior_rov_count (可选): 落库时的后验 real_rov_7d 快照;
@@ -196,6 +198,10 @@ def batch_save_demand_grades(items: list[dict[str, Any]], biz_dt: Optional[str]
             all_pool_ids = sorted({pid for ids in related_pool_id_lists for pid in ids})
             pool_rows = pool_repo.get_by_ids(all_pool_ids) if all_pool_ids else []
             pool_by_id = {int(r.id): r for r in pool_rows}
+            priority_by_biz_dt = {
+                resolved_dt: build_demand_priority_index(pool_repo.list_by_biz_dt(resolved_dt))
+                for resolved_dt in sorted({row["biz_dt"] for row in rows})
+            }
 
             final_rows: list[dict[str, Any]] = []
             saved_indices: list[int] = []
@@ -212,6 +218,11 @@ def batch_save_demand_grades(items: list[dict[str, Any]], biz_dt: Optional[str]
                     errors.append(
                         f"demand_name={row['demand_name']!r} 的 related_pool_ids 中 {missing} 未找到,已忽略"
                     )
+                priority = priority_by_biz_dt[row["biz_dt"]].get(row["demand_name"])
+                source_rank_score = priority.get("source_rank_score") if priority else None
+                row["score"] = (
+                    Decimal(str(source_rank_score)) if source_rank_score is not None else None
+                )
                 row["video_list"] = merge_video_ids(matched)
                 row["strategies"] = collect_strategies(matched)
                 final_rows.append(row)

+ 127 - 7
agents/demand_grade_agent/tools/build_grade_plan_context.py

@@ -3,6 +3,10 @@ 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
@@ -11,6 +15,100 @@ from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPo
 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,
@@ -29,21 +127,43 @@ def build_grade_plan_context(
             [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
+        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} for name in names],
+        "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": shared_traits.strip(),
-            "selection_method": "global_tree_heat_first",
+            "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)

+ 160 - 0
agents/demand_grade_agent/tools/demand_priority.py

@@ -0,0 +1,160 @@
+"""需求自身先验分:来源内排名归一化后再形成可比分。"""
+from __future__ import annotations
+
+from collections import defaultdict
+from typing import Any
+
+from supply_agent.ranking import rank_with_scores
+
+
+DEMAND_PRIORITY_SCORE_METHOD = {
+    "name": "source_rank_mean_v1",
+    "range": "0-100",
+    "steps": [
+        "同一需求在同一来源内的非零 weight 先取平均",
+        "每个来源内部独立按 weight 降序排名并归一化到 (0,1]",
+        "同一需求已有来源的归一分等权平均后乘 100",
+    ],
+    "missing_source_policy": "未出现的来源不补零",
+    "warning": "不同来源的原始 weight 不可直接相加;本分也不可与分类节点 total_score 直接相加",
+}
+
+
+def _average(values: list[float]) -> float | None:
+    return sum(values) / len(values) if values else None
+
+
+def build_demand_priority_index(pool_rows: list[Any]) -> dict[str, dict[str, Any]]:
+    """基于一天的完整需求池,构建需求自身的来源可比先验分。"""
+    weights: dict[tuple[str, str], list[float]] = defaultdict(list)
+    rows_by_name: dict[str, list[Any]] = defaultdict(list)
+    observed_strategies: dict[str, set[str]] = defaultdict(set)
+    for row in pool_rows:
+        demand_name = str(getattr(row, "demand_name", "") or "").strip()
+        strategy = str(getattr(row, "strategy", "") or "").strip()
+        if not demand_name:
+            continue
+        rows_by_name[demand_name].append(row)
+        if strategy:
+            observed_strategies[demand_name].add(strategy)
+        raw_weight = getattr(row, "weight", None)
+        if strategy and raw_weight is not None and float(raw_weight) != 0:
+            weights[(strategy, demand_name)].append(float(raw_weight))
+
+    source_averages = {
+        key: float(_average(values))
+        for key, values in weights.items()
+        if values
+    }
+    candidates_by_source: dict[str, list[tuple[str, float]]] = defaultdict(list)
+    for (strategy, demand_name), avg in source_averages.items():
+        candidates_by_source[strategy].append((demand_name, avg))
+    positions_by_source = {
+        strategy: rank_with_scores(candidates)
+        for strategy, candidates in candidates_by_source.items()
+    }
+
+    index: dict[str, dict[str, Any]] = {}
+    for demand_name, demand_rows in rows_by_name.items():
+        sources: list[dict[str, Any]] = []
+        for strategy in sorted(observed_strategies[demand_name]):
+            avg = source_averages.get((strategy, demand_name))
+            position = positions_by_source.get(strategy, {}).get(demand_name)
+            source_row_count = sum(
+                1
+                for row in demand_rows
+                if str(getattr(row, "strategy", "") or "").strip() == strategy
+            )
+            sources.append({
+                "strategy": strategy,
+                "raw_weight_avg": avg,
+                "raw_weight_row_count": source_row_count if avg is not None else 0,
+                "rank_in_source": float(position["rank"]) if position is not None else None,
+                "source_demand_count": int(position["total"]) if position is not None else 0,
+                "source_normalized_rank_score": (
+                    float(position["normalized_score"]) if position is not None else None
+                ),
+            })
+
+        normalized_scores = [
+            float(item["source_normalized_rank_score"])
+            for item in sources
+            if item["source_normalized_rank_score"] is not None
+        ]
+        prior_score = (
+            100 * sum(normalized_scores) / len(normalized_scores)
+            if normalized_scores
+            else None
+        )
+        rov_values = [
+            float(row.real_rov_7d)
+            for row in demand_rows
+            if getattr(row, "real_rov_7d", None) is not None
+        ]
+        vov_values = [
+            float(row.real_vov_7d)
+            for row in demand_rows
+            if getattr(row, "real_vov_7d", None) is not None
+        ]
+        index[demand_name] = {
+            "demand_name": demand_name,
+            "source_rank_score": round(prior_score, 4) if prior_score is not None else None,
+            "valid_source_count": len(normalized_scores),
+            "observed_source_count": len(sources),
+            "sources": sources,
+            "posterior_from_exact_pool_rows": {
+                "real_rov_7d": {
+                    "value": max(rov_values) if rov_values else None,
+                    "has_data": bool(rov_values),
+                },
+                "real_vov_7d": {
+                    "value": max(vov_values) if vov_values else None,
+                    "has_data": bool(vov_values),
+                },
+            },
+            "exact_pool_ids": sorted({int(row.id) for row in demand_rows}),
+        }
+
+    demand_positions = rank_with_scores([
+        (demand_name, float(item["source_rank_score"]))
+        for demand_name, item in index.items()
+        if item["source_rank_score"] is not None
+    ])
+    for demand_name, item in index.items():
+        position = demand_positions.get(demand_name)
+        item["global_demand_rank"] = float(position["rank"]) if position is not None else None
+        item["global_scored_demand_count"] = int(position["total"]) if position is not None else 0
+        item["score_method"] = DEMAND_PRIORITY_SCORE_METHOD["name"]
+    return index
+
+
+def format_demand_priority(item: dict[str, Any] | None) -> list[str]:
+    """将需求自身排名证据格式化为 Agent 易读文本。"""
+    if item is None:
+        return ["需求自身来源归一分=—(无需求池记录)"]
+    score = item.get("source_rank_score")
+    rank = item.get("global_demand_rank")
+    total = int(item.get("global_scored_demand_count") or 0)
+    score_text = "—" if score is None else f"{float(score):.2f}/100"
+    rank_text = "无排名" if rank is None else f"{float(rank):g}/{total}"
+    lines = [
+        f"需求自身来源归一分={score_text};全日需求自身排名={rank_text};"
+        f"有效来源={int(item.get('valid_source_count') or 0)}",
+    ]
+    for source in item.get("sources") or []:
+        raw = source.get("raw_weight_avg")
+        source_rank = source.get("rank_in_source")
+        normalized = source.get("source_normalized_rank_score")
+        raw_text = "—" if raw is None else f"{float(raw):.4f}"
+        source_rank_text = (
+            "—"
+            if source_rank is None
+            else f"{float(source_rank):g}/{source['source_demand_count']}"
+        )
+        normalized_text = "—" if normalized is None else f"{float(normalized):.4f}"
+        lines.append(
+            f"  来源={source['strategy']} raw_weight_avg={raw_text} "
+            f"来源内名次={source_rank_text} 来源归一分={normalized_text}"
+        )
+    lines.append("  口径:来源内先排名归一化,再对已有来源取均值;禁止直接相加原始 weight。")
+    return lines

+ 25 - 5
agents/demand_grade_agent/tools/query_score_distribution.py

@@ -1,14 +1,17 @@
-"""
-查询指定业务日 category_tree_weight 的分数分布,供批量分级前统一分档阈值。
-"""
+"""查询分类树、需求自身与后验的独立分布,供批量分级统一口径。"""
 from __future__ import annotations
 
 import logging
 from typing import Optional
 
+from agents.demand_grade_agent.tools.demand_priority import (
+    DEMAND_PRIORITY_SCORE_METHOD,
+    build_demand_priority_index,
+)
 from agents.demand_grade_agent.tools.shared import distribution_summary, normalize_biz_dt, to_float
 from supply_agent.tools import tool
 from supply_infra.db.repositories.category_tree_weight_repo import CategoryTreeWeightRepository
+from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
 from supply_infra.db.session import get_session
 
 logger = logging.getLogger(__name__)
@@ -31,10 +34,11 @@ def _format_dist(label: str, dist: dict) -> str:
 @tool
 def query_score_distribution(biz_dt: Optional[str] = None) -> str:
     """
-    查询指定业务日 category_tree_weight 的全局热度与后验分布。
+    查询指定业务日分类树全局热度、需求自身来源归一分与后验分布。
 
-    建议在批量分级任务开始时调用一次,参考分位数自行制定本批次统一的分档阈值
+    建议在批量分级任务开始时调用一次,分别参考各自分位数制定本批次统一的分档阈值
     (例如 total_score 前 10% 视为全局热度很高),避免同一批次内多次判断标准漂移。
+    需求自身分按来源内 rank 归一后对已有来源取均值,不直接合并跨来源 raw weight;
     后验 real_rov_7d_avg 的分布只统计 real_rov_7d_count>0(有真实验证数据)的子集,
     因为无验证数据的行 avg 无意义。
 
@@ -84,6 +88,14 @@ def query_score_distribution(biz_dt: Optional[str] = None) -> str:
                 for v in (to_float(w.real_rov_7d_avg) for w in weights if w.real_rov_7d_count > 0)
                 if v is not None
             ]
+            demand_priority_index = build_demand_priority_index(
+                MultiDemandPoolDiRepository(session).list_by_biz_dt(resolved_dt)
+            )
+            demand_priority_values = [
+                float(item["source_rank_score"])
+                for item in demand_priority_index.values()
+                if item["source_rank_score"] is not None
+            ]
 
         lines = [f"biz_dt={resolved_dt} 共 {node_count} 个树节点"]
         for field, label in _DIM_FIELDS:
@@ -95,6 +107,14 @@ def query_score_distribution(biz_dt: Optional[str] = None) -> str:
                 distribution_summary(posterior_values),
             )
         )
+        lines.append(
+            _format_dist(
+                "需求自身来源归一分(0-100,来源内排名后对已有来源取均值)",
+                distribution_summary(demand_priority_values),
+            )
+        )
+        lines.append(f"需求自身分口径={DEMAND_PRIORITY_SCORE_METHOD['name']};不同来源原始 weight 禁止直接相加。")
+        lines.append("注意:需求自身来源归一分与分类树 total_score 是两类独立证据,不得直接相加或共用阈值。")
 
         message = "\n".join(lines)
         logger.info("query_score_distribution completed: biz_dt=%s nodes=%d", resolved_dt, node_count)

+ 16 - 1
agents/demand_grade_agent/tools/search_related_pool_demands.py

@@ -7,6 +7,10 @@ import logging
 
 from sqlalchemy.orm import Session
 
+from agents.demand_grade_agent.tools.demand_priority import (
+    build_demand_priority_index,
+    format_demand_priority,
+)
 from agents.demand_grade_agent.tools.shared import (
     format_score,
     normalize_biz_dt,
@@ -23,6 +27,7 @@ def _search_one_related_pool_demands(
     session: Session,
     normalized: str,
     keyword: str,
+    priority_index: dict[str, dict],
 ) -> str:
     rows = MultiDemandPoolDiRepository(session).search_rows_by_name_fragment(normalized, keyword)
 
@@ -30,6 +35,9 @@ def _search_one_related_pool_demands(
         return f"biz_dt={normalized} 未找到与「{keyword}」同名/包含关系的需求词"
 
     lines = []
+    for demand_name in dict.fromkeys(str(row["demand_name"]) for row in rows):
+        lines.append(f"需求自身证据「{demand_name}」:")
+        lines.extend(format_demand_priority(priority_index.get(demand_name)))
     for row in rows:
         rov = format_score(row["real_rov_7d"])
         vov = format_score(row["real_vov_7d"])
@@ -75,9 +83,16 @@ def search_related_pool_demands(biz_dt: str, keywords: list[str]) -> str:
 
     try:
         with get_session() as session:
+            pool_repo = MultiDemandPoolDiRepository(session)
+            priority_index = build_demand_priority_index(pool_repo.list_by_biz_dt(normalized))
             sections: list[str] = []
             for keyword in keyword_list:
-                result = _search_one_related_pool_demands(session, normalized, keyword)
+                result = _search_one_related_pool_demands(
+                    session,
+                    normalized,
+                    keyword,
+                    priority_index,
+                )
                 sections.append(f"--- keyword: {keyword} ---\n{result}")
 
         message = "\n\n".join(sections)

+ 47 - 4
agents/demand_grade_agent/tools/tree_local.py

@@ -10,6 +10,7 @@ from agents.demand_grade_agent.tools.shared import (
     build_category_path,
     format_category_weight_lines,
 )
+from supply_agent.ranking import rank_with_scores
 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
@@ -77,15 +78,24 @@ def _describe_node(
     weight_by_id: dict[int, Any],
     *,
     biz_dt: str,
+    global_positions: dict[int, dict[str, float | int]],
 ) -> dict[str, Any]:
     category = by_id.get(category_id)
     weight = weight_by_id.get(category_id)
+    position = global_positions.get(category_id)
     return {
         "category_id": category_id,
         "name": category.name if category is not None else None,
         "path": build_category_path(category_id, by_id),
         "hung_word_count": int(weight.hung_word_count or 0) if weight is not None else 0,
         "heat": build_category_heat_summary(weight, biz_dt=biz_dt),
+        "global_tree_position": {
+            "rank": float(position["rank"]) if position is not None else None,
+            "scored_node_count": int(position["total"]) if position is not None else 0,
+            "normalized_rank_score": (
+                float(position["normalized_score"]) if position is not None else None
+            ),
+        },
     }
 
 
@@ -107,6 +117,11 @@ def build_local_heat_snapshot(
     else:
         by_id, children, weight_by_id = tree_state
     snapshots: list[dict[str, Any]] = []
+    global_positions = rank_with_scores([
+        (category_id, float(weight.total_score))
+        for category_id, weight in weight_by_id.items()
+        if weight is not None and weight.total_score is not None
+    ])
     for category_id in dict.fromkeys(int(value) for value in category_ids):
         category = by_id.get(category_id)
         if category is None:
@@ -129,7 +144,13 @@ def build_local_heat_snapshot(
         sibling_total = len(ranked_siblings)
         siblings = []
         for index, sibling_id in enumerate(ranked_siblings, start=1):
-            item = _describe_node(sibling_id, by_id, weight_by_id, biz_dt=biz_dt)
+            item = _describe_node(
+                sibling_id,
+                by_id,
+                weight_by_id,
+                biz_dt=biz_dt,
+                global_positions=global_positions,
+            )
             item["rank_among_siblings"] = index
             item["sibling_count"] = sibling_total
             item["is_self"] = sibling_id == category_id
@@ -138,9 +159,22 @@ def build_local_heat_snapshot(
         snapshots.append({
             "category_id": category_id,
             "biz_dt": biz_dt,
-            "self": _describe_node(category_id, by_id, weight_by_id, biz_dt=biz_dt),
+            "global_scored_node_count": len(global_positions),
+            "self": _describe_node(
+                category_id,
+                by_id,
+                weight_by_id,
+                biz_dt=biz_dt,
+                global_positions=global_positions,
+            ),
             "parent": (
-                _describe_node(parent_id, by_id, weight_by_id, biz_dt=biz_dt)
+                _describe_node(
+                    parent_id,
+                    by_id,
+                    weight_by_id,
+                    biz_dt=biz_dt,
+                    global_positions=global_positions,
+                )
                 if parent_id is not None
                 else None
             ),
@@ -171,6 +205,15 @@ def _append_node_block(
         biz_dt=biz_dt,
         hung_word_count=int(node.get("hung_word_count") or 0),
     ))
+    position = node.get("global_tree_position") or {}
+    rank = position.get("rank")
+    if rank is None:
+        lines.append("  整树位置=无排名(不是低热,表示缺少 total_score)")
+    else:
+        lines.append(
+            f"  整树位置={float(rank):g}/{int(position.get('scored_node_count') or 0)} "
+            f"整树排名归一分={float(position['normalized_rank_score']):.4f}"
+        )
 
 
 def format_local_heat_snapshot(snapshot: dict[str, Any], *, weight_by_id: dict[int, Any] | None = None) -> list[str]:
@@ -200,7 +243,7 @@ def render_local_heat_report(biz_dt: str, category_ids: list[int]) -> str:
         return "category_ids 不能为空"
     lines = [
         f"biz_dt={biz_dt} | 局部环境包含节点自身、父节点、全部兄弟节点的全局热度 total_score 与后验数据;",
-        "不得只参考部分兄弟节点,应以本报告列出的全部节点为准。",
+        "每个节点同时给出整棵树名次;不得只参考附近节点,也不得只参考部分兄弟节点。",
     ]
     for snapshot in snapshots:
         lines.extend(format_local_heat_snapshot(snapshot, weight_by_id=weight_by_id))

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

@@ -2,16 +2,24 @@
 
 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 (
+    HEAT_LEVEL_DEFINITION,
     format_heat_score,
+    format_rank,
+    global_heat_positions,
     has_hung_demand,
+    heat_level,
     load_tree_state,
     path,
 )
 
 __all__ = [
     "build_grade_plan_for_category_ids",
+    "HEAT_LEVEL_DEFINITION",
     "format_heat_score",
+    "format_rank",
+    "global_heat_positions",
     "has_hung_demand",
+    "heat_level",
     "load_tree_state",
     "path",
 ]

+ 180 - 33
agents/demand_grade_orchestrator_agent/common/plan_builder.py

@@ -1,17 +1,115 @@
-"""按分类节点构建统筹分级计划。"""
+"""按分类节点构建带批次热度等级的统筹分级计划。"""
 from __future__ import annotations
 
-from collections import defaultdict
+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],
@@ -19,46 +117,95 @@ def build_grade_plan_for_category_ids(
     *,
     max_nodes_per_group: int = 4,
 ) -> dict[str, Any]:
-    """为指定分类节点生成分组计划,不重复包含目标集之外的节点。"""
-    targets = {int(category_id) for category_id in category_ids}
+    """生成批次计划:先分热度等级,再把同等级相邻节点合为一批。"""
+    requested = {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)
-
+    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))
-    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}",
+
+    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,
-                "planning_reason": f"同属「{parent_path}」分支,按全局热度分从高到低编排({heat_scores})。",
-                "shared_traits": f"共同父节点={parent_path};{grouping_strategy.strip() or '按同分支、全局热度分顺序分组'}",
+                "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": sorted(targets - set(covered)),
-        "coverage_complete": not (targets - set(covered)),
+        "uncovered_category_ids": uncovered,
+        "coverage_complete": not uncovered,
     }

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

@@ -5,6 +5,7 @@ from collections import defaultdict
 from types import SimpleNamespace
 from typing import Any
 
+from supply_agent.ranking import rank_with_scores
 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
@@ -42,6 +43,45 @@ def has_hung_demand(weight: Any | None) -> bool:
     return weight is not None and int(weight.hung_word_count or 0) > 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},
+    "B": {"label": "中热", "min_global_rank_score": 0.50},
+    "C": {"label": "较低热", "min_global_rank_score": 0.25},
+    "D": {"label": "低热", "min_global_rank_score": 0.00},
+    "U": {"label": "数据不足", "min_global_rank_score": None},
+}
+
+
+def global_heat_positions(weights: dict[int, Any]) -> dict[int, dict[str, float | int]]:
+    """返回节点 ``total_score`` 在整棵有分节点中的排名位置。"""
+    return rank_with_scores([
+        (category_id, float(weight.total_score))
+        for category_id, weight in weights.items()
+        if weight is not None and weight.total_score is not None
+    ])
+
+
+def heat_level(position: dict[str, float | int] | None) -> str:
+    """按整树排名分映射批次热度等级;无分节点单列 U。"""
+    if position is None:
+        return "U"
+    score = float(position["normalized_score"])
+    for level in ("S", "A", "B", "C"):
+        threshold = HEAT_LEVEL_DEFINITION[level]["min_global_rank_score"]
+        if threshold is not None and score >= float(threshold):
+            return level
+    return "D"
+
+
+def format_rank(position: dict[str, float | int] | None) -> str:
+    if position is None:
+        return "无排名"
+    rank = float(position["rank"])
+    rank_text = str(int(rank)) if rank.is_integer() else f"{rank:.1f}"
+    return f"{rank_text}/{int(position['total'])}"
+
+
 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

+ 4 - 2
agents/demand_grade_orchestrator_agent/prompt/system_prompt.md

@@ -5,13 +5,15 @@
 ## 固定工作流
 
 1. 先调用 `query_global_heat_tree(biz_dt)`,查看有需求节点所在的全局树及其原始 `total_score` 热度分。不得跳过这一步。树中 `+` 仅表示该节点自身有需求,`null` 表示没有热度数据。
-2. 选择一个热点节点,或多个具有共同路径、相近热度或父节点背景的节点组成节点组;调用 `query_heat_node_group(biz_dt, category_ids)` 下钻验证,可多轮下钻。
-3. 总结该节点组的 `planning_reason` 与 `shared_traits`:必须说明其全局位置、父/兄弟关系,以及为什么这些节点适合被同一批处理。
+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 是当天完整执行计划。
 
 ## 规划原则
 
 - 优先规划高热节点簇、同父节点下共同上升的兄弟节点,或“局部突发但大盘偏冷”的对照节点组。
+- 每个批次必须有明确的 `batch_heat_level`:S/A/B/C/D 来自节点在整棵有分分类树中的排名分,U 表示热度数据不足。批次等级用于区分热度,不要求下游严格按等级串行执行。
+- 相邻/相似节点只有在同一热度等级时才合批,避免一个极热节点把冷节点所在批次整体抬高。
 - 下游会逐组解析节点下的待分级需求;不要反过来根据需求名称拼凑批次。
 - 热度为空或样本数为 0 是数据不足,不是低热;在规划原因中明确说明。
 - 不自行评级、不调用落库工具。只依据工具真实返回的路径、节点和分数。

+ 1 - 1
agents/demand_grade_orchestrator_agent/tools/build_full_day_grade_plan.py

@@ -10,7 +10,7 @@ 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,完整任务明细将单独落表。"""
+    """按整树热度等级和树邻接关系生成全天节点组;每组明确标注 S/A/B/C/D/U。"""
     payload = build_grade_plan_for_category_ids(
         biz_dt,
         sorted(get_required_hanging_category_ids(biz_dt)),

+ 16 - 4
agents/demand_grade_orchestrator_agent/tools/query_global_heat_tree.py

@@ -1,7 +1,14 @@
 """以层级文本展示当天有需求分类所在的全局树及原始热度分。"""
 from __future__ import annotations
 
-from agents.demand_grade_orchestrator_agent.common import format_heat_score, has_hung_demand, load_tree_state
+from agents.demand_grade_orchestrator_agent.common import (
+    format_heat_score,
+    format_rank,
+    global_heat_positions,
+    has_hung_demand,
+    heat_level,
+    load_tree_state,
+)
 from supply_agent.tools import tool
 
 
@@ -9,6 +16,7 @@ from supply_agent.tools import 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)}
     visible_cache: dict[int, bool] = {}
 
@@ -20,8 +28,8 @@ def query_global_heat_tree(biz_dt: str) -> str:
         return value
 
     lines = [
-        f"biz_dt={biz_dt} | 格式:[分类ID]分类名称[全局热度分] + | "
-        "全局热度分为 total_score,保留两位小数;无数据为 null;+ 表示该分类自身有需求"
+        f"biz_dt={biz_dt} | 格式:[分类ID]分类名称[total_score|整树名次|热度等级] + | "
+        "名次在整棵有分节点中计算;无数据为 null/U;+ 表示该分类自身有需求"
     ]
 
     def render(category_id: int, depth: int) -> None:
@@ -30,8 +38,12 @@ def query_global_heat_tree(biz_dt: str) -> str:
         category = by_id[category_id]
         weight = weights.get(category_id)
         score_text = format_heat_score(weight)
+        position = positions.get(category_id)
         suffix = " +" if category_id in demand_nodes else ""
-        lines.append(f"{'  ' * depth}[{category_id}]{category.name or ''}[{score_text}]{suffix}")
+        lines.append(
+            f"{'  ' * depth}[{category_id}]{category.name or ''}"
+            f"[{score_text}|{format_rank(position)}|{heat_level(position)}]{suffix}"
+        )
         for child in children.get(category_id, []):
             render(child, depth + 1)
 

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

@@ -1,24 +1,37 @@
-"""下钻节点,查看父子/兄弟间的原始全局热度。"""
+"""下钻节点,查看整树位置及父子/兄弟间的全局热度。"""
 from __future__ import annotations
 
-from agents.demand_grade_orchestrator_agent.common import format_heat_score, has_hung_demand, load_tree_state, path
+from agents.demand_grade_orchestrator_agent.common import (
+    format_heat_score,
+    format_rank,
+    global_heat_positions,
+    has_hung_demand,
+    heat_level,
+    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)
+    positions = global_heat_positions(weights)
     lines = [
-        f"biz_dt={biz_dt} | 格式:[分类ID]分类名称[全局热度分] + | "
-        "全局热度分为 total_score,保留两位小数;无数据为 null;+ 表示该分类自身有需求"
+        f"biz_dt={biz_dt} | 格式:[分类ID]分类名称[total_score|整树名次|热度等级] + | "
+        "名次在整棵有分节点中计算;无数据为 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 ""
-        return f"[{category_id}]{category.name or ''}[{format_heat_score(weight)}]{suffix}"
+        return (
+            f"[{category_id}]{category.name or ''}"
+            f"[{format_heat_score(weight)}|{format_rank(position)}|{heat_level(position)}]{suffix}"
+        )
 
     for category_id in dict.fromkeys(int(value) for value in category_ids):
         category = by_id.get(category_id)
@@ -29,6 +42,8 @@ def query_heat_node_group(biz_dt: str, category_ids: list[int]) -> str:
         lines.append(f"  路径:{path(category_id, by_id)}")
         if parent_id is not None:
             lines.append(f"  父节点:{describe(parent_id)}")
+        sibling_text = "、".join(describe(sibling) for sibling in children.get(parent_id, []))
+        lines.append(f"  全部兄弟节点(含自身):{sibling_text or '无'}")
         child_text = "、".join(describe(child) for child in children.get(category_id, []))
         lines.append(f"  直接子节点:{child_text or '无'}")
     return "\n".join(lines)

+ 47 - 0
supply_agent/ranking.py

@@ -0,0 +1,47 @@
+"""跨业务模块复用的排名归一化工具。"""
+from __future__ import annotations
+
+from collections.abc import Hashable
+from typing import TypeVar
+
+
+KeyT = TypeVar("KeyT", bound=Hashable)
+
+
+def rank_with_scores(items: list[tuple[KeyT, float]]) -> dict[KeyT, dict[str, float | int]]:
+    """按值降序排名,并返回同分平均名次与 ``(0, 1]`` 排名分。
+
+    不同量纲的数据应先分别调用本函数,不能先把原始值相加。排名分公式与
+    ``category_tree_weight`` 的四维全局排名口径一致:
+    ``score = (n - avg_rank + 1) / n``。
+    """
+    if not items:
+        return {}
+
+    sorted_items = sorted(items, key=lambda item: (-item[1], str(item[0])))
+    total = len(sorted_items)
+    result: dict[KeyT, dict[str, float | int]] = {}
+    start = 0
+    while start < total:
+        end = start
+        value = sorted_items[start][1]
+        while end < total and sorted_items[end][1] == value:
+            end += 1
+        average_rank = (start + 1 + end) / 2.0
+        normalized_score = (total - average_rank + 1) / total
+        for index in range(start, end):
+            result[sorted_items[index][0]] = {
+                "rank": average_rank,
+                "total": total,
+                "normalized_score": normalized_score,
+            }
+        start = end
+    return result
+
+
+def rank_to_scores(items: list[tuple[KeyT, float]]) -> dict[KeyT, float]:
+    """兼容原调用方:只返回排名归一分。"""
+    return {
+        key: float(position["normalized_score"])
+        for key, position in rank_with_scores(items).items()
+    }

+ 3 - 1
supply_infra/db/models/demand_grade.py

@@ -25,7 +25,9 @@ class DemandGrade(Base):
     )
     grade: Mapped[str] = mapped_column(String(4), nullable=False, comment="等级 S/A/B/C/D")
     score: Mapped[Decimal | None] = mapped_column(
-        Numeric(6, 2), nullable=True, comment="数值分(可选,辅助同级排序)"
+        Numeric(6, 2),
+        nullable=True,
+        comment="需求自身来源归一分0-100(来源内排名后对已有来源取均值)",
     )
     prior_total_score: Mapped[Decimal | None] = mapped_column(
         Numeric(16, 8), nullable=True, comment="落库时的先验 total_score 快照"

+ 61 - 4
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 select, update
+from sqlalchemy import case, select, update
 
 from supply_infra.db.models.demand_grade_plan import DemandGradePlan, DemandGradePlanGroup
 from supply_infra.db.repositories.base import BaseRepository
@@ -45,11 +45,14 @@ class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
         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)
+        unfinished_groups = claimable_groups + group_status.get("running", 0)
         return {
             "planned_groups": len(groups),
             "assigned_category_ids": assigned_category_ids,
             "group_status": group_status,
             "claimable_groups": claimable_groups,
+            "unfinished_groups": unfinished_groups,
+            "execution_complete": unfinished_groups == 0,
         }
 
     def create_plan(self, biz_dt: str, payload: dict[str, Any]) -> None:
@@ -62,6 +65,8 @@ class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
             "grouping_strategy": payload.get("grouping_strategy"),
             "total_hanging_nodes": payload.get("total_hanging_nodes"),
             "group_count": len(groups),
+            "heat_level_definition": payload.get("heat_level_definition", {}),
+            "batch_heat_level_counts": payload.get("batch_heat_level_counts", {}),
             "coverage_complete": coverage_complete,
             "covered_category_ids": payload.get("covered_category_ids", []),
             "uncovered_category_ids": payload.get("uncovered_category_ids", []),
@@ -73,30 +78,78 @@ class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
             plan_json=json.dumps(summary, ensure_ascii=False),
         ))
         for group_no, group in enumerate(groups, start=1):
+            shared_traits = json.dumps({
+                "description": str(group["shared_traits"]),
+                "batch_heat_level": group.get("batch_heat_level"),
+                "batch_heat_label": group.get("batch_heat_label"),
+                "batch_global_rank_score": group.get("batch_global_rank_score"),
+                "batch_total_score_avg": group.get("batch_total_score_avg"),
+                "category_global_positions": group.get("category_global_positions", []),
+            }, ensure_ascii=False)
             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"]),
+                planning_reason=str(group["planning_reason"]), shared_traits=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)
+            .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 []
+        stmt = (
+            select(DemandGradePlanGroup.id)
+            .where(
+                DemandGradePlanGroup.biz_dt == biz_dt,
+                DemandGradePlanGroup.status == status,
+            )
+            .order_by(DemandGradePlanGroup.id)
+        )
+        return [int(group_id) for group_id in self.session.scalars(stmt).all()]
+
+    def claim_group(self, biz_dt: str, group_id: int) -> dict[str, Any] | None:
+        """按固定 ID 原子领取任务;状态已变化时返回 None。"""
+        stmt = (
+            select(DemandGradePlanGroup)
+            .where(
+                DemandGradePlanGroup.id == int(group_id),
+                DemandGradePlanGroup.biz_dt == biz_dt,
+                DemandGradePlanGroup.status.in_(("pending", "failed")),
+            )
+            .with_for_update(skip_locked=True)
+        )
+        group = self.session.scalar(stmt)
+        if group is None:
+            return None
+        return self._mark_claimed(group)
+
+    @staticmethod
+    def _mark_claimed(group: DemandGradePlanGroup) -> dict[str, Any]:
         group.status = "running"
         group.attempts += 1
         group.started_at = datetime.now()
+        group.error_message = None
+        group.finished_at = None
         return {
             "id": int(group.id),
             "group_key": group.group_key,
@@ -109,7 +162,11 @@ class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
         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())
+            .values(
+                status="finished" if success else "failed",
+                error_message=error_message,
+                finished_at=datetime.now(),
+            )
         )
 
     def summarize(self, biz_dt: str) -> dict[str, int]:

+ 15 - 2
supply_infra/db/repositories/multi_demand_pool_di_repo.py

@@ -82,6 +82,15 @@ class MultiDemandPoolDiRepository(BaseRepository[MultiDemandPoolDi]):
         )
         return int(self.session.scalar(stmt) or 0)
 
+    def list_by_biz_dt(self, biz_dt: str) -> list[MultiDemandPoolDi]:
+        """返回业务日全部需求池行,供来源内排名等全局计算使用。"""
+        stmt = (
+            select(MultiDemandPoolDi)
+            .where(MultiDemandPoolDi.biz_dt == biz_dt)
+            .order_by(MultiDemandPoolDi.strategy, MultiDemandPoolDi.demand_name, MultiDemandPoolDi.id)
+        )
+        return list(self.session.scalars(stmt).all())
+
     def search_rows_by_name_fragment(self, biz_dt: str, keyword: str) -> list[dict]:
         """
         按业务日 + 需求名双向包含关系搜索明细行。
@@ -129,8 +138,12 @@ class MultiDemandPoolDiRepository(BaseRepository[MultiDemandPoolDi]):
         """按 id 批量查询完整行。"""
         if not ids:
             return []
-        stmt = select(MultiDemandPoolDi).where(MultiDemandPoolDi.id.in_(ids))
-        return list(self.session.scalars(stmt).all())
+        rows: list[MultiDemandPoolDi] = []
+        for start in range(0, len(ids), _BATCH_SIZE):
+            batch = ids[start : start + _BATCH_SIZE]
+            stmt = select(MultiDemandPoolDi).where(MultiDemandPoolDi.id.in_(batch))
+            rows.extend(self.session.scalars(stmt).all())
+        return rows
 
     def count_by_biz_dt(self, biz_dt: str) -> int:
         """统计指定业务日期去重行数(strategy + demand_id)。"""

+ 398 - 63
supply_infra/scheduler/jobs/grade_demand_pool.py

@@ -1,19 +1,30 @@
-"""统筹规划落库后,由多个 worker 领取节点组任务并调用分级 Agent。"""
+"""统筹落库后执行分级任务,并对任务状态与树上需求等级做补偿闭环。"""
 from __future__ import annotations
 
 import json
 import logging
 from concurrent.futures import ThreadPoolExecutor, as_completed
 from datetime import datetime
+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 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_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
 
@@ -21,6 +32,9 @@ 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:
@@ -34,97 +48,418 @@ 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) -> int:
-    """一个 worker 持续领取任务;同组需求过多时使用同一计划上下文分批处理。"""
+def _run_group(biz_dt: str, batch_size: int, group_id: int) -> int:
+    """领取并执行一个固定任务;失败留到下一轮,不在当前轮内再次领取。"""
+    with get_session() as session:
+        group = DemandGradePlanRepository(session).claim_group(biz_dt, group_id)
+    if group is None:
+        logger.warning("分级任务已被其他 worker 领取或状态已变化: biz_dt=%s group_id=%s", biz_dt, group_id)
+        return 0
+
     processed_batches = 0
-    while True:
+    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,
+                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:
+                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
+
         with get_session() as session:
-            group = DemandGradePlanRepository(session).claim_next_group(biz_dt)
-        if group is None:
-            return processed_batches
+            DemandGradePlanRepository(session).finish_group(group["id"], success=True)
+    except Exception as exc:
+        logger.exception(
+            "分级计划任务失败,留待下一轮重试: biz_dt=%s group=%s",
+            biz_dt,
+            group["group_key"],
+        )
         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))
+                DemandGradePlanRepository(session).finish_group(
+                    group["id"],
+                    success=False,
+                    error_message=str(exc),
+                )
+        except Exception:
+            logger.exception(
+                "记录分级计划任务失败状态时发生错误: biz_dt=%s group=%s",
+                biz_dt,
+                group["group_key"],
+            )
+    return processed_batches
 
 
-def grade_demand_pool(
-    biz_dt: str | None = None, *, batch_size: int = _DEFAULT_BATCH_SIZE, workers: int = _DEFAULT_WORKERS
-) -> dict:
-    """先统筹落库(按需增量),再并发消费当天待执行节点组。"""
-    resolved_biz_dt = _resolve_biz_dt(biz_dt)
+def _execute_group_phase(
+    biz_dt: str,
+    *,
+    status: str,
+    batch_size: int,
+    workers: int,
+) -> dict[str, int]:
+    """执行一个状态阶段;调用方先执行 failed,再执行 pending。"""
+    with get_session() as session:
+        group_ids = DemandGradePlanRepository(session).list_group_ids_by_status(biz_dt, status)
+    if not group_ids:
+        return {"attempted_groups": 0, "batches_run": 0, "workers": 0}
+
+    worker_count = max(1, min(int(workers), len(group_ids)))
+    batches_run = 0
+    with ThreadPoolExecutor(max_workers=worker_count) as executor:
+        futures = [
+            executor.submit(_run_group, biz_dt, batch_size, group_id)
+            for group_id in group_ids
+        ]
+        for future in as_completed(futures):
+            try:
+                batches_run += future.result()
+            except Exception:
+                # _run_group 已自行兜底;这里再兜一次,禁止 worker 异常中断定时任务。
+                logger.exception("分级任务 worker 出现未捕获错误: biz_dt=%s status=%s", biz_dt, status)
+    return {
+        "attempted_groups": len(group_ids),
+        "batches_run": batches_run,
+        "workers": worker_count,
+    }
+
+
+def _execute_plan_tasks_with_retries(
+    biz_dt: str,
+    *,
+    batch_size: int,
+    workers: int,
+    max_rounds: int = _MAX_PLAN_EXECUTION_ROUNDS,
+) -> 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)
+            break
+
+        failed_result = _execute_group_phase(
+            biz_dt,
+            status="failed",
+            batch_size=batch_size,
+            workers=workers,
+        )
+        # 严格等失败阶段结束后再处理从未执行的任务。
+        pending_result = _execute_group_phase(
+            biz_dt,
+            status="pending",
+            batch_size=batch_size,
+            workers=workers,
+        )
+        total_batches += failed_result["batches_run"] + pending_result["batches_run"]
+        max_workers_used = max(
+            max_workers_used,
+            failed_result["workers"],
+            pending_result["workers"],
+        )
+
+        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"],
+            )
+
+    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,
+        "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"],
+    }
+
+
+def _grade_demand_pool_impl(
+    resolved_biz_dt: str,
+    *,
+    batch_size: int,
+    workers: int,
+) -> 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)
 
-    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)
+    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)}
 
-    with get_session() as session:
-        snapshot = DemandGradePlanRepository(session).get_execution_snapshot(resolved_biz_dt)
+    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 not assignment_after["assignment_complete"]:
+    if assignment_after.get("assignment_complete") is False:
         logger.error(
             "统筹规划未覆盖当天全部有需求的树节点,将按部分计划继续执行: biz_dt=%s uncovered=%s",
             resolved_biz_dt,
-            assignment_after["unassigned_category_ids"],
+            assignment_after.get("unassigned_category_ids"),
         )
 
-    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):
-                batches_run += future.result()
-    else:
-        logger.info("当天无待执行节点组,跳过分级 worker: biz_dt=%s", resolved_biz_dt)
+    plan_execution = _execute_plan_tasks_with_retries(
+        resolved_biz_dt,
+        batch_size=max(1, int(batch_size)),
+        workers=max(1, int(workers)),
+    )
+    # 只在统筹结束并完成计划任务状态检查后,验证树上需求等级覆盖并补偿。
+    supplement = _supplement_missing_grades(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)
-
+    final_snapshot = plan_execution["final_snapshot"]
     result = {
+        "success": bool(plan_execution["execution_complete"] and supplement["coverage_complete"]),
         "biz_dt": resolved_biz_dt,
         "total": total,
         "graded_before": graded_before,
         "graded_after": graded_after,
-        "assignment_before": assignment_state,
+        "assignment_before": assignment_before,
         "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,
-        "batches_run": batches_run,
+        "plan_execution": plan_execution,
+        "supplement": supplement,
+        "workers": plan_execution["workers"],
+        "batches_run": plan_execution["batches_run"],
+        "supplement_batches_run": supplement["batches_run"],
         "run_at": datetime.now().isoformat(),
     }
     logger.info("Tree-first grade completed: %s", result)
     return result
+
+
+def grade_demand_pool(
+    biz_dt: str | None = None,
+    *,
+    batch_size: int = _DEFAULT_BATCH_SIZE,
+    workers: int = _DEFAULT_WORKERS,
+) -> dict[str, Any]:
+    """执行完整分级闭环;任何错误只记录日志并返回,不向上中断定时任务。"""
+    resolved_biz_dt = str(biz_dt or "")
+    try:
+        resolved_biz_dt = _resolve_biz_dt(biz_dt)
+        return _grade_demand_pool_impl(
+            resolved_biz_dt,
+            batch_size=batch_size,
+            workers=workers,
+        )
+    except Exception as exc:
+        logger.exception("需求分级任务发生未捕获错误,已阻止异常中断定时任务: biz_dt=%s", resolved_biz_dt)
+        return {
+            "success": False,
+            "biz_dt": resolved_biz_dt,
+            "error": str(exc),
+            "run_at": datetime.now().isoformat(),
+        }

+ 97 - 20
supply_infra/scheduler/jobs/run_supply_pipeline.py

@@ -6,14 +6,15 @@
 3. 需求池分级评估(同上 biz_dt)
 
 各子步骤内部已做去重(INSERT IGNORE、diff 同步、跳过已分级词等);
-本文件额外用进程内锁防止同一轮次并发重入。
+本文件额外用进程内锁防止同一轮次并发重入,并隔离各步骤异常:前一步失败时记录
+error 后继续后续步骤,最终返回失败结果而不向 APScheduler 抛异常。
 """
 from __future__ import annotations
 
 import logging
 import threading
 from datetime import datetime, timedelta
-from typing import Any
+from typing import Any, Callable
 from zoneinfo import ZoneInfo
 
 from supply_infra.config import get_infra_settings
@@ -46,6 +47,47 @@ def _resolve_dates(biz_dt: str | None) -> tuple[str, str]:
     return resolved_biz_dt, tree_partition
 
 
+def _run_step(
+    step_name: str,
+    action: Callable[[], Any],
+) -> tuple[Any, bool, str | None]:
+    """执行单个流水线步骤;失败只转成结果,不允许异常越过定时任务边界。"""
+    try:
+        payload = action()
+    except Exception as exc:
+        logger.exception("Supply pipeline step failed: step=%s", step_name)
+        return {"success": False, "error": str(exc)}, False, str(exc)
+
+    if isinstance(payload, dict) and payload.get("success") is False:
+        error = str(payload.get("error") or f"{step_name} returned success=False")
+        logger.error("Supply pipeline step reported failure: step=%s error=%s", step_name, error)
+        return payload, False, error
+    return payload, True, None
+
+
+def _preflight_failure_result(biz_dt: str | None, exc: Exception) -> dict[str, Any]:
+    """日期/配置解析失败时也记录结果并正常返回,避免调度线程出现未捕获异常。"""
+    started_at = datetime.now()
+    recorder = JobExecutionRecorder(
+        job_name=SUPPLY_PIPELINE_JOB_NAME,
+        job_id=SUPPLY_PIPELINE_JOB_ID,
+        biz_dt=str(biz_dt) if biz_dt is not None else None,
+    )
+    recorder.record_started()
+    result: dict[str, Any] = {
+        "run_id": recorder.run_id,
+        "biz_dt": str(biz_dt) if biz_dt is not None else None,
+        "success": False,
+        "error": str(exc),
+        "started_at": started_at.isoformat(),
+    }
+    finished_at = datetime.now()
+    result["finished_at"] = finished_at.isoformat()
+    result["duration_seconds"] = round((finished_at - started_at).total_seconds(), 2)
+    recorder.record_finished(success=False, result=result, error_message=str(exc))
+    return result
+
+
 def run_supply_pipeline(biz_dt: str | None = None) -> dict[str, Any]:
     """
     按顺序执行全局树同步 → 需求池同步 → 需求分级。
@@ -56,7 +98,11 @@ def run_supply_pipeline(biz_dt: str | None = None) -> dict[str, Any]:
     Returns:
         各步骤统计;若上一轮仍在执行则返回 skipped。
     """
-    resolved_biz_dt, tree_partition = _resolve_dates(biz_dt)
+    try:
+        resolved_biz_dt, tree_partition = _resolve_dates(biz_dt)
+    except Exception as exc:
+        logger.exception("Supply pipeline preflight failed: biz_dt=%s", biz_dt)
+        return _preflight_failure_result(biz_dt, exc)
 
     if not _pipeline_lock.acquire(blocking=False):
         logger.warning("Supply pipeline already running, skip this round")
@@ -93,36 +139,67 @@ def run_supply_pipeline(biz_dt: str | None = None) -> dict[str, Any]:
         "global_tree_partition": tree_partition,
         "started_at": started_at.isoformat(),
     }
-    error_message: str | None = None
+    errors: list[str] = []
+    step_status: dict[str, dict[str, Any]] = {}
     success = False
 
     try:
-        result["global_tree"] = sync_global_tree_odps_to_mysql(partition_date=tree_partition)
-        result["demand_pool"] = sync_multi_demand_pool_odps_to_mysql(
-            partition_date=resolved_biz_dt,
-        )
-        result["grade"] = grade_demand_pool(biz_dt=resolved_biz_dt)
-        success = True
-        result["success"] = True
+        steps: list[tuple[str, Callable[[], Any]]] = [
+            (
+                "global_tree",
+                lambda: sync_global_tree_odps_to_mysql(partition_date=tree_partition),
+            ),
+            (
+                "demand_pool",
+                lambda: sync_multi_demand_pool_odps_to_mysql(
+                    partition_date=resolved_biz_dt,
+                ),
+            ),
+            (
+                "grade",
+                lambda: grade_demand_pool(biz_dt=resolved_biz_dt),
+            ),
+        ]
+        for step_name, action in steps:
+            payload, step_success, step_error = _run_step(step_name, action)
+            result[step_name] = payload
+            step_status[step_name] = {
+                "success": step_success,
+                "error": step_error,
+            }
+            if step_error:
+                errors.append(f"{step_name}: {step_error}")
+
+        success = all(item["success"] for item in step_status.values())
+        result["success"] = success
     except Exception as exc:
+        # 保护流水线编排本身;正常子步骤异常应已由 _run_step 消化。
         result["success"] = False
-        error_message = str(exc)
+        errors.append(f"pipeline: {exc}")
         logger.exception(
-            "Supply pipeline failed: biz_dt=%s global_tree_partition=%s",
+            "Supply pipeline orchestration failed but will not escape scheduler: "
+            "biz_dt=%s global_tree_partition=%s",
             resolved_biz_dt,
             tree_partition,
         )
-        raise
     finally:
         finished_at = datetime.now()
+        result["steps"] = step_status
+        if errors:
+            result["errors"] = errors
         result["finished_at"] = finished_at.isoformat()
         result["duration_seconds"] = round((finished_at - started_at).total_seconds(), 2)
-        recorder.record_finished(
-            success=success,
-            result=result,
-            error_message=error_message,
-        )
-        _pipeline_lock.release()
+        try:
+            recorder.record_finished(
+                success=success,
+                result=result,
+                error_message=" | ".join(errors) if errors else None,
+            )
+        except Exception:
+            # recorder 当前已自行兜底;这里防止未来实现变化造成锁无法释放。
+            logger.exception("Unexpected failure while recording supply pipeline finish")
+        finally:
+            _pipeline_lock.release()
         logger.info("Supply pipeline finished: %s", result)
 
     return result

+ 97 - 39
supply_infra/scheduler/jobs/sync_multi_demand_pool_odps_to_mysql.py

@@ -21,7 +21,7 @@ import json
 import logging
 from datetime import datetime, timedelta
 from decimal import Decimal
-from typing import Any
+from typing import Any, Callable
 
 from agents.demand_belong_category_agent.run import main as classify_demand_words
 from supply_infra.db.repositories.demand_belong_category_repo import (
@@ -184,16 +184,35 @@ def _classify_words(biz_dt: str) -> dict:
         len(batches),
     )
 
+    failed_batches: list[dict[str, Any]] = []
     for idx, batch in enumerate(batches, start=1):
         logger.info("Classifying batch %d/%d (%d words)", idx, len(batches), len(batch))
-        classify_demand_words(batch)
+        try:
+            classify_demand_words(batch)
+        except Exception as exc:
+            logger.exception(
+                "Demand belong classification batch failed; continue remaining batches: "
+                "biz_dt=%s batch=%s/%s",
+                biz_dt,
+                idx,
+                len(batches),
+            )
+            failed_batches.append(
+                {
+                    "batch": idx,
+                    "size": len(batch),
+                    "error": str(exc),
+                }
+            )
 
     return {
+        "success": not failed_batches,
         "demand_names": len(demand_names),
         "words": len(word_set),
         "existing_filtered": len(existing),
         "pending": len(pending),
         "batches": len(batches),
+        "failed_batches": failed_batches,
     }
 
 
@@ -460,21 +479,8 @@ def compute_popularity_stats(biz_dt: str) -> dict[str, Any]:
     return result
 
 
-def sync_multi_demand_pool_odps_to_mysql(partition_date: str | None = None) -> dict:
-    """
-    从 ODPS 增量同步策略需求天级数据到 MySQL,并对新词做归属分类与热度统计。
-
-    Args:
-        partition_date: 分区日期 (YYYYMMDD),默认当天
-    """
-    if partition_date is None:
-        partition_date = datetime.now().strftime("%Y%m%d")
-
-    logger.info(
-        "Starting multi demand pool ODPS → MySQL sync for partition: %s",
-        partition_date,
-    )
-
+def _sync_pool_rows(partition_date: str) -> dict[str, Any]:
+    """同步当天需求池主数据;供完整任务按独立阶段捕获异常。"""
     odps = get_odps_client()
     odps_count = odps.count_multi_demand_pool(partition_date)
     with get_session() as session:
@@ -483,7 +489,8 @@ def sync_multi_demand_pool_odps_to_mysql(partition_date: str | None = None) -> d
     logger.info("Count check: odps=%d mysql=%d", odps_count, mysql_count)
 
     if odps_count == mysql_count:
-        sync_stats: dict[str, Any] = {
+        logger.info("Same count, skip ODPS data sync")
+        return {
             "skipped_same_count": True,
             "odps_count": odps_count,
             "mysql_count": mysql_count,
@@ -491,31 +498,82 @@ def sync_multi_demand_pool_odps_to_mysql(partition_date: str | None = None) -> d
             "inserted": 0,
             "deleted": 0,
         }
-        logger.info("Same count, skip ODPS data sync")
-    else:
-        sync_stats = {
-            "skipped_same_count": False,
-            "odps_count": odps_count,
-            "mysql_count": mysql_count,
-            **_sync_diff(partition_date),
-        }
+    return {
+        "skipped_same_count": False,
+        "odps_count": odps_count,
+        "mysql_count": mysql_count,
+        **_sync_diff(partition_date),
+    }
+
+
+def _run_sync_stage(
+    stage_name: str,
+    action: Callable[[], dict[str, Any]],
+) -> tuple[dict[str, Any], str | None]:
+    """执行需求池同步子阶段,错误转为结构化结果并允许后续阶段继续。"""
+    try:
+        payload = action()
+    except Exception as exc:
+        logger.exception("Multi demand pool stage failed; continue: stage=%s", stage_name)
+        return {"success": False, "error": str(exc)}, str(exc)
+
+    if payload.get("success") is False:
+        error = str(payload.get("error") or f"{stage_name} returned success=False")
+        logger.error(
+            "Multi demand pool stage reported failure; continue: stage=%s error=%s",
+            stage_name,
+            error,
+        )
+        return payload, error
+    return payload, None
+
+
+def sync_multi_demand_pool_odps_to_mysql(partition_date: str | None = None) -> dict:
+    """
+    从 ODPS 增量同步策略需求天级数据到 MySQL,并对新词做归属分类与热度统计。
+
+    Args:
+        partition_date: 分区日期 (YYYYMMDD),默认当天
+    """
+    if partition_date is None:
+        partition_date = datetime.now().strftime("%Y%m%d")
+
+    logger.info(
+        "Starting multi demand pool ODPS → MySQL sync for partition: %s",
+        partition_date,
+    )
+
+    stage_definitions: list[tuple[str, Callable[[], dict[str, Any]]]] = [
+        ("source_sync", lambda: _sync_pool_rows(partition_date)),
+        ("classify", lambda: _classify_words(partition_date)),
+        ("belong_pool_rel", sync_demand_belong_pool_rel),
+        ("real_metrics", lambda: enrich_real_rov_vov_7d(partition_date)),
+        ("popularity", lambda: compute_popularity_stats(partition_date)),
+        ("tree_weight", lambda: compute_category_tree_weight(partition_date)),
+        ("videos", lambda: sync_multi_demand_videos(limit=None)),
+    ]
+    stage_results: dict[str, dict[str, Any]] = {}
+    errors: list[dict[str, str]] = []
+    for stage_name, action in stage_definitions:
+        payload, error = _run_sync_stage(stage_name, action)
+        stage_results[stage_name] = payload
+        if error:
+            errors.append({"stage": stage_name, "error": error})
 
-    classify_stats = _classify_words(partition_date)
-    belong_pool_rel_stats = sync_demand_belong_pool_rel()
-    real_metric_stats = enrich_real_rov_vov_7d(partition_date)
-    popularity_stats = compute_popularity_stats(partition_date)
-    tree_weight_stats = compute_category_tree_weight(partition_date)
-    video_stats = sync_multi_demand_videos(limit=None)
+    sync_stats = stage_results["source_sync"]
 
     result = {
+        "success": not errors,
         "partition_date": partition_date,
-        **sync_stats,
-        "real_metrics": real_metric_stats,
-        "classify": classify_stats,
-        "belong_pool_rel": belong_pool_rel_stats,
-        "popularity": popularity_stats,
-        "tree_weight": tree_weight_stats,
-        "videos": video_stats,
+        **({} if errors and sync_stats.get("success") is False else sync_stats),
+        "source_sync": sync_stats,
+        "real_metrics": stage_results["real_metrics"],
+        "classify": stage_results["classify"],
+        "belong_pool_rel": stage_results["belong_pool_rel"],
+        "popularity": stage_results["popularity"],
+        "tree_weight": stage_results["tree_weight"],
+        "videos": stage_results["videos"],
+        "errors": errors,
         "synced_at": datetime.now().isoformat(),
     }
     logger.info("Multi demand pool sync completed: %s", result)

+ 1 - 27
supply_infra/scheduler/jobs/update_category_tree_rank_scores.py

@@ -12,6 +12,7 @@ import logging
 from decimal import Decimal
 from typing import Any
 
+from supply_agent.ranking import rank_to_scores
 from supply_infra.db.repositories.category_tree_weight_repo import (
     CategoryTreeWeightRepository,
 )
@@ -33,33 +34,6 @@ def _dec(value: float, places: int = 8) -> Decimal:
     return Decimal(str(round(float(value), places)))
 
 
-def rank_to_scores(items: list[tuple[int, float]]) -> dict[int, float]:
-    """
-    按 avg 降序做全局排名,映射为归一化分。
-
-    items: [(category_id, avg), ...],仅含 count>0 的节点。
-    同分节点取平均名次;score = (n - avg_rank + 1) / n,范围 (0, 1]。
-    """
-    if not items:
-        return {}
-
-    sorted_items = sorted(items, key=lambda x: (-x[1], x[0]))
-    n = len(sorted_items)
-    scores: dict[int, float] = {}
-    i = 0
-    while i < n:
-        j = i
-        avg = sorted_items[i][1]
-        while j < n and sorted_items[j][1] == avg:
-            j += 1
-        avg_rank = (i + 1 + j) / 2.0
-        score = (n - avg_rank + 1) / n
-        for k in range(i, j):
-            scores[sorted_items[k][0]] = score
-        i = j
-    return scores
-
-
 def _optional_dec(value: float | None, places: int = 8) -> Decimal | None:
     if value is None:
         return None