Kaynağa Gözat

Merge branch 'master' into feature/zhangbo

zhang 1 hafta önce
ebeveyn
işleme
d2154edd87
100 değiştirilmiş dosya ile 7619 ekleme ve 271 silme
  1. 2 2
      .env.example
  2. 2 1
      README.md
  3. 10 0
      agents/demand_grade_agent/__init__.py
  4. 30 0
      agents/demand_grade_agent/agent.py
  5. 88 0
      agents/demand_grade_agent/prompt/system_prompt.md
  6. 63 0
      agents/demand_grade_agent/run.py
  7. 55 0
      agents/demand_grade_agent/tools/__init__.py
  8. 304 0
      agents/demand_grade_agent/tools/batch_save_demand_grades.py
  9. 160 0
      agents/demand_grade_agent/tools/demand_priority.py
  10. 56 0
      agents/demand_grade_agent/tools/query_category_local_heat.py
  11. 81 0
      agents/demand_grade_agent/tools/query_category_path.py
  12. 167 0
      agents/demand_grade_agent/tools/query_demand_category_and_weight.py
  13. 104 0
      agents/demand_grade_agent/tools/query_demand_popularity_by_word.py
  14. 60 0
      agents/demand_grade_agent/tools/query_latest_biz_dt.py
  15. 133 0
      agents/demand_grade_agent/tools/query_score_distribution.py
  16. 116 0
      agents/demand_grade_agent/tools/search_related_pool_demands.py
  17. 279 0
      agents/demand_grade_agent/tools/shared.py
  18. 251 0
      agents/demand_grade_agent/tools/tree_local.py
  19. 5 0
      agents/demand_grade_orchestrator_agent/__init__.py
  20. 134 0
      agents/demand_grade_orchestrator_agent/_verify_logic.py
  21. 24 0
      agents/demand_grade_orchestrator_agent/agent.py
  22. 27 0
      agents/demand_grade_orchestrator_agent/common/__init__.py
  23. 123 0
      agents/demand_grade_orchestrator_agent/common/assignment.py
  24. 102 0
      agents/demand_grade_orchestrator_agent/common/plan_persist.py
  25. 198 0
      agents/demand_grade_orchestrator_agent/common/plan_record.py
  26. 99 0
      agents/demand_grade_orchestrator_agent/common/tree_state.py
  27. 32 0
      agents/demand_grade_orchestrator_agent/prompt/system_prompt.md
  28. 158 0
      agents/demand_grade_orchestrator_agent/run.py
  29. 16 0
      agents/demand_grade_orchestrator_agent/tools/__init__.py
  30. 68 0
      agents/demand_grade_orchestrator_agent/tools/query_global_heat_tree.py
  31. 52 0
      agents/demand_grade_orchestrator_agent/tools/query_heat_node_group.py
  32. 132 0
      agents/demand_grade_orchestrator_agent/tools/save_grade_plan.py
  33. 9 0
      agents/demand_video_expand_agent/__init__.py
  34. 30 0
      agents/demand_video_expand_agent/agent.py
  35. 45 0
      agents/demand_video_expand_agent/prompt/system_prompt.md
  36. 91 0
      agents/demand_video_expand_agent/run.py
  37. 27 0
      agents/demand_video_expand_agent/tools/__init__.py
  38. 222 0
      agents/demand_video_expand_agent/tools/batch_save_demand_expansions.py
  39. 1 3
      agents/generate_demand_agent/tools/batch_save_generated_demands.py
  40. 2 2
      agents/generate_demand_agent/tools/dim_constants.py
  41. 1 1
      agents/generate_demand_agent/tools/query_demand_words_by_category.py
  42. 42 1
      api/app.py
  43. 6 0
      api/run.py
  44. 9 12
      api/services/category_tree.py
  45. 25 0
      api/services/demand_grade.py
  46. 96 0
      api/services/demand_grade_videos.py
  47. 11 4
      api/services/demand_videos.py
  48. 36 0
      jobs/backfill_demand_video_expansion_point_desc.py
  49. 1 1
      jobs/backfill_multi_demand_video_list.py
  50. 117 0
      jobs/backfill_multi_demand_video_points_table.py
  51. 48 0
      jobs/expand_demand_from_video_points.py
  52. 43 0
      jobs/grade_demand_pool.py
  53. 10 2
      jobs/init_db.py
  54. 45 0
      jobs/retry_failed_grade_plan_items.py
  55. 1 1
      jobs/run_scheduler.py
  56. 22 0
      jobs/run_supply_pipeline.py
  57. 162 0
      scripts/backfill_demand_video_expansion_point_desc.py
  58. 100 0
      scripts/retry_failed_grade_plan_items.py
  59. 85 0
      scripts/run_grade_plan_groups.py
  60. 1 1
      supply_agent/agent/core.py
  61. 1 1
      supply_agent/config.py
  62. 1 1
      supply_agent/llm/client.py
  63. 47 0
      supply_agent/ranking.py
  64. 22 0
      supply_infra/db/models/__init__.py
  65. 22 22
      supply_infra/db/models/category_tree_weight.py
  66. 70 0
      supply_infra/db/models/demand_grade.py
  67. 42 0
      supply_infra/db/models/demand_grade_category_rel.py
  68. 76 0
      supply_infra/db/models/demand_grade_plan.py
  69. 12 12
      supply_infra/db/models/demand_popularity_stats.py
  70. 100 0
      supply_infra/db/models/demand_video_expansion.py
  71. 48 0
      supply_infra/db/models/multi_demand_video_point.py
  72. 51 0
      supply_infra/db/models/scheduler_job_execution.py
  73. 22 0
      supply_infra/db/repositories/__init__.py
  74. 32 0
      supply_infra/db/repositories/category_tree_weight_repo.py
  75. 22 0
      supply_infra/db/repositories/demand_belong_category_repo.py
  76. 34 0
      supply_infra/db/repositories/demand_belong_pool_rel_repo.py
  77. 62 0
      supply_infra/db/repositories/demand_grade_category_rel_repo.py
  78. 346 0
      supply_infra/db/repositories/demand_grade_plan_repo.py
  79. 103 0
      supply_infra/db/repositories/demand_grade_repo.py
  80. 18 0
      supply_infra/db/repositories/demand_popularity_stats_repo.py
  81. 117 0
      supply_infra/db/repositories/demand_video_expansion_repo.py
  82. 131 2
      supply_infra/db/repositories/multi_demand_pool_di_repo.py
  83. 117 0
      supply_infra/db/repositories/multi_demand_video_point_repo.py
  84. 50 0
      supply_infra/db/repositories/scheduler_job_execution_repo.py
  85. 9 3
      supply_infra/db/session.py
  86. 0 4
      supply_infra/scheduler/__init__.py
  87. 113 29
      supply_infra/scheduler/app.py
  88. 4 0
      supply_infra/scheduler/constants.py
  89. 120 0
      supply_infra/scheduler/job_execution.py
  90. 35 0
      supply_infra/scheduler/jobs/backfill_multi_demand_pool_video_list.py
  91. 18 4
      supply_infra/scheduler/jobs/compute_category_tree_weight.py
  92. 335 0
      supply_infra/scheduler/jobs/expand_demand_from_video_points.py
  93. 386 0
      supply_infra/scheduler/jobs/grade_demand_pool.py
  94. 216 0
      supply_infra/scheduler/jobs/run_supply_pipeline.py
  95. 117 75
      supply_infra/scheduler/jobs/sync_multi_demand_pool_odps_to_mysql.py
  96. 29 40
      supply_infra/scheduler/jobs/sync_multi_demand_videos.py
  97. 31 39
      supply_infra/scheduler/jobs/update_category_tree_rank_scores.py
  98. 82 0
      supply_infra/scheduler/plan_group_batch.py
  99. 150 0
      supply_infra/video_points.py
  100. 9 8
      web/src/api/demand.ts

+ 2 - 2
.env.example

@@ -2,8 +2,8 @@
 OPENROUTER_API_KEY=sk-or-v1-...
 
 # Default model (any OpenRouter-supported model)
-# Examples: openai/gpt-4o, anthropic/claude-sonnet-5, google/gemini-2.5-pro-preview
-OPENROUTER_MODEL=anthropic/claude-sonnet-5
+# Examples: google/gemini-2.5-flash, google/gemini-2.5-flash-lite, anthropic/claude-sonnet-5
+OPENROUTER_MODEL=google/gemini-2.5-flash
 
 # Agent defaults
 AGENT_MAX_ITERATIONS=20

+ 2 - 1
README.md

@@ -136,6 +136,7 @@ async for event in agent.astream("Your question"):
 
 通过 OpenRouter 可使用任意支持的模型,例如:
 
+- `google/gemini-2.5-flash`(默认)
 - `anthropic/claude-sonnet-5`
 - `openai/gpt-4o`
 - `google/gemini-2.5-pro-preview`
@@ -148,7 +149,7 @@ async for event in agent.astream("Your question"):
 | 变量 | 说明 | 默认值 |
 |------|------|--------|
 | `OPENROUTER_API_KEY` | OpenRouter API 密钥 | (必填) |
-| `OPENROUTER_MODEL` | 默认模型 | `anthropic/claude-sonnet-5` |
+| `OPENROUTER_MODEL` | 默认模型 | `google/gemini-2.5-flash` |
 | `AGENT_MAX_ITERATIONS` | 最大循环次数 | `20` |
 | `AGENT_TEMPERATURE` | 生成温度 | `0.7` |
 | `SKILLS_DIR` | Skills 目录 | `skills` |

+ 10 - 0
agents/demand_grade_agent/__init__.py

@@ -0,0 +1,10 @@
+"""
+demand_grade_agent — 需求分级评估 Agent
+
+职责:对 multi_demand_pool_di 中的现有需求,结合全局树先验热度
+(category_tree_weight)与后验真实效果(real_rov_7d),以及需求词粒度效果
+(demand_popularity_stats),划分 S/A/B/C/D 等级并落库到 demand_grade 表。
+"""
+from agents.demand_grade_agent.agent import create_demand_grade_agent
+
+__all__ = ["create_demand_grade_agent"]

+ 30 - 0
agents/demand_grade_agent/agent.py

@@ -0,0 +1,30 @@
+"""
+demand_grade_agent 工厂 — 组装需求分级评估 Agent。
+"""
+from __future__ import annotations
+
+from pathlib import Path
+
+from supply_agent import Agent
+from supply_agent.config import Settings
+from agents.demand_grade_agent.tools import register_all_tools
+
+_PROMPT_PATH = Path(__file__).parent / "prompt" / "system_prompt.md"
+DEMAND_GRADE_AGENT_SYSTEM_PROMPT = _PROMPT_PATH.read_text(encoding="utf-8")
+
+
+def create_demand_grade_agent(
+    settings: Settings | None = None,
+    *,
+    model: str | None = None,
+) -> Agent:
+    """创建 demand_grade_agent 实例,注册专属工具。"""
+    agent = Agent(
+        settings=settings,
+        name="demand_grade_agent",
+        model=model,
+        system_prompt=DEMAND_GRADE_AGENT_SYSTEM_PROMPT,
+        max_iterations=40,
+    )
+    register_all_tools(agent.tools)
+    return agent

+ 88 - 0
agents/demand_grade_agent/prompt/system_prompt.md

