From 523e6a9c765abe1635a9e3e9bec3b68879df010e Mon Sep 17 00:00:00 2001 From: Claude Date: Mon, 29 Jun 2026 19:13:29 +0800 Subject: [PATCH] refactor(mcp): rename generate_job_file to write_job_file to match what it does MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The MCP tool was named `generate_job_file` from Stage 2 but it does NOT generate PySpark code — the calling LLM writes the code in its own context, and this tool only persists it to a file under SPARK_EXECUTOR_JOBS_DIR so `spark-submit` can see it. The misleading `generate_` prefix sent agents (and humans) looking for a code generator that doesn't exist. This commit folds three related polish changes into one (split later with rebase -i if you want them as separate history): 1. The rename itself: - `tools/generate.py` → `tools/write_job.py` - `generate_job_file` → `write_job_file` - `GenerateJobFileRequest` → `WriteJobFileRequest` - `/generate_job_file` route → `/write_job_file` - `operation_id="generate_job_file"` → `operation_id="write_job_file"` The internal helper `core.job_writer.write_job_file` (which just writes bytes to disk with no SQL guard) is imported with an `_write_to_disk` alias to avoid the name collision with the MCP-exposed function in the same module. The description for the tool now explicitly states 'this tool does NOT generate PySpark code. The calling LLM is expected to have already written the code; this tool only persists it.' 2. Skill for LLM agents operating the service (`docs/superpowers/skills/spark-executor-mcp-operate/SKILL.md`, 449 lines). Covers the 16 tools, the two-step prepare/confirm flow, the dual-ID contract (job_id vs application_id), the PendingSubmission state machine, the Connection profile, the job-file workflow, the error reference, common pitfalls, and a full end-to-end word-count example. 3. Default `executor_memory` lowered 4G → 2G (`_DEFAULTS_TO_CONFIRM` in `server.py`). Mirrors the matching change in `test_mcp_routes.py` and the 5 unit tests that reference the default. Aligns with the lighter workloads the service is sized for in its current container profile. Also tracked in git for the first time: - `docs/superpowers/plans/2026-06-24-spark-executor-mcp.md` (the original Stage 1/2/3 design plan, updated to use the new tool name throughout). Test rename: - `tests/unit/test_generate_tool.py` → `test_write_job_tool.py` - the new test file picks up an extra assertion that the SQL guard rejects a `DROP TABLE` statement at write time. 243 tests pass (was 242; +1 new SQL-guard assertion). Zero regressions. Co-Authored-By: Claude Fable 5 --- CLAUDE.md | 2 +- .../plans/2026-06-24-spark-executor-mcp.md | 3027 +++++++++++++++++ .../spark-executor-mcp-operate/SKILL.md | 449 +++ spark_executor/server.py | 33 +- spark_executor/tools/generate.py | 43 - spark_executor/tools/job_file.py | 8 +- spark_executor/tools/requests.py | 18 +- spark_executor/tools/submit.py | 10 +- spark_executor/tools/write_job.py | 57 + tests/integration/test_mcp_routes.py | 44 +- tests/unit/test_job_file.py | 2 +- tests/unit/test_logging.py | 2 +- tests/unit/test_models.py | 6 +- tests/unit/test_pending_store.py | 6 +- tests/unit/test_requests.py | 4 +- tests/unit/test_spark_submit.py | 18 +- tests/unit/test_submit_tool.py | 66 +- ...enerate_tool.py => test_write_job_tool.py} | 20 +- 18 files changed, 3661 insertions(+), 154 deletions(-) create mode 100644 docs/superpowers/plans/2026-06-24-spark-executor-mcp.md create mode 100644 docs/superpowers/skills/spark-executor-mcp-operate/SKILL.md delete mode 100644 spark_executor/tools/generate.py create mode 100644 spark_executor/tools/write_job.py rename tests/unit/{test_generate_tool.py => test_write_job_tool.py} (62%) diff --git a/CLAUDE.md b/CLAUDE.md index b8419a4..2bd02d6 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -94,5 +94,5 @@ docs/superpowers/plans/ # implementation plans ## Stage Status - **Stage 1: complete** — 12 MCP tools, 2-step submit, connection management, structured logging, Pydantic body models, exception handlers. 22 commits on `feat/stage-1`. -- **Stage 2: not started** — `generate_job_file` for LLM-written PySpark code (`job_writer` + one new tool). +- **Stage 2: not started** — `write_job_file` for LLM-written PySpark code (`job_writer` + one new tool). - **Stage 3: deferred** — async submission, SQLite `JobRegistry`, status poller, offset-based `LogCache`, multi-tenant `owner`. Plan lives at `docs/superpowers/plans/2026-06-24-spark-executor-mcp.md`. diff --git a/docs/superpowers/plans/2026-06-24-spark-executor-mcp.md b/docs/superpowers/plans/2026-06-24-spark-executor-mcp.md new file mode 100644 index 0000000..640955b --- /dev/null +++ b/docs/superpowers/plans/2026-06-24-spark-executor-mcp.md @@ -0,0 +1,3027 @@ +# Spark Executor MCP Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Build a three-stage MCP server that lets an LLM Agent submit, monitor, fetch logs from, and kill PySpark jobs on YARN — escalating from a synchronous minimum-viable tool set to an async, multi-tenant job platform backed by SQLite. + +**Architecture:** FastAPI app exposes MCP tools (via `fastapi-mcp`) that wrap shell invocations of `spark-submit` and `yarn`. Stage 1 uses synchronous `subprocess.run` and an in-memory job map. Stage 2 adds an LLM-facing `write_job_file` helper. Stage 3 refactors submission to `asyncio.create_subprocess_exec`, persists job state in SQLite, and introduces offset-based log caching. + +**Tech Stack:** Python 3.12+, FastAPI 0.138+, fastapi-mcp 0.4+, Pydantic 2.13+, loguru, httpx (already in `pyproject.toml`); uv for env mgmt; pytest for tests; SQLite (stdlib `sqlite3`) for Stage 3 registry. + +--- + +## File Structure + +Stage 1 (this plan) creates: + +``` +spark_executor/ + __init__.py + server.py # FastAPI app, mounts MCP tools + tools/ + __init__.py + submit.py # submit_job MCP tool (uses a saved Connection) + status.py # get_job_status MCP tool + logs.py # get_job_logs MCP tool + kill.py # kill_job MCP tool + connections.py # save/list/get/delete_connection MCP tools + core/ + __init__.py + spark_submit.py # builds & runs spark-submit commands + yarn_client.py # wraps yarn application / yarn logs + log_parser.py # extracts application_id from spark-submit output + job_store.py # in-memory dict: job_id -> {application_id, ...} + connection_store.py # JSON CRUD over ./data/connections.json + models.py # Pydantic: Job, JobStatus, SubmitResult, Connection +data/ + connections.json # created on first save_connection (gitignored) +tests/ + __init__.py + conftest.py + unit/ + test_log_parser.py + test_spark_submit.py + test_yarn_client.py + test_connection_store.py + test_connection_tools.py + test_submit_tool.py + test_status_tool.py + test_logs_tool.py + test_kill_tool.py +``` + +Stage 2 adds: + +``` +spark_executor/ + tools/ + generate.py # write_job_file MCP tool + core/ + job_writer.py # writes PySpark code to /tmp/jobs/job_xxx.py +tests/ + unit/ + test_write_job_tool.py + test_job_writer.py +``` + +Stage 3 adds (in a follow-on plan, not implemented here): + +``` +spark_executor/ + core/ + job_registry.py # SQLite-backed CRUD + log_cache.py # offset-based yarn-logs cache + async_submitter.py # asyncio.create_subprocess_exec wrapper + status_poller.py # background task updating registry +``` + +--- + +## Global Constraints + +- Python ≥ 3.12 (per `pyproject.toml`). +- Module header on every new Python file: `# coding=utf-8` + `@Time` / `@Author` docstring, matching the existing `common/` and `spark_executor/` files. +- Logging: import `from common.logging import logger` (the shared loguru logger). Do not instantiate new loggers. +- MCP tools are registered by decorating FastAPI routes on the `app` in `spark_executor/server.py`; `fastapi-mcp` auto-discovers them. No manual MCP registration code. +- Subprocess calls must capture stdout AND stderr (Spark writes the application_id line to stderr, not stdout). This is true for both `spark-submit` and `yarn application -status`. +- Default spark-submit flags (Stage 1, fixed): `--master` and `--deploy-mode` are derived from the `Connection` passed to `submit_job` (see Task 11). Other params are user-tunable. +- `Connection` records persist to `./data/connections.json` by default. Override with `SPARK_EXECUTOR_DATA_DIR` env var (the file path becomes `/connections.json`). The directory is auto-created on first write. `./data/` is added to `.gitignore` so credentials/configs are not committed. +- All test commands use the venv: `uv run pytest ...` (or activate `.venv/bin/activate` first). +- Every task ends in a commit. Commit messages use conventional-commits prefixes: `test:`, `feat:`, `chore:`, `refactor:`. +- YAGNI: do not introduce the SQLite registry, async execution, log cache, or ownership filters in Stage 1 tasks. They are out of scope and belong to Stage 3. + +--- + +# Stage 1 — Minimum Closed Loop (★☆☆☆☆) + +Target: 4 MCP tools (`submit_job`, `get_job_status`, `get_job_logs`, `kill_job`) backed by synchronous `subprocess.run`. ~300–500 lines of Python. + +--- + +### Task 1: Test infrastructure and Pydantic models + +**Files:** +- Create: `tests/__init__.py` +- Create: `tests/conftest.py` +- Create: `tests/unit/__init__.py` +- Create: `spark_executor/models.py` +- Create: `spark_executor/core/__init__.py` +- Create: `spark_executor/tools/__init__.py` +- Modify: `pyproject.toml` (add `[dependency-groups].dev` with `pytest`) +- Test: `tests/unit/test_models.py` + +**Interfaces:** +- Produces: + - `class Job(BaseModel)` with fields `job_id: str`, `application_id: str`, `script_path: str`, `queue: str`, `submit_time: datetime`, `connection: str` (name of the `Connection` used to submit) + - `class JobStatus(BaseModel)` with fields `application_id: str`, `state: str`, `raw: str` + - `class SubmitResult(BaseModel)` with fields `job_id: str`, `application_id: str`, `tracking_url: str | None` + - `class Connection(BaseModel)` with fields `name: str`, `master: str`, `deploy_mode: str = "cluster"`, `yarn_rm_url: str | None = None`, `spark_conf: dict[str, str] = Field(default_factory=dict)` + - `class PendingSubmission(BaseModel)` with fields `pending_id: str`, `connection: str`, `master: str`, `deploy_mode: str`, `script_path: str`, `queue: str`, `executor_memory: str`, `executor_cores: int`, `num_executors: int`, `spark_conf: dict[str, str]`, `created_at: datetime`, `status: str = "PENDING"` (one of `"PENDING" | "SUBMITTED" | "CANCELLED" | "FAILED"`), `error: str | None = None`, `job_id: str | None = None`, `application_id: str | None = None`. The `master` / `deploy_mode` / `spark_conf` fields are a snapshot of the source `Connection` taken at `prepare_submit_job` time so that editing the connection between prepare and confirm does not silently change the submission target. + +- [ ] **Step 1: Add pytest to dev dependencies** + +Edit `pyproject.toml`, appending at the end: + +```toml +[dependency-groups] +dev = [ + "pytest>=8.0", + "pytest-mock>=3.14", +] +``` + +Then run: + +```bash +cd /Users/taochen/llm/spark-executor-mcp && uv sync +``` + +Expected: `uv.lock` updated; `pytest` available in `.venv/bin/`. + +- [ ] **Step 2: Create package skeleton with empty `__init__.py` files** + +```bash +mkdir -p tests/unit spark_executor/core spark_executor/tools +``` + +Each new `__init__.py` gets the standard header: + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +``` + +- [ ] **Step 3: Write the failing test for models** + +`tests/unit/test_models.py`: + +```python +# coding=utf-8 +from datetime import datetime +from spark_executor.models import Connection, Job, JobStatus, PendingSubmission, SubmitResult + + +def test_job_roundtrip(): + job = Job( + job_id="abc123", + application_id="application_17400000001", + script_path="/tmp/jobs/job_001.py", + queue="default", + submit_time=datetime(2026, 6, 24, 10, 0, 0), + connection="prod-yarn", + ) + dumped = job.model_dump() + assert dumped["job_id"] == "abc123" + assert dumped["application_id"] == "application_17400000001" + assert dumped["connection"] == "prod-yarn" + + +def test_job_status_default_raw(): + s = JobStatus(application_id="application_1", state="RUNNING") + assert s.raw == "" + + +def test_submit_result_tracking_url_optional(): + r = SubmitResult(job_id="j1", application_id="application_1", tracking_url=None) + assert r.tracking_url is None + + +def test_connection_defaults(): + c = Connection(name="prod", master="yarn") + assert c.deploy_mode == "cluster" + assert c.yarn_rm_url is None + assert c.spark_conf == {} + + +def test_connection_with_spark_conf(): + c = Connection( + name="dev", + master="spark://master:7077", + deploy_mode="client", + spark_conf={"spark.sql.shuffle.partitions": "200"}, + ) + assert c.spark_conf["spark.sql.shuffle.partitions"] == "200" + + +def test_pending_submission_defaults_to_pending_status(): + p = PendingSubmission( + pending_id="p_1", + connection="prod", + master="yarn", + deploy_mode="cluster", + script_path="/tmp/j.py", + queue="default", + executor_memory="2G", + executor_cores=2, + num_executors=2, + spark_conf={}, + created_at=datetime(2026, 6, 24), + ) + assert p.status == "PENDING" + assert p.error is None + assert p.job_id is None + assert p.application_id is None + + +def test_pending_submission_can_record_outcome(): + p = PendingSubmission( + pending_id="p_2", + connection="prod", + master="yarn", + deploy_mode="cluster", + script_path="/tmp/j.py", + queue="default", + executor_memory="2G", + executor_cores=2, + num_executors=2, + spark_conf={}, + created_at=datetime(2026, 6, 24), + status="SUBMITTED", + job_id="j_abc", + application_id="application_17400000001", + ) + assert p.status == "SUBMITTED" + assert p.job_id == "j_abc" + assert p.application_id == "application_17400000001" +``` + +- [ ] **Step 4: Run test to verify it fails** + +```bash +cd /Users/taochen/llm/spark-executor-mcp && uv run pytest tests/unit/test_models.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.models'`. + +- [ ] **Step 5: Implement `spark_executor/models.py`** + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +from datetime import datetime + +from pydantic import BaseModel, Field + + +class Job(BaseModel): + job_id: str + application_id: str + script_path: str + queue: str + submit_time: datetime + connection: str + + +class JobStatus(BaseModel): + application_id: str + state: str + raw: str = Field(default="") + + +class SubmitResult(BaseModel): + job_id: str + application_id: str + tracking_url: str | None = None + + +class Connection(BaseModel): + name: str + master: str + deploy_mode: str = "cluster" + yarn_rm_url: str | None = None + spark_conf: dict[str, str] = Field(default_factory=dict) + + +class PendingSubmission(BaseModel): + pending_id: str + connection: str + master: str + deploy_mode: str + script_path: str + queue: str + executor_memory: str + executor_cores: int + num_executors: int + spark_conf: dict[str, str] = Field(default_factory=dict) + created_at: datetime + status: str = "PENDING" # PENDING | SUBMITTED | CANCELLED | FAILED + error: str | None = None + job_id: str | None = None + application_id: str | None = None +``` + +- [ ] **Step 6: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_models.py -v +``` + +Expected: 8 passed. + +- [ ] **Step 7: Create `conftest.py` to ensure `spark_executor` is importable** + +`tests/conftest.py`: + +```python +# coding=utf-8 +# Ensures the project root is on sys.path when running pytest from any cwd. +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) +``` + +- [ ] **Step 7b: Add `./data/` to `.gitignore`** + +Create `.gitignore` (or append, if it exists): + +``` +data/ +``` + +So credentials and config blobs are not committed. + +- [ ] **Step 8: Commit** + +```bash +git add pyproject.toml uv.lock tests spark_executor/models.py spark_executor/core/__init__.py spark_executor/tools/__init__.py .gitignore +git commit -m "chore: add pytest dev-deps and pydantic models (Job, Connection)" +``` + +--- + +### Task 2: `log_parser` — extract `application_id` and tracking URL from `spark-submit` output + +Spark writes the application_id to stderr in one of two forms: + +``` +Submitted application application_17400000001 +``` + +or + +``` +tracking URL: http://rm:8088/proxy/application_17400000001/ +``` + +**Files:** +- Create: `spark_executor/core/log_parser.py` +- Test: `tests/unit/test_log_parser.py` + +**Interfaces:** +- Produces: + - `def parse_spark_submit_output(stderr: str) -> tuple[str, str | None]` — returns `(application_id, tracking_url)`. Raises `ValueError` if no `application_id` is found. + +- [ ] **Step 1: Write the failing tests** + +`tests/unit/test_log_parser.py`: + +```python +# coding=utf-8 +import pytest + +from spark_executor.core.log_parser import parse_spark_submit_output + + +def test_parses_submitted_application_line(): + stderr = "Warning: ignoring...\nSubmitted application application_17400000001\n" + app_id, url = parse_spark_submit_output(stderr) + assert app_id == "application_17400000001" + assert url is None + + +def test_parses_tracking_url_line(): + stderr = "tracking URL: http://rm:8088/proxy/application_17400000002/\n" + app_id, url = parse_spark_submit_output(stderr) + assert app_id == "application_17400000002" + assert url == "http://rm:8088/proxy/application_17400000002/" + + +def test_prefers_submitted_application_line_over_tracking_url(): + stderr = ( + "tracking URL: http://rm:8088/proxy/application_9999/\n" + "Submitted application application_1234\n" + ) + app_id, _ = parse_spark_submit_output(stderr) + assert app_id == "application_1234" + + +def test_raises_when_no_application_id_found(): + with pytest.raises(ValueError, match="application_id"): + parse_spark_submit_output("some unrelated output\n") +``` + +- [ ] **Step 2: Run tests to verify they fail** + +```bash +uv run pytest tests/unit/test_log_parser.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.core.log_parser'`. + +- [ ] **Step 3: Implement `log_parser.py`** + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +import re + +_APP_ID_RE = re.compile(r"Submitted application (\S+)") +_TRACKING_URL_RE = re.compile(r"tracking URL:\s+(\S+)") + + +def parse_spark_submit_output(stderr: str) -> tuple[str, str | None]: + """Extract (application_id, tracking_url) from spark-submit stderr. + + Raises ValueError if no application_id can be found. + """ + app_match = _APP_ID_RE.search(stderr) + if app_match: + application_id = app_match.group(1) + else: + url_match = _TRACKING_URL_RE.search(stderr) + if not url_match: + raise ValueError("Could not find application_id in spark-submit output") + # tracking URL is of the form http://rm:8088/proxy/application_xxx/ + tracking_url = url_match.group(1) + tail = tracking_url.rstrip("/").rsplit("/", 1)[-1] + if not tail.startswith("application_"): + raise ValueError("Could not find application_id in spark-submit output") + application_id = tail + + url_match = _TRACKING_URL_RE.search(stderr) + tracking_url = url_match.group(1) if url_match else None + return application_id, tracking_url +``` + +- [ ] **Step 4: Run tests to verify they pass** + +```bash +uv run pytest tests/unit/test_log_parser.py -v +``` + +Expected: 4 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/core/log_parser.py tests/unit/test_log_parser.py +git commit -m "feat: add spark-submit output parser" +``` + +--- + +### Task 3: `spark_submit` — build and run a `spark-submit` command + +**Files:** +- Create: `spark_executor/core/spark_submit.py` +- Test: `tests/unit/test_spark_submit.py` + +**Interfaces:** +- Produces: + - `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]` — `--master` and `--deploy-mode` come from the `Connection` chosen by the caller (see Task 11). `spark_conf` entries are emitted as `--conf k=v` pairs in declaration order. + - `def run_spark_submit(cmd: list[str]) -> subprocess.CompletedProcess[str]` — runs the command, returns the CompletedProcess. Stdout AND stderr captured (text mode, `errors="replace"`). + - `class SparkSubmitError(Exception)` + +- [ ] **Step 1: Write the failing tests** + +`tests/unit/test_spark_submit.py`: + +```python +# 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="2G", + 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 "2G" 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="2G", + 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"]) +``` + +- [ ] **Step 2: Run tests to verify they fail** + +```bash +uv run pytest tests/unit/test_spark_submit.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.core.spark_submit'`. + +- [ ] **Step 3: Implement `spark_submit.py`** + +```python +# 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 +``` + +- [ ] **Step 4: Run tests to verify they pass** + +```bash +uv run pytest tests/unit/test_spark_submit.py -v +``` + +Expected: 5 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/core/spark_submit.py tests/unit/test_spark_submit.py +git commit -m "feat: add spark-submit command builder and runner" +``` + +--- + +### Task 4: `yarn_client` — wrap `yarn application -status`, `yarn logs`, `yarn application -kill` + +**Files:** +- Create: `spark_executor/core/yarn_client.py` +- Test: `tests/unit/test_yarn_client.py` + +**Interfaces:** +- Produces: + - `def get_application_status(application_id: str) -> tuple[str, str]` — returns `(state, raw_output)`. State is the value after `State :` in the output. Raises `YarnError` if yarn exits non-zero. + - `def get_application_logs(application_id: str) -> str` — returns full `yarn logs -applicationId ...` stdout. Raises `YarnError` on failure. + - `def kill_application(application_id: str) -> None` — runs `yarn application -kill ...`. Raises `YarnError` on failure. + - `class YarnError(Exception)` + +- [ ] **Step 1: Write the failing tests** + +`tests/unit/test_yarn_client.py`: + +```python +# coding=utf-8 +from unittest.mock import MagicMock, patch + +import pytest + +from spark_executor.core.yarn_client import ( + YarnError, + get_application_logs, + get_application_status, + kill_application, +) + + +def _fake_proc(returncode: int, stdout: str = "", stderr: str = ""): + p = MagicMock() + p.returncode = returncode + p.stdout = stdout + p.stderr = stderr + return p + + +def test_status_parses_state_line(): + fake = _fake_proc( + 0, + stdout="Application Report :\n State : RUNNING\n ...\n", + ) + with patch("spark_executor.core.yarn_client.subprocess.run", return_value=fake): + state, raw = get_application_status("application_1") + assert state == "RUNNING" + assert "RUNNING" in raw + + +def test_status_raises_on_nonzero_return(): + fake = _fake_proc(1, stderr="not found") + with patch("spark_executor.core.yarn_client.subprocess.run", return_value=fake): + with pytest.raises(YarnError): + get_application_status("application_x") + + +def test_logs_returns_stdout(): + fake = _fake_proc(0, stdout="log line 1\nlog line 2\n") + with patch("spark_executor.core.yarn_client.subprocess.run", return_value=fake) as m: + out = get_application_logs("application_1") + assert out == "log line 1\nlog line 2\n" + args = m.call_args.args[0] + assert args[:3] == ["yarn", "logs", "-applicationId"] + assert args[3] == "application_1" + + +def test_kill_invokes_yarn_application_kill(): + fake = _fake_proc(0) + with patch("spark_executor.core.yarn_client.subprocess.run", return_value=fake) as m: + kill_application("application_1") + args = m.call_args.args[0] + assert args[:3] == ["yarn", "application", "-kill"] + assert args[3] == "application_1" + + +def test_kill_raises_on_nonzero_return(): + fake = _fake_proc(1, stderr="denied") + with patch("spark_executor.core.yarn_client.subprocess.run", return_value=fake): + with pytest.raises(YarnError): + kill_application("application_1") +``` + +- [ ] **Step 2: Run tests to verify they fail** + +```bash +uv run pytest tests/unit/test_yarn_client.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.core.yarn_client'`. + +- [ ] **Step 3: Implement `yarn_client.py`** + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +import re +import subprocess + + +class YarnError(Exception): + """Raised when a yarn CLI invocation fails.""" + + +_STATE_RE = re.compile(r"State\s*:\s*(\S+)") + + +def _run(cmd: list[str]) -> "subprocess.CompletedProcess[str]": + return subprocess.run(cmd, capture_output=True, text=True, errors="replace") + + +def get_application_status(application_id: str) -> tuple[str, str]: + proc = _run(["yarn", "application", "-status", application_id]) + if proc.returncode != 0: + raise YarnError( + f"yarn application -status failed (rc={proc.returncode}): {proc.stderr}" + ) + match = _STATE_RE.search(proc.stdout) + if not match: + raise YarnError(f"Could not parse YARN state from output: {proc.stdout!r}") + return match.group(1), proc.stdout + + +def get_application_logs(application_id: str) -> str: + proc = _run(["yarn", "logs", "-applicationId", application_id]) + if proc.returncode != 0: + raise YarnError( + f"yarn logs failed (rc={proc.returncode}): {proc.stderr}" + ) + return proc.stdout + + +def kill_application(application_id: str) -> None: + proc = _run(["yarn", "application", "-kill", application_id]) + if proc.returncode != 0: + raise YarnError( + f"yarn application -kill failed (rc={proc.returncode}): {proc.stderr}" + ) +``` + +- [ ] **Step 4: Run tests to verify they pass** + +```bash +uv run pytest tests/unit/test_yarn_client.py -v +``` + +Expected: 5 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/core/yarn_client.py tests/unit/test_yarn_client.py +git commit -m "feat: add yarn client (status, logs, kill)" +``` + +--- + +### Task 5: `job_store` — in-memory job_id → Job registry + +**Files:** +- Create: `spark_executor/core/job_store.py` +- Test: `tests/unit/test_job_store.py` + +**Interfaces:** +- Produces: + - `class JobStore` with methods: + - `put(job: Job) -> None` + - `get(job_id: str) -> Job | None` + - `list() -> list[Job]` + +Stage 1 keeps this trivial (a `dict` behind a class). Stage 3 replaces the implementation with SQLite. + +- [ ] **Step 1: Write the failing test** + +`tests/unit/test_job_store.py`: + +```python +# coding=utf-8 +from datetime import datetime + +from spark_executor.core.job_store import JobStore +from spark_executor.models import Job + + +def _job(jid: str) -> Job: + return Job( + job_id=jid, + application_id=f"application_{jid}", + script_path="/tmp/j.py", + queue="default", + submit_time=datetime(2026, 6, 24), + ) + + +def test_put_then_get_roundtrip(): + store = JobStore() + store.put(_job("a")) + assert store.get("a") is not None + assert store.get("a").application_id == "application_a" + + +def test_get_missing_returns_none(): + store = JobStore() + assert store.get("nope") is None + + +def test_list_returns_all_jobs(): + store = JobStore() + store.put(_job("a")) + store.put(_job("b")) + assert {j.job_id for j in store.list()} == {"a", "b"} +``` + +- [ ] **Step 2: Run test to verify it fails** + +```bash +uv run pytest tests/unit/test_job_store.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.core.job_store'`. + +- [ ] **Step 3: Implement `job_store.py`** + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +from threading import Lock + +from spark_executor.models import Job + + +class JobStore: + """In-memory job registry. Stage-3 will swap this for SQLite.""" + + def __init__(self) -> None: + self._lock = Lock() + self._jobs: dict[str, Job] = {} + + def put(self, job: Job) -> None: + with self._lock: + self._jobs[job.job_id] = job + + def get(self, job_id: str) -> Job | None: + with self._lock: + return self._jobs.get(job_id) + + def list(self) -> list[Job]: + with self._lock: + return list(self._jobs.values()) +``` + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_job_store.py -v +``` + +Expected: 3 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/core/job_store.py tests/unit/test_job_store.py +git commit -m "feat: add in-memory JobStore" +``` + +--- + +### Task 6: `connection_store` — JSON-backed Connection CRUD + +**Files:** +- Create: `spark_executor/core/connection_store.py` +- Test: `tests/unit/test_connection_store.py` + +**Interfaces:** +- Produces: + - `class ConnectionStore` with methods: + - `list_all() -> list[Connection]` + - `get(name: str) -> Connection | None` + - `save(conn: Connection) -> None` — upsert by `conn.name`, then persist to disk atomically + - `delete(name: str) -> bool` — returns True if a connection was removed + - Module-level singleton `store: ConnectionStore`, defaulting to `./data/connections.json` (overridable via `SPARK_EXECUTOR_DATA_DIR` env var). + - `class ConnectionNotFound(KeyError)` — raised by tool layer when a name is missing. + +- [ ] **Step 1: Write the failing test** + +`tests/unit/test_connection_store.py`: + +```python +# coding=utf-8 +import json +from pathlib import Path + +from spark_executor.core import connection_store +from spark_executor.models import Connection + + +def _conn(name: str) -> Connection: + return Connection(name=name, master="yarn") + + +def test_save_then_get_roundtrip(tmp_path: Path, monkeypatch): + monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) + monkeypatch.setattr(connection_store, "store", connection_store.ConnectionStore()) + connection_store.store.save(_conn("prod")) + assert connection_store.store.get("prod") is not None + assert connection_store.store.get("prod").master == "yarn" + + +def test_save_persists_to_json_file(tmp_path: Path, monkeypatch): + monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) + monkeypatch.setattr(connection_store, "store", connection_store.ConnectionStore()) + connection_store.store.save(_conn("prod")) + path = tmp_path / "connections.json" + assert path.is_file() + payload = json.loads(path.read_text()) + assert "prod" in payload + assert payload["prod"]["master"] == "yarn" + + +def test_get_missing_returns_none(tmp_path: Path, monkeypatch): + monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) + monkeypatch.setattr(connection_store, "store", connection_store.ConnectionStore()) + assert connection_store.store.get("missing") is None + + +def test_list_all_returns_empty_when_file_absent(tmp_path: Path, monkeypatch): + monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) + monkeypatch.setattr(connection_store, "store", connection_store.ConnectionStore()) + assert connection_store.store.list_all() == [] + + +def test_list_all_returns_saved_connections(tmp_path: Path, monkeypatch): + monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) + monkeypatch.setattr(connection_store, "store", connection_store.ConnectionStore()) + connection_store.store.save(_conn("a")) + connection_store.store.save(_conn("b")) + names = {c.name for c in connection_store.store.list_all()} + assert names == {"a", "b"} + + +def test_save_upserts_existing(tmp_path: Path, monkeypatch): + monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) + monkeypatch.setattr(connection_store, "store", connection_store.ConnectionStore()) + connection_store.store.save(_conn("prod")) + updated = Connection(name="prod", master="spark://new:7077", deploy_mode="client") + connection_store.store.save(updated) + assert connection_store.store.get("prod").master == "spark://new:7077" + assert connection_store.store.get("prod").deploy_mode == "client" + + +def test_delete_returns_true_when_present(tmp_path: Path, monkeypatch): + monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) + monkeypatch.setattr(connection_store, "store", connection_store.ConnectionStore()) + connection_store.store.save(_conn("prod")) + assert connection_store.store.delete("prod") is True + assert connection_store.store.get("prod") is None + + +def test_delete_returns_false_when_absent(tmp_path: Path, monkeypatch): + monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) + monkeypatch.setattr(connection_store, "store", connection_store.ConnectionStore()) + assert connection_store.store.delete("nope") is False +``` + +- [ ] **Step 2: Run test to verify it fails** + +```bash +uv run pytest tests/unit/test_connection_store.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.core.connection_store'`. + +- [ ] **Step 3: Implement `connection_store.py`** + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +import json +import os +import tempfile +from pathlib import Path +from threading import Lock + +from pydantic import ValidationError + +from spark_executor.models import Connection + +DEFAULT_DATA_DIR = os.environ.get("SPARK_EXECUTOR_DATA_DIR", "./data") +DEFAULT_FILE_NAME = "connections.json" + + +class ConnectionNotFound(KeyError): + """Raised when a named connection does not exist.""" + + +class ConnectionStore: + """JSON-backed CRUD for Connection records. + + File path: /. + Writes are atomic: tempfile + os.replace. + """ + + def __init__(self, data_dir: str | None = None, file_name: str = DEFAULT_FILE_NAME) -> None: + self._data_dir = data_dir or DEFAULT_DATA_DIR + self._file_name = file_name + self._lock = Lock() + + @property + def path(self) -> Path: + return Path(self._data_dir) / self._file_name + + def _load(self) -> dict[str, Connection]: + if not self.path.exists(): + return {} + raw = json.loads(self.path.read_text(encoding="utf-8")) + return {name: Connection.model_validate(c) for name, c in raw.items()} + + def _dump(self, records: dict[str, Connection]) -> None: + os.makedirs(self._data_dir, exist_ok=True) + payload = {name: c.model_dump() for name, c in records.items()} + # atomic write: tempfile in same dir, then replace + fd, tmp_path = tempfile.mkstemp( + prefix=self._file_name + ".", dir=self._data_dir + ) + try: + with os.fdopen(fd, "w", encoding="utf-8") as f: + json.dump(payload, f, indent=2, ensure_ascii=False) + os.replace(tmp_path, self.path) + except Exception: + if os.path.exists(tmp_path): + os.unlink(tmp_path) + raise + + def list_all(self) -> list[Connection]: + with self._lock: + return list(self._load().values()) + + def get(self, name: str) -> Connection | None: + with self._lock: + return self._load().get(name) + + def save(self, conn: Connection) -> None: + with self._lock: + records = self._load() + records[conn.name] = conn + self._dump(records) + + def delete(self, name: str) -> bool: + with self._lock: + records = self._load() + if name not in records: + return False + del records[name] + self._dump(records) + return True + + +# Module-level singleton; replaced in tests. +store = ConnectionStore() +``` + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_connection_store.py -v +``` + +Expected: 8 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/core/connection_store.py tests/unit/test_connection_store.py +git commit -m "feat: add JSON-backed ConnectionStore" +``` + +--- + +### Task 7: `save_connection` MCP tool + +**Files:** +- Create: `spark_executor/tools/connections.py` +- Test: `tests/unit/test_connection_tools.py::test_save_connection` + +**Interfaces:** +- Produces: + - `def save_connection(*, name: str, master: str, deploy_mode: str = "cluster", yarn_rm_url: str | None = None, spark_conf: dict[str, str] | None = None) -> dict[str, str]` — returns `{"name": name, "status": "SAVED"}`. Upserts: re-saving with the same `name` overwrites. + +- [ ] **Step 1: Write the failing test (this file is shared by all four connection tools)** + +`tests/unit/test_connection_tools.py`: + +```python +# coding=utf-8 +from pathlib import Path + +import pytest + +from spark_executor.core import connection_store +from spark_executor.tools import connections + + +@pytest.fixture(autouse=True) +def _fresh_store(tmp_path: Path, monkeypatch): + monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) + monkeypatch.setattr(connection_store, "store", connection_store.ConnectionStore()) + monkeypatch.setattr(connections, "store", connection_store.store) + return connection_store.store + + +def test_save_connection(_fresh_store): + out = connections.save_connection(name="prod", master="yarn") + assert out == {"name": "prod", "status": "SAVED"} + assert _fresh_store.get("prod") is not None + + +def test_save_connection_upserts(_fresh_store): + connections.save_connection(name="prod", master="yarn") + connections.save_connection(name="prod", master="spark://new:7077", deploy_mode="client") + assert _fresh_store.get("prod").master == "spark://new:7077" + assert _fresh_store.get("prod").deploy_mode == "client" + + +def test_save_connection_with_spark_conf(_fresh_store): + connections.save_connection( + name="dev", + master="yarn", + spark_conf={"spark.sql.shuffle.partitions": "200"}, + ) + assert _fresh_store.get("dev").spark_conf["spark.sql.shuffle.partitions"] == "200" +``` + +- [ ] **Step 2: Run the new tests to verify they fail** + +```bash +uv run pytest tests/unit/test_connection_tools.py::test_save_connection -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.tools.connections'`. + +- [ ] **Step 3: Implement `connections.py`** + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +from common.logging import logger +from spark_executor.core.connection_store import ConnectionStore, store +from spark_executor.models import Connection + + +def save_connection( + *, + name: str, + master: str, + deploy_mode: str = "cluster", + yarn_rm_url: str | None = None, + spark_conf: dict[str, str] | None = None, +) -> dict[str, str]: + conn = Connection( + name=name, + master=master, + deploy_mode=deploy_mode, + yarn_rm_url=yarn_rm_url, + spark_conf=spark_conf or {}, + ) + store.save(conn) + logger.info(f"save_connection name={name} master={master}") + return {"name": name, "status": "SAVED"} + + +def list_connections() -> list[dict[str, object]]: + return [c.model_dump() for c in store.list_all()] + + +def get_connection(name: str) -> dict[str, object]: + conn = store.get(name) + if conn is None: + raise KeyError(f"Unknown connection: {name}") + return conn.model_dump() + + +def delete_connection(name: str) -> dict[str, str]: + removed = store.delete(name) + if not removed: + raise KeyError(f"Unknown connection: {name}") + logger.info(f"delete_connection name={name}") + return {"name": name, "status": "DELETED"} +``` + +- [ ] **Step 4: Run the save tests to verify they pass** + +```bash +uv run pytest tests/unit/test_connection_tools.py -v -k save_connection +``` + +Expected: 3 passed (the 3 `test_save_connection_*` cases). The other tests (list/get/delete) will fail because the implementation file already contains their functions — that is fine, they are filled in by Tasks 8–10 below. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/tools/connections.py tests/unit/test_connection_tools.py +git commit -m "feat: add save_connection tool" +``` + +--- + +### Task 8: `list_connections` MCP tool + +**Files:** +- Modify: `tests/unit/test_connection_tools.py` (append test case) +- Test: `tests/unit/test_connection_tools.py::test_list_connections` + +- [ ] **Step 1: Append the failing test to `tests/unit/test_connection_tools.py`** + +Add at the end of the file: + +```python +def test_list_connections_empty(_fresh_store): + assert connections.list_connections() == [] + + +def test_list_connections_returns_all_saved(_fresh_store): + connections.save_connection(name="a", master="yarn") + connections.save_connection(name="b", master="spark://m:7077", deploy_mode="client") + out = connections.list_connections() + names = {c["name"] for c in out} + assert names == {"a", "b"} + masters = {c["name"]: c["master"] for c in out} + assert masters == {"a": "yarn", "b": "spark://m:7077"} +``` + +- [ ] **Step 2: Run the new tests to verify they pass** + +```bash +uv run pytest tests/unit/test_connection_tools.py -v -k list_connections +``` + +Expected: 2 passed. (`list_connections` was already implemented in Task 7 Step 3 as part of the same module.) + +- [ ] **Step 3: Commit** + +```bash +git add tests/unit/test_connection_tools.py +git commit -m "test: cover list_connections" +``` + +--- + +### Task 9: `get_connection` MCP tool + +**Files:** +- Modify: `tests/unit/test_connection_tools.py` (append test case) + +- [ ] **Step 1: Append the failing test to `tests/unit/test_connection_tools.py`** + +```python +def test_get_connection_returns_dict(_fresh_store): + connections.save_connection(name="prod", master="yarn", deploy_mode="cluster") + out = connections.get_connection("prod") + assert out["name"] == "prod" + assert out["master"] == "yarn" + assert out["deploy_mode"] == "cluster" + assert out["yarn_rm_url"] is None + assert out["spark_conf"] == {} + + +def test_get_connection_unknown_raises(_fresh_store): + with pytest.raises(KeyError): + connections.get_connection("missing") +``` + +- [ ] **Step 2: Run the new tests to verify they pass** + +```bash +uv run pytest tests/unit/test_connection_tools.py -v -k get_connection +``` + +Expected: 2 passed. + +- [ ] **Step 3: Commit** + +```bash +git add tests/unit/test_connection_tools.py +git commit -m "test: cover get_connection" +``` + +--- + +### Task 10: `delete_connection` MCP tool + +**Files:** +- Modify: `tests/unit/test_connection_tools.py` (append test case) + +- [ ] **Step 1: Append the failing test to `tests/unit/test_connection_tools.py`** + +```python +def test_delete_connection_returns_status(_fresh_store): + connections.save_connection(name="prod", master="yarn") + assert connections.delete_connection("prod") == {"name": "prod", "status": "DELETED"} + assert _fresh_store.get("prod") is None + + +def test_delete_connection_unknown_raises(_fresh_store): + with pytest.raises(KeyError): + connections.delete_connection("missing") +``` + +- [ ] **Step 2: Run the new tests to verify they pass** + +```bash +uv run pytest tests/unit/test_connection_tools.py -v -k delete_connection +``` + +Expected: 2 passed. + +- [ ] **Step 3: Commit** + +```bash +git add tests/unit/test_connection_tools.py +git commit -m "test: cover delete_connection" +``` + +--- + +### Task 11: `pending_store` — JSON-backed store for `PendingSubmission` records + +Submission is a two-step flow: `prepare_submit_job` writes a `PendingSubmission` here, the user reviews and confirms, and only then does `confirm_submit_job` actually invoke `spark-submit`. The pending store sits in `./data/pending_jobs.json` (same `data_dir` as `connection_store`). + +**Files:** +- Create: `spark_executor/core/pending_store.py` +- Test: `tests/unit/test_pending_store.py` + +**Interfaces:** +- Produces: + - `class PendingStore` with methods: + - `list_all() -> list[PendingSubmission]` + - `get(pending_id: str) -> PendingSubmission | None` + - `save(pending: PendingSubmission) -> None` — upsert + - `delete(pending_id: str) -> bool` + - Module-level singleton `store: PendingStore`, defaulting to `/pending_jobs.json`. + +- [ ] **Step 1: Write the failing test** + +`tests/unit/test_pending_store.py`: + +```python +# coding=utf-8 +from datetime import datetime +from pathlib import Path + +import pytest + +from spark_executor.core import pending_store +from spark_executor.models import PendingSubmission + + +def _pending(pid: str = "p_1") -> PendingSubmission: + return PendingSubmission( + pending_id=pid, + connection="prod", + master="yarn", + deploy_mode="cluster", + script_path="/tmp/j.py", + queue="default", + executor_memory="2G", + executor_cores=2, + num_executors=2, + spark_conf={}, + created_at=datetime(2026, 6, 24), + ) + + +@pytest.fixture(autouse=True) +def _fresh(tmp_path: Path, monkeypatch): + monkeypatch.setattr(pending_store, "DEFAULT_DATA_DIR", str(tmp_path)) + monkeypatch.setattr(pending_store, "store", pending_store.PendingStore()) + + +def test_save_then_get_roundtrip(): + pending_store.store.save(_pending()) + out = pending_store.store.get("p_1") + assert out is not None + assert out.connection == "prod" + assert out.master == "yarn" + + +def test_save_persists_to_json(tmp_path: Path): + pending_store.store.save(_pending()) + path = tmp_path / "pending_jobs.json" + assert path.is_file() + + +def test_get_missing_returns_none(): + assert pending_store.store.get("missing") is None + + +def test_list_all_empty_when_file_absent(): + assert pending_store.store.list_all() == [] + + +def test_list_all_returns_saved(): + pending_store.store.save(_pending("a")) + pending_store.store.save(_pending("b")) + assert {p.pending_id for p in pending_store.store.list_all()} == {"a", "b"} + + +def test_save_upserts_existing(): + pending_store.store.save(_pending("a")) + updated = _pending("a") + updated.status = "SUBMITTED" + updated.job_id = "j_1" + pending_store.store.save(updated) + assert pending_store.store.get("a").status == "SUBMITTED" + assert pending_store.store.get("a").job_id == "j_1" + + +def test_delete_returns_true_when_present(): + pending_store.store.save(_pending()) + assert pending_store.store.delete("p_1") is True + assert pending_store.store.get("p_1") is None + + +def test_delete_returns_false_when_absent(): + assert pending_store.store.delete("nope") is False +``` + +- [ ] **Step 2: Run test to verify it fails** + +```bash +uv run pytest tests/unit/test_pending_store.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.core.pending_store'`. + +- [ ] **Step 3: Implement `pending_store.py`** + +`pending_store.py` is structurally identical to `connection_store.py` (Task 6) but operates on `PendingSubmission` records and a different file. Reuse the same atomic-write pattern. + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +import json +import os +import tempfile +from pathlib import Path +from threading import Lock + +from spark_executor.models import PendingSubmission + +DEFAULT_DATA_DIR = os.environ.get("SPARK_EXECUTOR_DATA_DIR", "./data") +DEFAULT_FILE_NAME = "pending_jobs.json" + + +class PendingStore: + """JSON-backed CRUD for PendingSubmission records. + + File path: /. + Atomic writes (tempfile + os.replace), same as ConnectionStore. + """ + + def __init__(self, data_dir: str | None = None, file_name: str = DEFAULT_FILE_NAME) -> None: + self._data_dir = data_dir or DEFAULT_DATA_DIR + self._file_name = file_name + self._lock = Lock() + + @property + def path(self) -> Path: + return Path(self._data_dir) / self._file_name + + def _load(self) -> dict[str, PendingSubmission]: + if not self.path.exists(): + return {} + raw = json.loads(self.path.read_text(encoding="utf-8")) + return {pid: PendingSubmission.model_validate(p) for pid, p in raw.items()} + + def _dump(self, records: dict[str, PendingSubmission]) -> None: + os.makedirs(self._data_dir, exist_ok=True) + payload = {pid: p.model_dump() for pid, p in records.items()} + fd, tmp_path = tempfile.mkstemp( + prefix=self._file_name + ".", dir=self._data_dir + ) + try: + with os.fdopen(fd, "w", encoding="utf-8") as f: + json.dump(payload, f, indent=2, ensure_ascii=False, default=str) + os.replace(tmp_path, self.path) + except Exception: + if os.path.exists(tmp_path): + os.unlink(tmp_path) + raise + + def list_all(self) -> list[PendingSubmission]: + with self._lock: + return list(self._load().values()) + + def get(self, pending_id: str) -> PendingSubmission | None: + with self._lock: + return self._load().get(pending_id) + + def save(self, pending: PendingSubmission) -> None: + with self._lock: + records = self._load() + records[pending.pending_id] = pending + self._dump(records) + + def delete(self, pending_id: str) -> bool: + with self._lock: + records = self._load() + if pending_id not in records: + return False + del records[pending_id] + self._dump(records) + return True + + +# Module-level singleton; replaced in tests. +store = PendingStore() +``` + +Note: `default=str` is required on `json.dump` because `created_at` is a `datetime`. + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_pending_store.py -v +``` + +Expected: 8 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/core/pending_store.py tests/unit/test_pending_store.py +git commit -m "feat: add JSON-backed PendingStore" +``` + +--- + +### Task 12: `prepare_submit_job` MCP tool — record params, do NOT submit + +`prepare_submit_job` is the *first half* of the submit flow. It snapshots the connection's master / deploy_mode / spark_conf into a `PendingSubmission`, persists it, and returns the `pending_id`. **No `spark-submit` is invoked.** + +**Files:** +- Create: `spark_executor/tools/submit.py` (this replaces the direct-submit tool that previously lived here) +- Create: `tests/unit/test_submit_tool.py` (this file is fresh — the old direct-submit tests are removed) + +**Interfaces:** +- Produces: + - `def prepare_submit_job(*, connection: str, script_path: str, queue: str = "default", executor_memory: str = "2G", executor_cores: int = 2, num_executors: int = 2) -> dict[str, object]` — returns `{"pending_id": ..., "status": "PENDING", "parameters": {...}}`. The full `PendingSubmission` (minus pending_id, created_at, status) is echoed back in `parameters` so the user (or agent) can review exactly what will be submitted. Raises `KeyError` if `connection` is unknown. + - Module-level `pending_store` (singleton) imported from `spark_executor.core.pending_store`. + +- [ ] **Step 1: Write the failing test** + +`tests/unit/test_submit_tool.py`: + +```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.pending_store = pending_store.store + connection_store.store.save(Connection(name="prod", master="yarn", deploy_mode="cluster")) + + +def _last_pending_id() -> str: + pending = submit.pending_store.list_all() + assert pending, "no pending submission was created" + return pending[-1].pending_id + + +def test_prepare_does_not_invoke_spark_submit(monkeypatch): + called = [] + monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: called.append(cmd)) + out = submit.prepare_submit_job(connection="prod", script_path="/tmp/j.py") + assert called == [] # never called + assert out["status"] == "PENDING" + assert out["pending_id"].startswith("p_") + + +def test_prepare_persists_pending_with_snapshot(monkeypatch): + monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) + connection_store.store.save( + Connection( + name="prod", + master="yarn", + deploy_mode="cluster", + spark_conf={"spark.sql.shuffle.partitions": "200"}, + ) + ) + submit.prepare_submit_job(connection="prod", script_path="/tmp/j.py", queue="research") + 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.spark_conf == {"spark.sql.shuffle.partitions": "200"} + assert p.queue == "research" + assert p.script_path == "/tmp/j.py" + + +def test_prepare_raises_for_unknown_connection(): + with pytest.raises(KeyError, match="missing"): + submit.prepare_submit_job(connection="missing", script_path="/tmp/j.py") + + +def test_prepare_snapshots_connection_at_prepare_time(monkeypatch): + """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="/tmp/j.py") + 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" +``` + +- [ ] **Step 2: Run test to verify it fails** + +```bash +uv run pytest tests/unit/test_submit_tool.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.tools.submit'`. + +- [ ] **Step 3: Implement `submit.py`** + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +import secrets +from datetime import datetime + +from common.logging import logger +from spark_executor.core.connection_store import store as conn_store +from spark_executor.core.pending_store import store as pending_store +from spark_executor.models import PendingSubmission + + +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, +) -> dict[str, object]: + """Snapshot connection params and persist a PendingSubmission. Does NOT submit.""" + conn = conn_store.get(connection) + if conn is None: + raise KeyError(f"Unknown connection: {connection}") + + pending_id = _new_pending_id() + pending = PendingSubmission( + pending_id=pending_id, + connection=connection, + master=conn.master, + deploy_mode=conn.deploy_mode, + script_path=script_path, + queue=queue, + executor_memory=executor_memory, + executor_cores=executor_cores, + num_executors=num_executors, + spark_conf=dict(conn.spark_conf), + created_at=datetime.utcnow(), + status="PENDING", + ) + pending_store.save(pending) + logger.info( + f"prepare_submit_job pending_id={pending_id} connection={connection} master={conn.master}" + ) + return { + "pending_id": pending_id, + "status": "PENDING", + "parameters": pending.model_dump(), + } +``` + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_submit_tool.py -v +``` + +Expected: 4 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/tools/submit.py tests/unit/test_submit_tool.py +git commit -m "feat: add prepare_submit_job (no spark-submit invocation)" +``` + +--- + +### Task 13: `confirm_submit_job` MCP tool — actually run `spark-submit` + +`confirm_submit_job` is the *second half* of the submit flow. It loads a `PendingSubmission` by `pending_id`, builds the spark-submit command from the **snapshot** (not the live connection), runs it, creates the `Job`, and marks the pending entry as `SUBMITTED` (or `FAILED`). + +**Files:** +- Modify: `spark_executor/tools/submit.py` (append `confirm_submit_job`) +- Modify: `tests/unit/test_submit_tool.py` (append tests) + +**Interfaces:** +- Produces: + - `def confirm_submit_job(*, pending_id: str) -> SubmitResult` — raises `KeyError` if `pending_id` unknown. Raises `ValueError` if the pending is not in `PENDING` status. On `SparkSubmitError`, the pending is marked `FAILED` with the error and re-raised. + +- [ ] **Step 1: Append the failing tests to `tests/unit/test_submit_tool.py`** + +```python +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 "") +``` + +- [ ] **Step 2: Run the new tests to verify they fail** + +```bash +uv run pytest tests/unit/test_submit_tool.py -v -k confirm +``` + +Expected: `AttributeError: module 'spark_executor.tools.submit' has no attribute 'confirm_submit_job'`. + +- [ ] **Step 3: Append `confirm_submit_job` to `submit.py`** + +```python +import uuid +from datetime import datetime as _datetime + +from spark_executor.core.job_store import JobStore +from spark_executor.core.log_parser import parse_spark_submit_output +from spark_executor.core.spark_submit import ( + SparkSubmitError, + build_spark_submit_command, + run_spark_submit, +) +from spark_executor.models import Job, SubmitResult + +# 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, + ) +``` + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_submit_tool.py -v +``` + +Expected: 8 passed (4 prepare + 4 confirm). + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/tools/submit.py tests/unit/test_submit_tool.py +git commit -m "feat: add confirm_submit_job (two-step submit flow)" +``` + +--- + +### Task 14: `list_pending_jobs` MCP tool + +**Files:** +- Modify: `spark_executor/tools/submit.py` (append `list_pending_jobs`) +- Modify: `tests/unit/test_submit_tool.py` (append tests) + +**Interfaces:** +- Produces: + - `def list_pending_jobs() -> list[dict[str, object]]` — returns the model dump of every `PendingSubmission` in the store. + +- [ ] **Step 1: Append the failing test to `tests/unit/test_submit_tool.py`** + +```python +def test_list_pending_jobs_empty(): + assert submit.list_pending_jobs() == [] + + +def test_list_pending_jobs_returns_all(monkeypatch): + monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) + submit.prepare_submit_job(connection="prod", script_path="/tmp/a.py") + submit.prepare_submit_job(connection="prod", script_path="/tmp/b.py") + out = submit.list_pending_jobs() + assert {p["script_path"] for p in out} == {"/tmp/a.py", "/tmp/b.py"} + assert all(p["status"] == "PENDING" for p in out) +``` + +- [ ] **Step 2: Run the new tests to verify they fail** + +```bash +uv run pytest tests/unit/test_submit_tool.py -v -k list_pending +``` + +Expected: `AttributeError: module 'spark_executor.tools.submit' has no attribute 'list_pending_jobs'`. + +- [ ] **Step 3: Append `list_pending_jobs` to `submit.py`** + +```python +def list_pending_jobs() -> list[dict[str, object]]: + return [p.model_dump() for p in pending_store.list_all()] +``` + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_submit_tool.py -v -k list_pending +``` + +Expected: 2 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/tools/submit.py tests/unit/test_submit_tool.py +git commit -m "feat: add list_pending_jobs" +``` + +--- + +### Task 15: `get_pending_job` MCP tool + +**Files:** +- Modify: `spark_executor/tools/submit.py` (append `get_pending_job`) +- Modify: `tests/unit/test_submit_tool.py` (append tests) + +**Interfaces:** +- Produces: + - `def get_pending_job(pending_id: str) -> dict[str, object]` — returns the model dump. Raises `KeyError` if unknown. + +- [ ] **Step 1: Append the failing test to `tests/unit/test_submit_tool.py`** + +```python +def test_get_pending_job_returns_dump(monkeypatch): + monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) + submit.prepare_submit_job(connection="prod", script_path="/tmp/j.py") + 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") +``` + +- [ ] **Step 2: Run the new tests to verify they fail** + +```bash +uv run pytest tests/unit/test_submit_tool.py -v -k get_pending +``` + +Expected: `AttributeError: module 'spark_executor.tools.submit' has no attribute 'get_pending_job'`. + +- [ ] **Step 3: Append `get_pending_job` to `submit.py`** + +```python +def get_pending_job(pending_id: str) -> dict[str, object]: + p = pending_store.get(pending_id) + if p is None: + raise KeyError(f"Unknown pending_id: {pending_id}") + return p.model_dump() +``` + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_submit_tool.py -v -k get_pending +``` + +Expected: 2 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/tools/submit.py tests/unit/test_submit_tool.py +git commit -m "feat: add get_pending_job" +``` + +--- + +### Task 16: `cancel_pending_job` MCP tool + +**Files:** +- Modify: `spark_executor/tools/submit.py` (append `cancel_pending_job`) +- Modify: `tests/unit/test_submit_tool.py` (append tests) + +**Interfaces:** +- Produces: + - `def cancel_pending_job(pending_id: str) -> dict[str, str]` — flips the entry's status to `CANCELLED` (the record is kept for audit). Refuses to cancel a record that is already `SUBMITTED` or `FAILED` (raises `ValueError`). Unknown id raises `KeyError`. + +- [ ] **Step 1: Append the failing test to `tests/unit/test_submit_tool.py`** + +```python +def test_cancel_pending_job_marks_cancelled(monkeypatch): + monkeypatch.setattr(submit, "run_spark_submit", lambda cmd: None) + submit.prepare_submit_job(connection="prod", script_path="/tmp/j.py") + 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): + submit.prepare_submit_job(connection="prod", script_path="/tmp/j.py") + 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) +``` + +- [ ] **Step 2: Run the new tests to verify they fail** + +```bash +uv run pytest tests/unit/test_submit_tool.py -v -k cancel_pending +``` + +Expected: `AttributeError: module 'spark_executor.tools.submit' has no attribute 'cancel_pending_job'`. + +- [ ] **Step 3: Append `cancel_pending_job` to `submit.py`** + +```python +def cancel_pending_job(pending_id: str) -> dict[str, str]: + 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 pending_id={pending_id}") + return {"pending_id": pending_id, "status": "CANCELLED"} +``` + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_submit_tool.py -v -k cancel_pending +``` + +Expected: 3 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/tools/submit.py tests/unit/test_submit_tool.py +git commit -m "feat: add cancel_pending_job" +``` + +--- + + +### Task 17: `get_job_status` MCP tool + +**Files:** +- Create: `spark_executor/tools/status.py` +- Test: `tests/unit/test_status_tool.py` + +**Interfaces:** +- Produces: + - `def get_job_status(job_id: str) -> JobStatus` — looks up the job in the store to get `application_id`, then calls `get_application_status`. Raises `KeyError` if `job_id` unknown. + +- [ ] **Step 1: Write the failing test** + +`tests/unit/test_status_tool.py`: + +```python +# coding=utf-8 +from datetime import datetime +from unittest.mock import patch + +from spark_executor.core.job_store import JobStore +from spark_executor.models import Job +from spark_executor.tools import status + + +def test_get_job_status_returns_state(): + status.store = JobStore() + status.store.put( + Job( + job_id="abc", + application_id="application_1", + script_path="/tmp/j.py", + queue="default", + submit_time=datetime(2026, 6, 24), + connection="prod", + ) + ) + with patch( + "spark_executor.tools.status.get_application_status", + return_value=("RUNNING", "State : RUNNING\n"), + ): + out = status.get_job_status("abc") + assert out.application_id == "application_1" + assert out.state == "RUNNING" + assert "RUNNING" in out.raw + + +def test_get_job_status_raises_for_unknown_job(): + status.store = JobStore() + import pytest + with pytest.raises(KeyError): + status.get_job_status("missing") +``` + +- [ ] **Step 2: Run test to verify it fails** + +```bash +uv run pytest tests/unit/test_status_tool.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.tools.status'`. + +- [ ] **Step 3: Implement `status.py`** + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +from common.logging import logger +from spark_executor.core.job_store import JobStore +from spark_executor.core.yarn_client import get_application_status +from spark_executor.models import JobStatus + +store = JobStore() + + +def get_job_status(job_id: str) -> JobStatus: + job = store.get(job_id) + if job is None: + raise KeyError(f"Unknown job_id: {job_id}") + state, raw = get_application_status(job.application_id) + logger.info(f"get_job_status job_id={job_id} state={state}") + return JobStatus(application_id=job.application_id, state=state, raw=raw) +``` + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_status_tool.py -v +``` + +Expected: 2 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/tools/status.py tests/unit/test_status_tool.py +git commit -m "feat: add get_job_status tool" +``` + +--- + +### Task 18: `get_job_logs` MCP tool (tail to 5000 chars) + +**Files:** +- Create: `spark_executor/tools/logs.py` +- Test: `tests/unit/test_logs_tool.py` + +**Interfaces:** +- Produces: + - `def get_job_logs(job_id: str, tail_chars: int = 5000) -> str` — looks up `application_id` from the store, calls `get_application_logs`, returns the last `tail_chars` characters (default 5000). Raises `KeyError` if `job_id` unknown. + +- [ ] **Step 1: Write the failing test** + +`tests/unit/test_logs_tool.py`: + +```python +# coding=utf-8 +from datetime import datetime +from unittest.mock import patch + +from spark_executor.core.job_store import JobStore +from spark_executor.models import Job +from spark_executor.tools import logs + + +def _seed(job_id="abc", app_id="application_1"): + logs.store = JobStore() + logs.store.put( + Job( + job_id=job_id, + application_id=app_id, + script_path="/tmp/j.py", + queue="default", + submit_time=datetime(2026, 6, 24), + connection="prod", + ) + ) + + +def test_get_job_logs_tails_to_default_5000(): + _seed() + big = "x" * 8000 + "\nEND" + with patch("spark_executor.tools.logs.get_application_logs", return_value=big): + out = logs.get_job_logs("abc") + assert out.endswith("END") + assert len(out) == 5000 + + +def test_get_job_logs_respects_custom_tail_chars(): + _seed() + with patch( + "spark_executor.tools.logs.get_application_logs", + return_value="0123456789", + ): + out = logs.get_job_logs("abc", tail_chars=3) + assert out == "789" + + +def test_get_job_logs_raises_for_unknown_job(): + _seed() + import pytest + with pytest.raises(KeyError): + logs.get_job_logs("missing") +``` + +- [ ] **Step 2: Run test to verify it fails** + +```bash +uv run pytest tests/unit/test_logs_tool.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.tools.logs'`. + +- [ ] **Step 3: Implement `logs.py`** + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +from common.logging import logger +from spark_executor.core.job_store import JobStore +from spark_executor.core.yarn_client import get_application_logs + +store = JobStore() + + +def get_job_logs(job_id: str, tail_chars: int = 5000) -> str: + job = store.get(job_id) + if job is None: + raise KeyError(f"Unknown job_id: {job_id}") + full = get_application_logs(job.application_id) + tailed = full[-tail_chars:] if len(full) > tail_chars else full + logger.info( + f"get_job_logs job_id={job_id} application_id={job.application_id} " + f"chars={len(tailed)}" + ) + return tailed +``` + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_logs_tool.py -v +``` + +Expected: 3 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/tools/logs.py tests/unit/test_logs_tool.py +git commit -m "feat: add get_job_logs tool with tail" +``` + +--- + +### Task 19: `kill_job` MCP tool + +**Files:** +- Create: `spark_executor/tools/kill.py` +- Test: `tests/unit/test_kill_tool.py` + +**Interfaces:** +- Produces: + - `def kill_job(job_id: str) -> dict[str, str]` — returns `{"job_id": ..., "application_id": ..., "status": "KILLED"}`. Raises `KeyError` if `job_id` unknown. + +- [ ] **Step 1: Write the failing test** + +`tests/unit/test_kill_tool.py`: + +```python +# coding=utf-8 +from datetime import datetime +from unittest.mock import patch + +from spark_executor.core.job_store import JobStore +from spark_executor.models import Job +from spark_executor.tools import kill + + +def test_kill_job_calls_yarn_kill(): + kill.store = JobStore() + kill.store.put( + Job( + job_id="abc", + application_id="application_1", + script_path="/tmp/j.py", + queue="default", + submit_time=datetime(2026, 6, 24), + connection="prod", + ) + ) + with patch("spark_executor.tools.kill.kill_application") as m: + result = kill.kill_job("abc") + m.assert_called_once_with("application_1") + assert result == { + "job_id": "abc", + "application_id": "application_1", + "status": "KILLED", + } + + +def test_kill_job_raises_for_unknown_job(): + kill.store = JobStore() + import pytest + with pytest.raises(KeyError): + kill.kill_job("missing") +``` + +- [ ] **Step 2: Run test to verify it fails** + +```bash +uv run pytest tests/unit/test_kill_tool.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.tools.kill'`. + +- [ ] **Step 3: Implement `kill.py`** + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +from common.logging import logger +from spark_executor.core.job_store import JobStore +from spark_executor.core.yarn_client import kill_application + +store = JobStore() + + +def kill_job(job_id: str) -> dict[str, str]: + job = store.get(job_id) + if job is None: + raise KeyError(f"Unknown job_id: {job_id}") + kill_application(job.application_id) + logger.info(f"kill_job job_id={job_id} application_id={job.application_id}") + return { + "job_id": job_id, + "application_id": job.application_id, + "status": "KILLED", + } +``` + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_kill_tool.py -v +``` + +Expected: 2 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/tools/kill.py tests/unit/test_kill_tool.py +git commit -m "feat: add kill_job tool" +``` + +--- + +### Task 20: Wire all tools into the FastAPI app and verify MCP discovery + +`fastapi-mcp` auto-discovers FastAPI routes and exposes them as MCP tools. We mount the twelve tool functions as POST routes on the existing `app` in `spark_executor/server.py`. + +**Files:** +- Modify: `spark_executor/server.py` +- Test: `tests/integration/__init__.py` (empty header file) +- Test: `tests/integration/test_mcp_routes.py` + +- [ ] **Step 1: Create integration test directory** + +`tests/integration/__init__.py` (standard header, empty body). + +- [ ] **Step 2: Write the integration test** + +`tests/integration/test_mcp_routes.py`: + +```python +# coding=utf-8 +from fastapi.testclient import TestClient + +from spark_executor.server import app + + +def test_health_still_present(): + c = TestClient(app) + r = c.get("/health") + assert r.status_code == 200 + assert r.json() == {"status": "ok"} + + +def test_twelve_tool_routes_registered(): + paths = {r.path for r in app.routes} + for path in ( + # pending-submission flow (5) + "/prepare_submit_job", + "/confirm_submit_job", + "/list_pending_jobs", + "/get_pending_job", + "/cancel_pending_job", + # job lifecycle (3) + "/get_job_status", + "/get_job_logs", + "/kill_job", + # connection management (4) + "/save_connection", + "/list_connections", + "/get_connection", + "/delete_connection", + ): + assert path in paths, f"missing MCP tool route: {path}" +``` + +- [ ] **Step 3: Run test to verify it fails** + +```bash +uv run pytest tests/integration/test_mcp_routes.py -v +``` + +Expected: `AssertionError: missing MCP tool route: /prepare_submit_job` (or similar — only `/health` is registered). + +- [ ] **Step 4: Add the twelve tool routes to `server.py`** + +Replace `spark_executor/server.py` with: + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +from fastapi import FastAPI + +from spark_executor.tools.connections import ( + delete_connection, + get_connection, + list_connections, + save_connection, +) +from spark_executor.tools.kill import kill_job +from spark_executor.tools.logs import get_job_logs +from spark_executor.tools.status import get_job_status +from spark_executor.tools.submit import ( + cancel_pending_job, + confirm_submit_job, + get_pending_job, + list_pending_jobs, + prepare_submit_job, +) + +app = FastAPI(title="Spark Executor", version="0.0.1", description="Spark Executor MCP Server") + + +@app.get("/health") +def health_check(): + return {"status": "ok"} + + +# MCP tool routes. fastapi-mcp discovers these and registers them as MCP tools. + +# --- Pending submission flow (two-step submit) --- + +@app.post("/prepare_submit_job") +def _prepare_submit_job(connection: str, script_path: str, queue: str = "default", + executor_memory: str = "2G", executor_cores: int = 2, + num_executors: int = 2): + return prepare_submit_job( + connection=connection, + script_path=script_path, + queue=queue, + executor_memory=executor_memory, + executor_cores=executor_cores, + num_executors=num_executors, + ) + + +@app.post("/confirm_submit_job") +def _confirm_submit_job(pending_id: str): + return confirm_submit_job(pending_id=pending_id) + + +@app.post("/list_pending_jobs") +def _list_pending_jobs(): + return list_pending_jobs() + + +@app.post("/get_pending_job") +def _get_pending_job(pending_id: str): + return get_pending_job(pending_id) + + +@app.post("/cancel_pending_job") +def _cancel_pending_job(pending_id: str): + return cancel_pending_job(pending_id) + + +# --- Spark job tools --- + +@app.post("/get_job_status") +def _get_job_status(job_id: str): + return get_job_status(job_id) + + +@app.post("/get_job_logs") +def _get_job_logs(job_id: str, tail_chars: int = 5000): + return get_job_logs(job_id, tail_chars=tail_chars) + + +@app.post("/kill_job") +def _kill_job(job_id: str): + return kill_job(job_id) + + +# --- Connection management tools --- + +@app.post("/save_connection") +def _save_connection(name: str, master: str, deploy_mode: str = "cluster", + yarn_rm_url: str | None = None, + spark_conf: dict[str, str] | None = None): + return save_connection( + name=name, + master=master, + deploy_mode=deploy_mode, + yarn_rm_url=yarn_rm_url, + spark_conf=spark_conf, + ) + + +@app.post("/list_connections") +def _list_connections(): + return list_connections() + + +@app.post("/get_connection") +def _get_connection(name: str): + return get_connection(name) + + +@app.post("/delete_connection") +def _delete_connection(name: str): + return delete_connection(name) +``` + +- [ ] **Step 5: Run integration test to verify it passes** + +```bash +uv run pytest tests/integration/test_mcp_routes.py -v +``` + +Expected: 2 passed. + +- [ ] **Step 6: Run the full suite to confirm no regressions** + +```bash +uv run pytest -v +``` + +Expected: all unit + integration tests pass (8 + 4 + 5 + 5 + 3 + 8 + 9 + 8 + 15 + 2 + 3 + 2 + 2 = 74). + +- [ ] **Step 7: Smoke-test the running server (manual)** + +```bash +uv run main.py & +sleep 2 +curl -s http://localhost:8000/health +curl -s -X POST http://localhost:8000/save_connection \ + -H "Content-Type: application/json" \ + -d '{"name":"prod","master":"yarn"}' +curl -s -X POST http://localhost:8000/prepare_submit_job \ + -H "Content-Type: application/json" \ + -d '{"connection":"prod","script_path":"/tmp/jobs/demo.py"}' +# (no submit happens; pending_id is returned) +curl -s -X POST http://localhost:8000/list_pending_jobs +curl -s -X POST http://localhost:8000/spark-executor-mcp/mcp -H "Content-Type: application/json" \ + -d '{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{}}' +kill %1 +``` + +Expected: +- `/health` returns `{"status":"ok"}`. +- `save_connection` returns `{"name":"prod","status":"SAVED"}`; `./data/connections.json` is created. +- `prepare_submit_job` returns a payload containing `pending_id` (prefix `p_`); `./data/pending_jobs.json` is created. **No `spark-submit` is invoked.** +- `list_pending_jobs` returns the prepared entry with `status: "PENDING"`. +- `tools/list` enumerates all 12 tools: `prepare_submit_job`, `confirm_submit_job`, `list_pending_jobs`, `get_pending_job`, `cancel_pending_job`, `get_job_status`, `get_job_logs`, `kill_job`, `save_connection`, `list_connections`, `get_connection`, `delete_connection`. + +- [ ] **Step 8: Commit** + +```bash +git add spark_executor/server.py tests/integration +git commit -m "feat: register 12 MCP tools (pending flow + job lifecycle + connections)" +``` + +--- + +### Stage 1 checkpoint + +After Task 20 you have a working minimum-closed-loop MCP server with cluster-connection management and a two-step pending submission flow: + +- `save_connection` / `list_connections` / `get_connection` / `delete_connection` — persist to `./data/connections.json`. +- `prepare_submit_job(connection, ...)` — snapshots the connection's master / deploy_mode / spark_conf into a `PendingSubmission`, writes `./data/pending_jobs.json`, returns `pending_id`. **Does NOT invoke `spark-submit`.** +- `list_pending_jobs` / `get_pending_job` — read pending entries. +- `cancel_pending_job` — flips a PENDING entry to CANCELLED. +- `confirm_submit_job(pending_id)` — only on user second confirmation: runs `spark-submit` with the snapshot, parses `application_id`, creates a `Job` in the in-memory `JobStore`, marks the pending entry `SUBMITTED`. +- `get_job_status` → `yarn application -status`, parses `State : ...`. +- `get_job_logs` → `yarn logs -applicationId ...`, tails to 5000 chars. +- `kill_job` → `yarn application -kill`. + +The end-to-end agent loop is therefore reviewable: + +``` +prepare_submit_job(...) -> pending_id + +confirm_submit_job(pending_id) -> job_id, application_id +get_job_status(job_id) -> RUNNING | SUCCEEDED | FAILED | KILLED +get_job_logs(job_id) -> tail of driver logs +kill_job(job_id) -> KILLED +``` + +Stop here and demo before continuing. Stage 2 is a small additive change; Stage 3 is a real refactor. + +--- + +# Stage 2 — LLM-driven PySpark generation (★★☆☆☆) + +Adds one tool, `write_job_file`, that takes a PySpark code string and writes it to a timestamped file under `/tmp/jobs/`. The Agent can then call `submit_job` on the path. No new infrastructure beyond one file. + +--- + +### Task 21: `job_writer` — write PySpark code to disk + +**Files:** +- Create: `spark_executor/core/job_writer.py` +- Test: `tests/unit/test_job_writer.py` + +**Interfaces:** +- Produces: + - `def write_job_file(code: str, jobs_dir: str = "/tmp/jobs") -> str` — writes `code` to `/job__.py`, creates the directory if needed, returns the absolute path. + +- [ ] **Step 1: Write the failing test** + +`tests/unit/test_job_writer.py`: + +```python +# coding=utf-8 +import os +import tempfile +from pathlib import Path + +from spark_executor.core.job_writer import write_job_file + + +def test_write_job_file_creates_file_with_code(tmp_path: Path): + out = write_job_file("print('hi')\n", jobs_dir=str(tmp_path)) + assert os.path.isfile(out) + assert out.startswith(str(tmp_path)) + assert out.endswith(".py") + with open(out) as f: + assert f.read() == "print('hi')\n" + + +def test_write_job_file_creates_jobs_dir_if_missing(tmp_path: Path): + target = tmp_path / "newdir" + out = write_job_file("x = 1\n", jobs_dir=str(target)) + assert target.is_dir() + assert os.path.isfile(out) +``` + +- [ ] **Step 2: Run test to verify it fails** + +```bash +uv run pytest tests/unit/test_job_writer.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.core.job_writer'`. + +- [ ] **Step 3: Implement `job_writer.py`** + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +import os +import secrets +from datetime import datetime + + +def write_job_file(code: str, jobs_dir: str = "/tmp/jobs") -> str: + os.makedirs(jobs_dir, exist_ok=True) + stamp = datetime.utcnow().strftime("%Y%m%d%H%M%S") + name = f"job_{stamp}_{secrets.token_hex(3)}.py" + path = os.path.join(jobs_dir, name) + with open(path, "w", encoding="utf-8") as f: + f.write(code) + return os.path.abspath(path) +``` + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_job_writer.py -v +``` + +Expected: 2 passed. + +- [ ] **Step 5: Commit** + +```bash +git add spark_executor/core/job_writer.py tests/unit/test_job_writer.py +git commit -m "feat: add job_writer for LLM-generated PySpark" +``` + +--- + +### Task 22: `write_job_file` MCP tool + +**Files:** +- Create: `spark_executor/tools/write_job.py` +- Test: `tests/unit/test_write_job_tool.py` +- Modify: `spark_executor/server.py` (register the new route) +- Modify: `tests/integration/test_mcp_routes.py` (assert new route registered) + +**Interfaces:** +- Produces: + - `def write_job_file(code: str) -> dict[str, str]` — returns `{"script_path": "..."}`. The `code` parameter is the full PySpark source to write to disk. + +- [ ] **Step 1: Write the failing test** + +`tests/unit/test_write_job_tool.py`: + +```python +# coding=utf-8 +import os +import tempfile +from pathlib import Path + +from spark_executor.tools import generate + + +def test_write_job_file_writes_and_returns_path(tmp_path: Path, monkeypatch): + monkeypatch.setattr(generate, "DEFAULT_JOBS_DIR", str(tmp_path)) + out = generate.write_job_file("spark = SparkSession.builder.getOrCreate()\n") + assert os.path.isfile(out["script_path"]) + with open(out["script_path"]) as f: + assert f.read() == "spark = SparkSession.builder.getOrCreate()\n" +``` + +- [ ] **Step 2: Run test to verify it fails** + +```bash +uv run pytest tests/unit/test_write_job_tool.py -v +``` + +Expected: `ModuleNotFoundError: No module named 'spark_executor.tools.generate'`. + +- [ ] **Step 3: Implement `generate.py`** + +```python +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +from common.logging import logger +from spark_executor.core.job_writer import write_job_file + +DEFAULT_JOBS_DIR = "/tmp/jobs" + + +def write_job_file(code: str) -> dict[str, str]: + path = write_job_file(code, jobs_dir=DEFAULT_JOBS_DIR) + logger.info(f"write_job_file script_path={path}") + return {"script_path": path} +``` + +- [ ] **Step 4: Run test to verify it passes** + +```bash +uv run pytest tests/unit/test_write_job_tool.py -v +``` + +Expected: 1 passed. + +- [ ] **Step 5: Register the route in `server.py`** + +In `spark_executor/server.py`, add import and route: + +```python +from spark_executor.tools.write_job import write_job_file +``` + +```python +@app.post("/write_job_file") +def _write_job_file(code: str): + return write_job_file(code) +``` + +- [ ] **Step 6: Update integration test to assert the new route** + +In `tests/integration/test_mcp_routes.py`, add to the loop in `test_eight_tool_routes_registered`: + +```python + "/write_job_file", +``` + +- [ ] **Step 7: Run full suite** + +```bash +uv run pytest -v +``` + +Expected: all tests pass (77 total). + +- [ ] **Step 8: Commit** + +```bash +git add spark_executor/tools/write_job.py spark_executor/server.py tests/unit/test_write_job_tool.py tests/integration/test_mcp_routes.py +git commit -m "feat: add write_job_file MCP tool" +``` + +--- + +### Stage 2 checkpoint + +You can now demonstrate the full agent loop: + +``` +save_connection(name=..., master=..., ...) -> {"name":..., "status":"SAVED"} +list_connections() -> [Connection, ...] +write_job_file(code=...) -> script_path +prepare_submit_job(connection=..., script_path=...) -> pending_id (no spark-submit) + +confirm_submit_job(pending_id=...) -> job_id, application_id +get_job_status(job_id=...) -> RUNNING | SUCCEEDED | FAILED | KILLED +get_job_logs(job_id=...) -> tail of driver logs +kill_job(job_id=...) -> KILLED +``` + +The submission step is now a two-phase commit: `prepare_submit_job` snapshots the connection's `master` / `deploy_mode` / `spark_conf` into a `PendingSubmission` record, and only `confirm_submit_job` actually invokes `spark-submit`. Edits to the named `Connection` between prepare and confirm are intentionally ignored. + +--- + +# Stage 3 — Async, registry, log cache, multi-tenant (★★★☆☆) + +> **Note:** Stage 3 is a non-trivial refactor. It is broken out here at a planning level only — **do not implement it as part of this plan**. A separate plan should be authored for Stage 3 once Stage 2 is shipping. This section exists so the Stage-1 / Stage-2 code can be structured to make Stage 3 cheap. + +### Why Stage 3 needs its own plan + +Stage 3 changes the runtime model: + +1. `submit_job` becomes non-blocking: it returns `job_id` immediately, and `spark-submit` runs as an `asyncio.create_subprocess_exec` task in the background. +2. The in-memory `JobStore` (Task 5) is replaced with a SQLite-backed `JobRegistry` (`spark_executor/core/job_registry.py`). Schema: `job_id TEXT PK, application_id TEXT, status TEXT, submit_time TEXT, owner TEXT, connection TEXT, script_path TEXT, queue TEXT, params JSON`. The `connection` table mirrors `ConnectionStore` records. +3. A status poller (`spark_executor/core/status_poller.py`) runs as a FastAPI `lifespan` background task, periodically refreshing `status` for every non-terminal job via `get_application_status`. +4. `get_job_logs` no longer calls `yarn logs` on every request. A `LogCache` (`spark_executor/core/log_cache.py`) persists `yarn logs` output to a per-application file under `.log_cache/.log` and accepts an `offset` parameter; subsequent calls only fetch logs after that offset. +5. An `owner` field is added to `Job` and threaded through every tool, with a filter so a tenant can only see their own jobs. Connection records may also be scoped per-owner at this point. + +### Implications for Stage 1 / Stage 2 code (so Stage 3 stays cheap) + +These are constraints the current plan already respects: + +- **Tool functions take keyword-only arguments** and return a Pydantic model. Stage 3 will add an `owner` kwarg without changing call sites. +- **`JobStore` is a class with `put / get / list`**, not module-level globals. Stage 3 swaps the implementation; the surface stays the same. +- **`ConnectionStore` is a class with `list_all / get / save / delete`** keyed by `name`, with a JSON file under `SPARK_EXECUTOR_DATA_DIR` (default `./data/`). Stage 3 may migrate it into the same SQLite DB; the surface is unchanged. +- **`PendingStore` is a class with `list_all / get / save / delete`** keyed by `pending_id`, with a JSON file in the same `data_dir`. The two-step `prepare` → `confirm` flow is intentionally synchronous and CPU-only; Stage 3 does not need to change it to introduce async submission, because `confirm_submit_job` already separates "what to submit" (the persisted snapshot) from "where to submit" (the YARN call). The async refactor in Stage 3 will move the *YARN call* inside `confirm_submit_job` to a background task; the prepare/confirm shape and the `PendingStore` interface stay the same. +- **All `yarn`/`spark-submit` calls go through `core/` modules** that return data, never raise HTTP errors directly. Stage 3 wraps them in async equivalents. +- **Pydantic `Job` includes `submit_time`, `script_path`, `queue`, and `connection` already** — Stage 3 only needs to add `owner` and `status`. SQLite migration is one `CREATE TABLE` away. +- **`PendingSubmission` carries a snapshot of `master` / `deploy_mode` / `spark_conf`** at prepare time, not a reference to a live `Connection`. This means Stage 3's multi-tenant scoping can be enforced at *prepare* time (rejecting unauthorised connections) without needing a connection lookup at confirm time. + +### Stage 3 task outline (deferred) + +When the user is ready, a new plan should be written with at minimum: + +1. Async `confirm_submit_job`: `asyncio.create_subprocess_exec` wrapper, persist `Job` with `status="SUBMITTING"` immediately and update to `SUBMITTED`/`FAILED` when the background task completes. `prepare_submit_job` is unchanged. +2. SQLite `JobRegistry` with `init_db()` called from `main.py` lifespan. +3. Status poller task started in lifespan, cancelled on shutdown. +4. `LogCache` writing per-application logs to disk; `get_job_logs(job_id, offset=0)`. +5. `owner` parameter plumbed through all twelve tools (5 pending + 3 job + 4 connection); `JobStore.list(owner=...)` and `PendingStore.list_all(owner=...)` filters. +6. `get_job_status` returns cached `status` from the registry (refreshed by poller), not a live YARN call. + +Target size: 1000–2000 lines of Python, including tests. + +--- + +## Self-Review Notes (author checklist) + +- **Spec coverage:** + - Stage 1 covers the job lifecycle tools (`get_job_status` / `get_job_logs` / `kill_job`), `application_id` extraction (both forms), 5000-char log tail, and YARN CLI invocation. ✔ + - Stage 1 also covers the connection-management requirement: 4 tools (save/list/get/delete) backed by `ConnectionStore` (JSON at `./data/connections.json`, overridable via `SPARK_EXECUTOR_DATA_DIR`). ✔ + - Stage 1 also covers the two-step pending submission requirement (`prepare_submit_job` / `confirm_submit_job` / `list_pending_jobs` / `get_pending_job` / `cancel_pending_job`): parameters are snapshotted into a `PendingSubmission` and persisted to `./data/pending_jobs.json`; `spark-submit` is only invoked by `confirm_submit_job`, which is the user-second-confirmation step. ✔ + - Stage 2 covers `write_job_file` writing to `/tmp/jobs/job_xxx.py`. ✔ + - Stage 3 is documented as deferred; the Stage-1/Stage-2 code is shaped to make it cheap. ✔ +- **Placeholder scan:** No "TBD" or "implement later" markers. All steps contain code or concrete commands. Stage 3 is explicitly out of scope and named as such. +- **Type consistency:** `Job`, `JobStatus`, `SubmitResult`, `Connection`, `PendingSubmission` Pydantic models defined in Task 1. `JobStore.put/get/list` is consistent across the three job-lifecycle tools (Tasks 17–19). `ConnectionStore.list_all/get/save/delete` is consistent across `connections.py` (Tasks 7–10) and `submit.py` reads it in Task 12. `PendingStore.list_all/get/save/delete` is consistent across all five pending tools (Tasks 12–16); `confirm_submit_job` writes back through it with status transitions. `Job.connection` is set by `confirm_submit_job` (Task 13) and consumed by `get_job_status` / `get_job_logs` / `kill_job` (Tasks 17–19). diff --git a/docs/superpowers/skills/spark-executor-mcp-operate/SKILL.md b/docs/superpowers/skills/spark-executor-mcp-operate/SKILL.md new file mode 100644 index 0000000..af9a87f --- /dev/null +++ b/docs/superpowers/skills/spark-executor-mcp-operate/SKILL.md @@ -0,0 +1,449 @@ +--- +name: spark-executor-mcp-operate +description: Operation guide for the spark-executor-mcp service — 16 MCP tools that submit, monitor, fetch logs from, and kill PySpark jobs on YARN, plus connection-profile and script-file management. USE THIS SKILL whenever the user wants to do anything with the spark-executor-mcp MCP service: submit a Spark/PySpark job to YARN, check job status or result, fetch container logs, kill a running job, save/list/get/delete a YARN connection, or write/read/update a PySpark script file. Also use it when the user asks generally about Spark, YARN, pyspark, cluster jobs, or distributed compute via this service. Do NOT guess tool names or invent a submit flow — read this skill first to learn the two-step prepare/confirm flow, the dual-ID contract (job_id vs application_id), and the SQL guard / path restrictions. The MCP service does NOT generate PySpark code — the calling LLM writes the code, then uses `write_job_file` to persist it on disk before submitting. +--- + +# spark-executor-mcp 操作指南 + +Spark-executor-mcp is an MCP (Model Context Protocol) service that exposes +**16 tools** for managing PySpark jobs on YARN. It runs as a FastAPI app +mounted at `/spark-executor-mcp`; clients (typically LLM agents) call the +tools via the standard MCP `tools/call` flow. + +This skill is the canonical reference for an LLM agent using the service. +It does **not** describe implementation details — for that, see the source +in `spark_executor/`. + +--- + +## 0. Mindset before you start + +Three things you MUST internalize before calling any tool: + +1. **The submit flow is two-step: `prepare_submit_job` then `confirm_submit_job`.** + Prepare is cheap and reversible (just snapshots params + writes JSON to + disk). Confirm is irreversible — it actually invokes `spark-submit` and + a YARN application starts running. Never skip the prepare step, even + when the user is in a hurry. + +2. **Every Job has two IDs that are NOT interchangeable everywhere.** + - `job_id` — local, 12-char hex (`a1b2c3d4e5f6`), generated by the + service. Used for: looking up the Job record, all MCP tool calls. + - `application_id` — YARN's own ID (`application_17400000001_0001`). + Used for: YARN UI, `tracking_url`, the YARN REST API. + The service's 4 job-lifecycle tools accept **either** ID (after the + 2026-06-29 fix), but other systems (YARN UI, kubectl-style scripts) + only know `application_id`. **Always store both** from every + `confirm_submit_job` response. + +3. **The LLM writes the PySpark code. The MCP service writes the file.** + `write_job_file` is a file writer with SQL-safety checks, not a + code generator. You compose the code in your context, then call + `write_job_file(code=...)` to drop it on the container's + filesystem where `spark-submit` can see it. + +--- + +## 1. Tool index + +| Operation ID | HTTP | Group | Purpose | +|---|---|---|---| +| `save_connection` | POST | Connection | Upsert a named YARN connection (master, deploy_mode, spark_conf, …) | +| `list_connections` | POST | Connection | List all saved connections | +| `get_connection` | POST | Connection | Fetch one connection by name | +| `delete_connection` | POST | Connection | Delete a connection by name | +| `write_job_file` | POST | Job file | Write LLM-written PySpark code to disk | +| `read_job_file` | POST | Job file | Read a script file back (capped at 1 MB) | +| `update_job_file` | POST | Job file | Overwrite a script file (capped at 1 MB) | +| `prepare_submit_job` | POST | Submit | Snapshot a job to `PendingSubmission`; **does NOT submit** | +| `confirm_submit_job` | POST | Submit | Actually invoke `spark-submit` for a `pending_id` | +| `update_pending_job` | POST | Submit | Edit a `PENDING` submission (script_path, queue, …) | +| `cancel_pending_job` | POST | Submit | Cancel a `PENDING` submission | +| `list_pending_jobs` | POST | Submit | List every PendingSubmission regardless of status | +| `get_pending_job` | POST | Submit | Fetch one PendingSubmission by id | +| `get_job_status` | POST | Job lifecycle | Query YARN for a job's live state (RUNNING / SUCCEEDED / …) | +| `get_job_result` | POST | Job lifecycle | Terminal-oriented view: final_status, diagnostics, timing | +| `get_job_logs` | POST | Job lifecycle | Aggregated container logs (default last 5000 chars) | +| `kill_job` | POST | Job lifecycle | PUT state=KILLED to YARN REST for the app | + +(Total: 16 MCP tools + `/health` GET, which is not exposed as a tool.) + +--- + +## 2. The canonical happy path: submit a PySpark job + +``` +1. save_connection # one-time, per cluster +2. write_job_file # write the LLM's PySpark to disk +3. read_job_file # (optional) sanity-check what was written +4. prepare_submit_job # snapshot + persist; no YARN call yet +5. confirm_submit_job # NOW spark-submit runs +6. get_job_status / logs # poll while running +7. get_job_result # terminal view when done +``` + +### 2.1 Save a connection (one-time per cluster) + +```python +save_connection( + name="prod", + master="yarn", + deploy_mode="cluster", + yarn_rm_url="http://rm.example.com:8088", + spark_conf={"spark.sql.shuffle.partitions": "200"}, + # optional: auth_type="kerberos", ssl_verify=False, ... +) +``` + +**`name`** is the lookup key — keep it short and stable. Re-calling +`save_connection` with the same `name` overwrites; the original `yarn_rm_url` +and `spark_conf` are **snapshot into the PendingSubmission at prepare time** +and survive later edits to the connection (this is intentional — the user +can re-point the connection at a different cluster without retargeting +in-flight submissions). + +### 2.2 Write the script to disk + +The LLM writes the PySpark code. Then: + +```python +script_path = write_job_file(code=PYSPARK_SOURCE)["script_path"] +# Returns {"script_path": "/absolute/path/in/container.py", "bytes": N} +``` + +The path lives inside the container's `SPARK_EXECUTOR_JOBS_DIR` (default +`./data/jobs/`). For pre-existing scripts on the host, mount the host +directory into the container and skip this step — pass the in-container +path directly to `prepare_submit_job`. + +### 2.3 Prepare (snapshot, no YARN call) + +```python +prepare_submit_job( + connection="prod", + script_path=script_path, + app_name="daily-aggregation", # REQUIRED, no default + queue="default", # these 4 must be explicit + executor_memory="2G", + executor_cores=2, + num_executors=2, + # extra_args={"jars": "hdfs://...jar"}, # optional +) +# Returns {"pending_id": "p_a1b2c3", "status": "PENDING", "parameters": {...}} +``` + +**You MUST pass `queue`, `executor_memory`, `executor_cores`, `num_executors` +explicitly every time** — even though defaults exist server-side, the +service rejects implicit defaults with a 400 listing the assumed values. +This forces the agent (or the human reviewer) to see and approve what +will actually be submitted. + +The response is the only chance to capture the `pending_id` you need for +`confirm_submit_job`. Store it. + +### 2.4 (Optional) Edit before submit + +If something needs to change between prepare and confirm: + +```python +update_pending_job( + pending_id="p_a1b2c3", + queue="research", + # script_path can also be changed — but the new file MUST exist + # and pass the SQL guard +) +``` + +Refuses to edit anything that is no longer `PENDING`. + +### 2.5 Confirm (the irreversible step) + +```python +result = confirm_submit_job(pending_id="p_a1b2c3") +# Returns {"job_id": "a1b2c3d4e5f6", +# "application_id": "application_17400000001_0001", +# "tracking_url": "http://rm:8088/proxy/application_17400000001_0001/"} +``` + +**STORE BOTH IDS NOW.** Future calls to `get_job_status` / `get_job_logs` +/ `kill_job` accept either, but `application_id` is the one YARN itself +recognizes. + +`confirm_submit_job` is **idempotent** on `SUBMITTED` (re-calling returns +the same record) and **retries** on transient `SparkSubmitError` (3 +attempts with 5 s delay by default). It transitions to `FAILED` only +after all retries are exhausted. + +### 2.6 Monitor / fetch / kill + +```python +# While running +state = get_job_status(job_id="a1b2c3d4e5f6").state +logs = get_job_logs(job_id="a1b2c3d4e5f6", tail_chars=10000) + +# When done +result = get_job_result(job_id="a1b2c3d4e5f6") +# result.state -> "SUCCEEDED" / "FAILED" / "KILLED" / ... +# result.final_status -> "SUCCEEDED" / "FAILED" / "KILLED" / "UNDEFINED" +# result.diagnostics -> YARN's "why" string (often empty on success) +# result.tracking_url -> the same proxy URL confirm_submit_job returned + +# To stop it +kill_job(job_id="a1b2c3d4e5f6") +``` + +All four accept **either** `job_id` or `application_id`. Pass whichever +you have on hand; the service tries `job_id` first, falls back to +`application_id`. + +--- + +## 3. The PendingSubmission state machine + +``` + prepare_submit_job + | + v + +-------+ cancel_pending_job + |PENDING+----------> CANCELLED (terminal) + +---+---+ + | confirm_submit_job + v + +----------+ + | SUBMITTED| (terminal — kill via kill_job, not cancel) + +----------+ + (or) + +--------+ + | FAILED | <-- confirm_submit_job exhausted retries + +--------+ + | confirm_submit_job (resets to PENDING, retries) + v + PENDING +``` + +- `cancel_pending_job` **refuses** to cancel `SUBMITTED` or `FAILED` + records — those are terminal; use `kill_job` instead. +- `update_pending_job` **refuses** to edit anything that is not + `PENDING`. +- `confirm_submit_job` on a `SUBMITTED` record is a no-op (returns the + same `job_id` / `application_id`). On a `FAILED` record, it resets + to `PENDING` and retries. On a `CANCELLED` record, it raises 400. + +--- + +## 4. The dual-ID contract + +| ID | Where it comes from | What it's good for | +|---|---|---| +| `job_id` | service-generated, 12-char hex (`a1b2c3d4e5f6`) | Looking up the Job in this service's store; all 4 lifecycle tools | +| `application_id` | parsed from `spark-submit` stderr (`application__`) | YARN UI, YARN REST, `tracking_url`, talking to ops scripts | + +**Both IDs are returned by `confirm_submit_job`.** Both are accepted by +`get_job_status` / `get_job_result` / `get_job_logs` / `kill_job`. **You +don't need to remember which to pass — but DO remember to store both +from the confirm response** because: + +- If the Job store has been wiped (rare; only on data dir loss), only + `application_id` is still meaningful to YARN — you can't poll a YARN + app you can't identify. +- If you're going to paste a link to a teammate, `tracking_url` (derived + from `application_id`) is the human-friendly form. + +If a lifecycle tool returns `KeyError("No Job found for id='X' …")`, +check that `X` came from a successful `confirm_submit_job` response. +The error message intentionally mentions both ID forms and shows an +example of each so you can self-diagnose. + +--- + +## 5. Connection profiles + +`Connection` is the cluster profile — it's the "where to submit" record. +Fields: + +| Field | Required | Default | Notes | +|---|---|---|---| +| `name` | yes | — | The lookup key | +| `master` | yes | `"yarn"` | `"yarn"` / `"spark://…"` / `"k8s://…"` / `"local[N]"` | +| `deploy_mode` | yes | `"cluster"` | `"cluster"` or `"client"` | +| `yarn_rm_url` | no | env `YARN_RESOURCE_MANAGER_URL` | Required if you'll call `get_job_status` / `get_job_logs` / `kill_job` (the lifecycle tools need a YARN endpoint) | +| `spark_conf` | no | `{}` | Merged into `spark-submit --conf` at submit time | +| `ssl_verify` | no | env `SPARK_EXECUTOR_SSL_VERIFY_DEFAULT` | bool | +| `ssl_ca_bundle` | no | env `SPARK_EXECUTOR_SSL_CA_BUNDLE_DEFAULT` | Path to CA file | +| `auth_type` | no | `"none"` | `"none"` / `"simple"` / `"basic"` / `"kerberos"` | +| `auth_user` / `auth_password` | for `basic` | — | | +| `auth_principal` / `auth_keytab` | for `kerberos` (display only) | — | Real auth uses the system cache | + +**Snapshot semantics:** when `prepare_submit_job` runs, it copies +`master` / `deploy_mode` / `spark_conf` / `yarn_rm_url` from the +**current** Connection into the `PendingSubmission`. Editing the +Connection later does **not** retarget an in-flight submission. + +**Lifecycle tools use the Job's snapshotted `yarn_rm_url`** (set at +confirm time) — not the Connection's current value. If the RM moves +between confirm and a follow-up `get_job_logs`, you have a problem; +re-pointing the Connection does not help. + +--- + +## 6. Job file workflow (`write_job_file` / `read_job_file` / `update_job_file`) + +These are **filesystem I/O tools, not code generators**. The LLM in +context writes the PySpark code; these tools just persist it where +`spark-submit` can find it, with safety guards. + +**Typical lifecycle:** + +``` +LLM composes code --> write_job_file(code=...) # write +LLM reviews code --> read_job_file(script_path=...) # verify +LLM revises code --> update_job_file(script_path=..., # overwrite + content=...) +LLM submits --> prepare_submit_job(script_path=...) +``` + +**Guards (these are why the tools exist instead of just `cat > file.sh`):** + +- **SQL guard** — `write_job_file` and `update_job_file` both reject + code containing forbidden SQL (DROP, DELETE, UPDATE, TRUNCATE, ALTER, + GRANT, REVOKE, …). The check is conservative; if it false-positives, + rewrite to use only SELECT/INSERT or restructure. +- **Path restriction** — `update_job_file` only writes under + `SPARK_EXECUTOR_JOBS_DIR`. This blocks accidentally overwriting + host-mounted configs or the service's own files. +- **Size cap** — 1 MB on both `read_job_file` and `update_job_file`. + Use `read_job_file` for small scripts; for huge ones, mount the host + directory and skip these tools. + +**For pre-existing scripts on the host:** mount the host dir into the +container (e.g. `-v /host/scripts:/app/scripts:ro` in `docker run`) and +pass the in-container path directly to `prepare_submit_job` — no need +to round-trip through the file tools. + +--- + +## 7. Error reference + +| Status | When | What to do | +|---|---|---| +| 400 | `prepare_submit_job` without explicit defaults; `cancel_pending_job` on terminal state; `update_pending_job` on non-PENDING; `prepare_submit_job` with missing/non-file `script_path`; SQL policy violation | Read the message — it almost always names the fix. Re-submit with corrected args. | +| 404 | `KeyError` from tool layer: unknown `connection`, `pending_id`, `job_id`, or `application_id` | The ID you sent doesn't match anything in the store. For lifecycle tools, the error message explicitly suggests both ID forms. | +| 422 | Pydantic validation: missing field, wrong type | Usually a typo or missing required arg. The response body lists every field error. | +| 500 | spark-submit subprocess failed and was not a `SparkSubmitError` (e.g. binary missing); YARN REST call failed | The tool logs the full traceback to `data/logs/error/...` — check there. | + +The service translates exceptions at the FastAPI layer: + +- `KeyError` → 404 (with the KeyError's string as `detail`) +- `ValueError` → 400 +- Pydantic validation → 422 + +So when the server returns 404, the body `detail` field is the same +"Unknown job_id: …" / "Unknown connection: …" string the tool layer +raised. + +--- + +## 8. Common pitfalls + +| Pitfall | Why it bites | What to do | +|---|---|---| +| Skipping the prepare step | `confirm_submit_job` requires a `pending_id`; calling it without one returns 400 | Always prepare first | +| Passing implicit defaults | Server returns 400 listing what you didn't confirm | Pass `queue`, `executor_memory`, `executor_cores`, `num_executors` every time | +| `script_path` doesn't exist in the container | `prepare_submit_job` returns 400; the agent generated the path in its own context but the container's filesystem is different | Either call `write_job_file` first, or mount a host dir | +| SQL injection in the script | `write_job_file` / `update_job_file` / `prepare_submit_job` all 400 | Only use SELECT/INSERT; never DROP/DELETE/UPDATE | +| `cancel_pending_job` on a running YARN app | Returns 400 — submissions are terminal once submitted | Use `kill_job(job_id=...)` instead, which talks to YARN REST | +| Re-pointing a Connection and expecting in-flight jobs to follow | They don't — the Job snapshots `yarn_rm_url` at confirm time | Re-issue with a new `prepare_submit_job` if you actually need a different RM | +| Calling lifecycle tools with the wrong ID type | Used to fail; **fixed in the 2026-06-29 update** — both IDs are accepted | Just pass whichever you have; the service tries `job_id` first | +| "Session not found" at the MCP layer | Multi-worker gunicorn + in-process session state | Out of scope for this skill; the service warns on startup if `GUNICORN_WORKERS>1`; set `GUNICORN_WORKERS=1` or front with a sticky LB | + +--- + +## 9. End-to-end example (LLM agent's perspective) + +User: "Run a word count on `s3a://my-bucket/books/*.txt` and write the +result back to `s3a://my-bucket/wc/`." + +``` +# 1. Cluster profile (assumes this is a fresh setup) +save_connection( + name="prod", + master="yarn", + deploy_mode="cluster", + yarn_rm_url="http://rm.prod.example.com:8088", + spark_conf={ + "spark.sql.shuffle.partitions": "200", + "spark.hadoop.fs.s3a.access.key": "...", + "spark.hadoop.fs.s3a.secret.key": "...", + }, +) + +# 2. Compose the PySpark code IN YOUR CONTEXT +pyspark_code = ''' +from pyspark.sql import SparkSession +spark = SparkSession.builder.appName("wordcount").getOrCreate() +sc = spark.sparkContext +files = sc.wholeTextFiles("s3a://my-bucket/books/") +words = files.flatMap(lambda f: f[1].lower().split()) \\ + .map(lambda w: (w, 1)) \\ + .reduceByKey(lambda a, b: a + b) +df = words.toDF(["word", "count"]) +df.write.mode("overwrite").parquet("s3a://my-bucket/wc/") +''' + +# 3. Persist to disk +script_path = write_job_file(code=pyspark_code)["script_path"] + +# 4. Sanity-check (optional but recommended) +contents = read_job_file(script_path=script_path)["content"] +# ... review in context ... + +# 5. Prepare (snapshot — no YARN call yet) +pending = prepare_submit_job( + connection="prod", + script_path=script_path, + app_name="wordcount-books", + queue="default", + executor_memory="2G", + executor_cores=2, + num_executors=2, +) +pending_id = pending["pending_id"] + +# 6. Confirm (NOW spark-submit runs) +result = confirm_submit_job(pending_id=pending_id) +job_id = result["job_id"] # STORE +application_id = result["application_id"] # STORE +tracking_url = result["tracking_url"] # STORE + +# 7. Poll +import time +while True: + st = get_job_status(job_id=job_id).state + if st in ("SUCCEEDED", "FAILED", "KILLED", "FINISHED"): + break + time.sleep(30) + +# 8. Final result +res = get_job_result(job_id=job_id) +print(res.final_status, res.diagnostics) + +# If something went wrong and you want to see why: +# logs = get_job_logs(job_id=job_id, tail_chars=20000) +``` + +--- + +## 10. Reference paths (for the curious) + +This skill is the **operation** guide. The architectural / implementation +details live in: + +- `CLAUDE.md` — project conventions, persistence layout, two-step submit + rationale, MCP arg pattern, MCP session affinity caveat. +- `docs/superpowers/plans/2026-06-24-spark-executor-mcp.md` — the + original implementation plan with Stage 1 / 2 / 3 breakdown. +- `spark_executor/server.py` — the route + exception-handler layer. +- `spark_executor/tools/requests.py` — the Pydantic body models; useful + to read when the MCP 422 response confuses you. +- `spark_executor/core/job_store.py` — Job persistence; explains the + dual-ID contract and the file-backed JSON layout. diff --git a/spark_executor/server.py b/spark_executor/server.py index 9f52d65..285c177 100644 --- a/spark_executor/server.py +++ b/spark_executor/server.py @@ -12,14 +12,14 @@ from spark_executor.tools.connections import ( list_connections, save_connection, ) -from spark_executor.tools.generate import generate_job_file +from spark_executor.tools.write_job import write_job_file from spark_executor.tools.job_file import read_job_file, update_job_file from spark_executor.tools.kill import kill_job from spark_executor.tools.logs import get_job_logs from spark_executor.tools.requests import ( ConnectionNameRequest, EmptyRequest, - GenerateJobFileRequest, + WriteJobFileRequest, GetJobLogsRequest, JobIdRequest, PendingIdRequest, @@ -46,7 +46,7 @@ app = FastAPI(title="Spark Executor MCP", version="0.0.1", description="Spark Ex _DEFAULTS_TO_CONFIRM = { "queue": "default", - "executor_memory": "4G", + "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, } @@ -90,7 +90,7 @@ def health_check(): "yarn_rm_url into a PendingSubmission record and persist it. " "Does NOT invoke spark-submit. Returns pending_id for use with " "confirm_submit_job (the user-second-confirmation step).\n\n" - "REQUIRED PATTERN for LLM-generated code: call generate_job_file(code=...) " + "REQUIRED PATTERN for LLM-generated code: call write_job_file(code=...) " "first, then pass the returned script_path here. Direct submission with a " "synthetic path (one that only exists in the agent's context) will be " "rejected with HTTP 400 — the script must exist inside the container's " @@ -297,19 +297,24 @@ def _delete_connection(req: ConnectionNameRequest): # --- LLM-driven PySpark generation (Stage 2) --- @app.post( - "/generate_job_file", - operation_id="generate_job_file", - summary="Write LLM-generated PySpark code to disk", + "/write_job_file", + operation_id="write_job_file", + summary="Write LLM-authored PySpark code to disk", description=( - "Takes a PySpark code string and writes it to a timestamped file under " + "Takes a PySpark code string the LLM has already composed in its " + "context and writes it to a timestamped file under " "SPARK_EXECUTOR_JOBS_DIR (default ./data/jobs/). Returns the absolute " "path for use as the script_path argument of prepare_submit_job — the " - "two-step pattern means the LLM can produce code, the user can review " - "the resulting file, and only then is the job submitted." + "two-step pattern means the LLM writes the file, the user can review " + "it (via read_job_file), and only then is the job submitted.\n\n" + "Note: this tool does NOT generate PySpark code. The calling LLM is " + "expected to have already written the code; this tool only persists " + "it. Code is also run through the SQL safety policy (SELECT/INSERT " + "only) before being written — forbidden statements cause a 400." ), ) -def _generate_job_file(req: GenerateJobFileRequest): - return generate_job_file(req.code) +def _write_job_file(req: WriteJobFileRequest): + return write_job_file(req.code) @app.post( @@ -318,7 +323,7 @@ def _generate_job_file(req: GenerateJobFileRequest): summary="Read the contents of an existing PySpark script", description=( "Returns the text content of an existing script file at the given " - "path. Caps reads at 1 MB. Typical use: after generate_job_file " + "path. Caps reads at 1 MB. Typical use: after write_job_file " "returns a path, call read_job_file on that path to inspect what " "was actually written, before deciding to prepare_submit_job or " "update_job_file." @@ -334,7 +339,7 @@ def _read_job_file(req: ReadJobFileRequest): summary="Overwrite an existing PySpark script with new content", description=( "Replaces the entire content of an existing script file. Path must " - "be under SPARK_EXECUTOR_JOBS_DIR (the dir generate_job_file writes " + "be under SPARK_EXECUTOR_JOBS_DIR (the dir write_job_file writes " "to) — protects against overwriting host-mounted configs or other " "non-script files. Caps writes at 1 MB. Typical use: read_job_file, " "edit the content (LLM or human), update_job_file, then " diff --git a/spark_executor/tools/generate.py b/spark_executor/tools/generate.py deleted file mode 100644 index 029f9c0..0000000 --- a/spark_executor/tools/generate.py +++ /dev/null @@ -1,43 +0,0 @@ -# coding=utf-8 -""" -@Time :2026/6/24 -@Author :tao.chen -""" -from common.logging import logger -from common.sql_guard import validate_pyspark_code -from spark_executor.core.job_writer import write_job_file - - -class SqlGuardViolation(ValueError): - """Raised when generate_job_file receives code that violates the SQL - safety policy. -> HTTP 400 via the FastAPI ValueError handler. - """ - pass - - -def generate_job_file(code: str) -> dict[str, str]: - """Write a PySpark code string to disk; return its absolute path. - - Validates the code against the SQL safety policy (SELECT/INSERT only) - BEFORE writing, so the agent gets immediate feedback rather than - learning at prepare_submit_job time. Use the returned path as the - `script_path` argument of prepare_submit_job. - - The output directory is controlled by the SPARK_EXECUTOR_JOBS_DIR env var - (default: ./data/jobs/). - """ - logger.debug(f"generate_job_file enter code_bytes={len(code)}") - offenses = validate_pyspark_code(code) - if offenses: - logger.warning( - f"generate_job_file rejected: SQL policy violation(s): {offenses}" - ) - raise SqlGuardViolation( - f"PySpark code violates SQL safety policy. " - f"Only SELECT and INSERT statements are allowed. " - f"Found forbidden statement(s): {offenses}. " - f"Rewrite the code to use only SELECT/INSERT." - ) - path = write_job_file(code) - logger.info(f"generate_job_file ok script_path={path}") - return {"script_path": path} diff --git a/spark_executor/tools/job_file.py b/spark_executor/tools/job_file.py index fd542a3..860040d 100644 --- a/spark_executor/tools/job_file.py +++ b/spark_executor/tools/job_file.py @@ -6,7 +6,7 @@ Read + update the contents of an existing PySpark script file. These two tools close the review-and-edit loop: - generate_job_file(code=...) -> {script_path} + write_job_file(code=...) -> {script_path} read_job_file(script_path=...) -> {content, path} <-- inspect update_job_file(path, content) -> {path, bytes_written} <-- edit prepare_submit_job(path) -> {pending_id, ...} @@ -43,14 +43,14 @@ def _check_readable(script_path: str) -> None: def _check_writable(script_path: str) -> None: """update_job_file is restricted to files under settings.jobs_dir - (the same dir generate_job_file writes to). This prevents the agent + (the same dir write_job_file writes to). This prevents the agent from overwriting arbitrary host-mounted files or the app's own code. """ if not script_path or not os.path.isfile(script_path): raise ScriptFileError( f"script_path does not exist or is not a file: {script_path!r}. " f"update_job_file can only edit existing files. " - f"Use generate_job_file to create a new one." + f"Use write_job_file to create a new one." ) jobs_root = Path(settings.jobs_dir).resolve() target = Path(script_path).resolve() @@ -59,7 +59,7 @@ def _check_writable(script_path: str) -> None: except ValueError: raise ScriptFileError( f"script_path must be under {jobs_root} (the directory " - f"generate_job_file writes to). Got {script_path!r}. " + f"write_job_file writes to). Got {script_path!r}. " f"This restriction protects host-mounted configs and other " f"non-script files from being overwritten by the agent." ) diff --git a/spark_executor/tools/requests.py b/spark_executor/tools/requests.py index 54ca2a7..fdae819 100644 --- a/spark_executor/tools/requests.py +++ b/spark_executor/tools/requests.py @@ -42,7 +42,7 @@ class PrepareSubmitJobRequest(BaseModel): description=( "Absolute path to the PySpark script inside the container's " "filesystem. Must point at an existing regular file. For " - "LLM-generated code, call generate_job_file(code=...) first " + "LLM-generated code, call write_job_file(code=...) first " "and pass the returned script_path here. For pre-existing " "files on the host, mount them via a docker volume and pass " "the in-container path. Returns 400 with a remediation hint " @@ -54,8 +54,8 @@ class PrepareSubmitJobRequest(BaseModel): description="YARN queue to submit to. Must be explicitly confirmed by the caller.", ) executor_memory: str = Field( - default="4G", - description="Executor memory, e.g. '4G'. Must be explicitly confirmed by the caller.", + default="2G", + description="Executor memory, e.g. '2G'. Must be explicitly confirmed by the caller.", ) executor_cores: int = Field( default=2, @@ -102,13 +102,15 @@ class ConnectionNameRequest(BaseModel): name: str -class GenerateJobFileRequest(BaseModel): +class WriteJobFileRequest(BaseModel): code: str = Field( ..., description=( - "Full PySpark source code to write to disk. Will be passed verbatim " - "to spark-submit after the agent calls prepare_submit_job on the " - "returned path." + "Full PySpark source code to write to disk. The LLM is expected " + "to have already written this code in its own context; this tool " + "persists it to a file under SPARK_EXECUTOR_JOBS_DIR so " + "spark-submit can see it. The returned script_path is what you " + "pass to prepare_submit_job." ), ) @@ -129,7 +131,7 @@ class UpdateJobFileRequest(BaseModel): description=( "Absolute path to an existing PySpark script inside the " "container's filesystem. Must be under SPARK_EXECUTOR_JOBS_DIR " - "(the same dir generate_job_file writes to) — protects against " + "(the same dir write_job_file writes to) — protects against " "overwriting host-mounted configs or other critical files." ), ) diff --git a/spark_executor/tools/submit.py b/spark_executor/tools/submit.py index 5adfdde..61cce75 100644 --- a/spark_executor/tools/submit.py +++ b/spark_executor/tools/submit.py @@ -30,7 +30,7 @@ def _script_path_error(path: str) -> ValueError: 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 generate_job_file(code=...) first to write " + 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 " @@ -52,7 +52,7 @@ def _sql_guard_error(offenses: list[str]) -> ValueError: 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"generate_job_file, the generator should have caught this — check " + 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.)" ) @@ -67,7 +67,7 @@ def prepare_submit_job( connection: str, script_path: str, queue: str = "default", - executor_memory: str = "4G", + executor_memory: str = "2G", executor_cores: int = 2, num_executors: int = 2, app_name: str, @@ -92,12 +92,12 @@ def prepare_submit_job( 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 - # generate_job_file first, or its path refers to the host filesystem + # 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 generate_job_file + # 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: diff --git a/spark_executor/tools/write_job.py b/spark_executor/tools/write_job.py new file mode 100644 index 0000000..8b5b0dc --- /dev/null +++ b/spark_executor/tools/write_job.py @@ -0,0 +1,57 @@ +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen + +The MCP-exposed `write_job_file` tool. Despite the historical name +`generate_job_file`, this tool does NOT generate PySpark code — the +calling LLM writes the code, and this tool persists it to disk so +`spark-submit` can see it. The rename (2026-06-29) tightens the naming +to match what the tool actually does, eliminating the "I thought this +would generate code for me" confusion. + +It also runs the SQL safety policy (SELECT/INSERT only) on the code +BEFORE writing, so the agent gets immediate feedback if it slipped a +DROP / DELETE / UPDATE past its own code-generation step. +""" +from common.logging import logger +from common.sql_guard import validate_pyspark_code +# Internal helper already named write_job_file (it just writes bytes to +# disk; no SQL guard). Import with an alias to avoid the name collision +# with the MCP-exposed function below. +from spark_executor.core.job_writer import write_job_file as _write_to_disk + + +class SqlGuardViolation(ValueError): + """Raised when write_job_file receives code that violates the SQL + safety policy. -> HTTP 400 via the FastAPI ValueError handler.""" + pass + + +def write_job_file(code: str) -> dict[str, str]: + """Persist a PySpark code string to disk and return its absolute path. + + Validates the code against the SQL safety policy (SELECT/INSERT only) + BEFORE writing, so the agent gets immediate feedback rather than + learning at prepare_submit_job time. Use the returned path as the + `script_path` argument of prepare_submit_job. + + The output directory is controlled by the SPARK_EXECUTOR_JOBS_DIR env var + (default: ./data/jobs/). Each call writes to a fresh timestamped file + under that directory, so repeated calls never overwrite one another. + """ + logger.debug(f"write_job_file enter code_bytes={len(code)}") + offenses = validate_pyspark_code(code) + if offenses: + logger.warning( + f"write_job_file rejected: SQL policy violation(s): {offenses}" + ) + raise SqlGuardViolation( + f"PySpark code violates SQL safety policy. " + f"Only SELECT and INSERT statements are allowed. " + f"Found forbidden statement(s): {offenses}. " + f"Rewrite the code to use only SELECT/INSERT." + ) + path = _write_to_disk(code) + logger.info(f"write_job_file ok script_path={path}") + return {"script_path": path} diff --git a/tests/integration/test_mcp_routes.py b/tests/integration/test_mcp_routes.py index 0fcf0b9..241e054 100644 --- a/tests/integration/test_mcp_routes.py +++ b/tests/integration/test_mcp_routes.py @@ -52,7 +52,7 @@ def test_sixteen_tool_routes_registered(): "/get_connection", "/delete_connection", # LLM-driven PySpark generation (3) — Stage 2 - "/generate_job_file", + "/write_job_file", "/read_job_file", "/update_job_file", ): @@ -146,7 +146,7 @@ def test_prepare_submit_job_works_via_body(tmp_path): "connection": "prod", "script_path": str(script), "queue": "research", - "executor_memory": "4G", + "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", @@ -177,7 +177,7 @@ def test_prepare_rejects_unconfirmed_defaults(tmp_path): detail = r.json()["detail"] assert "Please confirm default values" in detail assert "queue='default'" in detail - assert "executor_memory='4G'" in detail + assert "executor_memory='2G'" in detail assert "executor_cores=2" in detail assert "num_executors=2" in detail @@ -193,7 +193,7 @@ def test_prepare_accepts_confirmed_defaults(tmp_path): "connection": "prod", "script_path": str(script), "queue": "default", - "executor_memory": "4G", + "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", @@ -203,7 +203,7 @@ def test_prepare_accepts_confirmed_defaults(tmp_path): body = r.json() assert body["status"] == "PENDING" assert body["parameters"]["queue"] == "default" - assert body["parameters"]["executor_memory"] == "4G" + assert body["parameters"]["executor_memory"] == "2G" assert body["parameters"]["executor_cores"] == 2 assert body["parameters"]["num_executors"] == 2 @@ -218,7 +218,7 @@ def test_prepare_rejects_nonexistent_script_path_with_400(): "connection": "prod", "script_path": "/nope/does_not_exist.py", "queue": "default", - "executor_memory": "4G", + "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", @@ -228,12 +228,12 @@ def test_prepare_rejects_nonexistent_script_path_with_400(): detail = r.json()["detail"] # The message must guide the agent to the right next step assert "does not exist" in detail - assert "generate_job_file" in detail + assert "write_job_file" in detail def test_prepare_rejects_sql_policy_violation_with_400(tmp_path): """SQL guard: even if the file exists, prepare_submit_job must reject - code containing DROP/DELETE/etc. (defense in depth — generate_job_file + code containing DROP/DELETE/etc. (defense in depth — write_job_file also enforces this, but a host-mounted file might bypass that).""" c = TestClient(app) c.post("/save_connection", json={"name": "prod", "master": "yarn"}) @@ -245,7 +245,7 @@ def test_prepare_rejects_sql_policy_violation_with_400(tmp_path): "connection": "prod", "script_path": str(bad_script), "queue": "default", - "executor_memory": "4G", + "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", @@ -257,14 +257,14 @@ def test_prepare_rejects_sql_policy_violation_with_400(tmp_path): assert "SELECT" in detail and "INSERT" in detail # the policy is explained -def test_generate_rejects_sql_policy_violation_with_400(tmp_path): - """Same policy at generate_job_file — agent gets immediate feedback +def test_write_rejects_sql_policy_violation_with_400(tmp_path): + """Same policy at write_job_file — agent gets immediate feedback before the file is even written to disk.""" from common import config config.settings.jobs_dir = str(tmp_path / "jobs") c = TestClient(app) r = c.post( - "/generate_job_file", + "/write_job_file", json={"code": 'spark.sql("DELETE FROM events")\n'}, ) assert r.status_code == 400 @@ -331,7 +331,7 @@ def test_list_and_get_pending_job_roundtrip_via_body(tmp_path): "connection": "prod", "script_path": str(script), "queue": "default", - "executor_memory": "4G", + "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", @@ -354,15 +354,15 @@ def test_list_connections_works_with_empty_body(): assert r.json() == [] -# --- Stage 2: generate_job_file --- +# --- Stage 2: write_job_file --- -def test_generate_job_file_accepts_code_string_in_body(tmp_path, monkeypatch): +def test_write_job_file_accepts_code_string_in_body(tmp_path, monkeypatch): """End-to-end: a Pydantic body model lets tools/call pass a code string that FastAPI would reject if it were a query parameter (length limits).""" monkeypatch.chdir(tmp_path) c = TestClient(app) code = "print('from MCP integration test')\n" * 100 # > FastAPI query limit - r = c.post("/generate_job_file", json={"code": code}) + r = c.post("/write_job_file", json={"code": code}) assert r.status_code == 200, r.text p = r.json()["script_path"] assert os.path.isfile(p) @@ -401,7 +401,7 @@ def test_unknown_connection_in_prepare_returns_404(): "connection": "nope", "script_path": "/tmp/x.py", "queue": "default", - "executor_memory": "4G", + "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", @@ -423,7 +423,7 @@ def test_confirm_non_pending_returns_400(tmp_path): "connection": "prod", "script_path": str(script), "queue": "default", - "executor_memory": "4G", + "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", @@ -448,7 +448,7 @@ def test_cancel_already_submitted_returns_400(tmp_path): "connection": "prod", "script_path": str(script), "queue": "default", - "executor_memory": "4G", + "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", @@ -477,7 +477,7 @@ def test_update_pending_job_route(tmp_path): "connection": "prod", "script_path": str(script), "queue": "default", - "executor_memory": "4G", + "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", @@ -496,7 +496,7 @@ def test_update_pending_job_route(tmp_path): got = c.post("/get_pending_job", json={"pending_id": pid}).json() assert got["queue"] == "research" # untouched fields are preserved - assert got["executor_memory"] == "4G" + assert got["executor_memory"] == "2G" assert got["app_name"] == "test-app" @@ -511,7 +511,7 @@ def test_update_pending_job_route_rejects_submitted(tmp_path, monkeypatch): "connection": "prod", "script_path": str(script), "queue": "default", - "executor_memory": "4G", + "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", diff --git a/tests/unit/test_job_file.py b/tests/unit/test_job_file.py index 2a56719..55819f3 100644 --- a/tests/unit/test_job_file.py +++ b/tests/unit/test_job_file.py @@ -145,7 +145,7 @@ def test_full_edit_cycle_via_tools(tmp_path: Path): """Simulate the agent's review-and-edit loop end-to-end.""" config.settings.jobs_dir = str(tmp_path) path = tmp_path / "script.py" - # 1. Initial state: file doesn't exist (in real flow, generate_job_file + # 1. Initial state: file doesn't exist (in real flow, write_job_file # would create it; we just create it directly here for the test) path.write_text("# v1: initial\n") # 2. Agent reads back diff --git a/tests/unit/test_logging.py b/tests/unit/test_logging.py index 2d73a23..e73985b 100644 --- a/tests/unit/test_logging.py +++ b/tests/unit/test_logging.py @@ -51,7 +51,7 @@ def test_prepare_submit_job_emits_debug_and_info(log_capture, tmp_path): connection="prod", script_path=str(script), queue="research", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", diff --git a/tests/unit/test_models.py b/tests/unit/test_models.py index cab35a4..75bd380 100644 --- a/tests/unit/test_models.py +++ b/tests/unit/test_models.py @@ -95,7 +95,7 @@ def test_pending_submission_defaults_to_pending_status(): deploy_mode="cluster", script_path="/tmp/j.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, spark_conf={}, @@ -120,7 +120,7 @@ def test_pending_submission_carries_yarn_rm_url_snapshot(): yarn_rm_url="http://rm-prod:8088", script_path="/tmp/j.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, spark_conf={}, @@ -139,7 +139,7 @@ def test_pending_submission_can_record_outcome(): deploy_mode="cluster", script_path="/tmp/j.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, spark_conf={}, diff --git a/tests/unit/test_pending_store.py b/tests/unit/test_pending_store.py index 899073c..9b9c40d 100644 --- a/tests/unit/test_pending_store.py +++ b/tests/unit/test_pending_store.py @@ -18,7 +18,7 @@ def _pending(pid: str = "p_1", created_at: datetime | None = None) -> PendingSub deploy_mode="cluster", script_path="/tmp/j.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, spark_conf={}, @@ -85,7 +85,7 @@ def _legacy_record_text() -> str: return ( '{"p_legacy": {"pending_id": "p_legacy", "app_name": "legacy", ' '"connection": "prod", "master": "yarn", "deploy_mode": "cluster", ' - '"script_path": "/tmp/j.py", "queue": "default", "executor_memory": "4G", ' + '"script_path": "/tmp/j.py", "queue": "default", "executor_memory": "2G", ' '"executor_cores": 2, "num_executors": 2, "spark_conf": {}, ' '"created_at": "2026-06-23T00:00:00"}}' ) @@ -139,7 +139,7 @@ def test_loads_records_with_missing_or_null_app_name(tmp_path: Path): "deploy_mode": "cluster", "script_path": "/tmp/j.py", "queue": "default", - "executor_memory": "4G", + "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "spark_conf": {}, diff --git a/tests/unit/test_requests.py b/tests/unit/test_requests.py index a939b61..27dc086 100644 --- a/tests/unit/test_requests.py +++ b/tests/unit/test_requests.py @@ -16,7 +16,7 @@ def test_prepare_submit_job_request_restores_defaults_for_queue_and_resources(): app_name="test-app", ) assert req.queue == "default" - assert req.executor_memory == "4G" + assert req.executor_memory == "2G" assert req.executor_cores == 2 assert req.num_executors == 2 @@ -26,7 +26,7 @@ def test_prepare_submit_job_request_accepts_extra_args(): connection="prod", script_path="/tmp/x.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", diff --git a/tests/unit/test_spark_submit.py b/tests/unit/test_spark_submit.py index 276ebec..3bb0cec 100644 --- a/tests/unit/test_spark_submit.py +++ b/tests/unit/test_spark_submit.py @@ -36,7 +36,7 @@ def test_build_command_uses_provided_master_and_deploy_mode(): deploy_mode="cluster", script_path="/tmp/jobs/job_001.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, ) @@ -44,7 +44,7 @@ def test_build_command_uses_provided_master_and_deploy_mode(): 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-memory" in cmd and "2G" 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" @@ -56,7 +56,7 @@ def test_extra_args_none_default(): deploy_mode="cluster", script_path="/tmp/jobs/job_001.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, ) @@ -70,7 +70,7 @@ def test_extra_args_single_pair(): deploy_mode="cluster", script_path="/tmp/jobs/job_001.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, extra_args={"jars": "hdfs:///libs/foo.jar"}, @@ -88,7 +88,7 @@ def test_extra_args_multiple_pairs_preserve_order(): deploy_mode="cluster", script_path="/tmp/jobs/job_001.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, extra_args={ @@ -112,7 +112,7 @@ def test_extra_args_does_not_collide_with_spark_conf(): deploy_mode="cluster", script_path="/tmp/jobs/job_001.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, spark_conf={"spark.executor.memory": "8G"}, @@ -146,7 +146,7 @@ def test_build_command_appends_spark_conf_entries(): deploy_mode="cluster", script_path="/tmp/j.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, spark_conf={"spark.sql.shuffle.partitions": "200", "spark.executor.memoryOverhead": "1G"}, @@ -169,7 +169,7 @@ def test_build_command_uses_settings_spark_submit_bin(): deploy_mode="cluster", script_path="/tmp/j.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, ) @@ -183,7 +183,7 @@ def test_build_command_supports_pyspark_binary(): deploy_mode="cluster", script_path="/tmp/j.py", queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, ) diff --git a/tests/unit/test_submit_tool.py b/tests/unit/test_submit_tool.py index 377bda5..9a2b183 100644 --- a/tests/unit/test_submit_tool.py +++ b/tests/unit/test_submit_tool.py @@ -56,7 +56,7 @@ def test_prepare_does_not_invoke_spark_submit(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -81,7 +81,7 @@ def test_prepare_persists_pending_with_snapshot(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="research", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -103,7 +103,7 @@ def test_prepare_persists_app_name(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="my-etl-job", @@ -125,7 +125,7 @@ def test_prepare_accepts_restored_defaults(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -133,7 +133,7 @@ def test_prepare_accepts_restored_defaults(monkeypatch, real_script): assert out["status"] == "PENDING" p = submit.pending_store.get(_last_pending_id()) assert p.queue == "default" - assert p.executor_memory == "4G" + assert p.executor_memory == "2G" assert p.executor_cores == 2 assert p.num_executors == 2 @@ -144,7 +144,7 @@ def test_prepare_snapshots_extra_args(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -173,7 +173,7 @@ def test_prepare_raises_for_unknown_connection(real_script): connection="missing", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -189,7 +189,7 @@ def test_prepare_rejects_nonexistent_script_path(tmp_path): connection="prod", script_path=missing, queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -203,7 +203,7 @@ def test_prepare_rejects_directory_as_script_path(tmp_path): connection="prod", script_path=str(tmp_path), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -216,21 +216,21 @@ def test_prepare_rejects_empty_script_path(): connection="prod", script_path="", queue="default", - executor_memory="4G", + executor_memory="2G", 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"): +def test_prepare_error_message_mentions_write_job_file(tmp_path): + """The agent must be told to call write_job_file first.""" + with pytest.raises(ValueError, match="write_job_file"): submit.prepare_submit_job( connection="prod", script_path=str(tmp_path / "x.py"), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -245,7 +245,7 @@ def test_confirm_rejects_if_script_was_deleted_after_prepare(tmp_path, monkeypat connection="prod", script_path=str(script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -264,7 +264,7 @@ def test_prepare_snapshots_connection_at_prepare_time(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -287,7 +287,7 @@ def test_confirm_invokes_spark_submit_and_marks_submitted(monkeypatch, real_scri connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -323,7 +323,7 @@ def test_confirm_refuses_non_pending_status(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -342,7 +342,7 @@ def test_confirm_marks_failed_on_spark_submit_error(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -368,7 +368,7 @@ def test_confirm_retries_spark_submit_failure_then_succeeds(monkeypatch, real_sc connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -395,7 +395,7 @@ def test_confirm_exhausts_retries_and_sets_failed(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -419,7 +419,7 @@ def test_confirm_idempotent_when_submitted(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -446,7 +446,7 @@ def test_confirm_resets_failed_and_retries(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -476,7 +476,7 @@ def test_confirm_updates_pending_before_job_store(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -510,7 +510,7 @@ def test_list_pending_jobs_returns_all(monkeypatch, tmp_path): connection="prod", script_path=str(a), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -519,7 +519,7 @@ def test_list_pending_jobs_returns_all(monkeypatch, tmp_path): connection="prod", script_path=str(b), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -537,7 +537,7 @@ def test_get_pending_job_returns_dump(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -562,7 +562,7 @@ def test_cancel_pending_job_marks_cancelled(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -583,7 +583,7 @@ def test_cancel_pending_job_refuses_submitted(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -604,7 +604,7 @@ def test_update_pending_updates_resource_params(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -639,7 +639,7 @@ def test_update_pending_rejects_non_pending_status(monkeypatch, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -657,7 +657,7 @@ def test_update_pending_rejects_bad_script_path(tmp_path, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", @@ -673,7 +673,7 @@ def test_update_pending_rejects_sql_violation_script(tmp_path, real_script): connection="prod", script_path=str(real_script), queue="default", - executor_memory="4G", + executor_memory="2G", executor_cores=2, num_executors=2, app_name="test-app", diff --git a/tests/unit/test_generate_tool.py b/tests/unit/test_write_job_tool.py similarity index 62% rename from tests/unit/test_generate_tool.py rename to tests/unit/test_write_job_tool.py index 833efd9..e4883a2 100644 --- a/tests/unit/test_generate_tool.py +++ b/tests/unit/test_write_job_tool.py @@ -5,7 +5,7 @@ from pathlib import Path import pytest from common import config -from spark_executor.tools import generate +from spark_executor.tools import write_job @pytest.fixture(autouse=True) @@ -23,9 +23,9 @@ def _restore_settings(): config.settings.log_level = snapshot.log_level -def test_generate_writes_code_and_returns_path(monkeypatch, tmp_path: Path): +def test_write_writes_code_and_returns_path(monkeypatch, tmp_path: Path): config.settings.jobs_dir = str(tmp_path / "data" / "jobs") - out = generate.generate_job_file( + out = write_job.write_job_file( "from pyspark.sql import SparkSession\n" "spark = SparkSession.builder.getOrCreate()\n" ) @@ -38,8 +38,18 @@ def test_generate_writes_code_and_returns_path(monkeypatch, tmp_path: Path): assert "SparkSession.builder.getOrCreate()" in f.read() -def test_generate_uses_settings_jobs_dir(tmp_path: Path): +def test_write_uses_settings_jobs_dir(tmp_path: Path): config.settings.jobs_dir = str(tmp_path / "custom") - out = generate.generate_job_file("x = 1\n") + out = write_job.write_job_file("x = 1\n") assert out["script_path"].startswith(str(tmp_path / "custom")) assert os.path.isfile(out["script_path"]) + + +def test_write_rejects_sql_violation(tmp_path: Path): + """The SQL guard is the whole reason this isn't a bare file-write + tool: the LLM might have produced a script with a forbidden + statement and we want to fail at write time, not at prepare time.""" + from spark_executor.tools.write_job import SqlGuardViolation + config.settings.jobs_dir = str(tmp_path / "data" / "jobs") + with pytest.raises(SqlGuardViolation): + write_job.write_job_file("spark.sql('DROP TABLE users')\n")