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