# coding=utf-8 import pytest from common.sql_guard import ( SqlGuardError, extract_sql_literals, find_forbidden_statements, validate_pyspark_code, ) # --- extract_sql_literals --- def test_extract_picks_up_spark_sql_double_quoted(): code = 'spark.sql("SELECT * FROM users")' out = extract_sql_literals(code) assert out == ["SELECT * FROM users"] def test_extract_picks_up_spark_sql_single_quoted(): code = "spark.sql('SELECT 1')" out = extract_sql_literals(code) assert out == ["SELECT 1"] def test_extract_ignores_english_sentence_starting_with_select(): """False positive tolerated for safety: 'select the best option' would be picked up. That's fine — it's a non-SQL literal that gets classified, but won't contain any forbidden keyword.""" code = 'log("select the best option")' out = extract_sql_literals(code) assert out == ["select the best option"] # And the validator returns no offenses for it (no SQL verbs in body) assert find_forbidden_statements(out[0]) == [] def test_extract_ignores_non_sql_strings(): code = 'msg = "hello world"\nlog(msg)' assert extract_sql_literals(code) == [] def test_extract_handles_fstring_prefix(): code = 'spark.sql(f"SELECT * FROM {table}")' out = extract_sql_literals(code) assert out == ["SELECT * FROM {table}"] def test_extract_collects_multiple_literals(): code = ''' df1 = spark.sql("SELECT * FROM t1") df2 = spark.sql("SELECT * FROM t2") ''' out = extract_sql_literals(code) assert out == ["SELECT * FROM t1", "SELECT * FROM t2"] # --- find_forbidden_statements --- def test_select_only_is_allowed(): assert find_forbidden_statements("SELECT * FROM users") == [] def test_select_with_where_is_allowed(): assert find_forbidden_statements( "SELECT id, name FROM users WHERE age > 18 ORDER BY id" ) == [] def test_insert_is_allowed(): assert find_forbidden_statements( "INSERT INTO events SELECT * FROM raw_events" ) == [] def test_drop_table_is_forbidden(): assert find_forbidden_statements("DROP TABLE users") == ["DROP"] def test_drop_database_is_forbidden(): assert find_forbidden_statements("DROP DATABASE prod") == ["DROP"] def test_delete_is_forbidden(): assert find_forbidden_statements("DELETE FROM users WHERE id = 1") == ["DELETE"] def test_update_is_forbidden(): assert find_forbidden_statements( "UPDATE users SET name = 'x' WHERE id = 1" ) == ["UPDATE"] def test_truncate_is_forbidden(): assert find_forbidden_statements("TRUNCATE TABLE events") == ["TRUNCATE"] def test_alter_is_forbidden(): assert find_forbidden_statements("ALTER TABLE users ADD COLUMN x INT") == ["ALTER"] def test_create_table_is_forbidden(): assert find_forbidden_statements( "CREATE TABLE foo (id INT, name STRING)" ) == ["CREATE"] def test_with_cte_select_is_allowed(): sql = "WITH active AS (SELECT * FROM users WHERE active) SELECT * FROM active" assert find_forbidden_statements(sql) == [] def test_with_cte_drop_is_forbidden(): """With CTE body containing a DROP. sqlparse may split this into two statements or keep it as one — either way, the DROP must be flagged.""" sql = "WITH x AS (DROP TABLE y) SELECT * FROM x" out = find_forbidden_statements(sql) # The DROP must be reported (label may be 'DROP' or 'WITH ... DROP' # depending on how sqlparse tokenizes; both prove it's caught). assert any("DROP" in o for o in out) assert any(o in {"DROP", "WITH ... DROP"} for o in out) def test_multi_statement_selects_all_allowed(): sql = "SELECT 1; SELECT 2; INSERT INTO t VALUES (1)" assert find_forbidden_statements(sql) == [] def test_multi_statement_one_bad_is_caught(): sql = "SELECT 1; DROP TABLE users; SELECT 2" assert find_forbidden_statements(sql) == ["DROP"] def test_comments_and_whitespace_dont_confuse_parser(): sql = """ -- this is a comment /* multi-line comment */ SELECT * FROM users """ assert find_forbidden_statements(sql) == [] # --- validate_pyspark_code (end-to-end through the Python string extractor) --- def test_validate_clean_pyspark_code(): code = ''' from pyspark.sql import SparkSession spark = SparkSession.builder.getOrCreate() df = spark.sql("SELECT * FROM users") df.show() ''' assert validate_pyspark_code(code) == [] def test_validate_rejects_drop_in_pyspark_code(): code = 'spark.sql("DROP TABLE users")' out = validate_pyspark_code(code) assert "DROP" in out def test_validate_rejects_delete_with_where(): code = 'spark.sql("DELETE FROM events WHERE id = 1")' out = validate_pyspark_code(code) assert "DELETE" in out def test_validate_rejects_update(): code = 'spark.sql("UPDATE users SET x = 1")' out = validate_pyspark_code(code) assert "UPDATE" in out def test_validate_collects_multiple_offenses(): code = ''' spark.sql("DROP TABLE a") spark.sql("DELETE FROM b") ''' out = validate_pyspark_code(code) assert "DROP" in out assert "DELETE" in out def test_validate_passes_when_no_sql_present(): code = ''' from pyspark.sql import SparkSession spark = SparkSession.builder.getOrCreate() df = spark.range(0, 100) df.show() ''' assert validate_pyspark_code(code) == [] def test_validate_passes_with_dataframe_dsl(): """DataFrame operations (filter, select, groupBy, agg) are not raw SQL.""" code = ''' df = spark.read.parquet("/data/foo") filtered = df.filter(df.age > 18).select("id", "name").groupBy("name").count() filtered.show() ''' assert validate_pyspark_code(code) == []