Team Ai
Apppublic

joelgilbert/NL2SQL

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
query_executor.py287 linesDownload Raw Back to database
1"""2Safe SQL query execution with error handling and result formatting.3"""4 5import logging6import time7from typing import Dict, List, Optional, Any8import psycopg2.extras9from psycopg2 import Error as PsycopgError10 11from database.connection import db12 13logger = logging.getLogger(__name__)14 15 16class QueryExecutor:17    """Executes SQL queries safely with comprehensive error handling."""18    19    def __init__(self):20        """Initialize query executor."""21        self.max_result_rows = 1000022    23    def execute_readonly_query(self, sql: str) -> Dict[str, Any]:24        """25        Execute a SELECT query with read-only permissions.26        27        Args:28            sql: SQL query to execute29            30        Returns:31            Dictionary with:32                - success (bool): Whether execution succeeded33                - data (List[Dict]): Query results as list of dictionaries34                - error (str): Error message if failed35                - execution_time (float): Execution time in seconds36                - row_count (int): Number of rows returned37                - truncated (bool): Whether results were truncated38        """39        start_time = time.time()40        41        try:42            with db.get_readonly_connection() as conn:43                cursor = conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor)44                45                # Execute query46                cursor.execute(sql)47                48                # Fetch results with row limit49                rows = cursor.fetchmany(self.max_result_rows + 1)50                51                # Check if results were truncated52                truncated = len(rows) > self.max_result_rows53                if truncated:54                    rows = rows[:self.max_result_rows]55                56                # Convert to list of dicts57                data = [dict(row) for row in rows]58                59                cursor.close()60                execution_time = time.time() - start_time61                62                logger.info(f"Query executed successfully: {len(data)} rows in {execution_time:.3f}s")63                64                return {65                    "success": True,66                    "data": data,67                    "error": None,68                    "execution_time": execution_time,69                    "row_count": len(data),70                    "truncated": truncated71                }72                73        except PsycopgError as e:74            execution_time = time.time() - start_time75            error_msg = str(e).strip()76            logger.error(f"Query execution failed: {error_msg}")77            78            return {79                "success": False,80                "data": [],81                "error": error_msg,82                "execution_time": execution_time,83                "row_count": 0,84                "truncated": False85            }86        87        except Exception as e:88            execution_time = time.time() - start_time89            error_msg = f"Unexpected error: {str(e)}"90            logger.error(error_msg)91            92            return {93                "success": False,94                "data": [],95                "error": error_msg,96                "execution_time": execution_time,97                "row_count": 0,98                "truncated": False99            }100    101    def execute_dba_query(self, sql: str, approved: bool = False) -> Dict[str, Any]:102        """103        Execute a query with DBA permissions (INSERT, UPDATE, DELETE).104        105        Args:106            sql: SQL query to execute107            approved: Whether the query has been approved by a human108            109        Returns:110            Dictionary with execution results (same format as execute_readonly_query)111        """112        if not approved:113            return {114                "success": False,115                "data": [],116                "error": "Query must be approved before execution",117                "execution_time": 0.0,118                "row_count": 0,119                "truncated": False120            }121        122        start_time = time.time()123        124        try:125            with db.get_dba_connection() as conn:126                cursor = conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor)127                128                # Execute query129                cursor.execute(sql)130                131                # Get affected rows count132                row_count = cursor.rowcount133                134                # For SELECT queries in DBA mode, fetch results135                if sql.strip().upper().startswith('SELECT'):136                    rows = cursor.fetchmany(self.max_result_rows + 1)137                    truncated = len(rows) > self.max_result_rows138                    if truncated:139                        rows = rows[:self.max_result_rows]140                    data = [dict(row) for row in rows]141                else:142                    data = []143                    truncated = False144                145                cursor.close()146                execution_time = time.time() - start_time147                148                logger.info(f"DBA query executed: {row_count} rows affected in {execution_time:.3f}s")149                150                return {151                    "success": True,152                    "data": data,153                    "error": None,154                    "execution_time": execution_time,155                    "row_count": row_count,156                    "truncated": truncated157                }158                159        except PsycopgError as e:160            execution_time = time.time() - start_time161            error_msg = str(e).strip()162            logger.error(f"DBA query execution failed: {error_msg}")163            164            return {165                "success": False,166                "data": [],167                "error": error_msg,168                "execution_time": execution_time,169                "row_count": 0,170                "truncated": False171            }172        173        except Exception as e:174            execution_time = time.time() - start_time175            error_msg = f"Unexpected error: {str(e)}"176            logger.error(error_msg)177            178            return {179                "success": False,180                "data": [],181                "error": error_msg,182                "execution_time": execution_time,183                "row_count": 0,184                "truncated": False185            }186    187    def explain_query(self, sql: str) -> Dict[str, Any]:188        """189        Get PostgreSQL EXPLAIN output for query cost estimation.190        191        Args:192            sql: SQL query to analyze193            194        Returns:195            Dictionary with:196                - success (bool): Whether EXPLAIN succeeded197                - plan (List[str]): Query execution plan198                - error (str): Error message if failed199        """200        try:201            with db.get_readonly_connection() as conn:202                cursor = conn.cursor()203                204                # Get EXPLAIN output205                cursor.execute(f"EXPLAIN {sql}")206                plan = cursor.fetchall()207                208                cursor.close()209                210                # Format plan as list of strings211                plan_lines = [row[0] for row in plan]212                213                return {214                    "success": True,215                    "plan": plan_lines,216                    "error": None217                }218                219        except PsycopgError as e:220            error_msg = str(e).strip()221            logger.error(f"EXPLAIN failed: {error_msg}")222            223            return {224                "success": False,225                "plan": [],226                "error": error_msg227            }228        229        except Exception as e:230            error_msg = f"Unexpected error: {str(e)}"231            logger.error(error_msg)232            233            return {234                "success": False,235                "plan": [],236                "error": error_msg237            }238    239    def estimate_affected_rows(self, sql: str) -> int:240        """241        Estimate number of rows that would be affected by UPDATE/DELETE query.242        243        Args:244            sql: UPDATE or DELETE query245            246        Returns:247            Estimated row count (0 if estimation fails)248        """249        try:250            # Convert UPDATE/DELETE to SELECT COUNT(*)251            sql_upper = sql.strip().upper()252            253            if sql_upper.startswith('UPDATE'):254                # Extract table and WHERE clause255                parts = sql.split('WHERE', 1)256                table_part = parts[0].replace('UPDATE', 'SELECT COUNT(*) FROM', 1)257                table_part = table_part.split('SET')[0].strip()258                259                if len(parts) > 1:260                    count_query = f"{table_part} WHERE {parts[1]}"261                else:262                    count_query = table_part263            264            elif sql_upper.startswith('DELETE'):265                # Replace DELETE FROM with SELECT COUNT(*) FROM266                count_query = sql.replace('DELETE FROM', 'SELECT COUNT(*) FROM', 1)267            268            else:269                return 0270            271            # Execute count query272            with db.get_readonly_connection() as conn:273                cursor = conn.cursor()274                cursor.execute(count_query)275                count = cursor.fetchone()[0]276                cursor.close()277                278                return count279                280        except Exception as e:281            logger.warning(f"Failed to estimate affected rows: {e}")282            return 0283 284 285# Global query executor instance286query_executor = QueryExecutor()287