kgrag/graph_search/cypher_reset.py
2026-07-29 18:10:19 +08:00

601 lines
25 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.

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)