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/.
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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}
|
||||
@@ -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."
|
||||
),
|
||||
)
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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"])
|
||||
Reference in New Issue
Block a user