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:
+22
@@ -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
@@ -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",
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
# coding=utf-8
|
||||
"""
|
||||
@Time :2026/6/24
|
||||
@Author :tao.chen
|
||||
"""
|
||||
@@ -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()
|
||||
@@ -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())
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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")
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
# coding=utf-8
|
||||
"""
|
||||
@Time :2026/6/24
|
||||
@Author :tao.chen
|
||||
"""
|
||||
@@ -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"}
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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"}
|
||||
@@ -0,0 +1,5 @@
|
||||
# coding=utf-8
|
||||
"""
|
||||
@Time :2026/6/24
|
||||
@Author :tao.chen
|
||||
"""
|
||||
@@ -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))
|
||||
@@ -0,0 +1,5 @@
|
||||
# coding=utf-8
|
||||
"""
|
||||
@Time :2026/6/24
|
||||
@Author :tao.chen
|
||||
"""
|
||||
@@ -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"]
|
||||
@@ -0,0 +1,5 @@
|
||||
# coding=utf-8
|
||||
"""
|
||||
@Time :2026/6/24
|
||||
@Author :tao.chen
|
||||
"""
|
||||
@@ -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
|
||||
@@ -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")
|
||||
@@ -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"}
|
||||
@@ -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")
|
||||
@@ -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")
|
||||
@@ -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
|
||||
@@ -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")
|
||||
@@ -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"
|
||||
@@ -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
|
||||
@@ -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"])
|
||||
@@ -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")
|
||||
@@ -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)
|
||||
@@ -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")
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user