Two user-reported bugs, same root cause: the in-memory JobStore + the
'job_id must be the 12-char hex' tool contract.
Bug 1: 'Unknown job_id' reported frequently
JobStore was a process-local dict (spark_executor/core/job_store.py).
Under gunicorn workers > 1, a job created by confirm_submit_job
landing on worker A was invisible to worker B, so a follow-up
get_job_status / get_job_result / get_job_logs / kill_job landing on
a different worker returned 'Unknown job_id'. Same multi-worker
problem that bit the MCP session layer; only the affected data was
different.
Bug 2: 'get_job_logs frequently confuses job_id and application_id'
confirm_submit_job returns BOTH identifiers in SubmitResult, but
get_job_logs (and friends) only accepted the local 12-char job_id
and never said so in their description. When the agent passed the
YARN application_id, the error message itself was misleading:
'Unknown job_id: application_17400000001_0001' — the agent had
passed an id, just the wrong kind.
This change fixes both at the root:
* JobStore is now JSON-backed at data/jobs.json (atomic tempfile +
os.replace), with cross-process safety via fcntl.flock on a sibling
.lock file. Stage 3's SQLite migration is still planned; the file
format is intentionally simple so it is a straight
'for j in read_all(): db.insert(j)'.
* New JobStore.get_either(uid) looks up by job_id first, then
application_id. All four job-lifecycle tools (get_job_status,
get_job_result, get_job_logs, kill_job) call get_either instead
of get(job_id), so the agent can pass either identifier and get
the same answer.
* The 'neither matched' KeyError now spells out both id forms and
what they look like, so the agent isn't left guessing.
* server.py tool descriptions for the four job tools explicitly
state 'job_id accepts BOTH identifiers' so this is visible to the
LLM at tool-selection time, not only at error time.
Tests:
* test_job_store.py: tmp_path isolation, persistence across
instances, human-readable JSON, corrupt-file resilience,
get_either (by job_id, by application_id, collision preference,
unknown), put idempotency.
* test_{logs,status,kill,result}_tool.py: per-test tmp_path fixture,
'accepts application_id' regression for each tool, and an
explicit assertion that the unknown-id error message mentions
BOTH id forms. test_result_raises_keyerror_for_unknown_job's
match pattern updated for the new message.
242 tests pass (was 226; +16 new). Zero regressions.
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
346 lines
12 KiB
Python
346 lines
12 KiB
Python
# coding=utf-8
|
|
"""
|
|
@Time :2026/6/24
|
|
@Author :tao.chen
|
|
"""
|
|
from fastapi import FastAPI, Request
|
|
from fastapi.responses import JSONResponse
|
|
|
|
from spark_executor.tools.connections import (
|
|
delete_connection,
|
|
get_connection,
|
|
list_connections,
|
|
save_connection,
|
|
)
|
|
from spark_executor.tools.generate import generate_job_file
|
|
from spark_executor.tools.job_file import read_job_file, update_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,
|
|
PrepareSubmitJobRequest,
|
|
ReadJobFileRequest,
|
|
SaveConnectionRequest,
|
|
UpdateJobFileRequest,
|
|
UpdatePendingJobRequest,
|
|
)
|
|
from spark_executor.tools.status import get_job_status
|
|
from spark_executor.tools.submit import (
|
|
cancel_pending_job,
|
|
confirm_submit_job,
|
|
get_pending_job,
|
|
list_pending_jobs,
|
|
prepare_submit_job,
|
|
update_pending_job,
|
|
)
|
|
|
|
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
|
|
# name) and ValueError for invalid state transitions (e.g. confirming a
|
|
# CANCELLED pending). Without these handlers FastAPI would map them to a bare
|
|
# 500 "Internal Server Error" which is useless to MCP clients.
|
|
|
|
@app.exception_handler(KeyError)
|
|
async def _keyerror_handler(_request: Request, exc: KeyError) -> JSONResponse:
|
|
return JSONResponse(status_code=404, content={"detail": str(exc)})
|
|
|
|
|
|
@app.exception_handler(ValueError)
|
|
async def _valueerror_handler(_request: Request, exc: ValueError) -> JSONResponse:
|
|
return JSONResponse(status_code=400, content={"detail": str(exc)})
|
|
|
|
|
|
@app.get("/health")
|
|
def health_check():
|
|
return {"status": "ok"}
|
|
|
|
|
|
# MCP tool routes. fastapi-mcp discovers these and registers them as MCP tools.
|
|
# Each route takes a single Pydantic body model so tools/call (which sends args
|
|
# as JSON body) works for every tool, including those with dict-typed params
|
|
# like spark_conf.
|
|
|
|
# --- Pending submission flow (two-step submit) ---
|
|
|
|
@app.post(
|
|
"/prepare_submit_job",
|
|
operation_id="prepare_submit_job",
|
|
summary="Prepare a Spark job submission (no spark-submit yet)",
|
|
description=(
|
|
"Snapshot the named Connection's master / deploy_mode / spark_conf / "
|
|
"yarn_rm_url into a PendingSubmission record and persist it. "
|
|
"Does NOT invoke spark-submit. Returns pending_id for use with "
|
|
"confirm_submit_job (the user-second-confirmation step).\n\n"
|
|
"REQUIRED PATTERN for LLM-generated code: call generate_job_file(code=...) "
|
|
"first, then pass the returned script_path here. Direct submission with a "
|
|
"synthetic path (one that only exists in the agent's context) will be "
|
|
"rejected with HTTP 400 — the script must exist inside the container's "
|
|
"filesystem. For pre-existing files, mount the host directory into the "
|
|
"container and pass the in-container path."
|
|
),
|
|
)
|
|
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())
|
|
|
|
|
|
@app.post(
|
|
"/confirm_submit_job",
|
|
operation_id="confirm_submit_job",
|
|
summary="Confirm and submit a previously-prepared job",
|
|
description=(
|
|
"Actually invoke spark-submit for the PendingSubmission identified "
|
|
"by pending_id. Requires status=PENDING. On success, transitions the "
|
|
"pending entry to SUBMITTED and creates a Job record. On failure, "
|
|
"marks the entry FAILED and re-raises."
|
|
),
|
|
)
|
|
def _confirm_submit_job(req: PendingIdRequest):
|
|
return confirm_submit_job(pending_id=req.pending_id)
|
|
|
|
|
|
@app.post(
|
|
"/list_pending_jobs",
|
|
operation_id="list_pending_jobs",
|
|
summary="List all pending submissions",
|
|
description="Return every PendingSubmission in any status (PENDING, SUBMITTED, CANCELLED, FAILED).",
|
|
)
|
|
def _list_pending_jobs(_req: EmptyRequest = EmptyRequest()):
|
|
return list_pending_jobs()
|
|
|
|
|
|
@app.post(
|
|
"/get_pending_job",
|
|
operation_id="get_pending_job",
|
|
summary="Get a single pending submission",
|
|
description="Return the PendingSubmission identified by pending_id, including its current status and outcome fields.",
|
|
)
|
|
def _get_pending_job(req: PendingIdRequest):
|
|
return get_pending_job(req.pending_id)
|
|
|
|
|
|
@app.post(
|
|
"/update_pending_job",
|
|
operation_id="update_pending_job",
|
|
summary="Update an unsubmitted pending submission",
|
|
description=(
|
|
"Modify parameters of a PENDING submission before confirm_submit_job. "
|
|
"Only the provided fields are changed. If script_path is changed, the "
|
|
"new file must exist and pass the SQL guard."
|
|
),
|
|
)
|
|
def _update_pending_job(req: UpdatePendingJobRequest):
|
|
return update_pending_job(**req.model_dump(exclude_none=True))
|
|
|
|
|
|
@app.post(
|
|
"/cancel_pending_job",
|
|
operation_id="cancel_pending_job",
|
|
summary="Cancel a pending submission",
|
|
description=(
|
|
"Flip a PENDING (or already-CANCELLED) PendingSubmission to CANCELLED. "
|
|
"Refuses to cancel entries that are SUBMITTED or FAILED — those are "
|
|
"terminal and must be killed via kill_job instead."
|
|
),
|
|
)
|
|
def _cancel_pending_job(req: PendingIdRequest):
|
|
return cancel_pending_job(req.pending_id)
|
|
|
|
|
|
# --- Spark job tools ---
|
|
|
|
@app.post(
|
|
"/get_job_status",
|
|
operation_id="get_job_status",
|
|
summary="Query YARN for a job's current status",
|
|
description=(
|
|
"Return the YARN application state (RUNNING / SUCCEEDED / FAILED / "
|
|
"KILLED / ACCEPTED / NEW / NEW_SAVING / SUBMITTED / etc.) plus the "
|
|
"raw YARN REST response body.\n\n"
|
|
"**job_id accepts BOTH identifiers** returned by "
|
|
"confirm_submit_job: the local job_id (12-char hex, e.g. "
|
|
"'a1b2c3d4e5f6') and the YARN application_id (e.g. "
|
|
"'application_17400000001_0001'). The lookup is by job_id first, "
|
|
"then by application_id."
|
|
),
|
|
)
|
|
def _get_job_status(req: JobIdRequest):
|
|
return get_job_status(req.job_id)
|
|
|
|
|
|
@app.post(
|
|
"/get_job_result",
|
|
operation_id="get_job_result",
|
|
summary="Query YARN for a job's terminal result view",
|
|
description=(
|
|
"Return a terminal-oriented view of a Spark job: final_status, "
|
|
"diagnostics, tracking_url, started_time, and finished_time. "
|
|
"This is distinct from get_job_status, which is for polling the "
|
|
"running YARN state and returns the raw YARN response.\n\n"
|
|
"**job_id accepts BOTH identifiers** returned by "
|
|
"confirm_submit_job: the local job_id (12-char hex, e.g. "
|
|
"'a1b2c3d4e5f6') and the YARN application_id (e.g. "
|
|
"'application_17400000001_0001'). The lookup is by job_id first, "
|
|
"then by application_id."
|
|
),
|
|
)
|
|
def _get_job_result(req: JobIdRequest):
|
|
return get_job_result(req.job_id)
|
|
|
|
|
|
@app.post(
|
|
"/get_job_logs",
|
|
operation_id="get_job_logs",
|
|
summary="Fetch aggregated container logs for a job",
|
|
description=(
|
|
"Pull aggregated logs from the YARN ResourceManager. Returns the last "
|
|
"tail_chars characters (default 5000). Requires yarn.log-aggregation-enable "
|
|
"to be true on the target cluster.\n\n"
|
|
"**job_id accepts BOTH identifiers** returned by "
|
|
"confirm_submit_job: the local job_id (12-char hex, e.g. "
|
|
"'a1b2c3d4e5f6') and the YARN application_id (e.g. "
|
|
"'application_17400000001_0001'). The lookup is by job_id first, "
|
|
"then by application_id."
|
|
),
|
|
)
|
|
def _get_job_logs(req: GetJobLogsRequest):
|
|
return get_job_logs(req.job_id, tail_chars=req.tail_chars)
|
|
|
|
|
|
@app.post(
|
|
"/kill_job",
|
|
operation_id="kill_job",
|
|
summary="Kill a running job",
|
|
description=(
|
|
"PUT state=KILLED to YARN REST API for the job's application_id. "
|
|
"**job_id accepts BOTH identifiers** returned by "
|
|
"confirm_submit_job: the local job_id (12-char hex) and the YARN "
|
|
"application_id. The lookup is by job_id first, then by application_id."
|
|
),
|
|
)
|
|
def _kill_job(req: JobIdRequest):
|
|
return kill_job(req.job_id)
|
|
|
|
|
|
# --- Connection management tools ---
|
|
|
|
@app.post(
|
|
"/save_connection",
|
|
operation_id="save_connection",
|
|
summary="Save or update a named Spark connection",
|
|
description=(
|
|
"Upsert a Connection record (master URL, deploy mode, optional YARN RM URL, "
|
|
"spark_conf K/V) keyed by name. Used by prepare_submit_job via the "
|
|
"connection parameter."
|
|
),
|
|
)
|
|
def _save_connection(req: SaveConnectionRequest):
|
|
# exclude_none so we don't overwrite the function's default with explicit None
|
|
return save_connection(**req.model_dump(exclude_none=True))
|
|
|
|
|
|
@app.post(
|
|
"/list_connections",
|
|
operation_id="list_connections",
|
|
summary="List all saved Spark connections",
|
|
description="Return every Connection in the registry (model_dump form).",
|
|
)
|
|
def _list_connections(_req: EmptyRequest = EmptyRequest()):
|
|
return list_connections()
|
|
|
|
|
|
@app.post(
|
|
"/get_connection",
|
|
operation_id="get_connection",
|
|
summary="Get a single connection by name",
|
|
description="Return the Connection record, or 404 if not found.",
|
|
)
|
|
def _get_connection(req: ConnectionNameRequest):
|
|
return get_connection(req.name)
|
|
|
|
|
|
@app.post(
|
|
"/delete_connection",
|
|
operation_id="delete_connection",
|
|
summary="Delete a saved connection",
|
|
description="Remove a Connection by name. 404 if not found.",
|
|
)
|
|
def _delete_connection(req: ConnectionNameRequest):
|
|
return delete_connection(req.name)
|
|
|
|
|
|
# --- LLM-driven PySpark generation (Stage 2) ---
|
|
|
|
@app.post(
|
|
"/generate_job_file",
|
|
operation_id="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)
|
|
|
|
|
|
@app.post(
|
|
"/read_job_file",
|
|
operation_id="read_job_file",
|
|
summary="Read the contents of an existing PySpark script",
|
|
description=(
|
|
"Returns the text content of an existing script file at the given "
|
|
"path. Caps reads at 1 MB. Typical use: after generate_job_file "
|
|
"returns a path, call read_job_file on that path to inspect what "
|
|
"was actually written, before deciding to prepare_submit_job or "
|
|
"update_job_file."
|
|
),
|
|
)
|
|
def _read_job_file(req: ReadJobFileRequest):
|
|
return read_job_file(req.script_path)
|
|
|
|
|
|
@app.post(
|
|
"/update_job_file",
|
|
operation_id="update_job_file",
|
|
summary="Overwrite an existing PySpark script with new content",
|
|
description=(
|
|
"Replaces the entire content of an existing script file. Path must "
|
|
"be under SPARK_EXECUTOR_JOBS_DIR (the dir generate_job_file writes "
|
|
"to) — protects against overwriting host-mounted configs or other "
|
|
"non-script files. Caps writes at 1 MB. Typical use: read_job_file, "
|
|
"edit the content (LLM or human), update_job_file, then "
|
|
"prepare_submit_job with the same path."
|
|
),
|
|
)
|
|
def _update_job_file(req: UpdateJobFileRequest):
|
|
return update_job_file(req.script_path, req.content)
|