Files
mcp-server/tests/unit/test_submit_tool.py
T
Claude f170c3045b feat(submit): confirm defaults instead of removing them
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.
2026-06-26 15:16:43 +08:00

454 lines
14 KiB
Python

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