Spaces:
Running
Running
Download src/tools.py from LightRT/text2sql_backend: direct link, hf CLI and curl.
- Browser
- Download file 3.35 kB
-
https://huggingface.co/spaces/LightRT/text2sql_backend/resolve/main/src/tools.py
- Command line
-
hf download hf://spaces/LightRT/text2sql_backend/src/tools.py
-
curl -L -o tools.py https://huggingface.co/spaces/LightRT/text2sql_backend/resolve/main/src/tools.py
3.35 kB
| 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 | |
| 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 | |
| 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)}" |