252 lines
9.4 KiB
Python
252 lines
9.4 KiB
Python
# -*- coding: utf-8 -*-
|
||
import os
|
||
import ast
|
||
import re
|
||
import asyncio
|
||
import httpx
|
||
from typing import List, Dict
|
||
from openai import AsyncOpenAI
|
||
|
||
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
|
||
from config import LLM_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_v2"],
|
||
api_key=EMBEDDING_CONFIG["api_key"],
|
||
)
|
||
return _embedder
|
||
|
||
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]]:
|
||
url = self.base_url + "/embeddings"
|
||
with httpx.Client(timeout=60) as client:
|
||
response = client.post(
|
||
url,
|
||
json={"model": self.model_name, "input": texts},
|
||
headers={"Authorization": f"Bearer {self.api_key}"},
|
||
)
|
||
return [item["embedding"] for item in response.json()["data"]]
|
||
async def open_api_chat_without_thinking(
|
||
query: str = None,
|
||
model: str = None,
|
||
json_output: bool = False,
|
||
system_prompt: str = None,
|
||
messages: list = None,
|
||
enable_thinking: bool = True,
|
||
temperature: float = 0.1
|
||
) -> str:
|
||
if model is None:
|
||
model = LLM_CONFIG["model"]
|
||
if system_prompt is None:
|
||
system_prompt = "我是一个专业的AI助手,请用中文回答问题。请直接给出最终答案,不要包含任何推理步骤、思考过程、解释性文字或前缀。"
|
||
|
||
|
||
# 构建消息
|
||
message_list = [{"role": "system", "content": system_prompt}]
|
||
|
||
if messages is not None:
|
||
message_list.extend(messages)
|
||
if query is not None:
|
||
message_list.append({"role": "user", "content": query})
|
||
|
||
client = AsyncOpenAI(
|
||
api_key=LLM_CONFIG["api_key"],
|
||
base_url=LLM_CONFIG["base_url"],
|
||
)
|
||
|
||
kwargs = {
|
||
"model": model,
|
||
"messages": message_list,
|
||
"temperature": temperature,
|
||
"stream": False,
|
||
"max_tokens":LLM_CONFIG["max_tokens"]
|
||
}
|
||
|
||
kwargs["extra_body"] = {"chat_template_kwargs": {"enable_thinking": False}}
|
||
|
||
if json_output:
|
||
kwargs["response_format"] = {"type": "json_object"}
|
||
|
||
print("model API. history message", message_list)
|
||
response = await client.chat.completions.create(**kwargs)
|
||
return response.choices[0].message.content
|
||
# ================== 辅助:解析 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表字段映射表:
|
||
"故障模式" : "fault",
|
||
"系统" : "system_name",
|
||
"设备" : "device_name",
|
||
"舷号" : "ship_number",
|
||
其中键表示实体类型,值表示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 open_api_chat_without_thinking(
|
||
system_prompt=SYSTEM_PROMPT,
|
||
query=PROMPT_TEMPLATE.format(results=results, user_question=user_question),
|
||
)
|
||
|
||
return res.strip()
|
||
|
||
|
||
# ================== 入口 ==================
|
||
if __name__ == "__main__":
|
||
answer = asyncio.run(get_extra_info("发动机这个月发生了多少次故障"))
|
||
print(111111111111111111)
|
||
print(answer)
|
||
|