From f170c3045b35c396f2a4f4947f2f1fe299110027 Mon Sep 17 00:00:00 2001 From: Claude Date: Fri, 26 Jun 2026 15:16:43 +0800 Subject: [PATCH] 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. --- spark_executor/server.py | 15 +++++++++ spark_executor/tools/requests.py | 8 ++--- spark_executor/tools/submit.py | 8 ++--- tests/integration/test_mcp_routes.py | 48 ++++++++++++++++++++++++++++ tests/unit/test_requests.py | 14 +++++++- tests/unit/test_submit_tool.py | 19 +++++++++++ 6 files changed, 103 insertions(+), 9 deletions(-) diff --git a/spark_executor/server.py b/spark_executor/server.py index 8fe42be..a60a08e 100644 --- a/spark_executor/server.py +++ b/spark_executor/server.py @@ -42,6 +42,14 @@ from spark_executor.tools.result import get_job_result app = FastAPI(title="Spark Executor MCP", version="0.0.1", description="Spark Executor MCP Server") +_DEFAULTS_TO_CONFIRM = { + "queue": "default", + "executor_memory": "4G", + "executor_cores": 2, + "num_executors": 2, +} + + # --- Exception handlers: translate tool-layer errors into proper HTTP statuses --- # # Tool functions raise KeyError for "unknown id" (job_id, pending_id, connection @@ -88,6 +96,13 @@ def health_check(): ), ) def _prepare_submit_job(req: PrepareSubmitJobRequest): + omitted = [f for f in _DEFAULTS_TO_CONFIRM if f not in req.model_fields_set] + if omitted: + details = ", ".join(f"{f}={_DEFAULTS_TO_CONFIRM[f]!r}" for f in omitted) + raise ValueError( + f"Please confirm default values: {details}. " + f"Resubmit with these fields explicitly set." + ) return prepare_submit_job(**req.model_dump()) diff --git a/spark_executor/tools/requests.py b/spark_executor/tools/requests.py index 8f342af..1529b9f 100644 --- a/spark_executor/tools/requests.py +++ b/spark_executor/tools/requests.py @@ -50,19 +50,19 @@ class PrepareSubmitJobRequest(BaseModel): ), ) queue: str = Field( - ..., + default="default", description="YARN queue to submit to. Must be explicitly confirmed by the caller.", ) executor_memory: str = Field( - ..., + default="4G", description="Executor memory, e.g. '4G'. Must be explicitly confirmed by the caller.", ) executor_cores: int = Field( - ..., + default=2, description="Number of cores per executor. Must be explicitly confirmed by the caller.", ) num_executors: int = Field( - ..., + default=2, description="Total number of executors. Must be explicitly confirmed by the caller.", ) extra_args: dict[str, str] | None = Field( diff --git a/spark_executor/tools/submit.py b/spark_executor/tools/submit.py index 5fb33b1..be6585d 100644 --- a/spark_executor/tools/submit.py +++ b/spark_executor/tools/submit.py @@ -64,10 +64,10 @@ def prepare_submit_job( *, connection: str, script_path: str, - queue: str, - executor_memory: str, - executor_cores: int, - num_executors: int, + queue: str = "default", + executor_memory: str = "4G", + executor_cores: int = 2, + num_executors: int = 2, app_name: str, extra_args: dict[str, str] | None = None, ) -> dict[str, object]: diff --git a/tests/integration/test_mcp_routes.py b/tests/integration/test_mcp_routes.py index c49a77e..1d6a6dd 100644 --- a/tests/integration/test_mcp_routes.py +++ b/tests/integration/test_mcp_routes.py @@ -104,6 +104,54 @@ def test_prepare_submit_job_works_via_body(tmp_path): 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) diff --git a/tests/unit/test_requests.py b/tests/unit/test_requests.py index 1e1f51e..a939b61 100644 --- a/tests/unit/test_requests.py +++ b/tests/unit/test_requests.py @@ -4,11 +4,23 @@ import pytest from spark_executor.tools.requests import PrepareSubmitJobRequest -def test_prepare_submit_job_request_requires_all_confirmed_parameters(): +def test_prepare_submit_job_request_requires_connection_script_path_and_app_name(): with pytest.raises(ValueError): PrepareSubmitJobRequest(connection="prod", script_path="/tmp/x.py") +def test_prepare_submit_job_request_restores_defaults_for_queue_and_resources(): + req = PrepareSubmitJobRequest( + connection="prod", + script_path="/tmp/x.py", + app_name="test-app", + ) + assert req.queue == "default" + assert req.executor_memory == "4G" + assert req.executor_cores == 2 + assert req.num_executors == 2 + + def test_prepare_submit_job_request_accepts_extra_args(): req = PrepareSubmitJobRequest( connection="prod", diff --git a/tests/unit/test_submit_tool.py b/tests/unit/test_submit_tool.py index 1a87e2e..29a07bc 100644 --- a/tests/unit/test_submit_tool.py +++ b/tests/unit/test_submit_tool.py @@ -112,6 +112,25 @@ def test_prepare_requires_explicit_app_name(monkeypatch, real_script): submit.prepare_submit_job(connection="prod", script_path=str(real_script)) +def test_prepare_accepts_restored_defaults(monkeypatch, real_script): + monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) + out = submit.prepare_submit_job( + connection="prod", + script_path=str(real_script), + queue="default", + executor_memory="4G", + executor_cores=2, + num_executors=2, + app_name="test-app", + ) + assert out["status"] == "PENDING" + p = submit.pending_store.get(_last_pending_id()) + assert p.queue == "default" + assert p.executor_memory == "4G" + assert p.executor_cores == 2 + assert p.num_executors == 2 + + def test_prepare_snapshots_extra_args(monkeypatch, real_script): monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) submit.prepare_submit_job(