@@ -0,0 +1,88 @@
+## 角色与任务
+你是需求优先级分级专家。你会在用户消息中收到一批需求词(来自 `multi_demand_pool_di` 策略需求池)
+和对应的 biz_dt,任务是结合其归属树节点的**全局热度**与**后验真实效果**,逐一划分 S/A/B/C/D
+五档优先级,并调用工具落库到 `demand_grade` 表,供下游选题/投放决策参考。
+
+你只做分级判断,不生成新需求词,也不修改需求池原始数据;只处理消息中给定的这些需求词,
+不需要自行查找或列举其他待处理需求。所有分类、热度、后验等证据必须通过工具从数据库查询。
+
+## 全局热度 / 后验含义
+- **全局热度**:`category_tree_weight.total_score`。反映该需求所在树节点在类目树里的历史热度排名,是"没有真实上线数据时"的兜底依据。
+  - 工具返回中 **`—` 表示无数据**,不是分数为 0。
+- **后验**:`real_rov_7d_avg` + `real_rov_7d_count`(词级工具还会返回 `real_vov_7d`)。
+  - `real_rov_7d_count > 0`:说明该节点/需求已有真实上线验证数据,**这是高置信信息,判级时应优先参考**,可以据此给出全档位(包括 S 或 D)。
+  - `real_rov_7d_count = 0`(或数据缺失):说明效果未知,只能用全局热度兜底判断。**无论全局热度多高,都不建议给到 S 级**(因为没有真实验证支撑),一般封顶在 A。
+
+## 四类证据必须分开
+- **分类节点全局/局部证据**:`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)。三类分布必须分别使用,不得共用数值阈值。
+- 有后验数据的需求:
+  - 后验效果处于同类中高位(如 real_rov_7d_avg ≥ p75)→ 可评 S 或 A
+  - 中等 → B
+  - 明显偏低(如低于 p25)→ C 或 D(即使全局热度很高,也应如实按后验降级,说明"热度高但验证效果不佳")
+- 无后验数据的需求:
+  - 全局热度 total_score 很高(如 ≥ p75)→ A(不给 S,注明"无验证数据")
+  - 中等 → B
+  - 偏低 → C
+  - 全局热度也很低、几乎无信号(total_score 为 —)→ D
+- 一个需求可能挂在多个树节点上:以最相关/得分最高的节点为主要依据,reason 中说明取用了哪个节点。
+- **局部热度校正**:结合 `query_category_local_heat` 返回的完整局部环境(自身节点、父节点、全部兄弟节点的全局热度与后验数据)做校正。
+  节点自身热度高、父节点热且在兄弟中靠前时,可上调同等全局分下的等级或 `score`;
+  节点自身偏冷、父节点与兄弟整体也偏冷时,应下调等级或 `score`。节点自身与局部环境冲突时,
+  不直接套规则:结合词级后验判断它是局部突发还是弱信号,并在 reason 中写明冲突。
+  **禁止只引用部分兄弟节点**;必须以工具返回的全部兄弟节点数据作为局部参照。
+- **整树位置校正**:每个节点还会返回 `global_tree_position`。兄弟内排名只能回答局部冷热,必须同时参考整树名次/排名分和 `query_score_distribution`;不能因为一个冷分支里排名第一就直接判为高热。
+- 父节点/兄弟节点只用于校正,不得覆盖明确的低后验:有充分真实后验且表现差时,仍应降级。
+
+## 同义/相似需求合并
+同一语义的需求可能因措辞不同而在需求池里表现为多条独立记录(例如「减脂期加餐」与
+「减脂加餐」)。判级前应调用 `search_related_pool_demands(biz_dt, keywords=[...])`
+**批量**搜索本批各需求词,把找到的相关记录一并纳入参考(尤其是它们各自的 weight / real_rov_7d),
+不要只看单条记录就下结论;返回的 `[id=...]` 就是 `multi_demand_pool_di.id`,落库时必须原样
+收集进 `related_pool_ids`(**必填字段**,用于把分级结果关联回原始需求行)。
+
+`video_list`(关联视频列表)与 `strategies`(来源策略)会由 `batch_save_demand_grades` 根据
+`related_pool_ids` 自动从对应的原始需求行推导写入,你不需要手工整理这两个字段——但必须保证
+`related_pool_ids` 完整、准确,否则这两个字段会推导缺失或不完整。`category_ids` 除了写入
+`demand_grade` 的展示快照字段,也会同步写入 `demand_grade_category_rel` 映射表,供前端按分类
+高效查询已分级需求,因此尽量把该需求真实归属的所有节点都列全。
+
+## 可用工具
+- `query_latest_biz_dt()`:若用户消息未给出明确 biz_dt 时调用,返回需求池/权重表/热度统计表各自最新业务日。
+- `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)`:分别查询分类树全局热度、需求自身来源归一分与后验分布,制定跨批次一致标准。
+- `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)`,分别确定分类树、需求自身与后验的参考区间(该分布来自数据表,不依赖对话历史,每轮调用结果一致,可保证跨批次标准统一)。
+3. **优先批量调用取数工具以减少往返**:
+   - 一次 `search_related_pool_demands(biz_dt, keywords=[...])` 覆盖本批所有需求词(或按 5~10 个一组分批);
+   - 一次 `query_demand_category_and_weight(demand_names=[...], biz_dt=...)` 批量取归属与权重;
+   - 根据归属到的 `category_id`,调用 `query_category_local_heat(biz_dt, category_ids=[...])` 获取父/兄弟局部热度;
+   - 必要时一次 `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`
+   - 归属分类节点:`query_demand_category_and_weight`
+   - 全局与局部环境(整树名次 + 父节点 + 全部兄弟节点完整权重):`query_category_local_heat`
+   reason 必须同时写清需求自身来源归一分/有效来源数、分类节点整树位置、局部判断和词级后验;`related_pool_ids` 取自 `search_related_pool_demands` 返回的 `[id=...]`。
+5. 全部处理完后,调用一次(或分 2~3 次)`batch_save_demand_grades` 落库,覆盖这批给定的所有需求词。
+   `score` 不需要自行计算或传入,保存工具会用当日全量需求池确定性重算 `source_rank_score` 并落库;禁止另造一套模型分。
+6. 简要汇报本批次的分级结果后结束本轮任务。
+
+## 原则
+- 禁止幻觉:只能引用工具真实返回的数值和类目路径,不能编造 category_id、分数或树节点名称;工具输出中的 `—` 表示无数据,不得当作 0 写入落库字段。
+- reason 必须具体:写清引用的全局热度/后验数值、是否合并了同义词、依据哪个树节点。
+- 找不到归属树节点或权重数据的需求:如实说明"无法评级/数据缺失",不要强行给出等级去凑数。
+- 有后验数据始终优先于纯全局热度判断;无后验数据时保持谨慎,不给最高档。

+ 63 - 0
agents/demand_grade_agent/run.py

@@ -0,0 +1,63 @@
+#!/usr/bin/env python3
+"""Run demand_grade_agent on one batch of demand names.
+
+批次选择(哪些词、多少个)由调用方(通常是调度任务)负责,本入口只处理
+传入的这一批,不做分页/自行遍历需求池。
+"""
+from __future__ import annotations
+
+from typing import Any
+
+from agents.demand_grade_agent import create_demand_grade_agent
+
+
+def build_grade_user_input(demands: list[dict[str, Any]], biz_dt: str) -> str:
+    """构建传给分级 Agent 的用户消息。"""
+    demand_lines = []
+    for item in demands:
+        pool_id = item.get("pool_id", item.get("id"))
+        if pool_id is None:
+            raise ValueError(f"demand 缺少 pool_id: {item!r}")
+        demand_lines.append(f"[{int(pool_id)}] {str(item['demand_name']).strip()}")
+
+    lines_text = "\n".join(demand_lines)
+    return f"""请对以下 {len(demand_lines)} 个需求词逐一评级(S/A/B/C/D)。
+只处理这一批,不要尝试查找或列举更多需求词。判级完成后调用 batch_save_demand_grades 落库。
+落库时 related_pool_ids 使用列表中对应行的 pool_id;分类、热度、后验等证据请通过工具从数据库查询。
+
+biz_dt={biz_dt}
+需求词列表 [pool_id] demand_name
+{lines_text}
+"""
+
+
+def main(
+    demands: list[dict[str, Any]],
+    biz_dt: str | None = None,
+) -> None:
+    agent = create_demand_grade_agent()
+    print(f"demand_grade_agent ready | model={agent.model}")
+    print(f"tools: {agent.tools.list_tools()}")
+    print()
+
+    if not demands:
+        raise ValueError("demands 不能为空")
+
+    batch_biz_dt = (biz_dt or "").strip()
+    if not batch_biz_dt:
+        raise ValueError("biz_dt 不能为空")
+
+    user_input = build_grade_user_input(demands, batch_biz_dt)
+    result = agent.run(user_input)
+    print(result.content)
+    print(f"\n[iterations={result.iterations}, tool_calls={result.tool_calls_made}]")
+
+
+if __name__ == "__main__":
+    main(
+        [
+            {"pool_id": 101, "demand_name": "减脂期加餐"},
+            {"pool_id": 102, "demand_name": "减脂期"},
+        ],
+        biz_dt="20260721",
+    )

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

@@ -0,0 +1,55 @@
+"""
+demand_grade_agent 工具包
+
+批次选词(哪些需求词、多少个)由外部调度任务负责,agent 只处理调用方在
+用户消息里显式给出的需求词,不提供“列出/分页遍历需求池”类工具。
+"""
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import Any
+
+from agents.demand_grade_agent.tools.batch_save_demand_grades import batch_save_demand_grades
+from agents.demand_grade_agent.tools.query_category_local_heat import query_category_local_heat
+from agents.demand_grade_agent.tools.query_category_path import query_category_path
+from agents.demand_grade_agent.tools.query_demand_category_and_weight import (
+    query_demand_category_and_weight,
+)
+from agents.demand_grade_agent.tools.query_demand_popularity_by_word import (
+    query_demand_popularity_by_word,
+)
+from agents.demand_grade_agent.tools.query_latest_biz_dt import query_latest_biz_dt
+from agents.demand_grade_agent.tools.query_score_distribution import query_score_distribution
+from agents.demand_grade_agent.tools.search_related_pool_demands import (
+    search_related_pool_demands,
+)
+from supply_agent.tools.registry import ToolRegistry
+
+ALL_TOOLS: list[Callable[..., Any]] = [
+    query_latest_biz_dt,
+    search_related_pool_demands,
+    query_demand_category_and_weight,
+    query_category_path,
+    query_category_local_heat,
+    query_demand_popularity_by_word,
+    query_score_distribution,
+    batch_save_demand_grades,
+]
+
+__all__ = [
+    "ALL_TOOLS",
+    "batch_save_demand_grades",
+    "query_category_local_heat",
+    "query_category_path",
+    "query_demand_category_and_weight",
+    "query_demand_popularity_by_word",
+    "query_latest_biz_dt",
+    "query_score_distribution",
+    "register_all_tools",
+    "search_related_pool_demands",
+]
+
+
+def register_all_tools(registry: ToolRegistry) -> ToolRegistry:
+    """将 demand_grade_agent 包内的所有工具注册到 ToolRegistry。"""
+    return registry.from_decorated(*ALL_TOOLS)

+ 304 - 0
agents/demand_grade_agent/tools/batch_save_demand_grades.py

@@ -0,0 +1,304 @@
+"""
+批量保存需求分级结果到 demand_grade 表。
+"""
+from __future__ import annotations
+
+import json
+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,
+    dump_int_list,
+    merge_video_ids,
+    normalize_biz_dt,
+)
+from supply_agent.tools import tool
+from supply_infra.db.repositories.demand_grade_category_rel_repo import (
+    DemandGradeCategoryRelRepository,
+)
+from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
+from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
+from supply_infra.db.session import get_session
+
+logger = logging.getLogger(__name__)
+_MAX_ERROR_DETAILS = 10
+
+
+def _format_errors(errors: list[str]) -> str:
+    if not errors:
+        return ""
+    if len(errors) <= _MAX_ERROR_DETAILS:
+        return ";".join(errors)
+    hidden = len(errors) - _MAX_ERROR_DETAILS
+    return ";".join(errors[:_MAX_ERROR_DETAILS]) + f";...另有 {hidden} 条类似错误"
+
+
+def _coerce_items(raw: Any) -> tuple[list[Any], str | None]:
+    """将 Agent 传入的 items 规范为列表,兼容误传 JSON 字符串。"""
+    if raw is None:
+        return [], "items 不能为空"
+    if isinstance(raw, str):
+        text = raw.strip()
+        if not text:
+            return [], "items 不能为空"
+        try:
+            raw = json.loads(text)
+        except json.JSONDecodeError:
+            return [], "items 必须是对象数组,不能把未解析的 JSON 字符串直接传入"
+    if isinstance(raw, dict):
+        return [raw], None
+    if not isinstance(raw, list):
+        return [], f"items 必须是数组,当前类型: {type(raw).__name__}"
+    return raw, None
+
+
+def _optional_decimal(value: Any, field: str, idx: int, errors: list[str]) -> Decimal | None:
+    if value is None or value == "":
+        return None
+    try:
+        return Decimal(str(value))
+    except Exception:
+        errors.append(f"第 {idx} 项 {field} 无效: {value!r}")
+        return None
+
+
+def _optional_int_list(value: Any, field: str, idx: int, errors: list[str]) -> list[int]:
+    if value is None:
+        return []
+    if not isinstance(value, list):
+        errors.append(f"第 {idx} 项 {field} 必须是数组: {value!r}")
+        return []
+    out: list[int] = []
+    for v in value:
+        try:
+            out.append(int(v))
+        except (TypeError, ValueError):
+            errors.append(f"第 {idx} 项 {field} 含无效元素: {v!r}")
+            return []
+    return out
+
+
+def _normalize_items(
+    items: list[dict[str, Any]],
+    *,
+    default_biz_dt: str | None,
+) -> tuple[list[dict[str, Any]], list[list[int]], list[list[int]], list[str]]:
+    """校验并规范化待落库行,返回 (rows, related_pool_id_lists, category_id_lists, errors)。"""
+    rows: list[dict[str, Any]] = []
+    related_pool_id_lists: list[list[int]] = []
+    category_id_lists: list[list[int]] = []
+    errors: list[str] = []
+    seen_keys: set[tuple[str, str]] = set()
+
+    for idx, item in enumerate(items):
+        if not isinstance(item, dict):
+            errors.append(f"第 {idx} 项不是对象")
+            continue
+
+        demand_name = str(item.get("demand_name") or "").strip()
+        if not demand_name:
+            errors.append(f"第 {idx} 项缺少 demand_name")
+            continue
+
+        grade = str(item.get("grade") or "").strip().upper()
+        if grade not in VALID_GRADES:
+            errors.append(f"第 {idx} 项 grade 无效(只能是 {'/'.join(VALID_GRADES)}): {item.get('grade')!r}")
+            continue
+
+        reason = str(item.get("reason") or "").strip()
+        if not reason:
+            errors.append(f"第 {idx} 项缺少 reason")
+            continue
+
+        item_biz_dt, err = normalize_biz_dt(item.get("biz_dt"))
+        if err:
+            errors.append(f"第 {idx} 项 {err}")
+            continue
+        resolved_biz_dt = item_biz_dt or default_biz_dt
+        if not resolved_biz_dt:
+            errors.append(f"第 {idx} 项缺少 biz_dt(且未传入默认 biz_dt)")
+            continue
+
+        dedupe_key = (resolved_biz_dt, demand_name)
+        if dedupe_key in seen_keys:
+            errors.append(f"第 {idx} 项在本次请求中重复: biz_dt={resolved_biz_dt}, demand_name={demand_name}")
+            continue
+
+        related_pool_ids = _optional_int_list(item.get("related_pool_ids"), "related_pool_ids", idx, errors)
+        if not related_pool_ids and item.get("pool_id") is not None:
+            try:
+                related_pool_ids = [int(item["pool_id"])]
+            except (TypeError, ValueError):
+                errors.append(f"第 {idx} 项 pool_id 无效: {item.get('pool_id')!r}")
+                continue
+        if not related_pool_ids:
+            errors.append(
+                f"第 {idx} 项缺少 related_pool_ids(必填,需先用 search_related_pool_demands 找到对应的 "
+                f"multi_demand_pool_di.id)"
+            )
+            continue
+
+        seen_keys.add(dedupe_key)
+
+        prior_raw = item.get("prior_total_score")
+        if prior_raw is None or prior_raw == "" or prior_raw == "—":
+            prior_total_score = None
+        else:
+            parsed_prior = _optional_decimal(prior_raw, "prior_total_score", idx, errors)
+            prior_total_score = None if parsed_prior is None or parsed_prior == 0 else parsed_prior
+        posterior_rov_avg = _optional_decimal(item.get("posterior_rov_avg"), "posterior_rov_avg", idx, errors)
+
+        posterior_rov_count = 0
+        if item.get("posterior_rov_count") is not None:
+            try:
+                posterior_rov_count = int(item.get("posterior_rov_count"))
+            except (TypeError, ValueError):
+                errors.append(f"第 {idx} 项 posterior_rov_count 无效: {item.get('posterior_rov_count')!r}")
+                continue
+
+        category_ids = _optional_int_list(item.get("category_ids"), "category_ids", idx, errors)
+
+        has_posterior = 1 if posterior_rov_count > 0 else 0
+
+        rows.append(
+            {
+                "biz_dt": resolved_biz_dt,
+                "demand_name": demand_name,
+                "category_ids": dump_int_list(category_ids),
+                "grade": grade,
+                # 保存阶段会基于当日全量需求池确定性重算,禁止由模型自由填写。
+                "score": None,
+                "prior_total_score": prior_total_score,
+                "posterior_rov_avg": posterior_rov_avg,
+                "posterior_rov_count": posterior_rov_count,
+                "has_posterior": has_posterior,
+                "related_pool_ids": dump_int_list(related_pool_ids),
+                "reason": reason,
+            }
+        )
+        related_pool_id_lists.append(related_pool_ids)
+        category_id_lists.append(category_ids)
+
+    return rows, related_pool_id_lists, category_id_lists, errors
+
+
+@tool
+def batch_save_demand_grades(items: list[dict[str, Any]], biz_dt: Optional[str] = None) -> str:
+    """
+    批量保存需求分级结果到 demand_grade 表(按 biz_dt+demand_name upsert,可重复调用覆盖修正)。
+
+    video_list(关联视频列表)与 strategies(来源策略列表)会自动从 related_pool_ids 对应的
+    multi_demand_pool_di 原始行推导写入,无需手工传入。category_ids 除了写入展示快照字段,
+    也会同步写入 demand_grade_category_rel 映射表,供前端按分类高效查询。
+
+    Args:
+        items: 待保存列表,每项字段:
+            - demand_name (必填): 需求名称
+            - grade (必填): S/A/B/C/D 之一
+            - reason (必填): 判断依据,需引用具体的先验/后验数值
+            - related_pool_ids (必填): 该需求对应的 multi_demand_pool_di.id 列表;若输入里给了
+              pool_id,也可直接写 "pool_id": 123 代替 related_pool_ids
+            - score: 无需传入;保存时按当日全量需求池自动计算需求自身来源归一分(0-100),
+              即各 strategy 内独立排名归一化后,对该需求已有来源取均值
+            - category_ids (可选): 归属的树节点 id 列表,会写入 demand_grade_category_rel 映射表
+            - prior_total_score (可选): 落库时的先验 total_score 快照
+            - posterior_rov_avg / posterior_rov_count (可选): 落库时的后验 real_rov_7d 快照;
+              count>0 时自动标记为「有后验数据」
+            - biz_dt (可选): 覆盖本项使用的业务日,不传则用调用时的 biz_dt 参数
+        biz_dt: 本次调用的默认业务日期 YYYYMMDD;items 内每项也可单独指定 biz_dt 覆盖。
+
+    Returns:
+        保存结果摘要,包含成功条数与校验失败说明。
+    """
+    if not items:
+        return "items 不能为空"
+
+    default_biz_dt, err = normalize_biz_dt(biz_dt)
+    if err:
+        return err
+
+    coerced_items, coerce_err = _coerce_items(items)
+    if coerce_err:
+        return coerce_err
+
+    rows, related_pool_id_lists, category_id_lists, errors = _normalize_items(
+        coerced_items, default_biz_dt=default_biz_dt
+    )
+    if not rows:
+        detail = _format_errors(errors) if errors else "无有效数据"
+        return f"没有可保存的数据: {detail}"
+
+    try:
+        with get_session() as session:
+            pool_repo = MultiDemandPoolDiRepository(session)
+            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] = []
+            for i, (row, pool_ids) in enumerate(zip(rows, related_pool_id_lists)):
+                matched = [pool_by_id[pid] for pid in pool_ids if pid in pool_by_id]
+                missing = [pid for pid in pool_ids if pid not in pool_by_id]
+                if not matched:
+                    errors.append(
+                        f"demand_name={row['demand_name']!r} 的 related_pool_ids={pool_ids} "
+                        f"均未在 multi_demand_pool_di 中找到,跳过该项"
+                    )
+                    continue
+                if missing:
+                    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)
+                saved_indices.append(i)
+
+            if not final_rows:
+                detail = _format_errors(errors) if errors else "无有效数据"
+                return f"没有可保存的数据: {detail}"
+
+            grade_repo = DemandGradeRepository(session)
+            affected = grade_repo.bulk_upsert(final_rows)
+
+            names_by_biz_dt: dict[str, list[str]] = {}
+            for i in saved_indices:
+                names_by_biz_dt.setdefault(rows[i]["biz_dt"], []).append(rows[i]["demand_name"])
+
+            id_by_biz_dt_name: dict[tuple[str, str], int] = {}
+            for bd, names in names_by_biz_dt.items():
+                for name, demand_grade_id in grade_repo.get_ids_by_names(bd, names).items():
+                    id_by_biz_dt_name[(bd, name)] = demand_grade_id
+
+            rel_repo = DemandGradeCategoryRelRepository(session)
+            for i in saved_indices:
+                key = (rows[i]["biz_dt"], rows[i]["demand_name"])
+                demand_grade_id = id_by_biz_dt_name.get(key)
+                if demand_grade_id is not None:
+                    rel_repo.replace_for_demand_grade(demand_grade_id, category_id_lists[i])
+
+        parts = [f"提交 {len(rows)} 条,成功写入/更新 {affected} 条({len(final_rows)} 条通过校验)"]
+        if errors:
+            parts.append(f"校验失败/警告 {len(errors)} 条: {_format_errors(errors)}")
+
+        message = "。".join(parts)
+        logger.info("batch_save_demand_grades completed: %s", message)
+        return message
+
+    except Exception as e:
+        logger.error("batch_save_demand_grades failed: %s", e, exc_info=True)
+        return f"批量保存需求分级失败: {e}"

+ 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

+ 56 - 0
agents/demand_grade_agent/tools/query_category_local_heat.py

@@ -0,0 +1,56 @@
+"""查询树节点自身、父节点、兄弟节点的局部热度与排名。"""
+from __future__ import annotations
+
+import logging
+from typing import Any
+
+from agents.demand_grade_agent.tools.shared import normalize_biz_dt
+from agents.demand_grade_agent.tools.tree_local import render_local_heat_report
+from supply_agent.tools import tool
+
+logger = logging.getLogger(__name__)
+
+
+def _normalize_ids(category_ids: list[Any]) -> tuple[list[int], str | None]:
+    if not category_ids:
+        return [], "category_ids 不能为空"
+    out: list[int] = []
+    seen: set[int] = set()
+    for raw in category_ids:
+        try:
+            cid = int(raw)
+        except (TypeError, ValueError):
+            return [], f"category_ids 含无效 id: {raw!r}"
+        if cid in seen:
+            continue
+        seen.add(cid)
+        out.append(cid)
+    if not out:
+        return [], "category_ids 不能为空"
+    return out, None
+
+
+@tool
+def query_category_local_heat(biz_dt: str, category_ids: list[int]) -> str:
+    """查询节点自身、父节点、全部兄弟节点的全局热度与后验数据及兄弟内排名。
+
+    只返回 total_score 与 real_rov_7d/real_vov_7d,不包含四维先验明细。
+    """
+    normalized_dt, err = normalize_biz_dt(biz_dt)
+    if err:
+        return err
+    ids, err = _normalize_ids(category_ids)
+    if err:
+        return err
+
+    try:
+        report = render_local_heat_report(normalized_dt, ids)
+        logger.info("query_category_local_heat completed: biz_dt=%s ids=%s", normalized_dt, ids)
+        return report
+    except Exception as exc:
+        logger.error("query_category_local_heat failed: %s", exc, exc_info=True)
+        return f"查询局部树节点热度失败: {exc}"
+
+
+if __name__ == "__main__":
+    print(query_category_local_heat("20260714", [313, 686]))

+ 81 - 0
agents/demand_grade_agent/tools/query_category_path.py

@@ -0,0 +1,81 @@
+"""
+按分类 id 查询从根到节点的类目路径,用于写 reason。
+"""
+from __future__ import annotations
+
+import logging
+from typing import Any
+
+from agents.demand_grade_agent.tools.shared import build_category_path
+from supply_agent.tools import tool
+from supply_infra.db.repositories.global_tree_category_repo import (
+    GlobalTreeCategoryRepository,
+)
+from supply_infra.db.session import get_session
+
+logger = logging.getLogger(__name__)
+
+
+def _normalize_ids(category_ids: list[Any]) -> tuple[list[int], str | None]:
+    if not category_ids:
+        return [], "category_ids 不能为空"
+    out: list[int] = []
+    seen: set[int] = set()
+    for raw in category_ids:
+        try:
+            cid = int(raw)
+        except (TypeError, ValueError):
+            return [], f"category_ids 含无效 id: {raw!r}"
+        if cid in seen:
+            continue
+        seen.add(cid)
+        out.append(cid)
+    if not out:
+        return [], "category_ids 不能为空"
+    return out, None
+
+
+@tool
+def query_category_path(category_ids: list[int]) -> str:
+    """
+    查询分类节点从根到自身的名称路径,用于写 reason。
+
+    Args:
+        category_ids: 分类 id 列表(可多个)。
+
+    Returns:
+        每行一条路径,例如:
+        [88] 美食 > 减脂饮食 > 加餐
+        [128] 美妆护肤 > 护肤 > 精华
+    """
+    ids, err = _normalize_ids(category_ids)
+    if err:
+        return err
+
+    try:
+        with get_session() as session:
+            categories = GlobalTreeCategoryRepository(session).list_active_categories()
+            by_id = {int(c.id): c for c in categories}
+
+            lines: list[str] = []
+            for cid in ids:
+                path = build_category_path(cid, by_id)
+                if path is None:
+                    lines.append(f"[{cid}] (未找到该分类)")
+                else:
+                    lines.append(f"[{cid}] {path}")
+
+        logger.info("query_category_path completed: ids=%s", ids)
+        return "\n".join(lines)
+
+    except Exception as e:
+        logger.error("query_category_path failed: %s", e, exc_info=True)
+        return f"查询分类路径失败: {e}"
+
+
+def main() -> None:
+    print(query_category_path(category_ids=[1, 2]))
+
+
+if __name__ == "__main__":
+    main()

+ 167 - 0
agents/demand_grade_agent/tools/query_demand_category_and_weight.py

@@ -0,0 +1,167 @@
+"""
+核心取数工具:需求名 → 归属树节点 → 该节点的先验热度与后验真实效果。
+"""
+from __future__ import annotations
+
+import logging
+from typing import Optional
+
+from sqlalchemy.orm import Session
+
+from agents.demand_grade_agent.tools.shared import (
+    build_category_path,
+    format_category_weight_lines,
+    normalize_biz_dt,
+    normalize_str_list,
+)
+from supply_agent.tools import tool
+from supply_infra.db.models.category_tree_weight import CategoryTreeWeight
+from supply_infra.db.models.global_tree_category import GlobalTreeCategory
+from supply_infra.db.repositories.category_tree_weight_repo import CategoryTreeWeightRepository
+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.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
+
+logger = logging.getLogger(__name__)
+
+
+def _format_weight_row(weight: CategoryTreeWeight, path: str | None) -> str:
+    return "\n".join(format_category_weight_lines(
+        weight,
+        category_id=int(weight.category_id),
+        name=None,
+        path=path,
+        biz_dt=weight.biz_dt,
+        hung_word_count=int(weight.hung_word_count or 0),
+    ))
+
+
+def _query_one_demand_category_and_weight(
+    session: Session,
+    demand_name: str,
+    normalized_dt: str | None,
+    by_id: dict[int, GlobalTreeCategory],
+) -> str:
+    belong_ids: set[int] = set()
+
+    pool_repo = MultiDemandPoolDiRepository(session)
+    if normalized_dt:
+        pool_rows = pool_repo.search_rows_by_name_fragment(normalized_dt, demand_name)
+        pool_ids = [row["id"] for row in pool_rows if row["demand_name"] == demand_name]
+    else:
+        pool_ids = []
+
+    if pool_ids:
+        rel_map = DemandBelongPoolRelRepository(session).get_belong_ids_by_pool_ids(pool_ids)
+        for ids in rel_map.values():
+            belong_ids.update(ids)
+
+    belong_repo = DemandBelongCategoryRepository(session)
+    if not belong_ids:
+        fuzzy = belong_repo.search_by_name_like(demand_name)
+        belong_ids.update(int(row.id) for row in fuzzy)
+
+    if not belong_ids:
+        return f"「{demand_name}」未找到归属的树节点(demand_belong_pool_rel 与 demand_belong_category 均无匹配)"
+
+    belong_rows = belong_repo.get_by_ids(list(belong_ids))
+    category_ids = sorted({row.category_id for row in belong_rows if row.category_id})
+    if not category_ids:
+        return f"「{demand_name}」匹配到需求归属记录,但均未挂载 category_id"
+
+    weight_repo = CategoryTreeWeightRepository(session)
+    weights = weight_repo.get_by_category_ids(category_ids, normalized_dt)
+    if not weights:
+        dt_note = f"biz_dt={normalized_dt}" if normalized_dt else "任意业务日"
+        return (
+            f"「{demand_name}」归属树节点 {category_ids},"
+            f"但 category_tree_weight 在 {dt_note} 无数据"
+        )
+
+    lines = [f"「{demand_name}」归属 {len(weights)} 个树节点:"]
+    for weight in sorted(
+        weights,
+        key=lambda w: (w.total_score is None, -(float(w.total_score) if w.total_score is not None else 0)),
+    ):
+        path = build_category_path(int(weight.category_id), by_id)
+        lines.append(_format_weight_row(weight, path))
+    return "\n".join(lines)
+
+
+@tool
+def query_demand_category_and_weight(
+    demand_names: list[str],
+    biz_dt: Optional[str] = None,
+) -> str:
+    """
+    需求名 → 归属树节点 → 全局热度 total_score + 后验真实效果(real_rov_7d)。
+
+    支持批量传入多个需求名,一次调用返回各词的归属与权重;每段结果前会标注原始 demand_name。
+
+    查找顺序(对每个 demand_name):
+    1. 先在 multi_demand_pool_di 中按需求名精确匹配找到池表行 id;
+    2. 经 demand_belong_pool_rel 反查这些池表行归属的 demand_belong_category;
+    3. 若第 2 步查不到关系,退化为按需求名模糊匹配 demand_belong_category.name;
+    4. 用归属到的 category_id 查询 category_tree_weight 的全局热度与后验数据。
+
+    一个需求可能挂在多个树节点上,会全部列出。
+
+    Args:
+        demand_names: 需求名称列表(通常来自 multi_demand_pool_di.demand_name),可一次传多个。
+        biz_dt: 业务日期 YYYYMMDD,可选。传入时会先在该日期的需求池里精确匹配需求名,
+                再经 demand_belong_pool_rel 反查归属;不传时直接对 demand_belong_category.name
+                做模糊匹配(跳过第 1/2 步),权重也取每个树节点各自的最新业务日数据。
+
+    Returns:
+        每个 demand_name 一段,段首标注 `--- demand_name: xxx ---`,例如:
+        --- demand_name: 减脂期加餐 ---
+        「减脂期加餐」归属 1 个树节点:
+        [category_id=88] 美食 > 减脂饮食 > 加餐
+          biz_dt=20260716  全局热度total_score=3.42
+          ...
+    """
+    normalized_dt, err = normalize_biz_dt(biz_dt)
+    if err:
+        return err
+
+    names, err = normalize_str_list(demand_names, "demand_names")
+    if err:
+        return err
+
+    try:
+        with get_session() as session:
+            categories = GlobalTreeCategoryRepository(session).list_active_categories()
+            by_id = {int(c.id): c for c in categories}
+
+            sections: list[str] = []
+            for demand_name in names:
+                result = _query_one_demand_category_and_weight(
+                    session, demand_name, normalized_dt, by_id
+                )
+                sections.append(f"--- demand_name: {demand_name} ---\n{result}")
+
+        message = "\n\n".join(sections)
+        logger.info(
+            "query_demand_category_and_weight completed: count=%d biz_dt=%s",
+            len(names),
+            normalized_dt,
+        )
+        return message
+
+    except Exception as e:
+        logger.error("query_demand_category_and_weight failed: %s", e, exc_info=True)
+        return f"查询需求归属与权重失败: {e}"
+
+
+def main() -> None:
+    print(query_demand_category_and_weight(demand_names=["减脂期加餐", "减脂加餐"]))
+
+
+if __name__ == "__main__":
+    main()

+ 104 - 0
agents/demand_grade_agent/tools/query_demand_popularity_by_word.py

@@ -0,0 +1,104 @@
+"""
+按需求词粒度直接查询 demand_popularity_stats,作为树节点粒度数据的交叉验证。
+"""
+from __future__ import annotations
+
+import logging
+from typing import Optional
+
+from sqlalchemy.orm import Session
+
+from agents.demand_grade_agent.tools.shared import (
+    format_dim_with_count,
+    normalize_biz_dt,
+    normalize_str_list,
+)
+from supply_agent.tools import tool
+from supply_infra.db.repositories.demand_popularity_stats_repo import (
+    DemandPopularityStatsRepository,
+)
+from supply_infra.db.session import get_session
+
+logger = logging.getLogger(__name__)
+
+
+def _query_one_demand_popularity_by_word(
+    session: Session,
+    keyword: str,
+    normalized_dt: str | None,
+) -> str:
+    rows = DemandPopularityStatsRepository(session).search_by_word_name(keyword, normalized_dt)
+
+    if not rows:
+        dt_note = f"biz_dt={normalized_dt}" if normalized_dt else "任意业务日"
+        return f"「{keyword}」在 demand_popularity_stats({dt_note})无匹配数据"
+
+    lines = []
+    for row in rows:
+        posterior_note = "有后验数据" if row.real_rov_7d_count > 0 else "无后验数据(效果未知)"
+        lines.append(
+            f"[biz_dt={row.biz_dt}] {row.demand_word_name}: "
+            f"后验real_rov_7d={format_dim_with_count(row.real_rov_7d_avg, row.real_rov_7d_count)} "
+            f"后验real_vov_7d={format_dim_with_count(row.real_vov_7d_avg, row.real_vov_7d_count)} "
+            f"({posterior_note})"
+        )
+    return "\n".join(lines)
+
+
+@tool
+def query_demand_popularity_by_word(
+    demand_word_names: list[str],
+    biz_dt: Optional[str] = None,
+) -> str:
+    """
+    按需求词名(精确+模糊)查询 demand_popularity_stats 的词粒度效果数据。
+
+    支持批量传入多个需求词,一次调用返回各词的统计;每段结果前会标注原始 demand_word_name。
+
+    demand_popularity_stats 比 category_tree_weight 粒度更细(按具体需求词而非整个树节点),
+    可用于交叉验证 query_demand_category_and_weight 给出的树节点级结论,
+    也可用于发现措辞相近但独立统计的近似词数据。
+
+    Args:
+        demand_word_names: 需求词名称或关键片段列表,可一次传多个。
+        biz_dt: 业务日期 YYYYMMDD,可选;不传则不限日期,按业务日降序列出(可能有多天历史数据)。
+
+    Returns:
+        每个 demand_word_name 一段,段首标注 `--- demand_word_name: xxx ---`,例如:
+        --- demand_word_name: 减脂期加餐 ---
+        [biz_dt=20260716] 减脂期加餐: 后验real_rov_7d=0.12(n=5) 后验real_vov_7d=0.08(n=5) (有后验数据)
+    """
+    normalized_dt, err = normalize_biz_dt(biz_dt)
+    if err:
+        return err
+
+    names, err = normalize_str_list(demand_word_names, "demand_word_names")
+    if err:
+        return err
+
+    try:
+        with get_session() as session:
+            sections: list[str] = []
+            for keyword in names:
+                result = _query_one_demand_popularity_by_word(session, keyword, normalized_dt)
+                sections.append(f"--- demand_word_name: {keyword} ---\n{result}")
+
+        message = "\n\n".join(sections)
+        logger.info(
+            "query_demand_popularity_by_word completed: count=%d biz_dt=%s",
+            len(names),
+            normalized_dt,
+        )
+        return message
+
+    except Exception as e:
+        logger.error("query_demand_popularity_by_word failed: %s", e, exc_info=True)
+        return f"查询需求词热度统计失败: {e}"
+
+
+def main() -> None:
+    print(query_demand_popularity_by_word(demand_word_names=["加餐", "减脂期加餐"]))
+
+
+if __name__ == "__main__":
+    main()

+ 60 - 0
agents/demand_grade_agent/tools/query_latest_biz_dt.py

@@ -0,0 +1,60 @@
+"""
+查询需求分级相关三张表各自的最新业务日,供未指定 biz_dt 时先探测可用日期。
+"""
+from __future__ import annotations
+
+import logging
+
+from supply_agent.tools import tool
+from supply_infra.db.repositories.category_tree_weight_repo import CategoryTreeWeightRepository
+from supply_infra.db.repositories.demand_popularity_stats_repo import (
+    DemandPopularityStatsRepository,
+)
+from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
+from supply_infra.db.session import get_session
+
+logger = logging.getLogger(__name__)
+
+
+@tool
+def query_latest_biz_dt() -> str:
+    """
+    查询需求池、类目树权重、需求词热度统计三张表各自的最新业务日。
+
+    三张表由不同 job 产出,可能不完全同步(权重/热度统计可能滞后于需求池)。
+    未指定 biz_dt 时应先调用本工具,再根据返回结果决定分级使用哪个业务日
+    (通常以需求池的最新日为准;若权重/热度统计当日还没跑出来,可退到它们各自的最新日)。
+
+    Returns:
+        三张表各自的最新业务日,例如:
+        multi_demand_pool_di=20260716
+        category_tree_weight=20260716
+        demand_popularity_stats=20260715
+        (某表暂无数据时显示为 无数据)
+    """
+    try:
+        with get_session() as session:
+            pool_dt = MultiDemandPoolDiRepository(session).get_latest_biz_dt()
+            weight_dt = CategoryTreeWeightRepository(session).get_latest_biz_dt()
+            stats_dt = DemandPopularityStatsRepository(session).get_latest_biz_dt()
+
+        lines = [
+            f"multi_demand_pool_di={pool_dt or '无数据'}",
+            f"category_tree_weight={weight_dt or '无数据'}",
+            f"demand_popularity_stats={stats_dt or '无数据'}",
+        ]
+        message = "\n".join(lines)
+        logger.info("query_latest_biz_dt completed: %s", message.replace("\n", " "))
+        return message
+
+    except Exception as e:
+        logger.error("query_latest_biz_dt failed: %s", e, exc_info=True)
+        return f"查询最新业务日失败: {e}"
+
+
+def main() -> None:
+    print(query_latest_biz_dt())
+
+
+if __name__ == "__main__":
+    main()

+ 133 - 0
agents/demand_grade_agent/tools/query_score_distribution.py

@@ -0,0 +1,133 @@
+"""查询分类树、需求自身与后验的独立分布,供批量分级统一口径。"""
+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__)
+
+_DIM_FIELDS: tuple[tuple[str, str], ...] = (
+    ("total_score", "全局热度total_score"),
+)
+
+
+def _format_dist(label: str, dist: dict) -> str:
+    if dist["count"] == 0:
+        return f"{label}: 无数据"
+    return (
+        f"{label} (n={dist['count']}): "
+        f"min={dist['min']:.4f} p25={dist['p25']:.4f} p50={dist['p50']:.4f} "
+        f"p75={dist['p75']:.4f} p90={dist['p90']:.4f} max={dist['max']:.4f}"
+    )
+
+
+@tool
+def query_score_distribution(biz_dt: Optional[str] = None) -> str:
+    """
+    查询指定业务日分类树全局热度、需求自身来源归一分与后验分布。
+
+    建议在批量分级任务开始时调用一次,分别参考各自分位数制定本批次统一的分档阈值
+    (例如 total_score 前 10% 视为全局热度很高),避免同一批次内多次判断标准漂移。
+    需求自身分按来源内 rank 归一后对已有来源取均值,不直接合并跨来源 raw weight;
+    后验 real_rov_7d_avg 的分布只统计 real_rov_7d_count>0(有真实验证数据)的子集,
+    因为无验证数据的行 avg 无意义。
+
+    Args:
+        biz_dt: 业务日期 YYYYMMDD,可选;不传则使用 category_tree_weight 最新业务日。
+
+    Returns:
+        每个维度一行分布摘要,例如:
+        全局热度total_score (n=1500): min=0.0000 p25=0.8500 ...
+        后验real_rov_7d_avg(仅count>0子集) (n=210): ...
+    """
+    normalized_dt, err = normalize_biz_dt(biz_dt)
+    if err:
+        return err
+
+    try:
+        with get_session() as session:
+            repo = CategoryTreeWeightRepository(session)
+            resolved_dt = normalized_dt or repo.get_latest_biz_dt()
+            if not resolved_dt:
+                return "category_tree_weight 暂无数据"
+
+            weights = repo.list_by_biz_dt(resolved_dt)
+            if not weights:
+                return f"biz_dt={resolved_dt} category_tree_weight 无数据"
+
+            node_count = len(weights)
+            dim_values: dict[str, list[float]] = {}
+            for field, _ in _DIM_FIELDS:
+                if field == "total_score":
+                    dim_values[field] = [
+                        v
+                        for v in (to_float(getattr(w, field)) for w in weights)
+                        if v is not None and v > 0
+                    ]
+                else:
+                    count_field = field.replace("_avg", "_count")
+                    dim_values[field] = [
+                        v
+                        for w in weights
+                        if int(getattr(w, count_field, 0) or 0) > 0
+                        for v in [to_float(getattr(w, field))]
+                        if v is not None
+                    ]
+            posterior_values = [
+                v
+                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:
+            lines.append(_format_dist(label, distribution_summary(dim_values[field])))
+
+        lines.append(
+            _format_dist(
+                f"后验real_rov_7d_avg(仅count>0子集,共{len(posterior_values)}个节点有验证数据)",
+                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)
+        return message
+
+    except Exception as e:
+        logger.error("query_score_distribution failed: %s", e, exc_info=True)
+        return f"查询分数分布失败: {e}"
+
+
+def main() -> None:
+    print(query_score_distribution())
+
+
+if __name__ == "__main__":
+    main()

+ 116 - 0
agents/demand_grade_agent/tools/search_related_pool_demands.py

@@ -0,0 +1,116 @@
+"""
+在需求池中按同名/包含关系搜索需求词,用于合并同语义、措辞不同的需求一起判断。
+"""
+from __future__ import annotations
+
+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,
+    normalize_str_list,
+)
+from supply_agent.tools import tool
+from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
+from supply_infra.db.session import get_session
+
+logger = logging.getLogger(__name__)
+
+
+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)
+
+    if not rows:
+        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"])
+        weight = format_score(row["weight"])
+        video_count = row["video_count"] if row["video_count"] is not None else "—"
+        lines.append(
+            f"[id={row['id']}|{row['strategy']}|weight={weight}"
+            f"|视频数={video_count}|真实ROV={rov}|真实VOV={vov}] {row['demand_name']}"
+        )
+    return "\n".join(lines)
+
+
+@tool
+def search_related_pool_demands(biz_dt: str, keywords: list[str]) -> str:
+    """
+    按同名/包含关系搜索需求池,合并同语义需求的多条记录一起判断。
+
+    支持批量传入多个 keyword,同一 biz_dt 下一次调用返回各关键词的匹配结果;
+    每段结果前会标注原始 keyword。
+
+    匹配规则为双向包含:keyword 是 demand_name 的子串,或 demand_name 是 keyword 的子串。
+    用于发现措辞不同但语义相同/高度相关的需求词(例如「减脂期加餐」与「减脂加餐」),
+    在判级时应把这些记录一起纳入参考,而不是只看单条记录。
+
+    Args:
+        biz_dt: 业务日期,格式 YYYYMMDD(必填)。
+        keywords: 需求词或关键片段列表,可一次传多个。
+
+    Returns:
+        每个 keyword 一段,段首标注 `--- keyword: xxx ---`,例如:
+        --- keyword: 加餐 ---
+        [id=101|strategy_a|weight=3.20|视频数=12|真实ROV=0.0410|真实VOV=0.0021] 减脂期加餐怎么吃
+    """
+    normalized, err = normalize_biz_dt(biz_dt)
+    if err:
+        return err
+    if not normalized:
+        return "biz_dt 不能为空"
+
+    keyword_list, err = normalize_str_list(keywords, "keywords")
+    if err:
+        return err
+
+    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,
+                    priority_index,
+                )
+                sections.append(f"--- keyword: {keyword} ---\n{result}")
+
+        message = "\n\n".join(sections)
+        logger.info(
+            "search_related_pool_demands completed: biz_dt=%s count=%d",
+            normalized,
+            len(keyword_list),
+        )
+        return message
+
+    except Exception as e:
+        logger.error("search_related_pool_demands failed: %s", e, exc_info=True)
+        return f"搜索关联需求词失败: {e}"
+
+
+def main() -> None:
+    print(search_related_pool_demands(biz_dt="20260714", keywords=["加餐", "减脂"]))
+
+
+if __name__ == "__main__":
+    main()

+ 279 - 0
agents/demand_grade_agent/tools/shared.py

@@ -0,0 +1,279 @@
+"""demand_grade_agent 工具共享辅助函数(非 @tool,不对外暴露为工具)。"""
+from __future__ import annotations
+
+import json
+from decimal import Decimal
+from typing import Any
+
+from supply_infra.db.models.global_tree_category import GlobalTreeCategory
+from supply_infra.db.models.multi_demand_pool_di import MultiDemandPoolDi
+
+_VIDEO_LIST_LIMIT = 10
+
+VALID_GRADES: tuple[str, ...] = ("S", "A", "B", "C", "D")
+
+
+def normalize_biz_dt(biz_dt: str | None) -> tuple[str | None, str | None]:
+    """校验并规范化 biz_dt(YYYYMMDD);空值返回 (None, None) 表示未指定。"""
+    if biz_dt is None:
+        return None, None
+    text = str(biz_dt).strip()
+    if not text:
+        return None, None
+    if len(text) != 8 or not text.isdigit():
+        return None, f"biz_dt 格式无效,应为 YYYYMMDD: {biz_dt!r}"
+    return text, None
+
+
+def to_float(value: Any) -> float | None:
+    """将 Decimal/None/数值统一转换为 float,None 原样返回。"""
+    if value is None:
+        return None
+    return float(value)
+
+
+def format_score(avg: Decimal | float | int | None) -> str:
+    """格式化分数为易读文本,None 显示为 —。"""
+    if avg is None:
+        return "—"
+    value = float(avg)
+    if value >= 100:
+        return f"{value:.0f}"
+    if value >= 1:
+        return f"{value:.2f}"
+    return f"{value:.4f}"
+
+
+def format_dim_with_count(
+    avg: Decimal | float | int | None,
+    count: int | None,
+) -> str:
+    """格式化带样本数的维度分;count<=0 时显示 —(无数据)。"""
+    sample_count = int(count or 0)
+    if sample_count <= 0:
+        return "—"
+    return f"{format_score(avg)}(n={sample_count})"
+
+
+def build_category_heat_summary(weight: Any | None, *, biz_dt: str | None = None) -> dict[str, Any]:
+    """构建分类节点的全局热度与后验摘要(不含四维先验明细)。"""
+    if weight is None:
+        return {
+            "biz_dt": biz_dt,
+            "global_heat": {"total_score": None},
+            "posterior": {
+                "real_rov_7d": {"avg": None, "count": 0},
+                "real_vov_7d": {"avg": None, "count": 0},
+            },
+            "has_weight_data": False,
+        }
+    return {
+        "biz_dt": getattr(weight, "biz_dt", biz_dt),
+        "global_heat": {"total_score": to_float(weight.total_score)},
+        "posterior": {
+            "real_rov_7d": {
+                "avg": to_float(weight.real_rov_7d_avg),
+                "count": int(weight.real_rov_7d_count or 0),
+            },
+            "real_vov_7d": {
+                "avg": to_float(weight.real_vov_7d_avg),
+                "count": int(weight.real_vov_7d_count or 0),
+            },
+        },
+        "has_weight_data": True,
+    }
+
+
+def format_category_weight_lines(
+    weight: Any | None,
+    *,
+    category_id: int,
+    name: str | None,
+    path: str | None,
+    biz_dt: str | None = None,
+    hung_word_count: int = 0,
+) -> list[str]:
+    """格式化单个分类节点的全局热度与后验文本块。"""
+    detail = build_category_heat_summary(weight, biz_dt=biz_dt)
+    global_heat = detail["global_heat"]
+    posterior = detail["posterior"]
+    posterior_note = (
+        "有后验验证数据"
+        if int(posterior["real_rov_7d"]["count"]) > 0
+        else "无后验验证数据(效果未知)"
+    )
+    return [
+        f"[{category_id}]{name or ''} {path or '(未知路径)'}",
+        f"  biz_dt={detail.get('biz_dt') or biz_dt or '—'}  hung_word_count={hung_word_count}",
+        f"  全局热度total_score={format_score(global_heat['total_score'])}",
+        (
+            "  后验real_rov_7d="
+            f"{format_dim_with_count(posterior['real_rov_7d']['avg'], posterior['real_rov_7d']['count'])} "
+            f"后验real_vov_7d="
+            f"{format_dim_with_count(posterior['real_vov_7d']['avg'], posterior['real_vov_7d']['count'])} "
+            f"({posterior_note})"
+        ),
+    ]
+
+
+def percentile(sorted_values: list[float], pct: float) -> float | None:
+    """对已排序(升序)的数值列表求分位数(线性插值),pct 取 0~100。"""
+    if not sorted_values:
+        return None
+    if len(sorted_values) == 1:
+        return sorted_values[0]
+
+    rank = (pct / 100) * (len(sorted_values) - 1)
+    lower_idx = int(rank)
+    upper_idx = min(lower_idx + 1, len(sorted_values) - 1)
+    frac = rank - lower_idx
+    return sorted_values[lower_idx] + (sorted_values[upper_idx] - sorted_values[lower_idx]) * frac
+
+
+def distribution_summary(values: list[float]) -> dict[str, float | int | None]:
+    """返回一组数值的 min/p25/p50/p75/p90/max/count 分布摘要。"""
+    if not values:
+        return {"count": 0, "min": None, "p25": None, "p50": None, "p75": None, "p90": None, "max": None}
+    ordered = sorted(values)
+    return {
+        "count": len(ordered),
+        "min": ordered[0],
+        "p25": percentile(ordered, 25),
+        "p50": percentile(ordered, 50),
+        "p75": percentile(ordered, 75),
+        "p90": percentile(ordered, 90),
+        "max": ordered[-1],
+    }
+
+
+def _normalize_parent_id(parent_id: int | None) -> int | None:
+    if parent_id is None or parent_id == 0:
+        return None
+    return parent_id
+
+
+def build_category_path(
+    category_id: int,
+    by_id: dict[int, GlobalTreeCategory],
+) -> str | None:
+    """从根到指定节点的名称路径,如「美妆护肤 > 护肤 > 防晒」。"""
+    cat = by_id.get(category_id)
+    if cat is None:
+        return None
+    names: list[str] = []
+    current: GlobalTreeCategory | None = cat
+    seen: set[int] = set()
+    while current is not None:
+        cid = int(current.id)
+        if cid in seen:
+            break
+        seen.add(cid)
+        names.append(current.name or str(cid))
+        parent_key = _normalize_parent_id(current.parent_id)
+        current = by_id.get(parent_key) if parent_key is not None else None
+    names.reverse()
+    return " > ".join(names)
+
+
+def normalize_str_list(raw: Any, field_name: str = "items") -> tuple[list[str], str | None]:
+    """将 JSON 文本/列表/单个字符串统一解析为非空 str 列表(去重保序)。"""
+    if raw is None:
+        return [], f"{field_name} 不能为空"
+
+    items: list[Any]
+    if isinstance(raw, str):
+        text = raw.strip()
+        if not text:
+            return [], f"{field_name} 不能为空"
+        try:
+            parsed = json.loads(text)
+            items = list(parsed) if isinstance(parsed, list) else [text]
+        except (ValueError, TypeError):
+            items = [text]
+    elif isinstance(raw, (list, tuple)):
+        items = list(raw)
+    else:
+        items = [raw]
+
+    out: list[str] = []
+    seen: set[str] = set()
+    for item in items:
+        text = str(item).strip() if item is not None else ""
+        if not text or text in seen:
+            continue
+        seen.add(text)
+        out.append(text)
+
+    if not out:
+        return [], f"{field_name} 不能为空"
+    return out, None
+
+
+def parse_int_list(raw: Any) -> list[int]:
+    """将 JSON 文本/列表统一解析为 int 列表,解析失败返回空列表。"""
+    if raw is None:
+        return []
+    if isinstance(raw, str):
+        try:
+            raw = json.loads(raw)
+        except (ValueError, TypeError):
+            return []
+    if not isinstance(raw, list):
+        return []
+    out: list[int] = []
+    for item in raw:
+        try:
+            out.append(int(item))
+        except (TypeError, ValueError):
+            continue
+    return out
+
+
+def dump_int_list(values: list[int] | None) -> str | None:
+    """将 int 列表序列化为 JSON 文本,空列表/None 返回 None。"""
+    if not values:
+        return None
+    return json.dumps(values, ensure_ascii=False)
+
+
+def _parse_video_ids(raw: Any) -> list[str]:
+    """解析单个 multi_demand_pool_di.video_list(JSON数组/逗号分隔文本)为 vid 字符串列表。"""
+    if raw is None:
+        return []
+    items: list[Any]
+    if isinstance(raw, str):
+        text = raw.strip()
+        if not text:
+            return []
+        try:
+            parsed = json.loads(text)
+            items = list(parsed) if isinstance(parsed, list) else [text]
+        except (ValueError, TypeError):
+            items = [part.strip() for part in text.split(",") if part.strip()]
+    elif isinstance(raw, (list, tuple)):
+        items = list(raw)
+    else:
+        return []
+    return [str(v).strip() for v in items if v is not None and str(v).strip()]
+
+
+def merge_video_ids(pool_rows: list[MultiDemandPoolDi], limit: int = _VIDEO_LIST_LIMIT) -> str | None:
+    """合并多条原始需求行的 video_list,去重保序,最多取前 limit 个,返回 JSON 文本或 None。"""
+    merged: list[str] = []
+    seen: set[str] = set()
+    for row in pool_rows:
+        for vid in _parse_video_ids(row.video_list):
+            if vid not in seen:
+                seen.add(vid)
+                merged.append(vid)
+    if not merged:
+        return None
+    return json.dumps(merged[:limit], ensure_ascii=False)
+
+
+def collect_strategies(pool_rows: list[MultiDemandPoolDi]) -> str | None:
+    """收集多条原始需求行的策略名,去重排序,返回 JSON 文本或 None。"""
+    strategies = sorted({row.strategy.strip() for row in pool_rows if row.strategy and row.strategy.strip()})
+    if not strategies:
+        return None
+    return json.dumps(strategies, ensure_ascii=False)

+ 251 - 0
agents/demand_grade_agent/tools/tree_local.py

@@ -0,0 +1,251 @@
+"""demand_grade_agent 专用的全局树局部热度查询辅助(不对外注册为工具)。"""
+from __future__ import annotations
+
+from collections import defaultdict
+from types import SimpleNamespace
+from typing import Any
+
+from agents.demand_grade_agent.tools.shared import (
+    build_category_heat_summary,
+    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
+
+
+def _normalize_parent_id(parent_id: int | None) -> int | None:
+    if parent_id is None or parent_id == 0:
+        return None
+    return int(parent_id)
+
+
+def _materialize_category(row: Any) -> SimpleNamespace:
+    """在 session 内抽出标量,避免 DetachedInstanceError。"""
+    return SimpleNamespace(
+        id=int(row.id),
+        name=row.name,
+        parent_id=row.parent_id,
+        level=row.level,
+    )
+
+
+def _materialize_weight(row: Any) -> SimpleNamespace:
+    """在 session 内抽出标量,避免 DetachedInstanceError。"""
+    return SimpleNamespace(
+        category_id=int(row.category_id),
+        biz_dt=row.biz_dt,
+        total_score=row.total_score,
+        hung_word_count=row.hung_word_count,
+        real_rov_7d_avg=row.real_rov_7d_avg,
+        real_rov_7d_count=row.real_rov_7d_count,
+        real_vov_7d_avg=row.real_vov_7d_avg,
+        real_vov_7d_count=row.real_vov_7d_count,
+    )
+
+
+def _score_value(weight: Any | None) -> float | None:
+    if weight is None or weight.total_score is None:
+        return None
+    return float(weight.total_score)
+
+
+def load_local_tree_state(biz_dt: str) -> tuple[dict[int, Any], dict[int | None, list[int]], dict[int, Any]]:
+    with get_session() as session:
+        categories = [
+            _materialize_category(row)
+            for row in GlobalTreeCategoryRepository(session).list_active_categories()
+        ]
+        weights = [
+            _materialize_weight(row)
+            for row in CategoryTreeWeightRepository(session).list_by_biz_dt(biz_dt)
+        ]
+    by_id = {row.id: row for row in categories}
+    children: dict[int | None, list[int]] = defaultdict(list)
+    for row in categories:
+        children[_normalize_parent_id(row.parent_id)].append(row.id)
+    for ids in children.values():
+        ids.sort()
+    weight_by_id = {row.category_id: row for row in weights}
+    return by_id, children, weight_by_id
+
+
+def _describe_node(
+    category_id: int,
+    by_id: dict[int, Any],
+    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
+            ),
+        },
+    }
+
+
+def _sibling_ids(category_id: int, parent_id: int | None, children: dict[int | None, list[int]]) -> list[int]:
+    if parent_id is not None:
+        return list(children.get(parent_id, []))
+    return list(children.get(None, []))
+
+
+def build_local_heat_snapshot(
+    biz_dt: str,
+    category_ids: list[int],
+    *,
+    tree_state: tuple[dict[int, Any], dict[int | None, list[int]], dict[int, Any]] | None = None,
+) -> list[dict[str, Any]]:
+    """构建局部环境快照:自身、父节点、全部兄弟节点的全局热度与后验。"""
+    if tree_state is None:
+        by_id, children, weight_by_id = load_local_tree_state(biz_dt)
+    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:
+            snapshots.append({
+                "category_id": category_id,
+                "error": "分类不存在",
+            })
+            continue
+
+        parent_id = _normalize_parent_id(category.parent_id)
+        sibling_ids = _sibling_ids(category_id, parent_id, children)
+        ranked_siblings = sorted(
+            sibling_ids,
+            key=lambda cid: (
+                _score_value(weight_by_id.get(cid)) is None,
+                -(_score_value(weight_by_id.get(cid)) or 0),
+                cid,
+            ),
+        )
+        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,
+                global_positions=global_positions,
+            )
+            item["rank_among_siblings"] = index
+            item["sibling_count"] = sibling_total
+            item["is_self"] = sibling_id == category_id
+            siblings.append(item)
+
+        snapshots.append({
+            "category_id": category_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,
+                    global_positions=global_positions,
+                )
+                if parent_id is not None
+                else None
+            ),
+            "siblings": siblings,
+            "sibling_count": sibling_total,
+        })
+    return snapshots
+
+
+def _append_node_block(
+    lines: list[str],
+    *,
+    title: str,
+    node: dict[str, Any] | None,
+    biz_dt: str,
+    weight_by_id: dict[int, Any] | None = None,
+) -> None:
+    if node is None:
+        lines.append(f"{title}: 无")
+        return
+    lines.append(f"{title}:")
+    weight = weight_by_id.get(int(node["category_id"])) if weight_by_id is not None else None
+    lines.extend(format_category_weight_lines(
+        weight,
+        category_id=int(node["category_id"]),
+        name=node.get("name"),
+        path=node.get("path"),
+        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]:
+    if snapshot.get("error"):
+        return [f"--- category_id={snapshot['category_id']} ---", f"错误: {snapshot['error']}"]
+
+    biz_dt = snapshot["biz_dt"]
+    lines = [f"--- category_id={snapshot['category_id']} ---"]
+    _append_node_block(lines, title="自身节点", node=snapshot["self"], biz_dt=biz_dt, weight_by_id=weight_by_id)
+    _append_node_block(lines, title="父节点", node=snapshot.get("parent"), biz_dt=biz_dt, weight_by_id=weight_by_id)
+
+    siblings = snapshot.get("siblings") or []
+    lines.append(f"全部兄弟节点(共 {len(siblings)} 个,按 total_score 降序,* 为当前节点):")
+    for sibling in siblings:
+        marker = "*" if sibling.get("is_self") else " "
+        rank_note = f" 兄弟排名 {sibling.get('rank_among_siblings')}/{sibling.get('sibling_count')}"
+        lines.append(f"{marker} {'-' * 8}{rank_note}")
+        _append_node_block(lines, title="  节点", node=sibling, biz_dt=biz_dt, weight_by_id=weight_by_id)
+    return lines
+
+
+def render_local_heat_report(biz_dt: str, category_ids: list[int]) -> str:
+    tree_state = load_local_tree_state(biz_dt)
+    _, _, weight_by_id = tree_state
+    snapshots = build_local_heat_snapshot(biz_dt, category_ids, tree_state=tree_state)
+    if not snapshots:
+        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))
+        lines.append("")
+    return "\n".join(lines).rstrip()

+ 5 - 0
agents/demand_grade_orchestrator_agent/__init__.py

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

+ 134 - 0
agents/demand_grade_orchestrator_agent/_verify_logic.py

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

+ 24 - 0
agents/demand_grade_orchestrator_agent/agent.py

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

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

@@ -0,0 +1,27 @@
+"""统筹规划 Agent 的共享数据访问与格式化逻辑。"""
+
+from agents.demand_grade_orchestrator_agent.common.plan_record import prepare_grade_groups
+from agents.demand_grade_orchestrator_agent.common.tree_state import (
+    HEAT_LEVEL_DEFINITION,
+    format_demand_count,
+    format_heat_score,
+    format_rank,
+    global_heat_positions,
+    has_hung_demand,
+    heat_level,
+    load_tree_state,
+    path,
+)
+
+__all__ = [
+    "prepare_grade_groups",
+    "HEAT_LEVEL_DEFINITION",
+    "format_demand_count",
+    "format_heat_score",
+    "format_rank",
+    "global_heat_positions",
+    "has_hung_demand",
+    "heat_level",
+    "load_tree_state",
+    "path",
+]

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

@@ -0,0 +1,123 @@
+"""当天节点分配状态查询(不含校验重试逻辑)。"""
+from __future__ import annotations
+
+from typing import Any
+
+from agents.demand_grade_orchestrator_agent.common.tree_state import has_hung_demand, load_tree_state
+from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
+from supply_infra.db.session import get_session
+
+MAX_DAILY_BATCHES = 200
+
+
+def get_required_hanging_category_ids(biz_dt: str) -> set[int]:
+    by_id, _children, weights = load_tree_state(biz_dt)
+    return {
+        category_id
+        for category_id, weight in weights.items()
+        if category_id in by_id and has_hung_demand(weight)
+    }
+
+
+def get_assigned_category_ids(biz_dt: str) -> set[int]:
+    with get_session() as session:
+        return DemandGradePlanRepository(session).get_assigned_category_ids(biz_dt)
+
+
+def get_existing_group_count(biz_dt: str) -> int:
+    with get_session() as session:
+        snapshot = DemandGradePlanRepository(session).get_execution_snapshot(biz_dt)
+    return int(snapshot["planned_groups"])
+
+
+def resolve_planning_state(biz_dt: str) -> dict[str, Any]:
+    """汇总当天有需求节点、已分批节点与剩余批次额度。"""
+    required = get_required_hanging_category_ids(biz_dt)
+    assigned = get_assigned_category_ids(biz_dt)
+    unassigned = sorted(required - assigned)
+    existing_groups = get_existing_group_count(biz_dt)
+    remaining_batch_quota = max(0, MAX_DAILY_BATCHES - existing_groups)
+    batch_limit_reached = existing_groups >= MAX_DAILY_BATCHES
+    has_unassigned_nodes = bool(unassigned)
+    can_plan_more = has_unassigned_nodes and not batch_limit_reached
+    skip_reason: str | None = None
+    if batch_limit_reached:
+        skip_reason = f"当天批次已达上限 {MAX_DAILY_BATCHES}"
+    elif not has_unassigned_nodes:
+        skip_reason = "当天有需求节点均已分批,无待分配节点"
+    return {
+        "biz_dt": biz_dt,
+        "total_hanging_nodes": len(required),
+        "required_category_ids": sorted(required),
+        "assigned_category_ids": sorted(assigned),
+        "unassigned_category_ids": unassigned,
+        "existing_groups": existing_groups,
+        "remaining_batch_quota": remaining_batch_quota,
+        "batch_limit_reached": batch_limit_reached,
+        "has_unassigned_nodes": has_unassigned_nodes,
+        "can_plan_more": can_plan_more,
+        "skip_reason": skip_reason,
+    }
+
+
+def get_unassigned_hanging_category_ids(biz_dt: str) -> set[int]:
+    """返回当天有需求且尚未进入任何批次的分类节点。"""
+    return get_required_hanging_category_ids(biz_dt) - get_assigned_category_ids(biz_dt)
+
+
+def _sync_group_positions(group: dict[str, Any]) -> dict[str, Any]:
+    category_ids = {int(value) for value in group.get("category_ids") or []}
+    positions = group.get("category_global_positions")
+    if not isinstance(positions, list):
+        return group
+    return {
+        **group,
+        "category_global_positions": [
+            item for item in positions
+            if int(item.get("category_id", -1)) in category_ids
+        ],
+    }
+
+
+def dedupe_cross_group_category_ids(plan: dict[str, Any]) -> list[int]:
+    """同计划内后批次若含前批已出现的节点,从后批中移除。"""
+    removed: list[int] = []
+    seen: set[int] = set()
+    cleaned_groups: list[dict[str, Any]] = []
+    for group in plan.get("groups") or []:
+        kept: list[int] = []
+        for category_id in group.get("category_ids") or []:
+            try:
+                cid = int(category_id)
+            except (TypeError, ValueError):
+                continue
+            if cid in seen:
+                removed.append(cid)
+                continue
+            seen.add(cid)
+            kept.append(cid)
+        if kept:
+            cleaned_groups.append(_sync_group_positions({**group, "category_ids": kept}))
+    plan["groups"] = cleaned_groups
+    return sorted(set(removed))
+
+
+def strip_assigned_category_ids(plan: dict[str, Any], assigned_ids: set[int]) -> list[int]:
+    """从计划中移除当天已分批的分类,避免重复落库。"""
+    removed: list[int] = []
+    cleaned_groups: list[dict[str, Any]] = []
+    for group in plan.get("groups") or []:
+        kept: list[int] = []
+        for category_id in group.get("category_ids") or []:
+            try:
+                cid = int(category_id)
+            except (TypeError, ValueError):
+                continue
+            if cid in assigned_ids:
+                removed.append(cid)
+            else:
+                kept.append(cid)
+        if kept:
+            cleaned_groups.append(_sync_group_positions({**group, "category_ids": kept}))
+    plan["groups"] = cleaned_groups
+    return sorted(set(removed))

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

@@ -0,0 +1,102 @@
+"""批次计划逐条入库(由 save_grade_plan 调用)。"""
+from __future__ import annotations
+
+import logging
+from typing import Any
+
+from agents.demand_grade_orchestrator_agent.common.assignment import (
+    MAX_DAILY_BATCHES,
+    get_assigned_category_ids,
+    get_existing_group_count,
+    get_required_hanging_category_ids,
+    get_unassigned_hanging_category_ids,
+)
+from agents.demand_grade_orchestrator_agent.common.tree_state import has_hung_demand, load_tree_state
+from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
+from supply_infra.db.session import get_session
+
+logger = logging.getLogger(__name__)
+
+
+def _filter_group_category_ids(biz_dt: str, category_ids: list[int]) -> list[int]:
+    """入库前再次剔除已分配与无需求节点。"""
+    assigned = get_assigned_category_ids(biz_dt)
+    _by_id, _children, weights = load_tree_state(biz_dt)
+    kept: list[int] = []
+    for category_id in category_ids:
+        if category_id in assigned:
+            continue
+        if not has_hung_demand(weights.get(category_id)):
+            continue
+        kept.append(category_id)
+    return kept
+
+
+def persist_groups_one_by_one(biz_dt: str, base_payload: dict[str, Any]) -> dict[str, Any]:
+    """逐批入库,单批失败不影响其余批次;额度用尽时停止。本函数不向外抛异常。"""
+    total_hanging_nodes = len(get_required_hanging_category_ids(biz_dt))
+    persisted_group_count = 0
+    persisted_groups: list[dict[str, Any]] = []
+    skipped_quota = 0
+    skipped_empty = 0
+    failed_groups: list[dict[str, Any]] = []
+
+    for group in base_payload.get("groups") or []:
+        try:
+            if get_existing_group_count(biz_dt) >= MAX_DAILY_BATCHES:
+                skipped_quota += 1
+                continue
+
+            category_ids = _filter_group_category_ids(biz_dt, list(group.get("category_ids") or []))
+            if not category_ids:
+                skipped_empty += 1
+                continue
+
+            unassigned_before = get_unassigned_hanging_category_ids(biz_dt)
+            single_group = {**group, "category_ids": category_ids}
+            single_payload = {
+                **base_payload,
+                "total_hanging_nodes": total_hanging_nodes,
+                "groups": [single_group],
+                "covered_category_ids": category_ids,
+                "uncovered_category_ids": sorted(unassigned_before - set(category_ids)),
+                "coverage_complete": not (unassigned_before - set(category_ids)),
+            }
+
+            with get_session() as session:
+                DemandGradePlanRepository(session).create_plan(biz_dt, single_payload)
+            persisted_group_count += 1
+            persisted_groups.append(single_group)
+        except Exception as exc:
+            logger.exception(
+                "单批入库失败,已跳过并继续: biz_dt=%s group_id=%s category_ids=%s",
+                biz_dt,
+                group.get("group_id"),
+                group.get("category_ids"),
+            )
+            failed_groups.append({
+                "group_id": group.get("group_id"),
+                "category_ids": group.get("category_ids"),
+                "error": str(exc),
+            })
+
+    try:
+        unassigned_after = sorted(get_unassigned_hanging_category_ids(biz_dt))
+        existing_groups = get_existing_group_count(biz_dt)
+    except Exception as exc:
+        logger.exception("读取入库后状态失败: biz_dt=%s", biz_dt)
+        unassigned_after = []
+        existing_groups = persisted_group_count
+
+    return {
+        "persisted_group_count": persisted_group_count,
+        "persisted_groups": persisted_groups,
+        "skipped_quota": skipped_quota,
+        "skipped_empty": skipped_empty,
+        "failed_groups": failed_groups,
+        "existing_groups": existing_groups,
+        "remaining_batch_quota": max(0, MAX_DAILY_BATCHES - existing_groups),
+        "unassigned_category_ids": unassigned_after,
+        "coverage_complete": not unassigned_after,
+        "total_hanging_nodes": total_hanging_nodes,
+    }

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

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

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

@@ -0,0 +1,99 @@
+"""全局树热度状态的加载与格式化。"""
+from __future__ import annotations
+
+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
+
+
+def load_tree_state(biz_dt: str) -> tuple[dict[int, Any], dict[int | None, list[int]], dict[int, Any]]:
+    with get_session() as session:
+        categories = [
+            SimpleNamespace(id=int(row.id), name=row.name, parent_id=row.parent_id, level=row.level)
+            for row in GlobalTreeCategoryRepository(session).list_active_categories()
+        ]
+        weights = [
+            SimpleNamespace(category_id=int(row.category_id), total_score=row.total_score, hung_word_count=row.hung_word_count)
+            for row in CategoryTreeWeightRepository(session).list_by_biz_dt(biz_dt)
+        ]
+    by_id = {row.id: row for row in categories}
+    children: dict[int | None, list[int]] = defaultdict(list)
+    for row in categories:
+        parent_id = int(row.parent_id) if row.parent_id not in (None, 0) else None
+        children[parent_id].append(row.id)
+    for ids in children.values():
+        ids.sort()
+    return by_id, children, {row.category_id: row for row in weights}
+
+
+def format_heat_score(weight: Any | None) -> str:
+    """格式化原始 total_score:只保留两位小数,无数据时明确为 null。"""
+    if weight is None or weight.total_score is None:
+        return "null"
+    return f"{float(weight.total_score):.2f}"
+
+
+def has_hung_demand(weight: Any | None) -> bool:
+    """判断分类自身是否挂有需求。"""
+    return weight is not None and int(weight.hung_word_count or 0) > 0
+
+
+def format_demand_count(weight: Any | None) -> str:
+    """格式化节点挂载需求数量;无需求时返回空字符串。"""
+    if not has_hung_demand(weight):
+        return ""
+    return str(int(weight.hung_word_count or 0))
+
+
+HEAT_LEVEL_DEFINITION: dict[str, dict[str, str | float | None]] = {
+    "S": {"label": "高热", "min_global_rank_score": 0.90},
+    "A": {"label": "较高热", "min_global_rank_score": 0.75},
+    "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
+    while current is not None:
+        names.append(current.name or str(current.id))
+        parent_id = int(current.parent_id) if current.parent_id not in (None, 0) else None
+        current = by_id.get(parent_id) if parent_id is not None else None
+    return " > ".join(reversed(names))

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

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

+ 158 - 0
agents/demand_grade_orchestrator_agent/run.py

@@ -0,0 +1,158 @@
+"""运行统筹规划 Agent(save_grade_plan 工具内直接入库)。"""
+from __future__ import annotations
+
+import json
+import logging
+from typing import Any
+
+from agents.demand_grade_orchestrator_agent import create_demand_grade_orchestrator_agent
+from agents.demand_grade_orchestrator_agent.common.assignment import MAX_DAILY_BATCHES, resolve_planning_state
+from supply_agent.types import Role
+from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
+from supply_infra.db.session import get_session
+
+logger = logging.getLogger(__name__)
+
+
+def _summarize_agent_saves(result: Any) -> dict[str, Any]:
+    """统计 Agent 会话中 save_grade_plan 成功入库的次数与批次数。"""
+    save_count = 0
+    persisted_groups = 0
+    last_payload: dict[str, Any] | None = None
+    for message in result.messages:
+        if message.role != Role.TOOL or message.name != "save_grade_plan":
+            continue
+        try:
+            payload = json.loads(message.content or "")
+        except json.JSONDecodeError:
+            continue
+        if not isinstance(payload, dict):
+            continue
+        if payload.get("ok") is True and payload.get("persisted") is True:
+            save_count += 1
+            persisted_groups += int(payload.get("persisted_group_count") or 0)
+            last_payload = payload
+    return {
+        "save_count": save_count,
+        "persisted_groups": persisted_groups,
+        "last_payload": last_payload,
+    }
+
+
+def _orchestrate_via_agent(
+    biz_dt: str,
+    *,
+    planning_state: dict[str, Any],
+    unassigned_count: int,
+) -> None:
+    remaining_quota = int(planning_state["remaining_batch_quota"])
+
+    agent = create_demand_grade_orchestrator_agent()
+    try:
+        result = agent.run(
+            f"""请为业务日 {biz_dt} 制定全局树需求分级计划。
+必须从 query_global_heat_tree 开始,经过至少一次 query_heat_node_group 下钻,
+由你自行划分批次后调用 save_grade_plan 提交(可多次调用,每轮提交一部分批次)。
+
+待分批节点数:{unassigned_count}(请通过 query_global_heat_tree 查看,勿依赖本消息枚举 ID)。
+剩余可新增批次数:{remaining_quota}(每日总上限 {MAX_DAILY_BATCHES},当前已存在 {planning_state["existing_groups"]} 个批次)。
+每批最多 20 个节点;节点较多时分多轮 save_grade_plan,每轮关注工具返回的 unassigned_category_ids 与 remaining_batch_quota。
+"""
+        )
+    except Exception:
+        logger.exception("统筹 Agent 运行异常: biz_dt=%s", biz_dt)
+        return
+
+    summary = _summarize_agent_saves(result)
+    if summary["save_count"] == 0:
+        logger.warning("统筹 Agent 未成功入库任何批次: biz_dt=%s", biz_dt)
+        return
+    logger.info(
+        "统筹 Agent 完成: biz_dt=%s save_count=%s persisted_groups=%s",
+        biz_dt,
+        summary["save_count"],
+        summary["persisted_groups"],
+    )
+    last = summary["last_payload"] or {}
+    if last.get("unassigned_category_ids"):
+        logger.info(
+            "统筹后仍有未分批节点: biz_dt=%s remaining=%s quota=%s",
+            biz_dt,
+            len(last["unassigned_category_ids"]),
+            last.get("remaining_batch_quota"),
+        )
+
+
+def orchestrate_daily_grade_plan(*, biz_dt: str) -> None:
+    """运行统筹规划(批次由 Agent 划分,save_grade_plan 工具入库)。"""
+    planning_state = resolve_planning_state(biz_dt)
+
+    if planning_state["batch_limit_reached"]:
+        logger.info(
+            "跳过统筹 Agent:当天批次已达上限 biz_dt=%s existing_groups=%s limit=%s",
+            biz_dt,
+            planning_state["existing_groups"],
+            MAX_DAILY_BATCHES,
+        )
+        return
+
+    unassigned_ids = planning_state["unassigned_category_ids"]
+    if not unassigned_ids:
+        logger.info(
+            "跳过统筹 Agent:当天有需求节点均已分批 biz_dt=%s total_hanging_nodes=%s",
+            biz_dt,
+            planning_state["total_hanging_nodes"],
+        )
+        return
+
+    logger.info(
+        "执行统筹 Agent: biz_dt=%s unassigned=%s remaining_quota=%s",
+        biz_dt,
+        len(unassigned_ids),
+        planning_state["remaining_batch_quota"],
+    )
+    _orchestrate_via_agent(
+        biz_dt,
+        planning_state=planning_state,
+        unassigned_count=len(unassigned_ids),
+    )
+
+
+def main(biz_dt: str) -> dict[str, Any]:
+    """手动测试:运行统筹规划并打印当天分配摘要。"""
+    planning_before = resolve_planning_state(biz_dt)
+    orchestrate_daily_grade_plan(biz_dt=biz_dt)
+    planning_after = resolve_planning_state(biz_dt)
+    with get_session() as session:
+        repo = DemandGradePlanRepository(session)
+        plan = repo.get_latest_plan(biz_dt)
+        groups = repo.list_groups_by_biz_dt(biz_dt)
+        group_status = repo.summarize(biz_dt)
+        plan_summary = json.loads(plan.plan_json) if plan is not None else {}
+        result = {
+            "biz_dt": biz_dt,
+            "planning_before": planning_before,
+            "planning_after": planning_after,
+            "plan_count": 1 if plan is not None else 0,
+            "total_hanging_nodes": planning_after["total_hanging_nodes"],
+            "group_count": len(groups),
+            "group_status": group_status,
+            "unassigned_category_ids": planning_after["unassigned_category_ids"],
+            "sample_groups": [
+                {
+                    "group_key": group.group_key,
+                    "category_ids": json.loads(group.category_ids),
+                    "status": group.status,
+                }
+                for group in groups[:3]
+            ],
+            "latest_plan_summary": plan_summary,
+        }
+    print(json.dumps(result, ensure_ascii=False, indent=2))
+    return result
+
+
+if __name__ == "__main__":
+    import sys
+
+    main(sys.argv[1] if len(sys.argv) > 1 else "20260714")

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

@@ -0,0 +1,16 @@
+"""需求分级统筹规划 Agent 的工具。"""
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import Any
+
+from agents.demand_grade_orchestrator_agent.tools.query_global_heat_tree import query_global_heat_tree
+from agents.demand_grade_orchestrator_agent.tools.query_heat_node_group import query_heat_node_group
+from agents.demand_grade_orchestrator_agent.tools.save_grade_plan import save_grade_plan
+from supply_agent.tools.registry import ToolRegistry
+
+ALL_TOOLS: list[Callable[..., Any]] = [query_global_heat_tree, query_heat_node_group, save_grade_plan]
+
+
+def register_all_tools(registry: ToolRegistry) -> ToolRegistry:
+    return registry.from_decorated(*ALL_TOOLS)

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

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

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

@@ -0,0 +1,52 @@
+"""下钻节点,查看整树位置及父子/兄弟间的全局热度(未分批节点标 +)。"""
+from __future__ import annotations
+
+from agents.demand_grade_orchestrator_agent.common.assignment import get_unassigned_hanging_category_ids
+from agents.demand_grade_orchestrator_agent.common import (
+    format_heat_score,
+    format_rank,
+    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)
+    unassigned_nodes = get_unassigned_hanging_category_ids(biz_dt)
+    lines = [
+        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)
+        is_unassigned = category_id in unassigned_nodes
+        suffix = " +" if is_unassigned and has_hung_demand(weight) else ""
+        return (
+            f"[{category_id}]{category.name or ''}"
+            f"[{format_heat_score(weight)}|{format_rank(position)}|{heat_level(position)}]{suffix}"
+        )
+
+    for category_id in dict.fromkeys(int(value) for value in category_ids):
+        category = by_id.get(category_id)
+        if category is None:
+            continue
+        parent_id = int(category.parent_id) if category.parent_id not in (None, 0) else None
+        lines.append(f"节点:{describe(category_id)}")
+        lines.append(f"  路径:{path(category_id, by_id)}")
+        if parent_id is not None:
+            lines.append(f"  父节点:{describe(parent_id)}")
+        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)

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

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

+ 9 - 0
agents/demand_video_expand_agent/__init__.py

@@ -0,0 +1,9 @@
+"""
+demand_video_expand_agent — S/A 需求视频点位拓展判断 Agent
+
+职责:对任务层已组装好的「需求 + 视频点位」做语义判断,
+筛选可作为拓展需求的点位并落库。不负责查库。
+"""
+from agents.demand_video_expand_agent.agent import create_demand_video_expand_agent
+
+__all__ = ["create_demand_video_expand_agent"]

+ 30 - 0
agents/demand_video_expand_agent/agent.py

@@ -0,0 +1,30 @@
+"""
+demand_video_expand_agent 工厂 — 组装需求视频点位拓展判断 Agent。
+"""
+from __future__ import annotations
+
+from pathlib import Path
+
+from supply_agent import Agent
+from supply_agent.config import Settings
+from agents.demand_video_expand_agent.tools import register_all_tools
+
+_PROMPT_PATH = Path(__file__).parent / "prompt" / "system_prompt.md"
+DEMAND_VIDEO_EXPAND_AGENT_SYSTEM_PROMPT = _PROMPT_PATH.read_text(encoding="utf-8")
+
+
+def create_demand_video_expand_agent(
+    settings: Settings | None = None,
+    *,
+    model: str | None = None,
+) -> Agent:
+    """创建 demand_video_expand_agent 实例,注册专属工具。"""
+    agent = Agent(
+        settings=settings,
+        name="demand_video_expand_agent",
+        model=model,
+        system_prompt=DEMAND_VIDEO_EXPAND_AGENT_SYSTEM_PROMPT,
+        max_iterations=10,
+    )
+    register_all_tools(agent.tools)
+    return agent

+ 45 - 0
agents/demand_video_expand_agent/prompt/system_prompt.md

@@ -0,0 +1,45 @@
+## 角色与任务
+你是需求拓展判断专家。用户会提供一个已评级的需求(demand_name、grade、demand_grade_id),
+以及其关联视频的全部点位列表(point_type / point_data / point_desc / video_id)。
+
+你只负责:
+1. 判断哪些点位可作为「拓展需求」
+2. 为每条候选写清 reason(为何与原需求相近)
+3. 调用 `batch_save_demand_expansions` 落库
+
+禁止调用任何查询工具;所有数据已在用户消息中给出。
+
+## 点位类型含义
+- **purpose**(目的点):用户观看该视频的目的,拓展价值最高
+- **key**(关键点):视频核心话题,常可作为子需求/细分需求
+- **inspiration**(灵感点):创作灵感/场景,仅在与原需求强相关时采纳
+
+同等相近时,优先级:**purpose > key > inspiration**。
+
+## 相近可采纳(满足其一即可,但 reason 须写清依据)
+1. **语义包含/被包含**:原需求是上位概念,点位是下位具体话题(或相反但意图一致)
+2. **同场景细分**:同一使用场景下的更细需求表达
+3. **同意图不同表达**:措辞不同但用户意图一致
+
+## 必须剔除
+- 与 demand_name 完全相同或仅差标点/空格
+- 过于宽泛、无法作为具体供给方向(如单独的「健康饮食」「减肥」且无场景)
+- 与原需求无关的点位(同一视频但语义脱节)
+- 纯描述性语句而非需求表达(如「视频展示了三种做法」)
+- point_data 为空或无意义片段
+
+## 工作流程
+1. 阅读 demand_name 与各点位
+2. 筛选可拓展候选,整理 items
+3. 调用 `batch_save_demand_expansions`,将用户消息中的 biz_dt、demand_grade_id、
+   demand_name、grade、run_id 原样传入
+4. 若无合适拓展,传 `items=[]`,并在回复中说明原因
+5. 简要总结:采纳几条、主要剔除原因
+
+## 落库字段说明
+items 每项:
+- `expanded_text`:来自 point_data,作为拓展需求文本
+- `point_type`:inspiration / purpose / key
+- `point_desc`:点位描述(有则带上)
+- `video_id`:来源视频
+- `reason`:为何可作为拓展(必填,具体可读)

+ 91 - 0
agents/demand_video_expand_agent/run.py

@@ -0,0 +1,91 @@
+#!/usr/bin/env python3
+"""单需求视频点位拓展判断 — 由任务层调用。"""
+from __future__ import annotations
+
+from dataclasses import dataclass
+
+from supply_agent.types import AgentResult, Role
+
+_POINT_TYPE_LABEL = {
+    "inspiration": "灵感点",
+    "purpose": "目的点",
+    "key": "关键点",
+}
+
+
+@dataclass
+class VideoPoint:
+    video_id: str
+    point_type: str
+    point_data: str | None
+    point_desc: str | None
+
+
+@dataclass
+class DemandExpandContext:
+    biz_dt: str
+    run_id: str
+    demand_grade_id: int
+    demand_name: str
+    grade: str
+    score: float | None
+    video_ids: list[str]
+    points: list[VideoPoint]
+
+
+def build_expand_user_input(ctx: DemandExpandContext) -> str:
+    """构建传给 Agent 的用户消息。"""
+    score_text = f"{ctx.score:.2f}" if ctx.score is not None else "—"
+    lines = [
+        f"biz_dt={ctx.biz_dt}",
+        f"run_id={ctx.run_id}",
+        f"demand_grade_id={ctx.demand_grade_id}",
+        f"demand_name={ctx.demand_name}",
+        f"grade={ctx.grade}",
+        f"score={score_text}",
+        f"video_count={len(ctx.video_ids)}",
+        "",
+        "以下是从关联视频中提取的点位,请判断哪些可作为拓展需求:",
+    ]
+    for index, point in enumerate(ctx.points, start=1):
+        type_label = _POINT_TYPE_LABEL.get(point.point_type, point.point_type)
+        data_text = point.point_data or "—"
+        desc_text = point.point_desc or "—"
+        lines.append(
+            f"[{index}|{type_label}|video={point.video_id}] "
+            f"{data_text} | 描述: {desc_text}"
+        )
+    lines.extend(
+        [
+            "",
+            "判断完成后调用 batch_save_demand_expansions 落库;",
+            "将上述 biz_dt、demand_grade_id、demand_name、grade、run_id 原样传入工具。",
+            "若无合适拓展,传 items=[]。",
+        ]
+    )
+    return "\n".join(lines)
+
+
+def judge_demand_expansion(ctx: DemandExpandContext) -> AgentResult:
+    """对单个需求执行拓展判断。"""
+    from agents.demand_video_expand_agent.agent import create_demand_video_expand_agent
+
+    agent = create_demand_video_expand_agent()
+    user_input = build_expand_user_input(ctx)
+    return agent.run(user_input)
+
+
+def extract_saved_count(result: AgentResult) -> int:
+    """从 Agent 工具返回消息中解析保存条数。"""
+    for msg in reversed(result.messages):
+        if msg.role != Role.TOOL or not msg.content:
+            continue
+        text = str(msg.content)
+        if "成功保存" not in text:
+            continue
+        try:
+            part = text.split("成功保存", 1)[1].strip()
+            return int(part.split("条", 1)[0].strip())
+        except (IndexError, ValueError):
+            continue
+    return 0

+ 27 - 0
agents/demand_video_expand_agent/tools/__init__.py

@@ -0,0 +1,27 @@
+"""
+demand_video_expand_agent 工具包
+"""
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import Any
+
+from agents.demand_video_expand_agent.tools.batch_save_demand_expansions import (
+    batch_save_demand_expansions,
+)
+from supply_agent.tools.registry import ToolRegistry
+
+ALL_TOOLS: list[Callable[..., Any]] = [
+    batch_save_demand_expansions,
+]
+
+__all__ = [
+    "ALL_TOOLS",
+    "batch_save_demand_expansions",
+    "register_all_tools",
+]
+
+
+def register_all_tools(registry: ToolRegistry) -> ToolRegistry:
+    """将 demand_video_expand_agent 包内的所有工具注册到 ToolRegistry。"""
+    return registry.from_decorated(*ALL_TOOLS)

+ 222 - 0
agents/demand_video_expand_agent/tools/batch_save_demand_expansions.py

@@ -0,0 +1,222 @@
+"""
+批量保存视频点位拓展判断结果。
+"""
+from __future__ import annotations
+
+import logging
+import re
+import uuid
+from typing import Any
+
+from supply_agent.tools import tool
+from supply_infra.db.models.multi_demand_video_point import POINT_TYPES
+from supply_infra.db.repositories.demand_video_expansion_repo import (
+    DemandVideoExpansionRepository,
+)
+from supply_infra.db.repositories.multi_demand_video_point_repo import (
+    MultiDemandVideoPointRepository,
+)
+from supply_infra.db.session import get_session
+
+logger = logging.getLogger(__name__)
+
+_VALID_GRADES = frozenset({"S", "A"})
+
+
+def _optional_str(value: Any) -> str | None:
+    if value is None:
+        return None
+    text = str(value).strip()
+    return text or None
+
+
+def _normalize_items(
+    items: list[dict[str, Any]],
+    *,
+    default_run_id: str,
+    biz_dt: str,
+    source_demand_grade_id: int,
+    source_demand_name: str,
+    source_grade: str,
+) -> tuple[list[dict[str, Any]], list[str]]:
+    rows: list[dict[str, Any]] = []
+    errors: list[str] = []
+    seen_keys: set[tuple[str, str]] = set()
+
+    for idx, item in enumerate(items):
+        if not isinstance(item, dict):
+            errors.append(f"第 {idx} 项不是对象")
+            continue
+
+        expanded_text = _optional_str(item.get("expanded_text"))
+        point_type = _optional_str(item.get("point_type"))
+        video_id = _optional_str(item.get("video_id"))
+        reason = _optional_str(item.get("reason"))
+
+        if not expanded_text:
+            errors.append(f"第 {idx} 项缺少 expanded_text")
+            continue
+        if not point_type or point_type not in POINT_TYPES:
+            allowed = " / ".join(POINT_TYPES)
+            errors.append(f"第 {idx} 项 point_type 无效(只能是 {allowed}): {point_type!r}")
+            continue
+        if not video_id:
+            errors.append(f"第 {idx} 项缺少 video_id")
+            continue
+        if not reason:
+            errors.append(f"第 {idx} 项缺少 reason")
+            continue
+
+        dedupe_key = (expanded_text, video_id)
+        if dedupe_key in seen_keys:
+            errors.append(f"第 {idx} 项在本次请求中重复: {expanded_text!r} / {video_id!r}")
+            continue
+        seen_keys.add(dedupe_key)
+
+        rows.append(
+            {
+                "biz_dt": biz_dt,
+                "run_id": _optional_str(item.get("run_id")) or default_run_id,
+                "source_demand_grade_id": source_demand_grade_id,
+                "source_demand_name": source_demand_name,
+                "source_grade": source_grade,
+                "expanded_text": expanded_text,
+                "point_type": point_type,
+                "point_desc": _optional_str(item.get("point_desc")),
+                "video_id": video_id,
+                "reason": reason,
+                "is_delete": 0,
+            }
+        )
+
+    return rows, errors
+
+
+def _normalize_demand_name(name: str) -> str:
+    text = re.sub(r"\s+", "", name.strip().lower())
+    return re.sub(r"[,,。..!!??;;::""''\"'、/\\|·—_()()【】\[\]《》<>]", "", text)
+
+
+def _fill_missing_point_descs(rows: list[dict[str, Any]], session) -> None:
+    """point_desc 为空时,从 multi_demand_video_point 按 video_id/point_type/point_data 补全。
+
+    查不到或源表 point_desc 也为空时,保持原值(仍为 None)。
+    """
+    missing = [row for row in rows if not row.get("point_desc")]
+    if not missing:
+        return
+
+    video_ids = sorted({str(row["video_id"]) for row in missing})
+    points_by_vid = MultiDemandVideoPointRepository(session).list_by_video_ids(video_ids)
+
+    lookup: dict[tuple[str, str, str], str] = {}
+    for vid, points in points_by_vid.items():
+        for point in points:
+            point_data = _optional_str(point.get("point_data"))
+            point_type = _optional_str(point.get("point_type"))
+            point_desc = _optional_str(point.get("point_desc"))
+            if not point_data or not point_type or not point_desc:
+                continue
+            lookup[(vid, point_type, point_data)] = point_desc
+
+    for row in missing:
+        key = (str(row["video_id"]), str(row["point_type"]), str(row["expanded_text"]))
+        desc = lookup.get(key)
+        if desc:
+            row["point_desc"] = desc
+        # 未命中或源表无描述:不改动,保留空值
+
+
+@tool
+def batch_save_demand_expansions(
+    items: list[dict[str, Any]],
+    biz_dt: str,
+    source_demand_grade_id: int,
+    source_demand_name: str,
+    source_grade: str,
+    run_id: str | None = None,
+) -> str:
+    """
+    保存视频点位拓展判断结果到 demand_video_expansion 表。
+
+    Args:
+        items: 拓展候选列表。每项必填:
+            - expanded_text: 拓展需求文本(来自 point_data)
+            - point_type: inspiration / purpose / key
+            - video_id: 来源视频 id
+            - reason: 为何与原需求相近、可作为拓展
+          选填:
+            - point_desc: 点位描述快照
+        biz_dt: 业务日 YYYYMMDD。
+        source_demand_grade_id: 来源 demand_grade.id(由用户消息提供,原样传入)。
+        source_demand_name: 来源需求名。
+        source_grade: 来源等级 S 或 A。
+        run_id: 任务批次 id;省略则自动生成。
+
+    Returns:
+        保存结果摘要。若无合适拓展,传 items=[] 即可。
+    """
+    biz_dt_text = _optional_str(biz_dt)
+    if not biz_dt_text or len(biz_dt_text) != 8 or not biz_dt_text.isdigit():
+        return f"biz_dt 格式无效,应为 YYYYMMDD: {biz_dt!r}"
+
+    demand_name = _optional_str(source_demand_name)
+    if not demand_name:
+        return "source_demand_name 不能为空"
+
+    grade = _optional_str(source_grade)
+    if grade not in _VALID_GRADES:
+        return f"source_grade 无效,只能是 S 或 A: {source_grade!r}"
+
+    try:
+        grade_id = int(source_demand_grade_id)
+    except (TypeError, ValueError):
+        return f"source_demand_grade_id 无效: {source_demand_grade_id!r}"
+
+    if not items:
+        return "无拓展候选,跳过落库"
+
+    default_run_id = _optional_str(run_id) or uuid.uuid4().hex
+    rows, errors = _normalize_items(
+        items,
+        default_run_id=default_run_id,
+        biz_dt=biz_dt_text,
+        source_demand_grade_id=grade_id,
+        source_demand_name=demand_name,
+        source_grade=grade,
+    )
+
+    normalized_demand = _normalize_demand_name(demand_name)
+    filtered_rows: list[dict[str, Any]] = []
+    skipped_same: list[str] = []
+    for row in rows:
+        if _normalize_demand_name(row["expanded_text"]) == normalized_demand:
+            skipped_same.append(row["expanded_text"])
+            continue
+        filtered_rows.append(row)
+
+    if not filtered_rows:
+        detail = ";".join(errors) if errors else "无有效数据"
+        same_note = ""
+        if skipped_same:
+            same_note = f";剔除与原需求相同 {len(skipped_same)} 条"
+        return f"没有可保存的数据: {detail}{same_note}"
+
+    try:
+        with get_session() as session:
+            _fill_missing_point_descs(filtered_rows, session)
+            saved = DemandVideoExpansionRepository(session).bulk_upsert(filtered_rows)
+
+        parts = [f"成功保存 {saved} 条拓展需求", f"biz_dt={biz_dt_text}", f"run_id={default_run_id}"]
+        if skipped_same:
+            parts.append(f"剔除与原需求相同 {len(skipped_same)} 条")
+        if errors:
+            parts.append(f"校验失败 {len(errors)} 条: " + ";".join(errors[:10]))
+
+        message = "。".join(parts)
+        logger.info("batch_save_demand_expansions completed: %s", message)
+        return message
+
+    except Exception as e:
+        logger.error("batch_save_demand_expansions failed: %s", e, exc_info=True)
+        return f"保存拓展需求失败: {e}"

+ 1 - 3
agents/generate_demand_agent/tools/batch_save_generated_demands.py

@@ -258,9 +258,7 @@ def batch_save_generated_demands(
                     row["dim_count"] = int(getattr(stats, f"{dim}_count", 0) or 0)
                 if row.get("dim_avg") is None:
                     avg = getattr(stats, f"{dim}_avg", None)
-                    row["dim_avg"] = (
-                        Decimal(str(avg)) if avg is not None else Decimal("0")
-                    )
+                    row["dim_avg"] = Decimal(str(avg)) if avg is not None else None
 
             inserted = GeneratedDemandRepository(session).bulk_insert(valid_rows)
             run_ids = sorted({str(r["run_id"]) for r in valid_rows})

+ 2 - 2
agents/generate_demand_agent/tools/dim_constants.py

@@ -93,14 +93,14 @@ def dim_score(
     weight: CategoryTreeWeight | None,
     dim: str,
 ) -> tuple[float | None, int]:
-    """返回 (avg, count);count>0 即有维度数据(avg 为 0 也算有)。"""
+    """返回 (avg, count);count>0 即有维度数据。"""
     if weight is None:
         return None, 0
     count = int(getattr(weight, f"{dim}_count", 0) or 0)
     if count <= 0:
         return None, 0
     avg = getattr(weight, f"{dim}_avg", None)
-    return (float(avg) if avg is not None else 0.0), count
+    return (float(avg) if avg is not None else None), count
 
 
 def build_children_map(

+ 1 - 1
agents/generate_demand_agent/tools/query_demand_words_by_category.py

@@ -58,7 +58,7 @@ def _word_dim_score(
     if count <= 0:
         return None, 0
     avg = getattr(stats, f"{dim}_avg", None)
-    return (float(avg) if avg is not None else 0.0), count
+    return (float(avg) if avg is not None else None), count
 
 
 @tool

+ 42 - 1
api/app.py

@@ -1,6 +1,7 @@
 """FastAPI application — category tree API on port 8080."""
 from __future__ import annotations
 
+from contextlib import asynccontextmanager
 from pathlib import Path
 
 from fastapi import FastAPI, HTTPException, Query
@@ -9,10 +10,23 @@ from fastapi.staticfiles import StaticFiles
 
 from api.services.category_tree import build_category_tree
 from api.services.demand_belong_category import list_demand_belong_categories
+from api.services.demand_grade import list_demand_grades
+from api.services.demand_grade_videos import list_videos_for_demand_grade
 from api.services.demand_videos import list_videos_for_demand_belong
 from api.services.oss_logs import list_demand_belong_oss_logs
+from supply_infra.db import init_db
+from supply_infra.scheduler.app import get_scheduler_status, start_scheduler, stop_scheduler
 
-app = FastAPI(title="SupplyAgent API", version="0.1.0")
+
+@asynccontextmanager
+async def lifespan(_app: FastAPI):
+    init_db()
+    start_scheduler()
+    yield
+    stop_scheduler()
+
+
+app = FastAPI(title="SupplyAgent API", version="0.1.0", lifespan=lifespan)
 
 app.add_middleware(
     CORSMiddleware,
@@ -33,6 +47,12 @@ def health() -> dict[str, str]:
     return {"status": "ok"}
 
 
+@app.get("/api/scheduler/status")
+def scheduler_status() -> dict:
+    """Return scheduler enabled/running state and next run times."""
+    return get_scheduler_status()
+
+
 @app.get("/api/category-tree")
 def category_tree(
     biz_dt: str | None = Query(
@@ -60,6 +80,27 @@ def demand_belong_videos(belong_id: int) -> dict:
     return result
 
 
+@app.get("/api/demand-grade")
+def demand_grade(
+    biz_dt: str | None = Query(
+        default=None,
+        description="业务日 YYYYMMDD;省略则取 demand_grade 最新一日",
+    ),
+) -> dict:
+    """Return demand_grade rows (one per category_id) for the given/latest biz_dt."""
+    items = list_demand_grades(biz_dt=biz_dt)
+    return {"items": items}
+
+
+@app.get("/api/demand-grade/{demand_grade_id}/videos")
+def demand_grade_videos(demand_grade_id: int) -> dict:
+    """Return expansion videos/points for a demand_grade row (by its biz_dt)."""
+    result = list_videos_for_demand_grade(demand_grade_id)
+    if result is None:
+        raise HTTPException(status_code=404, detail="demand_grade not found")
+    return result
+
+
 @app.get("/api/demand-belong-oss-logs")
 def demand_belong_oss_logs() -> dict:
     """Return demand_belong_category_agent oss_logs ordered by create_time desc."""

+ 6 - 0
api/run.py

@@ -1,10 +1,16 @@
 """Start the SupplyAgent API on port 8080."""
 from __future__ import annotations
 
+import logging
+
 import uvicorn
 
 
 def main() -> None:
+    logging.basicConfig(
+        level=logging.INFO,
+        format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
+    )
     uvicorn.run(
         "api.app:app",
         host="0.0.0.0",

+ 9 - 12
api/services/category_tree.py

@@ -1,7 +1,6 @@
 """Build nested category tree from global_tree_category (+ optional weights)."""
 from __future__ import annotations
 
-from decimal import Decimal
 from typing import Any
 
 from supply_infra.db.models.category_tree_weight import CategoryTreeWeight
@@ -48,19 +47,17 @@ def _normalize_parent_id(parent_id: int | None) -> int | None:
     return parent_id
 
 
-def _to_float(value: Decimal | float | int | None) -> float:
-    if value is None:
-        return 0.0
-    return float(value)
-
-
-def _weights_payload(row: CategoryTreeWeight | None) -> dict[str, float]:
+def _weights_payload(row: CategoryTreeWeight | None) -> dict[str, float | None]:
     if row is None:
-        return {dim: 0.0 for dim in (TOTAL_SCORE_KEY, *DIM_KEYS)}
+        return {dim: None for dim in (TOTAL_SCORE_KEY, *DIM_KEYS)}
     return {
-        TOTAL_SCORE_KEY: _to_float(row.total_score),
+        TOTAL_SCORE_KEY: float(row.total_score) if row.total_score is not None else None,
         **{
-            dim: _to_float(getattr(row, f"{dim}_avg", None))
+            dim: (
+                float(getattr(row, f"{dim}_avg"))
+                if getattr(row, f"{dim}_avg", None) is not None
+                else None
+            )
             for dim in DIM_KEYS
         },
     }
@@ -121,7 +118,7 @@ def build_category_tree(biz_dt: str | None = None) -> dict[str, Any]:
     """Load active categories with per-dim avg; nest as roots.
 
     Returns ``{biz_dt, dims, nodes}``. When no weight rows exist, ``biz_dt`` is
-    ``None`` and each node still carries zeroed ``weights`` / ``counts``.
+    ``None`` and each node still carries null ``weights`` / zero ``counts``.
     """
     with get_session() as session:
         categories = GlobalTreeCategoryRepository(session).list_active_categories()

+ 25 - 0
api/services/demand_grade.py

@@ -0,0 +1,25 @@
+"""Load demand_grade rows (expanded one row per category) for the category tree UI."""
+from __future__ import annotations
+
+from typing import Any
+
+from supply_infra.db.repositories.demand_grade_category_rel_repo import (
+    DemandGradeCategoryRelRepository,
+)
+from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
+from supply_infra.db.session import get_session
+
+
+def list_demand_grades(biz_dt: str | None = None) -> list[dict[str, Any]]:
+    """
+    Return demand_grade rows joined with demand_grade_category_rel, one row per
+    (demand_grade, category_id) pair — mirrors the shape the frontend already
+    groups by category_id.
+
+    biz_dt defaults to the latest business date present in demand_grade.
+    """
+    with get_session() as session:
+        resolved_biz_dt = biz_dt or DemandGradeRepository(session).get_latest_biz_dt()
+        if not resolved_biz_dt:
+            return []
+        return DemandGradeCategoryRelRepository(session).list_items_with_category(resolved_biz_dt)

+ 96 - 0
api/services/demand_grade_videos.py

@@ -0,0 +1,96 @@
+"""Resolve demand_video_expansion → videos + expansion points for the web UI."""
+from __future__ import annotations
+
+import json
+from typing import Any
+
+from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
+from supply_infra.db.repositories.demand_video_expansion_repo import (
+    DemandVideoExpansionRepository,
+    DemandVideoExpansionRunRepository,
+)
+from supply_infra.db.repositories.multi_demand_video_detail_repo import (
+    MultiDemandVideoDetailRepository,
+)
+from supply_infra.db.session import get_session
+
+
+def _parse_json_list(raw: str | None) -> list[Any]:
+    if not raw:
+        return []
+    try:
+        parsed = json.loads(raw)
+    except json.JSONDecodeError:
+        return []
+    return parsed if isinstance(parsed, list) else []
+
+
+def list_videos_for_demand_grade(demand_grade_id: int) -> dict[str, Any] | None:
+    """
+    按 demand_grade.id 返回该需求在对应 biz_dt 下的拓展视频与选题结果。
+
+    数据来源:
+    - demand_video_expansion_run:判断是否已完成拓展(含零结果)
+    - demand_video_expansion:真实视频实例与拓展点位
+    """
+    with get_session() as session:
+        grade = DemandGradeRepository(session).get_by_id(demand_grade_id)
+        if grade is None:
+            return None
+
+        biz_dt = str(grade.biz_dt)
+        run = DemandVideoExpansionRunRepository(session).get_by_demand_grade(
+            biz_dt, demand_grade_id
+        )
+        expansions = DemandVideoExpansionRepository(session).list_by_demand_grade(
+            biz_dt, demand_grade_id
+        )
+
+        vids: list[str] = []
+        seen_vids: set[str] = set()
+        points_by_vid: dict[str, list[dict[str, Any]]] = {}
+        for row in expansions:
+            vid = str(row.video_id).strip()
+            if not vid:
+                continue
+            if vid not in seen_vids:
+                seen_vids.add(vid)
+                vids.append(vid)
+            points_by_vid.setdefault(vid, []).append(
+                {
+                    "expanded_text": row.expanded_text,
+                    "point_type": row.point_type,
+                    "point_desc": row.point_desc,
+                    "reason": row.reason,
+                }
+            )
+
+        details = MultiDemandVideoDetailRepository(session).list_by_vids(vids)
+        videos: list[dict[str, Any]] = []
+        for vid in vids:
+            row = details.get(vid)
+            videos.append(
+                {
+                    "vid": vid,
+                    "title": row.title if row else None,
+                    "expansion_points": points_by_vid.get(vid, []),
+                }
+            )
+
+        expansion_status: str | None = None
+        expansion_saved_count = 0
+        if run is not None:
+            expansion_status = str(run.status)
+            expansion_saved_count = int(run.saved_count or 0)
+
+        return {
+            "demand_grade_id": grade.id,
+            "demand_name": grade.demand_name,
+            "biz_dt": biz_dt,
+            "category_ids": [int(c) for c in _parse_json_list(grade.category_ids)],
+            "grade": grade.grade,
+            "strategies": _parse_json_list(grade.strategies),
+            "expansion_status": expansion_status,
+            "expansion_saved_count": expansion_saved_count,
+            "videos": videos,
+        }

+ 11 - 4
api/services/demand_videos.py

@@ -10,6 +10,9 @@ from supply_infra.db.repositories.demand_belong_category_repo import (
 from supply_infra.db.repositories.multi_demand_video_detail_repo import (
     MultiDemandVideoDetailRepository,
 )
+from supply_infra.db.repositories.multi_demand_video_point_repo import (
+    MultiDemandVideoPointRepository,
+)
 from supply_infra.db.session import get_session
 
 
@@ -38,19 +41,23 @@ def list_videos_for_demand_belong(belong_id: int) -> dict[str, Any] | None:
 
         vids = _parse_video_ids(belong.video_list)
         details = MultiDemandVideoDetailRepository(session).list_by_vids(vids)
+        points_by_vid = MultiDemandVideoPointRepository(session).json_fields_by_video_ids(
+            vids
+        )
 
         videos: list[dict[str, Any]] = []
         for vid in vids:
             row = details.get(vid)
+            point_fields = points_by_vid.get(vid, {})
             videos.append(
                 {
                     "vid": vid,
                     "title": row.title if row else None,
-                    "inspiration_points_json": (
-                        row.inspiration_points_json if row else None
+                    "inspiration_points_json": point_fields.get(
+                        "inspiration_points_json"
                     ),
-                    "purpose_points_json": row.purpose_points_json if row else None,
-                    "key_points_json": row.key_points_json if row else None,
+                    "purpose_points_json": point_fields.get("purpose_points_json"),
+                    "key_points_json": point_fields.get("key_points_json"),
                 }
             )
 

+ 36 - 0
jobs/backfill_demand_video_expansion_point_desc.py

@@ -0,0 +1,36 @@
+#!/usr/bin/env python3
+"""手动补全 demand_video_expansion 缺失的 point_desc。
+
+用法:
+    python jobs/backfill_demand_video_expansion_point_desc.py
+    python jobs/backfill_demand_video_expansion_point_desc.py 20260721
+    python jobs/backfill_demand_video_expansion_point_desc.py 20260721 --dry-run
+"""
+from __future__ import annotations
+
+import logging
+import sys
+from pathlib import Path
+
+_ROOT = Path(__file__).resolve().parents[1]
+if str(_ROOT) not in sys.path:
+    sys.path.insert(0, str(_ROOT))
+
+from scripts.backfill_demand_video_expansion_point_desc import backfill_missing_point_descs
+
+logging.basicConfig(
+    level=logging.INFO,
+    format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
+)
+
+
+def main(biz_dt: str | None = None, *, dry_run: bool = False) -> dict:
+    result = backfill_missing_point_descs(biz_dt, dry_run=dry_run)
+    print(result)
+    return result
+
+
+if __name__ == "__main__":
+    args = sys.argv[1:]
+    biz_dt_arg = args[0] if args and not args[0].startswith("-") else None
+    main(biz_dt_arg, dry_run="--dry-run" in args)

+ 1 - 1
jobs/backfill_multi_demand_video_list.py

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

+ 117 - 0
jobs/backfill_multi_demand_video_points_table.py

@@ -0,0 +1,117 @@
+#!/usr/bin/env python3
+"""将 multi_demand_video_detail 三个 JSON 点位列迁移到 multi_demand_video_point 表。
+
+用法:
+    python jobs/backfill_multi_demand_video_points_table.py
+    python jobs/backfill_multi_demand_video_points_table.py 500   # 每批 500 条视频
+"""
+
+from __future__ import annotations
+
+import logging
+import sys
+
+from sqlalchemy import select
+
+from supply_infra.db.models.multi_demand_video_detail import MultiDemandVideoDetail
+from supply_infra.db.repositories.multi_demand_video_point_repo import (
+    MultiDemandVideoPointRepository,
+)
+from supply_infra.db.session import get_session
+from supply_infra.video_points import points_from_json_fields
+
+logging.basicConfig(
+    level=logging.INFO,
+    format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
+)
+logger = logging.getLogger(__name__)
+
+_DEFAULT_BATCH_SIZE = 200
+
+
+def backfill_video_points_table(batch_size: int = _DEFAULT_BATCH_SIZE) -> dict:
+    """从 detail 表 JSON 列回填点位表,跳过已迁移的 video_id。"""
+    chunk = max(1, int(batch_size))
+    total_rows = 0
+    total_points = 0
+    batches = 0
+    offset = 0
+
+    while True:
+        with get_session() as session:
+            stmt = (
+                select(MultiDemandVideoDetail)
+                .where(
+                    MultiDemandVideoDetail.inspiration_points_json.is_not(None)
+                    | MultiDemandVideoDetail.purpose_points_json.is_not(None)
+                    | MultiDemandVideoDetail.key_points_json.is_not(None)
+                )
+                .order_by(MultiDemandVideoDetail.id)
+                .offset(offset)
+                .limit(chunk)
+            )
+            rows = list(session.scalars(stmt).all())
+            if not rows:
+                break
+
+            vids = [str(row.vid) for row in rows if row.vid]
+            existing = MultiDemandVideoPointRepository(session).list_video_ids_with_points(
+                vids
+            )
+
+            points_by_vid: dict[str, list] = {}
+            for row in rows:
+                vid = str(row.vid).strip() if row.vid else ""
+                if not vid or vid in existing:
+                    continue
+                point_rows = points_from_json_fields(
+                    vid,
+                    inspiration_points_json=row.inspiration_points_json,
+                    purpose_points_json=row.purpose_points_json,
+                    key_points_json=row.key_points_json,
+                )
+                if point_rows:
+                    points_by_vid[vid] = point_rows
+
+            inserted = 0
+            if points_by_vid:
+                inserted = MultiDemandVideoPointRepository(session).replace_for_video_ids(
+                    points_by_vid
+                )
+
+        batch_count = len(rows)
+        migrated = len(points_by_vid)
+        total_rows += batch_count
+        total_points += inserted
+        batches += 1
+        offset += batch_count
+        logger.info(
+            "Batch %d: scanned=%d migrated_videos=%d inserted_points=%d offset=%d",
+            batches,
+            batch_count,
+            migrated,
+            inserted,
+            offset,
+        )
+
+        if batch_count < chunk:
+            break
+
+    result = {
+        "batches": batches,
+        "scanned_rows": total_rows,
+        "inserted_points": total_points,
+    }
+    logger.info("Backfill multi_demand_video_point completed: %s", result)
+    return result
+
+
+def main(batch_arg: str | None = None) -> dict:
+    batch_size = int(batch_arg) if batch_arg else _DEFAULT_BATCH_SIZE
+    result = backfill_video_points_table(batch_size=batch_size)
+    print(result)
+    return result
+
+
+if __name__ == "__main__":
+    main(sys.argv[1] if len(sys.argv) > 1 else None)

+ 48 - 0
jobs/expand_demand_from_video_points.py

@@ -0,0 +1,48 @@
+#!/usr/bin/env python3
+"""手动执行 S/A 需求视频点位拓展任务。
+
+用法:
+    python jobs/expand_demand_from_video_points.py                  # 最新/当天 biz_dt
+    python jobs/expand_demand_from_video_points.py 20260721          # 指定业务日
+    python jobs/expand_demand_from_video_points.py 20260721 5        # 指定业务日 + 5 并发
+    python jobs/expand_demand_from_video_points.py 20260721 --force   # 忽略已完成记录重跑
+"""
+from __future__ import annotations
+
+import logging
+import sys
+
+from supply_infra.scheduler.jobs.expand_demand_from_video_points import (
+    expand_demand_from_video_points,
+)
+
+logging.basicConfig(
+    level=logging.INFO,
+    format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
+)
+
+
+def main(
+    biz_dt: str | None = None,
+    workers_arg: str | None = None,
+    *,
+    skip_finished: bool = True,
+) -> dict:
+    workers = int(workers_arg) if workers_arg else 5
+    result = expand_demand_from_video_points(
+        biz_dt,
+        skip_finished=skip_finished,
+        workers=workers,
+    )
+    print(result)
+    return result
+
+
+if __name__ == "__main__":
+    args = sys.argv[1:]
+    biz_dt_arg = args[0] if args and not args[0].startswith("-") else None
+    workers_arg = None
+    if biz_dt_arg and len(args) > 1 and not args[1].startswith("-"):
+        workers_arg = args[1]
+    skip_finished = "--force" not in args
+    main(biz_dt_arg, workers_arg, skip_finished=skip_finished)

+ 43 - 0
jobs/grade_demand_pool.py

@@ -0,0 +1,43 @@
+#!/usr/bin/env python3
+"""手动执行树热度驱动的需求分级。
+
+统筹规划 Agent 先落库当天全量节点组计划,再由多个 worker 领取任务并调用分级 Agent。
+
+用法:
+    python jobs/grade_demand_pool.py                  # 默认业务日、5 个 worker
+    python jobs/grade_demand_pool.py 20260716         # 指定业务日
+    python jobs/grade_demand_pool.py 20260716 5       # 指定业务日 + 5 个并发 worker
+"""
+from __future__ import annotations
+
+import logging
+import sys
+
+from supply_infra.scheduler.jobs.grade_demand_pool import grade_demand_pool
+
+logging.basicConfig(
+    level=logging.INFO,
+    format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
+)
+
+
+def main(
+    biz_dt: str | None = None,
+    workers_arg: str | None = None,
+) -> dict:
+    workers = int(workers_arg) if workers_arg else 5
+
+    result = grade_demand_pool(
+        biz_dt,
+        workers=workers,
+        with_orchestrate=True,
+    )
+    print(result)
+    return result
+
+
+if __name__ == "__main__":
+    main(
+        sys.argv[1] if len(sys.argv) > 1 else None,
+        sys.argv[2] if len(sys.argv) > 2 else None,
+    )

+ 10 - 2
jobs/init_db.py

@@ -1,8 +1,16 @@
 #!/usr/bin/env python3
 """CLI entry point to initialize database tables."""
 
+from supply_infra.config import get_infra_settings
 from supply_infra.db import init_db
 
 if __name__ == "__main__":
-    init_db()
-    print("Database tables created.")
+    settings = get_infra_settings()
+    result = init_db()
+    print(
+        f"Database ready: {settings.mysql_host}:{settings.mysql_port}/{settings.mysql_database}"
+    )
+    if result["created"]:
+        print("Created tables:", ", ".join(result["created"]))
+    else:
+        print("No new tables created (all ORM tables already exist).")

+ 45 - 0
jobs/retry_failed_grade_plan_items.py

@@ -0,0 +1,45 @@
+#!/usr/bin/env python3
+"""手动重试 demand_grade_plan_group_item 中失败的分级任务。
+
+用法:
+    python jobs/retry_failed_grade_plan_items.py
+    python jobs/retry_failed_grade_plan_items.py 20260721
+    python jobs/retry_failed_grade_plan_items.py 20260721 5
+    python jobs/retry_failed_grade_plan_items.py 20260721 5 --dry-run
+"""
+from __future__ import annotations
+
+import logging
+import sys
+
+from supply_infra.scheduler.jobs.grade_demand_pool import retry_failed_plan_group_items
+
+logging.basicConfig(
+    level=logging.INFO,
+    format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
+)
+
+
+def main(
+    biz_dt: str | None = None,
+    workers_arg: str | None = None,
+    *,
+    dry_run: bool = False,
+) -> dict:
+    workers = int(workers_arg) if workers_arg else 5
+    result = retry_failed_plan_group_items(
+        biz_dt,
+        workers=workers,
+        dry_run=dry_run,
+    )
+    print(result)
+    return result
+
+
+if __name__ == "__main__":
+    args = sys.argv[1:]
+    biz_dt_arg = args[0] if args and not args[0].startswith("-") else None
+    workers_arg = None
+    if biz_dt_arg and len(args) > 1 and not args[1].startswith("-"):
+        workers_arg = args[1]
+    main(biz_dt_arg, workers_arg, dry_run="--dry-run" in args)

+ 1 - 1
jobs/run_scheduler.py

@@ -3,7 +3,7 @@
 
 import logging
 
-from supply_infra.scheduler import run_scheduler
+from supply_infra.scheduler.app import run_scheduler
 
 logging.basicConfig(
     level=logging.INFO,

+ 22 - 0
jobs/run_supply_pipeline.py

@@ -0,0 +1,22 @@
+#!/usr/bin/env python3
+"""手动执行供给数据流水线(全局树 → 需求池 → 分级 → 视频点位拓展)。"""
+
+import logging
+import sys
+
+from supply_infra.scheduler.jobs.run_supply_pipeline import run_supply_pipeline
+
+logging.basicConfig(
+    level=logging.INFO,
+    format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
+)
+
+
+def main() -> None:
+    biz_dt = sys.argv[1] if len(sys.argv) > 1 else None
+    result = run_supply_pipeline(biz_dt)
+    print(result)
+
+
+if __name__ == "__main__":
+    main()

+ 162 - 0
scripts/backfill_demand_video_expansion_point_desc.py

@@ -0,0 +1,162 @@
+#!/usr/bin/env python3
+"""补全 demand_video_expansion 表中缺失的 point_desc。
+
+从 multi_demand_video_point 按 (video_id, point_type, expanded_text=point_data) 匹配;
+查不到则保持空值。
+
+用法:
+  .venv/bin/python scripts/backfill_demand_video_expansion_point_desc.py
+  .venv/bin/python scripts/backfill_demand_video_expansion_point_desc.py --biz-dt 20260721
+  .venv/bin/python scripts/backfill_demand_video_expansion_point_desc.py --biz-dt 20260721 --dry-run
+"""
+from __future__ import annotations
+
+import argparse
+import json
+import logging
+import sys
+from pathlib import Path
+from typing import Any
+
+from sqlalchemy import or_, select, update
+
+_ROOT = Path(__file__).resolve().parents[1]
+if str(_ROOT) not in sys.path:
+    sys.path.insert(0, str(_ROOT))
+
+from agents.demand_video_expand_agent.tools.batch_save_demand_expansions import (
+    _fill_missing_point_descs,
+)
+from supply_infra.db.models.demand_video_expansion import DemandVideoExpansion
+from supply_infra.db.session import get_session
+
+logger = logging.getLogger(__name__)
+_BATCH_SIZE = 500
+
+
+def _list_rows_missing_point_desc(biz_dt: str | None) -> list[dict[str, Any]]:
+    stmt = select(DemandVideoExpansion).where(
+        DemandVideoExpansion.is_delete == 0,
+        or_(
+            DemandVideoExpansion.point_desc.is_(None),
+            DemandVideoExpansion.point_desc == "",
+        ),
+    )
+    if biz_dt:
+        stmt = stmt.where(DemandVideoExpansion.biz_dt == biz_dt)
+    stmt = stmt.order_by(DemandVideoExpansion.id)
+
+    with get_session() as session:
+        rows = session.scalars(stmt).all()
+        return [
+            {
+                "id": int(row.id),
+                "biz_dt": str(row.biz_dt),
+                "video_id": str(row.video_id),
+                "point_type": str(row.point_type),
+                "expanded_text": str(row.expanded_text),
+                "point_desc": row.point_desc,
+            }
+            for row in rows
+        ]
+
+
+def backfill_missing_point_descs(
+    biz_dt: str | None = None,
+    *,
+    dry_run: bool = False,
+) -> dict[str, Any]:
+    rows = _list_rows_missing_point_desc(biz_dt)
+    result: dict[str, Any] = {
+        "biz_dt": biz_dt,
+        "dry_run": dry_run,
+        "missing_total": len(rows),
+        "filled": 0,
+        "still_empty": 0,
+        "updated": 0,
+        "samples": [],
+    }
+    if not rows:
+        return result
+
+    with get_session() as session:
+        _fill_missing_point_descs(rows, session)
+
+    to_update: list[dict[str, Any]] = []
+    for row in rows:
+        if row.get("point_desc"):
+            to_update.append(row)
+            result["filled"] += 1
+            if len(result["samples"]) < 10:
+                result["samples"].append(
+                    {
+                        "id": row["id"],
+                        "video_id": row["video_id"],
+                        "point_type": row["point_type"],
+                        "expanded_text": row["expanded_text"],
+                        "point_desc": row["point_desc"][:80]
+                        if len(str(row["point_desc"])) > 80
+                        else row["point_desc"],
+                    }
+                )
+        else:
+            result["still_empty"] += 1
+
+    if dry_run or not to_update:
+        result["updated"] = 0
+        return result
+
+    with get_session() as session:
+        for i in range(0, len(to_update), _BATCH_SIZE):
+            batch = to_update[i : i + _BATCH_SIZE]
+            for row in batch:
+                session.execute(
+                    update(DemandVideoExpansion)
+                    .where(DemandVideoExpansion.id == int(row["id"]))
+                    .values(point_desc=row["point_desc"])
+                )
+            result["updated"] += len(batch)
+
+    return result
+
+
+def main(argv: list[str] | None = None) -> int:
+    parser = argparse.ArgumentParser(
+        description="补全 demand_video_expansion 缺失的 point_desc",
+    )
+    parser.add_argument("--biz-dt", default=None, help="业务日期 YYYYMMDD,默认全表")
+    parser.add_argument("--dry-run", action="store_true", help="仅统计,不写库")
+    parser.add_argument("--json", action="store_true", help="以 JSON 输出结果")
+    args = parser.parse_args(argv)
+
+    logging.basicConfig(
+        level=logging.INFO,
+        format="%(asctime)s %(levelname)s %(name)s: %(message)s",
+    )
+
+    result = backfill_missing_point_descs(args.biz_dt, dry_run=bool(args.dry_run))
+
+    if args.json:
+        print(json.dumps(result, ensure_ascii=False, indent=2, default=str))
+    else:
+        print("\n=== point_desc 补全 ===")
+        print(f"biz_dt={result.get('biz_dt') or '全部'}")
+        print(f"dry_run={result.get('dry_run')}")
+        print(f"缺失记录={result.get('missing_total')}")
+        print(f"可补全={result.get('filled')}")
+        print(f"仍为空={result.get('still_empty')}")
+        print(f"已更新={result.get('updated')}")
+        if result.get("samples"):
+            print("\n示例:")
+            for item in result["samples"]:
+                print(
+                    f"  id={item['id']} video={item['video_id']} "
+                    f"type={item['point_type']} text={item['expanded_text']!r} "
+                    f"desc={item['point_desc']!r}"
+                )
+
+    return 0
+
+
+if __name__ == "__main__":
+    raise SystemExit(main())

+ 100 - 0
scripts/retry_failed_grade_plan_items.py

@@ -0,0 +1,100 @@
+#!/usr/bin/env python3
+"""重试 demand_grade_plan_group_item 中 status=failed 的分级任务。
+
+用法:
+  .venv/bin/python scripts/retry_failed_grade_plan_items.py
+  .venv/bin/python scripts/retry_failed_grade_plan_items.py --biz-dt 20260721
+  .venv/bin/python scripts/retry_failed_grade_plan_items.py --biz-dt 20260721 --workers 5
+  .venv/bin/python scripts/retry_failed_grade_plan_items.py --biz-dt 20260721 --dry-run
+  .venv/bin/python scripts/retry_failed_grade_plan_items.py --group-id 12 --group-id 15
+"""
+from __future__ import annotations
+
+import argparse
+import json
+import logging
+import sys
+from pathlib import Path
+
+_ROOT = Path(__file__).resolve().parents[1]
+if str(_ROOT) not in sys.path:
+    sys.path.insert(0, str(_ROOT))
+
+from supply_infra.scheduler.jobs.grade_demand_pool import retry_failed_plan_group_items
+from supply_infra.scheduler.plan_group_batch import MAX_DEMANDS_PER_BATCH
+
+logger = logging.getLogger(__name__)
+
+
+def main(argv: list[str] | None = None) -> int:
+    parser = argparse.ArgumentParser(
+        description="重试 demand_grade_plan_group_item 中失败的分级任务",
+    )
+    parser.add_argument("--biz-dt", default=None, help="业务日期 YYYYMMDD,默认当天")
+    parser.add_argument("--workers", type=int, default=5, help="并发执行的 plan_group 数")
+    parser.add_argument(
+        "--max-demands-per-batch",
+        type=int,
+        default=MAX_DEMANDS_PER_BATCH,
+        help=f"每个 Agent 子批次最多处理的需求条数,默认 {MAX_DEMANDS_PER_BATCH}",
+    )
+    parser.add_argument(
+        "--group-id",
+        type=int,
+        action="append",
+        dest="group_ids",
+        help="仅重试指定 group_id,可重复传入",
+    )
+    parser.add_argument(
+        "--dry-run",
+        action="store_true",
+        help="仅列出将要重试的 failed 记录,不实际执行",
+    )
+    parser.add_argument("--json", action="store_true", help="以 JSON 打印结果")
+    args = parser.parse_args(argv)
+
+    logging.basicConfig(
+        level=logging.INFO,
+        format="%(asctime)s %(levelname)s %(name)s: %(message)s",
+    )
+
+    result = retry_failed_plan_group_items(
+        args.biz_dt,
+        workers=max(1, int(args.workers)),
+        max_demands_per_batch=max(1, min(int(args.max_demands_per_batch), MAX_DEMANDS_PER_BATCH)),
+        group_ids=args.group_ids,
+        dry_run=bool(args.dry_run),
+    )
+
+    if args.json:
+        print(json.dumps(result, ensure_ascii=False, indent=2, default=str))
+    else:
+        reset = result.get("reset") or {}
+        print("\n=== 失败任务重试 ===")
+        print(f"biz_dt={result.get('biz_dt')}")
+        print(f"dry_run={result.get('dry_run')}")
+        print(f"failed_items={result.get('failed_items', 0)}")
+        print(f"reset_items={reset.get('reset_items', 0)}")
+        print(f"reset_groups={reset.get('reset_groups', 0)}")
+        if reset.get("group_ids"):
+            print(f"group_ids={reset.get('group_ids')}")
+        if not result.get("dry_run"):
+            print(f"graded: {result.get('graded_before')} -> {result.get('graded_after')}")
+            print(f"remaining_failed={result.get('remaining_failed', 0)}")
+            print(f"group_status={result.get('group_status')}")
+            print(f"success={result.get('success')}")
+        elif reset.get("items"):
+            print("\n待重试明细:")
+            for item in reset["items"][:20]:
+                print(
+                    f"  item_id={item['item_id']} group_id={item['group_id']} "
+                    f"demand={item['demand_name']!r}"
+                )
+            if len(reset["items"]) > 20:
+                print(f"  ... 另有 {len(reset['items']) - 20} 条")
+
+    return 0 if result.get("success") else 1
+
+
+if __name__ == "__main__":
+    raise SystemExit(main())

+ 85 - 0
scripts/run_grade_plan_groups.py

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

+ 1 - 1
supply_agent/agent/core.py

@@ -58,7 +58,7 @@ class Agent:
         skills: SkillRegistry | None = None,
         max_iterations: int | None = None,
         temperature: float | None = None,
-        reasoning_effort: str | None = "medium",
+        reasoning_effort: str | None = None,
         logger: AgentLogger | None = None,
     ) -> None:
         self.name = name

+ 1 - 1
supply_agent/config.py

@@ -22,7 +22,7 @@ class Settings(BaseSettings):
     # OpenRouter
     openrouter_api_key: str = Field(..., alias="OPENROUTER_API_KEY")
     openrouter_model: str = Field(
-        default="anthropic/claude-sonnet-5",
+        default="google/gemini-2.5-flash",
         alias="OPENROUTER_MODEL",
     )
     openrouter_base_url: str = Field(

+ 1 - 1
supply_agent/llm/client.py

@@ -20,7 +20,7 @@ class LLMClient:
         settings: Settings,
         logger: AgentLogger | None = None,
         *,
-        reasoning_effort: str | None = "medium",
+        reasoning_effort: str | None = None,
     ) -> None:
         self.settings = settings
         self.model = settings.openrouter_model

+ 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()
+    }

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

@@ -3,23 +3,45 @@
 from supply_infra.db.models.category_tree_weight import CategoryTreeWeight
 from supply_infra.db.models.demand_belong_category import DemandBelongCategory
 from supply_infra.db.models.demand_belong_pool_rel import DemandBelongPoolRel
+from supply_infra.db.models.demand_grade import DemandGrade
+from supply_infra.db.models.demand_grade_category_rel import DemandGradeCategoryRel
+from supply_infra.db.models.demand_grade_plan import (
+    DemandGradePlan,
+    DemandGradePlanGroup,
+    DemandGradePlanGroupItem,
+)
 from supply_infra.db.models.demand_popularity_stats import DemandPopularityStats
+from supply_infra.db.models.demand_video_expansion import (
+    DemandVideoExpansion,
+    DemandVideoExpansionRun,
+)
 from supply_infra.db.models.generated_demand import GeneratedDemand
 from supply_infra.db.models.global_tree_category import GlobalTreeCategory
 from supply_infra.db.models.global_tree_element import GlobalTreeElement
 from supply_infra.db.models.multi_demand_pool_di import MultiDemandPoolDi
 from supply_infra.db.models.multi_demand_video_detail import MultiDemandVideoDetail
+from supply_infra.db.models.multi_demand_video_point import MultiDemandVideoPoint
 from supply_infra.db.models.oss_log import OssLog
+from supply_infra.db.models.scheduler_job_execution import SchedulerJobExecution
 
 __all__ = [
     "CategoryTreeWeight",
     "DemandBelongCategory",
     "DemandBelongPoolRel",
+    "DemandGrade",
+    "DemandGradeCategoryRel",
+    "DemandGradePlan",
+    "DemandGradePlanGroup",
+    "DemandGradePlanGroupItem",
     "DemandPopularityStats",
+    "DemandVideoExpansion",
+    "DemandVideoExpansionRun",
     "GeneratedDemand",
     "GlobalTreeCategory",
     "GlobalTreeElement",
     "MultiDemandPoolDi",
     "MultiDemandVideoDetail",
+    "MultiDemandVideoPoint",
     "OssLog",
+    "SchedulerJobExecution",
 ]

+ 22 - 22
supply_infra/db/models/category_tree_weight.py

@@ -26,58 +26,58 @@ class CategoryTreeWeight(Base):
         Integer, nullable=False, default=0, comment="挂载需求词数量"
     )
 
-    ext_pop_avg: Mapped[Decimal] = mapped_column(
-        Numeric(16, 8), nullable=False, default=Decimal("0"), comment="外部热度-加权平均分"
+    ext_pop_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="外部热度-加权平均分"
     )
     ext_pop_count: Mapped[int] = mapped_column(
         Integer, nullable=False, default=0, comment="外部热度-样本数"
     )
 
-    plat_sust_pop_avg: Mapped[Decimal] = mapped_column(
-        Numeric(16, 8), nullable=False, default=Decimal("0"), comment="平台持续热度-加权平均分"
+    plat_sust_pop_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="平台持续热度-加权平均分"
     )
     plat_sust_pop_count: Mapped[int] = mapped_column(
         Integer, nullable=False, default=0, comment="平台持续热度-样本数"
     )
 
-    plat_ly_pop_avg: Mapped[Decimal] = mapped_column(
-        Numeric(16, 8), nullable=False, default=Decimal("0"), comment="平台去年同期-加权平均分"
+    plat_ly_pop_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="平台去年同期-加权平均分"
     )
     plat_ly_pop_count: Mapped[int] = mapped_column(
         Integer, nullable=False, default=0, comment="平台去年同期-样本数"
     )
 
-    recent_pop_avg: Mapped[Decimal] = mapped_column(
-        Numeric(16, 8), nullable=False, default=Decimal("0"), comment="近期热度-加权平均分"
+    recent_pop_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="近期热度-加权平均分"
     )
     recent_pop_count: Mapped[int] = mapped_column(
         Integer, nullable=False, default=0, comment="近期热度-样本数"
     )
 
-    ext_pop_score: Mapped[Decimal] = mapped_column(
-        Numeric(16, 8), nullable=False, default=Decimal("0"), comment="外部热度-排名归一化分"
+    ext_pop_score: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="外部热度-排名归一化分"
     )
-    plat_sust_pop_score: Mapped[Decimal] = mapped_column(
-        Numeric(16, 8), nullable=False, default=Decimal("0"), comment="平台持续热度-排名归一化分"
+    plat_sust_pop_score: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="平台持续热度-排名归一化分"
     )
-    plat_ly_pop_score: Mapped[Decimal] = mapped_column(
-        Numeric(16, 8), nullable=False, default=Decimal("0"), comment="平台去年同期-排名归一化分"
+    plat_ly_pop_score: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="平台去年同期-排名归一化分"
     )
-    recent_pop_score: Mapped[Decimal] = mapped_column(
-        Numeric(16, 8), nullable=False, default=Decimal("0"), comment="近期热度-排名归一化分"
+    recent_pop_score: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="近期热度-排名归一化分"
     )
-    total_score: Mapped[Decimal] = mapped_column(
-        Numeric(16, 8), nullable=False, default=Decimal("0"), comment="四维排名分之和"
+    total_score: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="四维排名分之和"
     )
 
-    real_rov_7d_avg: Mapped[Decimal] = mapped_column(
-        Numeric(16, 8), nullable=False, default=Decimal("0"), comment="近7日真实ROV-加权平均分"
+    real_rov_7d_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="近7日真实ROV-加权平均分"
     )
     real_rov_7d_count: Mapped[int] = mapped_column(
         Integer, nullable=False, default=0, comment="近7日真实ROV-样本数"
     )
-    real_vov_7d_avg: Mapped[Decimal] = mapped_column(
-        Numeric(16, 8), nullable=False, default=Decimal("0"), comment="近7日真实VOV-加权平均分"
+    real_vov_7d_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="近7日真实VOV-加权平均分"
     )
     real_vov_7d_count: Mapped[int] = mapped_column(
         Integer, nullable=False, default=0, comment="近7日真实VOV-样本数"

+ 70 - 0
supply_infra/db/models/demand_grade.py

@@ -0,0 +1,70 @@
+from __future__ import annotations
+
+from datetime import datetime
+from decimal import Decimal
+
+from sqlalchemy import BigInteger, Integer, Numeric, String, Text, UniqueConstraint, func
+from sqlalchemy.orm import Mapped, mapped_column
+
+from supply_infra.db.base import Base
+
+
+class DemandGrade(Base):
+    """需求分级结果 — 对 multi_demand_pool_di 中的需求按先验/后验热度评级。"""
+
+    __tablename__ = "demand_grade"
+    __table_args__ = (UniqueConstraint("biz_dt", "demand_name", name="uk_demand_grade"),)
+
+    id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
+    biz_dt: Mapped[str] = mapped_column(String(32), nullable=False, comment="业务日期YYYYMMDD")
+    demand_name: Mapped[str] = mapped_column(
+        String(256), nullable=False, comment="需求名称(去重后的代表词)"
+    )
+    category_ids: Mapped[str | None] = mapped_column(
+        Text, nullable=True, comment="归属的全局树节点id列表JSON数组(展示快照,权威关系见demand_grade_category_rel)"
+    )
+    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="需求自身来源归一分0-100(来源内排名后对已有来源取均值)",
+    )
+    prior_total_score: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="落库时的先验 total_score 快照"
+    )
+    posterior_rov_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(16, 8), nullable=True, comment="落库时的后验 real_rov_7d_avg 快照"
+    )
+    posterior_rov_count: Mapped[int] = mapped_column(
+        Integer, nullable=False, default=0, comment="落库时的后验样本数快照"
+    )
+    has_posterior: Mapped[int] = mapped_column(
+        Integer, nullable=False, default=0, comment="是否有后验验证数据 0-无 1-有"
+    )
+    related_pool_ids: Mapped[str] = mapped_column(
+        Text,
+        nullable=False,
+        comment="关联的原始 multi_demand_pool_di.id 列表JSON数组(必填,用于回溯原始需求)",
+    )
+    video_list: Mapped[str | None] = mapped_column(
+        Text,
+        nullable=True,
+        comment="关联视频列表JSON(最多10个,由related_pool_ids对应原始行的video_list合并去重得到)",
+    )
+    strategies: Mapped[str | None] = mapped_column(
+        Text,
+        nullable=True,
+        comment="来源策略列表JSON数组(由related_pool_ids对应原始行的strategy去重得到)",
+    )
+    reason: Mapped[str] = mapped_column(Text, nullable=False, comment="判断依据")
+    create_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        comment="创建时间",
+    )
+    update_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        onupdate=func.now(),
+        comment="更新时间",
+    )

+ 42 - 0
supply_infra/db/models/demand_grade_category_rel.py

@@ -0,0 +1,42 @@
+from __future__ import annotations
+
+from datetime import datetime
+
+from sqlalchemy import BigInteger, Index, UniqueConstraint, func
+from sqlalchemy.orm import Mapped, mapped_column
+
+from supply_infra.db.base import Base
+
+
+class DemandGradeCategoryRel(Base):
+    """需求分级结果与全局树节点的归属关系,支撑按分类高效查询已分级需求。"""
+
+    __tablename__ = "demand_grade_category_rel"
+    __table_args__ = (
+        UniqueConstraint(
+            "demand_grade_id",
+            "category_id",
+            name="uk_demand_grade_category",
+        ),
+        Index("idx_demand_grade_id", "demand_grade_id"),
+        Index("idx_category_id", "category_id"),
+    )
+
+    id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
+    demand_grade_id: Mapped[int] = mapped_column(
+        BigInteger, nullable=False, comment="demand_grade.id"
+    )
+    category_id: Mapped[int] = mapped_column(
+        BigInteger, nullable=False, comment="global_tree_category.id"
+    )
+    create_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        comment="创建时间",
+    )
+    update_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        onupdate=func.now(),
+        comment="更新时间",
+    )

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

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

+ 12 - 12
supply_infra/db/models/demand_popularity_stats.py

@@ -25,38 +25,38 @@ class DemandPopularityStats(Base):
         String(128), nullable=True, comment="需求词名称"
     )
     biz_dt: Mapped[str] = mapped_column(String(32), nullable=False, comment="日期")
-    ext_pop_avg: Mapped[Decimal] = mapped_column(
-        Numeric(12, 2), nullable=False, default=Decimal("0.00"), comment="外部热度-平均值"
+    ext_pop_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(12, 2), nullable=True, comment="外部热度-平均值"
     )
     ext_pop_count: Mapped[int] = mapped_column(
         Integer, nullable=False, default=0, comment="外部热度-出现数量"
     )
-    plat_sust_pop_avg: Mapped[Decimal] = mapped_column(
-        Numeric(12, 2), nullable=False, default=Decimal("0.00"), comment="平台持续热度-平均值"
+    plat_sust_pop_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(12, 2), nullable=True, comment="平台持续热度-平均值"
     )
     plat_sust_pop_count: Mapped[int] = mapped_column(
         Integer, nullable=False, default=0, comment="平台持续热度-出现数量"
     )
-    plat_ly_pop_avg: Mapped[Decimal] = mapped_column(
-        Numeric(12, 2), nullable=False, default=Decimal("0.00"), comment="平台去年同期热度-平均值"
+    plat_ly_pop_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(12, 2), nullable=True, comment="平台去年同期热度-平均值"
     )
     plat_ly_pop_count: Mapped[int] = mapped_column(
         Integer, nullable=False, default=0, comment="平台去年同期热度-出现数量"
     )
-    recent_pop_avg: Mapped[Decimal] = mapped_column(
-        Numeric(12, 2), nullable=False, default=Decimal("0.00"), comment="近期热度-平均值"
+    recent_pop_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(12, 2), nullable=True, comment="近期热度-平均值"
     )
     recent_pop_count: Mapped[int] = mapped_column(
         Integer, nullable=False, default=0, comment="近期热度-出现数量"
     )
-    real_rov_7d_avg: Mapped[Decimal] = mapped_column(
-        Numeric(12, 4), nullable=False, default=Decimal("0.0000"), comment="近7日真实ROV-平均值"
+    real_rov_7d_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(12, 4), nullable=True, comment="近7日真实ROV-平均值"
     )
     real_rov_7d_count: Mapped[int] = mapped_column(
         Integer, nullable=False, default=0, comment="近7日真实ROV-出现数量"
     )
-    real_vov_7d_avg: Mapped[Decimal] = mapped_column(
-        Numeric(12, 4), nullable=False, default=Decimal("0.0000"), comment="近7日真实VOV-平均值"
+    real_vov_7d_avg: Mapped[Decimal | None] = mapped_column(
+        Numeric(12, 4), nullable=True, comment="近7日真实VOV-平均值"
     )
     real_vov_7d_count: Mapped[int] = mapped_column(
         Integer, nullable=False, default=0, comment="近7日真实VOV-出现数量"

+ 100 - 0
supply_infra/db/models/demand_video_expansion.py

@@ -0,0 +1,100 @@
+from __future__ import annotations
+
+from datetime import datetime
+
+from sqlalchemy import BigInteger, Index, Integer, String, Text, UniqueConstraint, func
+from sqlalchemy.orm import Mapped, mapped_column
+
+from supply_infra.db.base import Base
+
+
+class DemandVideoExpansion(Base):
+    """S/A 需求关联视频点位拓展结果。"""
+
+    __tablename__ = "demand_video_expansion"
+    __table_args__ = (
+        UniqueConstraint(
+            "biz_dt",
+            "source_demand_grade_id",
+            "expanded_text",
+            "video_id",
+            name="uk_demand_video_expansion",
+        ),
+        Index("idx_dve_biz_dt", "biz_dt"),
+        Index("idx_dve_source_grade", "source_demand_grade_id"),
+    )
+
+    id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
+    biz_dt: Mapped[str] = mapped_column(String(32), nullable=False, comment="业务日 YYYYMMDD")
+    run_id: Mapped[str] = mapped_column(String(64), nullable=False, comment="任务批次 run_id")
+    source_demand_grade_id: Mapped[int] = mapped_column(
+        BigInteger, nullable=False, comment="来源 demand_grade.id"
+    )
+    source_demand_name: Mapped[str] = mapped_column(
+        String(256), nullable=False, comment="来源需求名"
+    )
+    source_grade: Mapped[str] = mapped_column(String(4), nullable=False, comment="来源等级 S/A")
+    expanded_text: Mapped[str] = mapped_column(
+        String(512), nullable=False, comment="拓展需求文本(来自 point_data)"
+    )
+    point_type: Mapped[str] = mapped_column(
+        String(32), nullable=False, comment="点类型:inspiration / purpose / key"
+    )
+    point_desc: Mapped[str | None] = mapped_column(Text, nullable=True, comment="点位描述快照")
+    video_id: Mapped[str] = mapped_column(String(64), nullable=False, comment="来源视频 id")
+    reason: Mapped[str] = mapped_column(Text, nullable=False, comment="相近判断依据")
+    is_delete: Mapped[int] = mapped_column(
+        Integer, default=0, nullable=False, comment="是否删除 0-正常 1-删除"
+    )
+    create_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        comment="创建时间",
+    )
+    update_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        onupdate=func.now(),
+        comment="更新时间",
+    )
+
+
+class DemandVideoExpansionRun(Base):
+    """记录每个需求是否已完成拓展判断(含零结果),用于幂等跳过。"""
+
+    __tablename__ = "demand_video_expansion_run"
+    __table_args__ = (
+        UniqueConstraint(
+            "biz_dt",
+            "source_demand_grade_id",
+            name="uk_demand_video_expansion_run",
+        ),
+        Index("idx_dver_biz_dt", "biz_dt"),
+    )
+
+    id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
+    biz_dt: Mapped[str] = mapped_column(String(32), nullable=False, comment="业务日 YYYYMMDD")
+    run_id: Mapped[str] = mapped_column(String(64), nullable=False, comment="任务批次 run_id")
+    source_demand_grade_id: Mapped[int] = mapped_column(
+        BigInteger, nullable=False, comment="来源 demand_grade.id"
+    )
+    saved_count: Mapped[int] = mapped_column(
+        Integer, nullable=False, default=0, comment="落库拓展条数"
+    )
+    status: Mapped[str] = mapped_column(
+        String(16), nullable=False, default="finished", comment="finished / failed"
+    )
+    error_message: Mapped[str | None] = mapped_column(
+        Text, nullable=True, comment="失败原因"
+    )
+    create_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        comment="创建时间",
+    )
+    update_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        onupdate=func.now(),
+        comment="更新时间",
+    )

+ 48 - 0
supply_infra/db/models/multi_demand_video_point.py

@@ -0,0 +1,48 @@
+from __future__ import annotations
+
+from datetime import datetime
+
+from sqlalchemy import BigInteger, Index, String, Text, func
+from sqlalchemy.orm import Mapped, mapped_column
+
+from supply_infra.db.base import Base
+
+POINT_TYPE_INSPIRATION = "inspiration"
+POINT_TYPE_PURPOSE = "purpose"
+POINT_TYPE_KEY = "key"
+
+POINT_TYPES = (POINT_TYPE_INSPIRATION, POINT_TYPE_PURPOSE, POINT_TYPE_KEY)
+
+
+class MultiDemandVideoPoint(Base):
+    """需求池视频点位表 — 灵感点/目的点/关键点按行存储。"""
+
+    __tablename__ = "multi_demand_video_point"
+    __table_args__ = (
+        Index("idx_multi_demand_video_point_video_type", "video_id", "point_type"),
+    )
+
+    id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
+    video_id: Mapped[str] = mapped_column(String(64), nullable=False, comment="视频id")
+    point_type: Mapped[str] = mapped_column(
+        String(32),
+        nullable=False,
+        comment="点类型:inspiration / purpose / key",
+    )
+    point_data: Mapped[str | None] = mapped_column(
+        Text, nullable=True, comment="点"
+    )
+    point_desc: Mapped[str | None] = mapped_column(
+        Text, nullable=True, comment="点描述"
+    )
+    create_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        comment="创建时间",
+    )
+    update_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        onupdate=func.now(),
+        comment="更新时间",
+    )

+ 51 - 0
supply_infra/db/models/scheduler_job_execution.py

@@ -0,0 +1,51 @@
+from __future__ import annotations
+
+from datetime import datetime
+
+from sqlalchemy import BigInteger, Numeric, String, Text, func
+from sqlalchemy.orm import Mapped, mapped_column
+
+from supply_infra.db.base import Base
+
+
+class SchedulerJobExecution(Base):
+    """定时任务执行记录(开始/结束各记一条,同一 run_id 关联)。"""
+
+    __tablename__ = "scheduler_job_execution"
+
+    id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
+    run_id: Mapped[str] = mapped_column(
+        String(36),
+        nullable=False,
+        index=True,
+        comment="同一次执行的唯一标识",
+    )
+    job_name: Mapped[str] = mapped_column(String(128), nullable=False, comment="定时任务名称")
+    job_id: Mapped[str | None] = mapped_column(String(64), nullable=True, comment="调度器 job id")
+    status: Mapped[str] = mapped_column(
+        String(32),
+        nullable=False,
+        comment="状态: started/finished/failed/skipped",
+    )
+    event_time: Mapped[datetime] = mapped_column(nullable=False, comment="事件发生时间")
+    biz_dt: Mapped[str | None] = mapped_column(String(8), nullable=True, comment="业务日 YYYYMMDD")
+    started_at: Mapped[datetime | None] = mapped_column(nullable=True, comment="任务开始时间")
+    finished_at: Mapped[datetime | None] = mapped_column(nullable=True, comment="任务结束时间")
+    duration_seconds: Mapped[float | None] = mapped_column(
+        Numeric(10, 2),
+        nullable=True,
+        comment="执行耗时(秒)",
+    )
+    error_message: Mapped[str | None] = mapped_column(Text, nullable=True, comment="失败原因")
+    detail: Mapped[str | None] = mapped_column(Text, nullable=True, comment="执行结果摘要 JSON")
+    create_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        comment="创建时间",
+    )
+    update_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        onupdate=func.now(),
+        comment="更新时间",
+    )

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

@@ -10,9 +10,18 @@ from supply_infra.db.repositories.demand_belong_category_repo import (
 from supply_infra.db.repositories.demand_belong_pool_rel_repo import (
     DemandBelongPoolRelRepository,
 )
+from supply_infra.db.repositories.demand_grade_category_rel_repo import (
+    DemandGradeCategoryRelRepository,
+)
+from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
+from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
 from supply_infra.db.repositories.demand_popularity_stats_repo import (
     DemandPopularityStatsRepository,
 )
+from supply_infra.db.repositories.demand_video_expansion_repo import (
+    DemandVideoExpansionRepository,
+    DemandVideoExpansionRunRepository,
+)
 from supply_infra.db.repositories.generated_demand_repo import GeneratedDemandRepository
 from supply_infra.db.repositories.global_tree_category_repo import GlobalTreeCategoryRepository
 from supply_infra.db.repositories.global_tree_element_repo import GlobalTreeElementRepository
@@ -20,18 +29,31 @@ from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPo
 from supply_infra.db.repositories.multi_demand_video_detail_repo import (
     MultiDemandVideoDetailRepository,
 )
+from supply_infra.db.repositories.multi_demand_video_point_repo import (
+    MultiDemandVideoPointRepository,
+)
 from supply_infra.db.repositories.oss_log_repo import OssLogRepository
+from supply_infra.db.repositories.scheduler_job_execution_repo import (
+    SchedulerJobExecutionRepository,
+)
 
 __all__ = [
     "BaseRepository",
     "CategoryTreeWeightRepository",
     "DemandBelongCategoryRepository",
     "DemandBelongPoolRelRepository",
+    "DemandGradeCategoryRelRepository",
+    "DemandGradeRepository",
+    "DemandGradePlanRepository",
     "DemandPopularityStatsRepository",
+    "DemandVideoExpansionRepository",
+    "DemandVideoExpansionRunRepository",
     "GeneratedDemandRepository",
     "GlobalTreeCategoryRepository",
     "GlobalTreeElementRepository",
     "MultiDemandPoolDiRepository",
     "MultiDemandVideoDetailRepository",
+    "MultiDemandVideoPointRepository",
     "OssLogRepository",
+    "SchedulerJobExecutionRepository",
 ]

+ 32 - 0
supply_infra/db/repositories/category_tree_weight_repo.py

@@ -52,6 +52,38 @@ class CategoryTreeWeightRepository(BaseRepository[CategoryTreeWeight]):
         stmt = select(CategoryTreeWeight).where(CategoryTreeWeight.biz_dt == biz_dt)
         return list(self.session.scalars(stmt).all())
 
+    def get_by_category_ids(
+        self,
+        category_ids: list[int],
+        biz_dt: str | None = None,
+    ) -> list[CategoryTreeWeight]:
+        """按 category_id 列表查询权重行;未指定 biz_dt 时取每个节点各自的最新业务日。"""
+        if not category_ids:
+            return []
+
+        if biz_dt:
+            stmt = select(CategoryTreeWeight).where(
+                CategoryTreeWeight.category_id.in_(category_ids),
+                CategoryTreeWeight.biz_dt == biz_dt,
+            )
+            return list(self.session.scalars(stmt).all())
+
+        latest_dt_subq = (
+            select(
+                CategoryTreeWeight.category_id,
+                func.max(CategoryTreeWeight.biz_dt).label("max_biz_dt"),
+            )
+            .where(CategoryTreeWeight.category_id.in_(category_ids))
+            .group_by(CategoryTreeWeight.category_id)
+            .subquery()
+        )
+        stmt = select(CategoryTreeWeight).join(
+            latest_dt_subq,
+            (CategoryTreeWeight.category_id == latest_dt_subq.c.category_id)
+            & (CategoryTreeWeight.biz_dt == latest_dt_subq.c.max_biz_dt),
+        )
+        return list(self.session.scalars(stmt).all())
+
     def get_by_category_biz_dt(
         self, category_id: int, biz_dt: str
     ) -> CategoryTreeWeight | None:

+ 22 - 0
supply_infra/db/repositories/demand_belong_category_repo.py

@@ -16,6 +16,28 @@ class DemandBelongCategoryRepository(BaseRepository[DemandBelongCategory]):
 
     model = DemandBelongCategory
 
+    def get_by_ids(self, ids: Iterable[int]) -> list[DemandBelongCategory]:
+        """按 id 批量查询,跳过已软删除记录。"""
+        id_list = [int(i) for i in ids]
+        if not id_list:
+            return []
+        stmt = select(DemandBelongCategory).where(
+            DemandBelongCategory.id.in_(id_list),
+            DemandBelongCategory.is_delete == 0,
+        )
+        return list(self.session.scalars(stmt).all())
+
+    def search_by_name_like(self, keyword: str) -> list[DemandBelongCategory]:
+        """按名称模糊匹配(双向包含关系不在此处理,仅 LIKE %keyword%),跳过已软删除记录。"""
+        keyword = (keyword or "").strip()
+        if not keyword:
+            return []
+        stmt = select(DemandBelongCategory).where(
+            DemandBelongCategory.name.like(f"%{keyword}%"),
+            DemandBelongCategory.is_delete == 0,
+        )
+        return list(self.session.scalars(stmt).all())
+
     def get_existing_names(self, names: Iterable[str]) -> set[str]:
         """返回 names 中已存在于表内的名称(含软删除)。"""
         name_list = [n for n in names if n]

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

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

+ 62 - 0
supply_infra/db/repositories/demand_grade_category_rel_repo.py

@@ -0,0 +1,62 @@
+from __future__ import annotations
+
+from sqlalchemy import delete, select
+from sqlalchemy.dialects.mysql import insert
+
+from supply_infra.db.models.demand_grade import DemandGrade
+from supply_infra.db.models.demand_grade_category_rel import DemandGradeCategoryRel
+from supply_infra.db.repositories.base import BaseRepository
+
+
+class DemandGradeCategoryRelRepository(BaseRepository[DemandGradeCategoryRel]):
+    """demand_grade ↔ global_tree_category 归属关系 — 支撑按分类高效查询已分级需求。"""
+
+    model = DemandGradeCategoryRel
+
+    def replace_for_demand_grade(self, demand_grade_id: int, category_ids: list[int]) -> None:
+        """覆盖某个 demand_grade 的分类映射:先删除旧的,再写入新的(空列表则只删)。"""
+        self.session.execute(
+            delete(DemandGradeCategoryRel).where(
+                DemandGradeCategoryRel.demand_grade_id == demand_grade_id
+            )
+        )
+        unique_ids = sorted({int(c) for c in category_ids if c is not None})
+        if not unique_ids:
+            return
+        rows = [
+            {"demand_grade_id": demand_grade_id, "category_id": category_id}
+            for category_id in unique_ids
+        ]
+        stmt = insert(DemandGradeCategoryRel).values(rows).prefix_with("IGNORE")
+        self.session.execute(stmt)
+
+    def list_items_with_category(self, biz_dt: str) -> list[dict]:
+        """按 biz_dt 展开为一行一个 category_id,供前端按分类展示需求列表使用。"""
+        stmt = (
+            select(
+                DemandGrade.id,
+                DemandGrade.demand_name,
+                DemandGradeCategoryRel.category_id,
+                DemandGrade.grade,
+                DemandGrade.score,
+                DemandGrade.reason,
+                DemandGrade.strategies,
+                DemandGrade.biz_dt,
+            )
+            .join(DemandGradeCategoryRel, DemandGradeCategoryRel.demand_grade_id == DemandGrade.id)
+            .where(DemandGrade.biz_dt == biz_dt)
+            .order_by(DemandGrade.grade, DemandGrade.demand_name)
+        )
+        return [
+            {
+                "id": int(row.id),
+                "demand_name": row.demand_name,
+                "category_id": int(row.category_id),
+                "grade": row.grade,
+                "score": float(row.score) if row.score is not None else None,
+                "reason": row.reason,
+                "strategies": row.strategies,
+                "biz_dt": row.biz_dt,
+            }
+            for row in self.session.execute(stmt).all()
+        ]

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

@@ -0,0 +1,346 @@
+from __future__ import annotations
+
+import json
+import uuid
+from datetime import datetime
+from typing import Any
+
+from sqlalchemy import func, select, update
+
+from supply_infra.db.models.demand_grade_plan import (
+    DemandGradePlan,
+    DemandGradePlanGroup,
+    DemandGradePlanGroupItem,
+)
+from supply_infra.db.repositories.base import BaseRepository
+from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
+
+
+class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
+    model = DemandGradePlan
+
+    def get_latest_plan(self, biz_dt: str) -> DemandGradePlan | None:
+        return self.session.scalar(
+            select(DemandGradePlan)
+            .where(DemandGradePlan.biz_dt == biz_dt)
+            .order_by(DemandGradePlan.create_time.desc())
+            .limit(1)
+        )
+
+    def list_groups_by_biz_dt(self, biz_dt: str) -> list[DemandGradePlanGroup]:
+        return list(self.session.scalars(
+            select(DemandGradePlanGroup)
+            .where(DemandGradePlanGroup.biz_dt == biz_dt)
+            .order_by(DemandGradePlanGroup.id)
+        ).all())
+
+    def get_assigned_category_ids(self, biz_dt: str) -> set[int]:
+        assigned: set[int] = set()
+        for group in self.list_groups_by_biz_dt(biz_dt):
+            for category_id in json.loads(group.category_ids):
+                try:
+                    assigned.add(int(category_id))
+                except (TypeError, ValueError):
+                    continue
+        return assigned
+
+    def get_execution_snapshot(self, biz_dt: str) -> dict[str, Any]:
+        """在 session 内汇总当天计划分组执行情况,避免 ORM 脱离会话。"""
+        groups = self.list_groups_by_biz_dt(biz_dt)
+        group_status = self.summarize(biz_dt)
+        assigned_category_ids = sorted(self.get_assigned_category_ids(biz_dt))
+        claimable_groups = group_status.get("pending", 0)
+        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:
+        plan_id = str(uuid.uuid4())
+        groups = payload.get("groups") or []
+        coverage_complete = bool(payload.get("coverage_complete"))
+        # 分组明细已写入 demand_grade_plan_group,计划表仅保存可检索的紧凑摘要,避免 TEXT 溢出。
+        summary = {
+            "biz_dt": biz_dt,
+            "grouping_strategy": payload.get("grouping_strategy"),
+            "total_hanging_nodes": payload.get("total_hanging_nodes"),
+            "group_count": len(groups),
+            "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", []),
+        }
+        self.add(DemandGradePlan(
+            plan_id=plan_id, biz_dt=biz_dt, status="planned",
+            total_hanging_nodes=int(payload.get("total_hanging_nodes") or 0),
+            group_count=len(groups), coverage_complete=1 if coverage_complete else 0,
+            plan_json=json.dumps(summary, ensure_ascii=False),
+        ))
+        for group_no, group in enumerate(groups, start=1):
+            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=shared_traits,
+                status="pending",
+            ))
+            group_row = self.session.scalar(
+                select(DemandGradePlanGroup)
+                .where(DemandGradePlanGroup.plan_id == plan_id, DemandGradePlanGroup.group_no == group_no)
+                .limit(1)
+            )
+            if group_row is not None:
+                graded_names = DemandGradeRepository(self.session).get_existing_demand_names(biz_dt)
+                self.materialize_group_items(
+                    int(group_row.id),
+                    biz_dt,
+                    list(group.get("category_ids") or []),
+                    graded_names=graded_names,
+                )
+
+    def list_pending_group_ids(self, biz_dt: str) -> list[int]:
+        """返回当天待执行的 pending 任务 id。"""
+        stmt = (
+            select(DemandGradePlanGroup.id)
+            .where(
+                DemandGradePlanGroup.biz_dt == biz_dt,
+                DemandGradePlanGroup.status == "pending",
+            )
+            .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 == "pending",
+            )
+            .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,
+            "category_ids": json.loads(group.category_ids),
+        }
+
+    def finish_group(self, group_id: int, *, success: bool, error_message: str | None = None) -> None:
+        self.session.execute(
+            update(DemandGradePlanGroup)
+            .where(DemandGradePlanGroup.id == group_id)
+            .values(
+                status="finished" if success else "failed",
+                error_message=error_message,
+                finished_at=datetime.now(),
+            )
+        )
+
+    def summarize(self, biz_dt: str) -> dict[str, int]:
+        rows = self.session.execute(
+            select(DemandGradePlanGroup.status).where(DemandGradePlanGroup.biz_dt == biz_dt)
+        ).scalars().all()
+        return {status: sum(value == status for value in rows) for status in ("pending", "running", "finished", "failed")}
+
+    def group_item_count(self, group_id: int) -> int:
+        return int(self.session.scalar(
+            select(func.count())
+            .select_from(DemandGradePlanGroupItem)
+            .where(DemandGradePlanGroupItem.group_id == int(group_id))
+        ) or 0)
+
+    def materialize_group_items(
+        self,
+        group_id: int,
+        biz_dt: str,
+        category_ids: list[int],
+        *,
+        graded_names: set[str] | None = None,
+    ) -> int:
+        """将 category_ids 展开为组内需求明细;已存在明细时跳过。"""
+        if self.group_item_count(group_id) > 0:
+            return 0
+
+        graded = graded_names or set()
+        from supply_infra.scheduler.plan_group_batch import resolve_demands_for_category_ids
+
+        demands = resolve_demands_for_category_ids(biz_dt, category_ids)
+        created = 0
+        for sort_order, demand in enumerate(demands, start=1):
+            demand_name = str(demand["demand_name"])
+            status = "skipped" if demand_name in graded else "pending"
+            self.add(DemandGradePlanGroupItem(
+                group_id=int(group_id),
+                biz_dt=biz_dt,
+                pool_id=int(demand["pool_id"]),
+                demand_name=demand_name,
+                sort_order=sort_order,
+                status=status,
+            ))
+            created += 1
+        return created
+
+    def materialize_pending_groups(self, biz_dt: str, *, graded_names: set[str] | None = None) -> int:
+        """为当天尚未物化明细的 pending 计划组补写需求列表。"""
+        created = 0
+        for group in self.list_groups_by_biz_dt(biz_dt):
+            if group.status != "pending":
+                continue
+            if self.group_item_count(int(group.id)) > 0:
+                continue
+            category_ids = json.loads(group.category_ids)
+            created += self.materialize_group_items(
+                int(group.id),
+                biz_dt,
+                category_ids,
+                graded_names=graded_names,
+            )
+        return created
+
+    def list_pending_group_items(
+        self,
+        group_id: int,
+        *,
+        limit: int | None = None,
+    ) -> list[dict[str, Any]]:
+        stmt = (
+            select(DemandGradePlanGroupItem)
+            .where(
+                DemandGradePlanGroupItem.group_id == int(group_id),
+                DemandGradePlanGroupItem.status == "pending",
+            )
+            .order_by(DemandGradePlanGroupItem.sort_order, DemandGradePlanGroupItem.id)
+        )
+        if limit is not None:
+            stmt = stmt.limit(max(1, int(limit)))
+        rows = self.session.scalars(stmt).all()
+        return [
+            {
+                "item_id": int(row.id),
+                "pool_id": int(row.pool_id),
+                "demand_name": str(row.demand_name),
+            }
+            for row in rows
+        ]
+
+    def summarize_group_items(self, group_id: int) -> dict[str, int]:
+        rows = self.session.execute(
+            select(DemandGradePlanGroupItem.status)
+            .where(DemandGradePlanGroupItem.group_id == int(group_id))
+        ).scalars().all()
+        statuses = ("pending", "finished", "failed", "skipped")
+        return {status: sum(value == status for value in rows) for status in statuses}
+
+    def mark_group_items_status(
+        self,
+        item_ids: list[int],
+        *,
+        status: str,
+        error_message: str | None = None,
+    ) -> None:
+        if not item_ids:
+            return
+        values: dict[str, Any] = {"status": status, "error_message": error_message}
+        if status in {"finished", "failed", "skipped"}:
+            values["finished_at"] = datetime.now()
+        self.session.execute(
+            update(DemandGradePlanGroupItem)
+            .where(DemandGradePlanGroupItem.id.in_([int(item_id) for item_id in item_ids]))
+            .values(**values)
+        )
+
+    def list_failed_group_items(
+        self,
+        biz_dt: str | None = None,
+        *,
+        group_ids: list[int] | None = None,
+    ) -> list[dict[str, Any]]:
+        """返回 status=failed 的组内需求明细。"""
+        stmt = select(DemandGradePlanGroupItem).where(DemandGradePlanGroupItem.status == "failed")
+        if biz_dt:
+            stmt = stmt.where(DemandGradePlanGroupItem.biz_dt == biz_dt)
+        if group_ids:
+            stmt = stmt.where(
+                DemandGradePlanGroupItem.group_id.in_([int(group_id) for group_id in group_ids])
+            )
+        stmt = stmt.order_by(
+            DemandGradePlanGroupItem.biz_dt,
+            DemandGradePlanGroupItem.group_id,
+            DemandGradePlanGroupItem.sort_order,
+            DemandGradePlanGroupItem.id,
+        )
+        rows = self.session.scalars(stmt).all()
+        return [
+            {
+                "item_id": int(row.id),
+                "group_id": int(row.group_id),
+                "biz_dt": str(row.biz_dt),
+                "pool_id": int(row.pool_id),
+                "demand_name": str(row.demand_name),
+                "error_message": row.error_message,
+            }
+            for row in rows
+        ]
+
+    def reset_failed_items_to_pending(
+        self,
+        biz_dt: str | None = None,
+        *,
+        group_ids: list[int] | None = None,
+    ) -> dict[str, Any]:
+        """将 failed 明细重置为 pending,并将所属计划组重置为 pending 以便重新领取。"""
+        failed_rows = self.list_failed_group_items(biz_dt, group_ids=group_ids)
+        if not failed_rows:
+            return {"reset_items": 0, "reset_groups": 0, "group_ids": [], "items": []}
+
+        item_ids = [int(row["item_id"]) for row in failed_rows]
+        affected_group_ids = sorted({int(row["group_id"]) for row in failed_rows})
+
+        self.session.execute(
+            update(DemandGradePlanGroupItem)
+            .where(DemandGradePlanGroupItem.id.in_(item_ids))
+            .values(status="pending", error_message=None, finished_at=None)
+        )
+        self.session.execute(
+            update(DemandGradePlanGroup)
+            .where(DemandGradePlanGroup.id.in_(affected_group_ids))
+            .values(
+                status="pending",
+                error_message=None,
+                started_at=None,
+                finished_at=None,
+            )
+        )
+        return {
+            "reset_items": len(item_ids),
+            "reset_groups": len(affected_group_ids),
+            "group_ids": affected_group_ids,
+            "items": failed_rows,
+        }

+ 103 - 0
supply_infra/db/repositories/demand_grade_repo.py

@@ -0,0 +1,103 @@
+from __future__ import annotations
+
+from collections.abc import Iterable
+
+from sqlalchemy import func, select
+from sqlalchemy.dialects.mysql import insert
+
+from supply_infra.db.models.demand_grade import DemandGrade
+from supply_infra.db.repositories.base import BaseRepository
+
+_BATCH_SIZE = 500
+
+_UPSERT_COLUMNS = (
+    "category_ids",
+    "grade",
+    "score",
+    "prior_total_score",
+    "posterior_rov_avg",
+    "posterior_rov_count",
+    "has_posterior",
+    "related_pool_ids",
+    "video_list",
+    "strategies",
+    "reason",
+)
+
+
+class DemandGradeRepository(BaseRepository[DemandGrade]):
+    """需求分级结果表 — 按 (biz_dt, demand_name) 增量/更新写入。"""
+
+    model = DemandGrade
+
+    def get_existing_demand_names(self, biz_dt: str, names: Iterable[str] | None = None) -> set[str]:
+        """返回指定业务日已分级的需求名集合;传入 names 时只在其中查交集。"""
+        stmt = select(DemandGrade.demand_name).where(DemandGrade.biz_dt == biz_dt)
+        if names is not None:
+            name_list = [n for n in names if n]
+            if not name_list:
+                return set()
+            stmt = stmt.where(DemandGrade.demand_name.in_(name_list))
+        return {n for n in self.session.scalars(stmt).all() if n}
+
+    def count_by_biz_dt(self, biz_dt: str) -> int:
+        """统计指定业务日已分级的需求数。"""
+        stmt = select(func.count()).select_from(DemandGrade).where(DemandGrade.biz_dt == biz_dt)
+        return int(self.session.scalar(stmt) or 0)
+
+    def list_by_biz_dt(self, biz_dt: str) -> list[DemandGrade]:
+        """返回指定业务日的全部分级结果,按等级、需求名排序。"""
+        stmt = (
+            select(DemandGrade)
+            .where(DemandGrade.biz_dt == biz_dt)
+            .order_by(DemandGrade.grade, DemandGrade.demand_name)
+        )
+        return list(self.session.scalars(stmt).all())
+
+    def list_by_biz_dt_and_grades(
+        self,
+        biz_dt: str,
+        grades: Iterable[str] = ("S", "A"),
+    ) -> list[DemandGrade]:
+        """返回指定业务日、指定等级的分级结果,按等级、需求名排序。"""
+        grade_list = [g for g in grades if g]
+        if not grade_list:
+            return []
+        stmt = (
+            select(DemandGrade)
+            .where(DemandGrade.biz_dt == biz_dt, DemandGrade.grade.in_(grade_list))
+            .order_by(DemandGrade.grade, DemandGrade.demand_name)
+        )
+        return list(self.session.scalars(stmt).all())
+
+    def get_latest_biz_dt(self) -> str | None:
+        """返回 demand_grade 中最新的业务日期。"""
+        stmt = select(func.max(DemandGrade.biz_dt))
+        return self.session.scalar(stmt)
+
+    def get_ids_by_names(self, biz_dt: str, names: Iterable[str]) -> dict[str, int]:
+        """按 (biz_dt, demand_name) 反查 id,供 upsert 后写关联表使用。"""
+        name_list = [n for n in names if n]
+        if not name_list:
+            return {}
+        stmt = select(DemandGrade.demand_name, DemandGrade.id).where(
+            DemandGrade.biz_dt == biz_dt,
+            DemandGrade.demand_name.in_(name_list),
+        )
+        return {name: int(id_) for name, id_ in self.session.execute(stmt).all()}
+
+    def bulk_upsert(self, rows: list[dict]) -> int:
+        """按 (biz_dt, demand_name) 批量 upsert。"""
+        if not rows:
+            return 0
+
+        affected = 0
+        for i in range(0, len(rows), _BATCH_SIZE):
+            batch = rows[i : i + _BATCH_SIZE]
+            stmt = insert(DemandGrade).values(batch)
+            stmt = stmt.on_duplicate_key_update(
+                **{col: stmt.inserted[col] for col in _UPSERT_COLUMNS}
+            )
+            result = self.session.execute(stmt)
+            affected += result.rowcount or 0
+        return affected

+ 18 - 0
supply_infra/db/repositories/demand_popularity_stats_repo.py

@@ -47,6 +47,24 @@ class DemandPopularityStatsRepository(BaseRepository[DemandPopularityStats]):
         stmt = select(DemandPopularityStats).where(DemandPopularityStats.biz_dt == biz_dt)
         return list(self.session.scalars(stmt).all())
 
+    def search_by_word_name(
+        self,
+        keyword: str,
+        biz_dt: str | None = None,
+    ) -> list[DemandPopularityStats]:
+        """按 demand_word_name 精确+模糊搜索;未指定 biz_dt 时不限日期,按 biz_dt 降序返回。"""
+        keyword = (keyword or "").strip()
+        if not keyword:
+            return []
+
+        stmt = select(DemandPopularityStats).where(
+            DemandPopularityStats.demand_word_name.like(f"%{keyword}%")
+        )
+        if biz_dt:
+            stmt = stmt.where(DemandPopularityStats.biz_dt == biz_dt)
+        stmt = stmt.order_by(DemandPopularityStats.biz_dt.desc())
+        return list(self.session.scalars(stmt).all())
+
     def list_by_biz_dt_and_belong_ids(
         self, biz_dt: str, belong_ids: list[int]
     ) -> list[DemandPopularityStats]:

+ 117 - 0
supply_infra/db/repositories/demand_video_expansion_repo.py

@@ -0,0 +1,117 @@
+from __future__ import annotations
+
+from collections.abc import Iterable
+
+from sqlalchemy import select
+from sqlalchemy.dialects.mysql import insert
+
+from supply_infra.db.models.demand_video_expansion import (
+    DemandVideoExpansion,
+    DemandVideoExpansionRun,
+)
+from supply_infra.db.repositories.base import BaseRepository
+
+_BATCH_SIZE = 500
+
+
+class DemandVideoExpansionRepository(BaseRepository[DemandVideoExpansion]):
+    """需求视频点位拓展结果 repository。"""
+
+    model = DemandVideoExpansion
+
+    def list_by_demand_grade(
+        self, biz_dt: str, source_demand_grade_id: int
+    ) -> list[DemandVideoExpansion]:
+        """按业务日与来源需求 id 查询拓展结果(未删除)。"""
+        stmt = (
+            select(DemandVideoExpansion)
+            .where(
+                DemandVideoExpansion.biz_dt == biz_dt,
+                DemandVideoExpansion.source_demand_grade_id == int(source_demand_grade_id),
+                DemandVideoExpansion.is_delete == 0,
+            )
+            .order_by(
+                DemandVideoExpansion.video_id,
+                DemandVideoExpansion.point_type,
+                DemandVideoExpansion.id,
+            )
+        )
+        return list(self.session.scalars(stmt).all())
+
+    def bulk_upsert(self, rows: list[dict]) -> int:
+        """按唯一键批量 upsert,冲突时更新 reason / point_desc。"""
+        if not rows:
+            return 0
+
+        affected = 0
+        for i in range(0, len(rows), _BATCH_SIZE):
+            batch = rows[i : i + _BATCH_SIZE]
+            stmt = insert(DemandVideoExpansion).values(batch)
+            stmt = stmt.on_duplicate_key_update(
+                reason=stmt.inserted.reason,
+                point_desc=stmt.inserted.point_desc,
+                run_id=stmt.inserted.run_id,
+                is_delete=0,
+            )
+            result = self.session.execute(stmt)
+            affected += result.rowcount or 0
+        return affected
+
+
+class DemandVideoExpansionRunRepository(BaseRepository[DemandVideoExpansionRun]):
+    """拓展任务执行记录 repository。"""
+
+    model = DemandVideoExpansionRun
+
+    def list_finished_grade_ids(self, biz_dt: str) -> set[int]:
+        """返回指定业务日已成功完成拓展判断的 demand_grade.id 集合。"""
+        stmt = select(DemandVideoExpansionRun.source_demand_grade_id).where(
+            DemandVideoExpansionRun.biz_dt == biz_dt,
+            DemandVideoExpansionRun.status == "finished",
+        )
+        return {int(v) for v in self.session.scalars(stmt).all() if v is not None}
+
+    def upsert_run(
+        self,
+        *,
+        biz_dt: str,
+        run_id: str,
+        source_demand_grade_id: int,
+        saved_count: int,
+        status: str = "finished",
+        error_message: str | None = None,
+    ) -> None:
+        """记录单条需求的拓展判断完成状态。"""
+        row = {
+            "biz_dt": biz_dt,
+            "run_id": run_id,
+            "source_demand_grade_id": int(source_demand_grade_id),
+            "saved_count": int(saved_count),
+            "status": status,
+            "error_message": error_message,
+        }
+        stmt = insert(DemandVideoExpansionRun).values(row)
+        stmt = stmt.on_duplicate_key_update(
+            run_id=stmt.inserted.run_id,
+            saved_count=stmt.inserted.saved_count,
+            status=stmt.inserted.status,
+            error_message=stmt.inserted.error_message,
+        )
+        self.session.execute(stmt)
+
+    def list_by_biz_dt(self, biz_dt: str) -> list[DemandVideoExpansionRun]:
+        stmt = (
+            select(DemandVideoExpansionRun)
+            .where(DemandVideoExpansionRun.biz_dt == biz_dt)
+            .order_by(DemandVideoExpansionRun.source_demand_grade_id)
+        )
+        return list(self.session.scalars(stmt).all())
+
+    def get_by_demand_grade(
+        self, biz_dt: str, source_demand_grade_id: int
+    ) -> DemandVideoExpansionRun | None:
+        stmt = select(DemandVideoExpansionRun).where(
+            DemandVideoExpansionRun.biz_dt == biz_dt,
+            DemandVideoExpansionRun.source_demand_grade_id == int(source_demand_grade_id),
+        )
+        return self.session.scalars(stmt).first()

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

@@ -17,6 +17,134 @@ class MultiDemandPoolDiRepository(BaseRepository[MultiDemandPoolDi]):
 
     model = MultiDemandPoolDi
 
+    def get_latest_biz_dt(self) -> str | None:
+        """返回需求池中最新业务日;无数据时返回 None。"""
+        stmt = select(func.max(MultiDemandPoolDi.biz_dt))
+        return self.session.scalar(stmt)
+
+    def list_distinct_demand_name_summaries(
+        self,
+        biz_dt: str,
+        *,
+        limit: int = 50,
+        offset: int = 0,
+        exclude_names: list[str] | None = None,
+    ) -> list[dict]:
+        """
+        按业务日分页列出去重需求词及聚合统计。
+
+        返回每项包含:demand_name、row_count(出现行数)、strategies(策略列表)、
+        total_weight、total_video_count、max_real_rov_7d、max_real_vov_7d。
+        按 max_real_rov_7d 降序、total_weight 降序排列,优先暴露有后验数据/高权重的需求词。
+        """
+        stmt = (
+            select(
+                MultiDemandPoolDi.demand_name,
+                func.count().label("row_count"),
+                func.group_concat(MultiDemandPoolDi.strategy.distinct()).label("strategies"),
+                func.sum(MultiDemandPoolDi.weight).label("total_weight"),
+                func.sum(MultiDemandPoolDi.video_count).label("total_video_count"),
+                func.max(MultiDemandPoolDi.real_rov_7d).label("max_real_rov_7d"),
+                func.max(MultiDemandPoolDi.real_vov_7d).label("max_real_vov_7d"),
+            )
+            .where(MultiDemandPoolDi.biz_dt == biz_dt)
+        )
+        if exclude_names:
+            stmt = stmt.where(MultiDemandPoolDi.demand_name.notin_(exclude_names))
+        stmt = (
+            stmt.group_by(MultiDemandPoolDi.demand_name)
+            .order_by(
+                func.max(MultiDemandPoolDi.real_rov_7d).desc(),
+                func.sum(MultiDemandPoolDi.weight).desc(),
+            )
+            .limit(limit)
+            .offset(offset)
+        )
+
+        rows = self.session.execute(stmt).all()
+        return [
+            {
+                "demand_name": name,
+                "row_count": int(row_count or 0),
+                "strategies": (strategies or "").split(",") if strategies else [],
+                "total_weight": float(total_weight) if total_weight is not None else None,
+                "total_video_count": int(total_video_count) if total_video_count is not None else None,
+                "max_real_rov_7d": float(max_rov) if max_rov is not None else None,
+                "max_real_vov_7d": float(max_vov) if max_vov is not None else None,
+            }
+            for name, row_count, strategies, total_weight, total_video_count, max_rov, max_vov in rows
+        ]
+
+    def count_distinct_demand_names(self, biz_dt: str) -> int:
+        """统计指定业务日去重需求词总数。"""
+        stmt = select(func.count(func.distinct(MultiDemandPoolDi.demand_name))).where(
+            MultiDemandPoolDi.biz_dt == biz_dt
+        )
+        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]:
+        """
+        按业务日 + 需求名双向包含关系搜索明细行。
+
+        匹配 demand_name LIKE %keyword% 或 keyword LIKE %demand_name%(互相包含),
+        用于把同语义、措辞不同的需求词合并到一起判断。返回按 demand_name 去重后的明细。
+        """
+        keyword = (keyword or "").strip()
+        if not keyword:
+            return []
+
+        stmt = select(MultiDemandPoolDi).where(
+            MultiDemandPoolDi.biz_dt == biz_dt,
+            MultiDemandPoolDi.demand_name.like(f"%{keyword}%"),
+        )
+        rows = list(self.session.scalars(stmt).all())
+
+        if len(keyword) >= 2:
+            broad_stmt = select(MultiDemandPoolDi).where(
+                MultiDemandPoolDi.biz_dt == biz_dt,
+            )
+            seen_ids = {int(r.id) for r in rows}
+            for row in self.session.scalars(broad_stmt).all():
+                if int(row.id) in seen_ids:
+                    continue
+                name = row.demand_name or ""
+                if name and name in keyword:
+                    rows.append(row)
+                    seen_ids.add(int(row.id))
+
+        return [
+            {
+                "id": int(row.id),
+                "demand_name": row.demand_name,
+                "strategy": row.strategy,
+                "weight": None if row.weight is None or row.weight == 0 else float(row.weight),
+                "video_count": int(row.video_count) if row.video_count is not None else None,
+                "real_rov_7d": float(row.real_rov_7d) if row.real_rov_7d is not None else None,
+                "real_vov_7d": float(row.real_vov_7d) if row.real_vov_7d is not None else None,
+            }
+            for row in rows
+        ]
+
+    def get_by_ids(self, ids: list[int]) -> list[MultiDemandPoolDi]:
+        """按 id 批量查询完整行。"""
+        if not ids:
+            return []
+        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)。"""
         stmt = (
@@ -103,7 +231,7 @@ class MultiDemandPoolDiRepository(BaseRepository[MultiDemandPoolDi]):
             MultiDemandPoolDi.strategy.in_(strategies),
         )
         return [
-            (str(strategy), float(weight) if weight is not None else None)
+            (str(strategy), None if weight is None or weight == 0 else float(weight))
             for strategy, weight in self.session.execute(stmt).all()
         ]
 
@@ -176,7 +304,7 @@ class MultiDemandPoolDiRepository(BaseRepository[MultiDemandPoolDi]):
         return updated
 
     def update_video_fields(self, biz_dt: str, rows: list[dict]) -> int:
-        """按 (strategy, demand_id) 批量更新 video_list / video_count。"""
+        """按 (strategy, demand_id) 批量更新 video_list / video_count / weight。"""
         if not rows:
             return 0
 
@@ -192,6 +320,7 @@ class MultiDemandPoolDiRepository(BaseRepository[MultiDemandPoolDi]):
                 .values(
                     video_list=row.get("video_list"),
                     video_count=row.get("video_count"),
+                    weight=row.get("weight"),
                 )
             )
             result = self.session.execute(stmt)

+ 117 - 0
supply_infra/db/repositories/multi_demand_video_point_repo.py

@@ -0,0 +1,117 @@
+from __future__ import annotations
+
+from collections.abc import Iterable
+from typing import Any
+
+from sqlalchemy import delete, select
+
+from supply_infra.db.models.multi_demand_video_point import MultiDemandVideoPoint
+from supply_infra.db.repositories.base import BaseRepository
+from supply_infra.video_points import json_fields_from_point_rows
+
+_BATCH_SIZE = 1000
+
+
+class MultiDemandVideoPointRepository(BaseRepository[MultiDemandVideoPoint]):
+    """需求池视频点位表 — 按 video_id 批量替换与查询。"""
+
+    model = MultiDemandVideoPoint
+
+    def list_video_ids_with_points(self, video_ids: Iterable[str]) -> set[str]:
+        """返回 video_ids 中在点位表已有记录的视频 id。"""
+        vid_list = [v for v in video_ids if v]
+        if not vid_list:
+            return set()
+
+        existing: set[str] = set()
+        for i in range(0, len(vid_list), _BATCH_SIZE):
+            batch = vid_list[i : i + _BATCH_SIZE]
+            stmt = (
+                select(MultiDemandVideoPoint.video_id)
+                .where(MultiDemandVideoPoint.video_id.in_(batch))
+                .distinct()
+            )
+            existing.update(
+                str(v) for v in self.session.scalars(stmt).all() if v
+            )
+        return existing
+
+    def list_by_video_ids(
+        self, video_ids: Iterable[str]
+    ) -> dict[str, list[dict[str, Any]]]:
+        """按 video_id 批量查询点位行,返回 video_id → 行列表。"""
+        vid_list = [v for v in video_ids if v]
+        if not vid_list:
+            return {}
+
+        result: dict[str, list[dict[str, Any]]] = {}
+        for i in range(0, len(vid_list), _BATCH_SIZE):
+            batch = vid_list[i : i + _BATCH_SIZE]
+            stmt = (
+                select(MultiDemandVideoPoint)
+                .where(MultiDemandVideoPoint.video_id.in_(batch))
+                .order_by(
+                    MultiDemandVideoPoint.video_id,
+                    MultiDemandVideoPoint.point_type,
+                    MultiDemandVideoPoint.id,
+                )
+            )
+            for row in self.session.scalars(stmt).all():
+                result.setdefault(str(row.video_id), []).append(
+                    {
+                        "point_type": row.point_type,
+                        "point_data": row.point_data,
+                        "point_desc": row.point_desc,
+                    }
+                )
+        return result
+
+    def json_fields_by_video_ids(
+        self, video_ids: Iterable[str]
+    ) -> dict[str, dict[str, str | None]]:
+        """按 video_id 返回三个 JSON 列(API 兼容)。"""
+        rows_by_vid = self.list_by_video_ids(video_ids)
+        return {
+            vid: json_fields_from_point_rows(rows)
+            for vid, rows in rows_by_vid.items()
+        }
+
+    def replace_for_video_ids(
+        self, points_by_video_id: dict[str, list[dict[str, Any]]]
+    ) -> int:
+        """按 video_id 全量替换点位:先删后插。"""
+        if not points_by_video_id:
+            return 0
+
+        video_ids = [v for v in points_by_video_id if v]
+        if not video_ids:
+            return 0
+
+        for i in range(0, len(video_ids), _BATCH_SIZE):
+            batch = video_ids[i : i + _BATCH_SIZE]
+            self.session.execute(
+                delete(MultiDemandVideoPoint).where(
+                    MultiDemandVideoPoint.video_id.in_(batch)
+                )
+            )
+
+        insert_rows: list[dict[str, Any]] = []
+        for video_id, rows in points_by_video_id.items():
+            if not video_id or not rows:
+                continue
+            for row in rows:
+                insert_rows.append(
+                    {
+                        "video_id": video_id,
+                        "point_type": row["point_type"],
+                        "point_data": row.get("point_data"),
+                        "point_desc": row.get("point_desc"),
+                    }
+                )
+
+        inserted = 0
+        for i in range(0, len(insert_rows), _BATCH_SIZE):
+            batch = insert_rows[i : i + _BATCH_SIZE]
+            self.session.bulk_insert_mappings(MultiDemandVideoPoint, batch)
+            inserted += len(batch)
+        return inserted

+ 50 - 0
supply_infra/db/repositories/scheduler_job_execution_repo.py

@@ -0,0 +1,50 @@
+from __future__ import annotations
+
+from datetime import datetime
+
+from sqlalchemy import select
+
+from supply_infra.db.models.scheduler_job_execution import SchedulerJobExecution
+from supply_infra.db.repositories.base import BaseRepository
+
+
+class SchedulerJobExecutionRepository(BaseRepository[SchedulerJobExecution]):
+    model = SchedulerJobExecution
+
+    def log_event(
+        self,
+        *,
+        run_id: str,
+        job_name: str,
+        status: str,
+        event_time: datetime,
+        job_id: str | None = None,
+        biz_dt: str | None = None,
+        started_at: datetime | None = None,
+        finished_at: datetime | None = None,
+        duration_seconds: float | None = None,
+        error_message: str | None = None,
+        detail: str | None = None,
+    ) -> SchedulerJobExecution:
+        entity = SchedulerJobExecution(
+            run_id=run_id,
+            job_name=job_name,
+            job_id=job_id,
+            status=status,
+            event_time=event_time,
+            biz_dt=biz_dt,
+            started_at=started_at,
+            finished_at=finished_at,
+            duration_seconds=duration_seconds,
+            error_message=error_message,
+            detail=detail,
+        )
+        return self.add(entity)
+
+    def list_recent(self, *, limit: int = 50) -> list[SchedulerJobExecution]:
+        stmt = (
+            select(SchedulerJobExecution)
+            .order_by(SchedulerJobExecution.event_time.desc(), SchedulerJobExecution.id.desc())
+            .limit(limit)
+        )
+        return list(self.session.scalars(stmt).all())

+ 9 - 3
supply_infra/db/session.py

@@ -4,7 +4,7 @@ from collections.abc import Generator
 from contextlib import contextmanager
 from typing import Any
 
-from sqlalchemy import create_engine
+from sqlalchemy import create_engine, inspect
 from sqlalchemy.orm import Session, sessionmaker
 
 from supply_infra.config import get_infra_settings
@@ -29,11 +29,17 @@ def get_engine():
     return _engine
 
 
-def init_db() -> None:
+def init_db() -> dict[str, list[str]]:
     """Create all tables (dev / first-run). Import models before calling."""
     import supply_infra.db.models  # noqa: F401 — register all models
 
-    Base.metadata.create_all(bind=get_engine())
+    engine = get_engine()
+    inspector = inspect(engine)
+    before = set(inspector.get_table_names())
+    Base.metadata.create_all(bind=engine)
+    after = set(inspect(engine).get_table_names())
+    created = sorted(after - before)
+    return {"created": created}
 
 
 @contextmanager

+ 0 - 4
supply_infra/scheduler/__init__.py

@@ -1,5 +1 @@
 """Scheduler for periodic jobs."""
-
-from supply_infra.scheduler.app import create_scheduler, run_scheduler
-
-__all__ = ["create_scheduler", "run_scheduler"]

+ 113 - 29
supply_infra/scheduler/app.py

@@ -1,59 +1,143 @@
 from __future__ import annotations
 
 import logging
+import signal
+import time
+from typing import TYPE_CHECKING
 
-from apscheduler.schedulers.blocking import BlockingScheduler
+from apscheduler.schedulers.background import BackgroundScheduler
 from apscheduler.triggers.cron import CronTrigger
 
 from supply_infra.config import get_infra_settings
-from supply_infra.scheduler.jobs.sync_global_tree_odps_to_mysql import sync_global_tree_odps_to_mysql
-from supply_infra.scheduler.jobs.sync_multi_demand_pool_odps_to_mysql import (
-    sync_multi_demand_pool_odps_to_mysql,
-)
+from supply_infra.scheduler.constants import SUPPLY_PIPELINE_JOB_ID, SUPPLY_PIPELINE_JOB_NAME
+from supply_infra.scheduler.jobs.run_supply_pipeline import run_supply_pipeline
+
+if TYPE_CHECKING:
+    from apscheduler.schedulers.base import BaseScheduler
 
 logger = logging.getLogger(__name__)
 
+_scheduler: BackgroundScheduler | None = None
 
-def create_scheduler() -> BlockingScheduler:
-    """Create and configure the scheduler with all registered jobs."""
-    settings = get_infra_settings()
-    scheduler = BlockingScheduler(timezone=settings.scheduler_timezone)
+_PIPELINE_CRON_HOUR = 12
 
-    # 每天凌晨 2:30 从 ODPS 同步全局树元素与分类到 MySQL
-    scheduler.add_job(
-        sync_global_tree_odps_to_mysql,
-        trigger=CronTrigger(hour=2, minute=30),
-        id="sync_global_tree_odps_to_mysql",
-        name="ODPS → MySQL 全局树同步",
-        replace_existing=True,
-    )
 
-    # 每天 12:00 从 ODPS 同步策略需求天级表到 MySQL(当天 dt)
+def create_scheduler() -> BackgroundScheduler:
+    """Create and configure the scheduler with the chained supply pipeline job."""
+    settings = get_infra_settings()
+    scheduler = BackgroundScheduler(timezone=settings.scheduler_timezone)
+
+    # 每天 12:00 串行执行:全局树 → 需求池 → 分级
     scheduler.add_job(
-        sync_multi_demand_pool_odps_to_mysql,
-        trigger=CronTrigger(hour=12, minute=0),
-        id="sync_multi_demand_pool_odps_to_mysql",
-        name="ODPS → MySQL 策略需求池同步",
+        run_supply_pipeline,
+        trigger=CronTrigger(hour=_PIPELINE_CRON_HOUR, minute=0),
+        id=SUPPLY_PIPELINE_JOB_ID,
+        name=SUPPLY_PIPELINE_JOB_NAME,
         replace_existing=True,
+        max_instances=1,
+        coalesce=True,
+        misfire_grace_time=3600,
     )
 
     logger.info("Scheduler configured with %d job(s)", len(scheduler.get_jobs()))
     return scheduler
 
 
-def run_scheduler() -> None:
-    """Start the blocking scheduler (CLI entry point)."""
+def _log_next_runs(scheduler: BaseScheduler) -> None:
+    for job in scheduler.get_jobs():
+        # apscheduler 3.x: 在 scheduler.start() 之前,Job.next_run_time 访问会抛 AttributeError
+        # 这里用 getattr 兜底,保证启动阶段不因日志而中断。
+        next_run_time = getattr(job, "next_run_time", None)
+        logger.info("  - %s | next run: %s", job.name, next_run_time)
+
+
+def start_scheduler() -> BackgroundScheduler | None:
+    """Start the background scheduler (idempotent). Used by API lifespan."""
+    global _scheduler
+
     settings = get_infra_settings()
     if not settings.scheduler_enabled:
         logger.warning("Scheduler is disabled (SCHEDULER_ENABLED=false)")
-        return
+        return None
 
-    scheduler = create_scheduler()
+    if _scheduler is not None and _scheduler.running:
+        return _scheduler
+
+    _scheduler = create_scheduler()
     logger.info("Starting scheduler...")
-    for job in scheduler.get_jobs():
-        logger.info("  - %s | next run: %s", job.name, job.next_run_time)
+    _scheduler.start()
+    _log_next_runs(_scheduler)
+    return _scheduler
+
+
+def get_scheduler_status() -> dict:
+    """Return scheduler runtime status for health checks."""
+    settings = get_infra_settings()
+    if not settings.scheduler_enabled:
+        return {
+            "enabled": False,
+            "running": False,
+            "timezone": settings.scheduler_timezone,
+            "jobs": [],
+        }
+
+    if _scheduler is None or not _scheduler.running:
+        return {
+            "enabled": True,
+            "running": False,
+            "timezone": settings.scheduler_timezone,
+            "jobs": [],
+        }
+
+    jobs = []
+    for job in _scheduler.get_jobs():
+        next_run = getattr(job, "next_run_time", None)
+        jobs.append(
+            {
+                "id": job.id,
+                "name": job.name,
+                "next_run_time": next_run.isoformat() if next_run else None,
+            }
+        )
+
+    return {
+        "enabled": True,
+        "running": True,
+        "timezone": settings.scheduler_timezone,
+        "jobs": jobs,
+    }
+
+
+def stop_scheduler() -> None:
+    """Shut down the background scheduler if running."""
+    global _scheduler
+
+    if _scheduler is None or not _scheduler.running:
+        return
+
+    logger.info("Stopping scheduler...")
+    _scheduler.shutdown(wait=False)
+    _scheduler = None
+    logger.info("Scheduler stopped.")
+
+
+def run_scheduler() -> None:
+    """Start the scheduler and block until interrupted (CLI entry point)."""
+    scheduler = start_scheduler()
+    if scheduler is None:
+        return
+
+    def _handle_exit(signum: int, _frame: object) -> None:
+        logger.info("Received signal %s, shutting down scheduler...", signum)
+        stop_scheduler()
+        raise SystemExit(0)
+
+    signal.signal(signal.SIGINT, _handle_exit)
+    signal.signal(signal.SIGTERM, _handle_exit)
 
     try:
-        scheduler.start()
+        while scheduler.running:
+            time.sleep(3600)
     except (KeyboardInterrupt, SystemExit):
+        stop_scheduler()
         logger.info("Scheduler stopped.")

+ 4 - 0
supply_infra/scheduler/constants.py

@@ -0,0 +1,4 @@
+"""Scheduler job identifiers shared by app and job implementations."""
+
+SUPPLY_PIPELINE_JOB_ID = "run_supply_pipeline"
+SUPPLY_PIPELINE_JOB_NAME = "供给数据流水线"

+ 120 - 0
supply_infra/scheduler/job_execution.py

@@ -0,0 +1,120 @@
+"""定时任务执行记录写入辅助。"""
+from __future__ import annotations
+
+import json
+import logging
+import uuid
+from datetime import datetime
+from typing import Any
+
+from supply_infra.db.repositories.scheduler_job_execution_repo import (
+    SchedulerJobExecutionRepository,
+)
+from supply_infra.db.session import get_session
+
+logger = logging.getLogger(__name__)
+
+STATUS_STARTED = "started"
+STATUS_FINISHED = "finished"
+STATUS_FAILED = "failed"
+STATUS_SKIPPED = "skipped"
+
+
+def _serialize_detail(payload: Any) -> str | None:
+    if payload is None:
+        return None
+    return json.dumps(payload, ensure_ascii=False, default=str)
+
+
+class JobExecutionRecorder:
+    """在任务开始/结束时各写一条执行记录。"""
+
+    def __init__(
+        self,
+        *,
+        job_name: str,
+        job_id: str | None = None,
+        biz_dt: str | None = None,
+    ) -> None:
+        self.run_id = str(uuid.uuid4())
+        self.job_name = job_name
+        self.job_id = job_id
+        self.biz_dt = biz_dt
+        self.started_at = datetime.now()
+
+    def record_started(self) -> None:
+        try:
+            with get_session() as session:
+                SchedulerJobExecutionRepository(session).log_event(
+                    run_id=self.run_id,
+                    job_name=self.job_name,
+                    job_id=self.job_id,
+                    status=STATUS_STARTED,
+                    event_time=self.started_at,
+                    biz_dt=self.biz_dt,
+                    started_at=self.started_at,
+                )
+        except Exception:
+            logger.exception("Failed to record job start: job_name=%s run_id=%s", self.job_name, self.run_id)
+
+    def record_finished(
+        self,
+        *,
+        success: bool,
+        result: dict[str, Any] | None = None,
+        error_message: str | None = None,
+    ) -> None:
+        finished_at = datetime.now()
+        duration_seconds = round((finished_at - self.started_at).total_seconds(), 2)
+        status = STATUS_FINISHED if success else STATUS_FAILED
+
+        try:
+            with get_session() as session:
+                SchedulerJobExecutionRepository(session).log_event(
+                    run_id=self.run_id,
+                    job_name=self.job_name,
+                    job_id=self.job_id,
+                    status=status,
+                    event_time=finished_at,
+                    biz_dt=self.biz_dt,
+                    started_at=self.started_at,
+                    finished_at=finished_at,
+                    duration_seconds=duration_seconds,
+                    error_message=error_message,
+                    detail=_serialize_detail(result),
+                )
+        except Exception:
+            logger.exception(
+                "Failed to record job finish: job_name=%s run_id=%s status=%s",
+                self.job_name,
+                self.run_id,
+                status,
+            )
+
+
+def record_skipped(
+    *,
+    job_name: str,
+    job_id: str | None = None,
+    biz_dt: str | None = None,
+    reason: str,
+) -> None:
+    """记录因并发等原因跳过的执行。"""
+    now = datetime.now()
+    run_id = str(uuid.uuid4())
+    try:
+        with get_session() as session:
+            SchedulerJobExecutionRepository(session).log_event(
+                run_id=run_id,
+                job_name=job_name,
+                job_id=job_id,
+                status=STATUS_SKIPPED,
+                event_time=now,
+                biz_dt=biz_dt,
+                started_at=now,
+                finished_at=now,
+                duration_seconds=0,
+                error_message=reason,
+            )
+    except Exception:
+        logger.exception("Failed to record skipped job: job_name=%s reason=%s", job_name, reason)

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

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

+ 18 - 4
supply_infra/scheduler/jobs/compute_category_tree_weight.py

@@ -65,6 +65,12 @@ def _dec(value: float, places: int = 8) -> Decimal:
     return Decimal(str(round(float(value), places)))
 
 
+def _optional_dec(value: float | None, places: int = 8) -> Decimal | None:
+    if value is None:
+        return None
+    return _dec(value, places)
+
+
 def _normalize_parent_id(parent_id: int | None) -> int | None:
     if parent_id is None or parent_id == 0:
         return None
@@ -107,7 +113,10 @@ def _aggregate_hang_features(
             n = int(getattr(stats, f"{dim}_count") or 0)
             if n <= 0:
                 continue
-            avg = float(getattr(stats, f"{dim}_avg") or 0)
+            avg_raw = getattr(stats, f"{dim}_avg", None)
+            if avg_raw is None:
+                continue
+            avg = float(avg_raw)
             cell = acc[category_id][dim]
             cell["wsum"] += avg * n
             cell["nsum"] += n
@@ -193,8 +202,8 @@ def _row_from_state(
     }
     for dim in METRIC_KEYS:
         ds = dim_stats[dim]
-        row[f"{dim}_avg"] = _dec(ds.avg)
         row[f"{dim}_count"] = int(ds.count)
+        row[f"{dim}_avg"] = _optional_dec(ds.avg) if ds.count > 0 else None
     return row
 
 
@@ -202,8 +211,13 @@ def _materialize_stats_row(stats: Any) -> SimpleNamespace:
     """在 session 内抽出标量,避免 DetachedInstanceError。"""
     payload: dict[str, Any] = {"demand_category_id": int(stats.demand_category_id)}
     for dim in METRIC_KEYS:
-        payload[f"{dim}_count"] = int(getattr(stats, f"{dim}_count") or 0)
-        payload[f"{dim}_avg"] = float(getattr(stats, f"{dim}_avg") or 0)
+        count = int(getattr(stats, f"{dim}_count") or 0)
+        payload[f"{dim}_count"] = count
+        if count > 0:
+            avg_raw = getattr(stats, f"{dim}_avg", None)
+            payload[f"{dim}_avg"] = float(avg_raw) if avg_raw is not None else None
+        else:
+            payload[f"{dim}_avg"] = None
     return SimpleNamespace(**payload)
 
 

+ 335 - 0
supply_infra/scheduler/jobs/expand_demand_from_video_points.py

@@ -0,0 +1,335 @@
+"""
+从 S/A 级需求关联视频中挖掘拓展需求。
+
+任务层负责查库与组装上下文;Agent 仅做语义判断与落库。
+"""
+from __future__ import annotations
+
+import json
+import logging
+import uuid
+from concurrent.futures import ThreadPoolExecutor, as_completed
+from datetime import datetime
+from typing import Any
+from zoneinfo import ZoneInfo
+
+from agents.demand_video_expand_agent.run import (
+    DemandExpandContext,
+    VideoPoint,
+    extract_saved_count,
+    judge_demand_expansion,
+)
+from supply_infra.config import get_infra_settings
+from supply_infra.db.models.demand_grade import DemandGrade
+from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
+from supply_infra.db.repositories.demand_video_expansion_repo import (
+    DemandVideoExpansionRunRepository,
+)
+from supply_infra.db.repositories.multi_demand_video_point_repo import (
+    MultiDemandVideoPointRepository,
+)
+from supply_infra.db.session import get_session
+
+logger = logging.getLogger(__name__)
+
+_DEFAULT_WORKERS = 5
+
+
+def _resolve_biz_dt(biz_dt: str | None) -> str:
+    if biz_dt:
+        text = str(biz_dt).strip()
+        if len(text) == 8 and text.isdigit():
+            return text
+        raise ValueError(f"biz_dt 格式无效,应为 YYYYMMDD: {biz_dt!r}")
+
+    with get_session() as session:
+        latest = DemandGradeRepository(session).get_latest_biz_dt()
+    if latest:
+        return str(latest)
+
+    timezone = ZoneInfo(get_infra_settings().scheduler_timezone)
+    return datetime.now(timezone).strftime("%Y%m%d")
+
+
+def _parse_video_ids(raw: Any) -> list[str]:
+    if raw is None:
+        return []
+    items: list[Any]
+    if isinstance(raw, str):
+        text = raw.strip()
+        if not text:
+            return []
+        try:
+            parsed = json.loads(text)
+            items = list(parsed) if isinstance(parsed, list) else [text]
+        except (ValueError, TypeError):
+            items = [part.strip() for part in text.split(",") if part.strip()]
+    elif isinstance(raw, (list, tuple)):
+        items = list(raw)
+    else:
+        return []
+    out: list[str] = []
+    seen: set[str] = set()
+    for item in items:
+        vid = str(item).strip() if item is not None else ""
+        if not vid or vid in seen:
+            continue
+        seen.add(vid)
+        out.append(vid)
+    return out
+
+
+def _to_float(value: Any) -> float | None:
+    if value is None:
+        return None
+    return float(value)
+
+
+def load_expand_contexts(
+    session,
+    biz_dt: str,
+    *,
+    run_id: str,
+    skip_finished: bool = True,
+) -> tuple[list[DemandExpandContext], dict[str, int]]:
+    """加载待拓展判断的需求上下文,并返回跳过统计。"""
+    stats = {
+        "total_sa": 0,
+        "skipped_no_video": 0,
+        "skipped_no_points": 0,
+        "skipped_already_done": 0,
+    }
+
+    grades = DemandGradeRepository(session).list_by_biz_dt_and_grades(biz_dt, ("S", "A"))
+    stats["total_sa"] = len(grades)
+
+    finished_ids: set[int] = set()
+    if skip_finished:
+        finished_ids = DemandVideoExpansionRunRepository(session).list_finished_grade_ids(
+            biz_dt
+        )
+
+    rows_with_video: list[tuple[DemandGrade, list[str]]] = []
+    all_video_ids: set[str] = set()
+
+    for row in grades:
+        if int(row.id) in finished_ids:
+            stats["skipped_already_done"] += 1
+            continue
+
+        video_ids = _parse_video_ids(row.video_list)
+        if not video_ids:
+            stats["skipped_no_video"] += 1
+            continue
+
+        rows_with_video.append((row, video_ids))
+        all_video_ids.update(video_ids)
+
+    points_by_vid = MultiDemandVideoPointRepository(session).list_by_video_ids(all_video_ids)
+
+    contexts: list[DemandExpandContext] = []
+    for row, video_ids in rows_with_video:
+        points: list[VideoPoint] = []
+        for vid in video_ids:
+            for point in points_by_vid.get(vid, []):
+                points.append(
+                    VideoPoint(
+                        video_id=vid,
+                        point_type=str(point.get("point_type") or ""),
+                        point_data=point.get("point_data"),
+                        point_desc=point.get("point_desc"),
+                    )
+                )
+        if not points:
+            stats["skipped_no_points"] += 1
+            continue
+
+        contexts.append(
+            DemandExpandContext(
+                biz_dt=biz_dt,
+                run_id=run_id,
+                demand_grade_id=int(row.id),
+                demand_name=str(row.demand_name),
+                grade=str(row.grade),
+                score=_to_float(row.score),
+                video_ids=video_ids,
+                points=points,
+            )
+        )
+
+    return contexts, stats
+
+
+def _record_run(
+    *,
+    biz_dt: str,
+    run_id: str,
+    source_demand_grade_id: int,
+    saved_count: int,
+    status: str,
+    error_message: str | None = None,
+) -> None:
+    with get_session() as session:
+        DemandVideoExpansionRunRepository(session).upsert_run(
+            biz_dt=biz_dt,
+            run_id=run_id,
+            source_demand_grade_id=source_demand_grade_id,
+            saved_count=saved_count,
+            status=status,
+            error_message=error_message,
+        )
+
+
+def _process_single_expand(
+    ctx: DemandExpandContext,
+    *,
+    biz_dt: str,
+    run_id: str,
+) -> dict[str, Any]:
+    """并发 worker:对单个需求执行拓展判断并落库执行记录。"""
+    try:
+        agent_result = judge_demand_expansion(ctx)
+        saved_count = extract_saved_count(agent_result)
+        _record_run(
+            biz_dt=biz_dt,
+            run_id=run_id,
+            source_demand_grade_id=ctx.demand_grade_id,
+            saved_count=saved_count,
+            status="finished",
+        )
+        logger.info(
+            "expand demand done: grade_id=%s demand=%s saved=%d iterations=%d",
+            ctx.demand_grade_id,
+            ctx.demand_name,
+            saved_count,
+            agent_result.iterations,
+        )
+        return {
+            "success": True,
+            "demand_grade_id": ctx.demand_grade_id,
+            "demand_name": ctx.demand_name,
+            "saved_count": saved_count,
+            "iterations": agent_result.iterations,
+        }
+    except Exception as exc:
+        error_text = str(exc)
+        _record_run(
+            biz_dt=biz_dt,
+            run_id=run_id,
+            source_demand_grade_id=ctx.demand_grade_id,
+            saved_count=0,
+            status="failed",
+            error_message=error_text,
+        )
+        logger.exception(
+            "expand demand failed: grade_id=%s demand=%s",
+            ctx.demand_grade_id,
+            ctx.demand_name,
+        )
+        return {
+            "success": False,
+            "demand_grade_id": ctx.demand_grade_id,
+            "demand_name": ctx.demand_name,
+            "error": error_text,
+        }
+
+
+def expand_demand_from_video_points(
+    biz_dt: str | None = None,
+    *,
+    skip_finished: bool = True,
+    workers: int = _DEFAULT_WORKERS,
+) -> dict[str, Any]:
+    """
+    对指定业务日的 S/A 需求执行视频点位拓展判断。
+
+    程序负责查需求与点位;无视频或无点位则跳过;有数据则并发调用 Agent。
+    """
+    started_at = datetime.now()
+    run_id = uuid.uuid4().hex
+
+    try:
+        resolved_biz_dt = _resolve_biz_dt(biz_dt)
+    except Exception as exc:
+        logger.exception("expand_demand_from_video_points preflight failed")
+        return {
+            "success": False,
+            "error": str(exc),
+            "started_at": started_at.isoformat(),
+            "finished_at": datetime.now().isoformat(),
+        }
+
+    with get_session() as session:
+        contexts, preload_stats = load_expand_contexts(
+            session,
+            resolved_biz_dt,
+            run_id=run_id,
+            skip_finished=skip_finished,
+        )
+
+    logger.info(
+        "expand_demand_from_video_points start: biz_dt=%s run_id=%s workers=%s pending=%s",
+        resolved_biz_dt,
+        run_id,
+        workers,
+        len(contexts),
+    )
+
+    result: dict[str, Any] = {
+        "success": True,
+        "run_id": run_id,
+        "biz_dt": resolved_biz_dt,
+        "started_at": started_at.isoformat(),
+        "workers": 0,
+        **preload_stats,
+        "processed": 0,
+        "saved_total": 0,
+        "failed": 0,
+        "errors": [],
+    }
+
+    if not contexts:
+        finished_at = datetime.now()
+        result["finished_at"] = finished_at.isoformat()
+        result["duration_seconds"] = round((finished_at - started_at).total_seconds(), 2)
+        logger.info("expand_demand_from_video_points finished: %s", result)
+        return result
+
+    worker_count = max(1, min(int(workers), len(contexts)))
+    result["workers"] = worker_count
+
+    with ThreadPoolExecutor(max_workers=worker_count) as executor:
+        futures = [
+            executor.submit(_process_single_expand, ctx, biz_dt=resolved_biz_dt, run_id=run_id)
+            for ctx in contexts
+        ]
+        for future in as_completed(futures):
+            try:
+                item_result = future.result()
+            except Exception as exc:
+                logger.exception("expand demand worker 出现未捕获错误: biz_dt=%s", resolved_biz_dt)
+                result["failed"] += 1
+                result["errors"].append({"error": str(exc)})
+                continue
+
+            result["processed"] += 1
+            if item_result.get("success"):
+                result["saved_total"] += int(item_result.get("saved_count") or 0)
+                continue
+
+            result["failed"] += 1
+            result["errors"].append(
+                {
+                    "demand_grade_id": item_result.get("demand_grade_id"),
+                    "demand_name": item_result.get("demand_name"),
+                    "error": item_result.get("error"),
+                }
+            )
+
+    finished_at = datetime.now()
+    result["finished_at"] = finished_at.isoformat()
+    result["duration_seconds"] = round((finished_at - started_at).total_seconds(), 2)
+    result["success"] = result["failed"] == 0
+
+    logger.info("expand_demand_from_video_points finished: %s", result)
+    return result

+ 386 - 0
supply_infra/scheduler/jobs/grade_demand_pool.py

@@ -0,0 +1,386 @@
+"""统筹落库后执行分级任务。"""
+from __future__ import annotations
+
+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_orchestrator_agent.run import orchestrate_daily_grade_plan
+from supply_infra.config import get_infra_settings
+from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
+from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
+from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
+from supply_infra.db.session import get_session
+from supply_infra.scheduler.plan_group_batch import MAX_DEMANDS_PER_BATCH, split_even_batches
+
+logger = logging.getLogger(__name__)
+
+_DEFAULT_WORKERS = 5
+
+
+def _resolve_biz_dt(biz_dt: str | None) -> str:
+    if biz_dt:
+        return biz_dt
+    return datetime.now(ZoneInfo(get_infra_settings().scheduler_timezone)).strftime("%Y%m%d")
+
+
+def _materialize_pending_group_items(biz_dt: str) -> int:
+    """执行前为 pending 计划组物化待分级需求明细。"""
+    with get_session() as session:
+        graded_names = DemandGradeRepository(session).get_existing_demand_names(biz_dt)
+        return DemandGradePlanRepository(session).materialize_pending_groups(
+            biz_dt,
+            graded_names=graded_names,
+        )
+
+
+def _group_success(item_summary: dict[str, int], processed_batches: int) -> bool:
+    if item_summary.get("pending", 0) > 0:
+        return False
+    total = sum(item_summary.values())
+    if total == 0:
+        return True
+    if processed_batches > 0:
+        return True
+    return item_summary.get("failed", 0) == 0 and (
+        item_summary.get("finished", 0) > 0 or item_summary.get("skipped", 0) > 0
+    )
+
+
+def _run_group(biz_dt: str, group_id: int, *, max_demands_per_batch: int) -> int:
+    """领取并执行一个固定任务,从组内需求明细表按批读取。"""
+    with get_session() as session:
+        group = DemandGradePlanRepository(session).claim_group(biz_dt, group_id)
+    if group is None:
+        logger.warning("分级任务已被其他 worker 领取或状态已变化: biz_dt=%s group_id=%s", biz_dt, group_id)
+        return 0
+
+    gid = int(group["id"])
+    processed_batches = 0
+    batch_errors: list[str] = []
+    try:
+        with get_session() as session:
+            items = DemandGradePlanRepository(session).list_pending_group_items(gid)
+
+        for batch_items in split_even_batches(items, max_per_batch=max_demands_per_batch):
+            demands = [
+                {"pool_id": item["pool_id"], "demand_name": item["demand_name"]}
+                for item in batch_items
+            ]
+            item_ids = [int(item["item_id"]) for item in batch_items]
+            try:
+                grade_demand_words(demands, biz_dt=biz_dt)
+                with get_session() as session:
+                    DemandGradePlanRepository(session).mark_group_items_status(
+                        item_ids,
+                        status="finished",
+                    )
+                processed_batches += 1
+            except Exception as exc:
+                logger.exception(
+                    "分级子批次失败,跳过并继续本组其余需求: biz_dt=%s group=%s count=%s",
+                    biz_dt,
+                    group["group_key"],
+                    len(demands),
+                )
+                batch_errors.append(str(exc))
+                with get_session() as session:
+                    DemandGradePlanRepository(session).mark_group_items_status(
+                        item_ids,
+                        status="failed",
+                        error_message=str(exc),
+                    )
+
+        with get_session() as session:
+            repo = DemandGradePlanRepository(session)
+            item_summary = repo.summarize_group_items(gid)
+            repo.finish_group(
+                gid,
+                success=_group_success(item_summary, processed_batches),
+                error_message="; ".join(batch_errors) if batch_errors else None,
+            )
+    except Exception as exc:
+        logger.exception(
+            "分级计划任务失败: biz_dt=%s group=%s",
+            biz_dt,
+            group["group_key"],
+        )
+        try:
+            with get_session() as session:
+                DemandGradePlanRepository(session).finish_group(
+                    gid,
+                    success=False,
+                    error_message=str(exc),
+                )
+        except Exception:
+            logger.exception(
+                "记录分级计划任务失败状态时发生错误: biz_dt=%s group=%s",
+                biz_dt,
+                group["group_key"],
+            )
+        return 0
+    return processed_batches
+
+
+def _execute_plan_tasks(
+    biz_dt: str,
+    *,
+    workers: int,
+    max_demands_per_batch: int,
+) -> dict[str, Any]:
+    """并发执行当天全部 pending 计划任务,仅执行一轮。"""
+    with get_session() as session:
+        group_ids = DemandGradePlanRepository(session).list_pending_group_ids(biz_dt)
+    if not group_ids:
+        with get_session() as session:
+            final_snapshot = DemandGradePlanRepository(session).get_execution_snapshot(biz_dt)
+        return {
+            "attempted_groups": 0,
+            "groups_run": 0,
+            "workers": 0,
+            "final_snapshot": final_snapshot,
+            "execution_complete": final_snapshot["execution_complete"],
+        }
+
+    worker_count = max(1, min(int(workers), len(group_ids)))
+    groups_run = 0
+    with ThreadPoolExecutor(max_workers=worker_count) as executor:
+        futures = [
+            executor.submit(_run_group, biz_dt, group_id, max_demands_per_batch=max_demands_per_batch)
+            for group_id in group_ids
+        ]
+        for future in as_completed(futures):
+            try:
+                groups_run += future.result()
+            except Exception:
+                logger.exception("分级任务 worker 出现未捕获错误: biz_dt=%s", biz_dt)
+
+    with get_session() as session:
+        final_snapshot = DemandGradePlanRepository(session).get_execution_snapshot(biz_dt)
+    return {
+        "attempted_groups": len(group_ids),
+        "groups_run": groups_run,
+        "workers": worker_count,
+        "final_snapshot": final_snapshot,
+        "execution_complete": final_snapshot["execution_complete"],
+    }
+
+
+def execute_plan_tasks_until_complete(
+    biz_dt: str,
+    *,
+    workers: int,
+    max_demands_per_batch: int = MAX_DEMANDS_PER_BATCH,
+    max_rounds: int = 0,
+) -> dict[str, Any]:
+    """循环执行 pending 计划任务,直到全部完成或达到 max_rounds。"""
+    rounds: list[dict[str, Any]] = []
+    groups_run = 0
+    round_no = 0
+    while True:
+        if max_rounds > 0 and round_no >= max_rounds:
+            break
+        round_no += 1
+        result = _execute_plan_tasks(
+            biz_dt,
+            workers=workers,
+            max_demands_per_batch=max_demands_per_batch,
+        )
+        rounds.append(result)
+        groups_run += int(result.get("groups_run") or 0)
+        if int(result.get("attempted_groups") or 0) == 0:
+            break
+
+    final_snapshot = rounds[-1]["final_snapshot"] if rounds else None
+    if final_snapshot is None:
+        with get_session() as session:
+            final_snapshot = DemandGradePlanRepository(session).get_execution_snapshot(biz_dt)
+
+    return {
+        "rounds": round_no,
+        "groups_run": groups_run,
+        "workers": max(1, int(workers)),
+        "round_details": rounds,
+        "final_snapshot": final_snapshot,
+        "execution_complete": bool(final_snapshot["execution_complete"]),
+    }
+
+
+def _grade_demand_pool_impl(
+    resolved_biz_dt: str,
+    *,
+    workers: int,
+    max_demands_per_batch: int = MAX_DEMANDS_PER_BATCH,
+    with_orchestrate: bool = True,
+    max_rounds: int = 0,
+) -> dict[str, Any]:
+    with get_session() as session:
+        total = MultiDemandPoolDiRepository(session).count_distinct_demand_names(resolved_biz_dt)
+        graded_before = DemandGradeRepository(session).count_by_biz_dt(resolved_biz_dt)
+
+    if with_orchestrate:
+        try:
+            orchestrate_daily_grade_plan(biz_dt=resolved_biz_dt)
+        except Exception:
+            logger.exception("统筹 Agent 执行失败,继续处理数据库中已有任务: biz_dt=%s", resolved_biz_dt)
+
+    materialized = _materialize_pending_group_items(resolved_biz_dt)
+    logger.info("物化计划组需求明细: biz_dt=%s items=%s", resolved_biz_dt, materialized)
+
+    plan_execution = execute_plan_tasks_until_complete(
+        resolved_biz_dt,
+        workers=max(1, int(workers)),
+        max_demands_per_batch=max(1, min(int(max_demands_per_batch), MAX_DEMANDS_PER_BATCH)),
+        max_rounds=max_rounds,
+    )
+
+    with get_session() as session:
+        graded_after = DemandGradeRepository(session).count_by_biz_dt(resolved_biz_dt)
+    final_snapshot = plan_execution["final_snapshot"]
+    result = {
+        "success": bool(plan_execution["execution_complete"]),
+        "biz_dt": resolved_biz_dt,
+        "total": total,
+        "graded_before": graded_before,
+        "graded_after": graded_after,
+        "materialized_items": materialized,
+        "planned_category_count": len(final_snapshot["assigned_category_ids"]),
+        "planned_groups": final_snapshot["planned_groups"],
+        "group_status": final_snapshot["group_status"],
+        "plan_execution": plan_execution,
+        "workers": plan_execution["workers"],
+        "groups_run": plan_execution["groups_run"],
+        "run_at": datetime.now().isoformat(),
+    }
+    logger.info("Grade demand pool completed: %s", result)
+    return result
+
+
+def retry_failed_plan_group_items(
+    biz_dt: str | None = None,
+    *,
+    workers: int = _DEFAULT_WORKERS,
+    max_demands_per_batch: int = MAX_DEMANDS_PER_BATCH,
+    group_ids: list[int] | None = None,
+    dry_run: bool = False,
+) -> dict[str, Any]:
+    """将 demand_grade_plan_group_item 中 failed 记录重置后重新执行分级。"""
+    resolved_biz_dt = _resolve_biz_dt(biz_dt)
+    try:
+        with get_session() as session:
+            failed_items = DemandGradePlanRepository(session).list_failed_group_items(
+                resolved_biz_dt,
+                group_ids=group_ids,
+            )
+
+        if not failed_items:
+            with get_session() as session:
+                snapshot = DemandGradePlanRepository(session).get_execution_snapshot(resolved_biz_dt)
+            return {
+                "success": True,
+                "biz_dt": resolved_biz_dt,
+                "dry_run": dry_run,
+                "failed_items": 0,
+                "reset": {"reset_items": 0, "reset_groups": 0, "group_ids": [], "items": []},
+                "execution": None,
+                "final_snapshot": snapshot,
+                "run_at": datetime.now().isoformat(),
+            }
+
+        if dry_run:
+            affected_group_ids = sorted({int(item["group_id"]) for item in failed_items})
+            return {
+                "success": True,
+                "biz_dt": resolved_biz_dt,
+                "dry_run": True,
+                "failed_items": len(failed_items),
+                "reset": {
+                    "reset_items": len(failed_items),
+                    "reset_groups": len(affected_group_ids),
+                    "group_ids": affected_group_ids,
+                    "items": failed_items,
+                },
+                "execution": None,
+                "run_at": datetime.now().isoformat(),
+            }
+
+        with get_session() as session:
+            reset_result = DemandGradePlanRepository(session).reset_failed_items_to_pending(
+                resolved_biz_dt,
+                group_ids=group_ids,
+            )
+
+        with get_session() as session:
+            graded_before = DemandGradeRepository(session).count_by_biz_dt(resolved_biz_dt)
+
+        plan_execution = execute_plan_tasks_until_complete(
+            resolved_biz_dt,
+            workers=max(1, int(workers)),
+            max_demands_per_batch=max(1, min(int(max_demands_per_batch), MAX_DEMANDS_PER_BATCH)),
+        )
+
+        with get_session() as session:
+            graded_after = DemandGradeRepository(session).count_by_biz_dt(resolved_biz_dt)
+            remaining_failed = DemandGradePlanRepository(session).list_failed_group_items(
+                resolved_biz_dt,
+                group_ids=group_ids,
+            )
+
+        final_snapshot = plan_execution["final_snapshot"]
+        return {
+            "success": len(remaining_failed) == 0,
+            "biz_dt": resolved_biz_dt,
+            "dry_run": False,
+            "failed_items": len(failed_items),
+            "reset": reset_result,
+            "graded_before": graded_before,
+            "graded_after": graded_after,
+            "remaining_failed": len(remaining_failed),
+            "execution": plan_execution,
+            "final_snapshot": final_snapshot,
+            "group_status": final_snapshot["group_status"],
+            "run_at": datetime.now().isoformat(),
+        }
+    except Exception as exc:
+        logger.exception(
+            "重试 failed plan group items 发生未捕获错误: biz_dt=%s",
+            resolved_biz_dt,
+        )
+        return {
+            "success": False,
+            "biz_dt": resolved_biz_dt,
+            "error": str(exc),
+            "run_at": datetime.now().isoformat(),
+        }
+
+
+def grade_demand_pool(
+    biz_dt: str | None = None,
+    *,
+    workers: int = _DEFAULT_WORKERS,
+    max_demands_per_batch: int = MAX_DEMANDS_PER_BATCH,
+    with_orchestrate: bool = True,
+    max_rounds: int = 0,
+) -> dict[str, Any]:
+    """执行完整分级闭环;任何错误只记录日志并返回,不向上中断定时任务。"""
+    resolved_biz_dt = str(biz_dt or "")
+    try:
+        resolved_biz_dt = _resolve_biz_dt(biz_dt)
+        return _grade_demand_pool_impl(
+            resolved_biz_dt,
+            workers=workers,
+            max_demands_per_batch=max_demands_per_batch,
+            with_orchestrate=with_orchestrate,
+            max_rounds=max_rounds,
+        )
+    except Exception as exc:
+        logger.exception("需求分级任务发生未捕获错误,已阻止异常中断定时任务: biz_dt=%s", resolved_biz_dt)
+        return {
+            "success": False,
+            "biz_dt": resolved_biz_dt,
+            "error": str(exc),
+            "run_at": datetime.now().isoformat(),
+        }

+ 216 - 0
supply_infra/scheduler/jobs/run_supply_pipeline.py

@@ -0,0 +1,216 @@
+"""
+串联执行的供给数据流水线(可重复执行、幂等):
+
+1. ODPS → MySQL 全局树同步(T-1 分区)
+2. ODPS → MySQL 策略需求池同步(当天 biz_dt)
+3. 需求池分级评估(同上 biz_dt)
+4. S/A 需求视频点位拓展判断
+
+各子步骤内部已做去重(INSERT IGNORE、diff 同步、跳过已分级词等);
+本文件额外用进程内锁防止同一轮次并发重入,并隔离各步骤异常:前一步失败时记录
+error 后继续后续步骤,最终返回失败结果而不向 APScheduler 抛异常。
+"""
+from __future__ import annotations
+
+import logging
+import threading
+from datetime import datetime, timedelta
+from typing import Any, Callable
+from zoneinfo import ZoneInfo
+
+from supply_infra.config import get_infra_settings
+from supply_infra.scheduler.constants import (
+    SUPPLY_PIPELINE_JOB_ID,
+    SUPPLY_PIPELINE_JOB_NAME,
+)
+from supply_infra.scheduler.job_execution import JobExecutionRecorder, record_skipped
+from supply_infra.scheduler.jobs.expand_demand_from_video_points import (
+    expand_demand_from_video_points,
+)
+from supply_infra.scheduler.jobs.grade_demand_pool import grade_demand_pool
+from supply_infra.scheduler.jobs.sync_global_tree_odps_to_mysql import sync_global_tree_odps_to_mysql
+from supply_infra.scheduler.jobs.sync_multi_demand_pool_odps_to_mysql import (
+    sync_multi_demand_pool_odps_to_mysql,
+)
+
+logger = logging.getLogger(__name__)
+
+_pipeline_lock = threading.Lock()
+
+
+def _resolve_dates(biz_dt: str | None) -> tuple[str, str]:
+    """返回 (biz_dt, global_tree_partition_date),global_tree 使用 biz_dt 前一日。"""
+    if biz_dt:
+        resolved_biz_dt = biz_dt
+    else:
+        timezone = ZoneInfo(get_infra_settings().scheduler_timezone)
+        resolved_biz_dt = datetime.now(timezone).strftime("%Y%m%d")
+    tree_partition = (
+        datetime.strptime(resolved_biz_dt, "%Y%m%d") - timedelta(days=1)
+    ).strftime("%Y%m%d")
+    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]:
+    """
+    按顺序执行全局树同步 → 需求池同步 → 需求分级 → 视频点位拓展。
+
+    Args:
+        biz_dt: 业务日 YYYYMMDD;省略则取当天。
+
+    Returns:
+        各步骤统计;若上一轮仍在执行则返回 skipped。
+    """
+    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")
+        record_skipped(
+            job_name=SUPPLY_PIPELINE_JOB_NAME,
+            job_id=SUPPLY_PIPELINE_JOB_ID,
+            biz_dt=resolved_biz_dt,
+            reason="already_running",
+        )
+        return {
+            "skipped": True,
+            "reason": "already_running",
+            "run_at": datetime.now().isoformat(),
+        }
+
+    started_at = datetime.now()
+    recorder = JobExecutionRecorder(
+        job_name=SUPPLY_PIPELINE_JOB_NAME,
+        job_id=SUPPLY_PIPELINE_JOB_ID,
+        biz_dt=resolved_biz_dt,
+    )
+    recorder.record_started()
+
+    logger.info(
+        "Supply pipeline start: biz_dt=%s global_tree_partition=%s run_id=%s",
+        resolved_biz_dt,
+        tree_partition,
+        recorder.run_id,
+    )
+
+    result: dict[str, Any] = {
+        "run_id": recorder.run_id,
+        "biz_dt": resolved_biz_dt,
+        "global_tree_partition": tree_partition,
+        "started_at": started_at.isoformat(),
+    }
+    errors: list[str] = []
+    step_status: dict[str, dict[str, Any]] = {}
+    success = False
+
+    try:
+        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),
+            ),
+            (
+                "expand_video_points",
+                lambda: expand_demand_from_video_points(
+                    biz_dt=resolved_biz_dt,
+                    workers=5,
+                ),
+            ),
+        ]
+        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
+        errors.append(f"pipeline: {exc}")
+        logger.exception(
+            "Supply pipeline orchestration failed but will not escape scheduler: "
+            "biz_dt=%s global_tree_partition=%s",
+            resolved_biz_dt,
+            tree_partition,
+        )
+    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)
+        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

+ 117 - 75
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 (
@@ -111,6 +111,17 @@ def _normalize_video_list(raw: Any) -> tuple[str | None, int]:
     return json.dumps(truncated, ensure_ascii=False), len(truncated)
 
 
+def _normalize_weight(raw: Any) -> float | None:
+    """weight 为 0 或无效时视为无分数,存 null。"""
+    if raw is None:
+        return None
+    try:
+        value = float(raw)
+    except (TypeError, ValueError):
+        return None
+    return None if value == 0 else value
+
+
 def _to_mysql_rows(raw_rows: list[dict[str, Any]], biz_dt: str) -> list[dict[str, Any]]:
     """转换为 MySQL 行,按 (strategy, demand_id) 去重(保留最后一条)。"""
     by_key: dict[RowKey, dict[str, Any]] = {}
@@ -127,7 +138,7 @@ def _to_mysql_rows(raw_rows: list[dict[str, Any]], biz_dt: str) -> list[dict[str
             "strategy": str(strategy),
             "demand_id": str(demand_id),
             "demand_name": str(demand_name),
-            "weight": row.get("weight"),
+            "weight": _normalize_weight(row.get("weight")),
             "type": str(row["type"]) if row.get("type") is not None else None,
             "video_count": video_count,
             "video_list": video_list,
@@ -173,21 +184,40 @@ 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,
     }
 
 
 def _sync_diff(partition_date: str) -> dict[str, Any]:
-    """行数不同时拉取 ODPS,只同步 (strategy, demand_id) 差异,并回填已有行的 video 字段。"""
+    """行数不同时拉取 ODPS,只同步 (strategy, demand_id) 差异,并回填已有行的 video/weight 字段。"""
     odps = get_odps_client()
     raw_rows = odps.fetch_multi_demand_pool(partition_date)
     mysql_rows = _to_mysql_rows(raw_rows, partition_date)
@@ -227,30 +257,6 @@ def _sync_diff(partition_date: str) -> dict[str, Any]:
     }
 
 
-def backfill_video_list(partition_date: str) -> dict[str, Any]:
-    """从 ODPS 回填指定分区的 video_list / video_count(每条最多前 10 个 video_id)。"""
-    logger.info("Backfill video_list for partition: %s", partition_date)
-    odps = get_odps_client()
-    raw_rows = odps.fetch_multi_demand_pool(partition_date)
-    mysql_rows = _to_mysql_rows(raw_rows, partition_date)
-
-    with get_session() as session:
-        updated = MultiDemandPoolDiRepository(session).update_video_fields(
-            partition_date,
-            mysql_rows,
-        )
-
-    result = {
-        "partition_date": partition_date,
-        "fetched": len(raw_rows),
-        "unique_rows": len(mysql_rows),
-        "updated": updated,
-        "with_video": sum(1 for r in mysql_rows if r.get("video_list")),
-    }
-    logger.info("Backfill video_list completed: %s", result)
-    return result
-
-
 def _to_float(value: Any) -> float | None:
     if value is None:
         return None
@@ -349,15 +355,14 @@ def enrich_real_rov_vov_7d(biz_dt: str) -> dict[str, Any]:
 def _calc_metric_stats(
     weights: list[float],
     places: int = 2,
-) -> tuple[Decimal, int]:
+) -> tuple[Decimal | None, int]:
     """
     计算单策略 avg / count。
-    权重为 0 不参与平均值;全部为 0(或无有效权重)时 avg=0、count=0。
+    权重为 0 不参与平均值;全部为 0(或无有效权重)时 avg=null、count=0。
     """
     nonzero = [w for w in weights if w != 0]
     if not nonzero:
-        zero = Decimal("0").quantize(Decimal("0." + "0" * places))
-        return zero, 0
+        return None, 0
 
     avg_val = sum(nonzero) / len(nonzero)
     q = Decimal("0." + "0" * places)
@@ -367,12 +372,10 @@ def _calc_metric_stats(
     )
 
 
-def _empty_metric_stats() -> dict[str, Decimal | int]:
-    result: dict[str, Decimal | int] = {}
+def _empty_metric_stats() -> dict[str, Decimal | int | None]:
+    result: dict[str, Decimal | int | None] = {}
     for key in _ALL_METRIC_KEYS:
-        places = _METRIC_DECIMAL_PLACES[key]
-        zero = Decimal("0").quantize(Decimal("0." + "0" * places))
-        result[f"{key}_avg"] = zero
+        result[f"{key}_avg"] = None
         result[f"{key}_count"] = 0
     return result
 
@@ -388,7 +391,7 @@ def _build_stats_row(
     grouped: dict[str, list[float]] = {key: [] for key in _METRIC_KEYS}
     for strategy, weight in strategy_weights:
         metric = _STRATEGY_METRIC.get(strategy)
-        if metric is None or weight is None:
+        if metric is None or weight is None or weight == 0:
             continue
         grouped[metric].append(float(weight))
 
@@ -476,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:
@@ -499,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,
@@ -507,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)

+ 29 - 40
supply_infra/scheduler/jobs/sync_multi_demand_videos.py

@@ -13,22 +13,21 @@ from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPo
 from supply_infra.db.repositories.multi_demand_video_detail_repo import (
     MultiDemandVideoDetailRepository,
 )
+from supply_infra.db.repositories.multi_demand_video_point_repo import (
+    MultiDemandVideoPointRepository,
+)
 from supply_infra.db.session import get_session
 from supply_infra.odps.client import ODPSClient, get_odps_client
+from supply_infra.video_points import (
+    extract_points_json_from_decode,
+    points_from_decode_payload,
+)
 
 logger = logging.getLogger(__name__)
 
 _FINAL_TOPIC_KEY = "最终选题"
 _TARGET_POST_KEY = "target_post"
 _TITLE_KEY = "title"
-_POINT_KEY = "点"
-_POINT_DESC_KEY = "点描述"
-# decode_result 中的点位 key → 落库字段名
-_POINT_FIELD_MAP = {
-    "灵感点": "inspiration_points_json",
-    "目的点": "purpose_points_json",
-    "关键点": "key_points_json",
-}
 VIDEO_SYNC_BATCH_SIZE = 100
 
 
@@ -83,33 +82,9 @@ def _extract_final_topic_json(payload: dict[str, Any]) -> str | None:
     return json.dumps(final_topic, ensure_ascii=False)
 
 
-def _extract_points(payload: dict[str, Any], key: str) -> str | None:
-    """从 decode_result[key](灵感点/目的点/关键点)取每项的 点/点描述,重组为 JSON 文本。"""
-    items = payload.get(key)
-    if not isinstance(items, list):
-        return None
-
-    points: list[dict[str, Any]] = []
-    for item in items:
-        if not isinstance(item, dict):
-            continue
-        point = item.get(_POINT_KEY)
-        point_desc = item.get(_POINT_DESC_KEY)
-        if point is None and point_desc is None:
-            continue
-        points.append({_POINT_KEY: point, _POINT_DESC_KEY: point_desc})
-
-    if not points:
-        return None
-    return json.dumps(points, ensure_ascii=False)
-
-
 def _extract_all_points(payload: dict[str, Any]) -> dict[str, str | None]:
-    """从 decode_result 提取灵感点/目的点/关键点三个字段,返回 落库字段名 → JSON 文本。"""
-    return {
-        field: _extract_points(payload, key)
-        for key, field in _POINT_FIELD_MAP.items()
-    }
+    """从 decode_result 提取灵感点/目的点/关键点三个 JSON 列。"""
+    return extract_points_json_from_decode(payload)
 
 
 def _extract_title(payload: dict[str, Any]) -> str | None:
@@ -144,6 +119,7 @@ def _sync_one_batch(
     )
 
     insert_rows: list[dict[str, Any]] = []
+    point_rows_by_vid: dict[str, list[dict[str, Any]]] = {}
     skipped_no_topic = 0
     seen_vids: set[str] = set()
     for row in odps_rows:
@@ -170,11 +146,17 @@ def _sync_one_batch(
                 **_extract_all_points(payload),
             }
         )
+        point_rows = points_from_decode_payload(vid, payload)
+        if point_rows:
+            point_rows_by_vid[vid] = point_rows
 
     with get_session() as session:
-        inserted = MultiDemandVideoDetailRepository(session).bulk_insert_ignore(
-            insert_rows
-        )
+        detail_repo = MultiDemandVideoDetailRepository(session)
+        inserted = detail_repo.bulk_insert_ignore(insert_rows)
+        if point_rows_by_vid:
+            MultiDemandVideoPointRepository(session).replace_for_video_ids(
+                point_rows_by_vid
+            )
 
     odps_vids = {
         str(r.get("vid")).strip()
@@ -424,6 +406,7 @@ def backfill_video_points(
             decode_dt, batch_vids, batch_size=len(batch_vids)
         )
         points_by_vid: dict[str, dict[str, str | None]] = {}
+        point_rows_by_vid: dict[str, list[dict[str, Any]]] = {}
         for row in odps_rows:
             raw_vid = row.get("vid")
             if raw_vid is None:
@@ -436,11 +419,17 @@ def backfill_video_points(
                 total_skipped += 1
                 continue
             points_by_vid[vid] = _extract_all_points(payload)
+            point_rows = points_from_decode_payload(vid, payload)
+            if point_rows:
+                point_rows_by_vid[vid] = point_rows
 
         with get_session() as session:
-            updated = MultiDemandVideoDetailRepository(session).update_points(
-                points_by_vid
-            )
+            detail_repo = MultiDemandVideoDetailRepository(session)
+            updated = detail_repo.update_points(points_by_vid)
+            if point_rows_by_vid:
+                MultiDemandVideoPointRepository(session).replace_for_video_ids(
+                    point_rows_by_vid
+                )
         total_updated += updated
 
         odps_vids = {

+ 31 - 39
supply_infra/scheduler/jobs/update_category_tree_rank_scores.py

@@ -3,8 +3,8 @@ category_tree_weight 四维热度全局排名归一化打分。
 
 在 category_tree_weight 全部节点写入完成后执行:
 1. 对 ext_pop / plat_sust_pop / plat_ly_pop / recent_pop 各自独立做全局排名
-2. 将排名映射为 [1/n, 1] 区间的归一化分(count>0 的节点参与排名,其余为 0
-3. 四维分相加写入 total_score
+2. 将排名映射为 [1/n, 1] 区间的归一化分(count>0 的节点参与排名,其余为 null
+3. 四维分相加写入 total_score(无有效维度时为 null)
 """
 from __future__ import annotations
 
@@ -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,31 +34,10 @@ 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
+    return _dec(value, places)
 
 
 def _materialize_weight_row(row: Any) -> dict[str, Any]:
@@ -67,8 +47,13 @@ def _materialize_weight_row(row: Any) -> dict[str, Any]:
         "biz_dt": str(row.biz_dt),
     }
     for dim in POP_DIM_KEYS:
-        payload[f"{dim}_count"] = int(getattr(row, f"{dim}_count", 0) or 0)
-        payload[f"{dim}_avg"] = float(getattr(row, f"{dim}_avg", 0) or 0)
+        count = int(getattr(row, f"{dim}_count", 0) or 0)
+        payload[f"{dim}_count"] = count
+        if count > 0:
+            avg_raw = getattr(row, f"{dim}_avg", None)
+            payload[f"{dim}_avg"] = float(avg_raw) if avg_raw is not None else None
+        else:
+            payload[f"{dim}_avg"] = None
     return payload
 
 
@@ -83,7 +68,10 @@ def _build_rank_score_rows(
             count = int(row[f"{dim}_count"])
             if count <= 0:
                 continue
-            avg = float(row[f"{dim}_avg"])
+            avg_raw = row[f"{dim}_avg"]
+            if avg_raw is None:
+                continue
+            avg = float(avg_raw)
             candidates.append((int(row["category_id"]), avg))
         dim_scores[dim] = rank_to_scores(candidates)
 
@@ -92,21 +80,25 @@ def _build_rank_score_rows(
         category_id = int(row["category_id"])
         biz_dt = str(row["biz_dt"])
         score_values = {
-            f"{dim}_score": dim_scores[dim].get(category_id, 0.0)
+            f"{dim}_score": dim_scores[dim].get(category_id)
             for dim in POP_DIM_KEYS
         }
-        total = sum(score_values.values())
+        non_null_scores = [v for v in score_values.values() if v is not None]
+        total = sum(non_null_scores) if non_null_scores else None
         updates.append(
             {
                 "category_id": category_id,
                 "biz_dt": biz_dt,
-                **{col: _dec(score_values[col]) for col in (
-                    "ext_pop_score",
-                    "plat_sust_pop_score",
-                    "plat_ly_pop_score",
-                    "recent_pop_score",
-                )},
-                "total_score": _dec(total),
+                **{
+                    col: _optional_dec(score_values[col])
+                    for col in (
+                        "ext_pop_score",
+                        "plat_sust_pop_score",
+                        "plat_ly_pop_score",
+                        "recent_pop_score",
+                    )
+                },
+                "total_score": _optional_dec(total),
             }
         )
     return updates

+ 82 - 0
supply_infra/scheduler/plan_group_batch.py

@@ -0,0 +1,82 @@
+"""计划组需求明细:物化与读取。"""
+from __future__ import annotations
+
+from typing import Any
+
+from supply_infra.db.repositories.demand_belong_category_repo import DemandBelongCategoryRepository
+from supply_infra.db.repositories.demand_belong_pool_rel_repo import DemandBelongPoolRelRepository
+from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
+from supply_infra.db.session import get_session
+
+MAX_DEMANDS_PER_BATCH = 30
+
+
+def split_even_batches(items: list[Any], *, max_per_batch: int = MAX_DEMANDS_PER_BATCH) -> list[list[Any]]:
+    """按总数均分子批次:总数不超过上限则一批;否则递增组数直到每组不超过上限,余数从前组分配。"""
+    n = len(items)
+    if n == 0:
+        return []
+    cap = max(1, int(max_per_batch))
+    if n <= cap:
+        return [items]
+
+    k = 2
+    while True:
+        base = n // k
+        rem = n % k
+        max_size = base + 1 if rem > 0 else base
+        if max_size <= cap:
+            break
+        k += 1
+
+    batches: list[list[Any]] = []
+    idx = 0
+    for i in range(k):
+        size = base + (1 if i < rem else 0)
+        batches.append(items[idx : idx + size])
+        idx += size
+    return batches
+
+
+def _priority_sort_key(name: str, priority_index: dict) -> tuple:
+    return (
+        priority_index.get(name, {}).get("source_rank_score") is None,
+        -float(priority_index.get(name, {}).get("source_rank_score") or 0),
+        float(priority_index.get(name, {}).get("global_demand_rank") or float("inf")),
+        name,
+    )
+
+
+def resolve_demands_for_category_ids(
+    biz_dt: str,
+    category_ids: list[int],
+) -> list[dict[str, Any]]:
+    """按分类节点解析全部需求池记录(pool_id + demand_name)。"""
+    selected_ids = list(dict.fromkeys(int(value) for value in category_ids))
+    with get_session() as session:
+        belongs = DemandBelongCategoryRepository(session).list_by_category_ids(selected_ids)
+        pool_ids_by_belong = DemandBelongPoolRelRepository(session).get_pool_ids_by_belong_ids(
+            [int(row.id) for row in belongs]
+        )
+        pool_ids = sorted({pool_id for values in pool_ids_by_belong.values() for pool_id in values})
+        pool_repo = MultiDemandPoolDiRepository(session)
+        pool_rows = pool_repo.get_by_ids(pool_ids)
+        from agents.demand_grade_agent.tools.demand_priority import build_demand_priority_index
+
+        priority_index = build_demand_priority_index(pool_repo.list_by_biz_dt(biz_dt))
+        candidates: list[dict[str, Any]] = []
+        for row in pool_rows:
+            if row.biz_dt != biz_dt or not row.demand_name:
+                continue
+            candidates.append({
+                "pool_id": int(row.id),
+                "demand_name": str(row.demand_name),
+            })
+
+    candidates.sort(
+        key=lambda item: (
+            *_priority_sort_key(item["demand_name"], priority_index),
+            item["pool_id"],
+        )
+    )
+    return candidates

+ 150 - 0
supply_infra/video_points.py

@@ -0,0 +1,150 @@
+"""视频点位解析与序列化 — decode_result / JSON 字段 ↔ 行记录。"""
+from __future__ import annotations
+
+import json
+from typing import Any
+
+from supply_infra.db.models.multi_demand_video_point import (
+    POINT_TYPE_INSPIRATION,
+    POINT_TYPE_KEY,
+    POINT_TYPE_PURPOSE,
+)
+
+_POINT_KEY = "点"
+_POINT_DESC_KEY = "点描述"
+
+# decode_result 中的 key → point_type
+DECODE_RESULT_KEY_TO_POINT_TYPE = {
+    "灵感点": POINT_TYPE_INSPIRATION,
+    "目的点": POINT_TYPE_PURPOSE,
+    "关键点": POINT_TYPE_KEY,
+}
+
+# 原 JSON 列名 → point_type
+JSON_FIELD_TO_POINT_TYPE = {
+    "inspiration_points_json": POINT_TYPE_INSPIRATION,
+    "purpose_points_json": POINT_TYPE_PURPOSE,
+    "key_points_json": POINT_TYPE_KEY,
+}
+
+# point_type → 原 JSON 列名(API 兼容)
+POINT_TYPE_TO_JSON_FIELD = {v: k for k, v in JSON_FIELD_TO_POINT_TYPE.items()}
+
+
+def _item_to_row(
+    video_id: str, point_type: str, item: dict[str, Any]
+) -> dict[str, str | None] | None:
+    point = item.get(_POINT_KEY)
+    point_desc = item.get(_POINT_DESC_KEY)
+    if point is None and point_desc is None:
+        return None
+    return {
+        "video_id": video_id,
+        "point_type": point_type,
+        "point_data": str(point) if point is not None else None,
+        "point_desc": str(point_desc) if point_desc is not None else None,
+    }
+
+
+def points_from_decode_payload(
+    video_id: str, payload: dict[str, Any]
+) -> list[dict[str, str | None]]:
+    """从 decode_result 提取全部点位行。"""
+    rows: list[dict[str, str | None]] = []
+    for decode_key, point_type in DECODE_RESULT_KEY_TO_POINT_TYPE.items():
+        items = payload.get(decode_key)
+        if not isinstance(items, list):
+            continue
+        for item in items:
+            if not isinstance(item, dict):
+                continue
+            row = _item_to_row(video_id, point_type, item)
+            if row:
+                rows.append(row)
+    return rows
+
+
+def points_from_json_fields(
+    video_id: str,
+    *,
+    inspiration_points_json: str | None = None,
+    purpose_points_json: str | None = None,
+    key_points_json: str | None = None,
+) -> list[dict[str, str | None]]:
+    """从 multi_demand_video_detail 的三个 JSON 列提取点位行。"""
+    field_values = {
+        "inspiration_points_json": inspiration_points_json,
+        "purpose_points_json": purpose_points_json,
+        "key_points_json": key_points_json,
+    }
+    rows: list[dict[str, str | None]] = []
+    for field, point_type in JSON_FIELD_TO_POINT_TYPE.items():
+        text = field_values.get(field)
+        if not text:
+            continue
+        try:
+            items = json.loads(text)
+        except json.JSONDecodeError:
+            continue
+        if not isinstance(items, list):
+            continue
+        for item in items:
+            if not isinstance(item, dict):
+                continue
+            row = _item_to_row(video_id, point_type, item)
+            if row:
+                rows.append(row)
+    return rows
+
+
+def json_fields_from_point_rows(
+    rows: list[dict[str, Any]],
+) -> dict[str, str | None]:
+    """将点位行重组为三个 JSON 列(保持 API 兼容)。"""
+    grouped: dict[str, list[dict[str, str | None]]] = {
+        POINT_TYPE_INSPIRATION: [],
+        POINT_TYPE_PURPOSE: [],
+        POINT_TYPE_KEY: [],
+    }
+    for row in rows:
+        point_type = row.get("point_type")
+        if point_type not in grouped:
+            continue
+        grouped[point_type].append(
+            {
+                _POINT_KEY: row.get("point_data"),
+                _POINT_DESC_KEY: row.get("point_desc"),
+            }
+        )
+
+    result: dict[str, str | None] = {}
+    for point_type, field in POINT_TYPE_TO_JSON_FIELD.items():
+        items = grouped[point_type]
+        result[field] = (
+            json.dumps(items, ensure_ascii=False) if items else None
+        )
+    return result
+
+
+def extract_points_json_from_decode(payload: dict[str, Any]) -> dict[str, str | None]:
+    """从 decode_result 提取灵感点/目的点/关键点 JSON 列(原逻辑)。"""
+    result: dict[str, str | None] = {}
+    for decode_key, field in {
+        k: POINT_TYPE_TO_JSON_FIELD[v]
+        for k, v in DECODE_RESULT_KEY_TO_POINT_TYPE.items()
+    }.items():
+        items = payload.get(decode_key)
+        if not isinstance(items, list):
+            result[field] = None
+            continue
+        points: list[dict[str, Any]] = []
+        for item in items:
+            if not isinstance(item, dict):
+                continue
+            point = item.get(_POINT_KEY)
+            point_desc = item.get(_POINT_DESC_KEY)
+            if point is None and point_desc is None:
+                continue
+            points.append({_POINT_KEY: point, _POINT_DESC_KEY: point_desc})
+        result[field] = json.dumps(points, ensure_ascii=False) if points else None
+    return result

+ 9 - 8
web/src/api/demand.ts

@@ -1,17 +1,18 @@
-import type { DemandBelongResponse, DemandVideosResponse } from '../types/demand'
+import type { DemandGradeResponse, DemandGradeVideosResponse } from '../types/demand'
 
-export async function fetchDemandBelongCategory(): Promise<DemandBelongResponse> {
-  const res = await fetch('/api/demand-belong-category')
+export async function fetchDemandGrade(bizDt?: string | null): Promise<DemandGradeResponse> {
+  const query = bizDt ? `?biz_dt=${encodeURIComponent(bizDt)}` : ''
+  const res = await fetch(`/api/demand-grade${query}`)
   if (!res.ok) {
-    throw new Error(`加载需求归属失败: ${res.status} ${res.statusText}`)
+    throw new Error(`加载需求分级失败: ${res.status} ${res.statusText}`)
   }
   return res.json()
 }
 
-export async function fetchDemandBelongVideos(
-  belongId: number,
-): Promise<DemandVideosResponse> {
-  const res = await fetch(`/api/demand-belong-category/${belongId}/videos`)
+export async function fetchDemandGradeVideos(
+  demandGradeId: number,
+): Promise<DemandGradeVideosResponse> {
+  const res = await fetch(`/api/demand-grade/${demandGradeId}/videos`)
   if (!res.ok) {
     throw new Error(`加载关联视频失败: ${res.status} ${res.statusText}`)
   }

Bu fark içinde çok fazla dosya değişikliği olduğu için bazı dosyalar gösterilmiyor