Text-to-SQL / app.py
sagar312's picture
Upload 2 files
13a76e8 verified
Raw History Blame Contribute Delete
17 kB
"""
Text-to-SQL demo — Hugging Face Spaces (Gradio)
Stopgap public deployment using a pretrained community model
(Ellbendls/Qwen-2.5-3b-Text_to_SQL) while the project's own fine-tuned
adapter continues training separately on Kaggle. Swap MODEL_ID below (and
add PEFT adapter loading, see the main project's generation/sql_generator.py)
once the custom adapter is ready and validated.
Reuses this project's schema retrieval (hybrid dense+BM25, in-memory
Qdrant), AST validation (read-only guardrails, bounded LIMIT), and safe
execution (read-only transaction, statement timeout) logic — these are the
same tested patterns from the main project and the Kaggle test notebook,
not reimplemented from scratch.
Required Space secrets (Settings -> Repository secrets):
DB_HOST, DB_PORT, DB_NAME, DB_USER, DB_PASSWORD
"""
import os
import re
import uuid
import gradio as gr
import pandas as pd
import spaces
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from sqlalchemy import create_engine, inspect, text as sqltext
from qdrant_client import QdrantClient
from qdrant_client.models import Distance, VectorParams, PointStruct
from sentence_transformers import SentenceTransformer
from rank_bm25 import BM25Okapi
import sqlglot
from sqlglot import exp
MODEL_ID = "Ellbendls/Qwen-2.5-3b-Text_to_SQL"
DB_HOST = os.environ["DB_HOST"]
DB_PORT = os.environ.get("DB_PORT", "5432")
DB_NAME = os.environ.get("DB_NAME", "postgres")
DB_USER = os.environ["DB_USER"]
DB_PASSWORD = os.environ["DB_PASSWORD"]
DEFAULT_ROW_LIMIT = 100
STATEMENT_TIMEOUT_MS = 5000
MAX_RETRIES = 3
# ============================================================
# Schema introspection (same exclusion list as the main project's
# schema/introspector.py — Supabase's internal schemas are platform
# plumbing, never the user's actual business tables)
# ============================================================
_EXCLUDED_SCHEMAS = {
"information_schema", "pg_catalog", "pg_toast", "mysql", "sys", "performance_schema",
"auth", "storage", "realtime", "vault", "extensions",
"graphql", "graphql_public", "pgbouncer", "pgsodium", "pgsodium_masks",
"supabase_functions", "supabase_migrations", "net", "cron",
}
class TableSchema:
def __init__(self, name, schema, columns):
self.name = name
self.schema = schema
self.columns = columns # list of (name, type, is_primary_key)
def to_ddl(self):
lines = [f"CREATE TABLE {self.schema}.{self.name} ("]
col_lines = []
for col_name, col_type, is_pk in self.columns:
pk = " PRIMARY KEY" if is_pk else ""
col_lines.append(f" {col_name} {col_type}{pk}")
lines.append(",\n".join(col_lines))
lines.append(");")
return "\n".join(lines)
def to_retrieval_document(self):
parts = [f"{self.schema}.{self.name}"] + [c[0] for c in self.columns]
return " | ".join(parts)
def extract_schema(engine):
inspector = inspect(engine)
tables = []
for schema_name in inspector.get_schema_names():
if schema_name in _EXCLUDED_SCHEMAS:
continue
for table_name in inspector.get_table_names(schema=schema_name):
try:
pk_cols = set(
inspector.get_pk_constraint(table_name, schema=schema_name).get("constrained_columns", [])
)
columns = [
(col["name"], str(col["type"]), col["name"] in pk_cols)
for col in inspector.get_columns(table_name, schema=schema_name)
]
tables.append(TableSchema(name=table_name, schema=schema_name, columns=columns))
except Exception:
continue
return tables
# ============================================================
# Hybrid retrieval — embedded Qdrant (no separate server needed
# for a Space), dense (bge-small) + BM25, fused
# ============================================================
class SchemaRetriever:
def __init__(self, top_k_final=5):
self.client = QdrantClient(":memory:")
# CPU, not GPU -- this is a tiny 133M-param embedding model; keeping
# it off any accelerator leaves all compute for the actual LLM.
self.embedder = SentenceTransformer("BAAI/bge-small-en-v1.5", device="cpu")
self.collection = "schema_metadata"
self.top_k_final = top_k_final
self._tables_by_name = {}
self._bm25 = None
self._bm25_order = []
def index(self, tables):
self._tables_by_name = {f"{t.schema}.{t.name}": t for t in tables}
docs = [t.to_retrieval_document() for t in tables]
vectors = self.embedder.encode(docs, normalize_embeddings=True)
if self.client.collection_exists(self.collection):
self.client.delete_collection(self.collection)
self.client.create_collection(
collection_name=self.collection,
vectors_config=VectorParams(size=vectors.shape[1], distance=Distance.COSINE),
)
points = [
PointStruct(id=str(uuid.uuid4()), vector=vec.tolist(), payload={"table_name": f"{t.schema}.{t.name}"})
for t, vec in zip(tables, vectors)
]
self.client.upsert(collection_name=self.collection, points=points)
tokenized = [d.lower().split() for d in docs]
self._bm25 = BM25Okapi(tokenized)
self._bm25_order = [f"{t.schema}.{t.name}" for t in tables]
def build_schema_context(self, query):
query_vec = self.embedder.encode(query, normalize_embeddings=True)
dense_hits = self.client.query_points(collection_name=self.collection, query=query_vec.tolist(), limit=8).points
dense_scores = {h.payload["table_name"]: h.score for h in dense_hits}
bm25_raw = self._bm25.get_scores(query.lower().split())
max_bm25 = max(bm25_raw) if len(bm25_raw) and max(bm25_raw) > 0 else 1.0
bm25_scores = {name: s / max_bm25 for name, s in zip(self._bm25_order, bm25_raw)}
all_names = set(dense_scores) | set(bm25_scores)
fused = sorted(
all_names,
key=lambda n: 0.6 * dense_scores.get(n, 0) + 0.4 * bm25_scores.get(n, 0),
reverse=True,
)[: self.top_k_final]
tables = [self._tables_by_name[n] for n in fused if n in self._tables_by_name]
return "\n\n".join(t.to_ddl() for t in tables)
# ============================================================
# AST validator — identical guardrails to the main project's
# validation/validator.py (verified there against 7 test cases:
# injection attempts, hallucinated tables, the aggregate/LIMIT
# edge case, etc.)
# ============================================================
_FORBIDDEN = (exp.Drop, exp.Delete, exp.Insert, exp.Update, exp.Alter, exp.Create, exp.TruncateTable, exp.Grant)
def _has_aggregate_without_group_by(select):
has_agg = any(select.find(cls) for cls in (exp.Sum, exp.Count, exp.Avg, exp.Max, exp.Min))
return has_agg and select.args.get("group") is None
def _known_tables(schema_context):
tables = set()
for line in schema_context.splitlines():
s = line.strip()
if s.upper().startswith("CREATE TABLE"):
name = s.split()[2].rstrip("(").strip()
tables.add(name.split(".")[-1].lower())
return tables
def validate_and_prepare(raw_sql, schema_context, default_limit=DEFAULT_ROW_LIMIT):
try:
statements = [s for s in sqlglot.parse(raw_sql, read="postgres") if s is not None]
except Exception as e:
return None, f"Parse error: {e}"
if len(statements) != 1:
return None, f"Expected one statement, found {len(statements)}."
tree = statements[0]
for forbidden in _FORBIDDEN:
if tree.find(forbidden):
return None, f"Rejected: forbidden operation ({forbidden.__name__})."
root_select = tree if isinstance(tree, exp.Select) else tree.find(exp.Select)
if root_select is None:
return None, "Rejected: no SELECT found."
known_tables = _known_tables(schema_context)
referenced = {t.name.lower() for t in tree.find_all(exp.Table)}
unknown = referenced - known_tables
if unknown:
return None, f"Rejected: references unknown table(s): {unknown}"
if isinstance(tree, exp.Select):
if tree.args.get("limit") is None and not _has_aggregate_without_group_by(tree):
tree.set("limit", exp.Limit(expression=exp.Literal.number(default_limit)))
try:
return tree.sql(dialect="postgres"), None
except Exception as e:
return None, f"Render error: {e}"
# ============================================================
# Safe executor — read-only transaction, statement timeout,
# always rolled back (same pattern as execution/executor.py)
# ============================================================
def execute_sql(engine, sql, timeout_ms=STATEMENT_TIMEOUT_MS):
try:
with engine.connect() as conn:
conn.execute(sqltext(f"SET statement_timeout = {timeout_ms}"))
conn.execute(sqltext("SET TRANSACTION READ ONLY"))
try:
result = conn.execute(sqltext(sql))
columns = list(result.keys())
rows = [dict(zip(columns, row)) for row in result.fetchall()]
return True, rows, columns, None
finally:
conn.rollback()
except Exception as e:
return False, [], [], str(e)
# ============================================================
# Model loading — ZeroGPU (Spaces free tier). ZeroGPU dynamically grants
# GPU access only inside functions decorated with @spaces.GPU; outside
# those functions (including at this module-level load time) no GPU is
# visible, so the model is loaded in its default location (CPU) here and
# explicitly moved to CUDA only inside the decorated inference call below.
# ============================================================
print(f"Loading {MODEL_ID} ...")
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
model = AutoModelForCausalLM.from_pretrained(MODEL_ID, dtype=torch.bfloat16)
model.eval()
print("Model loaded.")
@spaces.GPU(duration=60)
def _generate_on_gpu(prompt: str) -> str:
# Runs only while ZeroGPU has actually granted this call a GPU. Moving
# the model to CUDA on every call (rather than once at startup) is the
# expected ZeroGPU pattern -- GPU access is per-call, not persistent,
# since the same physical GPU is dynamically shared across many
# different Spaces' requests.
model.to("cuda")
inputs = tokenizer(prompt, return_tensors="pt").to("cuda")
with torch.no_grad():
output_ids = model.generate(
**inputs, max_new_tokens=256, do_sample=False, pad_token_id=tokenizer.eos_token_id
)
return tokenizer.decode(output_ids[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True)
def generate_sql(question, schema_context, error_trace=None, previous_sql=None):
# Prompt shape follows the "Context / Question / SQL Query" convention
# common to Gretel-style synthetic text-to-SQL datasets (what this
# specific model was fine-tuned on). The model card doesn't document an
# exact retry-prompt format, so retry feedback is appended in plain
# language rather than a structure this model is confirmed to have
# seen during training -- a known limitation of using a pretrained
# stopgap model instead of one trained on this project's own data.
prompt = f"Context:\n{schema_context}\n\nQuestion: {question}\n\n"
if error_trace and previous_sql:
prompt += (
f"Your previous attempt failed with this error: {error_trace}\n"
f"Previous SQL: {previous_sql}\n"
"Fix the query.\n\n"
)
prompt += "SQL Query:\n"
raw = _generate_on_gpu(prompt)
text = raw.strip()
fence = re.search(r"```(?:sql)?\s*(.*?)```", text, re.DOTALL | re.IGNORECASE)
if fence:
text = fence.group(1).strip()
if ";" in text:
text = text[: text.find(";") + 1]
return text.strip()
# ============================================================
# Startup: connect + index schema once (not per-request)
# ============================================================
print(f"Connecting to {DB_HOST} ...")
_db_uri = f"postgresql+psycopg2://{DB_USER}:{DB_PASSWORD}@{DB_HOST}:{DB_PORT}/{DB_NAME}"
db_engine = create_engine(_db_uri, pool_pre_ping=True)
_tables = extract_schema(db_engine)
print(f"Found {len(_tables)} table(s): {[t.name for t in _tables]}")
retriever = SchemaRetriever()
retriever.index(_tables)
print("Schema indexed.")
# ============================================================
# Self-healing pipeline
# ============================================================
def answer_question(question):
if not question or not question.strip():
return "", "Please enter a question.", None
schema_context = retriever.build_schema_context(question)
error_trace, previous_sql = None, None
final_sql, success, rows, columns = None, False, [], []
raw_sql = None
log_lines = []
for attempt in range(MAX_RETRIES + 1):
raw_sql = generate_sql(question, schema_context, error_trace, previous_sql)
validated_sql, validation_error = validate_and_prepare(raw_sql, schema_context)
if validated_sql is None:
error_trace = validation_error
previous_sql = raw_sql
log_lines.append(f"Attempt {attempt + 1}: validation failed — {validation_error}")
continue
success, rows, columns, exec_error = execute_sql(db_engine, validated_sql)
if success:
final_sql = validated_sql
log_lines.append(f"Attempt {attempt + 1}: succeeded.")
break
error_trace = exec_error
previous_sql = validated_sql
log_lines.append(f"Attempt {attempt + 1}: execution failed — {exec_error}")
sql_display = final_sql or raw_sql or "(no SQL generated)"
df = pd.DataFrame(rows, columns=columns) if success and rows else None
status = (
f"✅ Success after {attempt + 1} attempt(s)."
if success
else f"❌ Failed after {MAX_RETRIES + 1} attempts. Last error: {error_trace}"
)
return sql_display, status, df
# ============================================================
# Gradio UI
# ============================================================
with gr.Blocks(title="Text-to-SQL Demo") as demo:
gr.Markdown(
"# Text-to-SQL Demo\n"
"Ask a question in plain English about the sample e-commerce database "
"(customers, products, orders, order_items) and get back real SQL and results.\n\n"
"**Note:** this demo currently runs a pretrained community model "
"(`Ellbendls/Qwen-2.5-3b-Text_to_SQL`) as a stopgap while a custom "
"fine-tuned model finishes training separately — expect imperfect "
"results. Safety guardrails (read-only execution, schema-validated "
"queries, bounded result size) are fully active regardless of "
"model quality; nothing but SELECT ever reaches the database."
)
question_input = gr.Textbox(
label="Your question", placeholder="e.g. What were our top 3 customer segments by revenue?"
)
submit_btn = gr.Button("Ask", variant="primary")
sql_output = gr.Code(label="Generated SQL", language="sql")
status_output = gr.Markdown()
results_output = gr.Dataframe(label="Results")
submit_btn.click(fn=answer_question, inputs=question_input, outputs=[sql_output, status_output, results_output])
question_input.submit(fn=answer_question, inputs=question_input, outputs=[sql_output, status_output, results_output])
gr.Examples(
examples=[
"What were our top 3 customer segments by revenue?",
"Which products have the highest profit margin?",
"How many orders were refunded?",
],
inputs=question_input,
)
if __name__ == "__main__":
# server_name="0.0.0.0" is standard practice for HF Spaces deployments
# (binds to all interfaces so the Spaces proxy can reach it) -- gradio
# usually auto-detects the Spaces environment via env vars and does this
# itself, but setting it explicitly costs nothing and removes one more
# variable if the launch-time self-test misbehaves again.
# ssr_mode=False: gradio 5's server-side rendering is explicitly marked
# experimental, and its Node.js subprocess is a known, documented cause
# of silent crashes right after a clean "Running on local URL..."
# startup message on Spaces (confirmed by multiple reports on HF's own
# forums with this exact "Stopping Node.js server..." signature).
# Disabling it falls back to standard client-side rendering.
demo.launch(server_name="0.0.0.0", ssr_mode=False)