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>
342 lines
13 KiB
Python
342 lines
13 KiB
Python
# coding=utf-8
|
|
"""
|
|
@Time :2026/6/24
|
|
@Author :tao.chen
|
|
"""
|
|
import os
|
|
import secrets
|
|
import time
|
|
import uuid
|
|
from datetime import datetime
|
|
|
|
from common.config import settings
|
|
from common.logging import logger
|
|
from common.sql_guard import validate_pyspark_code
|
|
from spark_executor.core.connection_store import store as conn_store
|
|
from spark_executor.core.job_store import JobStore
|
|
from spark_executor.core.log_parser import parse_spark_submit_output
|
|
from spark_executor.core.pending_store import store as pending_store
|
|
from spark_executor.core.spark_submit import (
|
|
SparkSubmitError,
|
|
build_spark_submit_command,
|
|
run_spark_submit,
|
|
)
|
|
from spark_executor.models import Job, PendingSubmission, SubmitResult
|
|
|
|
|
|
def _script_path_error(path: str) -> ValueError:
|
|
"""Instructive error for an invalid script_path. The MCP client receives
|
|
this via the ValueError -> 400 exception handler in server.py."""
|
|
return ValueError(
|
|
f"script_path does not exist or is not a file: {path!r}. "
|
|
f"Two ways to fix this:\n"
|
|
f" 1. (Recommended) Call write_job_file(code=...) first to write "
|
|
f"the PySpark source to disk, then pass the returned script_path.\n"
|
|
f" 2. If the file already exists on the host, mount it into the "
|
|
f"container (e.g. -v /host/path:/app/scripts:ro in docker run) and "
|
|
f"pass the in-container path here."
|
|
)
|
|
|
|
|
|
def _check_script_path(script_path: str) -> None:
|
|
"""Verify script_path points at an existing regular file. Raises ValueError
|
|
(-> 400 via the FastAPI exception handler) with an instructive message."""
|
|
if not script_path or not os.path.isfile(script_path):
|
|
raise _script_path_error(script_path)
|
|
|
|
|
|
def _sql_guard_error(offenses: list[str]) -> ValueError:
|
|
"""Instructive 400 error when the script's SQL fails the safety policy."""
|
|
return ValueError(
|
|
f"Script SQL violates the safety policy. "
|
|
f"Only SELECT and INSERT statements are allowed. "
|
|
f"Found forbidden statement(s): {offenses}. "
|
|
f"Edit the script and try again. (If the code was generated via "
|
|
f"write_job_file, the generator should have caught this — check "
|
|
f"for f-string-based SQL injection where the static analysis can't "
|
|
f"see the runtime value.)"
|
|
)
|
|
|
|
|
|
def _new_pending_id() -> str:
|
|
return "p_" + secrets.token_hex(6)
|
|
|
|
|
|
def prepare_submit_job(
|
|
*,
|
|
connection: str,
|
|
script_path: str,
|
|
queue: str = "default",
|
|
executor_memory: str = "2G",
|
|
executor_cores: int = 2,
|
|
num_executors: int = 2,
|
|
app_name: str,
|
|
extra_args: dict[str, str] | None = None,
|
|
) -> dict[str, object]:
|
|
"""Snapshot connection params and persist a PendingSubmission. Does NOT submit."""
|
|
logger.debug(
|
|
f"prepare_submit_job enter connection={connection} script_path={script_path} "
|
|
f"queue={queue} executor_memory={executor_memory} executor_cores={executor_cores} "
|
|
f"num_executors={num_executors} app_name={app_name}"
|
|
)
|
|
# Order of checks matters for the error the agent sees:
|
|
# 1. Unknown connection -> 404 (KeyError -> 404 in server.py)
|
|
# 2. Missing script file -> 400 (ValueError -> 400)
|
|
# 3. SQL policy violation -> 400 (ValueError -> 400)
|
|
# Connection is checked first because it is a more fundamental problem
|
|
# (the agent is asking about a cluster that doesn't exist), and the
|
|
# agent shouldn't have to fix the script path only to learn the
|
|
# connection name is wrong.
|
|
conn = conn_store.get(connection)
|
|
if conn is None:
|
|
raise KeyError(f"Unknown connection: {connection}")
|
|
# Fail fast: a non-existent path is the most common agent mistake (it
|
|
# generated the code in its own context but forgot to call
|
|
# write_job_file first, or its path refers to the host filesystem
|
|
# which is invisible inside the container). Better to surface this with
|
|
# a 400 + clear remediation than to let spark-submit fail later with
|
|
# an opaque FileNotFoundError -> 500.
|
|
_check_script_path(script_path)
|
|
# SQL safety: re-validate the script even though write_job_file
|
|
# already guards its own output. Catches files written by other means
|
|
# (host volume mounts, manual edits).
|
|
with open(script_path, encoding="utf-8") as f:
|
|
script_body = f.read()
|
|
offenses = validate_pyspark_code(script_body)
|
|
if offenses:
|
|
logger.warning(
|
|
f"prepare_submit_job rejected: SQL policy violation(s) in "
|
|
f"{script_path}: {offenses}"
|
|
)
|
|
raise _sql_guard_error(offenses)
|
|
|
|
pending_id = _new_pending_id()
|
|
pending = PendingSubmission(
|
|
pending_id=pending_id,
|
|
app_name=app_name,
|
|
connection=connection,
|
|
master=conn.master,
|
|
deploy_mode=conn.deploy_mode,
|
|
yarn_rm_url=conn.yarn_rm_url,
|
|
script_path=script_path,
|
|
queue=queue,
|
|
executor_memory=executor_memory,
|
|
executor_cores=executor_cores,
|
|
num_executors=num_executors,
|
|
spark_conf=dict(conn.spark_conf),
|
|
extra_args=dict(extra_args or {}),
|
|
created_at=datetime.utcnow(),
|
|
status="PENDING",
|
|
)
|
|
pending_store.save(pending)
|
|
logger.info(
|
|
f"prepare_submit_job ok pending_id={pending_id} connection={connection} "
|
|
f"master={conn.master} script_path={script_path}"
|
|
)
|
|
return {
|
|
"pending_id": pending_id,
|
|
"status": "PENDING",
|
|
"parameters": pending.model_dump(),
|
|
}
|
|
|
|
|
|
# Module-level job store singleton; replaced in tests.
|
|
job_store: JobStore = JobStore()
|
|
|
|
|
|
def confirm_submit_job(*, pending_id: str) -> SubmitResult:
|
|
"""Actually invoke spark-submit for a previously-prepared PendingSubmission."""
|
|
logger.debug(f"confirm_submit_job enter pending_id={pending_id}")
|
|
pending = pending_store.get(pending_id)
|
|
if pending is None:
|
|
raise KeyError(f"Unknown pending_id: {pending_id}")
|
|
|
|
if pending.status == "SUBMITTED":
|
|
logger.info(
|
|
f"confirm_submit_job idempotent pending_id={pending_id} "
|
|
f"job_id={pending.job_id} application_id={pending.application_id}"
|
|
)
|
|
return SubmitResult(
|
|
job_id=pending.job_id,
|
|
application_id=pending.application_id,
|
|
tracking_url=pending.tracking_url,
|
|
)
|
|
if pending.status == "CANCELLED":
|
|
raise ValueError(f"pending_id {pending_id} is CANCELLED, cannot confirm")
|
|
if pending.status == "FAILED":
|
|
pending.status = "PENDING"
|
|
pending.error = None
|
|
pending_store.save(pending)
|
|
logger.info(f"confirm_submit_job reset pending_id={pending_id} from FAILED to PENDING")
|
|
if pending.status != "PENDING":
|
|
raise ValueError(
|
|
f"pending_id {pending_id} is in status {pending.status!r}, not PENDING"
|
|
)
|
|
|
|
# Defense in depth: re-verify the script still exists. A user could
|
|
# delete the file between prepare and confirm (or an external cleanup
|
|
# job could remove it). 400 via the ValueError -> 400 handler.
|
|
_check_script_path(pending.script_path)
|
|
|
|
cmd = build_spark_submit_command(
|
|
master=pending.master,
|
|
deploy_mode=pending.deploy_mode,
|
|
script_path=pending.script_path,
|
|
queue=pending.queue,
|
|
executor_memory=pending.executor_memory,
|
|
executor_cores=pending.executor_cores,
|
|
num_executors=pending.num_executors,
|
|
spark_conf=pending.spark_conf,
|
|
extra_args=pending.extra_args,
|
|
)
|
|
logger.info(
|
|
f"confirm_submit_job start pending_id={pending_id} "
|
|
f"application_target={pending.master} script_path={pending.script_path}"
|
|
)
|
|
|
|
last_error: SparkSubmitError | None = None
|
|
max_attempts = settings.confirm_max_retries + 1
|
|
for attempt in range(1, max_attempts + 1):
|
|
try:
|
|
result = run_spark_submit(cmd)
|
|
except SparkSubmitError as exc:
|
|
last_error = exc
|
|
logger.warning(
|
|
f"confirm_submit_job attempt {attempt}/{max_attempts} failed "
|
|
f"pending_id={pending_id} err={exc}"
|
|
)
|
|
if attempt < max_attempts:
|
|
time.sleep(settings.confirm_retry_delay_seconds)
|
|
continue
|
|
except Exception:
|
|
# Unexpected failure (not a spark-submit error): fail fast without retry.
|
|
logger.exception(
|
|
f"confirm_submit_job unexpected error pending_id={pending_id}"
|
|
)
|
|
raise
|
|
|
|
application_id, tracking_url = parse_spark_submit_output(result.stderr)
|
|
job_id = uuid.uuid4().hex[:12]
|
|
|
|
# Persist pending as SUBMITTED before creating the in-memory Job so a
|
|
# job_store failure cannot leave the record as PENDING while the YARN
|
|
# app is already running.
|
|
pending.status = "SUBMITTED"
|
|
pending.job_id = job_id
|
|
pending.application_id = application_id
|
|
pending.tracking_url = tracking_url
|
|
pending_store.save(pending)
|
|
|
|
job_store.put(
|
|
Job(
|
|
job_id=job_id,
|
|
application_id=application_id,
|
|
script_path=pending.script_path,
|
|
queue=pending.queue,
|
|
submit_time=datetime.utcnow(),
|
|
connection=pending.connection,
|
|
yarn_rm_url=pending.yarn_rm_url,
|
|
)
|
|
)
|
|
|
|
logger.info(
|
|
f"confirm_submit_job ok pending_id={pending_id} job_id={job_id} "
|
|
f"application_id={application_id}"
|
|
)
|
|
return SubmitResult(
|
|
job_id=job_id,
|
|
application_id=application_id,
|
|
tracking_url=tracking_url,
|
|
)
|
|
|
|
# Exhausted all retries.
|
|
assert last_error is not None
|
|
pending.status = "FAILED"
|
|
pending.error = str(last_error)
|
|
pending_store.save(pending)
|
|
logger.error(
|
|
f"confirm_submit_job failed pending_id={pending_id} after {max_attempts} attempts "
|
|
f"err={last_error}"
|
|
)
|
|
raise last_error
|
|
|
|
|
|
def list_pending_jobs() -> list[dict[str, object]]:
|
|
logger.debug("list_pending_jobs enter")
|
|
return [p.model_dump() for p in pending_store.list_all()]
|
|
|
|
|
|
def get_pending_job(pending_id: str) -> dict[str, object]:
|
|
logger.debug(f"get_pending_job enter pending_id={pending_id}")
|
|
p = pending_store.get(pending_id)
|
|
if p is None:
|
|
raise KeyError(f"Unknown pending_id: {pending_id}")
|
|
return p.model_dump()
|
|
|
|
|
|
def cancel_pending_job(pending_id: str) -> dict[str, str]:
|
|
logger.debug(f"cancel_pending_job enter pending_id={pending_id}")
|
|
p = pending_store.get(pending_id)
|
|
if p is None:
|
|
raise KeyError(f"Unknown pending_id: {pending_id}")
|
|
if p.status in ("SUBMITTED", "FAILED"):
|
|
raise ValueError(
|
|
f"pending_id {pending_id} is in status {p.status!r} and cannot be cancelled"
|
|
)
|
|
p.status = "CANCELLED"
|
|
pending_store.save(p)
|
|
logger.info(f"cancel_pending_job ok pending_id={pending_id}")
|
|
return {"pending_id": pending_id, "status": "CANCELLED"}
|
|
|
|
|
|
def update_pending_job(
|
|
*,
|
|
pending_id: str,
|
|
script_path: str | None = None,
|
|
queue: str | None = None,
|
|
executor_memory: str | None = None,
|
|
executor_cores: int | None = None,
|
|
num_executors: int | None = None,
|
|
app_name: str | None = None,
|
|
extra_args: dict[str, str] | None = None,
|
|
) -> dict[str, object]:
|
|
logger.debug(f"update_pending_job enter pending_id={pending_id}")
|
|
pending = pending_store.get(pending_id)
|
|
if pending is None:
|
|
raise KeyError(f"Unknown pending_id: {pending_id}")
|
|
if pending.status != "PENDING":
|
|
raise ValueError(
|
|
f"pending_id {pending_id} is in status {pending.status!r}; "
|
|
f"only PENDING submissions can be updated"
|
|
)
|
|
|
|
if script_path is not None:
|
|
_check_script_path(script_path)
|
|
with open(script_path, encoding="utf-8") as f:
|
|
script_body = f.read()
|
|
offenses = validate_pyspark_code(script_body)
|
|
if offenses:
|
|
logger.warning(
|
|
f"update_pending_job rejected: SQL policy violation(s) in "
|
|
f"{script_path}: {offenses}"
|
|
)
|
|
raise _sql_guard_error(offenses)
|
|
pending.script_path = script_path
|
|
|
|
if queue is not None:
|
|
pending.queue = queue
|
|
if executor_memory is not None:
|
|
pending.executor_memory = executor_memory
|
|
if executor_cores is not None:
|
|
pending.executor_cores = executor_cores
|
|
if num_executors is not None:
|
|
pending.num_executors = num_executors
|
|
if app_name is not None:
|
|
pending.app_name = app_name
|
|
if extra_args is not None:
|
|
pending.extra_args = dict(extra_args)
|
|
|
|
pending_store.save(pending)
|
|
logger.info(f"update_pending_job ok pending_id={pending_id}")
|
|
return {"pending_id": pending_id, "status": pending.status, "parameters": pending.model_dump()}
|