Restore defaults for queue/executor_memory/executor_cores/num_executors in prepare_submit_job, but require the caller to explicitly confirm them. If any defaulted field is omitted, the route returns HTTP 400 listing the defaults and asks the caller to resubmit with explicit values. app_name remains required (no meaningful default). extra_args remains optional. Tests cover rejection of unconfirmed defaults and acceptance of explicitly confirmed defaults.
409 lines
13 KiB
Python
409 lines
13 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_sixteen_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 (3) — Stage 2
|
|
"/generate_job_file",
|
|
"/read_job_file",
|
|
"/update_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",
|
|
"executor_memory": "4G",
|
|
"executor_cores": 2,
|
|
"num_executors": 2,
|
|
"app_name": "test-app",
|
|
},
|
|
)
|
|
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_unconfirmed_defaults(tmp_path):
|
|
c = TestClient(app)
|
|
c.post("/save_connection", json={"name": "prod", "master": "yarn"})
|
|
script = tmp_path / "demo.py"
|
|
script.write_text("print('hi')\n")
|
|
r = c.post(
|
|
"/prepare_submit_job",
|
|
json={
|
|
"connection": "prod",
|
|
"script_path": str(script),
|
|
"app_name": "test-app",
|
|
},
|
|
)
|
|
assert r.status_code == 400
|
|
detail = r.json()["detail"]
|
|
assert "Please confirm default values" in detail
|
|
assert "queue='default'" in detail
|
|
assert "executor_memory='4G'" in detail
|
|
assert "executor_cores=2" in detail
|
|
assert "num_executors=2" in detail
|
|
|
|
|
|
def test_prepare_accepts_confirmed_defaults(tmp_path):
|
|
c = TestClient(app)
|
|
c.post("/save_connection", json={"name": "prod", "master": "yarn"})
|
|
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": "default",
|
|
"executor_memory": "4G",
|
|
"executor_cores": 2,
|
|
"num_executors": 2,
|
|
"app_name": "test-app",
|
|
},
|
|
)
|
|
assert r.status_code == 200
|
|
body = r.json()
|
|
assert body["status"] == "PENDING"
|
|
assert body["parameters"]["queue"] == "default"
|
|
assert body["parameters"]["executor_memory"] == "4G"
|
|
assert body["parameters"]["executor_cores"] == 2
|
|
assert body["parameters"]["num_executors"] == 2
|
|
|
|
|
|
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",
|
|
"queue": "default",
|
|
"executor_memory": "4G",
|
|
"executor_cores": 2,
|
|
"num_executors": 2,
|
|
"app_name": "test-app",
|
|
},
|
|
)
|
|
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),
|
|
"queue": "default",
|
|
"executor_memory": "4G",
|
|
"executor_cores": 2,
|
|
"num_executors": 2,
|
|
"app_name": "test-app",
|
|
},
|
|
)
|
|
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")) == []
|
|
|
|
|
|
# --- Stage 2: read_job_file / update_job_file ---
|
|
|
|
def test_read_job_file_returns_content_via_mcp(tmp_path):
|
|
c = TestClient(app)
|
|
script = tmp_path / "demo.py"
|
|
script.write_text("print('hello from read_job_file')\n")
|
|
r = c.post("/read_job_file", json={"script_path": str(script)})
|
|
assert r.status_code == 200
|
|
body = r.json()
|
|
assert body["content"] == "print('hello from read_job_file')\n"
|
|
assert body["path"] == str(script)
|
|
|
|
|
|
def test_read_job_file_404_for_missing_file(tmp_path):
|
|
c = TestClient(app)
|
|
r = c.post("/read_job_file", json={"script_path": str(tmp_path / "nope.py")})
|
|
assert r.status_code == 400
|
|
assert "does not exist" in r.json()["detail"]
|
|
|
|
|
|
def test_update_job_file_writes_and_reads_back(tmp_path):
|
|
from common import config
|
|
config.settings.jobs_dir = str(tmp_path)
|
|
c = TestClient(app)
|
|
script = tmp_path / "edit.py"
|
|
script.write_text("v1\n")
|
|
# Update
|
|
r1 = c.post("/update_job_file", json={"script_path": str(script), "content": "v2\n"})
|
|
assert r1.status_code == 200
|
|
assert r1.json()["bytes_written"] == 3
|
|
# Read back
|
|
r2 = c.post("/read_job_file", json={"script_path": str(script)})
|
|
assert r2.json()["content"] == "v2\n"
|
|
|
|
|
|
def test_update_job_file_rejects_paths_outside_jobs_dir(tmp_path):
|
|
from common import config
|
|
config.settings.jobs_dir = str(tmp_path / "jobs")
|
|
c = TestClient(app)
|
|
other = tmp_path / "elsewhere.py"
|
|
other.write_text("x\n")
|
|
r = c.post("/update_job_file", json={"script_path": str(other), "content": "y\n"})
|
|
assert r.status_code == 400
|
|
assert "must be under" in r.json()["detail"]
|
|
|
|
|
|
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),
|
|
"queue": "default",
|
|
"executor_memory": "4G",
|
|
"executor_cores": 2,
|
|
"num_executors": 2,
|
|
"app_name": "test-app",
|
|
},
|
|
).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",
|
|
"queue": "default",
|
|
"executor_memory": "4G",
|
|
"executor_cores": 2,
|
|
"num_executors": 2,
|
|
"app_name": "test-app",
|
|
},
|
|
)
|
|
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),
|
|
"queue": "default",
|
|
"executor_memory": "4G",
|
|
"executor_cores": 2,
|
|
"num_executors": 2,
|
|
"app_name": "test-app",
|
|
},
|
|
).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),
|
|
"queue": "default",
|
|
"executor_memory": "4G",
|
|
"executor_cores": 2,
|
|
"num_executors": 2,
|
|
"app_name": "test-app",
|
|
},
|
|
).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"]
|