diff --git a/spark_executor/core/spark_submit.py b/spark_executor/core/spark_submit.py new file mode 100644 index 0000000..559845d --- /dev/null +++ b/spark_executor/core/spark_submit.py @@ -0,0 +1,45 @@ +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +import subprocess + + +class SparkSubmitError(Exception): + """Raised when spark-submit exits with a non-zero return code.""" + + +def build_spark_submit_command( + *, + master: str, + deploy_mode: str, + script_path: str, + queue: str, + executor_memory: str, + executor_cores: int, + num_executors: int, + spark_conf: dict[str, str] | None = None, +) -> list[str]: + cmd = [ + "spark-submit", + "--master", master, + "--deploy-mode", deploy_mode, + "--queue", queue, + "--executor-memory", executor_memory, + "--executor-cores", str(executor_cores), + "--num-executors", str(num_executors), + ] + for key, value in (spark_conf or {}).items(): + cmd.extend(["--conf", f"{key}={value}"]) + cmd.append(script_path) + return cmd + + +def run_spark_submit(cmd: list[str]) -> "subprocess.CompletedProcess[str]": + result = subprocess.run(cmd, capture_output=True, text=True, errors="replace") + if result.returncode != 0: + raise SparkSubmitError( + f"spark-submit failed (rc={result.returncode}): {result.stderr}" + ) + return result diff --git a/tests/unit/test_spark_submit.py b/tests/unit/test_spark_submit.py new file mode 100644 index 0000000..67ed2e6 --- /dev/null +++ b/tests/unit/test_spark_submit.py @@ -0,0 +1,86 @@ +# coding=utf-8 +from unittest.mock import MagicMock, patch + +import pytest + +from spark_executor.core.spark_submit import ( + SparkSubmitError, + build_spark_submit_command, + run_spark_submit, +) + + +def test_build_command_uses_provided_master_and_deploy_mode(): + cmd = build_spark_submit_command( + master="yarn", + deploy_mode="cluster", + script_path="/tmp/jobs/job_001.py", + queue="default", + executor_memory="4G", + executor_cores=2, + num_executors=2, + ) + assert cmd[:2] == ["spark-submit", "--master"] + assert "yarn" in cmd + assert "cluster" in cmd + assert "--queue" in cmd and "default" in cmd + assert "--executor-memory" in cmd and "4G" in cmd + assert "--executor-cores" in cmd and "2" in cmd + assert "--num-executors" in cmd and "2" in cmd + assert cmd[-1] == "/tmp/jobs/job_001.py" + + +def test_build_command_supports_standalone_master(): + cmd = build_spark_submit_command( + master="spark://master:7077", + deploy_mode="client", + script_path="/tmp/j.py", + queue="default", + executor_memory="1G", + executor_cores=1, + num_executors=1, + ) + assert "spark://master:7077" in cmd + assert "client" in cmd + + +def test_build_command_appends_spark_conf_entries(): + cmd = build_spark_submit_command( + master="yarn", + deploy_mode="cluster", + script_path="/tmp/j.py", + queue="default", + executor_memory="4G", + executor_cores=2, + num_executors=2, + spark_conf={"spark.sql.shuffle.partitions": "200", "spark.executor.memoryOverhead": "1G"}, + ) + # every k=v is emitted as --conf k=v + assert "--conf" in cmd + assert "spark.sql.shuffle.partitions=200" in cmd + assert "spark.executor.memoryOverhead=1G" in cmd + # script_path is still last + assert cmd[-1] == "/tmp/j.py" + + +def test_run_spark_submit_returns_completed_process(monkeypatch): + fake = MagicMock() + fake.returncode = 0 + fake.stderr = "Submitted application application_1\n" + with patch("spark_executor.core.spark_submit.subprocess.run", return_value=fake) as m: + result = run_spark_submit(["spark-submit", "/tmp/x.py"]) + assert result is fake + m.assert_called_once() + # capture_output=True and text=True must be set + kwargs = m.call_args.kwargs + assert kwargs["capture_output"] is True + assert kwargs["text"] is True + + +def test_run_spark_submit_raises_on_nonzero_return(): + fake = MagicMock() + fake.returncode = 1 + fake.stderr = "boom" + with patch("spark_executor.core.spark_submit.subprocess.run", return_value=fake): + with pytest.raises(SparkSubmitError, match="spark-submit failed"): + run_spark_submit(["spark-submit", "/tmp/x.py"])