File size: 17,032 Bytes
13a76e8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
"""
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)