Team Ai
Apppublic

dvwn/nl2sql-api

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
sql_agent.py174 linesDownload Raw Back to nl2sql
1# Path: src/nl2sql/sql_agent.py2# SQL Agent for handling NL2SQL conversion with Auto-Correct functionality3from src.database.db_manager import get_db_connection, get_schema_context4from langchain_core.prompts import PromptTemplate5from src.nl2sql.hf_engine import get_llm6 7# Craft the Prompt Template to instruct LLM on its persona8SQL_PROMPT_TEMPLATE = """You are an expert SQLite developer.9Your task is to write a syntactically correct SQLite query to answer the user's question based strictly on the provided database schema.10Return ONLY the raw SQL query.11Do not include any explanations, markdown formatting, or code blocks.12 13Schema Context:14{schema}15 16User Question:17{question}18 19SQL Query:"""20 21REFINEMENT_PROMPT_TEMPLATE = """You are an expert SQLite developer.22You previously generated a SQL query to answer the user's question, but it resulted inan error when executed on the database.23 24Schema Context:25{schema}26 27User Question:28{question}29 30Previous Generated SQL:31{failed_sql}32 33Database Error Message:34{error_message}35 36Your task is to fix the SQL query based on the exact error message and the schema.37Pay close attention to column names, table relationships, and SQLite syntax.38Return ONLY the raw, corrected SQL query. 39Do not include any explanations, markdown formatting, or code blocks.40 41Corrected SQL Query:"""42 43# Generate text response44NL_RESPONSE_TEMPLATE = """You are a helpful data assisstant.45The user asked the following question: "{question}"46The database returned the following results: {results}47 48Provide a direct, natural language answer to the user's question using ONLY the provided data.49Keep it brief. Do not explain the SQL query or mention the database schema.50If the database returns more than 5 rows, DO NOT list the items individually. Instead, provide a brief summary sentence.51 52Answer:"""53 54prompt_template = PromptTemplate(55    input_variables = ["schema", "question"],56    template = SQL_PROMPT_TEMPLATE57)58 59refinement_prompt = PromptTemplate(60    input_variables = ["schema", "question", "failed_sql", "error_message"],61    template = REFINEMENT_PROMPT_TEMPLATE62)63 64nl_response_template = PromptTemplate(65    input_variables = ["question", "results"],66    template = NL_RESPONSE_TEMPLATE67)68 69# Clean the output70def clean_sql(raw_sql: str) -> str:71    """72    Utility to strip markdown formatting if the LLM hallucinated code blocks.73    Ensure the raw string can be directly executed by the SQLite cursor.74    """75    cleaned = raw_sql.strip()76    if cleaned.startswith("```sql"):77        cleaned = cleaned[6:]78    elif cleaned.startswith("```"):79        cleaned = cleaned[3:]80    81    if cleaned.endswith("```"):82        cleaned = cleaned[:-3]83    84    return cleaned.strip()85 86# Function to handle NL2SQL conversion87def nl2sql_agent(user_question: str, max_retries: int = 3, model_id: str = None) -> dict:88    """89    Complete flow execution with Auto-correction:90    Get Schema context -> Generate SQL query -> Execute SQL query -> If Error, Refine & Retry ->Return results91    """92    # Fetch database schema context using RAG93    print(f"Fetching RAG schema context for: '{user_question}'...")94    schema = get_schema_context(question = user_question)95 96    # Generate SQL query using the schema context and user question97    llm = get_llm(model_id=model_id)98 99    # LangChain Pipeline: Pipe prompt into LLM100    chain = prompt_template | llm101    refinement_chain = refinement_prompt | llm102    nl_chain = nl_response_template | llm103 104    current_sql = ""105    error_message = ""106 107    # Auto-correction Loop108    for attempt in range(1, max_retries + 1):109        if attempt == 1:110            print(f"Generating initial SQL query using {model_id or 'default model'}...")111            raw_response = chain.invoke({112                "schema": schema,113                "question": user_question114            })115        else:116            print(f"\n--- Attempt {attempt}/{max_retries}: Refining SQL query based on error ---")117            print(f"Feeding error back to LLM: {error_message}")118            raw_response = refinement_chain.invoke({119                "schema": schema,120                "question": user_question,121                "failed_sql": current_sql,122                "error_message": error_message123            })124 125        # Parse & clean the generated SQL query126        generated_sql = clean_sql(raw_response)127        current_sql = generated_sql128        print(f"Generated SQL: \n{generated_sql}")129 130        # Execute the generated SQL query and fetch results131        connection = get_db_connection()132        if not connection:133            return {134                "query": generated_sql,135                "error": "Could not establish database connection",136                "status": "failed"137            }138        139        try:140            cursor = connection.cursor()141            cursor.execute(generated_sql)142            results = cursor.fetchall()143            144            if attempt > 1:145                print(f"SQL query executed successfully after {attempt} attempts.")146            147            # Generate natural language response based on the results148            print("Generating natural language response based on query results...")149            nl_response = nl_chain.invoke({150                "question": user_question,151                "results": str(results)152            })153 154            return {155                "query": generated_sql,156                "results": results,157                "nl_response": nl_response,158                "status": "success",159                "attempts": attempt160            }161        except Exception as e:162            error_message = str(e)163            print(f"Error executing SQL: {error_message}")164 165            if attempt == max_retries:166                print("Max retries reached. Returning error.")167        finally:168            connection.close()169    170    return {171        "query": current_sql,172        "error": error_message,173        "status": f"Error executing SQL after {max_retries} attempts"174    }