weight_score_query_tools.py 6.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176
  1. #!/usr/bin/env python3
  2. # -*- coding: utf-8 -*-
  3. """
  4. 权重分查询工具
  5. 从 examples/demand/data/{execution_id} 目录读取权重分 JSON,
  6. 支持按元素/分类查询 TopN,以及按名称列表批量查询权重分。
  7. """
  8. import json
  9. import os
  10. from pathlib import Path
  11. from agent import tool
  12. from examples.demand.demand_agent_context import TopicBuildAgentContext
  13. from examples.demand.tool_logging import log_tool_input, log_tool_output
  14. _VALID_LEVELS = {"元素", "分类"}
  15. _VALID_DIMENSIONS = {"实质", "形式", "意图"}
  16. def _get_weight_file_path(level: str, dimension: str) -> Path:
  17. """根据参数构造权重数据文件路径。"""
  18. execution_id = TopicBuildAgentContext.get_execution_id()
  19. if execution_id is None:
  20. raise ValueError("未设置 execution_id,请先在 TopicBuildAgentContext 中设置")
  21. filename = f"{dimension}_{level}.json"
  22. configured_base_dir = os.getenv("DEMAND_WEIGHT_DATA_DIR")
  23. if configured_base_dir:
  24. base_dir = Path(configured_base_dir) / str(execution_id)
  25. else:
  26. base_dir = Path(__file__).parent / "data" / str(execution_id)
  27. return base_dir / filename
  28. def _load_weight_data(level: str, dimension: str) -> list[dict]:
  29. """读取并返回权重数据列表。"""
  30. file_path = _get_weight_file_path(level=level, dimension=dimension)
  31. if not file_path.exists():
  32. raise FileNotFoundError(f"权重数据文件不存在: {file_path}")
  33. with file_path.open("r", encoding="utf-8") as f:
  34. data = json.load(f)
  35. if not isinstance(data, list):
  36. raise ValueError(f"权重数据格式错误,期望 list,实际为: {type(data).__name__}")
  37. return data
  38. def _validate_params(level: str, dimension: str):
  39. """校验通用参数。"""
  40. if level not in _VALID_LEVELS:
  41. raise ValueError(f"level 参数非法: {level},可选值: {sorted(_VALID_LEVELS)}")
  42. if dimension not in _VALID_DIMENSIONS:
  43. raise ValueError(f"dimension 参数非法: {dimension},可选值: {sorted(_VALID_DIMENSIONS)}")
  44. @tool("查询元素或分类权重分排名区间。参数:level(元素/分类)、dimension(实质/形式/意图)、start(起始排名,含,从1开始)、end(结束排名,含,从1开始)。")
  45. def get_weight_score_topn(level: str, dimension: str, start: int = 1, end: int = 10) -> str:
  46. """查询元素或分类权重分排名区间。
  47. Args:
  48. level: 查询层级,元素 或 分类。
  49. dimension: 查询维度,实质 / 形式 / 意图。
  50. start: 起始排名(包含),从 1 开始。
  51. end: 结束排名(包含),从 1 开始。
  52. Returns:
  53. JSON 字符串,包含查询参数、总量和区间数据。
  54. """
  55. execution_id = TopicBuildAgentContext.get_execution_id()
  56. params = {
  57. "execution_id": execution_id,
  58. "level": level,
  59. "dimension": dimension,
  60. "start": start,
  61. "end": end,
  62. }
  63. log_tool_input("get_weight_score_topn", params)
  64. try:
  65. _validate_params(level=level, dimension=dimension)
  66. if start < 1 or end < 1:
  67. return log_tool_output(
  68. "get_weight_score_topn",
  69. f"错误: start/end 必须为大于等于 1 的整数,当前值: start={start}, end={end}",
  70. )
  71. if start > end:
  72. return log_tool_output(
  73. "get_weight_score_topn",
  74. f"错误: start 不能大于 end,当前值: start={start}, end={end}",
  75. )
  76. data = _load_weight_data(level=level, dimension=dimension)
  77. sorted_data = sorted(data, key=lambda x: float(x.get("score", 0)), reverse=True)
  78. # 用户输入为 1-based 且 end 为包含边界,需转换为 Python 切片
  79. ranged_items = sorted_data[start - 1 : end]
  80. result = {
  81. "level": level,
  82. "dimension": dimension,
  83. "start": start,
  84. "end": end,
  85. "total_count": len(data),
  86. "matched_count": len(ranged_items),
  87. "items": ranged_items,
  88. }
  89. return log_tool_output("get_weight_score_topn", json.dumps(result, ensure_ascii=False, indent=2))
  90. except Exception as e:
  91. return log_tool_output("get_weight_score_topn", f"查询失败: {e}")
  92. @tool("批量查询指定名称的权重分。参数:level(元素/分类)、dimension(实质/形式/意图)、names(名称列表)。")
  93. def get_weight_score_by_name(level: str, dimension: str, names: list[str]) -> str:
  94. """批量查询指定名称的权重分。
  95. Args:
  96. level: 查询层级,元素 或 分类。
  97. dimension: 查询维度,实质 / 形式 / 意图。
  98. names: 要查询的名称列表(元素名或分类名),顺序与返回 results 一一对应。
  99. Returns:
  100. JSON 字符串,含每个名称的 matched_count 与 items。
  101. """
  102. execution_id = TopicBuildAgentContext.get_execution_id()
  103. params = {
  104. "execution_id": execution_id,
  105. "level": level,
  106. "dimension": dimension,
  107. "names": names,
  108. }
  109. log_tool_input("get_weight_score_by_name", params)
  110. try:
  111. _validate_params(level=level, dimension=dimension)
  112. if not names:
  113. return log_tool_output("get_weight_score_by_name", "错误: names 不能为空列表")
  114. if not isinstance(names, list):
  115. return log_tool_output("get_weight_score_by_name", f"错误: names 必须为列表,当前类型: {type(names).__name__}")
  116. stripped: list[str] = []
  117. for i, n in enumerate(names):
  118. if n is None or (isinstance(n, str) and not n.strip()):
  119. return log_tool_output(
  120. "get_weight_score_by_name",
  121. f"错误: names[{i}] 不能为空",
  122. )
  123. stripped.append(str(n).strip())
  124. data = _load_weight_data(level=level, dimension=dimension)
  125. if level == "元素":
  126. key = "name"
  127. else:
  128. key = "category"
  129. results = []
  130. for target_name in stripped:
  131. matched = [item for item in data if str(item.get(key, "")).strip() == target_name]
  132. results.append(
  133. {
  134. "name": target_name,
  135. "matched_count": len(matched),
  136. "items": matched,
  137. }
  138. )
  139. result = {
  140. "level": level,
  141. "dimension": dimension,
  142. "query_count": len(stripped),
  143. "results": results,
  144. }
  145. return log_tool_output("get_weight_score_by_name", json.dumps(result, ensure_ascii=False, indent=2))
  146. except Exception as e:
  147. return log_tool_output("get_weight_score_by_name", f"查询失败: {e}")