Team Ai
Apppublic

LightRT/text2sql_backend

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes
tools.py103 linesDownload Raw Back to src
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)}"