Przeglądaj źródła

增加需求产出agent

xueyiming 2 tygodni temu
rodzic
commit
2000a1683a

+ 5 - 27
agents/generate_demand_agent/agent.py

@@ -3,37 +3,14 @@ generate_demand_agent 工厂 — 组装需求生成 Agent。
 """
 from __future__ import annotations
 
+from pathlib import Path
+
 from supply_agent import Agent
 from supply_agent.config import Settings
 from agents.generate_demand_agent.tools import register_all_tools
 
-GENERATE_DEMAND_AGENT_SYSTEM_PROMPT = """
-你是需求生成专家。从类目树落到局部需求组合产出。
-
-热度维度含义——必须按此理解与选用:
-- ext_pop:外部热度
-- plat_sust_pop:平台持续热度
-- plat_ly_pop:平台内去年同周期热度
-- recent_pop:平台内近期热度
-选用 source_dim 时写清对应含义;不同维度代表不同热度来源,不可混用概念。
-
-产出层级(严格按此四层,不可跳层或颠倒):
-1) source_dim:来源热度维度(选哪一种热度信号)
-2) overall_direction:整体方向(介于来源维度与汇总事件之间的大分类总结)
-3) summary_event:汇总事件(对单个或极少数强相关需求的轻量概括,禁止大杂烩式合集)
-4) demand_name:需求名称(必须原样取自 demand_belong_category.name,禁止新造词)
-
-说明:
-- 一个 overall_direction 下可有多个 summary_event。
-- summary_event 优先与 demand_name 一对一;仅当 2~3 个需求同属一个具体事件/场景时才可合并,禁止把大量弱相关需求塞进同一事件。
-- 宁可多写几个 summary_event,也不要做一个过度聚合的大事件。
-
-产出原则:
-- overall_direction 要像“大分类标签”,比维度细、比事件粗。
-- summary_event 要像“单一可识别的事件/场景”,聚焦一个传播点;不要写“XX合集”“XX大全”式过度聚合。
-- 多个 demand_name 只有在同一具体事件下才允许挂在同一 summary_event,且一般不超过 2 个。
-- 理由写清:路径、维度、方向、事件、选用了哪些需求名。
-""".strip()
+_PROMPT_PATH = Path(__file__).parent / "prompt" / "system_prompt.md"
+GENERATE_DEMAND_AGENT_SYSTEM_PROMPT = _PROMPT_PATH.read_text(encoding="utf-8")
 
 
 def create_generate_demand_agent(
@@ -47,6 +24,7 @@ def create_generate_demand_agent(
         name="generate_demand_agent",
         model=model,
         system_prompt=GENERATE_DEMAND_AGENT_SYSTEM_PROMPT,
+        max_iterations=30,
     )
     register_all_tools(agent.tools)
     return agent

+ 46 - 0
agents/generate_demand_agent/prompt/system_prompt.md

@@ -0,0 +1,46 @@
+## 角色与任务
+你是需求生成专家。从类目树落到局部需求组合产出。
+
+一次运行只处理**一个** source_dim,不可混用多维度概念。若用户要求多个维度,按维度分别完成整条流程。
+
+## 热度维度含义
+选用 source_dim 时必须按此理解:
+- ext_pop:外部热度
+- plat_sust_pop:平台持续热度
+- plat_ly_pop:平台内去年同周期热度
+- recent_pop:平台内近期热度
+
+## 产出层级(严格按此四层,不可跳层或颠倒)
+1) source_dim:来源热度维度
+2) overall_direction:整体方向(介于来源维度与汇总事件之间的大分类总结)
+3) summary_event:汇总事件(对单个或极少数强相关需求的轻量概括,禁止大杂烩式合集)
+4) demand_name:需求名称(必须原样取自工具返回的 demand_belong_category.name,禁止新造词)
+
+说明:
+- 一个 overall_direction 下可有多个 summary_event。
+- summary_event 优先与 demand_name 一对一;仅当 2~3 个需求同属一个具体事件/场景时才可合并,禁止把大量弱相关需求塞进同一事件。
+- 宁可多写几个 summary_event,也不要做一个过度聚合的大事件。
+
+## 可用工具
+- `query_latest_biz_dt()`:**未指定 biz_dt 时先调用**,返回两表均有数据的最新业务日(一个确切日期),再用于后续工具。
+- `query_category_tree_by_dim(dimension, biz_dt=None)`:按维度查看有数据的类目树;带「+」表示其下有带数据的叶子。`biz_dt` 格式 YYYYMMDD,省略则用最新业务日。
+- `query_category_leaves_by_dim(dimension, category_ids, biz_dt=None)`:下钻到指定分支下有数据的叶子节点。
+- `query_demand_words_by_category(category_ids, dimension, ..., biz_dt=None)`:查询类目(含子树)下可选用的挂载词及该维热度。
+- `query_category_path(category_ids)`:查询类目根到叶路径,用于写 reason。
+- `batch_save_generated_demands(items, biz_dt=None)`:校验并落库四层产出;`biz_dt` 作为本次落库默认业务日。
+
+## 单维度工作流程
+1. 若用户未指定 `biz_dt`,先调用 `query_latest_biz_dt()`,取返回的 `biz_dt=YYYYMMDD`;同一轮查询与落库保持该日期。
+2. 选定一个 source_dim,调用 `query_category_tree_by_dim(dim, biz_dt)`,找高分且带「+」的分支。
+3. 调用 `query_category_leaves_by_dim(dim, branch_ids, biz_dt)`,落到叶子热点。
+4. 调用 `query_demand_words_by_category(leaf_ids, dim, top_k=20, biz_dt=biz_dt)`,获取候选 demand_name。
+5. 对拟选用的类目调用 `query_category_path`,补全路径。
+6. 按四层组织产出:overall_direction → summary_event → demand_name;禁止跨维度、禁止造词。
+7. 调用 `batch_save_generated_demands(items, biz_dt=biz_dt)` 落库;items 必填 source_dim / overall_direction / summary_event / demand_name,并尽量带上 category_id、reason。
+
+## 产出原则
+- overall_direction 要像“大分类标签”,比维度细、比事件粗。
+- summary_event 要像“单一可识别的事件/场景”,聚焦一个传播点;不要写“XX合集”“XX大全”式过度聚合。
+- 多个 demand_name 只有在同一具体事件下才允许挂在同一 summary_event,且一般不超过 2 个。
+- 理由写清:路径、维度、方向、事件、选用了哪些需求名。
+- 只使用工具返回的真实类目与需求词,禁止幻觉编造。

+ 1 - 1
agents/generate_demand_agent/run.py

@@ -10,7 +10,7 @@ def main() -> None:
     print(f"tools: {agent.tools.list_tools()}")
     print()
     user_input = f'''
-    帮我基于不同维度,产出一批效果好的需求,每个维度都产出一批效果好的需求
+    帮我基于平台持续热度产出一批效果好的需求
        '''
     result = agent.run(user_input)
     print(result.content)

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

@@ -6,23 +6,39 @@ from __future__ import annotations
 from collections.abc import Callable
 from typing import Any
 
+from agents.generate_demand_agent.tools.batch_save_generated_demands import (
+    batch_save_generated_demands,
+)
 from agents.generate_demand_agent.tools.query_category_leaves_by_dim import (
     query_category_leaves_by_dim,
 )
+from agents.generate_demand_agent.tools.query_category_path import query_category_path
 from agents.generate_demand_agent.tools.query_category_tree_by_dim import (
     query_category_tree_by_dim,
 )
+from agents.generate_demand_agent.tools.query_latest_biz_dt import query_latest_biz_dt
+from agents.generate_demand_agent.tools.query_demand_words_by_category import (
+    query_demand_words_by_category,
+)
 from supply_agent.tools.registry import ToolRegistry
 
 ALL_TOOLS: list[Callable[..., Any]] = [
+    query_latest_biz_dt,
     query_category_tree_by_dim,
     query_category_leaves_by_dim,
+    query_demand_words_by_category,
+    query_category_path,
+    batch_save_generated_demands,
 ]
 
 __all__ = [
     "ALL_TOOLS",
+    "batch_save_generated_demands",
     "query_category_leaves_by_dim",
+    "query_category_path",
     "query_category_tree_by_dim",
+    "query_demand_words_by_category",
+    "query_latest_biz_dt",
     "register_all_tools",
 ]
 

+ 288 - 0
agents/generate_demand_agent/tools/batch_save_generated_demands.py

@@ -0,0 +1,288 @@
+"""
+批量保存单维度需求产生结果到 generated_demand 表。
+"""
+from __future__ import annotations
+
+import logging
+import uuid
+from decimal import Decimal
+from typing import Any
+
+from agents.generate_demand_agent.tools.dim_constants import (
+    DIM_KEYS,
+    DIM_LABEL,
+    build_category_path,
+    resolve_biz_dt,
+)
+from supply_agent.tools import tool
+from supply_infra.db.repositories.demand_belong_category_repo import (
+    DemandBelongCategoryRepository,
+)
+from supply_infra.db.repositories.demand_popularity_stats_repo import (
+    DemandPopularityStatsRepository,
+)
+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.session import get_session
+
+logger = logging.getLogger(__name__)
+
+
+def _optional_int(value: Any) -> int | None:
+    if value is None or value == "":
+        return None
+    return int(value)
+
+
+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,
+) -> tuple[list[dict[str, Any]], list[str]]:
+    """校验并规范化待插入行,返回 (rows, errors)。"""
+    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
+
+        source_dim = _optional_str(item.get("source_dim"))
+        overall_direction = _optional_str(item.get("overall_direction"))
+        summary_event = _optional_str(item.get("summary_event"))
+        demand_name = _optional_str(item.get("demand_name"))
+
+        if not source_dim:
+            errors.append(f"第 {idx} 项缺少 source_dim")
+            continue
+        if source_dim not in DIM_KEYS:
+            allowed = "、".join(f"{k}({DIM_LABEL[k]})" for k in DIM_KEYS)
+            errors.append(f"第 {idx} 项 source_dim 无效(只能是:{allowed}): {source_dim}")
+            continue
+        if not overall_direction:
+            errors.append(f"第 {idx} 项缺少 overall_direction")
+            continue
+        if not summary_event:
+            errors.append(f"第 {idx} 项缺少 summary_event")
+            continue
+        if not demand_name:
+            errors.append(f"第 {idx} 项缺少 demand_name")
+            continue
+
+        dedupe_key = (source_dim, demand_name)
+        if dedupe_key in seen_keys:
+            errors.append(
+                f"第 {idx} 项在本次请求中重复: source_dim={source_dim}, demand_name={demand_name}"
+            )
+            continue
+        seen_keys.add(dedupe_key)
+
+        try:
+            category_id = _optional_int(item.get("category_id"))
+        except (TypeError, ValueError):
+            errors.append(f"第 {idx} 项 category_id 无效: {item.get('category_id')!r}")
+            continue
+
+        try:
+            demand_belong_id = _optional_int(item.get("demand_belong_id"))
+        except (TypeError, ValueError):
+            errors.append(
+                f"第 {idx} 项 demand_belong_id 无效: {item.get('demand_belong_id')!r}"
+            )
+            continue
+
+        run_id = _optional_str(item.get("run_id")) or default_run_id
+        reason = _optional_str(item.get("reason"))
+        category_path = _optional_str(item.get("category_path"))
+        biz_dt = _optional_str(item.get("biz_dt"))
+
+        dim_avg: Decimal | None = None
+        dim_count: int | None = None
+        if item.get("dim_avg") is not None and item.get("dim_avg") != "":
+            try:
+                dim_avg = Decimal(str(item.get("dim_avg")))
+            except Exception:
+                errors.append(f"第 {idx} 项 dim_avg 无效: {item.get('dim_avg')!r}")
+                continue
+        if item.get("dim_count") is not None and item.get("dim_count") != "":
+            try:
+                dim_count = int(item.get("dim_count"))
+            except (TypeError, ValueError):
+                errors.append(f"第 {idx} 项 dim_count 无效: {item.get('dim_count')!r}")
+                continue
+
+        rows.append(
+            {
+                "source_dim": source_dim,
+                "overall_direction": overall_direction,
+                "summary_event": summary_event,
+                "demand_name": demand_name,
+                "demand_belong_id": demand_belong_id,
+                "category_id": category_id,
+                "category_path": category_path,
+                "dim_avg": dim_avg,
+                "dim_count": dim_count,
+                "reason": reason,
+                "biz_dt": biz_dt,
+                "run_id": run_id,
+                "is_delete": 0,
+            }
+        )
+
+    return rows, errors
+
+
+@tool
+def batch_save_generated_demands(
+    items: list[dict[str, Any]],
+    biz_dt: str | None = None,
+) -> str:
+    """
+    批量保存单维度需求产生结果。
+
+    写入 generated_demand 表。demand_name 必须已存在于 demand_belong_category;
+    不存在的词会被跳过。同一 run_id 内 (source_dim, demand_name) 去重。
+
+    Args:
+        items: 待保存列表。每项必填:
+            - source_dim: 四维之一(ext_pop / plat_sust_pop / plat_ly_pop / recent_pop)
+            - overall_direction: 整体方向
+            - summary_event: 汇总事件
+            - demand_name: 需求名(须来自 demand_belong_category.name)
+          选填:
+            - category_id, category_path, demand_belong_id, reason,
+              dim_avg, dim_count, biz_dt, run_id
+          若省略 run_id,本次调用自动生成同一 run_id。
+          若省略 demand_belong_id / category_id / 热度快照,将尽量从库中补全。
+        biz_dt: 业务日 YYYYMMDD;省略则使用 demand_popularity_stats 最新业务日。
+          会作为本次落库的默认 biz_dt,并用于补全热度快照。
+
+    Returns:
+        保存结果摘要。
+    """
+    if not items:
+        return "items 不能为空"
+
+    default_run_id = uuid.uuid4().hex
+    rows, errors = _normalize_items(items, default_run_id=default_run_id)
+    if not rows:
+        detail = ";".join(errors) if errors else "无有效数据"
+        return f"没有可保存的数据: {detail}"
+
+    try:
+        with get_session() as session:
+            stats_repo = DemandPopularityStatsRepository(session)
+            resolved_biz_dt, err = resolve_biz_dt(
+                biz_dt,
+                get_latest=stats_repo.get_latest_biz_dt,
+                has_data=stats_repo.has_biz_dt,
+                table_label="demand_popularity_stats 数据",
+            )
+            if err:
+                return err
+
+            belong_repo = DemandBelongCategoryRepository(session)
+            names = [r["demand_name"] for r in rows]
+            belong_by_name = belong_repo.get_by_names(names)
+
+            valid_rows: list[dict[str, Any]] = []
+            skipped_names: list[str] = []
+            for row in rows:
+                belong = belong_by_name.get(row["demand_name"])
+                if belong is None:
+                    skipped_names.append(row["demand_name"])
+                    continue
+                if row["demand_belong_id"] is None:
+                    row["demand_belong_id"] = int(belong.id)
+                if row["category_id"] is None:
+                    row["category_id"] = int(belong.category_id)
+                valid_rows.append(row)
+
+            if not valid_rows:
+                detail = "、".join(skipped_names)
+                extra = f";校验失败: {';'.join(errors)}" if errors else ""
+                return f"没有可保存的数据: demand_name 不存在于 demand_belong_category: {detail}{extra}"
+
+            # 补全路径
+            need_path_ids = [
+                int(r["category_id"])
+                for r in valid_rows
+                if r["category_id"] is not None and not r.get("category_path")
+            ]
+            if need_path_ids:
+                categories = GlobalTreeCategoryRepository(session).list_active_categories()
+                by_id = {int(c.id): c for c in categories}
+                for row in valid_rows:
+                    if row.get("category_path") or row["category_id"] is None:
+                        continue
+                    path = build_category_path(int(row["category_id"]), by_id)
+                    if path:
+                        row["category_path"] = path
+
+            # 补全 biz_dt 与热度快照
+            need_stats_ids = [
+                int(r["demand_belong_id"])
+                for r in valid_rows
+                if r["demand_belong_id"] is not None
+                and (r.get("dim_avg") is None or r.get("dim_count") is None)
+            ]
+            stats_by_belong: dict[int, Any] = {}
+            if need_stats_ids:
+                for stats in stats_repo.list_by_biz_dt_and_belong_ids(
+                    resolved_biz_dt, need_stats_ids
+                ):
+                    stats_by_belong[int(stats.demand_category_id)] = stats
+
+            for row in valid_rows:
+                if not row.get("biz_dt"):
+                    row["biz_dt"] = resolved_biz_dt
+                belong_id = row.get("demand_belong_id")
+                if belong_id is None:
+                    continue
+                stats = stats_by_belong.get(int(belong_id))
+                if stats is None:
+                    continue
+                dim = row["source_dim"]
+                if row.get("dim_count") is None:
+                    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")
+                    )
+
+            inserted = GeneratedDemandRepository(session).bulk_insert(valid_rows)
+            run_ids = sorted({str(r["run_id"]) for r in valid_rows})
+
+        parts = [
+            f"提交有效 {len(valid_rows)} 条,成功插入 {inserted} 条",
+            f"run_id={','.join(run_ids)}",
+            f"biz_dt={resolved_biz_dt}",
+        ]
+        if skipped_names:
+            parts.append(
+                f"因 demand_name 不存在跳过 {len(skipped_names)} 条: "
+                + "、".join(skipped_names[:20])
+                + ("…" if len(skipped_names) > 20 else "")
+            )
+        if errors:
+            parts.append(f"校验失败 {len(errors)} 条: " + ";".join(errors[:20]))
+
+        message = "。".join(parts)
+        logger.info("batch_save_generated_demands completed: %s", message)
+        return message
+
+    except Exception as e:
+        logger.error("batch_save_generated_demands failed: %s", e, exc_info=True)
+        return f"批量保存生成需求失败: {e}"

+ 90 - 0
agents/generate_demand_agent/tools/dim_constants.py

@@ -1,6 +1,7 @@
 """generate_demand_agent 分类树热度维度共享常量与工具函数。"""
 from __future__ import annotations
 
+from collections.abc import Callable
 from decimal import Decimal
 
 from supply_infra.db.models.category_tree_weight import CategoryTreeWeight
@@ -24,6 +25,53 @@ DIM_LABEL: dict[str, str] = {
 HAS_DATA_LEAF_MARK = "+"
 
 
+def normalize_biz_dt(biz_dt: str | None) -> tuple[str | None, str | None]:
+    """校验并规范化 biz_dt(YYYYMMDD);空值表示使用最新业务日。"""
+    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 resolve_biz_dt(
+    biz_dt: str | None,
+    *,
+    get_latest: Callable[[], str | None],
+    has_data: Callable[[str], bool] | None = None,
+    table_label: str = "热度数据",
+) -> tuple[str | None, str | None]:
+    """解析 biz_dt:未传则用最新业务日;传入则校验格式与数据是否存在。"""
+    normalized, err = normalize_biz_dt(biz_dt)
+    if err:
+        return None, err
+
+    resolved = normalized or get_latest()
+    if not resolved:
+        return None, f"暂无 {table_label}(请指定 biz_dt 或先写入数据)"
+
+    if normalized and has_data is not None and not has_data(normalized):
+        return None, f"biz_dt={normalized} 无 {table_label}"
+
+    return resolved, None
+
+
+def get_latest_common_biz_dt(session) -> str | None:
+    """返回 category_tree_weight 与 demand_popularity_stats 均有数据的最新 biz_dt。"""
+    from sqlalchemy import func, select
+
+    from supply_infra.db.models.demand_popularity_stats import DemandPopularityStats
+
+    stats_dates = select(DemandPopularityStats.biz_dt).distinct()
+    stmt = select(func.max(CategoryTreeWeight.biz_dt)).where(
+        CategoryTreeWeight.biz_dt.in_(stats_dates)
+    )
+    return session.scalar(stmt)
+
+
 def normalize_parent_id(parent_id: int | None) -> int | None:
     if parent_id is None or parent_id == 0:
         return None
@@ -94,3 +142,45 @@ def collect_descendant_leaves(
             stack.extend(int(c.id) for c in kids)
     leaves.sort()
     return leaves
+
+
+def collect_descendant_ids(
+    root_id: int,
+    children_map: dict[int | None, list[GlobalTreeCategory]],
+) -> list[int]:
+    """收集 root 及其全部子孙节点 id。"""
+    ids: list[int] = []
+    stack = [root_id]
+    seen: set[int] = set()
+    while stack:
+        cid = stack.pop()
+        if cid in seen:
+            continue
+        seen.add(cid)
+        ids.append(cid)
+        stack.extend(int(c.id) for c in children_map.get(cid, []))
+    ids.sort()
+    return ids
+
+
+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)

+ 17 - 7
agents/generate_demand_agent/tools/query_category_leaves_by_dim.py

@@ -13,6 +13,7 @@ from agents.generate_demand_agent.tools.dim_constants import (
     build_score_by_id,
     collect_descendant_leaves,
     format_score,
+    resolve_biz_dt,
 )
 from supply_agent.tools import tool
 from supply_infra.db.repositories.category_tree_weight_repo import (
@@ -49,6 +50,7 @@ def _normalize_ids(category_ids: list[Any]) -> tuple[list[int], str | None]:
 def query_category_leaves_by_dim(
     dimension: str,
     category_ids: list[int],
+    biz_dt: str | None = None,
 ) -> str:
     """
     按分类 id 与热度维度,查询各分类节点下有维度数据的叶子节点。
