568 lines
18 KiB
Python
568 lines
18 KiB
Python
"""
|
||
===========================================
|
||
临时图谱检索方案 - QuickGraphSearch
|
||
===========================================
|
||
功能:直接用用户问题embedding混合检索找最匹配节点
|
||
获取该节点的一跳邻居,用大模型选择最合适的
|
||
只输出入口节点和选中节点
|
||
===========================================
|
||
"""
|
||
|
||
import json
|
||
import os
|
||
import logging
|
||
import time
|
||
import re
|
||
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,
|
||
top_k: int = 20
|
||
) -> List[Dict[str, Any]]:
|
||
"""
|
||
直接对所有节点进行混合检索(embedding+全文),找最匹配的节点
|
||
|
||
Args:
|
||
session: Neo4j 异步会话
|
||
query: 用户查询
|
||
top_k: 返回节点数量
|
||
|
||
Returns:
|
||
List[Dict]: 匹配到的节点列表
|
||
"""
|
||
from neo4j_graphrag.retrievers import HybridRetriever
|
||
import jieba
|
||
|
||
uri = NEO4J_CONFIG.get("uri")
|
||
user = NEO4J_CONFIG.get("user")
|
||
password = NEO4J_CONFIG.get("password")
|
||
|
||
sync_driver = GraphDatabase.driver(uri, auth=(user, password))
|
||
|
||
try:
|
||
# 中文分词
|
||
try:
|
||
words = jieba.lcut(query)
|
||
filtered_words = [
|
||
word.strip()
|
||
for word in words
|
||
if word.strip() and re.fullmatch(r"[a-zA-Z0-9\u4e00-\u9fa5]+", word.strip())
|
||
]
|
||
if not filtered_words:
|
||
filtered_words = [query]
|
||
query_text = " OR ".join(filtered_words)
|
||
except Exception as e:
|
||
logger.warning(f"分词失败: {e},使用原始查询")
|
||
query_text = query
|
||
|
||
# 获取所有节点标签
|
||
all_labels = []
|
||
result = await session.run("CALL db.labels() YIELD label RETURN label")
|
||
async for record in result:
|
||
all_labels.append(record["label"])
|
||
|
||
logger.info(f"[混合检索] 所有节点标签: {all_labels}")
|
||
|
||
# 对每个标签进行检索
|
||
embeddings = EmbeddingWrapper()
|
||
query_vector = None
|
||
try:
|
||
vectors = embeddings.embed_documents([query])
|
||
if vectors and vectors[0]:
|
||
query_vector = vectors[0]
|
||
except Exception as e:
|
||
logger.warning(f"查询向量化失败: {e}")
|
||
|
||
all_results = []
|
||
for label in all_labels:
|
||
try:
|
||
vector_index = label.lower() + "_vector"
|
||
fulltext_index = label.lower() + "_fulltext"
|
||
|
||
retriever = HybridRetriever(
|
||
sync_driver,
|
||
vector_index_name=vector_index,
|
||
fulltext_index_name=fulltext_index
|
||
)
|
||
|
||
import asyncio
|
||
result = await asyncio.to_thread(
|
||
retriever.get_search_results,
|
||
query_text,
|
||
query_vector,
|
||
top_k // len(all_labels) + 1,
|
||
effective_search_ratio=2,
|
||
)
|
||
|
||
if hasattr(result, 'records'):
|
||
for record in result.records:
|
||
node_data = dict(record["node"])
|
||
score = record.get("score", 0.0)
|
||
name = node_data.get("名称") or node_data.get("name") or str(record["node"])
|
||
|
||
all_results.append({
|
||
"name": name,
|
||
"score": score,
|
||
"node": node_data,
|
||
"label": label
|
||
})
|
||
except Exception as e:
|
||
logger.warning(f"标签 {label} 检索失败: {e}")
|
||
continue
|
||
|
||
# 按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
|
||
|
||
finally:
|
||
sync_driver.close()
|
||
|
||
|
||
async def _rerank_nodes(
|
||
query: str,
|
||
nodes: List[Dict[str, Any]]
|
||
) -> List[Dict[str, Any]]:
|
||
"""
|
||
使用 rerank 对节点进行重排序
|
||
|
||
Args:
|
||
query: 用户查询
|
||
nodes: 节点列表
|
||
|
||
Returns:
|
||
List[Dict]: 重排序后的节点列表
|
||
"""
|
||
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", []))
|
||
props_str = "; ".join([f"{k}={v}" for k, v in props.items() if v and k not in ["embedding", "fulltext", "last_updated", "created_at"]])
|
||
doc_text = f"名称: {name}; 类型: {labels}; 属性: {props_str}"
|
||
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 _get_one_hop_neighbors(
|
||
session,
|
||
node_id: str
|
||
) -> List[Dict[str, Any]]:
|
||
"""
|
||
获取指定节点的一跳邻居节点
|
||
|
||
Args:
|
||
session: Neo4j 异步会话
|
||
node_id: 节点ID
|
||
|
||
Returns:
|
||
List[Dict]: 一跳邻居节点列表
|
||
"""
|
||
exclude_keys = {
|
||
"last_updated",
|
||
"created_at",
|
||
"knowledge_source",
|
||
"fulltext",
|
||
"embedding",
|
||
}
|
||
|
||
query = """
|
||
MATCH (n)
|
||
WHERE elementId(n) = $node_id
|
||
MATCH (n)-[r]-(m)
|
||
RETURN
|
||
elementId(m) AS id,
|
||
m.名称 AS name,
|
||
labels(m) AS labels,
|
||
properties(m) AS props,
|
||
elementId(r) AS rel_id,
|
||
type(r) AS rel_type,
|
||
properties(r) AS rel_props,
|
||
elementId(n) AS source_id,
|
||
n.名称 AS source_name,
|
||
labels(n) AS source_labels
|
||
LIMIT 100
|
||
"""
|
||
|
||
result = await session.run(query, node_id=node_id)
|
||
records = await result.data()
|
||
|
||
neighbors = []
|
||
for record in records:
|
||
props = record.get("props", {})
|
||
props_filtered = {k: v for k, v in props.items() if k not in exclude_keys}
|
||
|
||
rel_props = record.get("rel_props", {})
|
||
rel_props_filtered = {k: v for k, v in rel_props.items() if k not in exclude_keys}
|
||
|
||
neighbors.append({
|
||
"id": record.get("id"),
|
||
"name": record.get("name"),
|
||
"labels": record.get("labels", []),
|
||
"properties": props_filtered,
|
||
"relationship": {
|
||
"id": record.get("rel_id"),
|
||
"type": record.get("rel_type"),
|
||
"properties": rel_props_filtered,
|
||
"source_id": record.get("source_id"),
|
||
"source_name": record.get("source_name"),
|
||
"source_labels": record.get("source_labels", [])
|
||
}
|
||
})
|
||
|
||
return neighbors
|
||
|
||
|
||
def _format_node_for_llm(node: Dict[str, Any]) -> str:
|
||
"""
|
||
将节点信息格式化为供LLM理解的字符串(仅包含标签和中文名称)
|
||
|
||
Args:
|
||
node: 节点信息字典
|
||
|
||
Returns:
|
||
str: 格式化后的字符串
|
||
"""
|
||
name = node.get("name", "未知")
|
||
labels = ", ".join(node.get("labels", []))
|
||
rel = node.get("relationship", {})
|
||
|
||
parts = [f"节点名称: {name}"]
|
||
if labels:
|
||
parts.append(f"节点类型: {labels}")
|
||
if rel:
|
||
parts.append(f"通过关系: {rel.get('type', '未知')} 连接自: {rel.get('source_name', '未知')}")
|
||
|
||
return "\n".join(parts)
|
||
|
||
|
||
async def _select_best_node_with_llm(
|
||
query: str,
|
||
neighbors: List[Dict[str, Any]]
|
||
) -> Optional[Dict[str, Any]]:
|
||
"""
|
||
使用大模型从一跳邻居中选择最合适的节点
|
||
|
||
Args:
|
||
query: 用户查询
|
||
neighbors: 一跳邻居节点列表
|
||
|
||
Returns:
|
||
Optional[Dict]: 最合适的节点信息,None表示没有合适的
|
||
"""
|
||
if not neighbors:
|
||
return None
|
||
|
||
from modelsAPI.model_api import OpenaiAPI
|
||
|
||
# 格式化邻居节点信息供LLM参考
|
||
neighbors_text = []
|
||
for i, node in enumerate(neighbors, 1):
|
||
neighbors_text.append(f"【候选节点 {i}】\n{_format_node_for_llm(node)}\n")
|
||
|
||
prompt = f"""你是一个智能图谱检索助手,需要根据用户的查询,从提供的候选节点中选择最相关的一个节点。
|
||
|
||
用户查询: {query}
|
||
|
||
候选节点列表:
|
||
{''.join(neighbors_text)}
|
||
|
||
请仔细阅读用户查询和所有候选节点信息,选择与用户查询最相关的一个节点。
|
||
|
||
输出要求:
|
||
1. 请仅输出你选择的候选节点编号(1, 2, 3...),不要输出任何其他内容
|
||
2. 如果没有任何相关的节点,请输出"0"
|
||
3. 不要解释原因,只输出数字编号
|
||
|
||
你的选择:"""
|
||
|
||
try:
|
||
model = os.getenv("OPENAI_MODEL")
|
||
response = await OpenaiAPI.open_api_chat_async(prompt, model, temperature=0)
|
||
|
||
# 提取数字编号
|
||
match = re.search(r'\d+', response.strip())
|
||
if match:
|
||
selected_index = int(match.group()) - 1
|
||
if 0 <= selected_index < len(neighbors):
|
||
logger.info(f"LLM选择了候选节点 {selected_index + 1}: {neighbors[selected_index].get('name', '未知')}")
|
||
return neighbors[selected_index]
|
||
|
||
logger.warning(f"LLM未选择有效节点,原始响应: {response}")
|
||
return None
|
||
except Exception as e:
|
||
logger.error(f"使用LLM选择节点失败: {e}", exc_info=True)
|
||
return None
|
||
|
||
|
||
def _build_graph_result(
|
||
entry_node: Optional[Dict[str, Any]],
|
||
best_node: Optional[Dict[str, Any]],
|
||
query: str
|
||
) -> Dict[str, Any]:
|
||
"""
|
||
构建与operate方法一致的输出格式(只包含入口节点和选中节点)
|
||
|
||
Args:
|
||
entry_node: 入口节点
|
||
best_node: 选择的最佳节点
|
||
query: 用户查询
|
||
|
||
Returns:
|
||
Dict: 格式化后的结果
|
||
"""
|
||
exclude_keys = {
|
||
"last_updated",
|
||
"created_at",
|
||
"knowledge_source",
|
||
"fulltext",
|
||
"embedding",
|
||
}
|
||
|
||
nodes_g: List[Dict[str, Any]] = []
|
||
links_g: List[Dict[str, Any]] = []
|
||
node_id_set = set()
|
||
|
||
# 添加入口节点
|
||
if entry_node:
|
||
nid = entry_node.get("id")
|
||
if nid and nid not in node_id_set:
|
||
node_id_set.add(nid)
|
||
|
||
props = entry_node.get("props", {})
|
||
props_filtered = {k: v for k, v in props.items() if k not in exclude_keys}
|
||
|
||
name = props.get("名称") or props.get("name") or str(nid)
|
||
labels = entry_node.get("labels", [])
|
||
|
||
nodes_g.append({
|
||
"id": nid,
|
||
"name": name,
|
||
"labels": labels,
|
||
"properties": props_filtered,
|
||
})
|
||
|
||
# 添加最佳节点和关系
|
||
if best_node:
|
||
# 添加最佳节点
|
||
best_nid = best_node.get("id")
|
||
if best_nid and best_nid not in node_id_set:
|
||
node_id_set.add(best_nid)
|
||
|
||
props = best_node.get("properties", {})
|
||
name = best_node.get("name") or str(best_nid)
|
||
labels = best_node.get("labels", [])
|
||
|
||
nodes_g.append({
|
||
"id": best_nid,
|
||
"name": name,
|
||
"labels": labels,
|
||
"properties": props,
|
||
})
|
||
|
||
# 添加关系
|
||
rel = best_node.get("relationship", {})
|
||
if rel:
|
||
rel_id = rel.get("id")
|
||
src_id = rel.get("source_id")
|
||
tgt_id = best_nid
|
||
|
||
if src_id in node_id_set and tgt_id in node_id_set:
|
||
links_g.append({
|
||
"id": rel_id,
|
||
"label": rel.get("type", ""),
|
||
"source": src_id,
|
||
"target": tgt_id,
|
||
"properties": rel.get("properties", {}),
|
||
})
|
||
|
||
# 生成 results 文本(优先用best_node,其次用entry_node)
|
||
results_text = ""
|
||
if best_node:
|
||
best_props = best_node.get("properties", {})
|
||
best_info = {k: v for k, v in best_props.items() if k not in exclude_keys and v}
|
||
results_text = json.dumps(best_info, ensure_ascii=False) if best_info else ""
|
||
elif entry_node:
|
||
entry_props = entry_node.get("props", {})
|
||
entry_info = {k: v for k, v in entry_props.items() if k not in exclude_keys and v}
|
||
results_text = json.dumps(entry_info, ensure_ascii=False) if entry_info else ""
|
||
|
||
graph_item = {"nodes": nodes_g, "links": links_g, "results": results_text}
|
||
return {"success": True, "xxxx.graph": [graph_item]}
|
||
|
||
|
||
async def quick_graph_search(
|
||
query: str,
|
||
route_res: Optional[List[Dict[str, str]]] = None,
|
||
top_k: int = 10,
|
||
) -> Dict[str, Any]:
|
||
"""
|
||
临时图谱检索接口
|
||
|
||
Args:
|
||
query: 用户查询
|
||
route_res: 路由结果列表(本方案不使用)
|
||
top_k: 检索返回的节点数量上限
|
||
|
||
Returns:
|
||
Dict: 检索结果,包含 success, xxxx.graph 等字段
|
||
"""
|
||
start_time = time.time()
|
||
|
||
if not query or not str(query).strip():
|
||
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:
|
||
async with AsyncGraphDatabase.driver(uri, auth=(user, password)) as driver:
|
||
async with driver.session() as session:
|
||
logger.info(f"[QuickGraphSearch] 开始直接混合检索,查询: '{query}'")
|
||
|
||
# 1. 直接用用户问题混合检索所有节点
|
||
candidate_nodes = await _search_all_nodes_with_hybrid(session, query, top_k=20)
|
||
|
||
if not candidate_nodes:
|
||
return {"success": False, "error": "未找到任何匹配节点"}
|
||
|
||
logger.info(f"[QuickGraphSearch] 混合检索找到 {len(candidate_nodes)} 个候选节点")
|
||
|
||
# 2. Rerank 重排序
|
||
reranked_nodes = await _rerank_nodes(query, candidate_nodes)
|
||
|
||
if not reranked_nodes:
|
||
return {"success": False, "error": "Rerank失败"}
|
||
|
||
# 3. 取最匹配的节点作为入口
|
||
entry_node = reranked_nodes[0]
|
||
entry_id = entry_node.get("id")
|
||
logger.info(f"[QuickGraphSearch] 选择入口节点: {entry_node.get('name', '未知')}, score={entry_node.get('rerank_score', 0)}")
|
||
|
||
# 4. 获取一跳邻居
|
||
neighbors = await _get_one_hop_neighbors(session, entry_id)
|
||
|
||
logger.info(f"[QuickGraphSearch] 获取到 {len(neighbors)} 个一跳邻居")
|
||
|
||
if not neighbors:
|
||
# 如果没有邻居节点,只返回入口节点
|
||
return _build_graph_result(entry_node, None, query)
|
||
|
||
# 5. 用LLM选择最合适的节点
|
||
best_node = await _select_best_node_with_llm(query, neighbors)
|
||
|
||
# 6. 构建结果(只包含入口节点和选中节点)
|
||
result = _build_graph_result(entry_node, best_node, query)
|
||
|
||
elapsed_ms = (time.time() - start_time) * 1000
|
||
logger.info(f"[QuickGraphSearch] 完成,耗时: {elapsed_ms:.2f}ms")
|
||
|
||
return result
|
||
|
||
except Exception as e:
|
||
logger.error(f"QuickGraphSearch 失败: {e}", exc_info=True)
|
||
return {"success": False, "error": f"检索失败: {str(e)}"}
|
||
|