from langchain_chroma import Chroma from langchain_huggingface import HuggingFaceEmbeddings from langchain_community.retrievers import BM25Retriever from langchain_classic.retrievers import EnsembleRetriever from langchain_core.documents import Document from langchain_community.utilities import SQLDatabase from langchain_core.tools import tool from langgraph.runtime import get_runtime import re import chromadb import os from dotenv import load_dotenv load_dotenv() COLLECTION_NAME = "Text2SQL" chroma_client = chromadb.CloudClient( api_key=os.getenv("CHROMA_API_KEY"), tenant=os.getenv("CHROMA_TENANT"), database=os.getenv("CHROMA_DATABASE"), ) BLOCKED_KEYWORDS = ["INSERT", "UPDATE", "DELETE", "DROP", "ALTER","TRUNCATE", "CREATE", "GRANT", "REVOKE", "REPLACE"] embedding_model = HuggingFaceEmbeddings(model_name="sentence-transformers/all-MiniLM-L6-v2") vectorstore = Chroma(collection_name=COLLECTION_NAME,embedding_function=embedding_model,client=chroma_client) _bm25_cache = {} def execution_guardrail(sql_query: str) -> bool : query = sql_query.strip().rstrip(";") if ";" in query : return False if not query.upper().startswith("SELECT") : return False for keyword in BLOCKED_KEYWORDS : if re.search(rf"\b{keyword}\b" , query , re.IGNORECASE) : return False return True @tool def retrieve(query: str) -> str: """Retrieve the relevant database schema (tables and columns) needed to answer the user's question.""" runtime = get_runtime() user_id = runtime.context.user_id connection_url = runtime.context.connection_url semantic_retriever = vectorstore.as_retriever(search_kwargs={"k" : 10 , "filter" : {"user_id" : user_id}}) if user_id not in _bm25_cache : user_docs_raw = vectorstore.get(where={"user_id" : user_id}) user_documents = [Document(page_content=text , metadata=meta)for text , meta in zip(user_docs_raw['documents'] , user_docs_raw['metadatas'])] bm25_retriever = BM25Retriever.from_documents(user_documents) bm25_retriever.k = 10 _bm25_cache[user_id] = bm25_retriever bm25_retriever = _bm25_cache[user_id] ensemble_retriever = EnsembleRetriever(retrievers=[semantic_retriever , bm25_retriever],weights=[0.5 , 0.5]) results = ensemble_retriever.invoke(query) tables = [] for doc in results : table = doc.metadata['table_name'] if table not in tables : tables.append(table) db = SQLDatabase.from_uri(connection_url , sample_rows_in_table_info=0) dialect = db.dialect final_schemes = f"Dialect : {dialect}\n {db.get_table_info(table_names=tables)}\n" return final_schemes @tool def execute_query(sql_query: str) -> str : """Execute a validated, read-only SQL SELECT query against the connected database and return the raw results.""" runtime = get_runtime() connection_url = runtime.context.connection_url if not execution_guardrail(sql_query) : return "Error: This query was blocked. Only a single SELECT statement is allowed — no INSERT, UPDATE, DELETE, DROP, ALTER, or multi-statement queries." db = SQLDatabase.from_uri(connection_url) try : result = db.run(sql_query) return result except Exception as e : return f"Error : {str(e)}"