1845 lines
70 KiB
Python
1845 lines
70 KiB
Python
import os
|
||
import httpx
|
||
import json
|
||
import asyncio
|
||
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 AsyncGraphDatabase, 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 AsyncOpenAI
|
||
|
||
# ✅ 1. 导入配置
|
||
import sys
|
||
if hasattr(sys.stdout, "reconfigure"):
|
||
sys.stdout.reconfigure(encoding="utf-8", errors="replace")
|
||
sys.stderr.reconfigure(encoding="utf-8", errors="replace")
|
||
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||
from config import LLM_CONFIG, EMBEDDING_CONFIG, NEO4J_CONFIG, SEARCH_CONFIG
|
||
|
||
# ✅ 2. 【关键】全局清除代理环境变量,防止内网请求走代理
|
||
for _key in ['HTTP_PROXY', 'HTTPS_PROXY', 'http_proxy', 'https_proxy']:
|
||
if _key in os.environ:
|
||
del os.environ[_key]
|
||
print(f"[Init] Removed environment variable: {_key}")
|
||
|
||
|
||
# ================== 2. 基础配置 ==================
|
||
# 如果 config.py 中没有 Neo4j 配置,这里暂时保留硬编码,建议也移入 config.py
|
||
# ================== 2. 基础配置 ==================
|
||
# 从配置文件中加载参数
|
||
URI = os.getenv("NEO4J_URI", NEO4J_CONFIG["uri"])
|
||
AUTH = (NEO4J_CONFIG["username"], NEO4J_CONFIG["password"])
|
||
# URI = "bolt://192.168.1.164:7687"
|
||
# AUTH = ("neo4j", "zdht123@")
|
||
NAME_PROPERTY = SEARCH_CONFIG["name_property"]
|
||
FULLTEXT_PROPERTY = SEARCH_CONFIG["fulltext_property"]
|
||
EMBEDDING_PROPERTY = SEARCH_CONFIG["embedding_property"]
|
||
SEARCH_LABEL = SEARCH_CONFIG["search_label"]
|
||
VECTOR_INDEX_NAME = SEARCH_CONFIG["vector_index_name"]
|
||
FULLTEXT_INDEX_NAME = SEARCH_CONFIG["fulltext_index_name"]
|
||
FULLTEXT_FIELD_INDEX_NAME = SEARCH_CONFIG["fulltext_field_index_name"]
|
||
|
||
# 将列表转换为集合以提高查找效率
|
||
EXCLUDED_BUSINESS_LABELS: Set[str] = set(SEARCH_CONFIG["excluded_business_labels"])
|
||
# 确保 SEARCH_LABEL 也在排除列表中(如果配置中没写,这里手动添加以防万一)
|
||
EXCLUDED_BUSINESS_LABELS.add(SEARCH_LABEL)
|
||
|
||
VECTOR_DIMENSION = SEARCH_CONFIG["vector_dimension"]
|
||
_JIEBA_POS_WHITELIST = set(SEARCH_CONFIG["jieba_pos_whitelist"])
|
||
|
||
PROMPT_RESULT_TOP_K = int(SEARCH_CONFIG.get("cypher_prompt_top_k", 24))
|
||
PROMPT_RESULT_PER_LABEL = int(SEARCH_CONFIG.get("cypher_prompt_per_label", 5))
|
||
RELATION_ENTITY_TOP_K = int(SEARCH_CONFIG.get("relation_entity_top_k", 30))
|
||
ENABLE_EXTERNAL_RERANK = os.getenv("ENABLE_EXTERNAL_RERANK", "0").strip().lower() in {"1", "true", "yes", "on"}
|
||
|
||
_LLM_CLIENT: Optional[AsyncOpenAI] = None
|
||
_SCHEMA_CACHE: Optional[Tuple[str, List[Tuple[str, str, str]]]] = None
|
||
_SCHEMA_LOCK = asyncio.Lock()
|
||
|
||
# ✅ 2. 修改 LocalBgeM3Embeddings,使用配置并禁用代理
|
||
class LocalBgeM3Embeddings(Embedder):
|
||
def __init__(self, base_url: str = None, api_key: str = None, model_name: str = None):
|
||
# 使用默认配置,允许外部覆盖
|
||
self.base_url = base_url or EMBEDDING_CONFIG["base_url"]
|
||
self.api_key = api_key or EMBEDDING_CONFIG["api_key"]
|
||
self.model_name = model_name or EMBEDDING_CONFIG["model"]
|
||
|
||
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 = self.base_url
|
||
|
||
# ✅ 修复:临时清除代理环境变量,避免内网请求走代理
|
||
# 保存旧值
|
||
old_proxies = {
|
||
k: os.environ.pop(k, None)
|
||
for k in ('HTTP_PROXY', 'HTTPS_PROXY', 'http_proxy', 'https_proxy')
|
||
}
|
||
try:
|
||
# 创建客户端(不再传 proxies=None,因为旧版 httpx 不支持)
|
||
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"]]
|
||
finally:
|
||
# 恢复旧值
|
||
for k, v in old_proxies.items():
|
||
if v is not None:
|
||
os.environ[k] = v
|
||
|
||
return embedding_list
|
||
|
||
async def openai_chat_aysnc_nothink(query: str, timeout: int = None) -> str | None:
|
||
"""异步调用 OpenAI 兼容 API"""
|
||
global _LLM_CLIENT
|
||
if timeout is None:
|
||
timeout = LLM_CONFIG["timeout"]
|
||
|
||
try:
|
||
# 复用客户端,避免每次生成 Cypher 都重新建连接。
|
||
if _LLM_CLIENT is None:
|
||
_LLM_CLIENT = AsyncOpenAI(
|
||
api_key=LLM_CONFIG["api_key"],
|
||
base_url=LLM_CONFIG["base_url"],
|
||
timeout=timeout
|
||
)
|
||
|
||
response = await _LLM_CLIENT.chat.completions.create(
|
||
model=LLM_CONFIG["model"],
|
||
messages=[
|
||
{"role": "user", "content": query},
|
||
],
|
||
temperature=LLM_CONFIG["temperature"],
|
||
stream=False,
|
||
extra_body={"chat_template_kwargs": {"enable_thinking": False}}
|
||
)
|
||
return response.choices[0].message.content
|
||
|
||
except Exception as e:
|
||
print(f"❌ 调用 OpenAI API 时出错: {e}")
|
||
print(f" Base URL: {LLM_CONFIG['base_url']}")
|
||
print(f" Model: {LLM_CONFIG['model']}")
|
||
return None
|
||
|
||
|
||
# ================== 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(sync_driver, embedder) -> HybridCypherRetriever:
|
||
"""使用同步驱动构建检索器,因为HybridCypherRetriever不支持异步"""
|
||
return HybridCypherRetriever(
|
||
driver=sync_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)
|
||
|
||
|
||
def _load_schema_sync() -> Tuple[str, List[Tuple[str, str, str]]]:
|
||
graph = Neo4jGraph(url=URI, username=AUTH[0], password=AUTH[1], enhanced_schema=True)
|
||
schema = graph.schema
|
||
return schema, parse_schema_relationships(schema)
|
||
|
||
|
||
async def get_schema_triples() -> Tuple[str, List[Tuple[str, str, str]]]:
|
||
"""缓存 Neo4j schema,避免每次查询都触发 apoc.meta.graphSample。"""
|
||
global _SCHEMA_CACHE
|
||
if _SCHEMA_CACHE is not None:
|
||
return _SCHEMA_CACHE
|
||
|
||
async with _SCHEMA_LOCK:
|
||
if _SCHEMA_CACHE is None:
|
||
_SCHEMA_CACHE = await asyncio.to_thread(_load_schema_sync)
|
||
return _SCHEMA_CACHE
|
||
|
||
|
||
# ================== 7. Rerank 重排序 ==================
|
||
async 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,
|
||
}
|
||
|
||
async with httpx.AsyncClient() as client:
|
||
response = await client.post(
|
||
"http://rerank:9600/v1/rerank",
|
||
headers=headers,
|
||
json=data,
|
||
timeout=30.0,
|
||
)
|
||
|
||
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
|
||
|
||
|
||
async 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
|
||
|
||
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 = await 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:
|
||
keywords = list(dict.fromkeys(
|
||
w for w in jieba.cut(text) if len(w.strip()) >= min_len
|
||
))[:max_keywords]
|
||
|
||
return keywords
|
||
|
||
|
||
# ================== 8-NEW. fulltext 字段关键词检索 ==================
|
||
async 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` 属性字段做关键词全文检索。
|
||
"""
|
||
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_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
|
||
ORDER BY score DESC
|
||
LIMIT $top_k
|
||
"""
|
||
|
||
results: List[Dict] = []
|
||
try:
|
||
async with driver.session() as session:
|
||
result = await session.run(
|
||
cypher,
|
||
lucene_query=lucene_query,
|
||
score_threshold=score_threshold,
|
||
top_k=top_k,
|
||
)
|
||
rows = await result.data()
|
||
|
||
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]:
|
||
"""
|
||
合并两路检索结果
|
||
"""
|
||
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"
|
||
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 _primary_label(result: Dict) -> str:
|
||
labels = result.get("标签") or []
|
||
return labels[0] if labels else "未知标签"
|
||
|
||
|
||
def _prune_search_results(
|
||
search_results: List[Dict],
|
||
max_total: int = PROMPT_RESULT_TOP_K,
|
||
per_label: int = PROMPT_RESULT_PER_LABEL,
|
||
) -> List[Dict]:
|
||
"""
|
||
控制传入 LLM/关系查询的候选规模。
|
||
按标签保留高分实体,避免 top 结果被单一标签占满。
|
||
"""
|
||
if not search_results:
|
||
return []
|
||
|
||
label_priority = ["设备", "故障模式", "维修工作", "维修项目", "操作程序", "操作项目"]
|
||
sorted_results = sorted(search_results, key=lambda x: x.get("得分", 0), reverse=True)
|
||
selected: List[Dict] = []
|
||
selected_names: Set[str] = set()
|
||
|
||
def add_result(item: Dict) -> None:
|
||
name = item.get("名称")
|
||
if not name or name in selected_names or len(selected) >= max_total:
|
||
return
|
||
selected.append(item)
|
||
selected_names.add(name)
|
||
|
||
for label in label_priority:
|
||
count = 0
|
||
for item in sorted_results:
|
||
if label in (item.get("标签") or []):
|
||
add_result(item)
|
||
count += 1
|
||
if count >= per_label or len(selected) >= max_total:
|
||
break
|
||
|
||
for item in sorted_results:
|
||
if len(selected) >= max_total:
|
||
break
|
||
add_result(item)
|
||
|
||
return selected
|
||
|
||
|
||
async def search(
|
||
sentence: str,
|
||
sync_driver, # 同步驱动,用于HybridCypherRetriever
|
||
async_driver, # 异步驱动,用于fulltext检索
|
||
embedder,
|
||
top_k: int = 10,
|
||
hops: int = 3,
|
||
include_reverse: bool = True,
|
||
fulltext_top_k: int = 10,
|
||
fulltext_score_threshold: float = 0.0,
|
||
verbose: bool = True,
|
||
) -> Dict:
|
||
if verbose:
|
||
print(f"\n🔍 正在对整句话进行混合检索:「{sentence}」\n")
|
||
|
||
# hybrid 检索 - 使用同步驱动
|
||
retriever = _build_retriever(sync_driver, embedder)
|
||
hybrid_results: List[Dict] = []
|
||
try:
|
||
results = await asyncio.to_thread(
|
||
lambda: 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 = await search_fulltext_field(
|
||
query=sentence,
|
||
driver=async_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)
|
||
|
||
_full_schema, triples = await get_schema_triples()
|
||
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查询语句,语句中不要出现任何假设或推理,只能根据提供给你的实体关系进行查询。
|
||
7. **关系方向硬约束**:`#检索实体之间已在图数据库中确认存在的真实一跳关系` 中的方向必须严格照抄,绝对不能反转关系方向。例如真实关系是 `(维修工作)-[:表征]->(故障模式)`,如果从故障模式反查维修工作,必须写成 `(故障模式)<-[:表征]-(维修工作)`,不能写成 `(故障模式)-[:表征]->(维修工作)`。
|
||
8. **核心路径不要用 OPTIONAL MATCH**:回答用户问题所必需的主路径必须使用 `MATCH`。只有用户明确需要“可选补充信息”时,才允许使用 `OPTIONAL MATCH`。不要用 `OPTIONAL MATCH` 掩盖关系方向错误。
|
||
9. **优先使用真实一跳关系**:如果 `formatted_relations` 中已经给出实体之间的真实关系路径,优先基于这些关系拼接 Cypher;Schema 只作为标签和关系类型参考,不要用 Schema 臆造候选实体之间不存在的边。
|
||
10. **必须包含 RETURN**:Cypher 查询必须以 `RETURN` 子句结束,至少返回路径中关键节点的 `名称` 属性,例如 `RETURN d.名称 AS 设备名称, w.名称 AS 维修工作名称, f.名称 AS 故障模式名称`。绝对不能只输出 MATCH/WHERE 而没有 RETURN。
|
||
# 输出要求:
|
||
- 只输出纯 Cypher 查询语句本身,不要使用 JSON 包裹,不要使用 markdown 代码块(```),不要添加任何引号包裹整条语句,不要输出任何解释说明。
|
||
#用户问题如下:
|
||
{query}
|
||
#实体类型和实体关系如下:
|
||
{filtered_schema}
|
||
#根据用户问题,从数据库中检索的相关实体如下:
|
||
{search_results}
|
||
#检索实体之间已在图数据库中确认存在的真实一跳关系如下(这些关系路径真实可信,优先用于构建查询路径):
|
||
{formatted_relations}
|
||
|
||
示例1:
|
||
用户问题:连杆大端轴承发生损坏异常该如何解决?
|
||
实体类型和实体关系如下:
|
||
实体: ["主要附件仪表的规格型号及性能参数", "使用期间安全防护及注意事项", "功能", "器材保障", "图册", "基本情况", "备品备件", "子系统", "安全保护措施及故障处理", "安全警告", "工作原理", "布置情况", "技术参数", "技术指标参数", "接口", "接口情况", "操作使用", "操作员与设备之间操作员与其他设备或系统设备员之间关系", "操作程序", "操作项目", "故障模式", "环境条件", "系统", "组成", "结构特点", "维修安全防护及注意事项汇总表", "维修工作", "维修项目", "维护保养", "舰船", "舰艇", "设备", "设计单位", "调试"]
|
||
关系:
|
||
(:组成)-[:包含]->(:设备)
|
||
(:维修工作)-[:包含故障]->(:故障模式)
|
||
(:维修工作)-[:表征]->(:故障模式)
|
||
(:维修项目)-[:修理]->(:故障模式)
|
||
(:设备)-[:包含]->(:工作原理)
|
||
(:设备)-[:包含]->(:布置情况)
|
||
(:设备)-[:包含]->(:技术参数)
|
||
(:设备)-[:包含]->(:接口情况)
|
||
(:设备)-[:包含]->(:操作员与设备之间操作员与其他设备或系统设备员之间关系)
|
||
(:设备)-[:包含]->(:组成)
|
||
(:设备)-[:包含]->(:结构特点)
|
||
(:设备)-[:包含]->(:维修安全防护及注意事项汇总表)
|
||
(:设备)-[:包含]->(:维护保养)
|
||
(:设备)-[:包括]->(:组成)
|
||
(:设备)-[:发生故障使用维修工作]->(:维修工作)
|
||
(:设备)-[:安全警告]->(:安全警告)
|
||
(:设备)-[:实现]->(:功能)
|
||
(:设备)-[:实现]->(:基本情况)
|
||
(:设备)-[:支持操作]->(:操作使用)
|
||
(:设备)-[:维修]->(:维修工作)
|
||
(:设备)-[:设计]->(:设计单位)
|
||
(:设备)-[:调试]->(:调试)
|
||
(:设备)-[:适用]->(:环境条件)
|
||
(:设备)-[:配备]->(:接口)
|
||
(:设备)-[:配套]->(:图册)
|
||
#根据用户问题,从数据库中找到的相关实体如下:
|
||
{{'标签': ['维修项目'], '名称': '修复连杆大端轴承异常磨损', '得分': 1.0}}
|
||
{{'标签': ['维修项目'], '名称': '更换连杆大端轴承', '得分': 0.9839502279586991}}
|
||
{{'标签': ['维修工作'], '名称': '连杆大端轴承维修工作指导', '得分': 0.9601912347183726}}
|
||
{{'标签': ['设备'], '名称': '连杆大端轴承', '得分': 0.9378418041605234}}
|
||
{{'标签': ['维修项目'], '名称': '更换连杆小段轴承', '得分': 0.9370194814204257}}
|
||
{{'标签': ['设备'], '名称': '连杆小端轴承', '得分': 0.9056004631162252}}
|
||
{{'标签': ['故障模式'], '名称': '主轴承或连杆大端轴承磨损导致主机运行中振动异常增大', '得分': 0.890951410907792}}
|
||
{{'标签': ['故障模式'], '名称': '主轴承或连杆大端轴承磨损导致机座与机架组件运行中振动异常增大', '得分': 0.8895577902640475}}
|
||
{{'标签': ['设备'], '名称': '连杆小段轴承', '得分': 0.8785823193324869}}
|
||
{{'标签': ['维修工作'], '名称': '传动轴的维修工作', '得分': 0.8768441389219146}}
|
||
{{'标签': ['维修项目'], '名称': '修复连杆大端轴承异常磨损', '得分': 1.0}}
|
||
{{'标签': ['故障模式'], '名称': '多缸燃烧严重失衡导致轴承组件紧急停机', '得分': 2.8766837120056152}}
|
||
{{'标签': ['维修项目'], '名称': '拆检活塞头磨损情况', '得分': 0.8453663838751524}}
|
||
|
||
#检索实体之间已在图数据库中确认存在的真实一跳关系如下(这些关系路径真实可信,优先用于构建查询路径):
|
||
(设备:连杆大端轴承) -[:发生故障使用维修工作]-> (维修工作:连杆大端轴承维修工作指导)
|
||
(维修工作:连杆大端轴承维修工作指导) -[:包含故障]-> (故障模式:主轴承或连杆大端轴承磨损导致主机运行中振动异常增大)
|
||
(维修项目:修复连杆大端轴承异常磨损) -[:修理]-> (故障模式:主轴承或连杆大端轴承磨损导致机座与机架组件运行中振动异常增大)
|
||
(维修项目:更换连杆大端轴承) -[:修理]-> (故障模式:主轴承或连杆大端轴承磨损导致机座与机架组件运行中振动异常增大)
|
||
思考:
|
||
首先确定入口节点:问题中首先提到"连杆大端轴承",检索结果中存在标签为"设备"、名称为"连杆大端轴承"的实体,得分为0.9378,将其作为查询起点。
|
||
接着依据已确认的真实一跳关系,逐步拼接完整路径:
|
||
- 步骤1:(设备:连杆大端轴承) -[:发生故障使用维修工作]-> (维修工作:连杆大端轴承维修工作指导) ✅ 已确认
|
||
- 步骤2:(维修工作:连杆大端轴承维修工作指导) -[:包含故障]-> (故障模式:主轴承或连杆大端轴承磨损导致主机运行中振动异常增大) ✅ 已确认
|
||
- 步骤3:需要找到修理该故障模式的维修项目。真实关系中"修复连杆大端轴承异常磨损"和"更换连杆大端轴承"均指向另一故障模式节点,结合检索得分,优先选取得分最高的"修复连杆大端轴承异常磨损"(得分1.0)作为维修项目节点。
|
||
- 完整路径:设备 -[:发生故障使用维修工作]-> 维修工作 -[:包含故障]-> 故障模式 <-[:修理]- 维修项目
|
||
输出为:
|
||
MATCH (d:设备)-[:发生故障使用维修工作]->(w:维修工作)-[:包含故障]->(f:故障模式)<-[:修理]-(p:维修项目)
|
||
WHERE d.名称 = '连杆大端轴承'
|
||
AND w.名称 = '连杆大端轴承维修工作指导'
|
||
AND f.名称 = '主轴承或连杆大端轴承磨损导致主机运行中振动异常增大'
|
||
AND p.名称 = '修复连杆大端轴承异常磨损'
|
||
RETURN w.名称 AS 维修工作名称, f.名称 AS 故障模式, p.名称 AS 维修项目名称
|
||
|
||
请参考上述示例,对提供的问题进行cpyher语句的编写。
|
||
注意,上述示例中是以设备为起始节点,并不是所有问题都是以设备为起点,你要进行合理的选择正确的起始图谱节点,而不是一味模仿上述案例
|
||
输出为:
|
||
"""
|
||
|
||
|
||
def _format_search_results_for_prompt(search_results: list) -> str:
|
||
prompt_results = _prune_search_results(search_results)
|
||
lines = []
|
||
for idx, r in enumerate(prompt_results, 1):
|
||
labels = ",".join(r.get("标签", []))
|
||
score = r.get("得分", 0)
|
||
source = r.get("来源", "")
|
||
lines.append(
|
||
f"{idx}. 标签=[{labels}] | 名称={r.get('名称', '')} | 得分={score:.4f} | 来源={source}"
|
||
)
|
||
return "\n".join(lines)
|
||
|
||
|
||
def build_cypher_prompt(query: str, filtered_schema: str, search_results: list, formatted_relations: str) -> str:
|
||
prompt_results = _format_search_results_for_prompt(search_results)
|
||
|
||
return CYPHER_PROMPT.format(
|
||
query=query,
|
||
filtered_schema=filtered_schema,
|
||
search_results=prompt_results,
|
||
formatted_relations=formatted_relations,
|
||
)
|
||
|
||
def extract_cypher(llm_output: str | None) -> str:
|
||
if not llm_output or not isinstance(llm_output, str):
|
||
return ""
|
||
|
||
match = re.search(r"(MATCH[\s\S]+?RETURN[^\n]+(?:\n[^\n]+)*)", llm_output)
|
||
if match:
|
||
cypher = match.group(1).strip()
|
||
else:
|
||
match_without_return = re.search(r"(MATCH[\s\S]+)", llm_output)
|
||
cypher = match_without_return.group(1).strip() if match_without_return else llm_output.strip()
|
||
|
||
# 修复:将字面的 \n 转换为实际换行符
|
||
cypher = cypher.replace('\\n', '\n')
|
||
|
||
return cypher
|
||
|
||
|
||
def _default_return_clause(cypher: str) -> str:
|
||
alias_labels = _parse_alias_label(cypher)
|
||
if not alias_labels:
|
||
return "*"
|
||
|
||
label_alias = {
|
||
label: alias
|
||
for alias, label in alias_labels.items()
|
||
}
|
||
priority = [
|
||
("设备", "设备名称"),
|
||
("维修工作", "维修工作名称"),
|
||
("故障模式", "故障模式名称"),
|
||
("维修项目", "维修项目名称"),
|
||
("操作程序", "操作程序名称"),
|
||
("操作项目", "操作项目名称"),
|
||
]
|
||
|
||
parts = []
|
||
used_aliases: Set[str] = set()
|
||
for label, column in priority:
|
||
alias = label_alias.get(label)
|
||
if alias and alias not in used_aliases:
|
||
parts.append(f"{alias}.`名称` AS {column}")
|
||
used_aliases.add(alias)
|
||
|
||
for alias, label in alias_labels.items():
|
||
if alias in used_aliases:
|
||
continue
|
||
parts.append(f"{alias}.`名称` AS {label}名称")
|
||
used_aliases.add(alias)
|
||
|
||
return ", ".join(parts) if parts else "*"
|
||
|
||
|
||
def ensure_return_clause(cypher: str) -> str:
|
||
"""LLM 偶尔漏写 RETURN,这里兜底补齐,避免 Neo4j 语法错误。"""
|
||
if not cypher:
|
||
return cypher
|
||
if re.search(r"\bRETURN\b", cypher, re.IGNORECASE):
|
||
return cypher
|
||
return f"{cypher.rstrip()}\nRETURN {_default_return_clause(cypher)}"
|
||
|
||
# ================== 11. Cypher 诊断基础工具 ==================
|
||
async def _run_query(session, cypher: str, **params) -> list:
|
||
try:
|
||
result = await session.run(cypher, **params)
|
||
return await result.data()
|
||
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, str]]:
|
||
"""
|
||
解析 Cypher 中的一跳关系段,返回真实图方向。
|
||
tuple: (src_alias, src_label, rel, tgt_alias, tgt_label, text_direction)
|
||
text_direction 用于后续必要时翻转原始 Cypher 文本。
|
||
"""
|
||
alias_labels = _parse_alias_label(cypher)
|
||
segments = []
|
||
|
||
forward = re.compile(
|
||
r"(?=(\((\w+)(?::([^\)]+))?\)\s*-\s*\[:([^\]]+)\]\s*->\s*\((\w+)(?::([^\)]+))?\)))"
|
||
)
|
||
reverse = re.compile(
|
||
r"(?=(\((\w+)(?::([^\)]+))?\)\s*<-\s*\[:([^\]]+)\]\s*-\s*\((\w+)(?::([^\)]+))?\)))"
|
||
)
|
||
|
||
for match in forward.finditer(cypher):
|
||
_full, sa, sl, rel, ta, tl = match.groups()
|
||
sl = (sl or alias_labels.get(sa, "")).strip()
|
||
tl = (tl or alias_labels.get(ta, "")).strip()
|
||
if not sl or not tl:
|
||
continue
|
||
segments.append((
|
||
match.start(),
|
||
(sa.strip(), sl, rel.strip(), ta.strip(), tl, "forward"),
|
||
))
|
||
|
||
for match in reverse.finditer(cypher):
|
||
_full, left_alias, left_label, rel, right_alias, right_label = match.groups()
|
||
left_label = (left_label or alias_labels.get(left_alias, "")).strip()
|
||
right_label = (right_label or alias_labels.get(right_alias, "")).strip()
|
||
if not left_label or not right_label:
|
||
continue
|
||
segments.append((
|
||
match.start(),
|
||
(
|
||
right_alias.strip(), right_label, rel.strip(),
|
||
left_alias.strip(), left_label, "reverse",
|
||
),
|
||
))
|
||
|
||
return [segment for _pos, segment in sorted(segments, key=lambda x: x[0])]
|
||
|
||
|
||
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,
|
||
)
|
||
|
||
|
||
async 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 = await _run_query(session, q, src=src_name, tgt=tgt_name)
|
||
return rows[0]["c"] if rows else 0
|
||
|
||
|
||
async def _test_reversed_segment(
|
||
session,
|
||
src_label: str, src_name: Optional[str],
|
||
rel: str,
|
||
tgt_label: str, tgt_name: Optional[str],
|
||
) -> int:
|
||
return await _test_segment(
|
||
session=session,
|
||
src_label=tgt_label,
|
||
src_name=tgt_name,
|
||
rel=rel,
|
||
tgt_label=src_label,
|
||
tgt_name=src_name,
|
||
)
|
||
|
||
|
||
def _flip_segment_direction(
|
||
cypher: str,
|
||
src_alias: str,
|
||
src_label: str,
|
||
rel: str,
|
||
tgt_alias: str,
|
||
tgt_label: str,
|
||
text_direction: str,
|
||
) -> str:
|
||
if text_direction == "forward":
|
||
pattern = re.compile(
|
||
rf"\(\s*{re.escape(src_alias)}\s*(?::\s*{re.escape(src_label)}\s*)?\)"
|
||
rf"\s*-\s*\[:\s*{re.escape(rel)}\s*\]\s*->\s*"
|
||
rf"\(\s*{re.escape(tgt_alias)}\s*(?::\s*{re.escape(tgt_label)}\s*)?\)"
|
||
)
|
||
replacement = f"({src_alias}:{src_label})<-[:{rel}]-({tgt_alias}:{tgt_label})"
|
||
else:
|
||
pattern = re.compile(
|
||
rf"\(\s*{re.escape(tgt_alias)}\s*(?::\s*{re.escape(tgt_label)}\s*)?\)"
|
||
rf"\s*<-\s*\[:\s*{re.escape(rel)}\s*\]\s*-\s*"
|
||
rf"\(\s*{re.escape(src_alias)}\s*(?::\s*{re.escape(src_label)}\s*)?\)"
|
||
)
|
||
replacement = f"({tgt_alias}:{tgt_label})-[:{rel}]->({src_alias}:{src_label})"
|
||
|
||
return pattern.sub(replacement, cypher, count=1)
|
||
|
||
|
||
async 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 = await _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
|
||
|
||
|
||
async def _find_verified_source_replacement(
|
||
session,
|
||
src_label: str,
|
||
rel: str,
|
||
tgt_label: str,
|
||
tgt_name: Optional[str],
|
||
candidates: list,
|
||
verbose: bool,
|
||
) -> Optional[str]:
|
||
for i, cand in enumerate(candidates):
|
||
cand_name = cand.get("名称", "")
|
||
if not cand_name:
|
||
continue
|
||
cnt = await _test_segment(session, src_label, cand_name, rel, tgt_label, tgt_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
|
||
]
|
||
|
||
|
||
async 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 = await _run_query(session, q, src=src_name, tgt=tgt_name)
|
||
return rows[0]["c"] if rows else 0
|
||
|
||
|
||
async 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 = await _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 = await _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
|
||
|
||
|
||
def _return_alias_columns(return_clause: str) -> Dict[str, str]:
|
||
"""
|
||
提取 RETURN 中 alias.名称 AS 输出列 的映射。
|
||
例如 d.名称 AS 设备名称 -> {"d": "设备名称"}
|
||
"""
|
||
result: Dict[str, str] = {}
|
||
pattern = re.compile(
|
||
r"(\w+)\.`?名称`?\s+AS\s+([A-Za-z0-9_\u4e00-\u9fff]+)",
|
||
re.IGNORECASE,
|
||
)
|
||
for alias, column in pattern.findall(return_clause or ""):
|
||
result[alias] = column
|
||
return result
|
||
|
||
|
||
def _null_return_aliases(rows: list, return_clause: str) -> List[str]:
|
||
if not rows:
|
||
return []
|
||
alias_columns = _return_alias_columns(return_clause)
|
||
null_aliases = []
|
||
for alias, column in alias_columns.items():
|
||
if all(row.get(column) in (None, "") for row in rows):
|
||
null_aliases.append(alias)
|
||
return null_aliases
|
||
|
||
|
||
def _apply_name_replacements(cypher: str, alias_name: Dict[str, str], current_names: Dict[str, str]) -> str:
|
||
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}'")
|
||
return fixed_cypher
|
||
|
||
|
||
async def _repair_name_bindings_once(
|
||
session,
|
||
segments: List[Tuple[str, str, str, str, str, str]],
|
||
current_names: Dict[str, str],
|
||
search_results: list,
|
||
verbose: bool = True,
|
||
) -> bool:
|
||
"""
|
||
当某个中间节点被后续段修正后,重新校验全链路。
|
||
如果 target alias 也是下一段的 source,优先替换当前段的 source,避免把中间节点改回去。
|
||
"""
|
||
aliases_used_as_src = {seg[0] for seg in segments}
|
||
changed = False
|
||
|
||
for src_alias, src_label, rel, tgt_alias, tgt_label, _text_direction in segments:
|
||
src_name = current_names.get(src_alias)
|
||
tgt_name = current_names.get(tgt_alias)
|
||
cnt = await _test_segment(session, src_label, src_name, rel, tgt_label, tgt_name)
|
||
if cnt > 0:
|
||
continue
|
||
|
||
prefer_source = tgt_alias in aliases_used_as_src
|
||
attempts = ("source", "target") if prefer_source else ("target", "source")
|
||
|
||
if verbose:
|
||
print(
|
||
f" 🔁 重校验断链: ({src_alias}:{src_label} '{src_name}')"
|
||
f" -[:{rel}]-> ({tgt_alias}:{tgt_label} '{tgt_name}')"
|
||
)
|
||
|
||
for side in attempts:
|
||
if side == "source":
|
||
verified_src = await _find_verified_source_replacement(
|
||
session,
|
||
src_label=src_label,
|
||
rel=rel,
|
||
tgt_label=tgt_label,
|
||
tgt_name=tgt_name,
|
||
candidates=_candidates_for_label(search_results, src_label),
|
||
verbose=verbose,
|
||
)
|
||
if verified_src and verified_src != src_name:
|
||
current_names[src_alias] = verified_src
|
||
changed = True
|
||
if verbose:
|
||
print(f" ✅ 重校验修复源节点: '{src_name}' → '{verified_src}'")
|
||
break
|
||
else:
|
||
verified_tgt = await _find_verified_replacement(
|
||
session,
|
||
src_label=src_label,
|
||
src_name=src_name,
|
||
rel=rel,
|
||
tgt_label=tgt_label,
|
||
candidates=_candidates_for_label(search_results, tgt_label),
|
||
verbose=verbose,
|
||
)
|
||
if verified_tgt and verified_tgt != tgt_name:
|
||
current_names[tgt_alias] = verified_tgt
|
||
changed = True
|
||
if verbose:
|
||
print(f" ✅ 重校验修复目标节点: '{tgt_name}' → '{verified_tgt}'")
|
||
break
|
||
|
||
return changed
|
||
|
||
|
||
# ================== 12. 端到端链路诊断与修复 ==================
|
||
async def diagnose_and_fix_cypher(
|
||
cypher: str,
|
||
driver,
|
||
search_results: list,
|
||
schema_triples: List[Tuple[str, str, str]] = None,
|
||
verbose: bool = True,
|
||
) -> dict:
|
||
cypher = ensure_return_clause(cypher)
|
||
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")
|
||
|
||
async with driver.session() as session:
|
||
initial_result = await _run_query(session, cypher)
|
||
null_aliases = _null_return_aliases(initial_result, return_clause)
|
||
if initial_result and not null_aliases:
|
||
if verbose:
|
||
print("✅ 完整查询有结果,无需修复。")
|
||
return {
|
||
"status": "ok", "chain_verified": True,
|
||
"diagnosis": "查询正常,有结果返回。",
|
||
"broken_segments": [], "fixed_cypher": None,
|
||
"final_names": current_names,
|
||
}
|
||
elif initial_result and null_aliases and verbose:
|
||
print(
|
||
"⚠️ 查询虽然返回了行,但以下 RETURN 别名全为空,"
|
||
f"疑似 OPTIONAL MATCH 掩盖了断链: {null_aliases}\n"
|
||
)
|
||
|
||
if verbose:
|
||
print("⚠️ 开始逐段诊断...\n")
|
||
|
||
broken_segments: list = []
|
||
segment_reports: list = []
|
||
|
||
for seg_idx, (src_alias, src_label, rel, tgt_alias, tgt_label, text_direction) 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 = await _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
|
||
|
||
reverse_cnt = await _test_reversed_segment(
|
||
session, src_label, src_name, rel, tgt_label, tgt_name
|
||
)
|
||
if reverse_cnt > 0:
|
||
cypher = _flip_segment_direction(
|
||
cypher,
|
||
src_alias=src_alias,
|
||
src_label=src_label,
|
||
rel=rel,
|
||
tgt_alias=tgt_alias,
|
||
tgt_label=tgt_label,
|
||
text_direction=text_direction,
|
||
)
|
||
segments[seg_idx] = (
|
||
tgt_alias, tgt_label, rel, src_alias, src_label,
|
||
"reverse" if text_direction == "forward" else "forward",
|
||
)
|
||
if verbose:
|
||
print(
|
||
f" ✅ 关系方向修复: "
|
||
f"({src_label})-[:{rel}]->({tgt_label}) "
|
||
f"翻转为 ({src_label})<-[:{rel}]-({tgt_label}) "
|
||
f"(匹配数: {reverse_cnt})\n"
|
||
)
|
||
segment_reports.append(
|
||
f"🔧 {seg_desc}\n → 关系方向翻转修复,匹配数 {reverse_cnt}"
|
||
)
|
||
broken_segments.append({
|
||
"seg_idx": seg_idx + 1,
|
||
"pass": "direction",
|
||
"original_rel": rel,
|
||
"src_alias": src_alias,
|
||
"tgt_alias": tgt_alias,
|
||
})
|
||
continue
|
||
|
||
# Pass A
|
||
if verbose:
|
||
print(f" ❌ 断链!Pass A: 尝试替换节点名称,保持关系 [:{rel}] 不变...")
|
||
cands = _candidates_for_label(search_results, tgt_label)
|
||
verified = await _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
|
||
|
||
if verbose:
|
||
print(f" ❌ 目标节点替换失败。Pass A2: 尝试替换源节点,保持关系 [:{rel}] 和目标节点不变...")
|
||
src_cands = _candidates_for_label(search_results, src_label)
|
||
verified_src = await _find_verified_source_replacement(
|
||
session, src_label, rel, tgt_label, tgt_name, src_cands, verbose
|
||
)
|
||
if verified_src:
|
||
if verbose:
|
||
print(f" ✅ Pass A2 修复: 源节点 '{src_name}' → '{verified_src}'\n")
|
||
current_names[src_alias] = verified_src
|
||
segment_reports.append(
|
||
f"🔧 {seg_desc}\n → Pass A2: 源节点 '{src_name}' 替换为 '{verified_src}'"
|
||
)
|
||
broken_segments.append({
|
||
"seg_idx": seg_idx + 1, "pass": "A2",
|
||
"src_alias": src_alias, "src_label": src_label,
|
||
"original": src_name, "fixed": verified_src,
|
||
})
|
||
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 = await _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 = await _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, text_direction)
|
||
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 = await _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")
|
||
verify_result = await _run_query(session, new_cypher)
|
||
if verify_result:
|
||
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 = _apply_name_replacements(cypher, alias_name, current_names)
|
||
final_verify = await _run_query(session, fixed_cypher)
|
||
final_null_aliases = _null_return_aliases(final_verify, return_clause)
|
||
chain_ok = bool(final_verify) and not final_null_aliases
|
||
|
||
for _ in range(2):
|
||
if chain_ok:
|
||
break
|
||
changed = await _repair_name_bindings_once(
|
||
session=session,
|
||
segments=segments,
|
||
current_names=current_names,
|
||
search_results=search_results,
|
||
verbose=verbose,
|
||
)
|
||
if not changed:
|
||
break
|
||
fixed_cypher = _apply_name_replacements(cypher, alias_name, current_names)
|
||
final_verify = await _run_query(session, fixed_cypher)
|
||
final_null_aliases = _null_return_aliases(final_verify, return_clause)
|
||
chain_ok = bool(final_verify) and not final_null_aliases
|
||
|
||
if verbose:
|
||
print("=" * 60)
|
||
print("📋 逐段诊断报告:")
|
||
for line in segment_reports:
|
||
print(f" {line}")
|
||
print(f"\n🗂️ 最终名称绑定: {current_names}")
|
||
if final_null_aliases:
|
||
print(f"\n⚠️ 以下 RETURN 别名仍为空: {final_null_aliases}")
|
||
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,
|
||
}
|
||
|
||
|
||
async def retrieve_entities(
|
||
query: str,
|
||
sync_driver, # 同步驱动
|
||
async_driver, # 异步驱动
|
||
embedder,
|
||
hybrid_top_k: int = 100,
|
||
fulltext_top_k: int = 15,
|
||
rerank_top_k: int = 15,
|
||
use_rerank: bool = True,
|
||
hops: int = 3,
|
||
verbose: bool = True,
|
||
) -> Dict:
|
||
"""
|
||
统一检索入口:
|
||
1. hybrid 检索
|
||
2. fulltext 检索
|
||
3. merge
|
||
4. rerank
|
||
5. schema 裁剪
|
||
"""
|
||
|
||
# Step1:原始检索
|
||
result = await search(
|
||
sentence=query,
|
||
sync_driver=sync_driver,
|
||
async_driver=async_driver,
|
||
embedder=embedder,
|
||
top_k=hybrid_top_k,
|
||
hops=hops,
|
||
include_reverse=True,
|
||
fulltext_top_k=fulltext_top_k,
|
||
fulltext_score_threshold=0.0,
|
||
verbose=verbose,
|
||
)
|
||
|
||
search_results = result["search_results"]
|
||
|
||
# Step2:Rerank
|
||
# if use_rerank and ENABLE_EXTERNAL_RERANK and search_results:
|
||
# search_results = await apply_rerank(
|
||
# query=query,
|
||
# search_results=search_results,
|
||
# top_n=rerank_top_k,
|
||
# verbose=verbose,
|
||
# )
|
||
|
||
# search_results = _prune_search_results(
|
||
# search_results,
|
||
# max_total=max(PROMPT_RESULT_TOP_K, rerank_top_k),
|
||
# per_label=PROMPT_RESULT_PER_LABEL,
|
||
# )
|
||
result["search_results"] = search_results
|
||
|
||
if verbose:
|
||
print("\n📌 最终检索结果:")
|
||
for r in search_results:
|
||
print(
|
||
f"[{r['来源']:<15}] "
|
||
f"{r['名称']} "
|
||
f"score={r['得分']:.4f} "
|
||
f"标签={r['标签']}"
|
||
)
|
||
|
||
return result
|
||
|
||
|
||
async def query_with_cypher(
|
||
query: str,
|
||
retrieval_result: Dict,
|
||
driver,
|
||
verbose: bool = True,
|
||
) -> Dict:
|
||
"""
|
||
统一 Cypher 查询入口
|
||
"""
|
||
search_results = retrieval_result["search_results"]
|
||
filtered_schema = retrieval_result["filtered_schema"]
|
||
triples = retrieval_result["triples"]
|
||
formatted_relations = retrieval_result["formatted_relations"]
|
||
|
||
# Step1:构建Prompt
|
||
prompt = build_cypher_prompt(
|
||
query,
|
||
filtered_schema,
|
||
search_results,
|
||
formatted_relations,
|
||
)
|
||
# print("\n🤖 构建 Prompt:")
|
||
# print(prompt)
|
||
# Step2:LLM生成Cypher
|
||
llm_output = await openai_chat_aysnc_nothink(prompt)
|
||
|
||
# ✅ 检查 LLM 是否返回了有效内容
|
||
if not llm_output:
|
||
raise Exception("LLM 未返回有效内容,无法生成 Cypher 查询。请检查模型服务连通性。")
|
||
|
||
cypher = extract_cypher(llm_output)
|
||
|
||
# Step3:删除不存在名称
|
||
cypher = strip_unverified_name_conditions(cypher, search_results)
|
||
cypher = ensure_return_clause(cypher)
|
||
|
||
if verbose:
|
||
print("\n🤖 初始 Cypher:")
|
||
print(cypher if cypher else "(Empty)")
|
||
|
||
# ✅ 如果提取后为空,直接报错或返回空结果,不要继续执行诊断
|
||
if not cypher or not cypher.strip().upper().startswith("MATCH"):
|
||
raise Exception(f"未能从 LLM 输出中提取有效的 MATCH 语句。Output: {llm_output[:100]}...")
|
||
|
||
# Step4:链路诊断修复
|
||
diag = await diagnose_and_fix_cypher(
|
||
cypher=cypher,
|
||
driver=driver,
|
||
search_results=search_results,
|
||
schema_triples=triples,
|
||
verbose=verbose,
|
||
)
|
||
|
||
final_cypher = diag["fixed_cypher"] if diag["fixed_cypher"] else cypher
|
||
final_cypher = ensure_return_clause(final_cypher)
|
||
|
||
# ✅ 再次确保最终 Cypher 不为空
|
||
if not final_cypher or not final_cypher.strip():
|
||
raise Exception("修复后的 Cypher 仍为空,无法执行查询。")
|
||
|
||
if verbose:
|
||
print("\n💡 最终执行 Cypher:")
|
||
print(final_cypher)
|
||
|
||
# Step5:执行查询
|
||
async with driver.session() as session:
|
||
result = await session.run(final_cypher)
|
||
rows = await result.data()
|
||
|
||
final_results = [dict(r) for r in rows]
|
||
|
||
return {
|
||
"original_cypher": cypher,
|
||
"final_cypher": final_cypher,
|
||
"diagnosis": diag,
|
||
"results": final_results,
|
||
}
|
||
|
||
async def fetch_relations_among_results(
|
||
search_results: List[Dict],
|
||
async_driver,
|
||
verbose: bool = True,
|
||
) -> List[Dict]:
|
||
"""
|
||
查询检索结果中所有实体之间存在的一跳关系。
|
||
|
||
Args:
|
||
search_results: retrieve_entities 返回的检索结果列表
|
||
async_driver: 异步 Neo4j 驱动
|
||
verbose: 是否打印详细信息
|
||
|
||
Returns:
|
||
关系列表,每项包含 src_name, src_label, rel, tgt_name, tgt_label
|
||
"""
|
||
# 提取所有实体名称
|
||
names = [
|
||
r["名称"]
|
||
for r in search_results[:RELATION_ENTITY_TOP_K]
|
||
if r.get("名称") and r["名称"] != "未知名称"
|
||
]
|
||
|
||
if not names:
|
||
if verbose:
|
||
print("⚠️ 检索结果为空,跳过关系查询。")
|
||
return []
|
||
|
||
if verbose:
|
||
print(f"\n🔗 正在查询 {len(names)} 个实体之间的一跳关系...")
|
||
|
||
# 用 WHERE 过滤两端节点都在结果集中的边,一次查询搞定
|
||
cypher = """
|
||
MATCH (a)-[r]->(b)
|
||
WHERE a.`名称` IN $names
|
||
AND b.`名称` IN $names
|
||
RETURN
|
||
a.`名称` AS src_name,
|
||
labels(a)[0] AS src_label,
|
||
type(r) AS rel,
|
||
b.`名称` AS tgt_name,
|
||
labels(b)[0] AS tgt_label
|
||
ORDER BY src_name, rel, tgt_name
|
||
"""
|
||
|
||
relations: List[Dict] = []
|
||
try:
|
||
async with async_driver.session() as session:
|
||
result = await session.run(cypher, names=names)
|
||
rows = await result.data()
|
||
|
||
for row in rows:
|
||
relations.append({
|
||
"src_name": row["src_name"],
|
||
"src_label": row["src_label"],
|
||
"rel": row["rel"],
|
||
"tgt_name": row["tgt_name"],
|
||
"tgt_label": row["tgt_label"],
|
||
})
|
||
|
||
if verbose:
|
||
print(f"✅ 共找到 {len(relations)} 条一跳关系:\n")
|
||
# 按关系类型分组打印,更易读
|
||
from itertools import groupby
|
||
sorted_rels = sorted(relations, key=lambda x: x["rel"])
|
||
for rel_type, group in groupby(sorted_rels, key=lambda x: x["rel"]):
|
||
print(f" [{rel_type}]")
|
||
for item in group:
|
||
print(
|
||
f" ({item['src_label']}:{item['src_name']})"
|
||
f" -[:{item['rel']}]->"
|
||
f" ({item['tgt_label']}:{item['tgt_name']})"
|
||
)
|
||
print()
|
||
|
||
except Exception as e:
|
||
print(f"⚠️ 查询实体间关系时出错: {e}")
|
||
|
||
return relations
|
||
|
||
def format_relations(relations: List[Dict]) -> str:
|
||
lines = []
|
||
for r in relations:
|
||
lines.append(
|
||
f"({r['src_label']}:{r['src_name']}) "
|
||
f"-[:{r['rel']}]-> "
|
||
f"({r['tgt_label']}:{r['tgt_name']})"
|
||
)
|
||
return "\n".join(lines)
|
||
# ================== 13. 主程序 ==================
|
||
async def main():
|
||
# 初始化同步驱动(用于HybridCypherRetriever)
|
||
sync_driver = GraphDatabase.driver(URI, auth=AUTH)
|
||
|
||
# 初始化异步驱动(用于其他异步操作)
|
||
async_driver = AsyncGraphDatabase.driver(URI, auth=AUTH)
|
||
|
||
embedder = LocalBgeM3Embeddings(
|
||
base_url=os.getenv("LOCAL_EMBEDDING_BASE"),
|
||
api_key=os.getenv("LOCAL_API_KEY"),
|
||
)
|
||
|
||
query = "发动机的工作原理"
|
||
try:
|
||
# ① 检索
|
||
retrieval = await retrieve_entities(
|
||
query=query,
|
||
sync_driver=sync_driver,
|
||
async_driver=async_driver,
|
||
embedder=embedder,
|
||
use_rerank=True,
|
||
verbose=True,
|
||
)
|
||
relations = await fetch_relations_among_results(
|
||
search_results=retrieval["search_results"],
|
||
async_driver=async_driver,
|
||
verbose=True,
|
||
)
|
||
formatted_relations = format_relations(relations)
|
||
retrieval["formatted_relations"] = formatted_relations
|
||
print(retrieval["formatted_relations"])
|
||
# ② Cypher生成+修复+查询
|
||
result = await query_with_cypher(
|
||
query=query,
|
||
retrieval_result=retrieval,
|
||
driver=async_driver,
|
||
verbose=True,
|
||
)
|
||
|
||
print("\n" + "="*60)
|
||
print("🚀 最终 Cypher")
|
||
print("="*60)
|
||
print(result["final_cypher"])
|
||
|
||
print("\n📊 查询结果:")
|
||
for row in result["results"]:
|
||
print(row)
|
||
|
||
finally:
|
||
sync_driver.close()
|
||
await async_driver.close()
|
||
|
||
|
||
if __name__ == "__main__":
|
||
asyncio.run(main())
|
||
|