""" =========================================== 图册检索模块 - 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()