Team Ai
Apppublic

kavyabammidi/text2sql

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py110 linesDownload Raw Back to root
1from dotenv import load_dotenv2import streamlit as st3import os4import sqlite35import google.generativeai as genai6 7# Load environment variables8load_dotenv()9 10# Configure Google Gemini API11API_KEY = 'AIzaSyD3O00WtFIB3cVoLp36sT7reWuI2c09jr4'12if API_KEY is None:13    st.error("API Key is missing! Please check your .env file.")14else:15    genai.configure(api_key=API_KEY)16 17# Function to generate SQL query from natural language input18def get_gemini_response(question, prompt):19    try:20        model = genai.GenerativeModel("gemini-1.5-pro")  # Use latest model21        response = model.generate_content([prompt[0], question])22        return response.text.strip()  # Ensure clean output23    except Exception as e:24        st.error(f"Error fetching response from Gemini AI: {e}")25        return None26 27# Function to clean SQL query28def clean_sql_query(sql):29    # Correct the case for section values30    if "section=" in sql:31        section_value = sql.split("section=")[1].split("'")[1]32        if section_value.islower():33            sql = sql.replace(f"section='{section_value}'", f"section='{section_value.upper()}'")34    return sql35 36# Function to execute SQL query37def read_sql_query(sql, db):38    conn = None39    try:40        # Basic validation to ensure the query starts with SELECT41        if not sql.strip().lower().startswith("select"):42            st.error("Invalid SQL query. Only SELECT queries are allowed.")43            return []44        45        conn = sqlite3.connect(db)46        cur = conn.cursor()47        cur.execute(sql)48        rows = cur.fetchall()49        return rows50    except Exception as e:51        st.error(f"Error executing SQL query: {e}")52        return []53    finally:54        if conn:55            conn.close()56 57# SQL Prompt for Gemini AI58prompt = [59    """60    You are an expert in converting English questions into SQL queries. 61    The database is named 'student.db' and has a table named 'student' with columns: 62    name, class, section, and marks.63 64    Rules:65    1. Always use lowercase for table and column names (e.g., `student`, `name`, `class`, `section`, `marks`).66    2. Always enclose string values in single quotes and match the exact case as stored in the database (e.g., 'A', 'B', 'C' for section).67    3. Do not include any unnecessary subqueries or complex logic unless explicitly required.68    4. Ensure the SQL query is simple, efficient, and directly executable.69    5. **Only return the SQL query. Do not include any explanations, additional text, or formatting.**70    6.remove all symbols from the query you have give only the query 71    Example 1:72    Question: "What is the average marks of all students?"73    SQL Command: SELECT AVG(marks) FROM student;74 75    Example 2:76    Question: "List all students in section A."77    SQL Command: SELECT * FROM student WHERE section='A';78 79    Example 3:80    Question: "What is the highest marks in the Data Science class?"81    SQL Command: SELECT MAX(marks) FROM student WHERE class='datascience';82 83    Example 4:84    Question: "What is the second highest marks in the Data Science class?"85    SQL Command: SELECT marks FROM student WHERE class='datascience' ORDER BY marks DESC LIMIT 1 OFFSET 1;86    """87]88 89# Streamlit App UI90st.set_page_config(page_title="SQL Query Generator & Retriever")91st.header("AI-powered SQL Query Generator")92 93# User Input94question = st.text_input("Enter your query:", key="input")95submit = st.button("Generate SQL & Retrieve Data")96 97# Handling Submission98if submit:99    response = get_gemini_response(question, prompt)100 101    if response:102        response = clean_sql_query(response)  # Clean the SQL query103        st.subheader("Generated SQL Query:")104        st.code(response, language="sql")  # Display the generated SQL query105        106        data = read_sql_query(response, "student.db")107 108        st.subheader("Query Results:")109        for row in data:110            st.write(row)  # Display query results