joelgilbert/NL2SQL
0
1"""2Agent 2: SQL Generator - Generate PostgreSQL queries using SQLCoder-7B-2.3"""4 5import logging6import requests7import re8from typing import Optional9from tenacity import retry, stop_after_attempt, wait_exponential10 11from config.settings import settings12from config.prompts import PromptTemplates13 14logger = logging.getLogger(__name__)15 16 17class SQLGeneratorAgent:18 """Agent responsible for SQL query generation and correction."""19 20 def __init__(self):21 """Initialize SQL generator with Cloudflare Workers AI credentials."""22 self.account_id = settings.api.cloudflare_account_id23 self.auth_token = settings.api.cloudflare_auth_token24 self.api_url = f"https://api.cloudflare.com/client/v4/accounts/{self.account_id}/ai/run/@cf/defog/sqlcoder-7b-2"25 self.headers = {26 "Authorization": f"Bearer {self.auth_token}",27 "Content-Type": "application/json"28 }29 30 @retry(stop=stop_after_attempt(3), wait=wait_exponential(multiplier=1, min=2, max=10))31 def generate_sql(self, question: str, schema_context: str, examples: str) -> str:32 """33 Generate SQL query using SQLCoder-7B-2 via Cloudflare Workers AI.34 35 Args:36 question: Natural language question37 schema_context: Database schema information38 examples: Similar successful queries39 40 Returns:41 Generated SQL query string42 43 Raises:44 Exception: If SQL generation fails45 """46 try:47 prompt = PromptTemplates.sql_generation_prompt(question, schema_context, examples)48 49 payload = {50 "messages": [51 {"role": "system", "content": PromptTemplates.SQL_GENERATION_SYSTEM},52 {"role": "user", "content": prompt}53 ],54 "max_tokens": 500,55 "temperature": 0.256 }57 58 response = requests.post(59 self.api_url,60 headers=self.headers,61 json=payload,62 timeout=3063 )64 65 response.raise_for_status()66 67 result = response.json()68 69 # Extract SQL from response70 if 'result' in result and 'response' in result['result']:71 sql_response = result['result']['response']72 elif 'response' in result:73 sql_response = result['response']74 else:75 raise ValueError(f"Unexpected response format: {result}")76 77 # Clean and extract SQL78 sql = self.extract_sql_from_response(sql_response)79 80 logger.info(f"Generated SQL: {sql[:100]}...")81 return sql82 83 except requests.RequestException as e:84 logger.error(f"Cloudflare API request failed: {e}")85 raise86 87 except Exception as e:88 logger.error(f"SQL generation failed: {e}")89 raise90 91 def correct_sql(self, failed_sql: str, error_message: str, error_type: str, schema_context: str) -> str:92 """93 Attempt to fix SQL based on error classification.94 95 Args:96 failed_sql: SQL query that failed97 error_message: Error message from database98 error_type: Classified error type99 schema_context: Database schema information100 101 Returns:102 Corrected SQL query103 104 Raises:105 Exception: If correction fails106 """107 try:108 prompt = PromptTemplates.error_correction_prompt(109 failed_sql=failed_sql,110 error_message=error_message,111 error_type=error_type,112 schema=schema_context113 )114 115 payload = {116 "messages": [117 {"role": "system", "content": PromptTemplates.ERROR_CORRECTION_SYSTEM},118 {"role": "user", "content": prompt}119 ],120 "max_tokens": 500,121 "temperature": 0.1 # Lower temperature for corrections122 }123 124 response = requests.post(125 self.api_url,126 headers=self.headers,127 json=payload,128 timeout=30129 )130 131 response.raise_for_status()132 133 result = response.json()134 135 # Extract SQL from response136 if 'result' in result and 'response' in result['result']:137 sql_response = result['result']['response']138 elif 'response' in result:139 sql_response = result['response']140 else:141 raise ValueError(f"Unexpected response format: {result}")142 143 # Clean and extract SQL144 sql = self.extract_sql_from_response(sql_response)145 146 logger.info(f"Corrected SQL: {sql[:100]}...")147 return sql148 149 except Exception as e:150 logger.error(f"SQL correction failed: {e}")151 raise152 153 def extract_sql_from_response(self, response: str) -> str:154 """155 Extract clean SQL from LLM response.156 157 Args:158 response: Raw LLM response159 160 Returns:161 Clean SQL query162 """163 # Remove markdown code blocks164 sql = re.sub(r'```sql\n?', '', response)165 sql = re.sub(r'```\n?', '', sql)166 167 # Remove common prefixes168 prefixes = ['SQL:', 'Query:', 'Answer:', 'Here is the SQL:']169 for prefix in prefixes:170 if sql.strip().startswith(prefix):171 sql = sql.replace(prefix, '', 1)172 173 # Strip whitespace174 sql = sql.strip()175 176 # Remove trailing semicolon if present (we'll add it when executing)177 if sql.endswith(';'):178 sql = sql[:-1].strip()179 180 # Remove any explanatory text after the query181 # Look for common sentence starters182 explanation_markers = [183 '\n\nThis query',184 '\n\nThe above',185 '\n\nNote:',186 '\n\nExplanation:',187 ]188 189 for marker in explanation_markers:190 if marker in sql:191 sql = sql.split(marker)[0].strip()192 193 return sql194 195 def validate_sql_syntax(self, sql: str) -> bool:196 """197 Basic syntax validation for SQL.198 199 Args:200 sql: SQL query to validate201 202 Returns:203 True if syntax appears valid, False otherwise204 """205 if not sql or len(sql.strip()) == 0:206 return False207 208 sql_upper = sql.upper().strip()209 210 # Check if it starts with a valid SQL keyword211 valid_starts = ['SELECT', 'INSERT', 'UPDATE', 'DELETE', 'WITH']212 starts_valid = any(sql_upper.startswith(keyword) for keyword in valid_starts)213 214 if not starts_valid:215 return False216 217 # Check for balanced parentheses218 if sql.count('(') != sql.count(')'):219 logger.warning("Unbalanced parentheses in SQL")220 return False221 222 # Check for balanced quotes223 single_quotes = sql.count("'")224 if single_quotes % 2 != 0:225 logger.warning("Unbalanced single quotes in SQL")226 return False227 228 return True229 230 231# Global SQL generator agent instance232sql_generator = SQLGeneratorAgent()233 