The MCP tool was named `generate_job_file` from Stage 2 but it does
NOT generate PySpark code — the calling LLM writes the code in its own
context, and this tool only persists it to a file under
SPARK_EXECUTOR_JOBS_DIR so `spark-submit` can see it. The misleading
`generate_` prefix sent agents (and humans) looking for a code
generator that doesn't exist.
This commit folds three related polish changes into one (split later
with rebase -i if you want them as separate history):
1. The rename itself:
- `tools/generate.py` → `tools/write_job.py`
- `generate_job_file` → `write_job_file`
- `GenerateJobFileRequest` → `WriteJobFileRequest`
- `/generate_job_file` route → `/write_job_file`
- `operation_id="generate_job_file"` → `operation_id="write_job_file"`
The internal helper `core.job_writer.write_job_file` (which just
writes bytes to disk with no SQL guard) is imported with an
`_write_to_disk` alias to avoid the name collision with the
MCP-exposed function in the same module.
The description for the tool now explicitly states 'this tool
does NOT generate PySpark code. The calling LLM is expected to
have already written the code; this tool only persists it.'
2. Skill for LLM agents operating the service
(`docs/superpowers/skills/spark-executor-mcp-operate/SKILL.md`,
449 lines). Covers the 16 tools, the two-step prepare/confirm
flow, the dual-ID contract (job_id vs application_id), the
PendingSubmission state machine, the Connection profile, the
job-file workflow, the error reference, common pitfalls, and a
full end-to-end word-count example.
3. Default `executor_memory` lowered 4G → 2G
(`_DEFAULTS_TO_CONFIRM` in `server.py`). Mirrors the matching
change in `test_mcp_routes.py` and the 5 unit tests that
reference the default. Aligns with the lighter workloads the
service is sized for in its current container profile.
Also tracked in git for the first time:
- `docs/superpowers/plans/2026-06-24-spark-executor-mcp.md`
(the original Stage 1/2/3 design plan, updated to use the new
tool name throughout).
Test rename:
- `tests/unit/test_generate_tool.py` → `test_write_job_tool.py`
- the new test file picks up an extra assertion that the SQL guard
rejects a `DROP TABLE` statement at write time.
243 tests pass (was 242; +1 new SQL-guard assertion). Zero regressions.
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
534 lines
18 KiB
Python
534 lines
18 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
|
|
"/write_job_file",
|
|
"/read_job_file",
|
|
"/update_job_file",
|
|
):
|
|
assert path in paths, f"missing MCP tool route: {path}"
|
|
|
|
|
|
def test_seventeen_tool_routes_registered():
|
|
paths = {r.path for r in app.routes}
|
|
assert "/update_pending_job" in paths
|
|
|
|
|
|
# --- operation_id: pin clean MCP tool names (no auto-generated suffixes) ---
|
|
#
|
|
# fastapi-mcp uses each route's OpenAPI `operationId` as the MCP tool name
|
|
# exposed to LLM clients (fastapi_mcp/server.py:619-620). If we omit
|
|
# `operation_id=`, FastAPI auto-generates a name like
|
|
# `prepare_submit_job_prepare_submit_job_post` (function_name + path +
|
|
# method), which is ugly for the agent. This test asserts every MCP tool
|
|
# route has an explicit, clean operation_id matching its handler.
|
|
|
|
# Set of routes that are NOT MCP tools (no operation_id expected).
|
|
_NON_MCP_PATHS = {"/health", "/openapi.json", "/docs", "/docs/oauth2-redirect", "/redoc"}
|
|
|
|
|
|
def test_every_mcp_route_has_explicit_clean_operation_id():
|
|
c = TestClient(app)
|
|
openapi = c.get("/openapi.json").json()
|
|
paths = openapi["paths"]
|
|
|
|
seen: list[tuple[str, str, str]] = [] # (path, method, operation_id)
|
|
for path, methods in paths.items():
|
|
if path in _NON_MCP_PATHS:
|
|
continue
|
|
for method, op in methods.items():
|
|
if method.upper() not in {"GET", "POST", "PUT", "DELETE", "PATCH"}:
|
|
continue
|
|
op_id = op.get("operationId")
|
|
assert op_id, (
|
|
f"route {method.upper()} {path} has no operationId in OpenAPI — "
|
|
f"fastapi-mcp will fall back to an auto-generated name like "
|
|
f"`{path.strip('/').replace('/', '_')}_{method}_...` and expose "
|
|
f"that to LLM clients. Add an explicit `operation_id=` to the decorator."
|
|
)
|
|
# Reject FastAPI's auto-generated format: it always embeds the
|
|
# HTTP method as a `_post`/`_get` suffix, and repeats the path.
|
|
assert not op_id.endswith(f"_{method}"), (
|
|
f"route {method.upper()} {path} has auto-generated operationId "
|
|
f"{op_id!r} (ends with `_{method}`). Add an explicit "
|
|
f"`operation_id=` to the decorator."
|
|
)
|
|
seen.append((path, method.upper(), op_id))
|
|
|
|
# Sanity: we expect at least the 16 MCP tool routes from the project.
|
|
assert len(seen) >= 16, f"only found {len(seen)} MCP tool routes, expected >=16: {seen}"
|
|
|
|
# Every operation_id must be unique — fastapi-mcp's tool registry keys
|
|
# by name, so duplicates would silently shadow one another.
|
|
ids = [op_id for _, _, op_id in seen]
|
|
dupes = [x for x in set(ids) if ids.count(x) > 1]
|
|
assert not dupes, f"duplicate operation_ids: {dupes}"
|
|
|
|
|
|
# --- 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": "2G",
|
|
"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='2G'" 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": "2G",
|
|
"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"] == "2G"
|
|
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": "2G",
|
|
"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 "write_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 — write_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": "2G",
|
|
"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_write_rejects_sql_policy_violation_with_400(tmp_path):
|
|
"""Same policy at write_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(
|
|
"/write_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": "2G",
|
|
"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: write_job_file ---
|
|
|
|
def test_write_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("/write_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": "2G",
|
|
"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": "2G",
|
|
"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": "2G",
|
|
"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"]
|
|
|
|
|
|
# --- update_pending_job route ---
|
|
|
|
def test_update_pending_job_route(tmp_path):
|
|
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": "2G",
|
|
"executor_cores": 2,
|
|
"num_executors": 2,
|
|
"app_name": "test-app",
|
|
},
|
|
).json()
|
|
pid = prep["pending_id"]
|
|
|
|
r = c.post(
|
|
"/update_pending_job",
|
|
json={"pending_id": pid, "queue": "research"},
|
|
)
|
|
assert r.status_code == 200, r.text
|
|
assert r.json()["status"] == "PENDING"
|
|
assert r.json()["parameters"]["queue"] == "research"
|
|
|
|
got = c.post("/get_pending_job", json={"pending_id": pid}).json()
|
|
assert got["queue"] == "research"
|
|
# untouched fields are preserved
|
|
assert got["executor_memory"] == "2G"
|
|
assert got["app_name"] == "test-app"
|
|
|
|
|
|
def test_update_pending_job_route_rejects_submitted(tmp_path, monkeypatch):
|
|
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": "2G",
|
|
"executor_cores": 2,
|
|
"num_executors": 2,
|
|
"app_name": "test-app",
|
|
},
|
|
).json()
|
|
pid = prep["pending_id"]
|
|
fake_proc = type("P", (), {
|
|
"returncode": 0,
|
|
"stderr": "tracking URL: http://rm:8088/proxy/application_17400000001/\n",
|
|
})()
|
|
monkeypatch.setattr(submit, "run_spark_submit", lambda _cmd: fake_proc)
|
|
c.post("/confirm_submit_job", json={"pending_id": pid})
|
|
|
|
r = c.post(
|
|
"/update_pending_job",
|
|
json={"pending_id": pid, "queue": "research"},
|
|
)
|
|
assert r.status_code == 400
|
|
assert "only PENDING submissions can be updated" in r.json()["detail"]
|