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

1144 lines
45 KiB
Python
Raw Permalink 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
from modelsAPI.model_api import OpenaiAPI
# ================== 1. 本地 bge-m3 嵌入模型 ==================
class LocalBgeM3Embeddings(Embedder):
def __init__(self, base_url: str, api_key: str, model_name: str = "bge-m3"):
self.base_url = base_url.rstrip("/")
self.api_key = api_key
self.model_name = model_name
def embed_query(self, text: str) -> List[float]:
return self.embed_documents([text])[0]
def embed_documents(self, texts: List[str]) -> List[List[float]]:
embedding_base_url = os.path.join(self.base_url, "embeddings")
with httpx.Client(timeout=60) as client:
response = client.post(
embedding_base_url,
json={"model": self.model_name, "input": texts},
headers={"Authorization": f"Bearer {self.api_key}"},
)
embedding_list = [item["embedding"] for item in response.json()["data"]]
return embedding_list
# ================== 2. 基础配置 ==================
os.environ["LOCAL_EMBEDDING_BASE"] = "http://192.168.0.46:59700/v1"
os.environ["LOCAL_API_KEY"] = "none"
URI = "bolt://192.168.0.46:57687"
AUTH = ("neo4j", "zdht123@")
NAME_PROPERTY = "名称"
FULLTEXT_PROPERTY = "fulltext" # ← 新增fulltext 字段名
EMBEDDING_PROPERTY = "embedding"
SEARCH_LABEL = "Searchable"
VECTOR_INDEX_NAME = "searchable_vector" # 向量索引embedding 字段)
FULLTEXT_INDEX_NAME = "global_searchable_content_search" # 名称全文索引(名称 字段)
FULLTEXT_FIELD_INDEX_NAME = "searchable_fulltext" # fulltext 字段全文索引
EXCLUDED_BUSINESS_LABELS: Set[str] = {
"Entity", "Chunk", "Document", "_Bloom_Perspective_", SEARCH_LABEL,
}
VECTOR_DIMENSION = 1024
_JIEBA_POS_WHITELIST = {"n", "nr", "ns", "nt", "nz", "vn", "eng"}
# ================== 5. 检索结果解析辅助 ==================
def parse_record_string(s: str):
name_match = re.search(r"name='([^']*)'", s)
labels_match = re.search(r"labels=(\[[^\]]*\])", s)
score_match = re.search(r"score=([\d.]+)", s)
name = name_match.group(1) if name_match else "未知名称"
labels = ["未知标签"]
if labels_match:
try:
labels = ast.literal_eval(labels_match.group(1))
except (ValueError, SyntaxError):
pass
score = float(score_match.group(1)) if score_match else 0.0
return name, labels, score
def _build_retriever(driver, embedder) -> HybridCypherRetriever:
return HybridCypherRetriever(
driver=driver,
vector_index_name=VECTOR_INDEX_NAME,
fulltext_index_name=FULLTEXT_INDEX_NAME,
embedder=embedder,
retrieval_query=f"""
RETURN
node.`{NAME_PROPERTY}` AS name,
[lbl IN labels(node) WHERE lbl <> '{SEARCH_LABEL}'] AS labels,
score
"""
)
# ================== 6. Schema 解析与 N 跳标签扩展 ==================
def parse_schema_relationships(schema_str: str) -> List[Tuple[str, str, str]]:
pattern = re.compile(r"\(:([^)]+)\)-\[:([^\]]+)\]->\(:([^)]+)\)")
return [
(m.group(1).strip(), m.group(2).strip(), m.group(3).strip())
for m in pattern.finditer(schema_str)
]
def expand_labels_n_hops(
seed_labels: Set[str],
triples: List[Tuple[str, str, str]],
hops: int = 3,
include_reverse: bool = True,
skip_labels: Set[str] = None,
) -> Set[str]:
if skip_labels is None:
skip_labels = set()
adj: Dict[str, Set[str]] = {}
for src, _rel, tgt in triples:
if src in skip_labels or tgt in skip_labels:
continue
adj.setdefault(src, set()).add(tgt)
if include_reverse:
adj.setdefault(tgt, set()).add(src)
visited: Set[str] = set(seed_labels)
frontier: Set[str] = set(seed_labels) - skip_labels
for _ in range(hops):
next_frontier: Set[str] = set()
for node in frontier:
for nbr in adj.get(node, ()):
if nbr not in visited:
next_frontier.add(nbr)
if not next_frontier:
break
visited |= next_frontier
frontier = next_frontier
return visited
def build_compact_schema(
triples: List[Tuple[str, str, str]],
keep_labels: Set[str],
) -> str:
kept_triples = [
(src, rel, tgt) for src, rel, tgt in triples
if src in keep_labels and tgt in keep_labels
]
entities_in_use: Set[str] = set()
for src, _rel, tgt in kept_triples:
entities_in_use.add(src)
entities_in_use.add(tgt)
all_entities = sorted(keep_labels | entities_in_use)
entities_str = "[" + ", ".join(f'"{e}"' for e in all_entities) + "]"
unique_triples = sorted(set(kept_triples))
rels_lines = [f"(:{src})-[:{rel}]->(:{tgt})" for src, rel, tgt in unique_triples]
return "\n".join([f"实体: {entities_str}", "", "关系:"] + rels_lines)
# ================== 7. Rerank 重排序 ==================
def rerank_query(query: str, documents: List[str], top_n: int = 3) -> List[dict]:
headers = {
"Content-Type": "application/json",
"Authorization": f'Bearer {os.getenv("OPENAI_API_KEY", "none")}',
}
data = {
"model": "bge-rerank",
"query": query,
"top_n": top_n,
"documents": documents,
}
response = requests.post(
"http://192.168.0.46:59600/v1/rerank",
headers=headers,
data=json.dumps(data),
)
if response.status_code != 200:
raise Exception(f"Rerank Error {response.status_code}: {response.text}")
result_json = response.json()
results_list = []
for result in result_json["results"]:
results_list.append({
"index": result["index"],
"text": result["document"]["text"],
"scores": result["relevance_score"],
})
return results_list
def apply_rerank(
query: str,
search_results: List[Dict],
top_n: int = 5,
verbose: bool = True,
) -> List[Dict]:
"""
用 rerank 对 search_results 重新排序,返回重排后的 search_results 子集。
来源字段source会被透传保留方便追溯是哪一路检索命中的。
"""
if not search_results:
return search_results
# rerank 文档内容 = 名称,若有 fulltext 片段则拼接进去提升语义匹配质量
documents = [
(r.get("名称", "") + " " + r.get("fulltext_snippet", "")).strip()
for r in search_results
]
if verbose:
print(f"\n🔀 开始 Reranktop_n={top_n}...")
try:
reranked = rerank_query(query=query, documents=documents, top_n=top_n)
except Exception as e:
print(f"⚠️ Rerank 调用失败,跳过重排序,保持原顺序。原因: {e}")
return search_results[:top_n]
reranked_results = []
for item in reranked:
original = search_results[item["index"]]
reranked_results.append({
"标签": original["标签"],
"名称": original["名称"],
"得分": item["scores"],
"向量得分": original["得分"],
"来源": original.get("来源", "hybrid"), # ← 保留来源标注
"fulltext_snippet": original.get("fulltext_snippet", ""),
})
reranked_results.sort(key=lambda x: x["得分"], reverse=True)
if verbose:
print(f"✅ Rerank 完成,保留 {len(reranked_results)} 条结果:")
for i, r in enumerate(reranked_results, 1):
print(
f" [{i:2d}] rerank={r['得分']:.4f} | "
f"向量={r['向量得分']:.4f} | "
f"来源={r['来源']:<12} | "
f"{r['名称']} 标签={r['标签']}"
)
return reranked_results
# ================== 8-NEW. jieba 关键词提取 ==================
def extract_keywords_jieba(
text: str,
pos_whitelist: Set[str] = _JIEBA_POS_WHITELIST,
min_len: int = 2,
max_keywords: int = 8,
) -> List[str]:
"""
用 jieba 词性标注提取名词/关键词,去除停用词和过短词条。
返回:去重后的关键词列表(保持出现顺序)。
"""
words_pos = pseg.cut(text)
seen: Set[str] = set()
keywords: List[str] = []
for word, flag in words_pos:
w = word.strip()
# 词性白名单过滤 + 长度过滤 + 去重
if (
len(w) >= min_len
and any(flag.startswith(p) for p in pos_whitelist)
and w not in seen
):
seen.add(w)
keywords.append(w)
if len(keywords) >= max_keywords:
break
if not keywords:
# 兜底:直接分词,取长度 >= min_len 的词
keywords = list(dict.fromkeys(
w for w in jieba.cut(text) if len(w.strip()) >= min_len
))[:max_keywords]
return keywords
# ================== 8-NEW. fulltext 字段关键词检索 ==================
def search_fulltext_field(
query: str,
driver,
top_k: int = 10,
min_keyword_len: int = 2,
max_keywords: int = 8,
score_threshold: float = 0.0,
verbose: bool = True,
) -> List[Dict]:
"""
对节点的 `fulltext` 属性字段做关键词全文检索。
流程:
1. jieba 分词 → 提取名词/关键词
2. 拼接成 Lucene 查询式(空格分隔 = OR可按需改成 AND
3. 调用 Neo4j 全文索引 FULLTEXT_FIELD_INDEX_NAME
4. 返回与 search_results 格式一致的列表,附加 来源='fulltext' 标注
Returns:
List[Dict],每项含: 标签、名称、得分、来源、fulltext_snippet
"""
keywords = extract_keywords_jieba(
query,
pos_whitelist=_JIEBA_POS_WHITELIST,
min_len=min_keyword_len,
max_keywords=max_keywords,
)
if not keywords:
if verbose:
print("⚠️ jieba 未提取到有效关键词,跳过 fulltext 字段检索。")
return []
# Lucene 查询:各关键词 OR 连接;若需要更精确可改为 " AND ".join(...)
lucene_query = " OR ".join(keywords)
if verbose:
print(f"\n🔑 jieba 提取关键词: {keywords}")
print(f"📄 Lucene 查询式: {lucene_query}")
print(f"🔍 正在对 fulltext 字段做全文检索...")
cypher = f"""
CALL db.index.fulltext.queryNodes(
'{FULLTEXT_FIELD_INDEX_NAME}',
$lucene_query
)
YIELD node, score
WHERE score > $score_threshold
RETURN
node.`{NAME_PROPERTY}` AS name,
[lbl IN labels(node)
WHERE lbl <> '{SEARCH_LABEL}'] AS labels,
score,
left(node.`{FULLTEXT_PROPERTY}`, 200) AS snippet -- 截取前 200 字符作为片段
ORDER BY score DESC
LIMIT $top_k
"""
results: List[Dict] = []
try:
with driver.session() as session:
rows = list(session.run(
cypher,
lucene_query=lucene_query,
score_threshold=score_threshold,
top_k=top_k,
))
for row in rows:
results.append({
"标签": row["labels"] or ["未知标签"],
"名称": row["name"] or "未知名称",
"得分": row["score"],
"来源": "fulltext", # ← 标注来源
"fulltext_snippet": row["snippet"] or "",
})
if verbose:
print(f"✅ fulltext 字段检索完成,命中 {len(results)} 条:")
for i, r in enumerate(results, 1):
print(
f" [{i:2d}] score={r['得分']:.4f} | "
f"{r['名称']} 标签={r['标签']}"
)
except Exception as e:
print(f"⚠️ fulltext 字段检索出错(可能索引尚未就绪): {e}")
return results
# ================== 8. 检索(原有 hybrid + 新增 fulltext 合并) ==================
def _merge_search_results(
hybrid_results: List[Dict],
fulltext_results: List[Dict],
verbose: bool = True,
) -> List[Dict]:
"""
合并两路检索结果:
- hybrid_results 来源标注为 'hybrid'
- fulltext_results 来源已标注为 'fulltext'
去重策略:以「名称」为 key保留得分更高的那条来源冲突时记为 'hybrid+fulltext'
"""
merged: Dict[str, Dict] = {}
for r in hybrid_results:
r.setdefault("来源", "hybrid")
r.setdefault("fulltext_snippet", "")
key = r["名称"]
if key not in merged or r["得分"] > merged[key]["得分"]:
merged[key] = r
for r in fulltext_results:
key = r["名称"]
if key not in merged:
merged[key] = r
else:
existing = merged[key]
# 同一节点被两路命中:保留最高分,来源合并
if r["得分"] > existing["得分"]:
existing["得分"] = r["得分"]
if existing["来源"] != r["来源"]:
existing["来源"] = "hybrid+fulltext"
# 补充 fulltext_snippethybrid 路通常没有此字段)
if not existing.get("fulltext_snippet") and r.get("fulltext_snippet"):
existing["fulltext_snippet"] = r["fulltext_snippet"]
result_list = list(merged.values())
result_list.sort(key=lambda x: x["得分"], reverse=True)
if verbose:
hybrid_cnt = sum(1 for r in result_list if r["来源"] == "hybrid")
fulltext_cnt = sum(1 for r in result_list if r["来源"] == "fulltext")
both_cnt = sum(1 for r in result_list if r["来源"] == "hybrid+fulltext")
print(
f"\n🔀 合并结果: 共 {len(result_list)}"
f"(hybrid={hybrid_cnt}, fulltext={fulltext_cnt}, 两路命中={both_cnt})"
)
return result_list
def search(
sentence: str,
driver,
embedder,
top_k: int = 10,
hops: int = 3,
include_reverse: bool = True,
fulltext_top_k: int = 10, # ← 新增参数fulltext 检索条数
fulltext_score_threshold: float = 0.0, # ← 新增参数fulltext 分数阈值
verbose: bool = True,
) -> Dict:
if verbose:
print(f"\n🔍 正在对整句话进行混合检索:「{sentence}\n")
# ── 原有 hybrid 检索 ─────────────────────────────────────────────────────
retriever = _build_retriever(driver, embedder)
hybrid_results: List[Dict] = []
try:
results = retriever.search(query_text=sentence, top_k=top_k)
for item in results.items:
item_content = item.content
name = "未知名称"; labels = ["未知标签"]; score = 0.0
if hasattr(item_content, "keys"):
if "name" in item_content.keys(): name = item_content["name"]
if "labels" in item_content.keys(): labels = item_content["labels"]
if "score" in item_content.keys(): score = item_content["score"]
elif isinstance(item_content, str):
if item_content.startswith("<Record"):
name, labels, score = parse_record_string(item_content)
else:
try:
parsed = ast.literal_eval(item_content)
if isinstance(parsed, dict):
name = parsed.get("name", name)
labels = parsed.get("labels", labels)
score = parsed.get("score", score)
except (ValueError, SyntaxError):
name = item_content
if score == 0.0 and hasattr(item, "metadata") and item.metadata:
score = item.metadata.get("score", score)
hybrid_results.append({
"标签": labels, "名称": name, "得分": score,
"来源": "hybrid", "fulltext_snippet": "",
})
except Exception as e:
print(f"hybrid 检索时出错: {e}")
import traceback; traceback.print_exc()
# ── 新增fulltext 字段关键词检索 ────────────────────────────────────────
ft_results = search_fulltext_field(
query=sentence,
driver=driver,
top_k=fulltext_top_k,
score_threshold=fulltext_score_threshold,
verbose=verbose,
)
# ── 合并两路结果 ──────────────────────────────────────────────────────────
search_results = _merge_search_results(hybrid_results, ft_results, verbose=verbose)
# ── Schema 扩展(与原逻辑相同)────────────────────────────────────────────
EXCLUDE_LABELS: Set[str] = EXCLUDED_BUSINESS_LABELS.copy()
seed_labels: Set[str] = set()
for r in search_results:
for lbl in r.get("标签", []):
if lbl and lbl != "未知标签" and lbl not in EXCLUDE_LABELS:
seed_labels.add(lbl)
_neo4j_graph = Neo4jGraph(url=URI, username=AUTH[0], password=AUTH[1], enhanced_schema=True)
full_schema = _neo4j_graph.schema
triples = parse_schema_relationships(full_schema)
kept_labels = expand_labels_n_hops(
seed_labels=seed_labels, triples=triples,
hops=hops, include_reverse=include_reverse, skip_labels=EXCLUDE_LABELS,
)
kept_labels -= EXCLUDE_LABELS
filtered_schema = build_compact_schema(triples, kept_labels)
return {
"search_results": search_results,
"seed_labels": seed_labels,
"kept_labels": kept_labels,
"filtered_schema": filtered_schema,
"triples": triples,
}
# ================== 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,
}
# ================== 13. 主程序 ==================
if __name__ == "__main__":
driver = GraphDatabase.driver(URI, auth=AUTH)
embedder = LocalBgeM3Embeddings(
base_url=os.getenv("LOCAL_EMBEDDING_BASE"),
api_key=os.getenv("LOCAL_API_KEY"),
)
query = "发动机中喷油器工作不一致导致机座与机架组件功率输出不稳定如何解决"
try:
# ── 阶段二三路混合检索hybrid + fulltext 字段rerank 前候选池) ──
search_result = search(
sentence=query,
driver=driver,
embedder=OpenaiAPI.batch_embeddings,
top_k=40, # hybrid 检索条数
hops=3,
include_reverse=True,
fulltext_top_k=15, # ← fulltext 字段检索条数
fulltext_score_threshold=0.0,
verbose=True,
)
search_results = search_result["search_results"]
filtered_schema = search_result["filtered_schema"]
triples = search_result["triples"]
print("\n检索到的实体rerank 前,已合并两路):")
for r in search_results:
src = r.get("来源", "?")
print(f" [{src:<16}] {r['名称']} 得分={r['得分']:.4f} 标签={r['标签']}")
# ── 阶段三Rerank 重排序(统一对合并结果排序)────────────────────────
# # search_results = apply_rerank(
# # query=query,
# # search_results=search_results,
# # top_n=15,
# # verbose=True,
# # )
# # print("\n检索到的实体rerank 后):")
# # for r in search_results:
# # print(r)
# print("\n精简 Schema:\n", filtered_schema)
# ── 阶段四LLM 生成 Cypher ───────────────────────────────────────
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)
with driver.session() as session:
final_results = list(session.run(final_cypher))
print(f"\n📊 查询结果 ({len(final_results)} 条):")
for row in final_results:
print(dict(row))
finally:
driver.close()