gwdoc/tools/texttosql_utils.py

185 lines
6.9 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import os
import 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)