更新 graph_search/result_formatter.py

This commit is contained in:
ChenRui 2026-07-15 17:18:29 +08:00
parent 61ee346a43
commit 1b3cd3bfc3

View File

@ -5,9 +5,13 @@
"""
import logging
from typing import Dict, List, Any, Optional
import re
from typing import Dict, List, Any, Optional, Set
logger = logging.getLogger(__name__)
EXCLUDED_RESULT_NODE_LABELS = {"维修工作", "操作程序", "操作使用"}
_FILTERED_OUT = object()
class SimpleNode:
"""
@ -537,6 +541,110 @@ def format_entry_nodes_as_results(entry_nodes: dict) -> List[Dict[str, Any]]:
return results
def _normalize_label_values(labels: Any) -> Set[str]:
"""将 labels/标签/type 等字段统一成标签集合。"""
if labels is None:
return set()
if isinstance(labels, str):
parts = re.split(r"[,,、/\s]+", labels)
return {part.strip("[]'\" ") for part in parts if part.strip("[]'\" ")}
if isinstance(labels, (list, tuple, set, frozenset)):
return {str(label).strip() for label in labels if str(label).strip()}
return {str(labels).strip()} if str(labels).strip() else set()
def _dict_label_values(data: Dict[str, Any]) -> Set[str]:
label_keys = ("labels", "标签", "label", "类型", "type", "节点类型", "node_type")
labels = set()
for key in label_keys:
if key in data:
labels.update(_normalize_label_values(data.get(key)))
return labels
def _has_excluded_result_label(value: Any) -> bool:
"""判断一个结果对象是否明确属于需要从 results 隐藏的节点类型。"""
labels = set()
if isinstance(value, dict):
labels.update(_dict_label_values(value))
elif hasattr(value, "labels") and (hasattr(value, "element_id") or hasattr(value, "id")):
try:
labels.update(_normalize_label_values(value.labels))
except Exception:
pass
return bool(labels & EXCLUDED_RESULT_NODE_LABELS)
def _mentions_excluded_result_label(text: Any) -> bool:
if text is None:
return False
return any(label in str(text) for label in EXCLUDED_RESULT_NODE_LABELS)
def _is_empty_result_content(value: Any) -> bool:
return value in (None, {}, [])
def _filter_excluded_result_content(value: Any, parent_key: Any = None) -> Any:
"""
只用于 results 字段过滤掉维修工作操作程序操作使用这类节点内容
nodes/links 的原始返回不会经过这个函数
"""
if parent_key is not None and _mentions_excluded_result_label(parent_key):
return _FILTERED_OUT
if _has_excluded_result_label(value):
return _FILTERED_OUT
if isinstance(value, dict):
name_value = value.get("名称", value.get("name"))
if _mentions_excluded_result_label(name_value):
return _FILTERED_OUT
filtered = {}
for key, child in value.items():
if _mentions_excluded_result_label(key):
continue
filtered_child = _filter_excluded_result_content(child, key)
if filtered_child is not _FILTERED_OUT:
filtered[key] = filtered_child
return _FILTERED_OUT if _is_empty_result_content(filtered) else filtered
if isinstance(value, list):
filtered_items = []
for item in value:
filtered_item = _filter_excluded_result_content(item)
if filtered_item is not _FILTERED_OUT and not _is_empty_result_content(filtered_item):
filtered_items.append(filtered_item)
return filtered_items
return value
def _filter_results_for_output(results: Any) -> List[Any]:
filtered = _filter_excluded_result_content(results)
if filtered is _FILTERED_OUT or _is_empty_result_content(filtered):
return []
if isinstance(filtered, list):
return filtered
return [filtered]
def _filter_result_nodes_for_output(nodes: Any) -> List[Dict[str, Any]]:
if not isinstance(nodes, list):
return []
return [
node for node in nodes
if isinstance(node, dict) and not _has_excluded_result_label(node)
]
def _format_value(value: Any) -> str:
"""
@ -649,8 +757,9 @@ def serialize_graph_for_llm(graph_response: dict) -> str:
- results任意查询返回值聚合/记录/字符串等
"""
data = graph_response.get("data", {})
nodes = data.get("nodes", [])
results = data.get("results", [])
# 这里的过滤只影响 results 字段里的文本展示,不改动接口返回的 data.nodes/data.links。
nodes = _filter_result_nodes_for_output(data.get("nodes", []))
results = _filter_results_for_output(data.get("results", []))
output_lines = []
@ -695,3 +804,4 @@ def serialize_graph_for_llm(graph_response: dict) -> str:
return "图谱查询完成,但未返回有效数据。"
return "\n".join(output_lines).rstrip()