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)}"