""" =========================================== 临时图谱检索方案 - 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)}"}