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/.
This commit is contained in:
@@ -4,17 +4,40 @@
|
||||
@Author :tao.chen
|
||||
"""
|
||||
from common.logging import logger
|
||||
from common.sql_guard import validate_pyspark_code
|
||||
from spark_executor.core.job_writer import write_job_file
|
||||
|
||||
|
||||
class SqlGuardViolation(ValueError):
|
||||
"""Raised when generate_job_file receives code that violates the SQL
|
||||
safety policy. -> HTTP 400 via the FastAPI ValueError handler.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
def generate_job_file(code: str) -> dict[str, str]:
|
||||
"""Write a PySpark code string to disk; return its absolute path.
|
||||
|
||||
Use the returned path as the `script_path` argument of prepare_submit_job.
|
||||
Validates the code against the SQL safety policy (SELECT/INSERT only)
|
||||
BEFORE writing, so the agent gets immediate feedback rather than
|
||||
learning at prepare_submit_job time. Use the returned path as the
|
||||
`script_path` argument of prepare_submit_job.
|
||||
|
||||
The output directory is controlled by the SPARK_EXECUTOR_JOBS_DIR env var
|
||||
(default: ./data/jobs/).
|
||||
"""
|
||||
logger.debug(f"generate_job_file enter code_bytes={len(code)}")
|
||||
offenses = validate_pyspark_code(code)
|
||||
if offenses:
|
||||
logger.warning(
|
||||
f"generate_job_file rejected: SQL policy violation(s): {offenses}"
|
||||
)
|
||||
raise SqlGuardViolation(
|
||||
f"PySpark code violates SQL safety policy. "
|
||||
f"Only SELECT and INSERT statements are allowed. "
|
||||
f"Found forbidden statement(s): {offenses}. "
|
||||
f"Rewrite the code to use only SELECT/INSERT (or DataFrame DSL)."
|
||||
)
|
||||
path = write_job_file(code)
|
||||
logger.info(f"generate_job_file ok script_path={path}")
|
||||
return {"script_path": path}
|
||||
|
||||
@@ -9,6 +9,7 @@ import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from common.logging import logger
|
||||
from common.sql_guard import validate_pyspark_code
|
||||
from spark_executor.core.connection_store import store as conn_store
|
||||
from spark_executor.core.job_store import JobStore
|
||||
from spark_executor.core.log_parser import parse_spark_submit_output
|
||||
@@ -42,6 +43,19 @@ def _check_script_path(script_path: str) -> None:
|
||||
raise _script_path_error(script_path)
|
||||
|
||||
|
||||
def _sql_guard_error(offenses: list[str]) -> ValueError:
|
||||
"""Instructive 400 error when the script's SQL fails the safety policy."""
|
||||
return ValueError(
|
||||
f"Script SQL violates the safety policy. "
|
||||
f"Only SELECT and INSERT statements are allowed. "
|
||||
f"Found forbidden statement(s): {offenses}. "
|
||||
f"Edit the script and try again. (If the code was generated via "
|
||||
f"generate_job_file, the generator should have caught this — check "
|
||||
f"for f-string-based SQL injection where the static analysis can't "
|
||||
f"see the runtime value.)"
|
||||
)
|
||||
|
||||
|
||||
def _new_pending_id() -> str:
|
||||
return "p_" + secrets.token_hex(6)
|
||||
|
||||
@@ -64,6 +78,7 @@ def prepare_submit_job(
|
||||
# Order of checks matters for the error the agent sees:
|
||||
# 1. Unknown connection -> 404 (KeyError -> 404 in server.py)
|
||||
# 2. Missing script file -> 400 (ValueError -> 400)
|
||||
# 3. SQL policy violation -> 400 (ValueError -> 400)
|
||||
# Connection is checked first because it is a more fundamental problem
|
||||
# (the agent is asking about a cluster that doesn't exist), and the
|
||||
# agent shouldn't have to fix the script path only to learn the
|
||||
@@ -78,6 +93,18 @@ def prepare_submit_job(
|
||||
# a 400 + clear remediation than to let spark-submit fail later with
|
||||
# an opaque FileNotFoundError -> 500.
|
||||
_check_script_path(script_path)
|
||||
# SQL safety: re-validate the script even though generate_job_file
|
||||
# already guards its own output. Catches files written by other means
|
||||
# (host volume mounts, manual edits).
|
||||
with open(script_path, encoding="utf-8") as f:
|
||||
script_body = f.read()
|
||||
offenses = validate_pyspark_code(script_body)
|
||||
if offenses:
|
||||
logger.warning(
|
||||
f"prepare_submit_job rejected: SQL policy violation(s) in "
|
||||
f"{script_path}: {offenses}"
|
||||
)
|
||||
raise _sql_guard_error(offenses)
|
||||
|
||||
pending_id = _new_pending_id()
|
||||
pending = PendingSubmission(
|
||||
|
||||
Reference in New Issue
Block a user