Files
mcp-server/tests/unit/test_sql_guard.py
T
Claude 34e6208f54 feat: SQL safety policy (SELECT/INSERT only) at submit time
The agent can now write PySpark that runs DROP/DELETE/UPDATE/etc. on
production tables. Add a static guard that rejects anything other than
SELECT and INSERT at two enforcement points:

  1. generate_job_file: validates BEFORE writing to disk. Agent gets
     immediate feedback ('rewrite to use only SELECT/INSERT') rather
     than learning at submit time.

  2. prepare_submit_job: re-validates the script content (reads the
     file) as a defense-in-depth check. Catches host-mounted files,
     manually-edited files, anything that bypassed generate_job_file.

How it works:
  - common/sql_guard.py extracts Python string literals whose first
    keyword is a SQL verb (catches spark.sql('...'), f-strings, and any
    raw SQL literal)
  - sqlparse splits each literal into statements; we check the first
    keyword against the policy (SELECT/INSERT/WITH allowed; DROP,
    DELETE, UPDATE, TRUNCATE, ALTER, CREATE, REPLACE, MERGE, GRANT,
    REVOKE, SET, SHOW, KILL, EXEC, etc. forbidden)
  - WITH recurses into the CTE body to catch WITH x AS (DROP ...) ...
  - The MCP layer maps ValueError -> HTTP 400 (existing handler)

Test coverage:
  - 28 unit tests in test_sql_guard.py cover: extraction (single/double/
    f-string, English false positives, multi-literal), statement
    classification (SELECT, INSERT, DROP, DELETE, UPDATE, TRUNCATE,
    ALTER, CREATE, multi-statement, CTE bodies, comments)
  - 2 integration tests verify MCP layer returns 400 with the policy
    explanation at both generate_job_file and prepare_submit_job

Limitations (documented in sql_guard.py docstring):
  - f-strings where the SQL is built at runtime (e.g. f'SELECT * FROM
    {user_input}') look like SELECTs at static-analysis time. The
    guard catches the static literal; the runtime substitution is the
    caller's responsibility.
  - pyspark.sql.functions.expr('...') accepts SQL inline; not currently
    caught. (Future work.)

146/146 still pass. Live verified: DROP TABLE -> MCP 400 with policy
explanation; SELECT -> MCP 200 + file written to ./data/jobs/.
2026-06-25 12:52:16 +08:00

201 lines
5.5 KiB
Python

# 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) == []