""" =========================================== 图谱检索服务模块 - GraphService =========================================== 功能:整合路由、节点检索和图谱检索的完整服务 提供统一的接口供app.py调用 =========================================== """ import os import logging import time from typing import Dict, Any, Optional, List, Tuple # 加载 .env 文件中的环境变量 try: from dotenv import load_dotenv load_dotenv() except ImportError: # 如果 python-dotenv 未安装,跳过加载(环境变量可能已通过其他方式设置) pass logger = logging.getLogger(__name__) logger.setLevel(logging.INFO) # 导入固定的schema try: from graph_search.schema_sample import SCHEMA_SAMPLE except ImportError: logger.warning("无法导入SCHEMA_SAMPLE,将使用空schema") SCHEMA_SAMPLE = "" def remove_knowledge_source(data: Any) -> Any: """ 递归删除数据中的 knowledge_source 字段(用于 results 字段) Args: data: 要处理的数据(可以是字典、列表或基本类型) Returns: 清理后的数据(已删除所有 knowledge_source 字段) """ if isinstance(data, dict): # 创建新字典,排除 knowledge_source result = {} for key, value in data.items(): if key != 'knowledge_source': # 递归处理嵌套的字典、列表 result[key] = remove_knowledge_source(value) return result elif isinstance(data, list): # 递归处理列表中的每个元素 return [remove_knowledge_source(item) for item in data] else: # 基本类型直接返回 return data # ========== ========== 导入依赖 ========== ========== try: from neo4j import GraphDatabase from langchain_community.graphs.neo4j_graph import Neo4jGraph from graph_search.graph_retrieval import GraphRetrieval from graph_search.route_label import RouteLabel from graph_search.node_retrieval import NodeRetrieval from graph_search.result_formatter import format_results, format_entry_nodes_as_results, remove_sensitive_fields, serialize_graph_for_llm from graph_search.init_prompts import get_all_ontology_labels, get_all_relationship_types from modelsAPI.model_api import OpenaiAPI GRAPH_SERVICE_AVAILABLE = True except ImportError as e: GRAPH_SERVICE_AVAILABLE = False logger.warning(f"图谱服务依赖未安装: {e}") class GraphService: """ 图谱检索服务类 功能:整合路由、节点检索和图谱检索的完整流程 """ def __init__(self,neo4j_url: str = None,neo4j_user: str = None,neo4j_password: str = None,driver = None ): """ 初始化图谱服务 Args: neo4j_url: Neo4j连接URL neo4j_user: Neo4j用户名 neo4j_password: Neo4j密码 driver: 可选的Neo4j driver实例,如果提供则复用外部driver(如app.py中的driver) """ if not GRAPH_SERVICE_AVAILABLE: raise ImportError("图谱服务依赖未安装") # ========== 从环境变量获取配置 ========== self.neo4j_url = neo4j_url or os.getenv("NEO4J_URI") or os.getenv("NEO4J_URL", "bolt://localhost:7687") self.neo4j_user = neo4j_user or os.getenv("NEO4J_USERNAME") or os.getenv("NEO4J_USER", "neo4j") self.neo4j_password = neo4j_password or os.getenv("NEO4J_PASSWORD", "neo4j") # ========== 初始化LLM API ========== self.llm_api = OpenaiAPI() # ========== 延迟初始化Neo4j连接 ========== self._graph_retrieval = None self._driver = driver self._get_driver(driver) # 如果提供了外部driver,直接使用;否则延迟初始化 self._neo4j_graph = None # Neo4jGraph 实例,用于获取 schema self._external_driver = driver is not None # 标记是否使用了外部driver # ========== 定义可选节点标签和关系类型 ========== self.node_types = get_all_ontology_labels(self._driver) self.relationship_types = get_all_relationship_types(self._driver) def _get_driver(self, driver): """获取Neo4j driver实例(延迟初始化,可复用)""" if driver is None: try: self._driver = GraphDatabase.driver( self.neo4j_url, auth=(self.neo4j_user, self.neo4j_password) ) self._external_driver = False # 标记为内部创建的driver except Exception as e: logger.error(f"Neo4j Driver实例初始化失败: {e}", exc_info=True) raise def _get_neo4j_graph(self): """获取Neo4jGraph实例(延迟初始化,用于获取schema)""" if self._neo4j_graph is None: try: self._neo4j_graph = Neo4jGraph( url=self.neo4j_url, username=self.neo4j_user, password=self.neo4j_password, enhanced_schema=True # 使用增强的schema信息,提供更详细的节点标签、关系类型和属性等 ) except Exception as e: logger.error(f"Neo4jGraph实例初始化失败: {e}", exc_info=True) raise return self._neo4j_graph def _fetch_schema_from_db_original(self): """ 原来的方式:通过 Neo4jGraph 自动获取 Neo4j schema 返回: neo4j_schema: str - 供 LLM 使用的 schema 文本(Neo4jGraph 返回的完整 schema) structured_schema: dict - 空字典(不进行解析,Schema验证会自动跳过) """ try: neo4j_graph = self._get_neo4j_graph() schema_text = neo4j_graph.schema if not schema_text: logger.warning("Neo4jGraph schema 返回为空") return "", {} return schema_text, {} except Exception as e: logger.warning(f"通过 Neo4jGraph 获取 schema 失败: {e}", exc_info=True) return "", {} def _fetch_schema_from_db_fixed(self): """ 固定方式:使用 schema_sample.py 中定义的固定 SCHEMA_SAMPLE 返回: neo4j_schema: str - 供 LLM 使用的 schema 文本 structured_schema: dict - 空字典(不进行解析,Schema验证会自动跳过) """ try: schema_text = SCHEMA_SAMPLE if not schema_text: logger.warning("SCHEMA_SAMPLE 为空") return "", {} return schema_text, {} except Exception as e: logger.warning(f"获取固定schema失败: {e}", exc_info=True) return "", {} def _fetch_schema_from_db(self): """ 获取schema(当前使用固定方式,可通过注释切换到原来的方式) 返回: neo4j_schema: str - 供 LLM 使用的 schema 文本 structured_schema: dict - 空字典(不进行解析,Schema验证会自动跳过) """ # ========== 当前使用的方式(可通过注释切换) ========== return self._fetch_schema_from_db_fixed() # ========== 原来的方式(切换时取消注释,注释上面一行) ========== # return self._fetch_schema_from_db_original() def _get_graph_retrieval(self) -> GraphRetrieval: """获取GraphRetrieval实例(延迟初始化)""" if self._graph_retrieval is None: # ========== 创建LLM包装器(适配GraphRetrieval的接口) ========== class LLMWrapper: """LLM包装器,将OpenaiAPI包装成GraphRetrieval需要的格式""" def __init__(self, llm_api): self.llm_api = llm_api async def ainvoke(self, prompt): """异步调用LLM""" if isinstance(prompt, str): query = prompt elif hasattr(prompt, 'messages'): # 如果是ChatPromptTemplate格式,转换为字符串 messages = prompt.messages query = "\n".join([msg.content for msg in messages if hasattr(msg, 'content')]) else: query = str(prompt) model = os.getenv("OPENAI_MODEL") content = await self.llm_api.open_api_chat_async(query, model, temperature=0) # 返回类似LLM输出的对象 class LLMOutput: def __init__(self, content): self.content = content return LLMOutput(content) llm_wrapper = LLMWrapper(self.llm_api) # ========== 获取 schema 信息 ========== # 使用 Neo4jGraph 获取 schema neo4j_schema, structured_schema = self._fetch_schema_from_db() if not neo4j_schema: logger.warning("GraphService: Neo4jGraph schema 为空,图谱检索可能无法正常工作") self._graph_retrieval = GraphRetrieval( driver=self._driver, neo4j_schema=neo4j_schema, llm=llm_wrapper, allowed_node_types=self.node_types, allowed_relationship_types=self.relationship_types ) return self._graph_retrieval async def extract_entry_nodes(self, query: str, top_n: int = 10) -> Tuple[Dict[str, List[Dict[str, Any]]], str, Dict[str, Any]]: """ 提取入口节点 功能: 1. 路由标签识别 2. 节点检索 Args: query: 用户查询 top_n: 每个标签检索的节点数量 Returns: tuple: (入口节点字典, 重写后的查询, 分类信息) 如果未重写,则返回原始query 分类信息包含 query_type, has_time, time_type """ extract_start = time.time() # ========== 1. 路由标签识别 ========== route_start = time.time() # 获取schema信息,用于帮助模型识别节点类型 # 使用 Neo4jGraph 获取 schema schema = None try: schema, _ = self._fetch_schema_from_db() if not schema: logger.warning("获取schema失败,将不使用schema信息") except Exception as e: logger.warning(f"获取schema失败,将不使用schema信息: {e}") schema = None route_label = RouteLabel(self.llm_api, self.node_types, schema=schema) route_res, rewritten_query, query_classification = await route_label.route(query) logger.info("输出route_res") logger.info(route_res) route_elapsed = (time.time() - route_start) * 1000 logger.info(f"[时间统计] 路由标签识别完成,耗时: {route_elapsed:.2f}ms") # 使用重写后的查询覆盖原始查询 if rewritten_query and rewritten_query != query: query = rewritten_query if not route_res: logger.warning("未能识别到入口节点") return {}, query, query_classification # ========== 2. 节点检索 ========== node_retrieval_start = time.time() try: # 创建嵌入模型包装器 class EmbeddingWrapper: """嵌入模型包装器""" def embed_documents(self, texts: List[str]) -> List[List[float]]: """批量获取嵌入向量(一次性处理所有文本)""" embed_start = time.time() if not texts: return [] try: # 一次性批量处理所有文本 result = OpenaiAPI.batch_embeddings(texts) # 调用服务器嵌入模型 embed_elapsed = (time.time() - embed_start) * 1000 # 从结果中提取embedding向量 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: 批量向量化返回格式异常,期望列表,实际: {type(result)}") return [None] * len(texts) except Exception as e: logger.error(f"EmbeddingWrapper: 批量向量化失败: {e}", exc_info=True) return [None] * len(texts) embeddings = EmbeddingWrapper() node_retrieval = NodeRetrieval(self._driver, embeddings) entry_nodes = await node_retrieval.retrieve_nodes(route_res, top_n) node_retrieval_elapsed = (time.time() - node_retrieval_start) * 1000 logger.info(f"[时间统计] 节点检索完成,耗时: {node_retrieval_elapsed:.2f}ms") extract_elapsed = (time.time() - extract_start) * 1000 logger.info(f"[时间统计] 入口节点提取完成,总耗时: {extract_elapsed:.2f}ms (路由: {route_elapsed:.2f}ms, 检索: {node_retrieval_elapsed:.2f}ms)") logger.info(f"{entry_nodes}") logger.info(f"{query}") logger.info(f"{query_classification}") return entry_nodes, query, query_classification except Exception as e: logger.error(f"入口节点提取失败: {e}", exc_info=True) return {}, query, query_classification # 注意:不再关闭driver,因为它是可复用的 # ========== ========== 图谱检索:简单版本(只判断统计/非统计) ========== ========== async def search( self, query: str, entry_nodes: Optional[Dict[str, List[Dict[str, Any]]]] = None, top_k: int = 10, max_attempts: int = 3, query_type: str = "auto", # "auto" | "aggregate" | "detail" ) -> Dict[str, Any]: """ 执行图谱检索(完整流程) Args: query: 用户查询 entry_nodes: 入口节点(可选,不提供则自动提取) top_k: 节点检索数量 max_attempts: Cypher生成最大尝试次数 Returns: Dict: 检索结果,包含code, message, data, meta信息,符合app.py要求 """ total_start = time.time() logger.info(f"[时间统计] ========== 开始图谱检索,查询: '{query}' ==========") try: # ========== 1. 入口节点处理(如果已提供,直接使用;否则提取) ========== entry_nodes_start = time.time() query_classification = None # 初始化分类信息 if not entry_nodes: logger.info(f"[时间统计] 入口节点未提供,开始自动提取...") entry_nodes, rewritten_query, query_classification = await self.extract_entry_nodes(query = query, top_n = 10) entry_nodes_elapsed = (time.time() - entry_nodes_start) * 1000 # 使用重写后的查询覆盖原始查询 if rewritten_query and rewritten_query != query: query = rewritten_query if not entry_nodes: total_elapsed = (time.time() - total_start) * 1000 logger.warning(f"[时间统计] 未能识别到入口节点,总耗时: {total_elapsed:.2f}ms") return { "code": 200, "message": "success", "data": { "nodes": [], "links": [], "results": [] }, "meta": { "elapsed_ms": round(total_elapsed, 2), "error": "未能识别到入口节点", "query_type": "no_entry_nodes", "query_classification": query_classification } } logger.info(f"[时间统计] 入口节点提取完成,耗时: {entry_nodes_elapsed:.2f}ms") else: entry_nodes_elapsed = (time.time() - entry_nodes_start) * 1000 # ========== 2. 执行图谱检索 ========== graph_search_start = time.time() graph_retrieval = self._get_graph_retrieval() # 如果从重写中获得了分类信息,优先使用分类信息确定 query_type if query_classification: classification_query_type = query_classification.get("query_type", "detail") effective_query_type = classification_query_type else: # 如果没有分类信息,使用传入的 query_type 参数 effective_query_type = query_type if query_type in ("auto", "aggregate", "detail") else "auto" raw_results = await graph_retrieval.search( query=query, entry_nodes=entry_nodes, max_attempts=max_attempts, query_type=effective_query_type, query_classification=query_classification # 传递分类信息 ) graph_search_elapsed = (time.time() - graph_search_start) * 1000 logger.info(f"[时间统计] 图谱检索完成,耗时: {graph_search_elapsed:.2f}ms") # ========== 3. 格式化结果 ========== # 判断是否为统计类查询(优先使用分类信息) if query_classification: is_aggregate = query_classification.get("query_type") == "aggregate" else: is_aggregate = effective_query_type == "aggregate" or (effective_query_type == "auto" and raw_results and raw_results[0].get('_aggregate_value') is not None) # 从 rerank 后的结果中提取路径信息(节点和关系) path_data = {"nodes": [], "links": []} try: # raw_results 已经是 rerank 后的结果(字典列表,包含 rerank_score) if len(raw_results) > top_k: path_results_for_format = raw_results[:top_k] else: path_results_for_format = raw_results # 使用限制后的结果来格式化路径信息 path_data = format_results(path_results_for_format) except Exception as e: logger.warning(f"格式化结果失败: {e}", exc_info=True) # 过滤原始结果,递归删除 embedding、fulltext 和 path,并展开嵌套的 result 字段 # raw_results 已经是 rerank 后的结果,需要去除 path 字段,保留其他信息作为 results filtered_results = [] if is_aggregate: # 统计类查询:results 只包含统计值(一个值) if raw_results and raw_results[0].get('_aggregate_value') is not None: # 从元数据中提取统计值 aggregate_value = raw_results[0]['_aggregate_value'] # 删除 knowledge_source filtered_results = [remove_knowledge_source(aggregate_value)] else: # 如果没有 _aggregate_value,尝试从第一个结果的 result 字段提取 if raw_results: first_result = raw_results[0] if isinstance(first_result, dict) and 'result' in first_result: aggregate_value = first_result['result'] # 删除 knowledge_source filtered_results = [remove_knowledge_source(aggregate_value)] else: # 降级:使用第一个结果,删除 knowledge_source filtered_results = [remove_knowledge_source(first_result)] else: # 非统计类查询:保持原有逻辑 for result in raw_results: # 创建结果副本,去除 path 相关字段 result_copy = result.copy() # 删除 path 和 paths 字段(这些已经提取到 nodes 和 links 中了) result_copy.pop('path', None) result_copy.pop('paths', None) result_copy.pop('_aggregate_value', None) # 删除统计类查询的元数据 # 先递归删除敏感字段(embedding、fulltext) cleaned_result = remove_sensitive_fields(result_copy) # 如果结果中有 result 字段,则展开它 if isinstance(cleaned_result, dict) and 'result' in cleaned_result: result_value = cleaned_result['result'] # 如果 result 的值是列表,则展开列表中的所有元素 if isinstance(result_value, list): # 删除敏感字段和 knowledge_source filtered_results.extend([remove_knowledge_source(remove_sensitive_fields(item)) for item in result_value]) else: # 如果 result 的值是单个对象,直接使用,删除 knowledge_source filtered_results.append(remove_knowledge_source(remove_sensitive_fields(result_value))) else: # 没有 result 字段,直接使用清理后的结果,删除 knowledge_source filtered_results.append(remove_knowledge_source(cleaned_result)) # ========== 4. 格式化结果用于 LLM(在返回前应用格式化字符串到 results)========== # 构建临时的 graph_response 字典用于序列化 temp_graph_response = { "data": { "nodes": path_data.get("nodes", []), "results": filtered_results } } # 使用 serialize_graph_for_llm 格式化结果 formatted_results_str = serialize_graph_for_llm(temp_graph_response) # 将格式化后的字符串应用到 results(替换原有的 results) filtered_results = formatted_results_str # 构建响应 total_elapsed = (time.time() - total_start) * 1000 response = { "code": 200, "message": "success", "data": { "nodes": path_data.get("nodes", []), "links": path_data.get("links", []), "results": filtered_results }, "meta": { "result_count": len(filtered_results), "query_type": "graph_search", "elapsed_ms": round(total_elapsed, 2), "cypher": graph_retrieval.last_cypher if hasattr(graph_retrieval, 'last_cypher') else None, "query_classification": query_classification # 添加分类信息到返回结果中 } } logger.info(f"[时间统计] ========== 图谱检索完成 ==========") logger.info(f"[时间统计] 总耗时: {total_elapsed:.2f}ms") logger.info(f"[时间统计] - 入口节点提取: {entry_nodes_elapsed:.2f}ms") logger.info(f"[时间统计] - 图谱检索: {graph_search_elapsed:.2f}ms") return response except Exception as e: total_elapsed = (time.time() - total_start) * 1000 logger.error(f"[时间统计] 图谱检索异常,总耗时: {total_elapsed:.2f}ms, 错误: {e}", exc_info=True) return { "code": 500, "message": "图谱检索失败", "data": { "nodes": [], "links": [], "results": [] }, "meta": { "elapsed_ms": round(total_elapsed, 2), "error": str(e), "query_type": "error" } } def close(self): """ 关闭连接,释放资源 注意:在应用关闭时调用此方法 注意:如果使用的是外部driver(如app.py传入的),则不会关闭driver,由外部管理 """ try: if self._driver and not self._external_driver: # 只关闭内部创建的driver,不关闭外部传入的driver self._driver.close() self._driver = None logger.info("Neo4j Driver连接已关闭") elif self._external_driver: logger.info("使用的是外部driver,不关闭连接(由外部管理)") if self._neo4j_graph: # Neo4jGraph 内部使用 driver,如果 driver 已关闭,这里不需要额外操作 # 但为了清晰,我们重置引用 self._neo4j_graph = None logger.info("Neo4jGraph 实例已释放") except Exception as e: logger.warning(f"关闭连接时出错: {e}") def __del__(self): """析构函数,确保资源被释放""" self.close()