Team Ai
Apppublic

Sharath7693/text2sql_rag

sourceHugging Facemitupdated 2y agoView on Hugging Face
1likes
app.py144 linesDownload Raw Back to root
1import gradio as gr2import pdfplumber3import docx4import json5import re6import sqlalchemy7import requests8from tenacity import retry, stop_after_attempt, wait_exponential9from langchain_postgres.vectorstores import PGVector10from langchain_core.documents import Document11from langchain_community.embeddings import HuggingFaceEmbeddings12from langchain_postgres import PGVector13 14# API Keys and Database Connection15GROQ_API_KEY = "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"16NEON_CONNECTION_STRING = "postgresql://neondb_owner:npg_TPtUn1ArS6bo@ep-crimson-king-a12tfmdw-pooler.ap-southeast-1.aws.neon.tech/neondb?sslmode=require"17embeddings = HuggingFaceEmbeddings(model_name="BAAI/bge-small-en")18 19# Extract text from various document types20def extract_text_from_doc(file_path):21    if file_path.endswith(".pdf"):22        with pdfplumber.open(file_path) as pdf:23            return "\n".join([page.extract_text() for page in pdf.pages if page.extract_text()])24    elif file_path.endswith(".docx"):25        doc = docx.Document(file_path)26        return "\n".join([p.text for p in doc.paragraphs])27    elif file_path.endswith(".txt"):28        with open(file_path, "r", encoding="utf-8") as f:29            return f.read()30    return ""31 32# Database Connection33engine = sqlalchemy.create_engine(url=NEON_CONNECTION_STRING, pool_pre_ping=True, pool_recycle=300)34vector_store = PGVector(embeddings=embeddings, connection=engine, use_jsonb=True, collection_name="text-to-sql-context")35 36# Retry for API Calls37@retry(stop=stop_after_attempt(5), wait=wait_exponential(multiplier=1, min=2, max=10))38def call_groq_api(prompt):39    headers = {40        "Content-Type": "application/json",41        "Authorization": f"Bearer {GROQ_API_KEY}",42    }43    data = {44        "model": "llama-3.3-70b-versatile",45        "messages": [{"role": "user", "content": prompt}]46    }47    response = requests.post("https://api.groq.com/openai/v1/chat/completions", headers=headers, json=data)48 49    if response.status_code != 200:50        raise Exception(f"Groq API error: {response.text}")51 52    result = response.json()53    return result.get("choices", [{}])[0].get("message", {}).get("content", "").strip()54 55# Remove extra text and extract only JSON56def extract_json(text):57    match = re.search(r"\[.*\]", text, re.DOTALL)58    if match:59        return match.group(0)  # Extract only the JSON array60    return None  # Invalid format61 62# SQL Query Generation Prompt63generation_prompt = """64Generate 50 SQL queries based on this schema:65<schema>66{SCHEMA}67</schema>68Provide JSON output with 'question' and 'query'.69"""70 71# Process Schema, Generate Queries, and Answer User's Question72def process_and_query(file, question):73    schema_text = extract_text_from_doc(file.name)74    75    # Generate Queries76    response = call_groq_api(generation_prompt.format(SCHEMA=schema_text))77    78    # Extract only valid JSON part79    json_response = extract_json(response)80    if not json_response:81        return f"Error: Unexpected response format from Groq API: {response}"82    83    try:84        qa_pairs = json.loads(json_response)85    except json.JSONDecodeError:86        return f"Error: Could not parse JSON: {json_response}"87    88    # Store Schema and Queries in Vector DB89    schema_doc = Document(page_content=schema_text, metadata={"id": "schema", "topic": "ddl"})90    query_docs = [Document(page_content=json.dumps(pair), metadata={"id": f"query-{i}", "topic": "query"}) for i, pair in enumerate(qa_pairs)]91    92    vector_store.add_documents([schema_doc] + query_docs, ids=[doc.metadata["id"] for doc in [schema_doc] + query_docs])93    94    # Retrieve Relevant Schema and Queries95    relevant_ddl = vector_store.similarity_search(query=question, k=5, filter={"topic": {"$eq": "ddl"}})96    similar_queries = vector_store.similarity_search(query=question, k=3, filter={"topic": {"$eq": "query"}})97    98    schema = "\n".join([doc.page_content for doc in relevant_ddl])99    examples = "\n".join([json.loads(doc.page_content)["question"] + "\nSQL: " + json.loads(doc.page_content)["query"] for doc in similar_queries])100    101    query_prompt = f"""102    You are an SQL expert. Generate a valid SQL query based on the schema and example queries.103 104    1. Some DDL statements describing tables, columns and indexes in the database:105    <schema>106    {schema}107    </schema>108 109    2. Some example pairs demonstrating how to convert natural language text into a corresponding SQL query for this schema:  110    <examples>111    {examples}112    </examples>113 114    3. The actual natural language question to convert into an SQL query:115    <question>116    {question}117    </question>118 119    Follow the instructions below:120    1. Your task is to generate an SQL query that will retrieve the data needed to answer the question, based on the database schema. 121    2. First, carefully study the provided schema and examples to understand the structure of the database and how the examples map natural language to SQL for this schema.122    3. Your answer should have two parts: 123    - Inside <scratchpad> XML tag, write out step-by-step reasoning to explain how you are generating the query based on the schema, example, and question. 124    - Then, inside <sql> XML tag, output your generated SQL. 125 126    SQL Query:127    """128 129    query_response = call_groq_api(query_prompt)130    return query_response131 132# Gradio UI133with gr.Blocks() as app:134    gr.Markdown("# Text-to-SQL Converter")135    136    file_input = gr.File(label="Upload Schema File")137    question_input = gr.Textbox(label="Ask a SQL-related Question")138    submit_button = gr.Button("Process & Generate SQL")139    query_output = gr.Textbox(label="Generated SQL Query")140    141    submit_button.click(process_and_query, inputs=[file_input, question_input], outputs=query_output)142 143app.launch(share=True)144