Spaces:
Sleeping
Sleeping
| import os | |
| import sqlite3 | |
| import re | |
| import json | |
| DB_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "scratch") | |
| def get_db_path(user_id: str, suffix: str = "") -> str: | |
| """Returns the path to the user's SQLite session database.""" | |
| # Ensure scratch dir exists | |
| os.makedirs(DB_DIR, exist_ok=True) | |
| clean_id = "".join(c for c in user_id if c.isalnum() or c in ("-", "_")) | |
| return os.path.join(DB_DIR, f"db_{clean_id}{suffix}.sqlite") | |
| def is_query_safe(sql: str) -> tuple[bool, str]: | |
| """ | |
| Checks if a query is safe to execute. | |
| Blocks command chaining, ATTACH, PRAGMA, and database administration commands. | |
| """ | |
| cleaned = sql.strip().upper() | |
| # 1. Block command chaining (multiple queries separated by semicolon) | |
| # Strip trailing semicolons first, then check if any semicolon remains | |
| temp = cleaned.rstrip(';') | |
| if ';' in temp: | |
| return False, "Query chaining (using ';') is disabled for security." | |
| # 2. Block file attachment and direct system configurations | |
| blocked_keywords = [ | |
| r"\bATTACH\b", r"\bDETACH\b", r"\bPRAGMA\b", r"\bLOAD_EXTENSION\b", | |
| r"\bSHUTDOWN\b", r"\bGRANT\b", r"\bREVOKE\b" | |
| ] | |
| for pattern in blocked_keywords: | |
| if re.search(pattern, cleaned): | |
| return False, f"SQL command blocked for security: contains restricted keyword." | |
| return True, "" | |
| def get_db_schema(conn: sqlite3.Connection) -> dict: | |
| """Extracts column definitions for all tables in the database.""" | |
| cursor = conn.cursor() | |
| cursor.execute("SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%';") | |
| tables = [row[0] for row in cursor.fetchall()] | |
| schema = {} | |
| for table in tables: | |
| cursor.execute(f"PRAGMA table_info({table});") | |
| # PRAGMA returns: (cid, name, type, notnull, dflt_value, pk) | |
| columns = [{"name": r[1], "type": r[2]} for r in cursor.fetchall()] | |
| schema[table] = columns | |
| return schema | |
| def init_user_db(user_id: str, challenge: dict) -> dict: | |
| """ | |
| Initializes a fresh session database for the user with DDL and mock data. | |
| """ | |
| db_path = get_db_path(user_id) | |
| # Remove existing db if it exists | |
| if os.path.exists(db_path): | |
| try: | |
| os.remove(db_path) | |
| except Exception as e: | |
| return {"success": False, "error": f"Failed to reset database session: {str(e)}"} | |
| conn = None | |
| try: | |
| conn = sqlite3.connect(db_path) | |
| cursor = conn.cursor() | |
| # Execute DDL | |
| ddl = challenge.get("ddl", "") | |
| # Split DDL by semicolon to run statements sequentially | |
| for statement in ddl.split(';'): | |
| stmt = statement.strip() | |
| if stmt: | |
| cursor.execute(stmt) | |
| # Execute Mock Data INSERTS | |
| inserts = challenge.get("inserts", "") | |
| for statement in inserts.split(';'): | |
| stmt = statement.strip() | |
| if stmt: | |
| cursor.execute(stmt) | |
| conn.commit() | |
| # Extract schema for frontend visualization | |
| schema = get_db_schema(conn) | |
| return { | |
| "success": True, | |
| "schema": schema, | |
| "message": "Database initialized successfully." | |
| } | |
| except Exception as e: | |
| return {"success": False, "error": f"Database initialization failed: {str(e)}"} | |
| finally: | |
| if conn: | |
| conn.close() | |
| def run_user_query(user_id: str, query: str, challenge: dict) -> dict: | |
| """ | |
| Executes a query against the user's session database and validates results. | |
| """ | |
| # 1. Input Safety Validation | |
| is_safe, err_msg = is_query_safe(query) | |
| if not is_safe: | |
| return {"success": False, "error": err_msg} | |
| db_path = get_db_path(user_id) | |
| if not os.path.exists(db_path): | |
| return {"success": False, "error": "Database session not initialized. Type 'db init' first."} | |
| conn = None | |
| try: | |
| conn = sqlite3.connect(db_path) | |
| # Enable column-name dictionary rows | |
| conn.row_factory = sqlite3.Row | |
| cursor = conn.cursor() | |
| # Run query with execution timeout (sqlite3 doesn't have a direct query timeout in execute, | |
| # but we can set busy_timeout, or just rely on local speed. Since it is local SQLite with small mock data, | |
| # execution time is sub-millisecond unless there is an infinite loop CTE. | |
| # SQLite detects circular CTEs, but we will wrap execute in a try block) | |
| cursor.execute(query) | |
| # Determine if statement returns rows | |
| is_select = challenge.get("challenge_type", "SELECT").upper() == "SELECT" | |
| rows = [] | |
| columns = [] | |
| rows_affected = cursor.rowcount | |
| if cursor.description: | |
| columns = [col[0] for col in cursor.description] | |
| db_rows = cursor.fetchall() | |
| # Convert SQLite Row objects to list of dicts | |
| rows = [dict(r) for r in db_rows] | |
| conn.commit() | |
| # 2. Validation Engine | |
| validation_success = False | |
| challenge_type = challenge.get("challenge_type", "SELECT").upper() | |
| if challenge_type == "SELECT": | |
| validation_success = validate_select_query(user_id, query, challenge.get("validation_query", "")) | |
| elif challenge_type == "DML": | |
| # For UPDATE/DELETE, we check if the user's table states match the target table states | |
| validation_success = validate_dml_query(user_id, query, challenge) | |
| return { | |
| "success": True, | |
| "columns": columns, | |
| "rows": rows, | |
| "rows_affected": rows_affected if rows_affected >= 0 else 0, | |
| "is_correct": validation_success, | |
| "schema": get_db_schema(conn) | |
| } | |
| except Exception as e: | |
| if conn: | |
| conn.rollback() | |
| return {"success": False, "error": f"SQL execution error: {str(e)}"} | |
| finally: | |
| if conn: | |
| conn.close() | |
| def validate_select_query(user_id: str, user_query: str, golden_query: str) -> bool: | |
| """ | |
| Validates a SELECT query by running both user query and golden query | |
| and comparing output sets. | |
| """ | |
| db_path = get_db_path(user_id) | |
| conn = None | |
| try: | |
| conn = sqlite3.connect(db_path) | |
| cursor = conn.cursor() | |
| # Run user query | |
| cursor.execute(user_query) | |
| user_res = cursor.fetchall() | |
| # Run golden query | |
| cursor.execute(golden_query) | |
| golden_res = cursor.fetchall() | |
| # Compare row sets (ignoring row order for general checking, unless they differ) | |
| # We check set equality of row tuples | |
| return set(user_res) == set(golden_res) | |
| except Exception: | |
| return False | |
| finally: | |
| if conn: | |
| conn.close() | |
| def validate_dml_query(user_id: str, user_query: str, challenge: dict) -> bool: | |
| """ | |
| Validates UPDATE/DELETE challenges by comparing database states. | |
| We initialize a reference database, run the golden DML query on it, | |
| and verify all tables in both databases are identical. | |
| """ | |
| db_user_path = get_db_path(user_id) | |
| db_ref_path = get_db_path(user_id, "_ref") | |
| # 1. Initialize reference database | |
| if os.path.exists(db_ref_path): | |
| os.remove(db_ref_path) | |
| conn_ref = None | |
| conn_user = None | |
| try: | |
| # Spin up reference DB matching original state | |
| conn_ref = sqlite3.connect(db_ref_path) | |
| cursor_ref = conn_ref.cursor() | |
| # Setup tables and inserts | |
| for stmt in challenge.get("ddl", "").split(';'): | |
| if stmt.strip(): | |
| cursor_ref.execute(stmt) | |
| for stmt in challenge.get("inserts", "").split(';'): | |
| if stmt.strip(): | |
| cursor_ref.execute(stmt) | |
| # Run golden validation DML query on reference DB | |
| cursor_ref.execute(challenge.get("validation_query", "")) | |
| conn_ref.commit() | |
| # 2. Fetch all tables | |
| cursor_ref.execute("SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%';") | |
| tables = [row[0] for row in cursor_ref.fetchall()] | |
| # Connect to user's updated database | |
| conn_user = sqlite3.connect(db_user_path) | |
| cursor_user = conn_user.cursor() | |
| # Compare contents of every table | |
| for table in tables: | |
| # Check user table exists | |
| cursor_user.execute(f"SELECT name FROM sqlite_master WHERE type='table' AND name='{table}';") | |
| if not cursor_user.fetchone(): | |
| return False | |
| # Get table content from reference DB | |
| cursor_ref.execute(f"SELECT * FROM {table};") | |
| ref_rows = cursor_ref.fetchall() | |
| # Get table content from user DB | |
| cursor_user.execute(f"SELECT * FROM {table};") | |
| user_rows = cursor_user.fetchall() | |
| if set(ref_rows) != set(user_rows): | |
| return False | |
| return True | |
| except Exception: | |
| return False | |
| finally: | |
| if conn_ref: | |
| conn_ref.close() | |
| if conn_user: | |
| conn_user.close() | |
| # Clean up reference DB file | |
| if os.path.exists(db_ref_path): | |
| try: | |
| os.remove(db_ref_path) | |
| except Exception: | |
| pass | |