@@ -64,6 +66,7 @@ def query_category_leaves_by_dim(
             - plat_ly_pop:平台去年同期热度
             - recent_pop:近期热度
         category_ids: 要查询的分类 id 列表(可多个)。
+        biz_dt: 业务日 YYYYMMDD;省略则使用 category_tree_weight 最新业务日。
 
     Returns:
         按分类分组的叶子列表,例如:
@@ -90,14 +93,21 @@ def query_category_leaves_by_dim(
             children_map = build_children_map(categories)
 
             weight_repo = CategoryTreeWeightRepository(session)
-            biz_dt = weight_repo.get_latest_biz_dt()
-            if not biz_dt:
-                return "暂无 category_tree_weight 数据"
-
-            score_by_id = build_score_by_id(weight_repo.list_by_biz_dt(biz_dt), dim)
+            resolved_biz_dt, err = resolve_biz_dt(
+                biz_dt,
+                get_latest=weight_repo.get_latest_biz_dt,
+                has_data=weight_repo.has_biz_dt,
+                table_label="category_tree_weight 数据",
+            )
+            if err:
+                return err
+
+            score_by_id = build_score_by_id(
+                weight_repo.list_by_biz_dt(resolved_biz_dt), dim
+            )
 
             lines = [
-                f"维度={dim}({DIM_LABEL[dim]})| biz_dt={biz_dt}",
+                f"维度={dim}({DIM_LABEL[dim]})| biz_dt={resolved_biz_dt}",
                 "",
             ]
             total_leaves = 0
@@ -134,7 +144,7 @@ def query_category_leaves_by_dim(
             dim,
             ids,
             total_leaves,
-            biz_dt,
+            resolved_biz_dt,
         )
         return "\n".join(lines).rstrip()
 

+ 82 - 0
agents/generate_demand_agent/tools/query_category_path.py

@@ -0,0 +1,82 @@
+"""
+按分类 id 查询从根到节点的类目路径。
+"""
+from __future__ import annotations
+
+import logging
+from typing import Any
+
+from agents.generate_demand_agent.tools.dim_constants 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:
+        每行一条路径,例如:
+        [256] 美妆护肤 > 护肤 > 防晒 > 物理防晒
+        [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()

+ 18 - 11
agents/generate_demand_agent/tools/query_category_tree_by_dim.py

@@ -17,6 +17,7 @@ from agents.generate_demand_agent.tools.dim_constants import (
     build_score_by_id,
     collect_descendant_leaves,
     format_score,
+    resolve_biz_dt,
 )
 from supply_agent.tools import tool
 from supply_infra.db.models.global_tree_category import GlobalTreeCategory
@@ -117,13 +118,13 @@ def _format_weighted_tree(
 
 
 @tool
-def query_category_tree_by_dim(dimension: str) -> str:
+def query_category_tree_by_dim(dimension: str, biz_dt: str | None = None) -> str:
     """
     按热度维度查询有数据的全局分类树。
 
-    联查 global_tree_category 与 category_tree_weight(最新 biz_dt),
-    仅返回指定维度 count>0 的节点及其祖先。每行格式为 [id]名称(score);
-    若该节点下存在有数据的叶子节点,末尾带 + 标记。
+    联查 global_tree_category 与 category_tree_weight,仅返回指定维度 count>0
+    的节点及其祖先。每行格式为 [id]名称(score);若该节点下存在有数据的叶子节点,
+    末尾带 + 标记。
 
     Args:
         dimension: 必选热度维度,只能是其一:
@@ -131,6 +132,7 @@ def query_category_tree_by_dim(dimension: str) -> str:
             - plat_sust_pop:平台持续热度
             - plat_ly_pop:平台去年同期热度
             - recent_pop:近期热度
+        biz_dt: 业务日 YYYYMMDD;省略则使用 category_tree_weight 最新业务日。
 
     Returns:
         层级分明的完整有数据分类树,例如:
@@ -148,22 +150,27 @@ def query_category_tree_by_dim(dimension: str) -> str:
         with get_session() as session:
             categories = GlobalTreeCategoryRepository(session).list_active_categories()
             weight_repo = CategoryTreeWeightRepository(session)
-            biz_dt = weight_repo.get_latest_biz_dt()
-            if not biz_dt:
-                return "暂无 category_tree_weight 数据"
-
-            weights = weight_repo.list_by_biz_dt(biz_dt)
+            resolved_biz_dt, err = resolve_biz_dt(
+                biz_dt,
+                get_latest=weight_repo.get_latest_biz_dt,
+                has_data=weight_repo.has_biz_dt,
+                table_label="category_tree_weight 数据",
+            )
+            if err:
+                return err
+
+            weights = weight_repo.list_by_biz_dt(resolved_biz_dt)
             score_by_id = build_score_by_id(weights, dim)
             tree_text = _format_weighted_tree(categories, score_by_id)
 
         header = (
-            f"维度={dim}({DIM_LABEL[dim]})| biz_dt={biz_dt} | "
+            f"维度={dim}({DIM_LABEL[dim]})| biz_dt={resolved_biz_dt} | "
             f"「{HAS_DATA_LEAF_MARK}」表示其下有带数据的叶子节点"
         )
         logger.info(
             "query_category_tree_by_dim completed: dim=%s biz_dt=%s",
             dim,
-            biz_dt,
+            resolved_biz_dt,
         )
         return f"{header}\n{tree_text}"
 

+ 246 - 0
agents/generate_demand_agent/tools/query_demand_words_by_category.py

@@ -0,0 +1,246 @@
+"""
+按分类 id + 热度维度,查询可选用的挂载需求词及词级热度。
+"""
+from __future__ import annotations
+
+import logging
+from typing import Any
+
+from agents.generate_demand_agent.tools.dim_constants import (
+    DIM_KEYS,
+    DIM_LABEL,
+    build_children_map,
+    collect_descendant_ids,
+    format_score,
+    resolve_biz_dt,
+)
+from supply_agent.tools import tool
+from supply_infra.db.models.demand_popularity_stats import DemandPopularityStats
+from supply_infra.db.repositories.demand_belong_category_repo import (
+    DemandBelongCategoryRepository,
+)
+from supply_infra.db.repositories.demand_popularity_stats_repo import (
+    DemandPopularityStatsRepository,
+)
+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
+
+
+def _word_dim_score(
+    stats: DemandPopularityStats | None, dim: str
+) -> tuple[float | None, int]:
+    if stats is None:
+        return None, 0
+    count = int(getattr(stats, f"{dim}_count", 0) or 0)
+    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
+
+
+@tool
+def query_demand_words_by_category(
+    category_ids: list[int],
+    dimension: str,
+    include_descendants: bool = True,
+    min_count: int = 1,
+    top_k: int = 30,
+    biz_dt: str | None = None,
+) -> str:
+    """
+    按分类 id 与热度维度,查询挂载需求词及该维词级热度。
+
+    demand_name 必须来自本工具返回的词名(对应 demand_belong_category.name)。
+    include_descendants=True 时包含子树内所有挂载词;False 时仅查节点自身直接挂载。
+    每个入参 category 下最多返回 top_k 个词,按该维 avg 降序。
+
+    Args:
+        category_ids: 分类 id 列表(可多个)。
+        dimension: 热度维度,只能是其一:
+            - ext_pop:外部热度
+            - plat_sust_pop:平台持续热度
+            - plat_ly_pop:平台去年同期热度
+            - recent_pop:近期热度
+        include_descendants: 是否包含子树挂载词,默认 True。
+        min_count: 该维 count 下限,默认 1。
+        top_k: 每个入参 category 最多返回词数,默认 30。
+        biz_dt: 业务日 YYYYMMDD;省略则使用 demand_popularity_stats 最新业务日。
+
+    Returns:
+        按入参分类分组的挂载词列表,含 score / count / belong_id。
+    """
+    dim = str(dimension or "").strip()
+    if dim not in DIM_KEYS:
+        allowed = "、".join(f"{k}({DIM_LABEL[k]})" for k in DIM_KEYS)
+        return f"dimension 无效,只能是:{allowed}"
+
+    ids, err = _normalize_ids(category_ids)
+    if err:
+        return err
+
+    try:
+        min_count_int = max(0, int(min_count))
+    except (TypeError, ValueError):
+        return f"min_count 无效: {min_count!r}"
+    try:
+        top_k_int = max(1, int(top_k))
+    except (TypeError, ValueError):
+        return f"top_k 无效: {top_k!r}"
+
+    try:
+        with get_session() as session:
+            categories = GlobalTreeCategoryRepository(session).list_active_categories()
+            by_id = {int(c.id): c for c in categories}
+            children_map = build_children_map(categories)
+
+            stats_repo = DemandPopularityStatsRepository(session)
+            resolved_biz_dt, err = resolve_biz_dt(
+                biz_dt,
+                get_latest=stats_repo.get_latest_biz_dt,
+                has_data=stats_repo.has_biz_dt,
+                table_label="demand_popularity_stats 数据",
+            )
+            if err:
+                return err
+
+            belong_repo = DemandBelongCategoryRepository(session)
+
+            # 收集本次需要查询的全部挂载类目
+            scope_by_root: dict[int, list[int]] = {}
+            all_scope_ids: set[int] = set()
+            for cid in ids:
+                if include_descendants:
+                    scope = collect_descendant_ids(cid, children_map)
+                else:
+                    scope = [cid]
+                scope_by_root[cid] = scope
+                all_scope_ids.update(scope)
+
+            belongs = belong_repo.list_by_category_ids(sorted(all_scope_ids))
+            belongs_by_cat: dict[int, list] = {}
+            for row in belongs:
+                if not row.name:
+                    continue
+                belongs_by_cat.setdefault(int(row.category_id), []).append(row)
+
+            belong_ids = [int(r.id) for rows in belongs_by_cat.values() for r in rows]
+            stats_rows = stats_repo.list_by_biz_dt_and_belong_ids(
+                resolved_biz_dt, belong_ids
+            )
+            stats_by_belong = {
+                int(s.demand_category_id): s for s in stats_rows
+            }
+
+            lines = [
+                f"维度={dim}({DIM_LABEL[dim]})| biz_dt={resolved_biz_dt}"
+                f" | include_descendants={include_descendants}"
+                f" | min_count={min_count_int} | top_k={top_k_int}",
+                "",
+            ]
+            total_words = 0
+
+            for cid in ids:
+                cat = by_id.get(cid)
+                if cat is None:
+                    lines.append(f"[{cid}]")
+                    lines.append("  (未找到该分类)")
+                    lines.append("")
+                    continue
+
+                cname = cat.name or ""
+                lines.append(f"[{cid}]{cname}")
+
+                candidates: list[tuple[float, int, int, str, int]] = []
+                # (avg, count, belong_id, name, hang_category_id)
+                for scope_cid in scope_by_root[cid]:
+                    for belong in belongs_by_cat.get(scope_cid, []):
+                        avg, count = _word_dim_score(
+                            stats_by_belong.get(int(belong.id)), dim
+                        )
+                        if count < min_count_int:
+                            continue
+                        score = avg if avg is not None else 0.0
+                        candidates.append(
+                            (
+                                score,
+                                count,
+                                int(belong.id),
+                                str(belong.name),
+                                scope_cid,
+                            )
+                        )
+
+                candidates.sort(key=lambda x: (-x[0], -x[1], x[2]))
+                selected = candidates[:top_k_int]
+                if not selected:
+                    lines.append("  (无符合条件的挂载词)")
+                    lines.append("")
+                    continue
+
+                total_words += len(selected)
+                for avg, count, belong_id, name, hang_cid in selected:
+                    hang_note = ""
+                    if include_descendants and hang_cid != cid:
+                        hang_cat = by_id.get(hang_cid)
+                        hang_name = hang_cat.name if hang_cat else ""
+                        hang_note = f" hang=[{hang_cid}]{hang_name}"
+                    lines.append(
+                        f"  - {name}({format_score(avg)}) count={count}"
+                        f"  belong_id={belong_id}{hang_note}"
+                    )
+                if len(candidates) > top_k_int:
+                    lines.append(
+                        f"  (另有 {len(candidates) - top_k_int} 个词未展示,已按 avg 截断)"
+                    )
+                lines.append("")
+
+        logger.info(
+            "query_demand_words_by_category completed: dim=%s ids=%s words=%d biz_dt=%s",
+            dim,
+            ids,
+            total_words,
+            resolved_biz_dt,
+        )
+        return "\n".join(lines).rstrip()
+
+    except Exception as e:
+        logger.error("query_demand_words_by_category failed: %s", e, exc_info=True)
+        return f"查询分类挂载词失败: {e}"
+
+
+def main() -> None:
+    """本地手动验证。"""
+    print(
+        query_demand_words_by_category(
+            category_ids=[1, 2],
+            dimension="recent_pop",
+            top_k=10,
+        )
+    )
+
+
+if __name__ == "__main__":
+    main()

+ 54 - 0
agents/generate_demand_agent/tools/query_latest_biz_dt.py

@@ -0,0 +1,54 @@
+"""
+查询热度相关表的最新共有业务日,供未指定 biz_dt 时先探测可用日期。
+"""
+from __future__ import annotations
+
+import logging
+
+from agents.generate_demand_agent.tools.dim_constants import get_latest_common_biz_dt
+from supply_agent.tools import tool
+from supply_infra.db.session import get_session
+
+logger = logging.getLogger(__name__)
+
+
+@tool
+def query_latest_biz_dt() -> str:
+    """
+    查询两表均有数据的最新业务日(biz_dt)。
+
+    当用户未指定 biz_dt、或不确定该用哪一天时,**应先调用本工具**,
+    再将其返回的 biz_dt 传给后续查询与落库工具。
+
+    共有业务日 = category_tree_weight 与 demand_popularity_stats 都存在的最新日期。
+
+    Returns:
+        一个确切的 biz_dt,例如:
+        biz_dt=20260715
+        若无共有日期则返回错误说明。
+    """
+    try:
+        with get_session() as session:
+            biz_dt = get_latest_common_biz_dt(session)
+
+        if not biz_dt:
+            message = "暂无两表共有业务日(category_tree_weight 与 demand_popularity_stats 无重叠日期)"
+            logger.info("query_latest_biz_dt completed: %s", message)
+            return message
+
+        message = f"biz_dt={biz_dt}"
+        logger.info("query_latest_biz_dt completed: %s", message)
+        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()

+ 1 - 0
pyproject.toml

@@ -21,6 +21,7 @@ dependencies = [
     "oss2>=2.18.0",
     "fastapi>=0.115.0",
     "uvicorn[standard]>=0.32.0",
+    "markdown>=3.6",
 ]
 
 [project.optional-dependencies]

+ 1 - 0
requirements.txt

@@ -6,6 +6,7 @@ pydantic-settings>=2.0
 httpx>=0.27.0
 rich>=13.0
 python-dotenv>=1.0.0
+markdown>=3.6
 
 # Database
 sqlalchemy>=2.0

+ 292 - 268
supply_agent/logging/visualize.py

@@ -7,8 +7,23 @@ import json
 from pathlib import Path
 from typing import Any
 
+import markdown as _markdown_lib
+
 from supply_agent.logging.parser import load_run_events, summarize_run
 
+_MARKDOWN_EXTENSIONS = ["fenced_code", "tables", "sane_lists"]
+
+
+def _render_markdown(text: Any, *, soft_breaks: bool = False, boxed: bool = True) -> str:
+    """Render text as Markdown; ``boxed`` wraps it in a bordered, scrollable box."""
+    if not text:
+        return ""
+    extensions = [*_MARKDOWN_EXTENSIONS, "nl2br"] if soft_breaks else _MARKDOWN_EXTENSIONS
+    body = _markdown_lib.markdown(str(text), extensions=extensions)
+    if not boxed:
+        return f'<div class="markdown-body">{body}</div>'
+    return f'<div class="markdown-box"><div class="markdown-body">{body}</div></div>'
+
 
 def _esc(text: Any) -> str:
     if text is None:
@@ -104,50 +119,46 @@ def _role_badge(role: str) -> str:
 def _render_messages(
     messages: list[dict[str, Any]] | None,
     *,
-    index_offset: int = 0,
     default_open: bool | None = None,
 ) -> str:
     if not messages:
         return '<p class="muted">无消息</p>'
     parts: list[str] = []
-    for i, msg in enumerate(messages):
+    for msg in messages:
         role = msg.get("role", "?")
         body_parts: list[str] = []
         if msg.get("content"):
             body_parts.append(_pre(msg["content"], klass="code prose"))
-        if msg.get("tool_calls"):
-            body_parts.append(
-                '<div class="sublabel">决定调用的工具</div>'
-                + _render_tool_call_cards(msg["tool_calls"])
-            )
-        if msg.get("tool_call_id"):
-            body_parts.append(
-                f'<div class="meta-line">tool_call_id: <code>{_esc(msg["tool_call_id"])}</code></div>'
-            )
-        if msg.get("name"):
-            body_parts.append(
-                f'<div class="meta-line">name: <code>{_esc(msg["name"])}</code></div>'
-            )
-        # Tool result content: try pretty JSON (already handled by _pre via content)        if default_open is None:
+
+        if role in ("user", "tool"):
+            open_attr = ""
+        elif default_open is None:
             open_attr = "open" if role != "system" else ""
         else:
             open_attr = "open" if default_open else ""
-        preview = _preview(msg.get("content"))
-        if not preview and msg.get("tool_calls"):
-            names = []
-            for tc in msg["tool_calls"]:
-                name, _, _ = _normalize_tool_call(tc)
-                names.append(name)
-            preview = "调用: " + ", ".join(names)
-        elif not preview and msg.get("name"):
-            preview = f"tool → {msg['name']}"
+
+        if role == "tool":
+            tool_name = msg.get("name") or "?"
+            summary_html = (
+                f'{_role_badge(role)} <span class="msg-preview">{_esc(tool_name)}</span>'
+            )
+        else:
+            preview = _preview(msg.get("content"))
+            if not preview and msg.get("tool_calls"):
+                names = []
+                for tc in msg["tool_calls"]:
+                    name, _, _ = _normalize_tool_call(tc)
+                    names.append(name)
+                preview = "调用: " + ", ".join(names)
+            summary_html = (
+                f'{_role_badge(role)} <span class="msg-preview">{_esc(preview)}</span>'
+            )
+
         parts.append(
             f"""
             <details class="msg" {open_attr}>
-              <summary>{_role_badge(role)} <span class="msg-idx">#{index_offset + i + 1}</span>
-                <span class="msg-preview">{_esc(preview)}</span>
-              </summary>
-              <div class="msg-body">{"".join(body_parts) or '<p class="muted">(empty)</p>'}</div>
+              <summary>{summary_html}</summary>
+              <div class="msg-body">{"".join(body_parts)}</div>
             </details>
             """
         )
@@ -161,68 +172,53 @@ def _preview(text: Any, limit: int = 80) -> str:
     return s if len(s) <= limit else s[: limit - 1] + "…"
 
 
-def _render_tools_schema(tools: list[dict[str, Any]] | None) -> str:
+def _tool_names(tools: list[dict[str, Any]] | None) -> list[str]:
     if not tools:
-        return '<p class="muted">本次请求未附带 tools</p>'
-    names = []
+        return []
+    names: list[str] = []
     for t in tools:
         fn = t.get("function") or t
         names.append(fn.get("name", "?"))
-    chips = "".join(f'<span class="chip">{_esc(n)}</span>' for n in names)
-    return f"""
-    <div class="chip-row">{chips}</div>
-    <details>
-      <summary>查看完整 tools schema</summary>
-      {_pre(tools)}
-    </details>
-    """
+    return names
+
 
+def _extract_system_prompt_and_tools(
+    events: list[dict[str, Any]],
+) -> tuple[str | None, list[str]]:
+    """Pull the system prompt and tool list from the first llm_input event."""
+    for ev in events:
+        if ev.get("event") != "llm_input":
+            continue
+        data = ev.get("data") or {}
+        messages = data.get("messages") or []
+        sys_msg = next((m for m in messages if m.get("role") == "system"), None)
+        content = sys_msg.get("content") if sys_msg else None
+        return content, _tool_names(data.get("tools"))
+    return None, []
 
-def _render_llm_input(
+
+def _render_llm_input_section(
     data: dict[str, Any],
-    seq: int,
-    iteration: Any,
     *,
     prev_messages: list[dict[str, Any]] | None = None,
-) -> str:
+) -> tuple[str, list[dict[str, Any]]]:
+    """Render the LLM-input body; returns (html, messages) so callers can track history."""
     messages = list(data.get("messages") or [])
-    history, latest = _split_history_and_latest(messages, prev_messages)
-    history_count = len(history)
+    _history, latest = _split_history_and_latest(messages, prev_messages)
     latest_count = len(latest)
 
-    history_html = ""
-    if history:
-        history_html = f"""
-        <details class="history-block">
-          <summary>查看历史输入({history_count} 条)</summary>
-          <div class="history-body">
-            {_render_messages(history, default_open=False)}
-          </div>
-        </details>
-        """
-
-    return f"""
-    <article class="card card-input" id="step-{seq}">
-      <header class="card-header">
-        <div class="step-num">Step {seq}</div>
-        <div class="card-title">LLM 输入</div>
-        <div class="card-tags">
-          <span class="tag">iteration {iteration}</span>
-          <span class="tag">{_esc(data.get("model", ""))}</span>
-          <span class="tag">temp {_esc(data.get("temperature", ""))}</span>
-          <span class="tag">新增 {latest_count}</span>
-          <span class="tag">共 {len(messages)} msgs</span>
-        </div>
-      </header>
-      <div class="card-body">
-        <h4>本轮新增输入</h4>
-        {_render_messages(latest, index_offset=history_count, default_open=True)}
-        {history_html}
-        <h4>Available Tools</h4>
-        {_render_tools_schema(data.get("tools"))}
+    html_out = f"""
+    <section class="substep">
+      <div class="substep-header">
+        <span class="substep-title">LLM 输入</span>
+        <span class="tag">新增 {latest_count}</span>
       </div>
-    </article>
+      <div class="substep-body">
+        {_render_messages(latest, default_open=True)}
+      </div>
+    </section>
     """
+    return html_out, messages
 
 
 def _render_reasoning(reasoning: Any) -> str:
@@ -231,7 +227,7 @@ def _render_reasoning(reasoning: Any) -> str:
     return f"""
     <div class="reasoning">
       <div class="sublabel">思考过程</div>
-      {_pre(reasoning, klass="code prose reasoning-text")}
+      {_render_markdown(reasoning, soft_breaks=True, boxed=False)}
     </div>
     """
 
@@ -254,48 +250,25 @@ def _normalize_tool_call(tc: dict[str, Any]) -> tuple[str, Any, str]:
     return name, _parse_tool_arguments(raw_args), str(tc.get("id") or "")
 
 
-def _render_tool_call_cards(tool_calls: list[dict[str, Any]]) -> str:
-    """Render tool calls as name + args cards; raw JSON in a collapsible."""
+def _render_output_tool_calls(tool_calls: list[dict[str, Any]]) -> str:
+    """Show each tool the model decided to call, collapsed by default; expand for args."""
     if not tool_calls:
         return ""
-
-    cards: list[str] = []
-    for i, tc in enumerate(tool_calls, 1):
-        name, args, tc_id = _normalize_tool_call(tc)
-        cards.append(
+    parts: list[str] = []
+    for tc in tool_calls:
+        name, args, _ = _normalize_tool_call(tc)
+        parts.append(
             f"""
-            <div class="planned-tool">
-              <div class="planned-tool-head">
-                <span class="chip">{_esc(name)}</span>
-                <span class="planned-tool-idx">#{i}</span>
-                {f'<code class="planned-tool-id">{_esc(tc_id)}</code>' if tc_id else ""}
-              </div>
-              <div class="sublabel">参数</div>
-              {_pre(args)}
-            </div>
+            <details class="msg">
+              <summary><span class="badge role-tool">tool</span> <span class="msg-preview">{_esc(name)}</span></summary>
+              <div class="msg-body">{_pre(args)}</div>
+            </details>
             """
         )
-
-    return f"""
-    <div class="planned-tool-list">{"".join(cards)}</div>
-    <details class="raw-block">
-      <summary>查看原始 tool_calls JSON</summary>
-      {_pre(tool_calls)}
-    </details>
-    """
-
-
-def _render_planned_tool_calls(tool_calls: list[dict[str, Any]]) -> str:
-    """Show each planned tool as name + args; omit entirely when empty."""
-    if not tool_calls:
-        return ""
-    return f"""
-    <h4>模型决定调用的工具</h4>
-    {_render_tool_call_cards(tool_calls)}
-    """
+    return '<div class="msg-list">' + "".join(parts) + "</div>"
 
 
-def _render_llm_output(data: dict[str, Any], seq: int, iteration: Any) -> str:
+def _render_llm_output_section(data: dict[str, Any]) -> str:
     tool_calls = data.get("tool_calls") or []
     reasoning = data.get("reasoning")
     content = data.get("content")
@@ -313,82 +286,60 @@ def _render_llm_output(data: dict[str, Any], seq: int, iteration: Any) -> str:
         """
 
     content_html = (
-        f"<h4>模型输出文本</h4>{_pre(content, klass='code prose')}" if content else ""
+        f"<h4>模型输出文本</h4>{_render_markdown(content, soft_breaks=True)}"
+        if content
+        else ""
     )
 
-    tags: list[str] = [f'<span class="tag">iteration {iteration}</span>']
+    tags: list[str] = []
     if reasoning:
         tags.append('<span class="tag">有思考</span>')
     if tool_calls:
         tags.append(f'<span class="tag">{len(tool_calls)} tool calls</span>')
 
-    raw_html = ""
-    if data.get("raw"):
-        raw_html = f"""
-        <details class="raw-block">
-          <summary>原始 API 响应 (raw)</summary>
-          {_pre(data["raw"])}
-        </details>
-        """
-
     return f"""
-    <article class="card card-output" id="step-{seq}">
-      <header class="card-header">
-        <div class="step-num">Step {seq}</div>
-        <div class="card-title">LLM 输出</div>
-        <div class="card-tags">
-          {"".join(tags)}
-        </div>
-      </header>
-      <div class="card-body">
+    <section class="substep">
+      <div class="substep-header">
+        <span class="substep-title">LLM 输出</span>
+        {"".join(tags)}
+      </div>
+      <div class="substep-body">
         {_render_reasoning(reasoning)}
         {content_html}
-        {_render_planned_tool_calls(tool_calls)}
+        {_render_output_tool_calls(tool_calls)}
         {usage_html}
-        {raw_html}
       </div>
-    </article>
+    </section>
     """
 
 
-def _render_tool_call(data: dict[str, Any], seq: int, iteration: Any) -> str:
-    is_error = bool(data.get("is_error"))
-    status = "error" if is_error else "ok"
-    args = data.get("arguments_parsed", data.get("arguments"))
-    result = data.get("result_parsed", data.get("result"))
-    tool_id = data.get("tool_call_id") or ""
-
+def _render_step_card(
+    step_no: int,
+    iteration: Any,
+    title: str,
+    inner_html: str,
+) -> str:
     return f"""
-    <article class="card card-tool {status}" id="step-{seq}">
+    <article class="card card-step" id="step-{step_no}">
       <header class="card-header">
-        <div class="step-num">Step {seq}</div>
-        <div class="card-title">工具调用 · {_esc(data.get("tool", "?"))}</div>
+        <div class="step-num">Step {step_no}</div>
+        <div class="card-title">{title}</div>
         <div class="card-tags">
           <span class="tag">iteration {iteration}</span>
-          <span class="tag tag-{status}">{"ERROR" if is_error else "OK"}</span>
-          {f'<span class="tag"><code>{_esc(tool_id)}</code></span>' if tool_id else ""}
         </div>
       </header>
-      <div class="card-body tool-io">
-        <div class="io-block">
-          <div class="sublabel">输入 (arguments)</div>
-          {_pre(args)}
-        </div>
-        <div class="io-arrow" aria-hidden="true">→</div>
-        <div class="io-block">
-          <div class="sublabel">输出 (result)</div>
-          {_pre(result)}
-        </div>
+      <div class="card-body step-body">
+        {inner_html}
       </div>
     </article>
     """
 
 
-def _render_skill(data: dict[str, Any], seq: int, iteration: Any) -> str:
+def _render_skill(data: dict[str, Any], step_no: int, iteration: Any) -> str:
     return f"""
-    <article class="card card-skill" id="step-{seq}">
+    <article class="card card-skill" id="step-{step_no}">
       <header class="card-header">
-        <div class="step-num">Step {seq}</div>
+        <div class="step-num">Step {step_no}</div>
         <div class="card-title">技能加载</div>
         <div class="card-tags">
           <span class="tag">iteration {iteration}</span>
@@ -399,42 +350,87 @@ def _render_skill(data: dict[str, Any], seq: int, iteration: Any) -> str:
     """
 
 
+_HIDDEN_EVENT_TYPES = ("run_start", "run_end", "tool_call")
+
+
+def _group_visible_events(
+    events: list[dict[str, Any]],
+) -> list[list[dict[str, Any]]]:
+    """
+    Group renderable events into steps.
+
+    A consecutive ``llm_input`` immediately followed by ``llm_output`` forms a
+    single step (one round-trip); everything else is its own step.
+    """
+    visible = [ev for ev in events if ev.get("event") not in _HIDDEN_EVENT_TYPES]
+    groups: list[list[dict[str, Any]]] = []
+    i = 0
+    n = len(visible)
+    while i < n:
+        ev = visible[i]
+        if (
+            ev.get("event") == "llm_input"
+            and i + 1 < n
+            and visible[i + 1].get("event") == "llm_output"
+        ):
+            groups.append([ev, visible[i + 1]])
+            i += 2
+        else:
+            groups.append([ev])
+            i += 1
+    return groups
+
+
+def _event_iteration(ev: dict[str, Any]) -> Any:
+    data = ev.get("data") or {}
+    return ev.get("iteration", data.get("iteration", "—"))
+
+
 def _render_timeline(events: list[dict[str, Any]]) -> str:
     parts: list[str] = []
     prev_llm_messages: list[dict[str, Any]] | None = None
-    for ev in events:
+    for step_no, group in enumerate(_group_visible_events(events), start=1):
+        if len(group) == 2:
+            input_ev, output_ev = group
+            input_data = input_ev.get("data") or {}
+            output_data = output_ev.get("data") or {}
+            input_html, messages = _render_llm_input_section(
+                input_data, prev_messages=prev_llm_messages
+            )
+            prev_llm_messages = messages
+            output_html = _render_llm_output_section(output_data)
+            parts.append(
+                _render_step_card(
+                    step_no,
+                    _event_iteration(input_ev),
+                    "LLM 输入 · 输出",
+                    input_html + output_html,
+                )
+            )
+            continue
+
+        ev = group[0]
         etype = ev.get("event")
         data = ev.get("data") or {}
-        seq = ev.get("seq", 0)
-        iteration = ev.get("iteration", data.get("iteration", "—"))
+        iteration = _event_iteration(ev)
 
-        if etype == "run_start":
-            continue
-        if etype == "run_end":
-            continue
         if etype == "llm_input":
-            messages = list(data.get("messages") or [])
-            parts.append(
-                _render_llm_input(
-                    data,
-                    seq,
-                    iteration,
-                    prev_messages=prev_llm_messages,
-                )
+            input_html, messages = _render_llm_input_section(
+                data, prev_messages=prev_llm_messages
             )
             prev_llm_messages = messages
+            parts.append(_render_step_card(step_no, iteration, "LLM 输入", input_html))
         elif etype == "llm_output":
-            parts.append(_render_llm_output(data, seq, iteration))
-        elif etype == "tool_call":
-            parts.append(_render_tool_call(data, seq, iteration))
+            output_html = _render_llm_output_section(data)
+            parts.append(_render_step_card(step_no, iteration, "LLM 输出", output_html))
         elif etype == "skill_loaded":
-            parts.append(_render_skill(data, seq, iteration))
+            parts.append(_render_skill(data, step_no, iteration))
         else:
             parts.append(
                 f"""
-                <article class="card" id="step-{seq}">
+                <article class="card" id="step-{step_no}">
                   <header class="card-header">
-                    <div class="step-num">Step {seq}</div>
+                    <div class="step-num">Step {step_no}</div>
                     <div class="card-title">{_esc(etype)}</div>
                   </header>
                   <div class="card-body">{_pre(data)}</div>
@@ -449,24 +445,25 @@ def _nav_items(events: list[dict[str, Any]]) -> str:
     labels = {
         "llm_input": "LLM 输入",
         "llm_output": "LLM 输出",
-        "tool_call": "工具",
         "skill_loaded": "技能",
     }
-    for ev in events:
-        etype = ev.get("event")
-        if etype in ("run_start", "run_end"):
-            continue
-        data = ev.get("data") or {}
-        seq = ev.get("seq", 0)
-        label = labels.get(etype, etype or "?")
-        detail = ""
-        if etype == "tool_call":
-            detail = f" · {_esc(data.get('tool', ''))}"
-        elif etype in ("llm_input", "llm_output"):
-            detail = f" · iter {ev.get('iteration', data.get('iteration', ''))}"
+    for step_no, group in enumerate(_group_visible_events(events), start=1):
+        if len(group) == 2:
+            ev = group[0]
+            label = "LLM 输入 · 输出"
+            nav_class = "nav-step-pair"
+            detail = f" · iter {_event_iteration(ev)}"
+        else:
+            ev = group[0]
+            etype = ev.get("event")
+            label = labels.get(etype, etype or "?")
+            nav_class = f"nav-{_esc(etype or '')}"
+            detail = ""
+            if etype in ("llm_input", "llm_output"):
+                detail = f" · iter {_event_iteration(ev)}"
         items.append(
-            f'<a class="nav-item nav-{_esc(etype or "")}" href="#step-{seq}">'
-            f'<span class="nav-seq">{seq}</span>{label}{detail}</a>'
+            f'<a class="nav-item {nav_class}" href="#step-{step_no}">'
+            f'<span class="nav-seq">{step_no}</span>{label}{detail}</a>'
         )
     return "\n".join(items)
 
@@ -482,9 +479,6 @@ _CSS = """
   --accent: #2563eb;
   --input: #2563eb;
   --output: #059669;
-  --tool: #d97706;
-  --tool-ok: #059669;
-  --tool-err: #dc2626;
   --skill: #7c3aed;
   --code-bg: #f8fafc;
   --mono: "JetBrains Mono", "SF Mono", "Fira Code", ui-monospace, monospace;
@@ -548,7 +542,7 @@ a:hover { text-decoration: underline; }
 }
 .nav-llm_input { border-left: 3px solid var(--input); }
 .nav-llm_output { border-left: 3px solid var(--output); }
-.nav-tool_call { border-left: 3px solid var(--tool); }
+.nav-step-pair { border-left: 3px solid var(--input); }
 .nav-skill_loaded { border-left: 3px solid var(--skill); }
 
 .main { padding: 1.5rem 2rem 3rem; max-width: 1100px; }
@@ -601,10 +595,7 @@ a:hover { text-decoration: underline; }
   overflow: hidden;
   box-shadow: 0 1px 2px rgba(26, 35, 50, 0.04);
 }
-.card-input { border-top: 3px solid var(--input); }
-.card-output { border-top: 3px solid var(--output); }
-.card-tool { border-top: 3px solid var(--tool); }
-.card-tool.error { border-top-color: var(--tool-err); }
+.card-step { border-top: 3px solid var(--input); }
 .card-skill { border-top: 3px solid var(--skill); }
 .card-header {
   display: flex;
@@ -634,8 +625,6 @@ a:hover { text-decoration: underline; }
   padding: 0.12rem 0.5rem;
   color: var(--muted);
 }
-.tag-ok { color: var(--tool-ok); border-color: #86efac; background: #f0fdf4; }
-.tag-error { color: var(--tool-err); border-color: #fecaca; background: #fef2f2; }
 .card-body { padding: 1rem; }
 .card-body h4 {
   margin: 1rem 0 0.5rem;
@@ -645,6 +634,23 @@ a:hover { text-decoration: underline; }
   letter-spacing: 0.05em;
 }
 .card-body h4:first-child { margin-top: 0; }
+.card-body.step-body { padding: 0; }
+.substep { padding: 1rem; }
+.substep + .substep { border-top: 1px solid var(--border); }
+.substep-header {
+  display: flex;
+  flex-wrap: wrap;
+  align-items: center;
+  gap: 0.5rem;
+  margin-bottom: 0.65rem;
+}
+.substep-title {
+  font-weight: 600;
+  font-size: 0.8rem;
+  color: var(--muted);
+  text-transform: uppercase;
+  letter-spacing: 0.05em;
+}
 .sublabel {
   font-size: 0.72rem;
   color: var(--muted);
@@ -675,20 +681,6 @@ a:hover { text-decoration: underline; }
   padding: 0.75rem;
   margin-bottom: 1rem;
 }
-.reasoning.empty { opacity: 0.75; }
-.reasoning-text { max-height: 360px; border-color: #fcd34d; background: #fffef5; }
-.tool-io {
-  display: grid;
-  grid-template-columns: 1fr auto 1fr;
-  gap: 0.75rem;
-  align-items: start;
-}
-.io-arrow {
-  color: var(--muted);
-  font-size: 1.4rem;
-  padding-top: 1.6rem;
-}
-.io-block { min-width: 0; }
 .msg-list { display: flex; flex-direction: column; gap: 0.4rem; }
 .msg {
   border: 1px solid var(--border);
@@ -705,7 +697,6 @@ a:hover { text-decoration: underline; }
 }
 .msg summary::-webkit-details-marker { display: none; }
 .msg-body { padding: 0 0.65rem 0.65rem; }
-.msg-idx { font-family: var(--mono); font-size: 0.7rem; color: var(--muted); }
 .msg-preview { color: var(--muted); font-size: 0.78rem; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; flex: 1; }
 .badge {
   font-size: 0.65rem;
@@ -730,35 +721,6 @@ a:hover { text-decoration: underline; }
   padding: 0.15rem 0.5rem;
   border-radius: 999px;
 }
-.planned-tool-list {
-  display: flex;
-  flex-direction: column;
-  gap: 0.65rem;
-  margin-bottom: 0.5rem;
-}
-.planned-tool {
-  border: 1px solid #fed7aa;
-  background: #fffbeb;
-  border-radius: 8px;
-  padding: 0.65rem 0.75rem;
-}
-.planned-tool-head {
-  display: flex;
-  align-items: center;
-  gap: 0.5rem;
-  margin-bottom: 0.45rem;
-  flex-wrap: wrap;
-}
-.planned-tool-idx {
-  font-family: var(--mono);
-  font-size: 0.7rem;
-  color: var(--muted);
-}
-.planned-tool-id {
-  font-size: 0.68rem;
-  color: var(--muted);
-  margin-left: auto;
-}
 .usage {
   display: flex;
   flex-wrap: wrap;
@@ -768,39 +730,99 @@ a:hover { text-decoration: underline; }
   font-size: 0.72rem;
   color: var(--muted);
 }
-.raw-block { margin-top: 0.75rem; }
-.history-block {
-  margin: 0.85rem 0 0.5rem;
+.sysprompt-block { margin: 0.85rem 0; }
+.markdown-box {
+  background: var(--code-bg);
   border: 1px solid var(--border);
-  border-radius: 8px;
-  background: #f8fafc;
-  padding: 0.35rem 0.65rem 0.65rem;
+  border-radius: 6px;
+  padding: 0.75rem 1rem;
+  max-height: 420px;
+  overflow: auto;
 }
-.history-block > summary {
-  cursor: pointer;
-  font-size: 0.85rem;
+.markdown-body { font-size: 0.88rem; line-height: 1.65; color: var(--text); }
+.markdown-body > *:first-child { margin-top: 0; }
+.markdown-body > *:last-child { margin-bottom: 0; }
+.markdown-body p { margin: 0.5rem 0; }
+.markdown-body h1,
+.markdown-body h2,
+.markdown-body h3,
+.markdown-body h4,
+.markdown-body h5,
+.markdown-body h6 {
+  margin: 1rem 0 0.5rem;
+  line-height: 1.35;
+  color: var(--text);
+}
+.markdown-body h1 { font-size: 1.15rem; }
+.markdown-body h2 { font-size: 1.05rem; }
+.markdown-body h3 { font-size: 0.98rem; }
+.markdown-body h4 { font-size: 0.9rem; }
+.markdown-body ul,
+.markdown-body ol { margin: 0.4rem 0; padding-left: 1.4rem; }
+.markdown-body li { margin: 0.2rem 0; }
+.markdown-body li > p { margin: 0.2rem 0; }
+.markdown-body code {
+  background: #eef2f7;
+  border-radius: 4px;
+  padding: 0.1rem 0.35rem;
+}
+.markdown-body pre {
+  background: #0f172a;
+  color: #e2e8f0;
+  border-radius: 6px;
+  padding: 0.65rem 0.85rem;
+  overflow: auto;
+  margin: 0.5rem 0;
+}
+.markdown-body pre code { background: transparent; padding: 0; color: inherit; }
+.markdown-body blockquote {
+  border-left: 3px solid var(--border);
+  margin: 0.5rem 0;
+  padding: 0.1rem 0.75rem;
   color: var(--muted);
-  padding: 0.35rem 0;
-  font-weight: 500;
 }
-.history-body { margin-top: 0.35rem; }
-.meta-line { font-size: 0.8rem; color: var(--muted); margin-bottom: 0.35rem; }
+.markdown-body table { border-collapse: collapse; margin: 0.5rem 0; font-size: 0.82rem; }
+.markdown-body th,
+.markdown-body td { border: 1px solid var(--border); padding: 0.3rem 0.55rem; }
+.markdown-body a { text-decoration: underline; }
+.markdown-body strong { font-weight: 600; }
+.markdown-body hr { border: none; border-top: 1px solid var(--border); margin: 0.75rem 0; }
 code { font-family: var(--mono); font-size: 0.85em; }
 
 @media (max-width: 900px) {
   .layout { grid-template-columns: 1fr; }
   .sidebar { position: relative; height: auto; max-height: 40vh; }
-  .tool-io { grid-template-columns: 1fr; }
-  .io-arrow { display: none; }
 }
 """
 
 
+def _render_system_prompt_block(system_prompt: str | None) -> str:
+    if not system_prompt:
+        return ""
+    return f"""
+    <div class="sysprompt-block">
+      <div class="sublabel">System Prompt</div>
+      {_render_markdown(system_prompt)}
+    </div>
+    """
+
+
+def _render_available_tools_block(tool_names: list[str]) -> str:
+    if not tool_names:
+        return ""
+    chips = "".join(f'<span class="chip">{_esc(n)}</span>' for n in tool_names)
+    return f"""
+    <div class="sublabel">可用工具({len(tool_names)})</div>
+    <div class="chip-row">{chips}</div>
+    """
+
+
 def render_html(events: list[dict[str, Any]]) -> str:
     """Render a full standalone HTML page for the given events."""
     meta = summarize_run(events)
     skills = meta.get("skills_used") or []
     skills_str = ", ".join(skills) if skills else "—"
+    system_prompt, tool_names = _extract_system_prompt_and_tools(events)
 
     return f"""<!DOCTYPE html>
 <html lang="zh-CN">
@@ -831,19 +853,21 @@ def render_html(events: list[dict[str, Any]]) -> str:
           <div class="stat"><div class="label">Events</div><div class="value">{_esc(meta.get("event_count"))}</div></div>
           <div class="stat"><div class="label">Skills</div><div class="value">{_esc(skills_str)}</div></div>
         </div>
+        {_render_available_tools_block(tool_names)}
+        {_render_system_prompt_block(system_prompt)}
         <div class="sublabel">用户输入</div>
         <div class="user-prompt">{_esc(meta.get("user_input") or "—")}</div>
         {f'''
         <div class="final">
           <h3>最终回答</h3>
-          <div class="user-prompt" style="background:transparent;border:none;padding:0">{_esc(meta.get("final_content") or "")}</div>
+          {_render_markdown(meta.get("final_content"), soft_breaks=True, boxed=False)}
         </div>
         ''' if meta.get("final_content") else ""}
       </section>
 
       <section class="timeline">
         <h2 style="margin:0 0 0.5rem;font-size:1.1rem">执行时间线</h2>
-        <p class="muted" style="margin:0 0 1rem">按步骤展示:LLM 输入 → 思考/输出 → 工具调用(含完整入参与返回)</p>
+        <p class="muted" style="margin:0 0 1rem">按步骤展示:LLM 输入 → 思考/输出</p>
         {_render_timeline(events)}
       </section>
     </main>

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

@@ -3,6 +3,7 @@
 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_popularity_stats import DemandPopularityStats
+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
@@ -12,6 +13,7 @@ __all__ = [
     "CategoryTreeWeight",
     "DemandBelongCategory",
     "DemandPopularityStats",
+    "GeneratedDemand",
     "GlobalTreeCategory",
     "GlobalTreeElement",
     "MultiDemandPoolDi",

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

@@ -10,6 +10,7 @@ from supply_infra.db.repositories.demand_belong_category_repo import (
 from supply_infra.db.repositories.demand_popularity_stats_repo import (
     DemandPopularityStatsRepository,
 )
+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
 from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
@@ -20,6 +21,7 @@ __all__ = [
     "CategoryTreeWeightRepository",
     "DemandBelongCategoryRepository",
     "DemandPopularityStatsRepository",
+    "GeneratedDemandRepository",
     "GlobalTreeCategoryRepository",
     "GlobalTreeElementRepository",
     "MultiDemandPoolDiRepository",

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

@@ -32,6 +32,13 @@ class CategoryTreeWeightRepository(BaseRepository[CategoryTreeWeight]):
         stmt = select(func.max(CategoryTreeWeight.biz_dt))
         return self.session.scalar(stmt)
 
+    def has_biz_dt(self, biz_dt: str) -> bool:
+        """指定业务日是否存在权重数据。"""
+        stmt = select(CategoryTreeWeight.id).where(
+            CategoryTreeWeight.biz_dt == biz_dt
+        ).limit(1)
+        return self.session.scalars(stmt).first() is not None
+
     def list_by_biz_dt(self, biz_dt: str) -> list[CategoryTreeWeight]:
         """返回指定业务日的全部节点权重行。"""
         stmt = select(CategoryTreeWeight).where(CategoryTreeWeight.biz_dt == biz_dt)

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

@@ -40,6 +40,22 @@ class DemandBelongCategoryRepository(BaseRepository[DemandBelongCategory]):
         )
         return list(self.session.scalars(stmt).all())
 
+    def list_by_category_ids(
+        self, category_ids: list[int]
+    ) -> list[DemandBelongCategory]:
+        """返回挂载在指定类目上的未删除需求词。"""
+        if not category_ids:
+            return []
+        stmt = (
+            select(DemandBelongCategory)
+            .where(
+                DemandBelongCategory.is_delete == 0,
+                DemandBelongCategory.category_id.in_(category_ids),
+            )
+            .order_by(DemandBelongCategory.category_id, DemandBelongCategory.id)
+        )
+        return list(self.session.scalars(stmt).all())
+
     def list_active_id_name(self) -> list[tuple[int, str]]:
         """返回所有未删除记录的 (id, name),跳过空名称。"""
         stmt = (
@@ -53,6 +69,24 @@ class DemandBelongCategoryRepository(BaseRepository[DemandBelongCategory]):
         )
         return [(int(row_id), name) for row_id, name in self.session.execute(stmt).all() if name]
 
+    def get_by_names(self, names: Iterable[str]) -> dict[str, DemandBelongCategory]:
+        """按 name 批量查询未删除记录,返回 name → row。"""
+        name_list = [n for n in names if n]
+        if not name_list:
+            return {}
+
+        result: dict[str, DemandBelongCategory] = {}
+        for i in range(0, len(name_list), _BATCH_SIZE):
+            batch = name_list[i : i + _BATCH_SIZE]
+            stmt = select(DemandBelongCategory).where(
+                DemandBelongCategory.is_delete == 0,
+                DemandBelongCategory.name.in_(batch),
+            )
+            for row in self.session.scalars(stmt).all():
+                if row.name:
+                    result[row.name] = row
+        return result
+
     def bulk_insert_ignore(self, rows: list[dict]) -> int:
         """批量插入,MySQL 按 name 唯一索引忽略已存在行。"""
         if not rows:

+ 29 - 1
supply_infra/db/repositories/demand_popularity_stats_repo.py

@@ -1,6 +1,6 @@
 from __future__ import annotations
 
-from sqlalchemy import select
+from sqlalchemy import func, select
 from sqlalchemy.dialects.mysql import insert
 
 from supply_infra.db.models.demand_popularity_stats import DemandPopularityStats
@@ -30,11 +30,39 @@ class DemandPopularityStatsRepository(BaseRepository[DemandPopularityStats]):
 
     model = DemandPopularityStats
 
+    def get_latest_biz_dt(self) -> str | None:
+        """返回热度统计表中最新业务日;无数据时返回 None。"""
+        stmt = select(func.max(DemandPopularityStats.biz_dt))
+        return self.session.scalar(stmt)
+
+    def has_biz_dt(self, biz_dt: str) -> bool:
+        """指定业务日是否存在热度统计数据。"""
+        stmt = select(DemandPopularityStats.id).where(
+            DemandPopularityStats.biz_dt == biz_dt
+        ).limit(1)
+        return self.session.scalars(stmt).first() is not None
+
     def list_by_biz_dt(self, biz_dt: str) -> list[DemandPopularityStats]:
         """返回指定业务日的全部热度统计行。"""
         stmt = select(DemandPopularityStats).where(DemandPopularityStats.biz_dt == biz_dt)
         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]:
+        """按业务日 + demand_belong_category.id 列表查询热度行。"""
+        if not belong_ids:
+            return []
+        rows: list[DemandPopularityStats] = []
+        for i in range(0, len(belong_ids), _BATCH_SIZE):
+            batch = belong_ids[i : i + _BATCH_SIZE]
+            stmt = select(DemandPopularityStats).where(
+                DemandPopularityStats.biz_dt == biz_dt,
+                DemandPopularityStats.demand_category_id.in_(batch),
+            )
+            rows.extend(self.session.scalars(stmt).all())
+        return rows
+
     def upsert_rows(self, rows: list[dict]) -> int:
         """按 (demand_category_id, biz_dt) 批量 upsert。"""
         if not rows: