Files
mcp-server/spark_executor/tools/submit.py
T
ClaudeandClaude Fable 5 523e6a9c76 refactor(mcp): rename generate_job_file to write_job_file to match what it does
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>
2026-06-29 19:13:29 +08:00

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()}