Team Ai
Apppublic

maikheb/nl2sql

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
app.py147 linesDownload Raw Back to root
1import streamlit as st
2import pandas as pd
3import sqlite3
4import tempfile
5from sqlalchemy import create_engine
6import os
7import re
8
9# --- Ensure vectorstore directory exists relative to src/app.py ---
10vectorstore_dir = os.path.abspath(os.path.join(os.path.dirname(__file__), "../vectorstore"))
11os.makedirs(vectorstore_dir, exist_ok=True)
12
13# If you need the FAISS index path, use:
14index_path = os.path.join(vectorstore_dir, "schema_index.faiss")
15
16from utils.schema_extractor import extract_schema_sqlite, extract_schema_rdbms
17from utils.embeddings import build_or_load_index
18from utils.llm_sql_generator import (
19    generate_sql_from_prompt,
20    generate_sql_schema_only,
21)
22from langchain_sql_pipeline import generate_sql_with_langchain
23from utils.er_diagram import render_er_diagram
24
25# --- SQL cleaning utility ---
26def clean_sql(raw_sql):
27    """
28    Removes markdown code block markers from the generated SQL.
29    Handles ```sql ... ```
30    """
31    # Remove ```sql or ``` from start
32    sql = re.sub(r"^```(?:sql)?\s*", "", raw_sql)
33    # Remove trailing ```
34    sql = re.sub(r"```$", "", sql)
35    return sql.strip()
36
37# --- Page setup ---
38st.set_page_config(
39    page_title="Text-to-SQL RAG Demo",
40    layout="wide",
41)
42
43# Hide Streamlit chrome
44st.markdown(
45    """
46    <style>
47      #MainMenu, header, footer { visibility: hidden; }
48      .css-12oz5g7, .block-container { background: #000 !important; }
49    </style>
50    """,
51    unsafe_allow_html=True,
52)
53
54# --- Sidebar for DB setup ---
55with st.sidebar:
56    st.header("Database Setup")
57    db_type = st.selectbox("Type", ["SQLite", "PostgreSQL", "MySQL"])
58    schema_data = None
59    db_connector = None
60
61    if db_type == "SQLite":
62        uploaded = st.file_uploader("Upload .db/.sqlite/.sql", type=["db", "sqlite", "sql"])
63        if uploaded:
64            tf = tempfile.NamedTemporaryFile(delete=False, suffix=".sqlite")
65            tf.write(uploaded.read())
66            tf.close()
67            schema_data = extract_schema_sqlite(tf.name)
68            db_connector = lambda q: pd.read_sql_query(q, sqlite3.connect(tf.name))
69
70    else:
71        with st.expander("Enter credentials"):
72            host = st.text_input("Host", key="host")
73            port = st.text_input("Port", "5432" if db_type == "PostgreSQL" else "3306", key="port")
74            user = st.text_input("User", key="user")
75            pwd = st.text_input("Password", type="password", key="pwd")
76            name = st.text_input("Database", key="dbname")
77        if st.button("Connect", key="connect_btn"):
78            uri = (
79                f"postgresql://{user}:{pwd}@{host}:{port}/{name}"
80                if db_type == "PostgreSQL"
81                else f"mysql+pymysql://{user}:{pwd}@{host}:{port}/{name}"
82            )
83            try:
84                schema_data = extract_schema_rdbms(uri)
85                engine = create_engine(uri)
86                db_connector = lambda q: pd.read_sql_query(q, engine)
87                st.success("Connected")
88            except Exception as e:
89                st.error(f"{e}")
90
91# --- Main panel ---
92st.title("Text-to-SQL Generator")
93
94if schema_data:
95    # Schema diagram
96    st.subheader("Schema Diagram")
97    st.graphviz_chart(render_er_diagram(schema_data))
98
99    # Question + mode
100    st.subheader("Ask Your Database")
101    q_col, m_col = st.columns((3, 1))
102    with q_col:
103        question = st.text_input(
104            "Your Question",
105            placeholder="e.g. List rock-genre tracks",
106            key="user_question",
107            label_visibility="collapsed"
108        )
109    with m_col:
110        mode = st.selectbox(
111            "Mode",
112            ["LangChain RAG", "Manual FAISS", "Schema Only"],
113            help="How to generate the SQL",
114            key="mode_select"
115        )
116
117    # Generate button
118    generate = st.button("Generate SQL", use_container_width=True, key="generate_btn")
119
120    if generate and question:
121        with st.spinner("Generating…"):
122            if mode == "LangChain RAG":
123                raw = generate_sql_with_langchain(question, schema_data)
124            elif mode == "Manual FAISS":
125                idx, meta = build_or_load_index(schema_data)
126                raw = generate_sql_from_prompt(question, idx, meta, schema_data)
127            else:
128                raw = generate_sql_schema_only(question, schema_data)
129        sql = clean_sql(raw)
130
131        # Show SQL
132        st.subheader("Generated SQL")
133        st.code(sql, language="sql")
134
135        # Show results
136        if db_connector:
137            st.subheader("Results")
138            try:
139                df = db_connector(sql)
140                st.metric("Rows returned", len(df))
141                st.dataframe(df, use_container_width=True)
142            except Exception as e:
143                st.error(f"Execution failed: {e}")
144
145else:
146    st.info("Use the sidebar to upload or connect to a database.")
147