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