Team Ai
Apppublic

joelgilbert/NL2SQL

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
error_handler.py269 linesDownload Raw Back to utils
1"""2SQL error classification and handling.3"""4 5import logging6import re7from typing import Dict, Tuple8from psycopg2 import Error as PsycopgError9 10logger = logging.getLogger(__name__)11 12 13class ErrorHandler:14    """Classifies and handles SQL execution errors."""15    16    ERROR_PATTERNS = {17        "column_name": [18            r'column "(\w+)" does not exist',19            r'no such column: (\w+)',20            r'unknown column',21        ],22        "table_name": [23            r'relation "(\w+)" does not exist',24            r'no such table: (\w+)',25            r'table.*not found',26        ],27        "syntax": [28            r'syntax error at or near "(\w+)"',29            r'syntax error',30            r'invalid syntax',31        ],32        "type_mismatch": [33            r'type mismatch',34            r'invalid input syntax for type',35            r'operator does not exist',36            r'cannot cast',37        ],38        "permission": [39            r'permission denied',40            r'must be owner of',41            r'access denied',42        ],43        "timeout": [44            r'timeout',45            r'canceling statement due to statement timeout',46        ]47    }48    49    def classify_error(self, error_message: str) -> str:50        """51        Classify error type based on error message.52        53        Args:54            error_message: Error message from database55            56        Returns:57            Error type: column_name, table_name, syntax, type_mismatch, 58                       permission, timeout, or other59        """60        error_lower = error_message.lower()61        62        for error_type, patterns in self.ERROR_PATTERNS.items():63            for pattern in patterns:64                if re.search(pattern, error_lower):65                    logger.info(f"Classified error as: {error_type}")66                    return error_type67        68        logger.info("Error classified as: other")69        return "other"70    71    def generate_correction_context(72        self,73        error_type: str,74        error_message: str,75        schema: str76    ) -> str:77        """78        Generate specific guidance for error correction.79        80        Args:81            error_type: Classified error type82            error_message: Original error message83            schema: Database schema information84            85        Returns:86            Guidance text for SQL correction87        """88        guidance_templates = {89            "column_name": f"""The query references a column that doesn't exist.90 91Error: {error_message}92 93Please review the schema below and use only the columns that actually exist.94Ensure column names are spelled correctly and match the case if necessary.95 96Available Schema:97{schema}""",98            99            "table_name": f"""The query references a table that doesn't exist.100 101Error: {error_message}102 103Please review the schema below and use only the tables that actually exist.104Ensure table names are spelled correctly.105 106Available Schema:107{schema}""",108            109            "syntax": f"""There's a SQL syntax error in the query.110 111Error: {error_message}112 113Common syntax issues:114- Missing or extra commas115- Unmatched parentheses116- Incorrect keyword order117- Missing required clauses118 119Please fix the syntax according to PostgreSQL standards.120 121Schema for reference:122{schema}""",123            124            "type_mismatch": f"""There's a data type mismatch in the query.125 126Error: {error_message}127 128Common type issues:129- Comparing incompatible types (e.g., text vs integer)130- Missing type casts (use ::type or CAST(column AS type))131- Incorrect aggregate function usage132 133Please ensure proper type casting and compatible comparisons.134 135Schema for reference:136{schema}""",137            138            "permission": f"""The query requires elevated permissions.139 140Error: {error_message}141 142The query is attempting an operation that requires DBA privileges.143If in read-only mode, revise the query to use only SELECT operations.144 145Schema for reference:146{schema}""",147            148            "timeout": f"""The query exceeded the timeout limit.149 150Error: {error_message}151 152The query is taking too long to execute. Consider:153- Adding WHERE clause to limit rows154- Simplifying complex joins or subqueries155- Using indexes effectively156 157Schema for reference:158{schema}""",159            160            "other": f"""An error occurred while executing the query.161 162Error: {error_message}163 164Please review the error message and schema to correct the query.165 166Schema for reference:167{schema}"""168        }169        170        return guidance_templates.get(error_type, guidance_templates["other"])171    172    def parse_postgres_error(self, error: Exception) -> Dict[str, str]:173        """174        Extract structured information from psycopg2 errors.175        176        Args:177            error: Exception from psycopg2178            179        Returns:180            Dictionary with error_code, error_type, detail, hint181        """182        error_dict = {183            "error_code": "",184            "error_type": "unknown",185            "detail": str(error),186            "hint": ""187        }188        189        if isinstance(error, PsycopgError):190            if hasattr(error, 'pgcode'):191                error_dict["error_code"] = error.pgcode or ""192            193            if hasattr(error, 'pgerror'):194                error_dict["detail"] = error.pgerror or str(error)195            196            # Extract hint if available197            error_str = str(error)198            hint_match = re.search(r'HINT: (.*?)(\n|$)', error_str)199            if hint_match:200                error_dict["hint"] = hint_match.group(1)201        202        # Classify error type203        error_dict["error_type"] = self.classify_error(error_dict["detail"])204        205        return error_dict206    207    def should_retry(self, error_type: str, attempt: int, max_attempts: int = 3) -> bool:208        """209        Determine if error is retryable and within attempt limits.210        211        Args:212            error_type: Classified error type213            attempt: Current attempt number (1-indexed)214            max_attempts: Maximum number of attempts215            216        Returns:217            True if should retry, False otherwise218        """219        # Don't retry if max attempts reached220        if attempt >= max_attempts:221            return False222        223        # Retryable error types224        retryable_types = ["column_name", "table_name", "syntax", "type_mismatch"]225        226        # Don't retry permission or timeout errors227        non_retryable_types = ["permission", "timeout"]228        229        if error_type in non_retryable_types:230            return False231        232        if error_type in retryable_types:233            return True234        235        # For "other" errors, allow one retry236        return attempt < 2237    238    def format_error_for_user(self, error_type: str, error_message: str) -> str:239        """240        Format error message for user-friendly display.241        242        Args:243            error_type: Classified error type244            error_message: Original error message245            246        Returns:247            User-friendly error message248        """249        user_messages = {250            "column_name": "❌ The query references a column that doesn't exist in the database. Please check the column names.",251            "table_name": "❌ The query references a table that doesn't exist in the database. Please check the table names.",252            "syntax": "❌ There's a syntax error in the generated SQL query. This is usually due to incorrect SQL formatting.",253            "type_mismatch": "❌ There's a data type mismatch in the query. The query is trying to compare incompatible data types.",254            "permission": "❌ This operation requires DBA privileges. Please switch to DBA mode or revise your request.",255            "timeout": "❌ The query took too long to execute and was cancelled. Try narrowing down your request.",256            "other": "❌ An error occurred while executing the query."257        }258        259        base_message = user_messages.get(error_type, user_messages["other"])260        261        # Add truncated error detail262        error_snippet = error_message[:200] if len(error_message) > 200 else error_message263        264        return f"{base_message}\n\n**Technical Details:** {error_snippet}"265 266 267# Global error handler instance268error_handler = ErrorHandler()269