Merge pull request 'Feat/stage 1' (#1) from feat/stage-1 into main

Reviewed-on: https://gitea-production-a772.up.railway.app/taochen/mcp-server/pulls/1
This commit is contained in:
2026-06-24 07:15:11 +00:00
38 changed files with 2145 additions and 5 deletions
+22
View File
@@ -0,0 +1,22 @@
# Python
__pycache__/
*.py[cod]
*.egg-info/
.eggs/
build/
dist/
# Virtual env
.venv/
# IDE
.idea/
.vscode/
# Data directory (persisted Spark connections / pending jobs)
data/
# Test / coverage
.pytest_cache/
.coverage
htmlcov/
+43 -2
View File
@@ -1,10 +1,51 @@
# coding=utf-8
"""
@Time :2026/6/24
@Author :tao.chen
@Author :tao.chen
Process-wide loguru configuration. Import `logger` from here in every
module instead of instantiating new loggers.
Levels used in this project:
DEBUG - entry/exit of public tools, subprocess commands, file I/O paths
INFO - business events (job submitted, status changed, connection saved)
WARNING - recoverable problems (transient YARN issues, retry-able)
ERROR - raised exceptions (caller will see the traceback)
"""
import sys
from pathlib import Path
from loguru import logger
Path("data/logs/debug").mkdir(parents=True, exist_ok=True)
Path("data/logs/info").mkdir(parents=True, exist_ok=True)
logger.remove()
logger.add(sys.stderr, level="DEBUG")
logger.add(
sys.stderr,
level="DEBUG",
format=(
"<green>{time:HH:mm:ss.SSS}</green> | "
"<level>{level: <7}</level> | "
"<cyan>{name}</cyan>:<cyan>{function}</cyan>:<cyan>{line}</cyan> - "
"<level>{message}</level>"
),
enqueue=True,
colorize=True,
)
logger.add(
"data/logs/debug/{time:YYYY-MM-DD}.log",
level="DEBUG",
enqueue=True,
retention="30 days",
compression="gz",
colorize=True,
)
logger.add(
"data/logs/info/{time:YYYY-MM-DD}.log",
level="INFO",
enqueue=True,
retention="30 days",
compression="gz",
encoding="utf-8",
)
+6
View File
@@ -12,6 +12,12 @@ dependencies = [
"uvicorn>=0.49.0",
]
[dependency-groups]
dev = [
"pytest>=8.0",
"pytest-mock>=3.14",
]
[[tool.uv.index]]
url = "https://pypi.tuna.tsinghua.edu.cn/simple/"
default = true
+5
View File
@@ -0,0 +1,5 @@
# coding=utf-8
"""
@Time :2026/6/24
@Author :tao.chen
"""
+91
View File
@@ -0,0 +1,91 @@
# 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 common.logging import logger
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: <DEFAULT_DATA_DIR>/<DEFAULT_FILE_NAME>.
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():
logger.debug(f"ConnectionStore._load: file {self.path} absent, returning empty")
return {}
raw = json.loads(self.path.read_text(encoding="utf-8"))
logger.debug(f"ConnectionStore._load: loaded {len(raw)} records from {self.path}")
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)
logger.debug(f"ConnectionStore._dump: wrote {len(records)} records to {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)
logger.info(f"connection saved name={conn.name} master={conn.master}")
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)
logger.info(f"connection deleted name={name}")
return True
# Module-level singleton; replaced in tests.
store = ConnectionStore()
+28
View File
@@ -0,0 +1,28 @@
# 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())
+41
View File
@@ -0,0 +1,41 @@
# coding=utf-8
"""
@Time :2026/6/24
@Author :tao.chen
"""
import re
from common.logging import logger
_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:
logger.error("Could not find application_id in spark-submit output")
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_"):
logger.error("Could not find application_id in spark-submit output (bad tracking URL)")
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
logger.debug(
f"parse_spark_submit_output -> application_id={application_id} "
f"tracking_url={tracking_url}"
)
return application_id, tracking_url
+89
View File
@@ -0,0 +1,89 @@
# 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 common.logging import logger
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: <DEFAULT_DATA_DIR>/<DEFAULT_FILE_NAME>.
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():
logger.debug(f"PendingStore._load: file {self.path} absent, returning empty")
return {}
raw = json.loads(self.path.read_text(encoding="utf-8"))
logger.debug(f"PendingStore._load: loaded {len(raw)} records from {self.path}")
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)
logger.debug(f"PendingStore._dump: wrote {len(records)} records to {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)
logger.info(
f"pending saved pending_id={pending.pending_id} "
f"status={pending.status} connection={pending.connection}"
)
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)
logger.info(f"pending deleted pending_id={pending_id}")
return True
# Module-level singleton; replaced in tests.
store = PendingStore()
+54
View File
@@ -0,0 +1,54 @@
# coding=utf-8
"""
@Time :2026/6/24
@Author :tao.chen
"""
import subprocess
from common.logging import logger
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)
logger.debug(f"build_spark_submit_command -> {cmd}")
return cmd
def run_spark_submit(cmd: list[str]) -> "subprocess.CompletedProcess[str]":
logger.debug(f"run_spark_submit exec: {cmd}")
result = subprocess.run(cmd, capture_output=True, text=True, errors="replace")
logger.debug(
f"run_spark_submit done rc={result.returncode} "
f"stdout_len={len(result.stdout)} stderr_len={len(result.stderr)}"
)
if result.returncode != 0:
logger.error(f"spark-submit failed (rc={result.returncode}): {result.stderr[:500]}")
raise SparkSubmitError(
f"spark-submit failed (rc={result.returncode}): {result.stderr}"
)
return result
+70
View File
@@ -0,0 +1,70 @@
# coding=utf-8
"""
@Time :2026/6/24
@Author :tao.chen
"""
import re
import subprocess
from common.logging import logger
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]":
logger.debug(f"yarn _run exec: {cmd}")
result = subprocess.run(cmd, capture_output=True, text=True, errors="replace")
logger.debug(
f"yarn _run done rc={result.returncode} "
f"stdout_len={len(result.stdout)} stderr_len={len(result.stderr)}"
)
return result
def get_application_status(application_id: str) -> tuple[str, str]:
proc = _run(["yarn", "application", "-status", application_id])
if proc.returncode != 0:
logger.error(
f"yarn application -status failed (rc={proc.returncode}) for "
f"{application_id}: {proc.stderr[:500]}"
)
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}")
state = match.group(1)
logger.info(f"yarn status {application_id} -> {state}")
return state, proc.stdout
def get_application_logs(application_id: str) -> str:
proc = _run(["yarn", "logs", "-applicationId", application_id])
if proc.returncode != 0:
logger.error(
f"yarn logs failed (rc={proc.returncode}) for {application_id}: {proc.stderr[:500]}"
)
raise YarnError(
f"yarn logs failed (rc={proc.returncode}): {proc.stderr}"
)
logger.info(f"yarn logs {application_id} -> {len(proc.stdout)} chars")
return proc.stdout
def kill_application(application_id: str) -> None:
proc = _run(["yarn", "application", "-kill", application_id])
if proc.returncode != 0:
logger.error(
f"yarn application -kill failed (rc={proc.returncode}) for "
f"{application_id}: {proc.stderr[:500]}"
)
raise YarnError(
f"yarn application -kill failed (rc={proc.returncode}): {proc.stderr}"
)
logger.info(f"yarn kill {application_id} -> ok")
+55
View File
@@ -0,0 +1,55 @@
# 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
+120 -3
View File
@@ -1,12 +1,129 @@
# coding=utf-8
"""
@Time :2026/6/24
@Author :tao.chen
@Author :tao.chen
"""
from fastapi import FastAPI
from fastapi import FastAPI, Request
from fastapi.responses import JSONResponse
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.requests import (
ConnectionNameRequest,
EmptyRequest,
GetJobLogsRequest,
JobIdRequest,
PendingIdRequest,
PrepareSubmitJobRequest,
SaveConnectionRequest,
)
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")
# --- Exception handlers: translate tool-layer errors into proper HTTP statuses ---
#
# Tool functions raise KeyError for "unknown id" (job_id, pending_id, connection
# name) and ValueError for invalid state transitions (e.g. confirming a
# CANCELLED pending). Without these handlers FastAPI would map them to a bare
# 500 "Internal Server Error" which is useless to MCP clients.
@app.exception_handler(KeyError)
async def _keyerror_handler(_request: Request, exc: KeyError) -> JSONResponse:
return JSONResponse(status_code=404, content={"detail": str(exc)})
@app.exception_handler(ValueError)
async def _valueerror_handler(_request: Request, exc: ValueError) -> JSONResponse:
return JSONResponse(status_code=400, content={"detail": str(exc)})
@app.get("/health")
def health_check():
return {"status": "ok"}
return {"status": "ok"}
# MCP tool routes. fastapi-mcp discovers these and registers them as MCP tools.
# Each route takes a single Pydantic body model so tools/call (which sends args
# as JSON body) works for every tool, including those with dict-typed params
# like spark_conf.
# --- Pending submission flow (two-step submit) ---
@app.post("/prepare_submit_job")
def _prepare_submit_job(req: PrepareSubmitJobRequest):
return prepare_submit_job(**req.model_dump())
@app.post("/confirm_submit_job")
def _confirm_submit_job(req: PendingIdRequest):
return confirm_submit_job(pending_id=req.pending_id)
@app.post("/list_pending_jobs")
def _list_pending_jobs(_req: EmptyRequest = EmptyRequest()):
return list_pending_jobs()
@app.post("/get_pending_job")
def _get_pending_job(req: PendingIdRequest):
return get_pending_job(req.pending_id)
@app.post("/cancel_pending_job")
def _cancel_pending_job(req: PendingIdRequest):
return cancel_pending_job(req.pending_id)
# --- Spark job tools ---
@app.post("/get_job_status")
def _get_job_status(req: JobIdRequest):
return get_job_status(req.job_id)
@app.post("/get_job_logs")
def _get_job_logs(req: GetJobLogsRequest):
return get_job_logs(req.job_id, tail_chars=req.tail_chars)
@app.post("/kill_job")
def _kill_job(req: JobIdRequest):
return kill_job(req.job_id)
# --- Connection management tools ---
@app.post("/save_connection")
def _save_connection(req: SaveConnectionRequest):
# exclude_none so we don't overwrite the function's default with explicit None
return save_connection(**req.model_dump(exclude_none=True))
@app.post("/list_connections")
def _list_connections(_req: EmptyRequest = EmptyRequest()):
return list_connections()
@app.post("/get_connection")
def _get_connection(req: ConnectionNameRequest):
return get_connection(req.name)
@app.post("/delete_connection")
def _delete_connection(req: ConnectionNameRequest):
return delete_connection(req.name)
+5
View File
@@ -0,0 +1,5 @@
# coding=utf-8
"""
@Time :2026/6/24
@Author :tao.chen
"""
+52
View File
@@ -0,0 +1,52 @@
# 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]:
logger.debug(
f"save_connection enter name={name} master={master} deploy_mode={deploy_mode} "
f"yarn_rm_url={yarn_rm_url} spark_conf_keys={list((spark_conf or {}).keys())}"
)
conn = Connection(
name=name,
master=master,
deploy_mode=deploy_mode,
yarn_rm_url=yarn_rm_url,
spark_conf=spark_conf or {},
)
store.save(conn)
return {"name": name, "status": "SAVED"}
def list_connections() -> list[dict[str, object]]:
logger.debug("list_connections enter")
return [c.model_dump() for c in store.list_all()]
def get_connection(name: str) -> dict[str, object]:
logger.debug(f"get_connection enter name={name}")
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]:
logger.debug(f"delete_connection enter name={name}")
removed = store.delete(name)
if not removed:
raise KeyError(f"Unknown connection: {name}")
return {"name": name, "status": "DELETED"}
+24
View File
@@ -0,0 +1,24 @@
# 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]:
logger.debug(f"kill_job enter job_id={job_id}")
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 ok job_id={job_id} application_id={job.application_id}")
return {
"job_id": job_id,
"application_id": job.application_id,
"status": "KILLED",
}
+24
View File
@@ -0,0 +1,24 @@
# 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:
logger.debug(f"get_job_logs enter job_id={job_id} tail_chars={tail_chars}")
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 ok job_id={job_id} application_id={job.application_id} "
f"full_chars={len(full)} returned_chars={len(tailed)}"
)
return tailed
+50
View File
@@ -0,0 +1,50 @@
# coding=utf-8
"""
@Time :2026/6/24
@Author :tao.chen
Pydantic request models for the FastAPI route layer. The underlying tool
functions in tools/*.py still take keyword arguments; these models exist only
so fastapi-mcp can call the routes via tools/call (which sends args as a
JSON body) without 422-ing on dict-typed parameters like spark_conf.
"""
from pydantic import BaseModel, Field
class EmptyRequest(BaseModel):
"""Used for tools that take no arguments (list_connections, list_pending_jobs)."""
pass
class SaveConnectionRequest(BaseModel):
name: str
master: str
deploy_mode: str = "cluster"
yarn_rm_url: str | None = None
spark_conf: dict[str, str] | None = None
class PrepareSubmitJobRequest(BaseModel):
connection: str
script_path: str
queue: str = "default"
executor_memory: str = "4G"
executor_cores: int = 2
num_executors: int = 2
class PendingIdRequest(BaseModel):
pending_id: str
class JobIdRequest(BaseModel):
job_id: str
class GetJobLogsRequest(BaseModel):
job_id: str
tail_chars: int = 5000
class ConnectionNameRequest(BaseModel):
name: str
+21
View File
@@ -0,0 +1,21 @@
# 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:
logger.debug(f"get_job_status enter job_id={job_id}")
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 ok job_id={job_id} application_id={job.application_id} state={state}")
return JobStatus(application_id=job.application_id, state=state, raw=raw)
+165
View File
@@ -0,0 +1,165 @@
# coding=utf-8
"""
@Time :2026/6/24
@Author :tao.chen
"""
import secrets
import uuid
from datetime import datetime
from common.logging import logger
from spark_executor.core.connection_store import store as conn_store
from spark_executor.core.job_store import JobStore
from spark_executor.core.log_parser import parse_spark_submit_output
from spark_executor.core.pending_store import store as pending_store
from spark_executor.core.spark_submit import (
SparkSubmitError,
build_spark_submit_command,
run_spark_submit,
)
from spark_executor.models import Job, PendingSubmission, SubmitResult
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 = "4G",
executor_cores: int = 2,
num_executors: int = 2,
) -> dict[str, object]:
"""Snapshot connection params and persist a PendingSubmission. Does NOT submit."""
logger.debug(
f"prepare_submit_job enter connection={connection} script_path={script_path} "
f"queue={queue} executor_memory={executor_memory} executor_cores={executor_cores} "
f"num_executors={num_executors}"
)
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 ok pending_id={pending_id} connection={connection} "
f"master={conn.master} script_path={script_path}"
)
return {
"pending_id": pending_id,
"status": "PENDING",
"parameters": pending.model_dump(),
}
# Module-level job store singleton; replaced in tests.
job_store: JobStore = JobStore()
def confirm_submit_job(*, pending_id: str) -> SubmitResult:
"""Actually invoke spark-submit for a previously-prepared PendingSubmission."""
logger.debug(f"confirm_submit_job enter pending_id={pending_id}")
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 start pending_id={pending_id} "
f"application_target={pending.master} script_path={pending.script_path}"
)
try:
result = run_spark_submit(cmd)
except SparkSubmitError as exc:
pending.status = "FAILED"
pending.error = str(exc)
pending_store.save(pending)
logger.error(f"confirm_submit_job failed pending_id={pending_id} err={exc}")
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 ok pending_id={pending_id} job_id={job_id} "
f"application_id={application_id}"
)
return SubmitResult(
job_id=job_id,
application_id=application_id,
tracking_url=tracking_url,
)
def list_pending_jobs() -> list[dict[str, object]]:
logger.debug("list_pending_jobs enter")
return [p.model_dump() for p in pending_store.list_all()]
def get_pending_job(pending_id: str) -> dict[str, object]:
logger.debug(f"get_pending_job enter pending_id={pending_id}")
p = pending_store.get(pending_id)
if p is None:
raise KeyError(f"Unknown pending_id: {pending_id}")
return p.model_dump()
def cancel_pending_job(pending_id: str) -> dict[str, str]:
logger.debug(f"cancel_pending_job enter pending_id={pending_id}")
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 ok pending_id={pending_id}")
return {"pending_id": pending_id, "status": "CANCELLED"}
+5
View File
@@ -0,0 +1,5 @@
# coding=utf-8
"""
@Time :2026/6/24
@Author :tao.chen
"""
+8
View File
@@ -0,0 +1,8 @@
# 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))
+5
View File
@@ -0,0 +1,5 @@
# coding=utf-8
"""
@Time :2026/6/24
@Author :tao.chen
"""
+178
View File
@@ -0,0 +1,178 @@
# coding=utf-8
from pathlib import Path
import pytest
from fastapi.testclient import TestClient
from spark_executor.core import connection_store, pending_store
from spark_executor.core.connection_store import ConnectionStore
from spark_executor.core.pending_store import PendingStore
from spark_executor.server import app
from spark_executor.tools import connections, submit
@pytest.fixture(autouse=True)
def _fresh_data(tmp_path: Path, monkeypatch):
"""Reset both stores to a fresh tmp dir and rebind the singletons that
the tool modules captured at import time."""
monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path))
monkeypatch.setattr(connection_store, "store", ConnectionStore())
monkeypatch.setattr(pending_store, "DEFAULT_DATA_DIR", str(tmp_path))
monkeypatch.setattr(pending_store, "store", PendingStore())
# Rebind imports inside the tool modules (captured at import time)
connections.store = connection_store.store
submit.conn_store = connection_store.store
submit.pending_store = pending_store.store
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}"
# --- End-to-end body-based calls (the gap the route-registration test missed) ---
def test_save_connection_accepts_dict_spark_conf_in_body():
"""The original query-param signature returned 422 for spark_conf dicts;
body models make tools/call roundtrip cleanly."""
c = TestClient(app)
r = c.post(
"/save_connection",
json={
"name": "prod",
"master": "yarn",
"deploy_mode": "cluster",
"spark_conf": {"spark.sql.shuffle.partitions": "200"},
},
)
assert r.status_code == 200, r.text
assert r.json() == {"name": "prod", "status": "SAVED"}
def test_prepare_submit_job_works_via_body():
c = TestClient(app)
c.post("/save_connection", json={"name": "prod", "master": "yarn"})
r = c.post(
"/prepare_submit_job",
json={"connection": "prod", "script_path": "/tmp/demo.py", "queue": "research"},
)
assert r.status_code == 200, r.text
body = r.json()
assert body["status"] == "PENDING"
assert body["pending_id"].startswith("p_")
assert body["parameters"]["queue"] == "research"
assert body["parameters"]["master"] == "yarn" # snapshotted from connection
def test_list_and_get_pending_job_roundtrip_via_body():
c = TestClient(app)
c.post("/save_connection", json={"name": "prod", "master": "yarn"})
prep = c.post(
"/prepare_submit_job",
json={"connection": "prod", "script_path": "/tmp/a.py"},
).json()
pid = prep["pending_id"]
listed = c.post("/list_pending_jobs", json={}).json()
assert any(p["pending_id"] == pid for p in listed)
got = c.post("/get_pending_job", json={"pending_id": pid}).json()
assert got["script_path"] == "/tmp/a.py"
assert got["status"] == "PENDING"
def test_list_connections_works_with_empty_body():
c = TestClient(app)
r = c.post("/list_connections", json={})
assert r.status_code == 200
assert r.json() == []
# --- Exception handlers: KeyError -> 404, ValueError -> 400 ---
def test_unknown_job_id_returns_404():
c = TestClient(app)
r = c.post("/get_job_status", json={"job_id": "missing"})
assert r.status_code == 404
assert "missing" in r.json()["detail"]
def test_unknown_pending_id_returns_404():
c = TestClient(app)
r = c.post("/get_pending_job", json={"pending_id": "p_nope"})
assert r.status_code == 404
assert "p_nope" in r.json()["detail"]
def test_unknown_connection_name_returns_404():
c = TestClient(app)
r = c.post("/get_connection", json={"name": "nope"})
assert r.status_code == 404
assert "nope" in r.json()["detail"]
def test_unknown_connection_in_prepare_returns_404():
c = TestClient(app)
r = c.post(
"/prepare_submit_job",
json={"connection": "nope", "script_path": "/tmp/x.py"},
)
assert r.status_code == 404
assert "nope" in r.json()["detail"]
def test_confirm_non_pending_returns_400():
"""A CANCELLED pending should refuse confirm; FastAPI should surface ValueError as 400."""
c = TestClient(app)
c.post("/save_connection", json={"name": "prod", "master": "yarn"})
prep = c.post(
"/prepare_submit_job",
json={"connection": "prod", "script_path": "/tmp/x.py"},
).json()
pid = prep["pending_id"]
c.post("/cancel_pending_job", json={"pending_id": pid})
r = c.post("/confirm_submit_job", json={"pending_id": pid})
assert r.status_code == 400
assert "CANCELLED" in r.json()["detail"]
def test_cancel_already_submitted_returns_400():
"""Cancelling a SUBMITTED pending should refuse with ValueError -> 400."""
c = TestClient(app)
c.post("/save_connection", json={"name": "prod", "master": "yarn"})
prep = c.post(
"/prepare_submit_job",
json={"connection": "prod", "script_path": "/tmp/x.py"},
).json()
pid = prep["pending_id"]
# Simulate a SUBMITTED state by mutating the pending directly
p = pending_store.store.get(pid)
p.status = "SUBMITTED"
pending_store.store.save(p)
r = c.post("/cancel_pending_job", json={"pending_id": pid})
assert r.status_code == 400
assert "SUBMITTED" in r.json()["detail"]
+5
View File
@@ -0,0 +1,5 @@
# coding=utf-8
"""
@Time :2026/6/24
@Author :tao.chen
"""
+74
View File
@@ -0,0 +1,74 @@
# 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
+85
View File
@@ -0,0 +1,85 @@
# 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
# --- save_connection ---
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"
# --- list_connections ---
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"}
# --- get_connection ---
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")
# --- delete_connection ---
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")
+35
View File
@@ -0,0 +1,35 @@
# 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),
connection="prod",
)
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"}
+36
View File
@@ -0,0 +1,36 @@
# 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")
+32
View File
@@ -0,0 +1,32 @@
# 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")
+63
View File
@@ -0,0 +1,63 @@
# coding=utf-8
import io
from pathlib import Path
import pytest
from spark_executor.core import connection_store, pending_store
from spark_executor.core.connection_store import ConnectionStore
from spark_executor.core.pending_store import PendingStore
from spark_executor.models import Connection
from spark_executor.tools import connections, submit
@pytest.fixture(autouse=True)
def _fresh_stores(tmp_path: Path, monkeypatch):
monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path))
monkeypatch.setattr(connection_store, "store", ConnectionStore())
monkeypatch.setattr(pending_store, "DEFAULT_DATA_DIR", str(tmp_path))
monkeypatch.setattr(pending_store, "store", PendingStore())
connections.store = connection_store.store
submit.conn_store = connection_store.store
submit.pending_store = pending_store.store
submit.job_store = submit.job_store.__class__() # fresh in-memory job store
@pytest.fixture
def log_capture():
"""Attach an in-memory sink to loguru so tests can assert on emitted lines."""
from common.logging import logger
buf = io.StringIO()
handler_id = logger.add(buf, level="DEBUG", format="{level}|{message}")
yield buf
logger.remove(handler_id)
def test_save_connection_emits_info_log(log_capture):
connections.save_connection(name="prod", master="yarn")
text = log_capture.getvalue()
assert "INFO" in text
assert "save_connection enter" in text
assert "DEBUG" in text
assert "connection saved" in text
def test_prepare_submit_job_emits_debug_and_info(log_capture):
connections.save_connection(name="prod", master="yarn", deploy_mode="cluster")
log_capture.truncate(0); log_capture.seek(0)
submit.prepare_submit_job(connection="prod", script_path="/tmp/j.py", queue="research")
text = log_capture.getvalue()
assert "DEBUG|prepare_submit_job enter" in text
assert "INFO|prepare_submit_job ok" in text
assert "script_path=/tmp/j.py" in text
assert "queue=research" in text
def test_get_unknown_pending_job_emits_debug(log_capture):
log_capture.truncate(0); log_capture.seek(0)
import pytest as _pytest
with _pytest.raises(KeyError):
submit.get_pending_job("p_doesnotexist")
text = log_capture.getvalue()
assert "DEBUG|get_pending_job enter" in text
assert "p_doesnotexist" in text
+47
View File
@@ -0,0 +1,47 @@
# 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")
+87
View File
@@ -0,0 +1,87 @@
# 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="4G",
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="4G",
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"
+78
View File
@@ -0,0 +1,78 @@
# 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="4G",
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
+86
View File
@@ -0,0 +1,86 @@
# coding=utf-8
from unittest.mock import MagicMock, patch
import pytest
from spark_executor.core.spark_submit import (
SparkSubmitError,
build_spark_submit_command,
run_spark_submit,
)
def test_build_command_uses_provided_master_and_deploy_mode():
cmd = build_spark_submit_command(
master="yarn",
deploy_mode="cluster",
script_path="/tmp/jobs/job_001.py",
queue="default",
executor_memory="4G",
executor_cores=2,
num_executors=2,
)
assert cmd[:2] == ["spark-submit", "--master"]
assert "yarn" in cmd
assert "cluster" in cmd
assert "--queue" in cmd and "default" in cmd
assert "--executor-memory" in cmd and "4G" in cmd
assert "--executor-cores" in cmd and "2" in cmd
assert "--num-executors" in cmd and "2" in cmd
assert cmd[-1] == "/tmp/jobs/job_001.py"
def test_build_command_supports_standalone_master():
cmd = build_spark_submit_command(
master="spark://master:7077",
deploy_mode="client",
script_path="/tmp/j.py",
queue="default",
executor_memory="1G",
executor_cores=1,
num_executors=1,
)
assert "spark://master:7077" in cmd
assert "client" in cmd
def test_build_command_appends_spark_conf_entries():
cmd = build_spark_submit_command(
master="yarn",
deploy_mode="cluster",
script_path="/tmp/j.py",
queue="default",
executor_memory="4G",
executor_cores=2,
num_executors=2,
spark_conf={"spark.sql.shuffle.partitions": "200", "spark.executor.memoryOverhead": "1G"},
)
# every k=v is emitted as --conf k=v
assert "--conf" in cmd
assert "spark.sql.shuffle.partitions=200" in cmd
assert "spark.executor.memoryOverhead=1G" in cmd
# script_path is still last
assert cmd[-1] == "/tmp/j.py"
def test_run_spark_submit_returns_completed_process(monkeypatch):
fake = MagicMock()
fake.returncode = 0
fake.stderr = "Submitted application application_1\n"
with patch("spark_executor.core.spark_submit.subprocess.run", return_value=fake) as m:
result = run_spark_submit(["spark-submit", "/tmp/x.py"])
assert result is fake
m.assert_called_once()
# capture_output=True and text=True must be set
kwargs = m.call_args.kwargs
assert kwargs["capture_output"] is True
assert kwargs["text"] is True
def test_run_spark_submit_raises_on_nonzero_return():
fake = MagicMock()
fake.returncode = 1
fake.stderr = "boom"
with patch("spark_executor.core.spark_submit.subprocess.run", return_value=fake):
with pytest.raises(SparkSubmitError, match="spark-submit failed"):
run_spark_submit(["spark-submit", "/tmp/x.py"])
+36
View File
@@ -0,0 +1,36 @@
# 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")
+190
View File
@@ -0,0 +1,190 @@
# coding=utf-8
from pathlib import Path
from unittest.mock import patch
import pytest
from spark_executor.core import connection_store, pending_store
from spark_executor.core.pending_store import PendingStore
from spark_executor.models import Connection
from spark_executor.tools import submit
@pytest.fixture(autouse=True)
def _fresh(tmp_path: Path, monkeypatch):
monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path))
monkeypatch.setattr(connection_store, "store", connection_store.ConnectionStore())
monkeypatch.setattr(pending_store, "DEFAULT_DATA_DIR", str(tmp_path))
monkeypatch.setattr(pending_store, "store", PendingStore())
submit.conn_store = connection_store.store
submit.pending_store = pending_store.store
connection_store.store.save(Connection(name="prod", master="yarn", deploy_mode="cluster"))
def _last_pending_id() -> str:
pending = submit.pending_store.list_all()
assert pending, "no pending submission was created"
return pending[-1].pending_id
# --- prepare_submit_job ---
def test_prepare_does_not_invoke_spark_submit(monkeypatch):
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"
# --- confirm_submit_job ---
def test_confirm_invokes_spark_submit_and_marks_submitted(monkeypatch):
submit.prepare_submit_job(connection="prod", script_path="/tmp/j.py")
pid = _last_pending_id()
fake_proc = type("P", (), {
"returncode": 0,
"stderr": "tracking URL: http://rm:8088/proxy/application_17400000001/\n",
})()
with patch("spark_executor.tools.submit.run_spark_submit", return_value=fake_proc) as m:
result = submit.confirm_submit_job(pending_id=pid)
cmd = m.call_args.args[0]
assert "yarn" in cmd
assert "cluster" in cmd
assert cmd[-1] == "/tmp/j.py"
assert result.application_id == "application_17400000001"
# pending updated
p = submit.pending_store.get(pid)
assert p.status == "SUBMITTED"
assert p.application_id == "application_17400000001"
assert p.job_id is not None
def test_confirm_raises_for_unknown_pending_id():
with pytest.raises(KeyError, match="missing"):
submit.confirm_submit_job(pending_id="missing")
def test_confirm_refuses_non_pending_status(monkeypatch):
submit.prepare_submit_job(connection="prod", script_path="/tmp/j.py")
pid = _last_pending_id()
# Mark it CANCELLED first
p = submit.pending_store.get(pid)
p.status = "CANCELLED"
submit.pending_store.save(p)
with pytest.raises(ValueError, match="CANCELLED"):
submit.confirm_submit_job(pending_id=pid)
def test_confirm_marks_failed_on_spark_submit_error(monkeypatch):
from spark_executor.core.spark_submit import SparkSubmitError
submit.prepare_submit_job(connection="prod", script_path="/tmp/j.py")
pid = _last_pending_id()
def _raise(_cmd):
raise SparkSubmitError("boom")
monkeypatch.setattr(submit, "run_spark_submit", _raise)
with pytest.raises(SparkSubmitError):
submit.confirm_submit_job(pending_id=pid)
p = submit.pending_store.get(pid)
assert p.status == "FAILED"
assert "boom" in (p.error or "")
# --- list_pending_jobs ---
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)
# --- get_pending_job ---
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")
# --- cancel_pending_job ---
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)
+63
View File
@@ -0,0 +1,63 @@
# 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")
Generated
+67
View File
@@ -349,6 +349,15 @@ wheels = [
{ url = "https://pypi.tuna.tsinghua.edu.cn/packages/1e/5e/d4e9f1a599fb8e573b7b87160658329fbf28d19eac2718f51fc3def3aa5a/idna-3.18-py3-none-any.whl", hash = "sha256:7f952cbe720b688055e3f87de14f5c3e5fdaa8bc3928985c4077ca689de849a2", size = 65455 },
]
[[package]]
name = "iniconfig"
version = "2.3.0"
source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple/" }
sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/72/34/14ca021ce8e5dfedc35312d08ba8bf51fdd999c576889fc2c24cb97f4f10/iniconfig-2.3.0.tar.gz", hash = "sha256:c76315c77db068650d49c5b56314774a7804df16fee4402c1f19d6d15d8c4730", size = 20503 }
wheels = [
{ url = "https://pypi.tuna.tsinghua.edu.cn/packages/cb/b1/3846dd7f199d53cb17f49cba7e651e9ce294d8497c8c150530ed11865bb8/iniconfig-2.3.0-py3-none-any.whl", hash = "sha256:f631c04d2c48c52b84d0d0549c99ff3859c98df65b3101406327ecc7d53fbf12", size = 7484 },
]
[[package]]
name = "jsonschema"
version = "4.26.0"
@@ -435,6 +444,24 @@ wheels = [
{ url = "https://pypi.tuna.tsinghua.edu.cn/packages/b3/38/89ba8ad64ae25be8de66a6d463314cf1eb366222074cfda9ee839c56a4b4/mdurl-0.1.2-py3-none-any.whl", hash = "sha256:84008a41e51615a49fc9966191ff91509e3c40b939176e643fd50a5c2196b8f8", size = 9979 },
]
[[package]]
name = "packaging"
version = "26.2"
source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple/" }
sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/d7/f1/e7a6dd94a8d4a5626c03e4e99c87f241ba9e350cd9e6d75123f992427270/packaging-26.2.tar.gz", hash = "sha256:ff452ff5a3e828ce110190feff1178bb1f2ea2281fa2075aadb987c2fb221661", size = 228134 }
wheels = [
{ url = "https://pypi.tuna.tsinghua.edu.cn/packages/df/b2/87e62e8c3e2f4b32e5fe99e0b86d576da1312593b39f47d8ceef365e95ed/packaging-26.2-py3-none-any.whl", hash = "sha256:5fc45236b9446107ff2415ce77c807cee2862cb6fac22b8a73826d0693b0980e", size = 100195 },
]
[[package]]
name = "pluggy"
version = "1.6.0"
source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple/" }
sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/f9/e2/3e91f31a7d2b083fe6ef3fa267035b518369d9511ffab804f839851d2779/pluggy-1.6.0.tar.gz", hash = "sha256:7dcc130b76258d33b90f61b658791dede3486c3e6bfb003ee5c9bfb396dd22f3", size = 69412 }
wheels = [
{ url = "https://pypi.tuna.tsinghua.edu.cn/packages/54/20/4d324d65cc6d9205fabedc306948156824eb9f0ee1633355a8f7ec5c66bf/pluggy-1.6.0-py3-none-any.whl", hash = "sha256:e920276dd6813095e9377c0bc5566d94c932c33b27a3e3945d8389c374dd4746", size = 20538 },
]
[[package]]
name = "pycparser"
version = "3.0"
@@ -567,6 +594,34 @@ crypto = [
{ name = "cryptography" },
]
[[package]]
name = "pytest"
version = "9.1.1"
source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple/" }
dependencies = [
{ name = "colorama", marker = "sys_platform == 'win32'" },
{ name = "iniconfig" },
{ name = "packaging" },
{ name = "pluggy" },
{ name = "pygments" },
]
sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/e4/47/b9efed96c114afcfa3c9d3fe98a76a1d14c74a9e266d397cf6eb64be5e01/pytest-9.1.1.tar.gz", hash = "sha256:1088fbde8f2b49d95a549a195707afa7a76a3ce9bcadc26b6d71f0ffda5fe313", size = 1636369 }
wheels = [
{ url = "https://pypi.tuna.tsinghua.edu.cn/packages/24/25/1de2678b631f5a49215c6c96fff41ba892b0a34df68d6d80292b1b48aa7f/pytest-9.1.1-py3-none-any.whl", hash = "sha256:37a86b45efb9a47a61a36449063e8e18d0cab3161329fc099eb21783169c4f0c", size = 386536 },
]
[[package]]
name = "pytest-mock"
version = "3.15.1"
source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple/" }
dependencies = [
{ name = "pytest" },
]
sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/68/14/eb014d26be205d38ad5ad20d9a80f7d201472e08167f0bb4361e251084a9/pytest_mock-3.15.1.tar.gz", hash = "sha256:1849a238f6f396da19762269de72cb1814ab44416fa73a8686deac10b0d87a0f", size = 34036 }
wheels = [
{ url = "https://pypi.tuna.tsinghua.edu.cn/packages/5a/cc/06253936f4a7fa2e0f48dfe6d851d9c56df896a9ab09ac019d70b760619c/pytest_mock-3.15.1-py3-none-any.whl", hash = "sha256:0a25e2eb88fe5168d535041d09a4529a188176ae608a6d249ee65abc0949630d", size = 10095 },
]
[[package]]
name = "python-dotenv"
version = "1.2.2"
@@ -779,6 +834,12 @@ dependencies = [
{ name = "uvicorn" },
]
[package.dev-dependencies]
dev = [
{ name = "pytest" },
{ name = "pytest-mock" },
]
[package.metadata]
requires-dist = [
{ name = "fastapi", specifier = ">=0.138.0" },
@@ -790,6 +851,12 @@ requires-dist = [
{ name = "uvicorn", specifier = ">=0.49.0" },
]
[package.metadata.requires-dev]
dev = [
{ name = "pytest", specifier = ">=8.0" },
{ name = "pytest-mock", specifier = ">=3.14" },
]
[[package]]
name = "sse-starlette"
version = "3.4.5"