kgrag/graph_search/quick_graph_search421.py
2026-07-29 18:10:19 +08:00

522 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
===========================================
临时图谱检索方案 - QuickGraphSearch
===========================================
功能直接用用户问题embedding混合检索找最匹配节点
获取该节点的一跳邻居,用大模型选择最合适的
只输出入口节点和选中节点
===========================================
"""
import json
import os
import logging
import time
import re
from typing import Any, Dict, List, Optional
from dotenv import load_dotenv
from neo4j import AsyncGraphDatabase, GraphDatabase
load_dotenv()
logger = logging.getLogger(__name__)
NEO4J_CONFIG = {
"uri": os.getenv("NEO4J_URI", "bolt://192.168.0.46:57687"),
"user": os.getenv("NEO4J_USER", "neo4j"),
"password": os.getenv("NEO4J_PASSWORD", "zdht123@"),
}
class EmbeddingWrapper:
"""嵌入模型包装器"""
def embed_documents(self, texts: List[str]) -> List[List[float]]:
"""批量获取嵌入向量(一次性处理所有文本)"""
from modelsAPI.model_api import OpenaiAPI
if not texts:
return []
try:
result = OpenaiAPI.batch_embeddings(texts)
if result and isinstance(result, list):
embedding_list = [item.get('embedding') if isinstance(item, dict) else item for item in result]
if embedding_list and len(embedding_list) == len(texts):
return embedding_list
else:
logger.error(f"EmbeddingWrapper: 批量向量化返回数量不匹配,期望: {len(texts)}, 实际: {len(embedding_list) if embedding_list else 0}")
return [None] * len(texts)
else:
logger.error(f"EmbeddingWrapper: 批量向量化返回格式异常")
return [None] * len(texts)
except Exception as e:
logger.error(f"EmbeddingWrapper: 批量向量化失败: {e}", exc_info=True)
return [None] * len(texts)
async def _search_all_nodes_with_hybrid(
session,
query: str,
top_k: int = 20
) -> List[Dict[str, Any]]:
"""
直接对所有节点进行混合检索embedding+全文),找最匹配的节点
使用 NodeRetrieval 替代 HybridRetriever避免兼容性问题
Args:
session: Neo4j 异步会话
query: 用户查询
top_k: 返回节点数量
Returns:
List[Dict]: 匹配到的节点列表
"""
from graph_search.node_retrieval import NodeRetrieval
uri = NEO4J_CONFIG.get("uri")
user = NEO4J_CONFIG.get("user")
password = NEO4J_CONFIG.get("password")
sync_driver = GraphDatabase.driver(uri, auth=(user, password))
try:
# 获取所有节点标签
all_labels = ["设备","系统","子系统","故障现象"]
logger.info(f"[混合检索] 所有节点标签: {all_labels}")
# 使用 NodeRetrieval 进行检索
embeddings = EmbeddingWrapper()
node_retrieval = NodeRetrieval(sync_driver, embeddings)
# 为每个标签构造路由结果
route_res = [{"label": label, "entity": query} for label in all_labels]
# 执行检索
retrieved_nodes = await node_retrieval.retrieve_nodes(route_res, top_k=top_k)
# 汇总结果
all_results = []
for label, nodes in retrieved_nodes.items():
for node in nodes:
all_results.append({
"name": node.get("name"),
"score": node.get("score", 0.0),
"node": node.get("node", {}),
"label": label
})
# 按score排序取top_k
all_results.sort(key=lambda x: x.get("score", 0), reverse=True)
all_results = all_results[:top_k]
logger.info(f"[混合检索] 找到 {len(all_results)} 个候选节点")
# 获取完整节点信息
final_nodes = []
for item in all_results:
node_data = item["node"]
label = item["label"]
name = item["name"]
get_node_query = f"""
MATCH (n:{label})
WHERE n.名称 = $name OR n.name = $name
RETURN elementId(n) AS id, n.名称 AS name, labels(n) AS labels, properties(n) AS props
LIMIT 1
"""
result = await session.run(get_node_query, name=name)
record = await result.single()
if record:
final_nodes.append({
"id": record.get("id"),
"name": record.get("name"),
"labels": record.get("labels"),
"props": record.get("props", {}),
"score": item.get("score", 0.0)
})
return final_nodes
finally:
sync_driver.close()
async def _rerank_nodes(
query,
nodes
):
"""
使用 rerank 对节点进行重排序
Args:
query: 用户查询
nodes: 节点列表
Returns:
List[Dict]: 重排序后的节点列表
"""
if not nodes:
return []
from modelsAPI.model_api import OpenaiAPI
# 格式化节点文本
documents = []
for node in nodes:
props = node.get("props", {})
name = node.get("name", "")
labels = ", ".join(node.get("labels", []))
doc_text = f"名称: {name}; 类型: {labels}"
documents.append(doc_text)
# Rerank
try:
rerank_results = OpenaiAPI.rerank_query(query, documents, top_n=len(documents))
# 按score排序
sorted_rerank = sorted(rerank_results, key=lambda x: x.get('scores', 0), reverse=True)
# 重新组织节点
reranked_nodes = []
for rr in sorted_rerank:
idx = rr.get('index', 0)
if 0 <= idx < len(nodes):
node = nodes[idx].copy()
node['rerank_score'] = rr.get('scores', 0)
reranked_nodes.append(node)
return reranked_nodes
except Exception as e:
logger.error(f"Rerank失败: {e}", exc_info=True)
# 失败时返回原顺序
for node in nodes:
node['rerank_score'] = 0
return nodes
async def _get_one_hop_neighbors(
session,
node_id: str
) -> List[Dict[str, Any]]:
"""
获取指定节点的一跳邻居节点
Args:
session: Neo4j 异步会话
node_id: 节点ID
Returns:
List[Dict]: 一跳邻居节点列表
"""
exclude_keys = {
"last_updated",
"created_at",
"knowledge_source",
"fulltext",
"embedding",
}
query = """
MATCH (n)
WHERE elementId(n) = $node_id
MATCH (n)-[r]-(m)
RETURN
elementId(m) AS id,
m.名称 AS name,
labels(m) AS labels,
properties(m) AS props,
elementId(r) AS rel_id,
type(r) AS rel_type,
properties(r) AS rel_props,
elementId(n) AS source_id,
n.名称 AS source_name,
labels(n) AS source_labels
LIMIT 100
"""
result = await session.run(query, node_id=node_id)
records = await result.data()
neighbors = []
for record in records:
props = record.get("props", {})
props_filtered = {k: v for k, v in props.items() if k not in exclude_keys}
rel_props = record.get("rel_props", {})
rel_props_filtered = {k: v for k, v in rel_props.items() if k not in exclude_keys}
neighbors.append({
"id": record.get("id"),
"name": record.get("name"),
"labels": record.get("labels", []),
"properties": props_filtered,
"relationship": {
"id": record.get("rel_id"),
"type": record.get("rel_type"),
"properties": rel_props_filtered,
"source_id": record.get("source_id"),
"source_name": record.get("source_name"),
"source_labels": record.get("source_labels", [])
}
})
return neighbors
def _format_node_for_llm(node: Dict[str, Any]) -> str:
"""
将节点信息格式化为供LLM理解的字符串仅包含标签和中文名称
Args:
node: 节点信息字典
Returns:
str: 格式化后的字符串
"""
name = node.get("name", "未知")
labels = ", ".join(node.get("labels", []))
rel = node.get("relationship", {})
parts = [f"节点名称: {name}"]
if labels:
parts.append(f"节点类型: {labels}")
if rel:
parts.append(f"通过关系: {rel.get('type', '未知')} 连接自: {rel.get('source_name', '未知')}")
return "\n".join(parts)
async def _select_best_node_with_llm(
query: str,
neighbors: List[Dict[str, Any]]
) -> Optional[Dict[str, Any]]:
"""
使用大模型从一跳邻居中选择最合适的节点
Args:
query: 用户查询
neighbors: 一跳邻居节点列表
Returns:
Optional[Dict]: 最合适的节点信息None表示没有合适的
"""
if not neighbors:
return None
from modelsAPI.model_api import OpenaiAPI
# 格式化邻居节点信息供LLM参考
neighbors_text = []
for i, node in enumerate(neighbors, 1):
neighbors_text.append(f"【候选节点 {i}\n{_format_node_for_llm(node)}\n")
prompt = f"""你是一个智能图谱检索助手,需要根据用户的查询,从提供的候选节点中选择最相关的一个节点。
用户查询: {query}
候选节点列表:
{''.join(neighbors_text)}
请仔细阅读用户查询和所有候选节点信息,选择与用户查询最相关的一个节点。
输出要求:
1. 请仅输出你选择的候选节点编号1, 2, 3...),不要输出任何其他内容
2. 如果没有任何相关的节点,请输出"0"
3. 不要解释原因,只输出数字编号
你的选择:"""
try:
model = os.getenv("OPENAI_MODEL")
response = await OpenaiAPI.open_api_chat_async(prompt, model, temperature=0)
# 提取数字编号
match = re.search(r'\d+', response.strip())
if match:
selected_index = int(match.group()) - 1
if 0 <= selected_index < len(neighbors):
logger.info(f"LLM选择了候选节点 {selected_index + 1}: {neighbors[selected_index].get('name', '未知')}")
return neighbors[selected_index]
logger.warning(f"LLM未选择有效节点原始响应: {response}")
return None
except Exception as e:
logger.error(f"使用LLM选择节点失败: {e}", exc_info=True)
return None
def _build_graph_result(
entry_node: Optional[Dict[str, Any]],
best_node: Optional[Dict[str, Any]],
query: str
) -> Dict[str, Any]:
"""
构建与operate方法一致的输出格式只包含入口节点和选中节点
Args:
entry_node: 入口节点
best_node: 选择的最佳节点
query: 用户查询
Returns:
Dict: 格式化后的结果
"""
exclude_keys = {
"last_updated",
"created_at",
"knowledge_source",
"fulltext",
"embedding",
}
nodes_g: List[Dict[str, Any]] = []
links_g: List[Dict[str, Any]] = []
node_id_set = set()
# 添加入口节点
if entry_node:
nid = entry_node.get("id")
if nid and nid not in node_id_set:
node_id_set.add(nid)
props = entry_node.get("props", {})
props_filtered = {k: v for k, v in props.items() if k not in exclude_keys}
name = props.get("名称") or props.get("name") or str(nid)
labels = entry_node.get("labels", [])
nodes_g.append({
"id": nid,
"name": name,
"labels": labels,
"properties": props_filtered,
})
# 添加最佳节点和关系
if best_node:
# 添加最佳节点
best_nid = best_node.get("id")
if best_nid and best_nid not in node_id_set:
node_id_set.add(best_nid)
props = best_node.get("properties", {})
name = best_node.get("name") or str(best_nid)
labels = best_node.get("labels", [])
nodes_g.append({
"id": best_nid,
"name": name,
"labels": labels,
"properties": props,
})
# 添加关系
rel = best_node.get("relationship", {})
if rel:
rel_id = rel.get("id")
src_id = rel.get("source_id")
tgt_id = best_nid
if src_id in node_id_set and tgt_id in node_id_set:
links_g.append({
"id": rel_id,
"label": rel.get("type", ""),
"source": src_id,
"target": tgt_id,
"properties": rel.get("properties", {}),
})
# 生成 results 文本优先用best_node其次用entry_node
results_text = ""
if best_node:
best_props = best_node.get("properties", {})
best_info = {k: v for k, v in best_props.items() if k not in exclude_keys and v}
results_text = json.dumps(best_info, ensure_ascii=False) if best_info else ""
elif entry_node:
entry_props = entry_node.get("props", {})
entry_info = {k: v for k, v in entry_props.items() if k not in exclude_keys and v}
results_text = json.dumps(entry_info, ensure_ascii=False) if entry_info else ""
graph_item = {"nodes": nodes_g, "links": links_g, "results": results_text}
return {"success": True, "xxxx.graph": [graph_item]}
async def quick_graph_search(
query: str,
route_res: Optional[List[Dict[str, str]]] = None,
top_k: int = 10,
) -> Dict[str, Any]:
"""
临时图谱检索接口
Args:
query: 用户查询
route_res: 路由结果列表(本方案不使用)
top_k: 检索返回的节点数量上限
Returns:
Dict: 检索结果,包含 success, xxxx.graph 等字段
"""
start_time = time.time()
if not query or not str(query).strip():
return {"success": False, "error": "查询不能为空"}
uri = NEO4J_CONFIG.get("uri")
user = NEO4J_CONFIG.get("user")
password = NEO4J_CONFIG.get("password")
if not uri or not user or not password:
return {"success": False, "error": "Neo4j 配置不完整,请检查 NEO4J_CONFIG"}
try:
async with AsyncGraphDatabase.driver(uri, auth=(user, password)) as driver:
async with driver.session() as session:
logger.info(f"[QuickGraphSearch] 开始直接混合检索,查询: '{query}'")
# 1. 直接用用户问题混合检索所有节点
candidate_nodes = await _search_all_nodes_with_hybrid(session, query, top_k=20)
if not candidate_nodes:
return {"success": False, "error": "未找到任何匹配节点"}
logger.info(f"[QuickGraphSearch] 混合检索找到 {len(candidate_nodes)} 个候选节点")
# 2. Rerank 重排序
reranked_nodes = await _rerank_nodes(query, candidate_nodes)
if not reranked_nodes:
return {"success": False, "error": "Rerank失败"}
# 3. 取最匹配的节点作为入口
entry_node = reranked_nodes[0]
entry_id = entry_node.get("id")
logger.info(f"[QuickGraphSearch] 选择入口节点: {entry_node.get('name', '未知')}, score={entry_node.get('rerank_score', 0)}")
# 4. 获取一跳邻居
neighbors = await _get_one_hop_neighbors(session, entry_id)
logger.info(f"[QuickGraphSearch] 获取到 {len(neighbors)} 个一跳邻居")
if not neighbors:
# 如果没有邻居节点,只返回入口节点
return _build_graph_result(entry_node, None, query)
# 5. 用LLM选择最合适的节点
best_node = await _select_best_node_with_llm(query, neighbors)
# 6. 构建结果(只包含入口节点和选中节点)
result = _build_graph_result(entry_node, best_node, query)
elapsed_ms = (time.time() - start_time) * 1000
logger.info(f"[QuickGraphSearch] 完成,耗时: {elapsed_ms:.2f}ms")
return result
except Exception as e:
logger.error(f"QuickGraphSearch 失败: {e}", exc_info=True)
return {"success": False, "error": f"检索失败: {str(e)}"}