584 lines
26 KiB
Python
584 lines
26 KiB
Python
"""
|
||
===========================================
|
||
图谱检索服务模块 - 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()
|