Files
mcp-server/tests/integration/test_mcp_routes.py
T
Claude f170c3045b feat(submit): confirm defaults instead of removing them
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.
2026-06-26 15:16:43 +08:00

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