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

364 lines
12 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.

"""
===========================================
图册检索模块 - atlas_retrieval
===========================================
功能:根据系统、子系统、设备,使用混合检索匹配节点
然后找到关联的图册节点提取图册信息和PDF URL
===========================================
"""
import json
import os
import logging
import time
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,
sync_driver,
top_k: int = 20
) -> List[Dict[str, Any]]:
"""
直接对所有节点进行混合检索embedding+全文),找最匹配的节点
(与 quick_graph_search.py 保持一致的逻辑)
使用 NodeRetrieval 替代 HybridRetriever避免兼容性问题
"""
from graph_search.node_retrieval import NodeRetrieval
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
except Exception as e:
logger.error(f"混合检索失败: {e}", exc_info=True)
return []
async def _rerank_nodes(
query,
nodes
):
"""
使用 rerank 对节点进行重排序
(与 quick_graph_search.py 保持一致)
"""
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 _match_nodes_with_retrieval(
session,
node_names: List[str],
sync_driver,
top_k: int = 10
) -> List[Dict[str, Any]]:
"""
使用混合检索+Rerank匹配节点与 quick_graph_search.py 保持一致的逻辑)
Args:
session: Neo4j 异步会话
node_names: 节点名称列表
sync_driver: Neo4j 同步 driver
top_k: 检索返回的节点数量上限
Returns:
匹配到的节点信息列表
"""
try:
all_matched_nodes = []
# 对每个节点名称分别进行检索+rerank
for node_name in node_names:
if not node_name:
continue
logger.info(f"[图册检索] 开始检索节点: '{node_name}'")
# 1. 混合检索
candidate_nodes = await _search_all_nodes_with_hybrid(
session, node_name, sync_driver, top_k=20
)
if not candidate_nodes:
logger.warning(f"[图册检索] 未找到节点: '{node_name}'")
continue
logger.info(f"[图册检索] 混合检索找到 {len(candidate_nodes)} 个候选节点")
# 2. Rerank 重排序
reranked_nodes = await _rerank_nodes(node_name, candidate_nodes)
if not reranked_nodes:
continue
# 3. 取最匹配的节点
best_node = reranked_nodes[0]
logger.info(f"[图册检索] 选择节点: {best_node.get('name', '未知')}, rerank_score={best_node.get('rerank_score', 0)}")
all_matched_nodes.append(best_node)
# 去重按节点ID
seen_ids = set()
unique_nodes = []
for node in all_matched_nodes:
nid = node.get("id")
if nid and nid not in seen_ids:
seen_ids.add(nid)
unique_nodes.append(node)
logger.info(f"[图册检索] 匹配到 {len(unique_nodes)} 个唯一节点")
return unique_nodes
except Exception as e:
logger.error(f"_match_nodes_with_retrieval 失败: {e}", exc_info=True)
return []
async def find_atlas_by_nodes(
node_names: Optional[List[str]] = None,
top_k: int = 10,
driver=None,
) -> Dict[str, Any]:
"""
根据节点名称列表检索关联的图册节点
Args:
node_names: 节点名称列表(可以是系统、子系统或设备)
top_k: 混合检索返回的候选节点数量
driver: Neo4j driver 实例(可选,如果不提供则创建新连接)
Returns:
Dict: 检索结果,层级结构为:节点名称 -> 图册名称 -> PDF URL
"""
start_time = time.time()
node_names = node_names or []
if not node_names:
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:
driver = AsyncGraphDatabase.driver(uri, auth=(user, password))
async with driver.session() as session:
logger.info(f"[图册检索] 开始检索 - 节点名称: {node_names}")
sync_driver = GraphDatabase.driver(uri, auth=(user, password))
try:
matched_nodes = await _match_nodes_with_retrieval(
session, node_names, sync_driver, top_k=top_k
)
finally:
sync_driver.close()
if not matched_nodes:
return {"success": False, "error": "未找到匹配的节点"}
logger.info(f"[图册检索] 匹配到 {len(matched_nodes)} 个源节点")
node_ids = [node["id"] for node in matched_nodes]
atlas_query = """
MATCH (source)
WHERE elementId(source) IN $node_ids
MATCH (source)-[*1..1]-(atlas:图册)
RETURN elementId(source) AS source_id,
source.名称 AS source_name,
elementId(atlas) AS atlas_id,
atlas.名称 AS atlas_name,
properties(atlas) AS atlas_props
"""
result = await session.run(atlas_query, node_ids=node_ids)
records = await result.data()
print(11111111111111111111111)
print(records)
if not records:
return {
"success": False,
"error": "未找到关联的图册节点",
"matched_nodes": matched_nodes
}
result_data = {}
for record in records:
source_name = record.get("source_name")
atlas_name = record.get("atlas_name")
atlas_props = record.get("atlas_props", {})
pdf_url = None
for key, value in atlas_props.items():
if "" in key or "pdf" in key.lower() or "url" in key.lower() or "文件" in key:
if isinstance(value, str) and value.startswith(("http", "/")):
pdf_url = value
break
if source_name not in result_data:
result_data[source_name] = {}
if atlas_name not in result_data[source_name]:
result_data[source_name][atlas_name] = []
if pdf_url and pdf_url not in result_data[source_name][atlas_name]:
result_data[source_name][atlas_name].append(pdf_url)
total_atlas_count = sum(len(atlas_dict) for atlas_dict in result_data.values())
total_pdf_count = sum(len(pdf_list) for atlas_dict in result_data.values() for pdf_list in atlas_dict.values())
logger.info(f"[图册检索] 找到 {total_atlas_count} 个图册节点,共 {total_pdf_count} 个PDF URL")
elapsed_ms = (time.time() - start_time) * 1000
logger.info(f"[时间统计] 图册检索完成,耗时: {elapsed_ms:.2f}ms")
return {
"success": True,
"data": result_data
}
except Exception as e:
logger.error(f"[图册检索] 检索失败: {e}", exc_info=True)
return {"success": False, "error": f"检索失败: {str(e)}"}
finally:
if driver is not None:
await driver.close()