maikheb/nl2sql
0
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 