Browse Source

perf(demand): streamline mysql batch generation

SamLee 1 month ago
parent
commit
0e193045d3

+ 59 - 0
examples/demand/demand_mysql.md

@@ -0,0 +1,59 @@
+---
+model: deepseek-v4-flash-260425
+temperature: 0.3
+max_iterations: 20
+---
+
+# DemandAgent MySQL 批量写库任务
+
+你是 DemandAgent,目标是为「%merge_level2%」生成约 %count% 条可写入 `demand_content` 的需求。
+
+本入口用于云端 MySQL 单表落库,必须高效执行,不要展开整棵分类树,不要查询权重文件。
+
+## 硬约束
+
+- 只允许基于 PG Pattern V2 的 `scope="topic"` 频繁项集生成需求。
+- 每个 DemandItem 都必须带 `evidence_refs`。
+- `evidence_refs.source_kind` 固定为 `"pattern_itemset"`。
+- `evidence_refs.source_tool` 固定为 `"get_itemset_detail"`。
+- `itemset_ids` 必须来自 `get_frequent_itemsets` 和 `get_itemset_detail` 返回的真实 itemset。
+- `source_post_id` 必须从 `get_itemset_detail` 返回的 `post_ids` 中选择。
+- `seed_terms` 必须来自 itemset 的 `items` 分类名或分类路径。
+- 不要写 `source_certainty` 或 `validation_status`,这些由代码 DB 强校验后补齐。
+
+## 执行步骤
+
+1. 调用一次 `think_and_plan`,说明会直接用频繁项集生成需求。
+2. 调用 `get_frequent_itemsets(top_n=%count% * 3, min_support=5, sort_by="absolute_support")`。
+3. 从返回结果中挑选语义清晰、彼此尽量不重复的 itemset,数量尽量接近 %count%。
+4. 调用一次 `get_itemset_detail(itemset_ids=[...])` 查询这些 itemset 的详情。
+5. 调用一次 `create_demand_items(demand_items=[...])` 批量创建需求。
+6. 最后只用一句话总结,不要再调用其他工具。
+
+## DemandItem 格式
+
+```json
+{
+  "element_names": ["需求词1", "需求词2"],
+  "reason": "说明该需求来自哪个 itemset,以及哪些分类为什么能表达用户需求",
+  "desc": "用户希望看到什么内容",
+  "type": "pattern",
+  "evidence_refs": {
+    "source_kind": "pattern_itemset",
+    "source_tool": "get_itemset_detail",
+    "itemset_ids": [123],
+    "source_post_id": "从 post_ids 中选择一个真实帖子",
+    "case_ids": {
+      "pattern_itemset": ["同 source_post_id"]
+    },
+    "seed_terms": ["从 items 中提取的真实分类词"]
+  }
+}
+```
+
+## 选择标准
+
+- 优先选择 `absolute_support` 高、分类组合语义明确的 itemset。
+- 过滤纯形式、过于抽象、难以表达用户需求的组合。
+- 如果候选不足 %count% 条,可以少于 %count%,不要编造需求。
+- 不要重复创建语义相同的需求。

+ 32 - 1
examples/demand/demand_pattern_tools.py

