Files
mcp-server/tests/integration/test_mcp_routes.py
T
ClaudeandClaude Fable 5 523e6a9c76 refactor(mcp): rename generate_job_file to write_job_file to match what it does
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>
2026-06-29 19:13:29 +08:00

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"]