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