kgrag/graph_search/atlas_retrieval421.py
2026-06-30 13:35:52 +08:00

256 lines
9.5 KiB
Python
Raw Permalink 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 _match_nodes_with_retrieval(
session,
node_names: List[str],
sync_driver,
top_k: int = 10
) -> List[Dict[str, Any]]:
"""
使用 NodeRetrieval 混合检索匹配节点(同时搜索系统、子系统、设备)
Args:
session: Neo4j 异步会话
node_names: 节点名称列表
sync_driver: Neo4j 同步 driver用于 NodeRetrieval
top_k: 检索返回的节点数量上限
Returns:
匹配到的节点信息列表
"""
from graph_search.node_retrieval import NodeRetrieval
try:
embeddings = EmbeddingWrapper()
node_retrieval = NodeRetrieval(sync_driver, embeddings)
route_res = []
for node_name in node_names:
if node_name:
route_res.append({"label": "系统", "entity": node_name})
route_res.append({"label": "子系统", "entity": node_name})
route_res.append({"label": "设备", "entity": node_name})
if not route_res:
logger.warning("没有提供有效的节点名称")
return []
retrieved_nodes = await node_retrieval.retrieve_nodes(route_res, top_k=top_k)
matched_nodes_map = {}
for label, candidates in retrieved_nodes.items():
if candidates:
best = candidates[0]
node_data = best.get("node", {})
best_name = node_data.get("名称") or node_data.get("name") or best.get("name", "")
if best_name in matched_nodes_map:
if best.get("score", 0.0) > matched_nodes_map[best_name]["score"]:
matched_nodes_map[best_name] = {
"best": best,
"score": best.get("score", 0.0)
}
else:
matched_nodes_map[best_name] = {
"best": best,
"score": best.get("score", 0.0)
}
matched_nodes = []
for best_name, data in matched_nodes_map.items():
best = data["best"]
query = """
MATCH (n)
WHERE (n:系统 OR n:子系统 OR n:设备)
AND (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(query, name=best_name)
record = await result.single()
if record:
matched_nodes.append({
"id": record.get("id"),
"name": record.get("name"),
"labels": record.get("labels"),
"props": record.get("props", {}),
"score": best.get("score", 0.0)
})
logger.info(f"NodeRetrieval 匹配到 {len(matched_nodes)} 个节点")
return matched_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()
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()