"""TRUST-SQL — Text-to-SQL over *unknown* schemas, on ZeroGPU. Faithful re-implementation of the four-phase tool-integrated agent loop from `JaneEyre0530/TrustSQL` (`trustsql_eval/`): the model never sees the schema. It must explore it with read-only metadata queries, propose a verified schema, generate + execute a candidate SQL, and only then confirm the final answer. """ import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") import spaces # noqa: E402 (must precede torch) import json # noqa: E402 import re # noqa: E402 import sqlite3 # noqa: E402 import time # noqa: E402 from pathlib import Path # noqa: E402 from threading import Thread # noqa: E402 import gradio as gr # noqa: E402 import pandas as pd # noqa: E402 import torch # noqa: E402 from transformers import ( # noqa: E402 AutoModelForCausalLM, AutoTokenizer, StoppingCriteria, StoppingCriteriaList, TextIteratorStreamer, ) # -------------------------------------------------------------------------------------- # Model # -------------------------------------------------------------------------------------- MODEL_ID = "AIJian/TrustSQL-8B" HERE = Path(__file__).parent DB_ROOT = HERE / "databases" # Exact system prompt shipped by the authors (trustsql_eval/prompt_template.txt). SYSTEM_PROMPT = (HERE / "prompt_template.txt").read_text(encoding="utf-8").strip() tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) model = AutoModelForCausalLM.from_pretrained( MODEL_ID, dtype=torch.bfloat16, attn_implementation="sdpa", ).to("cuda") model.eval() EOS_IDS = [151645, 151643] # <|im_end|>, <|endoftext|> MAX_CONTEXT = 40960 MAX_OBS_TOKENS = 2048 # trustsql_eval default SQL_TIMEOUT = 15.0 MAX_ROWS = 100 # trustsql_eval `_execute_sql_sync` default # -------------------------------------------------------------------------------------- # Sample databases (BIRD-Dev, CC BY-SA 4.0) # -------------------------------------------------------------------------------------- SAMPLE_DBS = { "california_schools": "California public schools — SAT scores, free-meal rates (3 tables)", "superhero": "Superhero attributes, powers, publishers (9 tables)", "student_club": "University club members, events, budgets, expenses (8 tables)", "toxicology": "Molecules, atoms, bonds and carcinogenicity labels (4 tables)", "formula_1": "Formula 1 races, drivers, constructors, lap times (13 tables)", } DB_CHOICES = [f"{k} — {v}" for k, v in SAMPLE_DBS.items()] def _db_id_from_choice(choice: str) -> str: return (choice or DB_CHOICES[0]).split(" — ")[0].strip() def _resolve_db(db_choice: str, uploaded_db): """Return (db_id, sqlite_path). An uploaded file always wins.""" if uploaded_db: path = uploaded_db if isinstance(uploaded_db, str) else getattr(uploaded_db, "name", None) if path and os.path.exists(path): return Path(path).stem, path db_id = _db_id_from_choice(db_choice) return db_id, str(DB_ROOT / db_id / f"{db_id}.sqlite") # -------------------------------------------------------------------------------------- # The one tool the agent gets: read-only SQL execution # -------------------------------------------------------------------------------------- ALLOWED_SQL_PREFIXES = ("SELECT", "PRAGMA", "EXPLAIN", "WITH") def _strip_sql_comments(sql: str) -> str: s = sql.strip() while s.startswith("--") or s.startswith("/*"): if s.startswith("--"): nl = s.find("\n") if nl == -1: return "" s = s[nl + 1 :].strip() else: end = s.find("*/") if end == -1: return "" s = s[end + 2 :].strip() return s def _is_readonly(sql: str): s = _strip_sql_comments(sql) if not s: return False, "Empty SQL query" if s.upper().startswith(ALLOWED_SQL_PREFIXES): return True, None return False, f"SQL must start with {ALLOWED_SQL_PREFIXES}, got: {s.split()[0]}" def run_sql(db_path: str, sql: str, max_rows: int = MAX_ROWS): """Execute read-only SQL. Returns (text_result, column_names, rows).""" ok, err = _is_readonly(sql) if not ok: return f"Error: {err}", [], [] if not os.path.exists(db_path): return f"Error: Database file not found: {db_path}", [], [] conn = None try: conn = sqlite3.connect(f"file:{db_path}?mode=ro", uri=True, check_same_thread=False) conn.execute(f"PRAGMA busy_timeout = {int(SQL_TIMEOUT * 1000)}") cur = conn.cursor() cur.execute(sql) rows = cur.fetchall() if not rows: return "Query executed successfully. No results returned.", ( [d[0] for d in cur.description] if cur.description else [] ), [] cols = [d[0] for d in cur.description] lines = ["\t".join(cols)] for i, row in enumerate(rows): if i >= max_rows: lines.append(f"... ({len(rows) - max_rows} more rows)") break lines.append("\t".join("NULL" if v is None else str(v) for v in row)) return "\n".join(lines), cols, rows[:max_rows] except sqlite3.Error as e: return f"Error: SQLite error: {e}", [], [] except Exception as e: # pragma: no cover return f"Error: Unexpected error: {e}", [], [] finally: if conn is not None: try: conn.close() except Exception: pass def db_schema_preview(db_choice: str, uploaded_db=None) -> str: """Human-readable DDL dump of a database (for the UI only — never shown to the model).""" db_id, path = _resolve_db(db_choice, uploaded_db) if not os.path.exists(path): return f"-- database `{db_id}` not found" text, _, rows = run_sql( path, "SELECT name, sql FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%'", max_rows=200, ) if not rows: return f"-- `{db_id}`: {text}" out = [f"-- database: {db_id} ({len(rows)} tables)", ""] for name, ddl in rows: out.append((ddl or f"-- {name}").strip() + ";") out.append("") return "\n".join(out) # -------------------------------------------------------------------------------------- # Prompt construction (matches AIJian/TrustSQL-data + trustsql_eval/prompt_builders.py) # -------------------------------------------------------------------------------------- def build_user_message(db_id: str, question: str, external_knowledge: str) -> str: parts = ["", "**Task Configuration**", "**Database Engine:** SQLite", f"**Database:** {db_id}"] if external_knowledge and external_knowledge.strip(): parts.append(f"**External Knowledge:** {external_knowledge.strip()}") parts.append(f"**User Question:** {question.strip()}?") parts.append("") return "\n".join(parts) def progress_prefix(current_round: int, max_rounds: int) -> str: """Verbatim port of MessageProcessor._format_progress_prefix.""" base = f"This is turn {current_round + 1} of {max_rounds}.\n\n" remaining = max_rounds - (current_round + 1) if remaining == 0: return "" if remaining == 1: return base + ( "Only 1 turn remaining after this.\n" "You MUST provide the final answer in the next turn.\n\n" "Use confirm_answer with your best SQL query.\n" "If you don't have a complete solution, provide your best attempt.\n\n" ) if remaining == 2: return base + ("Only 2 turns remaining after this.\nStart preparing your final SQL query.\n\n") return base FORMAT_HELP = ( "Invalid format detected. Your response is missing required components.\n\n" "Option 1: EXPLORE SCHEMA\n" "Purpose: Investigate database structure\n" "Required format:\n" "Your reasoning process\n" "explore_schema\n" '{"name": "execute_sql_query", "arguments": {"db_id": "...", "sql": "..."}}\n\n' "Option 2: PROPOSE SCHEMA\n" "Purpose: Document your understanding of relevant tables and columns\n" "Required format:\n" "Your reasoning process\n" "propose_schema\n" '{"tables": [...], "columns": {...}}\n\n' "Option 3: GENERATE SQL\n" "Purpose: Create SQL query and VERIFY it works by executing\n" "Your reasoning process\n" "generate_sql\n" '{"name": "execute_sql_query", "arguments": {"db_id": "...", "sql": "..."}}\n\n' "Option 4: FINAL ANSWER\n" "Purpose: Provide verified SQL query as final result\n" "ONLY use this AFTER successfully executing and verifying your SQL.\n" "Required format:\n" "Your reasoning process\n" "confirm_answer\n" "```sql\nYOUR_SQL\n```\n\n" ) def fix_tool_tag(content: str) -> str: content = re.sub(r"(.*?)", r"\1", content, flags=re.S) content = re.sub(r"(.*?)", r"\1", content, flags=re.S) return content def extract_tag(text: str, tag: str): m = re.search(rf"<{tag}>(.*?)", text, re.S | re.I) return m.group(1) if m else None def extract_final_sql(answer_body: str) -> str: for pat in (r"```sql\s*(.*?)```", r"'''sql\s*(.*?)'''", r"```\s*(.*?)```", r"'''\s*(.*?)'''"): m = re.search(pat, answer_body, re.S | re.I) if m: return m.group(1).strip() return answer_body.strip() def truncate_observation(text: str, max_tokens: int = MAX_OBS_TOKENS) -> str: ids = tokenizer(text, add_special_tokens=False)["input_ids"] if len(ids) <= max_tokens: return text return tokenizer.decode(ids[:max_tokens]) + "\n... (result truncated due to length)" # -------------------------------------------------------------------------------------- # Pretty-printing a turn for the chat transcript # -------------------------------------------------------------------------------------- ACTION_ICON = { "explore_schema": "🔍", "propose_schema": "📋", "generate_sql": "🛠️", "confirm_answer": "✅", } def _fence(text: str, lang: str = "") -> str: return f"```{lang}\n{str(text).replace('```', '`` `')}\n```" def _quote(text: str) -> str: text = text.strip() return "\n".join("> " + line for line in text.splitlines()) if text else "" def render_assistant(raw: str) -> str: body = fix_tool_tag(raw) think = extract_tag(body, "think") action = (extract_tag(body, "action") or "").strip().lower() blocks = [] if action: blocks.append(f"### {ACTION_ICON.get(action, '⚙️')} `{action}`") elif not think: blocks.append("### ⚙️ raw response") if think: blocks.append("💭 **Reasoning**\n\n" + _quote(think)) schema = extract_tag(body, "schema") if schema is not None: try: pretty = json.dumps(json.loads(schema), indent=2) except Exception: pretty = schema.strip() blocks.append("**Proposed schema**\n\n" + _fence(pretty, "json")) answer = extract_tag(body, "answer") if answer is not None: blocks.append("**Final SQL**\n\n" + _fence(extract_final_sql(answer), "sql")) tool_call = extract_tag(body, "tool_call") if tool_call is not None: sql = None try: payload = json.loads(tool_call.strip()) sql = (payload.get("arguments") or {}).get("sql") except Exception: pass if sql: blocks.append("**Tool call** · `execute_sql_query`\n\n" + _fence(sql, "sql")) else: blocks.append("**Tool call**\n\n" + _fence(tool_call.strip(), "json")) if not blocks: return _fence(raw) return "\n\n".join(blocks) def render_observation(text: str) -> str: return "📥 **Observation**\n\n" + _fence(text) # -------------------------------------------------------------------------------------- # Generation # -------------------------------------------------------------------------------------- class _Deadline(StoppingCriteria): def __init__(self, deadline: float): self.deadline = deadline def __call__(self, input_ids, scores, **kwargs) -> bool: return time.time() > self.deadline def stream_turn(messages, max_new_tokens: int, temperature: float, top_p: float, deadline: float): """Yield incremental text for one assistant turn; last yield is the full turn.""" prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) enc = tokenizer(prompt, return_tensors="pt", add_special_tokens=False) n_in = enc["input_ids"].shape[-1] budget = max(64, min(int(max_new_tokens), MAX_CONTEXT - n_in - 8)) enc = {k: v.to(model.device) for k, v in enc.items()} streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True) do_sample = float(temperature) > 0.0 kwargs = dict( **enc, streamer=streamer, max_new_tokens=budget, do_sample=do_sample, eos_token_id=EOS_IDS, pad_token_id=151643, stopping_criteria=StoppingCriteriaList([_Deadline(deadline)]), ) if do_sample: kwargs.update(temperature=float(temperature), top_p=float(top_p), top_k=20) thread = Thread(target=model.generate, kwargs=kwargs) thread.start() acc = "" last = 0.0 for chunk in streamer: acc += chunk now = time.time() if now - last > 0.25: last = now yield acc, False thread.join() yield acc.strip(), True def _estimate_duration(*args, **kwargs) -> int: """GPU budget, sized from measured throughput (~8 s/turn at 1536 max_new_tokens).""" max_turns, max_new = 8, 1536 try: if len(args) >= 5 and args[4] is not None: max_turns = int(args[4]) if len(args) >= 6 and args[5] is not None: max_new = int(args[5]) except Exception: pass per_turn = 6.0 + 6.0 * (max_new / 1536.0) return int(min(280, 15 + max_turns * per_turn)) # -------------------------------------------------------------------------------------- # The agent loop # -------------------------------------------------------------------------------------- @spaces.GPU(duration=_estimate_duration) def run_agent( question: str, db_choice: str = DB_CHOICES[0], external_knowledge: str = "", uploaded_db=None, max_turns: int = 8, max_new_tokens: int = 1536, temperature: float = 0.7, top_p: float = 0.9, ): """Run the TRUST-SQL agent on an unknown SQLite database and return the final SQL. The model receives only the database *name* and the question — never the schema. It explores metadata with read-only queries, proposes a verified schema, executes a candidate query, and confirms the final SQL. Args: question: the natural-language question to answer. db_choice: which bundled BIRD-Dev sample database to query. external_knowledge: optional domain hint / evidence string (BIRD "evidence" field). uploaded_db: optional path to a user-supplied SQLite file; overrides `db_choice`. max_turns: maximum agent turns before giving up. max_new_tokens: token budget per agent turn. temperature: sampling temperature; 0 means greedy decoding. top_p: nucleus sampling cutoff. """ t0 = time.time() budget = _estimate_duration(question, db_choice, external_knowledge, uploaded_db, max_turns) deadline = t0 + budget - 18 max_turns = int(max_turns) db_id, db_path = _resolve_db(db_choice, uploaded_db) user_msg = build_user_message(db_id, question or "", external_knowledge or "") messages = [ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": user_msg}, ] chat = [{"role": "user", "content": f"**Question**\n\n{question}\n\n_Database: `{db_id}` (schema unknown to the model)_"}] empty_df = pd.DataFrame() if not (question or "").strip(): yield chat + [{"role": "assistant", "content": "Please enter a question."}], "", empty_df, "⚠️ No question provided." return if not os.path.exists(db_path): yield chat, "", empty_df, f"❌ Database not found: `{db_path}`" return yield chat, "", empty_df, f"⏳ Turn 1/{max_turns} — exploring `{db_id}`…" final_sql = "" status = "" for turn in range(max_turns): if time.time() > deadline: status = f"⏱️ Stopped after {turn} turn(s): GPU time budget reached." break base = list(chat) raw = "" for text, done in stream_turn(messages, max_new_tokens, temperature, top_p, deadline): raw = text chat = base + [ { "role": "assistant", "content": (render_assistant(text) if done else _fence(text)), } ] yield chat, final_sql, empty_df, f"⏳ Turn {turn + 1}/{max_turns} — generating…" if not raw: status = "❌ The model returned an empty response." break raw = fix_tool_tag(raw) # ---- confirm_answer -> terminate ------------------------------------------- answer = extract_tag(raw, "answer") if answer is not None: messages.append({"role": "assistant", "content": raw}) final_sql = extract_final_sql(answer) status = f"✅ Confirmed after {turn + 1} turn(s)." break messages.append({"role": "assistant", "content": raw}) prefix = progress_prefix(turn, max_turns) # ---- propose_schema -> acknowledgement -------------------------------------- schema = extract_tag(raw, "schema") if schema is not None: try: data = json.loads(schema) tables = data.get("tables", []) or [] cols = data.get("columns", {}) or {} n_cols = sum(len(v) for v in cols.values()) if isinstance(cols, dict) else len(cols) feedback = ( prefix + f"Schema acknowledged: {len(tables)} table(s), {n_cols} column(s). " "You may now proceed to generate SQL.\n" ) except Exception: feedback = prefix + "Schema acknowledged. You may proceed to generate SQL.\n" messages.append({"role": "user", "content": feedback}) chat = chat + [{"role": "user", "content": render_observation(feedback)}] yield chat, final_sql, empty_df, f"⏳ Turn {turn + 2}/{max_turns}…" continue # ---- tool call -> execute ---------------------------------------------------- tool_call = extract_tag(raw, "tool_call") obs = None if tool_call is None: obs = prefix + FORMAT_HELP else: try: payload = json.loads(tool_call.strip()) name = payload.get("name", "") arguments = payload.get("arguments", {}) or {} if name != "execute_sql_query": obs = prefix + f"Error: Unknown function: {name}" elif not str(arguments.get("sql", "")).strip(): obs = prefix + "Error: SQL query is empty" else: result, _, _ = run_sql(db_path, arguments["sql"]) obs = prefix + truncate_observation(result) except json.JSONDecodeError as e: obs = ( prefix + f"Tool call parsing error:\nJSON parsing failed at line {e.lineno}, column {e.colno}: {e.msg}\n\n" "Please fix the JSON format and try again.\n\n" "Required format:\n" '{"name": "execute_sql_query", "arguments": {"db_id": "...", "sql": "..."}}\n' ) except Exception as e: # pragma: no cover obs = prefix + f"Error: Tool execution error: {e}" messages.append({"role": "user", "content": obs}) chat = chat + [{"role": "user", "content": render_observation(obs)}] yield chat, final_sql, empty_df, f"⏳ Turn {turn + 2}/{max_turns}…" else: status = f"⚠️ Reached the {max_turns}-turn limit without a confirmed answer." # ---- execute the confirmed SQL for display -------------------------------------- df = empty_df if final_sql: text, cols, rows = run_sql(db_path, final_sql, max_rows=100) if cols and rows: df = pd.DataFrame(rows, columns=cols) elif cols: df = pd.DataFrame(columns=cols) else: chat = chat + [{"role": "assistant", "content": "⚠️ Final SQL did not execute:\n\n" + _fence(text)}] else: status = status or "⚠️ No SQL was confirmed." elapsed = time.time() - t0 yield chat, final_sql, df, f"{status} · {elapsed:.0f}s on GPU" # -------------------------------------------------------------------------------------- # UI # -------------------------------------------------------------------------------------- EXAMPLES = [ [ "What is the highest eligible free rate for K-12 students in the schools in Alameda County?", DB_CHOICES[0], "Eligible free rate for K-12 = `Free Meal Count (K-12)` / `Enrollment (K-12)`", ], [ "Among the schools with the SAT test takers of over 500, please list the schools that are magnet schools or offer a magnet program.", DB_CHOICES[0], "Magnet schools or offer a magnet program means that Magnet = 1", ], [ "How many superheroes have blue eyes?", DB_CHOICES[1], "blue eyes refers to colour = 'Blue' and eye_colour_id = colour.id", ], [ "Please list all the superpowers of 3-D Man.", DB_CHOICES[1], "3-D Man refers to superhero_name = '3-D Man'; superpowers refers to power_name", ], [ "What is the event that has the highest attendance of the students from the Student_Club?", DB_CHOICES[2], "event with highest attendance refers to MAX(COUNT(link_to_event))", ], [ "In the non-carcinogenic molecules, how many contain chlorine atoms?", DB_CHOICES[3], "non-carcinogenic molecules refers to label = '-'; chlorine atoms refers to element = 'cl'", ], [ "Please give the name of the race held on the circuits in Germany.", DB_CHOICES[4], "Germany is a name of country;", ], ] CSS = """ #col-container { max-width: 1180px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); } """ INTRO = """# 🔎 TRUST-SQL — Text-to-SQL over **unknown** schemas [`AIJian/TrustSQL-8B`](https://huggingface.co/AIJian/TrustSQL-8B) · [paper](https://huggingface.co/papers/2603.16448) · [code](https://github.com/JaneEyre0530/TrustSQL) The schema is **not** in the prompt. The agent gets one tool — read-only SQL — and has to discover the database itself, following the authors' four-phase protocol: `explore_schema → propose_schema → generate_sql → confirm_answer`. """ with gr.Blocks(title="TRUST-SQL") as demo: with gr.Column(elem_id="col-container"): gr.Markdown(INTRO) with gr.Row(): with gr.Column(scale=3): question = gr.Textbox( label="Question", placeholder="e.g. Which school has the highest average SAT math score?", lines=2, ) with gr.Column(scale=1, min_width=140): run_btn = gr.Button("Run agent", variant="primary", size="lg") with gr.Row(): db_choice = gr.Dropdown( label="Database (BIRD-Dev sample)", choices=DB_CHOICES, value=DB_CHOICES[0], scale=2, ) external_knowledge = gr.Textbox( label="External knowledge (optional hint)", placeholder="e.g. charter schools refers to `Charter School (Y/N)` = 1", lines=1, scale=3, ) status = gr.Markdown("") with gr.Row(): with gr.Column(scale=3): chatbot = gr.Chatbot( label="Agent trajectory", height=620, resizable=True, group_consecutive_messages=False, ) with gr.Column(scale=2): final_sql = gr.Code(label="Confirmed SQL", language="sql", lines=8) result_df = gr.Dataframe(label="Execution result", wrap=True) with gr.Accordion("Peek at the database (the agent never sees this)", open=False): schema_box = gr.Code(label="DDL", language="sql", lines=14) peek_btn = gr.Button("Show schema", size="sm") with gr.Accordion("Advanced settings", open=False): uploaded_db = gr.File( label="Use your own SQLite database (.sqlite / .db) — overrides the dropdown", file_types=[".sqlite", ".db", ".sqlite3"], type="filepath", ) with gr.Row(): max_turns = gr.Slider(3, 12, value=8, step=1, label="Max agent turns") max_new_tokens = gr.Slider(256, 3072, value=1536, step=128, label="Max new tokens / turn") with gr.Row(): temperature = gr.Slider(0.0, 1.0, value=0.7, step=0.05, label="Temperature (0 = greedy)") top_p = gr.Slider(0.1, 1.0, value=0.9, step=0.05, label="Top-p") gr.Markdown( "Defaults mirror `trustsql_eval` (temperature 0.7 / top-p 0.9). " "Only `SELECT` / `PRAGMA` / `EXPLAIN` / `WITH` statements are ever executed, " "against a read-only connection." ) gr.Examples( examples=EXAMPLES, inputs=[question, db_choice, external_knowledge], outputs=[chatbot, final_sql, result_df, status], fn=run_agent, cache_examples=True, cache_mode="lazy", label="BIRD-Dev examples (question + official evidence hint)", ) gr.Markdown( "Sample databases are the **BIRD-Dev** SQLite databases " "([BIRD-SQL](https://bird-bench.github.io/), CC BY-SA 4.0); the example questions and " "hints are their official dev-set questions and `evidence` strings. " "The system prompt and agent loop are ported verbatim from " "[`JaneEyre0530/TrustSQL`](https://github.com/JaneEyre0530/TrustSQL) (Apache-2.0)." ) run_btn.click( fn=run_agent, inputs=[ question, db_choice, external_knowledge, uploaded_db, max_turns, max_new_tokens, temperature, top_p, ], outputs=[chatbot, final_sql, result_df, status], api_name="run_agent", ) question.submit( fn=run_agent, inputs=[ question, db_choice, external_knowledge, uploaded_db, max_turns, max_new_tokens, temperature, top_p, ], outputs=[chatbot, final_sql, result_df, status], api_name=False, ) peek_btn.click( fn=db_schema_preview, inputs=[db_choice, uploaded_db], outputs=schema_box, api_name="schema" ) db_choice.change(fn=db_schema_preview, inputs=[db_choice, uploaded_db], outputs=schema_box, api_name=False) if __name__ == "__main__": demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)