Explorar o código

Fix category elements top_n compatibility

SamLee hai 4 semanas
pai
achega
a4e2814604
Modificáronse 1 ficheiros con 14 adicións e 2 borrados
  1. 14 2
      examples/demand/demand_pattern_tools.py

+ 14 - 2
examples/demand/demand_pattern_tools.py

@@ -489,13 +489,14 @@ def search_categories(keyword: str, source_type: str = None) -> str:
     "\n- 通过 point_types 了解元素在灵感点/目的点/关键点中的分布"
     "\n- 通过 point_types 了解元素在灵感点/目的点/关键点中的分布"
     "\n- 获取元素名称后,传给 get_element_co_occurrences 查元素级共现"
     "\n- 获取元素名称后,传给 get_element_co_occurrences 查元素级共现"
     "\n- 从频繁项集的分类出发,落地到可用于选题的具体元素")
     "\n- 从频繁项集的分类出发,落地到可用于选题的具体元素")
-def get_category_elements(category_id: int, account_name=None, merge_leve2=None, platform=None) -> str:
+def get_category_elements(category_id: int, account_name=None, merge_leve2=None, platform=None, top_n: int = None) -> str:
     """获取某个分类节点下的元素列表(按名称去重聚合),按出现次数降序。
     """获取某个分类节点下的元素列表(按名称去重聚合),按出现次数降序。
 
 
     Args:
     Args:
         category_id: 分类节点ID。
         category_id: 分类节点ID。
         account_name: 按账号名筛选,支持单个字符串或列表(多个取OR)。
         account_name: 按账号名筛选,支持单个字符串或列表(多个取OR)。
         merge_leve2: 按二级品类筛选,支持单个字符串或列表(多个取OR)。
         merge_leve2: 按二级品类筛选,支持单个字符串或列表(多个取OR)。
+        top_n: 最多返回多少个元素;兼容 LLM 按旧工具习惯传入 top_n。
 
 
     Returns:
     Returns:
         元素列表的JSON字符串,每个元素含 name、element_type、occurrence_count、post_count。
         元素列表的JSON字符串,每个元素含 name、element_type、occurrence_count、post_count。
@@ -503,11 +504,22 @@ def get_category_elements(category_id: int, account_name=None, merge_leve2=None,
     execution_id = TopicBuildAgentContext.get_execution_id()
     execution_id = TopicBuildAgentContext.get_execution_id()
     merge_leve2 = _resolve_scope_arg(merge_leve2, "merge_leve2")
     merge_leve2 = _resolve_scope_arg(merge_leve2, "merge_leve2")
     platform = _resolve_scope_arg(platform, "platform")
     platform = _resolve_scope_arg(platform, "platform")
-    params = {"category_id": category_id, "account_name": account_name, "merge_leve2": merge_leve2, "platform": platform}
+    params = {
+        "category_id": category_id,
+        "account_name": account_name,
+        "merge_leve2": merge_leve2,
+        "platform": platform,
+        "top_n": top_n,
+    }
     _log_tool_input("get_category_elements", params)
     _log_tool_input("get_category_elements", params)
 
 
     data = pattern_service.get_category_elements(category_id, execution_id=execution_id,
     data = pattern_service.get_category_elements(category_id, execution_id=execution_id,
                                                  account_name=account_name, merge_leve2=merge_leve2, platform=platform)
                                                  account_name=account_name, merge_leve2=merge_leve2, platform=platform)
+    if top_n is not None:
+        try:
+            data = data[:max(int(top_n), 0)]
+        except (TypeError, ValueError):
+            pass
     result = json.dumps({
     result = json.dumps({
         "category_id": category_id,
         "category_id": category_id,
         "element_count": len(data),
         "element_count": len(data),