From 994c2d67abc73b98697b143685ef3a5c6d69429b Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 25 Jun 2026 10:19:59 +0800 Subject: [PATCH] feat: add generate_job_file MCP tool (Stage 2 Task 22) Exposes a single new MCP tool: generate_job_file(code) -> {script_path}. Wiring follows the Stage 1 conventions (Pydantic body model for the route, loguru DEBUG/INFO logging, summary+description for the startup log, MCP arg-pattern via the body model so tools/call roundtrips long code strings without 422-ing on FastAPI query length limits). Flow: 1. LLM calls generate_job_file(code=...) -> {script_path} 2. LLM calls prepare_submit_job(connection=..., script_path=...) -> {pending_id} 3. User reviews the file + pending record 4. LLM calls confirm_submit_job(pending_id=...) -> spark-submit runs The output directory is SPARK_EXECUTOR_JOBS_DIR (default ./data/jobs/, gitignored, persists across container restarts via the existing ./data volume mount in docker-compose.yml). 3 new tests: - Unit: env-var override, default fallback in tmp cwd - Integration: end-to-end body call with a > FastAPI-query-limit code string (the canary test that would have caught the Stage 1 query-params-422 bug) Live MCP smoke verified: tools/list shows 14 tools (13 from Stage 1 + the new one), tools/call generate_job_file returns the absolute path under ./data/jobs/. --- spark_executor/server.py | 19 +++++++++++++++++++ spark_executor/tools/generate.py | 20 ++++++++++++++++++++ spark_executor/tools/requests.py | 11 +++++++++++ tests/integration/test_mcp_routes.py | 21 ++++++++++++++++++++- tests/unit/test_generate_tool.py | 27 +++++++++++++++++++++++++++ 5 files changed, 97 insertions(+), 1 deletion(-) create mode 100644 spark_executor/tools/generate.py create mode 100644 tests/unit/test_generate_tool.py diff --git a/spark_executor/server.py b/spark_executor/server.py index 7734204..daf5aee 100644 --- a/spark_executor/server.py +++ b/spark_executor/server.py @@ -12,11 +12,13 @@ from spark_executor.tools.connections import ( list_connections, save_connection, ) +from spark_executor.tools.generate import generate_job_file from spark_executor.tools.kill import kill_job from spark_executor.tools.logs import get_job_logs from spark_executor.tools.requests import ( ConnectionNameRequest, EmptyRequest, + GenerateJobFileRequest, GetJobLogsRequest, JobIdRequest, PendingIdRequest, @@ -201,3 +203,20 @@ def _get_connection(req: ConnectionNameRequest): ) def _delete_connection(req: ConnectionNameRequest): return delete_connection(req.name) + + +# --- LLM-driven PySpark generation (Stage 2) --- + +@app.post( + "/generate_job_file", + summary="Write LLM-generated PySpark code to disk", + description=( + "Takes a PySpark code string and writes it to a timestamped file under " + "SPARK_EXECUTOR_JOBS_DIR (default ./data/jobs/). Returns the absolute " + "path for use as the script_path argument of prepare_submit_job — the " + "two-step pattern means the LLM can produce code, the user can review " + "the resulting file, and only then is the job submitted." + ), +) +def _generate_job_file(req: GenerateJobFileRequest): + return generate_job_file(req.code) diff --git a/spark_executor/tools/generate.py b/spark_executor/tools/generate.py new file mode 100644 index 0000000..d671822 --- /dev/null +++ b/spark_executor/tools/generate.py @@ -0,0 +1,20 @@ +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +from common.logging import logger +from spark_executor.core.job_writer import write_job_file + + +def generate_job_file(code: str) -> dict[str, str]: + """Write a PySpark code string to disk; return its absolute path. + + Use the returned path as the `script_path` argument of prepare_submit_job. + The output directory is controlled by the SPARK_EXECUTOR_JOBS_DIR env var + (default: ./data/jobs/). + """ + logger.debug(f"generate_job_file enter code_bytes={len(code)}") + path = write_job_file(code) + logger.info(f"generate_job_file ok script_path={path}") + return {"script_path": path} diff --git a/spark_executor/tools/requests.py b/spark_executor/tools/requests.py index d7c0d55..2242a68 100644 --- a/spark_executor/tools/requests.py +++ b/spark_executor/tools/requests.py @@ -48,3 +48,14 @@ class GetJobLogsRequest(BaseModel): class ConnectionNameRequest(BaseModel): name: str + + +class GenerateJobFileRequest(BaseModel): + code: str = Field( + ..., + description=( + "Full PySpark source code to write to disk. Will be passed verbatim " + "to spark-submit after the agent calls prepare_submit_job on the " + "returned path." + ), + ) diff --git a/tests/integration/test_mcp_routes.py b/tests/integration/test_mcp_routes.py index c77a975..4495f33 100644 --- a/tests/integration/test_mcp_routes.py +++ b/tests/integration/test_mcp_routes.py @@ -1,4 +1,5 @@ # coding=utf-8 +import os from pathlib import Path import pytest @@ -32,7 +33,7 @@ def test_health_still_present(): assert r.json() == {"status": "ok"} -def test_twelve_tool_routes_registered(): +def test_thirteen_tool_routes_registered(): paths = {r.path for r in app.routes} for path in ( # pending-submission flow (5) @@ -50,6 +51,8 @@ def test_twelve_tool_routes_registered(): "/list_connections", "/get_connection", "/delete_connection", + # LLM-driven PySpark generation (1) — Stage 2 + "/generate_job_file", ): assert path in paths, f"missing MCP tool route: {path}" @@ -112,6 +115,22 @@ def test_list_connections_works_with_empty_body(): 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(): diff --git a/tests/unit/test_generate_tool.py b/tests/unit/test_generate_tool.py new file mode 100644 index 0000000..e95fc80 --- /dev/null +++ b/tests/unit/test_generate_tool.py @@ -0,0 +1,27 @@ +# coding=utf-8 +import os +from pathlib import Path + +from spark_executor.tools import generate + + +def test_generate_writes_code_and_returns_path(monkeypatch, tmp_path: Path): + monkeypatch.chdir(tmp_path) + out = generate.generate_job_file( + "from pyspark.sql import SparkSession\n" + "spark = SparkSession.builder.getOrCreate()\n" + ) + assert "script_path" in out + p = out["script_path"] + assert os.path.isabs(p) + assert p.startswith(str(tmp_path / "data" / "jobs")) + assert p.endswith(".py") + with open(p) as f: + assert "SparkSession.builder.getOrCreate()" in f.read() + + +def test_generate_uses_env_var_when_set(monkeypatch, tmp_path: Path): + monkeypatch.setenv("SPARK_EXECUTOR_JOBS_DIR", str(tmp_path / "custom")) + out = generate.generate_job_file("x = 1\n") + assert out["script_path"].startswith(str(tmp_path / "custom")) + assert os.path.isfile(out["script_path"])