更新 tools/texttosql_utils.py
This commit is contained in:
parent
75a6238937
commit
572d8f1491
@ -0,0 +1,184 @@
|
||||
import os
|
||||
import ast
|
||||
import re
|
||||
import asyncio
|
||||
import httpx
|
||||
from typing import List, Dict
|
||||
|
||||
from neo4j import GraphDatabase
|
||||
from neo4j_graphrag.retrievers import HybridCypherRetriever
|
||||
from neo4j_graphrag.embeddings.base import Embedder
|
||||
from modelsAPI.model_api import LocalBgeM3Embeddings, OpenaiAPI
|
||||
from config import MAP_INS, NEO4J_CONFIG, INDEX_CONFIG, EMBEDDING_CONFIG
|
||||
|
||||
# ================== 配置 ==================
|
||||
URI = NEO4J_CONFIG["uri"]
|
||||
AUTH = (NEO4J_CONFIG["user"], NEO4J_CONFIG["password"])
|
||||
NAME_PROPERTY = INDEX_CONFIG["NAME_PROPERTY"]
|
||||
FULLTEXT_PROPERTY = INDEX_CONFIG["FULLTEXT_PROPERTY"]
|
||||
EMBEDDING_PROPERTY = INDEX_CONFIG["EMBEDDING_PROPERTY"]
|
||||
SEARCH_LABEL = INDEX_CONFIG["SEARCH_LABEL"]
|
||||
VECTOR_INDEX_NAME = INDEX_CONFIG["VECTOR_INDEX_NAME"]
|
||||
FULLTEXT_INDEX_NAME = INDEX_CONFIG["FULLTEXT_INDEX_NAME"]
|
||||
map_ins = MAP_INS
|
||||
|
||||
# ================== 全局复用:driver / embedder 只初始化一次 ==================
|
||||
_driver = None
|
||||
_embedder = None
|
||||
|
||||
def get_driver():
|
||||
global _driver
|
||||
if _driver is None:
|
||||
_driver = GraphDatabase.driver(URI, auth=AUTH)
|
||||
return _driver
|
||||
|
||||
def get_embedder():
|
||||
global _embedder
|
||||
if _embedder is None:
|
||||
_embedder = LocalBgeM3Embeddings(
|
||||
base_url=EMBEDDING_CONFIG["base_url"],
|
||||
api_key=EMBEDDING_CONFIG["api_key"],
|
||||
)
|
||||
return _embedder
|
||||
|
||||
|
||||
# ================== 辅助:解析 Record 字符串 ==================
|
||||
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
|
||||
|
||||
|
||||
# ================== 构建 Retriever ==================
|
||||
def build_retriever(sync_driver, embedder: Embedder) -> 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
|
||||
""",
|
||||
)
|
||||
|
||||
|
||||
# ================== 同步检索(在线程池中执行,不阻塞事件循环)==================
|
||||
def _hybrid_search_sync(sentence: str, top_k: int = 10) -> List[Dict]:
|
||||
retriever = build_retriever(get_driver(), get_embedder())
|
||||
results_raw = retriever.search(query_text=sentence, top_k=top_k)
|
||||
|
||||
hybrid_results: List[Dict] = []
|
||||
for item in results_raw.items:
|
||||
content = item.content
|
||||
name = "未知名称"
|
||||
labels = ["未知标签"]
|
||||
score = 0.0
|
||||
|
||||
if hasattr(content, "keys"):
|
||||
name = content.get("name", name)
|
||||
labels = content.get("labels", labels)
|
||||
score = content.get("score", score)
|
||||
elif isinstance(content, str):
|
||||
if content.startswith("<Record"):
|
||||
name, labels, score = parse_record_string(content)
|
||||
else:
|
||||
try:
|
||||
parsed = ast.literal_eval(content)
|
||||
if isinstance(parsed, dict):
|
||||
name = parsed.get("name", name)
|
||||
labels = parsed.get("labels", labels)
|
||||
score = parsed.get("score", score)
|
||||
except (ValueError, SyntaxError):
|
||||
name = 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": "",
|
||||
})
|
||||
|
||||
return hybrid_results
|
||||
|
||||
|
||||
# ================== Prompt 模板 ==================
|
||||
PROMPT_TEMPLATE = """
|
||||
请根据给出的问题,以及问题相关的检索结果,以及实体类型和sql表字段映射表,筛选出与问题相关的信息并输出。
|
||||
要求:
|
||||
1. 请仔细分析问题和检索结果,确定问题中出现的实体类型和sql表字段。
|
||||
2. 输出内容可以为充分分析后的一段文本,即你对这个问题中所用到的实体类型和sql表字段的描述。
|
||||
3. 给出的信息尽量精简,不能省略任何重要信息。
|
||||
实体类型和sql表字段映射表:
|
||||
{map_ins}
|
||||
其中键表示实体类型,值表示sql表字段。
|
||||
|
||||
检索结果:
|
||||
{results}
|
||||
|
||||
问题:{user_question}
|
||||
|
||||
|
||||
案例1:
|
||||
用户问题为:扫气箱内发现润滑油泄漏发生了多少次?
|
||||
检索知识库为:
|
||||
[ 1] score=1.0000 | 活塞杆填料函密封失效导致扫气箱内发现润滑油泄漏 标签=['故障模式']
|
||||
[ 2] score=0.9298 | 排除喷油器雾化不良故障 标签=['维修项目']
|
||||
[ 3] score=0.9115 | 执行发动机燃油泄漏应急封堵操作 标签=['操作项目']
|
||||
[ 4] score=0.9082 | 活塞环磨损或断裂导致机座与机架组件润滑油消耗异常增加 标签=['故障模式']
|
||||
[ 5] score=0.9058 | 检测气缸套内径磨损量 标签=['维修项目']
|
||||
[ 6] score=0.9049 | 拆检活塞头磨损情况 标签=['维修项目']
|
||||
[ 7] score=0.9048 | 检查排气阀杆密封面磨损 标签=['维修项目']
|
||||
[ 8] score=0.9015 | 润滑油压力骤降导致轴承组件紧急停机 标签=['故障模式']
|
||||
[ 9] score=0.8985 | 喷油器组件的维修工作 标签=['维修工作']
|
||||
[10] score=0.8979 | 发动机的安全警告 标签=['安全警告']
|
||||
输出为:
|
||||
问题中扫气箱内发现润滑油泄漏是一个故障现象,属于表中fault字段。
|
||||
"""
|
||||
|
||||
SYSTEM_PROMPT = "你是一个专业信息筛选助手,根据给出的问题和问题检索到的相关实体关系,筛选出与问题相关的信息。"
|
||||
|
||||
|
||||
# ================== 主函数(传参调用)==================
|
||||
async def get_extra_info(user_question: str, top_k: int = 20) -> str:
|
||||
"""
|
||||
参数:
|
||||
user_question: 用户问题
|
||||
top_k: 检索返回条数
|
||||
返回:
|
||||
LLM 筛选结果字符串
|
||||
"""
|
||||
# 同步检索放入线程池,不阻塞事件循环
|
||||
results = await asyncio.to_thread(_hybrid_search_sync, user_question, top_k)
|
||||
|
||||
# 构建 prompt 并异步调用 LLM
|
||||
res = await OpenaiAPI.open_api_chat_without_thinking(
|
||||
system_prompt=SYSTEM_PROMPT,
|
||||
query=PROMPT_TEMPLATE.format(results=results, map_ins=map_ins,user_question=user_question),
|
||||
)
|
||||
|
||||
return res.strip()
|
||||
|
||||
|
||||
# ================== 入口 ==================
|
||||
if __name__ == "__main__":
|
||||
answer = asyncio.run(get_extra_info("发动机这个月发生了多少次故障"))
|
||||
print(111111111111111111)
|
||||
print(answer)
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user