joelgilbert/NL2SQL
0
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 