# coding=utf-8 from pathlib import Path from unittest.mock import patch import pytest from spark_executor.core import connection_store, pending_store from spark_executor.core.pending_store import PendingStore from spark_executor.models import Connection from spark_executor.tools import submit @pytest.fixture(autouse=True) def _fresh(tmp_path: Path, monkeypatch): monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) monkeypatch.setattr(connection_store, "store", connection_store.ConnectionStore()) monkeypatch.setattr(pending_store, "DEFAULT_DATA_DIR", str(tmp_path)) monkeypatch.setattr(pending_store, "store", PendingStore()) submit.conn_store = connection_store.store submit.pending_store = pending_store.store connection_store.store.save(Connection( name="prod", master="yarn", deploy_mode="cluster", yarn_rm_url="http://rm:8088", )) @pytest.fixture def real_script(tmp_path: Path) -> Path: """A real .py file on disk that satisfies prepare_submit_job's path check.""" p = tmp_path / "demo.py" p.write_text("print('hi from test')\n") return p def _last_pending_id() -> str: pending = submit.pending_store.list_all() assert pending, "no pending submission was created" return pending[-1].pending_id # --- prepare_submit_job --- def test_prepare_does_not_invoke_spark_submit(monkeypatch, real_script): called = [] monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: called.append(cmd)) out = submit.prepare_submit_job( connection="prod", script_path=str(real_script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) assert called == [] # never called assert out["status"] == "PENDING" assert out["pending_id"].startswith("p_") def test_prepare_persists_pending_with_snapshot(monkeypatch, real_script): monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) connection_store.store.save( Connection( name="prod", master="yarn", deploy_mode="cluster", yarn_rm_url="http://rm:8088", spark_conf={"spark.sql.shuffle.partitions": "200"}, ) ) submit.prepare_submit_job( connection="prod", script_path=str(real_script), queue="research", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) p = submit.pending_store.get(_last_pending_id()) assert p is not None assert p.connection == "prod" assert p.master == "yarn" assert p.deploy_mode == "cluster" assert p.yarn_rm_url == "http://rm:8088" assert p.spark_conf == {"spark.sql.shuffle.partitions": "200"} assert p.queue == "research" assert p.script_path == str(real_script) def test_prepare_persists_app_name(monkeypatch, real_script): monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) submit.prepare_submit_job( connection="prod", script_path=str(real_script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="my-etl-job", ) p = submit.pending_store.get(_last_pending_id()) assert p is not None assert p.app_name == "my-etl-job" def test_prepare_requires_explicit_app_name(monkeypatch, real_script): monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) with pytest.raises(TypeError, match="app_name"): submit.prepare_submit_job(connection="prod", script_path=str(real_script)) def test_prepare_accepts_restored_defaults(monkeypatch, real_script): monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) out = submit.prepare_submit_job( connection="prod", script_path=str(real_script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) assert out["status"] == "PENDING" p = submit.pending_store.get(_last_pending_id()) assert p.queue == "default" assert p.executor_memory == "4G" assert p.executor_cores == 2 assert p.num_executors == 2 def test_prepare_snapshots_extra_args(monkeypatch, real_script): monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) submit.prepare_submit_job( connection="prod", script_path=str(real_script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", extra_args={"jars": "hdfs:///lib/foo.jar"}, ) pid = _last_pending_id() p = submit.pending_store.get(pid) assert p is not None assert p.extra_args == {"jars": "hdfs:///lib/foo.jar"} fake_proc = type("P", (), { "returncode": 0, "stderr": "tracking URL: http://rm:8088/proxy/application_17400000002/\n", })() with patch("spark_executor.tools.submit.run_spark_submit", return_value=fake_proc) as m: submit.confirm_submit_job(pending_id=pid) cmd = m.call_args.args[0] assert ["--jars", "hdfs:///lib/foo.jar"] in [ cmd[i : i + 2] for i in range(len(cmd) - 1) ] def test_prepare_raises_for_unknown_connection(real_script): with pytest.raises(KeyError, match="missing"): submit.prepare_submit_job( connection="missing", script_path=str(real_script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) # --- script_path validation --- def test_prepare_rejects_nonexistent_script_path(tmp_path): missing = str(tmp_path / "does_not_exist.py") with pytest.raises(ValueError, match="does not exist or is not a file"): submit.prepare_submit_job( connection="prod", script_path=missing, queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) def test_prepare_rejects_directory_as_script_path(tmp_path): """A directory is not a file, even if it exists.""" with pytest.raises(ValueError, match="does not exist or is not a file"): submit.prepare_submit_job( connection="prod", script_path=str(tmp_path), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) def test_prepare_rejects_empty_script_path(): with pytest.raises(ValueError, match="does not exist or is not a file"): submit.prepare_submit_job( connection="prod", script_path="", queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) def test_prepare_error_message_mentions_generate_job_file(tmp_path): """The agent must be told to call generate_job_file first.""" with pytest.raises(ValueError, match="generate_job_file"): submit.prepare_submit_job( connection="prod", script_path=str(tmp_path / "x.py"), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) def test_confirm_rejects_if_script_was_deleted_after_prepare(tmp_path, monkeypatch): """Defense in depth: file present at prepare, gone by confirm -> 400.""" script = tmp_path / "demo.py" script.write_text("print('hi')\n") submit.prepare_submit_job( connection="prod", script_path=str(script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) pid = _last_pending_id() # Simulate the file being deleted (e.g. cleanup job) between prepare and confirm script.unlink() with pytest.raises(ValueError, match="does not exist or is not a file"): submit.confirm_submit_job(pending_id=pid) def test_prepare_snapshots_connection_at_prepare_time(monkeypatch, real_script): """Editing the connection between prepare and confirm must NOT silently retarget.""" monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) submit.prepare_submit_job( connection="prod", script_path=str(real_script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) p = submit.pending_store.get(_last_pending_id()) # Now mutate the connection connection_store.store.save( Connection(name="prod", master="spark://attacker:7077", deploy_mode="client") ) # Snapshot is unchanged 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, real_script): submit.prepare_submit_job( connection="prod", script_path=str(real_script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) 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] == str(real_script) 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 # job carries the connection's yarn_rm_url snapshot assert submit.job_store.get(p.job_id).yarn_rm_url == "http://rm:8088" 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, real_script): submit.prepare_submit_job( connection="prod", script_path=str(real_script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) 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, real_script): from spark_executor.core.spark_submit import SparkSubmitError submit.prepare_submit_job( connection="prod", script_path=str(real_script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) 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 "") # --- list_pending_jobs --- def test_list_pending_jobs_empty(): assert submit.list_pending_jobs() == [] def test_list_pending_jobs_returns_all(monkeypatch, tmp_path): a = tmp_path / "a.py" b = tmp_path / "b.py" a.write_text("a\n"); b.write_text("b\n") monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) submit.prepare_submit_job( connection="prod", script_path=str(a), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) submit.prepare_submit_job( connection="prod", script_path=str(b), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) out = submit.list_pending_jobs() assert {p["script_path"] for p in out} == {str(a), str(b)} assert all(p["status"] == "PENDING" for p in out) # --- get_pending_job --- def test_get_pending_job_returns_dump(monkeypatch, real_script): monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) submit.prepare_submit_job( connection="prod", script_path=str(real_script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) pid = _last_pending_id() out = submit.get_pending_job(pid) assert out["pending_id"] == pid assert out["connection"] == "prod" assert out["status"] == "PENDING" def test_get_pending_job_unknown_raises(): with pytest.raises(KeyError): submit.get_pending_job("missing") # --- cancel_pending_job --- def test_cancel_pending_job_marks_cancelled(monkeypatch, real_script): monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) submit.prepare_submit_job( connection="prod", script_path=str(real_script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) pid = _last_pending_id() out = submit.cancel_pending_job(pid) assert out == {"pending_id": pid, "status": "CANCELLED"} assert submit.pending_store.get(pid).status == "CANCELLED" def test_cancel_pending_job_unknown_raises(): with pytest.raises(KeyError): submit.cancel_pending_job("missing") def test_cancel_pending_job_refuses_submitted(monkeypatch, real_script): submit.prepare_submit_job( connection="prod", script_path=str(real_script), queue="default", executor_memory="4G", executor_cores=2, num_executors=2, app_name="test-app", ) pid = _last_pending_id() p = submit.pending_store.get(pid) p.status = "SUBMITTED" submit.pending_store.save(p) with pytest.raises(ValueError, match="SUBMITTED"): submit.cancel_pending_job(pid)