1144 lines
45 KiB
Python
1144 lines
45 KiB
Python
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🔀 开始 Rerank(top_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_snippet(hybrid 路通常没有此字段)
|
||
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()
|