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.
240 lines
8.8 KiB
Python
240 lines
8.8 KiB
Python
# coding=utf-8
|
|
"""
|
|
@Time :2026/6/24
|
|
@Author :tao.chen
|
|
"""
|
|
import os
|
|
import secrets
|
|
import uuid
|
|
from datetime import datetime
|
|
|
|
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 generate_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"generate_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 = "4G",
|
|
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
|
|
# generate_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 generate_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 != "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}"
|
|
)
|
|
try:
|
|
result = run_spark_submit(cmd)
|
|
except SparkSubmitError as exc:
|
|
pending.status = "FAILED"
|
|
pending.error = str(exc)
|
|
pending_store.save(pending)
|
|
logger.error(f"confirm_submit_job failed pending_id={pending_id} err={exc}")
|
|
raise
|
|
|
|
application_id, tracking_url = parse_spark_submit_output(result.stderr)
|
|
job_id = uuid.uuid4().hex[:12]
|
|
|
|
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,
|
|
)
|
|
)
|
|
|
|
pending.status = "SUBMITTED"
|
|
pending.job_id = job_id
|
|
pending.application_id = application_id
|
|
pending_store.save(pending)
|
|
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,
|
|
)
|
|
|
|
|
|
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"}
|