# -*- 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(" 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)