import os import httpx import json import requests import warnings from typing import List, Set, Tuple, Dict, Optional import ast warnings.filterwarnings("ignore", category=UserWarning, module='jieba._compat') import re import jieba import jieba.posseg as pseg # ← 新增:词性标注,用于提取名词/关键词 from neo4j import GraphDatabase from neo4j_graphrag.retrievers import HybridCypherRetriever from neo4j_graphrag.embeddings.base import Embedder from langchain_community.graphs.neo4j_graph import Neo4jGraph from openai import OpenAI # ================== 9. Prompt 模板 ================== CYPHER_PROMPT = """ 你是一名专业的Cypher语法编写助手。你需要根据用户的问题、给定的Schema(实体类型和关系)以及向量检索到的相关实体,编写一条精准的Cypher查询语句。 # 编写请严格遵守以下规则: 1. **语法正确性**:确保编写的Cypher语法完全正确,符合Neo4j标准,不能包含任何语法错误。 2. **入口节点**:根据检索到的结果,从结果选取入口实体,作为查询的起点,入口节点一定要和用户问题中描述的实体相关。例如入口节点为"发射控制单元",与问题"发射控制单元有哪些子设备"相关。 2. **路径完整性**:根据用户问题的语义,结合Schema中定义的关系,构建完整且逻辑通顺的查询路径(例如:设备 -> 维修工作 -> 故障模式 -> 维修项目)。 3. **最高分实体绑定原则(核心规则)**:在编写 `WHERE` 过滤条件时,必须严格遵循"最高分优先"原则。请仔细分析 `search_results`,**依次为路径中的每一个节点(如设备、维修工作、故障模式、维修项目等)匹配并绑定其对应的得分(Score)最高的实体名称**。从而构建出一条全链路得分最高的精准查询路径。 4. **无冗余查询**:只返回与用户问题直接相关的节点属性,不要包含无关的查询或返回字段。 5. 检索的相关实体关系是根据用户问题进行语义匹配,从而获取到的实体关系(标签表示实体类型,名称表示实体名称,得分表示该实体与问题的相关性),来源于图数据库真实保存,不能包含任何假设或推理。提供给你参考的实体类型和实体关系则主要是为了帮助你理解用户问题的语义,根据检索的实体关系构建Cypher查询语句,而不是为了直接查询数据库。 6. 根据给你的实体关系,仔细分析用户的问题,撰写一条精准的Cypher查询语句,语句中不要出现任何假设或推理,只能根据提供给你的实体关系进行查询。 #用户问题如下: {query} #实体类型和实体关系如下: {filtered_schema} #根据用户问题,从数据库中检索的相关实体如下: {search_results} 示例1: 用户问题:连杆大端轴承发生损坏异常该如何解决? 实体类型和实体关系如下: 实体: ["主要附件仪表的规格型号及性能参数", "使用期间安全防护及注意事项", "功能", "器材保障", "图册", "基本情况", "备品备件", "子系统", "安全保护措施及故障处理", "安全警告", "工作原理", "布置情况", "技术参数", "技术指标参数", "接口", "接口情况", "操作使用", "操作员与设备之间操作员与其他设备或系统设备员之间关系", "操作程序", "操作项目", "故障模式", "环境条件", "系统", "组成", "结构特点", "维修安全防护及注意事项汇总表", "维修工作", "维修项目", "维护保养", "舰船", "舰艇", "设备", "设计单位", "调试"] 关系: (:组成)-[:包含]->(:设备) (:维修工作)-[:包含故障]->(:故障模式) (:维修工作)-[:表征]->(:故障模式) (:维修项目)-[:修理]->(:故障模式) (:设备)-[:包含]->(:工作原理) (:设备)-[:包含]->(:布置情况) (:设备)-[:包含]->(:技术参数) (:设备)-[:包含]->(:接口情况) (:设备)-[:包含]->(:操作员与设备之间操作员与其他设备或系统设备员之间关系) (:设备)-[:包含]->(:组成) (:设备)-[:包含]->(:结构特点) (:设备)-[:包含]->(:维修安全防护及注意事项汇总表) (:设备)-[:包含]->(:维护保养) (:设备)-[:包括]->(:组成) (:设备)-[:发生故障使用维修工作]->(:维修工作) (:设备)-[:安全警告]->(:安全警告) (:设备)-[:实现]->(:功能) (:设备)-[:实现]->(:基本情况) (:设备)-[:支持操作]->(:操作使用) (:设备)-[:维修]->(:维修工作) (:设备)-[:设计]->(:设计单位) (:设备)-[:调试]->(:调试) (:设备)-[:适用]->(:环境条件) (:设备)-[:配备]->(:接口) (:设备)-[:配套]->(:图册) #根据用户问题,从数据库中找到的相关实体如下: {{'标签': ['维修项目'], '名称': '修复连杆大端轴承异常磨损', '得分': 1.0}} {{'标签': ['维修项目'], '名称': '更换连杆大端轴承', '得分': 0.9839502279586991}} {{'标签': ['维修工作'], '名称': '连杆大端轴承维修工作指导', '得分': 0.9601912347183726}} {{'标签': ['设备'], '名称': '连杆大端轴承', '得分': 0.9378418041605234}} {{'标签': ['维修项目'], '名称': '更换连杆小段轴承', '得分': 0.9370194814204257}} {{'标签': ['设备'], '名称': '连杆小端轴承', '得分': 0.9056004631162252}} {{'标签': ['故障模式'], '名称': '主轴承或连杆大端轴承磨损导致主机运行中振动异常增大', '得分': 0.890951410907792}} {{'标签': ['故障模式'], '名称': '主轴承或连杆大端轴承磨损导致机座与机架组件运行中振动异常增大', '得分': 0.8895577902640475}} {{'标签': ['设备'], '名称': '连杆小段轴承', '得分': 0.8785823193324869}} {{'标签': ['维修工作'], '名称': '传动轴的维修工作', '得分': 0.8768441389219146}} 输出为: MATCH (d:设备)-[:发生故障使用维修工作]->(w:维修工作)-[:包含故障]->(f:故障模式)<-[:修理]-(p:维修项目) WHERE d.名称 = '连杆大端轴承' AND w.名称 = '连杆大端轴承维修工作指导' AND f.名称 = '主轴承或连杆大端轴承磨损导致主机运行中振动异常增大' AND p.名称 = '修复连杆大端轴承异常磨损' RETURN w.名称 AS 维修工作名称, f.名称 AS 故障模式, p.名称 AS 维修项目名称 请根据上述规则、Schema定义以及检索到的相关实体,为用户问题"{query}"编写Cypher查询语句。 输出为: """ def build_cypher_prompt(query: str, filtered_schema: str, search_results: list) -> str: return CYPHER_PROMPT.format( query=query, filtered_schema=filtered_schema, search_results=search_results, ) # ================== 10. LLM 调用 ================== def call_llm(prompt: str, model: str = "Qwen3.5-35B-A3B") -> str: response = OpenAI( api_key="none", base_url="http://192.168.0.46:59800/v1", ).chat.completions.create( model=model, messages=[ {"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": prompt}, ], temperature=0.1, stream=False, extra_body={"chat_template_kwargs": {"enable_thinking": False}}, ) return response.choices[0].message.content def extract_cypher(llm_output: str) -> str: match = re.search(r"(MATCH[\s\S]+?RETURN[^\n]+(?:\n[^\n]+)*)", llm_output) return match.group(1).strip() if match else llm_output.strip() # ================== 11. Cypher 诊断基础工具 ================== def _run_query(session, cypher: str, **params) -> list: try: return list(session.run(cypher, **params)) except Exception: return [] def _parse_alias_label(cypher: str) -> Dict[str, str]: return { alias: label.strip() for alias, label in re.findall(r"\((\w+):([^\)]+)\)", cypher) } def _parse_alias_name(cypher: str) -> Dict[str, str]: return { alias: name for alias, name in re.findall(r"(\w+)\.`?名称`?\s*=\s*'([^']+)'", cypher) } def _parse_segments(cypher: str) -> List[Tuple[str, str, str, str, str]]: return [ (sa.strip(), sl.strip(), r.strip(), ta.strip(), tl.strip()) for sa, sl, r, ta, tl in re.findall( r"\((\w+):([^\)]+)\)-\[:([^\]]+)\]->\((\w+):([^\)]+)\)", cypher ) ] def _candidates_for_label(search_results: list, label: str) -> list: return sorted( [r for r in search_results if label in r.get("标签", [])], key=lambda x: x.get("得分", 0), reverse=True, ) def _test_segment( session, src_label: str, src_name: Optional[str], rel: str, tgt_label: str, tgt_name: Optional[str], ) -> int: wheres = [] if src_name: wheres.append("a.`名称` = $src") if tgt_name: wheres.append("b.`名称` = $tgt") where_clause = ("WHERE " + " AND ".join(wheres)) if wheres else "" q = f"MATCH (a:{src_label})-[:{rel}]->(b:{tgt_label}) {where_clause} RETURN count(*) AS c" rows = _run_query(session, q, src=src_name, tgt=tgt_name) return rows[0]["c"] if rows else 0 def _find_verified_replacement( session, src_label: str, src_name: Optional[str], rel: str, tgt_label: str, candidates: list, verbose: bool, ) -> Optional[str]: for i, cand in enumerate(candidates): cand_name = cand.get("名称", "") if not cand_name: continue cnt = _test_segment(session, src_label, src_name, rel, tgt_label, cand_name) icon = "✅" if cnt > 0 else " " if verbose: print(f" [{i+1}] {icon} '{cand_name}' (得分={cand.get('得分',0):.4f}) → 匹配:{cnt}") if cnt > 0: return cand_name return None # ================== 11b. 关系回退 + 最短路径工具 ================== def _relations_between_labels( triples: List[Tuple[str, str, str]], src_label: str, tgt_label: str, ) -> List[str]: return [ rel for s, rel, t in triples if s == src_label and t == tgt_label ] def _test_segment_with_rel( session, src_label: str, src_name: Optional[str], rel: str, tgt_label: str, tgt_name: Optional[str], ) -> int: wheres = [] if src_name: wheres.append("a.`名称` = $src") if tgt_name: wheres.append("b.`名称` = $tgt") where_clause = ("WHERE " + " AND ".join(wheres)) if wheres else "" q = ( f"MATCH (a:{src_label})-[:{rel}]->(b:{tgt_label}) " f"{where_clause} RETURN count(*) AS c" ) rows = _run_query(session, q, src=src_name, tgt=tgt_name) return rows[0]["c"] if rows else 0 def _find_shortest_path( session, src_label: str, src_name: str, tgt_label: str, tgt_name: str, max_hops: int = 4, verbose: bool = True, ) -> Optional[List[dict]]: if tgt_name: q = f""" MATCH p = allShortestPaths( (a:{src_label})-[*1..{max_hops}]->(b:{tgt_label}) ) WHERE a.`名称` = $src AND b.`名称` = $tgt RETURN [r IN relationships(p) | type(r)] AS rels, [n IN nodes(p) | labels(n)[0]] AS node_labels, [n IN nodes(p) | n.`名称`] AS node_names LIMIT 1 """ rows = _run_query(session, q, src=src_name, tgt=tgt_name) else: q = f""" MATCH p = allShortestPaths( (a:{src_label})-[*1..{max_hops}]->(b:{tgt_label}) ) WHERE a.`名称` = $src RETURN [r IN relationships(p) | type(r)] AS rels, [n IN nodes(p) | labels(n)[0]] AS node_labels, [n IN nodes(p) | n.`名称`] AS node_names LIMIT 1 """ rows = _run_query(session, q, src=src_name) if not rows: if verbose: print(f" ⚠️ allShortestPaths: 未找到路径 " f"({src_label}:'{src_name}') → ({tgt_label}:'{tgt_name}')") return None row = rows[0] rels, labels_, names = row["rels"], row["node_labels"], row["node_names"] if verbose: path_str = " → ".join( f"[:{r}]→({labels_[i+1]}:'{names[i+1]}')" for i, r in enumerate(rels) ) print(f" ✅ 最短路径: ({src_label}:'{src_name}') → {path_str}") return [ {"rel": rel, "tgt_label": labels_[i + 1], "tgt_name": names[i + 1]} for i, rel in enumerate(rels) ] def _rebuild_cypher_from_path( src_alias: str, src_label: str, src_name: str, steps: List[dict], return_clause: str, ) -> str: aliases = [src_alias] + [f"n{i+1}" for i in range(len(steps))] labels_ = [src_label] + [s["tgt_label"] for s in steps] names = [src_name] + [s["tgt_name"] for s in steps] rels = [s["rel"] for s in steps] path_parts = [f"({aliases[0]}:{labels_[0]})"] for i, rel in enumerate(rels): path_parts.append(f"-[:{rel}]->({aliases[i+1]}:{labels_[i+1]})") match_clause = "MATCH " + "".join(path_parts) where_parts = [ f"{alias}.`名称` = '{name}'" for alias, name in zip(aliases, names) if name ] where_clause = "WHERE " + "\n AND ".join(where_parts) if where_parts else "" return f"{match_clause}\n{where_clause}\nRETURN {return_clause}" def strip_unverified_name_conditions(cypher: str, search_results: list) -> str: valid_names = {r["名称"] for r in search_results} pattern = re.compile(r"(\w+)\.`?名称`?\s*=\s*'([^']+)'") def remove_condition(cypher_str: str, alias: str, name: str) -> str: patterns_to_try = [ re.compile( r"\s*AND\s+" + re.escape(alias) + r"\.`?名称`?\s*=\s*'" + re.escape(name) + r"'", re.IGNORECASE, ), re.compile( re.escape(alias) + r"\.`?名称`?\s*=\s*'" + re.escape(name) + r"'\s*AND\s*", re.IGNORECASE, ), re.compile( r"WHERE\s+" + re.escape(alias) + r"\.`?名称`?\s*=\s*'" + re.escape(name) + r"'", re.IGNORECASE, ), ] for p in patterns_to_try: new_str = p.sub("", cypher_str) if new_str != cypher_str: new_str = re.sub(r"WHERE\s*(RETURN|$)", r"\1", new_str, flags=re.IGNORECASE) return new_str.strip() return cypher_str stripped = cypher for match in pattern.finditer(cypher): alias, name = match.group(1), match.group(2) if name not in valid_names: print(f"⚠️ 名称 '{name}' 不在检索结果中,移除条件: {alias}.名称 = '{name}'") stripped = remove_condition(stripped, alias, name) return stripped # ================== 12. 端到端链路诊断与修复 ================== def diagnose_and_fix_cypher( cypher: str, driver, search_results: list, schema_triples: List[Tuple[str, str, str]] = None, verbose: bool = True, ) -> dict: alias_label = _parse_alias_label(cypher) alias_name = _parse_alias_name(cypher) segments = _parse_segments(cypher) current_names = dict(alias_name) return_match = re.search(r"RETURN\s+(.+)$", cypher, re.IGNORECASE | re.DOTALL) return_clause = return_match.group(1).strip() if return_match else "*" if verbose: print("\n" + "=" * 60) print("🩺 端到端链路诊断(Pass A: 换名 / Pass B: 换关系 / Pass C: 最短路径)") print("=" * 60) print(f"📝 原始 Cypher:\n{cypher}\n") print(f"🗂️ 初始名称绑定: {current_names}\n") with driver.session() as session: if _run_query(session, cypher): if verbose: print("✅ 完整查询有结果,无需修复。") return { "status": "ok", "chain_verified": True, "diagnosis": "查询正常,有结果返回。", "broken_segments": [], "fixed_cypher": None, "final_names": current_names, } if verbose: print("⚠️ 完整查询结果为空,开始逐段诊断...\n") broken_segments: list = [] segment_reports: list = [] for seg_idx, (src_alias, src_label, rel, tgt_alias, tgt_label) in enumerate(segments): src_name = current_names.get(src_alias) tgt_name = current_names.get(tgt_alias) seg_desc = ( f"段[{seg_idx+1}] " f"({src_alias}:{src_label} '{src_name}')" f" -[:{rel}]-> " f"({tgt_alias}:{tgt_label} '{tgt_name}')" ) if verbose: print(f" 验证 {seg_desc}") cnt = _test_segment(session, src_label, src_name, rel, tgt_label, tgt_name) if cnt > 0: if verbose: print(f" ✅ 通过 (匹配数: {cnt})\n") segment_reports.append(f"✅ {seg_desc}") continue # Pass A if verbose: print(f" ❌ 断链!Pass A: 尝试替换节点名称,保持关系 [:{rel}] 不变...") cands = _candidates_for_label(search_results, tgt_label) verified = _find_verified_replacement( session, src_label, src_name, rel, tgt_label, cands, verbose ) if verified: if verbose: print(f" ✅ Pass A 修复: '{tgt_name}' → '{verified}'\n") current_names[tgt_alias] = verified segment_reports.append( f"🔧 {seg_desc}\n → Pass A: 节点名 '{tgt_name}' 替换为 '{verified}'" ) broken_segments.append({ "seg_idx": seg_idx + 1, "pass": "A", "tgt_alias": tgt_alias, "tgt_label": tgt_label, "original": tgt_name, "fixed": verified, }) continue # Pass B if verbose: print(f" ❌ Pass A 失败。Pass B: 遍历 schema 中 {src_label}→{tgt_label} 的所有关系...") alt_rels = ( _relations_between_labels(schema_triples, src_label, tgt_label) if schema_triples else [] ) alt_rels = [r for r in alt_rels if r != rel] fixed_rel = None fixed_tgt = tgt_name for alt_rel in alt_rels: cnt2 = _test_segment_with_rel( session, src_label, src_name, alt_rel, tgt_label, tgt_name ) if verbose: icon = "✅" if cnt2 > 0 else " " print(f" {icon} rel='{alt_rel}', tgt='{tgt_name}' → 匹配:{cnt2}") if cnt2 > 0: fixed_rel = alt_rel break for cand in cands: cname = cand.get("名称", "") if not cname or cname == tgt_name: continue cnt3 = _test_segment_with_rel( session, src_label, src_name, alt_rel, tgt_label, cname ) if verbose: icon = "✅" if cnt3 > 0 else " " print(f" {icon} rel='{alt_rel}', tgt='{cname}' → 匹配:{cnt3}") if cnt3 > 0: fixed_rel = alt_rel fixed_tgt = cname break if fixed_rel: break if fixed_rel: cypher = cypher.replace(f"[:{rel}]", f"[:{fixed_rel}]", 1) segments[seg_idx] = (src_alias, src_label, fixed_rel, tgt_alias, tgt_label) if fixed_tgt != tgt_name: current_names[tgt_alias] = fixed_tgt if verbose: print( f" ✅ Pass B 修复: " f"rel '[:{rel}]' → '[:{fixed_rel}]'" + (f", 节点名 '{tgt_name}' → '{fixed_tgt}'" if fixed_tgt != tgt_name else "") + "\n" ) segment_reports.append( f"🔧 {seg_desc}\n" f" → Pass B: rel '[:{rel}]'→'[:{fixed_rel}]'" + (f", 节点名 '{tgt_name}'→'[:{fixed_tgt}]'" if fixed_tgt != tgt_name else "") ) broken_segments.append({ "seg_idx": seg_idx + 1, "pass": "B", "original_rel": rel, "fixed_rel": fixed_rel, "tgt_alias": tgt_alias, "original": tgt_name, "fixed": fixed_tgt, }) continue # Pass C if verbose: print(f" ❌ Pass B 失败。Pass C: allShortestPaths 探测真实路径...") first_alias = segments[0][0] first_label = alias_label.get(first_alias, segments[0][1]) first_name = current_names.get(first_alias, "") sp_steps = None if first_name: sp_steps = _find_shortest_path( session, src_label=first_label, src_name=first_name, tgt_label=tgt_label, tgt_name=tgt_name or "", max_hops=4, verbose=verbose, ) if sp_steps: new_cypher = _rebuild_cypher_from_path( src_alias=first_alias, src_label=first_label, src_name=first_name, steps=sp_steps, return_clause=return_clause, ) if verbose: print(f" ✅ Pass C 重建 Cypher:\n{new_cypher}\n") if _run_query(session, new_cypher): if verbose: print("=" * 60) print("✅ Pass C 链路修复成功(最短路径)") print(f"\n💡 修正后的 Cypher:\n{new_cypher}") print("=" * 60) return { "status": "fixed_by_shortest_path", "chain_verified": True, "diagnosis": "\n".join(segment_reports) + "\n✅ Pass C: 最短路径修复", "broken_segments": broken_segments, "fixed_cypher": new_cypher, "final_names": current_names, } if verbose: print(f" ❌ Pass C 失败,三种策略均无法修复此段。\n") segment_reports.append(f"❌ {seg_desc}\n → Pass A/B/C 均失败") broken_segments.append({ "seg_idx": seg_idx + 1, "pass": "failed", "tgt_alias": tgt_alias, "tgt_label": tgt_label, "original": tgt_name, "fixed": None, }) fixed_cypher = cypher for alias, original in alias_name.items(): new = current_names.get(alias) if new and new != original: fixed_cypher = fixed_cypher.replace(f"'{original}'", f"'{new}'") chain_ok = bool(_run_query(session, fixed_cypher)) if verbose: print("=" * 60) print("📋 逐段诊断报告:") for line in segment_reports: print(f" {line}") print(f"\n🗂️ 最终名称绑定: {current_names}") status = "✅ 链路修复成功" if chain_ok else "⚠️ 仍有断链段(建议扩大 top_k 或检查数据)" print(f"\n{status}") print(f"\n💡 修正后的 Cypher:\n{fixed_cypher}") print("=" * 60) return { "status": "fixed" if chain_ok else "partial", "chain_verified": chain_ok, "diagnosis": "\n".join(segment_reports), "broken_segments": broken_segments, "fixed_cypher": fixed_cypher, "final_names": current_names, } def get_response(query,filtered_schema,search_results,triples,driver): prompt = build_cypher_prompt(query, filtered_schema, search_results) print(prompt) print("=" * 60) llm_output = call_llm(prompt) cypher = extract_cypher(llm_output) cypher = strip_unverified_name_conditions(cypher, search_results) print("\n" + "=" * 60) print("🤖 LLM 生成的 Cypher:") print("=" * 60) print(cypher) # ── 阶段五:端到端链路诊断与修复 ───────────────────────────────── diag = diagnose_and_fix_cypher( cypher=cypher, driver=driver, search_results=search_results, schema_triples=triples, verbose=True, ) # ── 阶段六:执行最终 Cypher 并输出结果 ─────────────────────────── final_cypher = diag["fixed_cypher"] if diag["fixed_cypher"] else cypher print("\n" + "=" * 60) print("🚀 执行最终 Cypher:") print("=" * 60) print(final_cypher)