dvwn/nl2sql-api
0
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 }