"""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}>(.*?){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)