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

584 lines
26 KiB
Python
Raw 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.

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