LightRT/text2sql_backend
0
1from langchain_chroma import Chroma2from langchain_huggingface import HuggingFaceEmbeddings3from langchain_community.retrievers import BM25Retriever4from langchain_classic.retrievers import EnsembleRetriever5from langchain_core.documents import Document6from langchain_community.utilities import SQLDatabase7from langchain_core.tools import tool8from langgraph.runtime import get_runtime9import re10import chromadb11import os12from dotenv import load_dotenv13 14load_dotenv()15 16COLLECTION_NAME = "Text2SQL"17 18chroma_client = chromadb.CloudClient(19 api_key=os.getenv("CHROMA_API_KEY"),20 tenant=os.getenv("CHROMA_TENANT"),21 database=os.getenv("CHROMA_DATABASE"),22)23 24BLOCKED_KEYWORDS = ["INSERT", "UPDATE", "DELETE", "DROP", "ALTER","TRUNCATE", "CREATE", "GRANT", "REVOKE", "REPLACE"]25 26embedding_model = HuggingFaceEmbeddings(model_name="sentence-transformers/all-MiniLM-L6-v2")27 28vectorstore = Chroma(collection_name=COLLECTION_NAME,embedding_function=embedding_model,client=chroma_client)29 30_bm25_cache = {}31 32def execution_guardrail(sql_query: str) -> bool :33 query = sql_query.strip().rstrip(";")34 35 if ";" in query :36 return False37 38 if not query.upper().startswith("SELECT") :39 return False40 41 for keyword in BLOCKED_KEYWORDS :42 if re.search(rf"\b{keyword}\b" , query , re.IGNORECASE) :43 return False44 45 return True46 47@tool48def retrieve(query: str) -> str:49 """Retrieve the relevant database schema (tables and columns) needed to answer the user's question."""50 51 runtime = get_runtime()52 user_id = runtime.context.user_id53 connection_url = runtime.context.connection_url54 55 semantic_retriever = vectorstore.as_retriever(search_kwargs={"k" : 10 , "filter" : {"user_id" : user_id}})56 57 if user_id not in _bm25_cache :58 user_docs_raw = vectorstore.get(where={"user_id" : user_id})59 60 user_documents = [Document(page_content=text , metadata=meta)for text , meta in zip(user_docs_raw['documents'] , user_docs_raw['metadatas'])]61 62 bm25_retriever = BM25Retriever.from_documents(user_documents)63 bm25_retriever.k = 1064 65 _bm25_cache[user_id] = bm25_retriever66 67 bm25_retriever = _bm25_cache[user_id]68 69 ensemble_retriever = EnsembleRetriever(retrievers=[semantic_retriever , bm25_retriever],weights=[0.5 , 0.5])70 71 results = ensemble_retriever.invoke(query)72 73 tables = []74 for doc in results :75 table = doc.metadata['table_name']76 if table not in tables :77 tables.append(table)78 79 db = SQLDatabase.from_uri(connection_url , sample_rows_in_table_info=0)80 81 dialect = db.dialect82 83 final_schemes = f"Dialect : {dialect}\n {db.get_table_info(table_names=tables)}\n"84 85 return final_schemes86 87@tool88def execute_query(sql_query: str) -> str :89 """Execute a validated, read-only SQL SELECT query against the connected database and return the raw results."""90 91 runtime = get_runtime()92 connection_url = runtime.context.connection_url93 94 if not execution_guardrail(sql_query) :95 return "Error: This query was blocked. Only a single SELECT statement is allowed — no INSERT, UPDATE, DELETE, DROP, ALTER, or multi-statement queries."96 97 db = SQLDatabase.from_uri(connection_url)98 99 try :100 result = db.run(sql_query)101 return result102 except Exception as e :103 return f"Error : {str(e)}"