Team Ai
Apppublic

joelgilbert/NL2SQL

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
sql_generator.py233 linesDownload Raw Back to agents
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