diff --git a/tools/texttosql_utils.py b/tools/texttosql_utils.py index e69de29..bb9da46 100644 --- a/tools/texttosql_utils.py +++ b/tools/texttosql_utils.py @@ -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(" 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) +