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/.
257 lines
8.9 KiB
Python
257 lines
8.9 KiB
Python
# coding=utf-8
|
|
import os
|
|
from pathlib import Path
|
|
|
|
import pytest
|
|
from fastapi.testclient import TestClient
|
|
|
|
from spark_executor.core import connection_store, pending_store
|
|
from spark_executor.core.connection_store import ConnectionStore
|
|
from spark_executor.core.pending_store import PendingStore
|
|
from spark_executor.server import app
|
|
from spark_executor.tools import connections, submit
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _fresh_data(tmp_path: Path, monkeypatch):
|
|
"""Reset both stores to a fresh tmp dir and rebind the singletons that
|
|
the tool modules captured at import time."""
|
|
monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path))
|
|
monkeypatch.setattr(connection_store, "store", ConnectionStore())
|
|
monkeypatch.setattr(pending_store, "DEFAULT_DATA_DIR", str(tmp_path))
|
|
monkeypatch.setattr(pending_store, "store", PendingStore())
|
|
# Rebind imports inside the tool modules (captured at import time)
|
|
connections.store = connection_store.store
|
|
submit.conn_store = connection_store.store
|
|
submit.pending_store = pending_store.store
|
|
|
|
|
|
def test_health_still_present():
|
|
c = TestClient(app)
|
|
r = c.get("/health")
|
|
assert r.status_code == 200
|
|
assert r.json() == {"status": "ok"}
|
|
|
|
|
|
def test_thirteen_tool_routes_registered():
|
|
paths = {r.path for r in app.routes}
|
|
for path in (
|
|
# pending-submission flow (5)
|
|
"/prepare_submit_job",
|
|
"/confirm_submit_job",
|
|
"/list_pending_jobs",
|
|
"/get_pending_job",
|
|
"/cancel_pending_job",
|
|
# job lifecycle (3)
|
|
"/get_job_status",
|
|
"/get_job_logs",
|
|
"/kill_job",
|
|
# connection management (4)
|
|
"/save_connection",
|
|
"/list_connections",
|
|
"/get_connection",
|
|
"/delete_connection",
|
|
# LLM-driven PySpark generation (1) — Stage 2
|
|
"/generate_job_file",
|
|
):
|
|
assert path in paths, f"missing MCP tool route: {path}"
|
|
|
|
|
|
# --- End-to-end body-based calls (the gap the route-registration test missed) ---
|
|
|
|
def test_save_connection_accepts_dict_spark_conf_in_body():
|
|
"""The original query-param signature returned 422 for spark_conf dicts;
|
|
body models make tools/call roundtrip cleanly."""
|
|
c = TestClient(app)
|
|
r = c.post(
|
|
"/save_connection",
|
|
json={
|
|
"name": "prod",
|
|
"master": "yarn",
|
|
"deploy_mode": "cluster",
|
|
"spark_conf": {"spark.sql.shuffle.partitions": "200"},
|
|
},
|
|
)
|
|
assert r.status_code == 200, r.text
|
|
assert r.json() == {"name": "prod", "status": "SAVED"}
|
|
|
|
|
|
def test_prepare_submit_job_works_via_body(tmp_path):
|
|
c = TestClient(app)
|
|
c.post("/save_connection", json={"name": "prod", "master": "yarn"})
|
|
# The script must exist (Stage 2 cleanup) — write a real file first.
|
|
script = tmp_path / "demo.py"
|
|
script.write_text("print('hi')\n")
|
|
r = c.post(
|
|
"/prepare_submit_job",
|
|
json={"connection": "prod", "script_path": str(script), "queue": "research"},
|
|
)
|
|
assert r.status_code == 200, r.text
|
|
body = r.json()
|
|
assert body["status"] == "PENDING"
|
|
assert body["pending_id"].startswith("p_")
|
|
assert body["parameters"]["queue"] == "research"
|
|
assert body["parameters"]["master"] == "yarn" # snapshotted from connection
|
|
|
|
|
|
def test_prepare_rejects_nonexistent_script_path_with_400():
|
|
"""MCP clients should see a clean 400 with a remediation hint, not a 500."""
|
|
c = TestClient(app)
|
|
c.post("/save_connection", json={"name": "prod", "master": "yarn"})
|
|
r = c.post(
|
|
"/prepare_submit_job",
|
|
json={"connection": "prod", "script_path": "/nope/does_not_exist.py"},
|
|
)
|
|
assert r.status_code == 400
|
|
detail = r.json()["detail"]
|
|
# The message must guide the agent to the right next step
|
|
assert "does not exist" in detail
|
|
assert "generate_job_file" in detail
|
|
|
|
|
|
def test_prepare_rejects_sql_policy_violation_with_400(tmp_path):
|
|
"""SQL guard: even if the file exists, prepare_submit_job must reject
|
|
code containing DROP/DELETE/etc. (defense in depth — generate_job_file
|
|
also enforces this, but a host-mounted file might bypass that)."""
|
|
c = TestClient(app)
|
|
c.post("/save_connection", json={"name": "prod", "master": "yarn"})
|
|
bad_script = tmp_path / "evil.py"
|
|
bad_script.write_text('spark.sql("DROP TABLE users")\n')
|
|
r = c.post(
|
|
"/prepare_submit_job",
|
|
json={"connection": "prod", "script_path": str(bad_script)},
|
|
)
|
|
assert r.status_code == 400
|
|
detail = r.json()["detail"]
|
|
assert "DROP" in detail
|
|
assert "SELECT" in detail and "INSERT" in detail # the policy is explained
|
|
|
|
|
|
def test_generate_rejects_sql_policy_violation_with_400(tmp_path):
|
|
"""Same policy at generate_job_file — agent gets immediate feedback
|
|
before the file is even written to disk."""
|
|
from common import config
|
|
config.settings.jobs_dir = str(tmp_path / "jobs")
|
|
c = TestClient(app)
|
|
r = c.post(
|
|
"/generate_job_file",
|
|
json={"code": 'spark.sql("DELETE FROM events")\n'},
|
|
)
|
|
assert r.status_code == 400
|
|
detail = r.json()["detail"]
|
|
assert "DELETE" in detail
|
|
# And the file should NOT have been written
|
|
assert list(tmp_path.glob("*.py")) == []
|
|
|
|
|
|
def test_list_and_get_pending_job_roundtrip_via_body(tmp_path):
|
|
c = TestClient(app)
|
|
c.post("/save_connection", json={"name": "prod", "master": "yarn"})
|
|
script = tmp_path / "a.py"
|
|
script.write_text("print('hi')\n")
|
|
prep = c.post(
|
|
"/prepare_submit_job",
|
|
json={"connection": "prod", "script_path": str(script)},
|
|
).json()
|
|
pid = prep["pending_id"]
|
|
|
|
listed = c.post("/list_pending_jobs", json={}).json()
|
|
assert any(p["pending_id"] == pid for p in listed)
|
|
|
|
got = c.post("/get_pending_job", json={"pending_id": pid}).json()
|
|
assert got["script_path"] == str(script)
|
|
assert got["status"] == "PENDING"
|
|
|
|
|
|
def test_list_connections_works_with_empty_body():
|
|
c = TestClient(app)
|
|
r = c.post("/list_connections", json={})
|
|
assert r.status_code == 200
|
|
assert r.json() == []
|
|
|
|
|
|
# --- Stage 2: generate_job_file ---
|
|
|
|
def test_generate_job_file_accepts_code_string_in_body(tmp_path, monkeypatch):
|
|
"""End-to-end: a Pydantic body model lets tools/call pass a code string
|
|
that FastAPI would reject if it were a query parameter (length limits)."""
|
|
monkeypatch.chdir(tmp_path)
|
|
c = TestClient(app)
|
|
code = "print('from MCP integration test')\n" * 100 # > FastAPI query limit
|
|
r = c.post("/generate_job_file", json={"code": code})
|
|
assert r.status_code == 200, r.text
|
|
p = r.json()["script_path"]
|
|
assert os.path.isfile(p)
|
|
with open(p) as f:
|
|
assert f.read() == code
|
|
|
|
|
|
# --- Exception handlers: KeyError -> 404, ValueError -> 400 ---
|
|
|
|
def test_unknown_job_id_returns_404():
|
|
c = TestClient(app)
|
|
r = c.post("/get_job_status", json={"job_id": "missing"})
|
|
assert r.status_code == 404
|
|
assert "missing" in r.json()["detail"]
|
|
|
|
|
|
def test_unknown_pending_id_returns_404():
|
|
c = TestClient(app)
|
|
r = c.post("/get_pending_job", json={"pending_id": "p_nope"})
|
|
assert r.status_code == 404
|
|
assert "p_nope" in r.json()["detail"]
|
|
|
|
|
|
def test_unknown_connection_name_returns_404():
|
|
c = TestClient(app)
|
|
r = c.post("/get_connection", json={"name": "nope"})
|
|
assert r.status_code == 404
|
|
assert "nope" in r.json()["detail"]
|
|
|
|
|
|
def test_unknown_connection_in_prepare_returns_404():
|
|
c = TestClient(app)
|
|
r = c.post(
|
|
"/prepare_submit_job",
|
|
json={"connection": "nope", "script_path": "/tmp/x.py"},
|
|
)
|
|
assert r.status_code == 404
|
|
assert "nope" in r.json()["detail"]
|
|
|
|
|
|
def test_confirm_non_pending_returns_400(tmp_path):
|
|
"""A CANCELLED pending should refuse confirm; FastAPI should surface ValueError as 400."""
|
|
c = TestClient(app)
|
|
c.post("/save_connection", json={"name": "prod", "master": "yarn"})
|
|
script = tmp_path / "x.py"
|
|
script.write_text("print('hi')\n")
|
|
prep = c.post(
|
|
"/prepare_submit_job",
|
|
json={"connection": "prod", "script_path": str(script)},
|
|
).json()
|
|
pid = prep["pending_id"]
|
|
c.post("/cancel_pending_job", json={"pending_id": pid})
|
|
r = c.post("/confirm_submit_job", json={"pending_id": pid})
|
|
assert r.status_code == 400
|
|
assert "CANCELLED" in r.json()["detail"]
|
|
|
|
|
|
def test_cancel_already_submitted_returns_400(tmp_path):
|
|
"""Cancelling a SUBMITTED pending should refuse with ValueError -> 400."""
|
|
c = TestClient(app)
|
|
c.post("/save_connection", json={"name": "prod", "master": "yarn"})
|
|
script = tmp_path / "x.py"
|
|
script.write_text("print('hi')\n")
|
|
prep = c.post(
|
|
"/prepare_submit_job",
|
|
json={"connection": "prod", "script_path": str(script)},
|
|
).json()
|
|
pid = prep["pending_id"]
|
|
# Simulate a SUBMITTED state by mutating the pending directly
|
|
p = pending_store.store.get(pid)
|
|
p.status = "SUBMITTED"
|
|
pending_store.store.save(p)
|
|
r = c.post("/cancel_pending_job", json={"pending_id": pid})
|
|
assert r.status_code == 400
|
|
assert "SUBMITTED" in r.json()["detail"]
|