@@ -20,6 +20,7 @@ Pattern 数据查询工具
 【重要】所有作为 Agent tool 注册的函数,必须包含完整的 docstring 签名。
 """
 import json
+import os
 from typing import Any
 
 from agent import tool
@@ -78,6 +79,33 @@ def _normalize_itemset_ids(itemset_ids: Any) -> list[int]:
         return []
 
 
+def _env_int(name: str, default: int) -> int:
+    try:
+        return int(os.getenv(name, str(default)))
+    except ValueError:
+        return default
+
+
+def _is_mysql_demand_content_entrypoint() -> bool:
+    return os.getenv("DEMAND_MYSQL_ENTRYPOINT") == "run_existing_execution_mysql"
+
+
+def _compact_itemset_detail_for_mysql(data: list[dict[str, Any]]) -> list[dict[str, Any]]:
+    if not _is_mysql_demand_content_entrypoint():
+        return data
+    max_post_ids = max(_env_int("DEMAND_ITEMSET_DETAIL_MAX_POST_IDS", 20), 1)
+    compacted: list[dict[str, Any]] = []
+    for raw_itemset in data:
+        itemset = dict(raw_itemset)
+        for key in ("post_ids", "matched_post_ids"):
+            post_ids = itemset.get(key)
+            if isinstance(post_ids, list) and len(post_ids) > max_post_ids:
+                itemset[f"{key}_total"] = len(post_ids)
+                itemset[key] = post_ids[:max_post_ids]
+        compacted.append(itemset)
+    return compacted
+
+
 # ============================================================================
 # 执行 & 配置 & 分类树
 # ============================================================================
@@ -197,6 +225,9 @@ def get_itemset_detail(itemset_ids) -> str:
         项集详情列表的JSON字符串,每项含 id, dimension_mode, target_depth, items, post_ids, absolute_support。
     """
     itemset_ids = _normalize_itemset_ids(itemset_ids)
+    if _is_mysql_demand_content_entrypoint():
+        max_ids = max(_env_int("DEMAND_ITEMSET_DETAIL_MAX_IDS", 12), 1)
+        itemset_ids = itemset_ids[:max_ids]
     params = {"itemset_ids": itemset_ids}
     _log_tool_input("get_itemset_detail", params)
 
@@ -208,7 +239,7 @@ def get_itemset_detail(itemset_ids) -> str:
     if not data:
         return _log_tool_output("get_itemset_detail", f"未找到 itemset_ids={itemset_ids} 的项集")
 
-    result = json.dumps(data, ensure_ascii=False, indent=2)
+    result = json.dumps(_compact_itemset_detail_for_mysql(data), ensure_ascii=False, indent=2)
     return _log_tool_output("get_itemset_detail", result)
 
 

+ 11 - 1
examples/demand/run.py

@@ -212,6 +212,15 @@ def register_selected_tools(tool_names: list[str]) -> None:
 
 def _enabled_tools_for_run(configured_tools: list[str]) -> list[str]:
     tools = configured_tools.copy()
+    if _is_mysql_demand_content_mode():
+        allowed = {
+            "think_and_plan",
+            "get_frequent_itemsets",
+            "get_itemset_detail",
+            "create_demand_item",
+            "create_demand_items",
+        }
+        return [tool_name for tool_name in tools if tool_name in allowed]
     if _is_local_json_mode():
         # 本地批量输出只需要 DemandItem JSON。长摘要会显著放大最后一轮上下文,
         # 且不参与下游 CFA 验证契约。
@@ -753,7 +762,8 @@ async def run_once(
     enabled_tools = _enabled_tools_for_run(ENABLED_TOOLS)
     register_selected_tools(enabled_tools)
 
-    prompt = SimplePrompt(base_dir / "demand.md")
+    prompt_file = "demand_mysql.md" if _is_mysql_demand_content_mode() else "demand.md"
+    prompt = SimplePrompt(base_dir / prompt_file)
 
     run_config = copy.deepcopy(RUN_CONFIG)
     model = resolve_model(prompt, run_config)

+ 2 - 0
examples/demand/run_existing_execution_mysql.py

@@ -47,6 +47,8 @@ def _configure_mysql_env(run_label: str) -> None:
     os.environ["DEMAND_OUTPUT_MODE"] = "mysql_demand_content"
     os.environ["DEMAND_MYSQL_ENTRYPOINT"] = "run_existing_execution_mysql"
     os.environ["DEMAND_RUN_LABEL"] = run_label
+    os.environ.setdefault("DEMAND_ITEMSET_DETAIL_MAX_IDS", "12")
+    os.environ.setdefault("DEMAND_ITEMSET_DETAIL_MAX_POST_IDS", "20")
 
 
 def _validate_execution_success(execution_id: int) -> None: