364 lines
12 KiB
Python
364 lines
12 KiB
Python
"""
|
||
===========================================
|
||
图册检索模块 - 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()
|
||
|