Spaces:
Running
Running
File size: 3,348 Bytes
0fb7e57 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 | 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)}" |