Restore defaults for queue/executor_memory/executor_cores/num_executors in prepare_submit_job, but require the caller to explicitly confirm them. If any defaulted field is omitted, the route returns HTTP 400 listing the defaults and asks the caller to resubmit with explicit values. app_name remains required (no meaningful default). extra_args remains optional. Tests cover rejection of unconfirmed defaults and acceptance of explicitly confirmed defaults.
294 lines
9.7 KiB
Python
294 lines
9.7 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,
|
|
)
|
|
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,
|
|
)
|
|
|
|
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",
|
|
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",
|
|
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",
|
|
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",
|
|
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(
|
|
"/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",
|
|
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."
|
|
),
|
|
)
|
|
def _get_job_status(req: JobIdRequest):
|
|
return get_job_status(req.job_id)
|
|
|
|
|
|
@app.post(
|
|
"/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."
|
|
),
|
|
)
|
|
def _get_job_result(req: JobIdRequest):
|
|
return get_job_result(req.job_id)
|
|
|
|
|
|
@app.post(
|
|
"/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."
|
|
),
|
|
)
|
|
def _get_job_logs(req: GetJobLogsRequest):
|
|
return get_job_logs(req.job_id, tail_chars=req.tail_chars)
|
|
|
|
|
|
@app.post(
|
|
"/kill_job",
|
|
summary="Kill a running job",
|
|
description="PUT state=KILLED to YARN REST API for the job's application_id.",
|
|
)
|
|
def _kill_job(req: JobIdRequest):
|
|
return kill_job(req.job_id)
|
|
|
|
|
|
# --- Connection management tools ---
|
|
|
|
@app.post(
|
|
"/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",
|
|
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",
|
|
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",
|
|
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",
|
|
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",
|
|
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",
|
|
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)
|