From 2f57c98aa697f8cd3113f5dd0a9efd6828acf4d4 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 24 Jun 2026 14:40:33 +0800 Subject: [PATCH] feat: add confirm_submit_job (two-step submit flow) --- spark_executor/tools/submit.py | 73 +++++++++++++++++++++++++++++++++- tests/unit/test_submit_tool.py | 53 ++++++++++++++++++++++++ 2 files changed, 124 insertions(+), 2 deletions(-) diff --git a/spark_executor/tools/submit.py b/spark_executor/tools/submit.py index 5e717d9..36c42bf 100644 --- a/spark_executor/tools/submit.py +++ b/spark_executor/tools/submit.py @@ -8,9 +8,17 @@ from datetime import datetime from common.logging import logger 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 run_spark_submit # re-exported for monkeypatch in tests -from spark_executor.models import PendingSubmission +from spark_executor.core.spark_submit import ( + SparkSubmitError, + build_spark_submit_command, + run_spark_submit, +) +from spark_executor.models import Job, PendingSubmission, SubmitResult +import uuid +from datetime import datetime as _datetime def _new_pending_id() -> str: @@ -55,3 +63,64 @@ def prepare_submit_job( "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.""" + 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" + ) + + 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, + ) + logger.info(f"confirm_submit_job pending_id={pending_id} cmd={cmd}") + try: + result = run_spark_submit(cmd) + except SparkSubmitError as exc: + pending.status = "FAILED" + pending.error = str(exc) + pending_store.save(pending) + 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, + ) + ) + + pending.status = "SUBMITTED" + pending.job_id = job_id + pending.application_id = application_id + pending_store.save(pending) + logger.info( + f"confirm_submit_job pending_id={pending_id} job_id={job_id} application_id={application_id}" + ) + return SubmitResult( + job_id=job_id, + application_id=application_id, + tracking_url=tracking_url, + ) diff --git a/tests/unit/test_submit_tool.py b/tests/unit/test_submit_tool.py index 8d50f42..0241bd8 100644 --- a/tests/unit/test_submit_tool.py +++ b/tests/unit/test_submit_tool.py @@ -77,3 +77,56 @@ def test_prepare_snapshots_connection_at_prepare_time(monkeypatch): p2 = submit.pending_store.get(p.pending_id) assert p2.master == "yarn" assert p2.deploy_mode == "cluster" + + +# --- confirm_submit_job --- + +def test_confirm_invokes_spark_submit_and_marks_submitted(monkeypatch): + submit.prepare_submit_job(connection="prod", script_path="/tmp/j.py") + pid = _last_pending_id() + fake_proc = type("P", (), { + "returncode": 0, + "stderr": "tracking URL: http://rm:8088/proxy/application_17400000001/\n", + })() + with patch("spark_executor.tools.submit.run_spark_submit", return_value=fake_proc) as m: + result = submit.confirm_submit_job(pending_id=pid) + cmd = m.call_args.args[0] + assert "yarn" in cmd + assert "cluster" in cmd + assert cmd[-1] == "/tmp/j.py" + assert result.application_id == "application_17400000001" + # pending updated + p = submit.pending_store.get(pid) + assert p.status == "SUBMITTED" + assert p.application_id == "application_17400000001" + assert p.job_id is not None + + +def test_confirm_raises_for_unknown_pending_id(): + with pytest.raises(KeyError, match="missing"): + submit.confirm_submit_job(pending_id="missing") + + +def test_confirm_refuses_non_pending_status(monkeypatch): + submit.prepare_submit_job(connection="prod", script_path="/tmp/j.py") + pid = _last_pending_id() + # Mark it CANCELLED first + p = submit.pending_store.get(pid) + p.status = "CANCELLED" + submit.pending_store.save(p) + with pytest.raises(ValueError, match="CANCELLED"): + submit.confirm_submit_job(pending_id=pid) + + +def test_confirm_marks_failed_on_spark_submit_error(monkeypatch): + from spark_executor.core.spark_submit import SparkSubmitError + submit.prepare_submit_job(connection="prod", script_path="/tmp/j.py") + pid = _last_pending_id() + def _raise(_cmd): + raise SparkSubmitError("boom") + monkeypatch.setattr(submit, "run_spark_submit", _raise) + with pytest.raises(SparkSubmitError): + submit.confirm_submit_job(pending_id=pid) + p = submit.pending_store.get(pid) + assert p.status == "FAILED" + assert "boom" in (p.error or "")