440 lines
14 KiB
Python
440 lines
14 KiB
Python
"""
|
||
===========================================
|
||
临时图谱检索方案 - QuickGraphSearch
|
||
===========================================
|
||
功能:直接用用户问题embedding混合检索找最匹配节点
|
||
获取该节点的一跳邻居,用rerank选择最合适的
|
||
只输出入口节点和选中节点
|
||
===========================================
|
||
"""
|
||
|
||
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: 节点列表,每个节点需要有 name 字段
|
||
|
||
Returns:
|
||
List[Dict]: 重排序后的节点列表
|
||
"""
|
||
if not nodes:
|
||
return []
|
||
|
||
from modelsAPI.model_api import OpenaiAPI
|
||
|
||
# 格式化节点文本(只使用节点名称)
|
||
documents = []
|
||
for node in nodes:
|
||
name = node.get("name", "")
|
||
doc_text = name
|
||
documents.append(doc_text)
|
||
|
||
# Rerank(只取top 10,减少token消耗)
|
||
try:
|
||
rerank_results = OpenaiAPI.rerank_query(query, documents, top_n=min(len(documents), 10))
|
||
|
||
# 按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 _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失败"}
|
||
|
||
entry_node = reranked_nodes[0]
|
||
entry_id = entry_node.get("id")
|
||
logger.info(f"[QuickGraphSearch] 选择入口节点: {entry_node.get('name', '未知')}, rerank_score={entry_node.get('rerank_score', 0)}")
|
||
|
||
# 3. 获取入口节点的一跳邻居
|
||
neighbors = await _get_one_hop_neighbors(session, entry_id)
|
||
|
||
logger.info(f"[QuickGraphSearch] 获取到 {len(neighbors)} 个一跳邻居")
|
||
|
||
# 4. 对一跳邻居进行 rerank 选择最佳节点
|
||
best_node = None
|
||
if neighbors:
|
||
# 对邻居节点进行 rerank
|
||
reranked_neighbors = await _rerank_nodes(query, neighbors)
|
||
|
||
if reranked_neighbors:
|
||
best_node = reranked_neighbors[0]
|
||
logger.info(f"[QuickGraphSearch] 选择最佳一跳节点: {best_node.get('name', '未知')}, rerank_score={best_node.get('rerank_score', 0)}")
|
||
|
||
# 5. 构建结果(只包含入口节点和选中节点)
|
||
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)}"}
|
||
|