Spaces:
Running on Zero
Running on Zero
Download app.py from sagar312/Text-to-SQL: direct link, hf CLI and curl.
- Browser
- Download file 17 kB
-
https://huggingface.co/spaces/sagar312/Text-to-SQL/resolve/main/app.py
- Command line
-
hf download hf://spaces/sagar312/Text-to-SQL/app.py
-
curl -L -o app.py https://huggingface.co/spaces/sagar312/Text-to-SQL/resolve/main/app.py
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.") | |
| 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) | |