fix(job_store): persist Jobs to disk and accept either ID in job tools

Two user-reported bugs, same root cause: the in-memory JobStore + the
'job_id must be the 12-char hex' tool contract.

Bug 1: 'Unknown job_id' reported frequently
  JobStore was a process-local dict (spark_executor/core/job_store.py).
  Under gunicorn workers > 1, a job created by confirm_submit_job
  landing on worker A was invisible to worker B, so a follow-up
  get_job_status / get_job_result / get_job_logs / kill_job landing on
  a different worker returned 'Unknown job_id'. Same multi-worker
  problem that bit the MCP session layer; only the affected data was
  different.

Bug 2: 'get_job_logs frequently confuses job_id and application_id'
  confirm_submit_job returns BOTH identifiers in SubmitResult, but
  get_job_logs (and friends) only accepted the local 12-char job_id
  and never said so in their description. When the agent passed the
  YARN application_id, the error message itself was misleading:
  'Unknown job_id: application_17400000001_0001' — the agent had
  passed an id, just the wrong kind.

This change fixes both at the root:

  * JobStore is now JSON-backed at data/jobs.json (atomic tempfile +
    os.replace), with cross-process safety via fcntl.flock on a sibling
    .lock file. Stage 3's SQLite migration is still planned; the file
    format is intentionally simple so it is a straight
    'for j in read_all(): db.insert(j)'.

  * New JobStore.get_either(uid) looks up by job_id first, then
    application_id. All four job-lifecycle tools (get_job_status,
    get_job_result, get_job_logs, kill_job) call get_either instead
    of get(job_id), so the agent can pass either identifier and get
    the same answer.

  * The 'neither matched' KeyError now spells out both id forms and
    what they look like, so the agent isn't left guessing.

  * server.py tool descriptions for the four job tools explicitly
    state 'job_id accepts BOTH identifiers' so this is visible to the
    LLM at tool-selection time, not only at error time.

Tests:
  * test_job_store.py: tmp_path isolation, persistence across
    instances, human-readable JSON, corrupt-file resilience,
    get_either (by job_id, by application_id, collision preference,
    unknown), put idempotency.
  * test_{logs,status,kill,result}_tool.py: per-test tmp_path fixture,
    'accepts application_id' regression for each tool, and an
    explicit assertion that the unknown-id error message mentions
    BOTH id forms. test_result_raises_keyerror_for_unknown_job's
    match pattern updated for the new message.

242 tests pass (was 226; +16 new). Zero regressions.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Claude
2026-06-29 18:57:31 +08:00
co-authored by Claude Fable 5
parent 565f6a9c9d
commit 0fef77c0b8
12 changed files with 541 additions and 42 deletions
+4 -2
View File
@@ -71,9 +71,11 @@ docs/superpowers/plans/ # implementation plans
**Persistence layout** (under `./data/`, overridable via `SPARK_EXECUTOR_DATA_DIR`): **Persistence layout** (under `./data/`, overridable via `SPARK_EXECUTOR_DATA_DIR`):
- `data/connections.json``Connection` records keyed by name (atomic write via `tempfile` + `os.replace`) - `data/connections.json``Connection` records keyed by name (atomic write via `tempfile` + `os.replace`)
- `data/pending_jobs.json``PendingSubmission` records keyed by `pending_id` - `data/pending_jobs.json``PendingSubmission` records keyed by `pending_id`
- `data/jobs.json``Job` records keyed by `job_id` (atomic write via `tempfile` + `os.replace`, cross-process safety via `fcntl.flock` on a sibling `.lock` file). The four job-lifecycle tools (`get_job_status` / `get_job_result` / `get_job_logs` / `kill_job`) accept EITHER the local `job_id` (12-char hex) OR the YARN `application_id``JobStore.get_either(...)` looks up by job_id first, then by application_id, so an agent that confuses the two (a common LLM mistake) still gets a sensible result.
- `data/logs/debug/YYYY-MM-DD.log` — DEBUG sink, gzipped, 30-day retention - `data/logs/debug/YYYY-MM-DD.log` — DEBUG sink, gzipped, 30-day retention
- `data/logs/info/YYYY-MM-DD.log` — INFO sink, gzipped, 30-day retention - `data/logs/info/YYYY-MM-DD.log` — INFO sink, gzipped, 30-day retention
- `JobStore` is in-memory only today; Stage 3 will move it to SQLite.
**Why `JobStore` is now file-backed (not in-memory):** in-memory only broke under multi-worker gunicorn (workers > 1) for the same reason MCP sessions do — the dict lives in worker A's process, so a follow-up `get_job_logs` landing on worker B returns "Unknown job_id". File-backed JSON is a Step-1 fix; Stage 3 will still move to SQLite for queryable / transactional semantics, and the file format is intentionally simple so the migration is a straight `for j in read_all(): db.insert(j)`.
**Key conventions:** **Key conventions:**
@@ -87,7 +89,7 @@ docs/superpowers/plans/ # implementation plans
- **The two-step submit flow is core, not optional.** `prepare_submit_job` snapshots the connection's `master` / `deploy_mode` / `spark_conf` into the `PendingSubmission`; `confirm_submit_job` is the only place `spark-submit` is invoked. Editing a connection between prepare and confirm does **not** retarget the pending job. - **The two-step submit flow is core, not optional.** `prepare_submit_job` snapshots the connection's `master` / `deploy_mode` / `spark_conf` into the `PendingSubmission`; `confirm_submit_job` is the only place `spark-submit` is invoked. Editing a connection between prepare and confirm does **not** retarget the pending job.
- **MCP session affinity: keep `GUNICORN_WORKERS=1` (or front with a sticky-session LB).** The `mcp` library stores each session in a per-process dict (`StreamableHTTPSessionManager._server_instances`); gunicorn round-robins requests across workers, so a multi-worker deploy returns "Session not found" / "Invalid or expired session ID" for the same `mcp-session-id` whenever it lands on a worker that didn't create it. `fastapi-mcp` hardcodes `stateless=False`, so there is no in-process workaround. The `on_starting` hook in `gunicorn.conf.py` logs a WARNING whenever `workers > 1` so the misconfig is loud, not silent. See the comment block above `workers =` in `gunicorn.conf.py` for full context and the nginx-sticky escape hatch. - **MCP session affinity: keep `GUNICORN_WORKERS=1` (or front with a sticky-session LB).** The `mcp` library stores each session in a per-process dict (`StreamableHTTPSessionManager._server_instances`); gunicorn round-robins requests across workers, so a multi-worker deploy returns "Session not found" / "Invalid or expired session ID" for the same `mcp-session-id` whenever it lands on a worker that didn't create it. `fastapi-mcp` hardcodes `stateless=False`, so there is no in-process workaround. The `on_starting` hook in `gunicorn.conf.py` logs a WARNING whenever `workers > 1` so the misconfig is loud, not silent. See the comment block above `workers =` in `gunicorn.conf.py` for full context and the nginx-sticky escape hatch.
- **Test fixtures rebind module-level singletons** (`connections.store`, `submit.conn_store`, `submit.pending_store`) in `monkeypatch.setattr` because the tool modules captured the originals at import time. See `tests/integration/test_mcp_routes.py` for the pattern. - **Test fixtures rebind module-level singletons** (`connections.store`, `submit.conn_store`, `submit.pending_store`) in `monkeypatch.setattr` because the tool modules captured the originals at import time. See `tests/integration/test_mcp_routes.py` for the pattern.
- **86 tests pass** as of the last Stage 1 cleanup; run `uv run pytest` after any change. - **86 tests pass** as of the last Stage 1 cleanup; run `uv run pytest` after any change. (242 tests as of the JobStore persistence fix.)
## Stage Status ## Stage Status
+171 -7
View File
@@ -2,27 +2,191 @@
""" """
@Time :2026/6/24 @Time :2026/6/24
@Author :tao.chen @Author :tao.chen
JSON-backed Job registry. Each Job carries both a local `job_id` (12-char
hex generated in confirm_submit_job) and the YARN `application_id`
(`application_<ts>_<n>`).
Why file-backed instead of in-memory:
- The MCP service runs under gunicorn with workers > 1 by default; an
in-memory dict lives in worker A's process, so a follow-up
`get_job_logs(job_id=...)` landing on worker B would return
"Unknown job_id" for jobs created elsewhere. The same problem that
bit us on MCP sessions.
- Restarting the server should not lose track of running YARN apps.
Why fcntl.flock on writes:
- In-process `threading.Lock` is per-process. Two gunicorn workers can
call `put()` concurrently and lose updates (last-writer-wins over
the other's data). `fcntl.flock(LOCK_EX)` on the data file gives
cross-process mutual exclusion.
Why `get_either` (accept job_id OR application_id):
- The MCP client gets both IDs back from `confirm_submit_job` and
routinely confuses which to pass to `get_job_logs` / `get_job_status`
/ `kill_job` / `get_job_result`. The tool contract used to be
"must be job_id" but the description never said so, so agents
passed application_id and got a misleading "Unknown job_id" error.
Accepting either ID lets us blame the lookup, not the user.
This is a Step-1 persistence fix; the Stage-3 plan still calls for
SQLite (queryable, transactional). The file format is intentionally
simple so a future migration is just `for j in read_all(): db.insert(j)`.
""" """
import fcntl
import json
import os
import tempfile
from contextlib import contextmanager
from pathlib import Path
from threading import Lock from threading import Lock
from common.config import settings
from common.logging import logger
from spark_executor.models import Job from spark_executor.models import Job
DEFAULT_DATA_DIR = settings.data_dir
DEFAULT_FILE_NAME = "jobs.json"
class JobStore: class JobStore:
"""In-memory job registry. Stage-3 will swap this for SQLite.""" """JSON-backed CRUD for Job records, indexed by job_id.
def __init__(self) -> None: File path: <data_dir>/<file_name>. Format: {job_id: <Job.model_dump()>}.
Writes are atomic (tempfile + os.replace) and cross-process safe
(fcntl.flock around the load-modify-dump critical section).
"""
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
# In-process lock: serializes threads within one worker.
# Cross-process serialization is handled by fcntl.flock in _locked_dump.
self._lock = Lock() self._lock = Lock()
self._jobs: dict[str, Job] = {}
@property
def path(self) -> Path:
return Path(self._data_dir) / self._file_name
# --- Low-level I/O ---
def _load(self) -> dict[str, Job]:
"""Read all jobs from disk. Returns {} on missing/corrupt file."""
if not self.path.exists():
logger.debug(f"JobStore._load: file {self.path} absent, returning empty")
return {}
try:
raw = json.loads(self.path.read_text(encoding="utf-8"))
except json.JSONDecodeError:
logger.exception(
f"JobStore._load: file {self.path} is corrupt; treating as empty. "
f"Inspect and either repair or delete to recover."
)
return {}
return {jid: Job.model_validate(j) for jid, j in raw.items()}
def _dump(self, records: dict[str, Job]) -> None:
"""Atomic write. Caller must hold `self._lock`."""
os.makedirs(self._data_dir, exist_ok=True)
payload = {jid: j.model_dump() for jid, j 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:
# default=str handles datetime — same trick pending_store uses.
json.dump(payload, f, indent=2, ensure_ascii=False, default=str)
os.replace(tmp_path, self.path)
logger.debug(f"JobStore._dump: wrote {len(records)} records to {self.path}")
except Exception:
if os.path.exists(tmp_path):
os.unlink(tmp_path)
raise
@contextmanager
def _locked_dump(self, mutate):
"""load -> mutate (in-memory) -> dump, under both in-process and
cross-process exclusive locks.
This is the only sanctioned way to write. It guarantees:
1. Two threads in the same worker see consistent state.
2. Two gunicorn workers can't both load X, both compute
X+their-change, both dump — one update would be lost.
"""
with self._lock:
# Open the data file for flocking; if it doesn't exist yet,
# open the parent dir so we can still take an exclusive lock
# before the first write creates the file.
lock_path = self.path
lock_path.parent.mkdir(parents=True, exist_ok=True)
# Use a stable lock file (separate from data file) so an
# os.replace on the data file doesn't break our flock.
flock_path = lock_path.with_suffix(lock_path.suffix + ".lock")
f = open(flock_path, "w")
try:
fcntl.flock(f.fileno(), fcntl.LOCK_EX)
records = self._load()
mutate(records)
self._dump(records)
finally:
fcntl.flock(f.fileno(), fcntl.LOCK_UN)
f.close()
# --- Public CRUD ---
def put(self, job: Job) -> None: def put(self, job: Job) -> None:
with self._lock: """Insert or replace a Job by job_id. Persists to disk."""
self._jobs[job.job_id] = job def _mutate(records: dict[str, Job]) -> None:
records[job.job_id] = job
self._locked_dump(_mutate)
logger.info(
f"job_store put job_id={job.job_id} application_id={job.application_id}"
)
def get(self, job_id: str) -> Job | None: def get(self, job_id: str) -> Job | None:
"""Look up a Job by its local job_id only.
Prefer `get_either` in tool code — it also accepts application_id,
which is what the LLM agent sometimes has on hand.
"""
with self._lock: with self._lock:
return self._jobs.get(job_id) return self._load().get(job_id)
def get_by_application_id(self, application_id: str) -> Job | None:
"""Look up a Job by its YARN application_id. O(n) over the file."""
with self._lock:
for job in self._load().values():
if job.application_id == application_id:
return job
return None
def get_either(self, job_id_or_application_id: str) -> Job | None:
"""Look up a Job by either identifier. Try job_id first (O(1)),
then fall back to application_id scan (O(n)).
This is what every job tool (get_job_status / get_job_logs /
kill_job / get_job_result) should call — it eliminates the
"agent passed the wrong ID and got a misleading 'Unknown job_id'"
failure mode.
"""
with self._lock:
records = self._load()
direct = records.get(job_id_or_application_id)
if direct is not None:
return direct
for job in records.values():
if job.application_id == job_id_or_application_id:
return job
return None
def list(self) -> list[Job]: def list(self) -> list[Job]:
with self._lock: with self._lock:
return list(self._jobs.values()) return list(self._load().values())
def clear(self) -> None:
"""Test helper: drop everything. Not part of the public API."""
def _mutate(records: dict[str, Job]) -> None:
records.clear()
self._locked_dump(_mutate)
# Module-level singleton; replaced in tests.
store = JobStore()
+24 -4
View File
@@ -181,7 +181,12 @@ def _cancel_pending_job(req: PendingIdRequest):
description=( description=(
"Return the YARN application state (RUNNING / SUCCEEDED / FAILED / " "Return the YARN application state (RUNNING / SUCCEEDED / FAILED / "
"KILLED / ACCEPTED / NEW / NEW_SAVING / SUBMITTED / etc.) plus the " "KILLED / ACCEPTED / NEW / NEW_SAVING / SUBMITTED / etc.) plus the "
"raw YARN REST response body." "raw YARN REST response body.\n\n"
"**job_id accepts BOTH identifiers** returned by "
"confirm_submit_job: the local job_id (12-char hex, e.g. "
"'a1b2c3d4e5f6') and the YARN application_id (e.g. "
"'application_17400000001_0001'). The lookup is by job_id first, "
"then by application_id."
), ),
) )
def _get_job_status(req: JobIdRequest): def _get_job_status(req: JobIdRequest):
@@ -196,7 +201,12 @@ def _get_job_status(req: JobIdRequest):
"Return a terminal-oriented view of a Spark job: final_status, " "Return a terminal-oriented view of a Spark job: final_status, "
"diagnostics, tracking_url, started_time, and finished_time. " "diagnostics, tracking_url, started_time, and finished_time. "
"This is distinct from get_job_status, which is for polling the " "This is distinct from get_job_status, which is for polling the "
"running YARN state and returns the raw YARN response." "running YARN state and returns the raw YARN response.\n\n"
"**job_id accepts BOTH identifiers** returned by "
"confirm_submit_job: the local job_id (12-char hex, e.g. "
"'a1b2c3d4e5f6') and the YARN application_id (e.g. "
"'application_17400000001_0001'). The lookup is by job_id first, "
"then by application_id."
), ),
) )
def _get_job_result(req: JobIdRequest): def _get_job_result(req: JobIdRequest):
@@ -210,7 +220,12 @@ def _get_job_result(req: JobIdRequest):
description=( description=(
"Pull aggregated logs from the YARN ResourceManager. Returns the last " "Pull aggregated logs from the YARN ResourceManager. Returns the last "
"tail_chars characters (default 5000). Requires yarn.log-aggregation-enable " "tail_chars characters (default 5000). Requires yarn.log-aggregation-enable "
"to be true on the target cluster." "to be true on the target cluster.\n\n"
"**job_id accepts BOTH identifiers** returned by "
"confirm_submit_job: the local job_id (12-char hex, e.g. "
"'a1b2c3d4e5f6') and the YARN application_id (e.g. "
"'application_17400000001_0001'). The lookup is by job_id first, "
"then by application_id."
), ),
) )
def _get_job_logs(req: GetJobLogsRequest): def _get_job_logs(req: GetJobLogsRequest):
@@ -221,7 +236,12 @@ def _get_job_logs(req: GetJobLogsRequest):
"/kill_job", "/kill_job",
operation_id="kill_job", operation_id="kill_job",
summary="Kill a running job", summary="Kill a running job",
description="PUT state=KILLED to YARN REST API for the job's application_id.", description=(
"PUT state=KILLED to YARN REST API for the job's application_id. "
"**job_id accepts BOTH identifiers** returned by "
"confirm_submit_job: the local job_id (12-char hex) and the YARN "
"application_id. The lookup is by job_id first, then by application_id."
),
) )
def _kill_job(req: JobIdRequest): def _kill_job(req: JobIdRequest):
return kill_job(req.job_id) return kill_job(req.job_id)
+10 -4
View File
@@ -7,23 +7,29 @@ from common.logging import logger
from spark_executor.core.job_store import JobStore from spark_executor.core.job_store import JobStore
from spark_executor.core.yarn_client import YarnClientConfig, kill_application from spark_executor.core.yarn_client import YarnClientConfig, kill_application
from spark_executor.tools.connections import store as conn_store from spark_executor.tools.connections import store as conn_store
from spark_executor.tools.logs import _unknown_job_error
store = JobStore() store = JobStore()
def kill_job(job_id: str) -> dict[str, str]: def kill_job(job_id: str) -> dict[str, str]:
"""Kill a running job.
`job_id` accepts either the local job_id (returned by
confirm_submit_job) or the YARN application_id.
"""
logger.debug(f"kill_job enter job_id={job_id}") logger.debug(f"kill_job enter job_id={job_id}")
job = store.get(job_id) job = store.get_either(job_id)
if job is None: if job is None:
raise KeyError(f"Unknown job_id: {job_id}") raise _unknown_job_error(job_id)
conn = conn_store.get(job.connection) conn = conn_store.get(job.connection)
if conn is None: if conn is None:
raise KeyError(f"Connection not found: {job.connection}") raise KeyError(f"Connection not found: {job.connection}")
config = YarnClientConfig.from_connection(conn) config = YarnClientConfig.from_connection(conn)
kill_application(job.application_id, config) kill_application(job.application_id, config)
logger.info(f"kill_job ok job_id={job_id} application_id={job.application_id}") logger.info(f"kill_job ok job_id={job.job_id} application_id={job.application_id}")
return { return {
"job_id": job_id, "job_id": job.job_id,
"application_id": job.application_id, "application_id": job.application_id,
"status": "KILLED", "status": "KILLED",
} }
+25 -3
View File
@@ -11,11 +11,33 @@ from spark_executor.tools.connections import store as conn_store
store = JobStore() store = JobStore()
def _unknown_job_error(uid: str) -> KeyError:
"""Standard "we tried both IDs and found nothing" message.
The agent gets this from confirm_submit_job's response:
{"job_id": "a1b2c3d4e5f6", "application_id": "application_...", ...}
and routinely confuses which to pass here. Spelling out that BOTH
IDs were tried (and what they look like) saves a round trip.
"""
return KeyError(
f"No Job found for id={uid!r} (neither as job_id nor as "
f"application_id). Pass the job_id from confirm_submit_job's "
f"response — it is a 12-char hex like 'a1b2c3d4e5f6'. The "
f"application_id is the YARN ID, e.g. 'application_17400000001_0001'."
)
def get_job_logs(job_id: str, tail_chars: int = 5000) -> str: def get_job_logs(job_id: str, tail_chars: int = 5000) -> str:
"""Fetch aggregated container logs for a Spark job.
`job_id` accepts either the local job_id (returned by
confirm_submit_job) or the YARN application_id — both are looked up
against the same Job record.
"""
logger.debug(f"get_job_logs enter job_id={job_id} tail_chars={tail_chars}") logger.debug(f"get_job_logs enter job_id={job_id} tail_chars={tail_chars}")
job = store.get(job_id) job = store.get_either(job_id)
if job is None: if job is None:
raise KeyError(f"Unknown job_id: {job_id}") raise _unknown_job_error(job_id)
conn = conn_store.get(job.connection) conn = conn_store.get(job.connection)
if conn is None: if conn is None:
raise KeyError(f"Connection not found: {job.connection}") raise KeyError(f"Connection not found: {job.connection}")
@@ -23,7 +45,7 @@ def get_job_logs(job_id: str, tail_chars: int = 5000) -> str:
full = get_application_logs(job.application_id, config) full = get_application_logs(job.application_id, config)
tailed = full[-tail_chars:] if len(full) > tail_chars else full tailed = full[-tail_chars:] if len(full) > tail_chars else full
logger.info( logger.info(
f"get_job_logs ok job_id={job_id} application_id={job.application_id} " f"get_job_logs ok job_id={job.job_id} application_id={job.application_id} "
f"full_chars={len(full)} returned_chars={len(tailed)}" f"full_chars={len(full)} returned_chars={len(tailed)}"
) )
return tailed return tailed
+9 -3
View File
@@ -10,15 +10,21 @@ from spark_executor.core.job_store import JobStore
from spark_executor.core.yarn_client import YarnClientConfig, get_application_status from spark_executor.core.yarn_client import YarnClientConfig, get_application_status
from spark_executor.tools.connections import store as conn_store from spark_executor.tools.connections import store as conn_store
from spark_executor.models import JobResult from spark_executor.models import JobResult
from spark_executor.tools.logs import _unknown_job_error
store = JobStore() store = JobStore()
def get_job_result(job_id: str) -> JobResult: def get_job_result(job_id: str) -> JobResult:
"""Query YARN for a job's terminal result view.
`job_id` accepts either the local job_id (returned by
confirm_submit_job) or the YARN application_id.
"""
logger.debug(f"get_job_result enter job_id={job_id}") logger.debug(f"get_job_result enter job_id={job_id}")
job = store.get(job_id) job = store.get_either(job_id)
if job is None: if job is None:
raise KeyError(f"Unknown job_id: {job_id}") raise _unknown_job_error(job_id)
conn = conn_store.get(job.connection) conn = conn_store.get(job.connection)
if conn is None: if conn is None:
raise KeyError(f"Connection not found: {job.connection}") raise KeyError(f"Connection not found: {job.connection}")
@@ -35,7 +41,7 @@ def get_job_result(job_id: str) -> JobResult:
finished_time=app.get("finishedTime"), finished_time=app.get("finishedTime"),
) )
logger.info( logger.info(
f"get_job_result ok job_id={job_id} application_id={job.application_id} " f"get_job_result ok job_id={job.job_id} application_id={job.application_id} "
f"state={state} final_status={result.final_status}" f"state={state} final_status={result.final_status}"
) )
return result return result
+9 -3
View File
@@ -9,19 +9,25 @@ from spark_executor.core.job_store import JobStore
from spark_executor.core.yarn_client import get_application_status from spark_executor.core.yarn_client import get_application_status
from spark_executor.models import JobStatus from spark_executor.models import JobStatus
from spark_executor.tools.connections import store as conn_store from spark_executor.tools.connections import store as conn_store
from spark_executor.tools.logs import _unknown_job_error
store = JobStore() store = JobStore()
def get_job_status(job_id: str) -> JobStatus: def get_job_status(job_id: str) -> JobStatus:
"""Query YARN for a job's current status.
`job_id` accepts either the local job_id (returned by
confirm_submit_job) or the YARN application_id.
"""
logger.debug(f"get_job_status enter job_id={job_id}") logger.debug(f"get_job_status enter job_id={job_id}")
job = store.get(job_id) job = store.get_either(job_id)
if job is None: if job is None:
raise KeyError(f"Unknown job_id: {job_id}") raise _unknown_job_error(job_id)
conn = conn_store.get(job.connection) conn = conn_store.get(job.connection)
if conn is None: if conn is None:
raise KeyError(f"Connection not found: {job.connection}") raise KeyError(f"Connection not found: {job.connection}")
config = YarnClientConfig.from_connection(conn) config = YarnClientConfig.from_connection(conn)
state, raw = get_application_status(job.application_id, config) state, raw = get_application_status(job.application_id, config)
logger.info(f"get_job_status ok job_id={job_id} application_id={job.application_id} state={state}") logger.info(f"get_job_status ok job_id={job.job_id} application_id={job.application_id} state={state}")
return JobStatus(application_id=job.application_id, state=state, raw=raw) return JobStatus(application_id=job.application_id, state=state, raw=raw)
+107 -8
View File
@@ -1,14 +1,17 @@
# coding=utf-8 # coding=utf-8
from datetime import datetime from datetime import datetime
from pathlib import Path
import pytest
from spark_executor.core.job_store import JobStore from spark_executor.core.job_store import JobStore
from spark_executor.models import Job from spark_executor.models import Job
def _job(jid: str) -> Job: def _job(jid: str, app_id: str | None = None) -> Job:
return Job( return Job(
job_id=jid, job_id=jid,
application_id=f"application_{jid}", application_id=app_id or f"application_{jid}",
script_path="/tmp/j.py", script_path="/tmp/j.py",
queue="default", queue="default",
submit_time=datetime(2026, 6, 24), submit_time=datetime(2026, 6, 24),
@@ -16,20 +19,116 @@ def _job(jid: str) -> Job:
) )
def test_put_then_get_roundtrip(): @pytest.fixture
store = JobStore() def store(tmp_path: Path) -> JobStore:
"""Per-test file-backed JobStore rooted in tmp_path. No cross-test leakage."""
return JobStore(data_dir=str(tmp_path))
# --- Basic CRUD (was the entire file before the persistence fix) ---
def test_put_then_get_roundtrip(store):
store.put(_job("a")) store.put(_job("a"))
assert store.get("a") is not None assert store.get("a") is not None
assert store.get("a").application_id == "application_a" assert store.get("a").application_id == "application_a"
def test_get_missing_returns_none(): def test_get_missing_returns_none(store):
store = JobStore()
assert store.get("nope") is None assert store.get("nope") is None
def test_list_returns_all_jobs(): def test_list_returns_all_jobs(store):
store = JobStore()
store.put(_job("a")) store.put(_job("a"))
store.put(_job("b")) store.put(_job("b"))
assert {j.job_id for j in store.list()} == {"a", "b"} assert {j.job_id for j in store.list()} == {"a", "b"}
# --- Persistence: writes go to disk and survive a fresh instance ---
def test_data_persists_across_instances(tmp_path: Path):
"""The whole reason this used to be in-memory: gunicorn workers don't
share memory. A second JobStore pointed at the same data_dir MUST see
the records the first one wrote, otherwise we're back to the
"Unknown job_id" bug from the multi-worker setup."""
writer = JobStore(data_dir=str(tmp_path))
writer.put(_job("a1b2c3d4e5f6", app_id="application_17400000001_0001"))
reader = JobStore(data_dir=str(tmp_path))
assert reader.get("a1b2c3d4e5f6") is not None
assert reader.get("a1b2c3d4e5f6").application_id == "application_17400000001_0001"
def test_data_file_is_human_readable_json(tmp_path: Path):
"""If we're going to disk at all, the file should be inspectable
without the application running — saves an ops engineer a forensics
trip at 2am."""
import json
store = JobStore(data_dir=str(tmp_path))
store.put(_job("abc", app_id="application_1"))
raw = json.loads((tmp_path / "jobs.json").read_text(encoding="utf-8"))
assert "abc" in raw
assert raw["abc"]["application_id"] == "application_1"
assert raw["abc"]["connection"] == "prod"
def test_corrupt_file_does_not_crash(store, tmp_path):
"""A partial write (e.g. killed mid-dump, full disk) shouldn't take
the whole tool surface down — log it and treat as empty."""
(tmp_path / "jobs.json").write_text("{not valid json", encoding="utf-8")
# Should not raise; should return None / empty.
assert store.get("anything") is None
assert store.list() == []
# And we should still be able to write through it.
store.put(_job("x"))
assert store.get("x") is not None
# --- get_either: accept job_id OR application_id ---
def test_get_either_finds_by_job_id(store):
store.put(_job("a1b2c3d4e5f6", app_id="application_1"))
job = store.get_either("a1b2c3d4e5f6")
assert job is not None
assert job.application_id == "application_1"
def test_get_either_finds_by_application_id(store):
"""The fix for 'agent passed the wrong id and got Unknown job_id':
if the caller has a YARN application_id on hand, look it up by that."""
store.put(_job("a1b2c3d4e5f6", app_id="application_17400000001_0001"))
job = store.get_either("application_17400000001_0001")
assert job is not None
assert job.job_id == "a1b2c3d4e5f6"
def test_get_either_prefers_job_id_on_collision(store):
"""If a job_id and an application_id collide (unlikely but possible —
e.g. someone seeds both), the direct job_id lookup wins."""
store.put(_job("collide", app_id="not-collide"))
store.put(_job("other", app_id="collide"))
job = store.get_either("collide")
assert job is not None
assert job.job_id == "collide"
def test_get_either_returns_none_for_unknown(store):
assert store.get_either("nothing-here") is None
def test_get_by_application_id_is_distinct(store):
"""get_by_application_id should NOT match by job_id (it's the explicit
application_id-only lookup). get_either is the forgiving one."""
store.put(_job("a1b2c3d4e5f6", app_id="application_1"))
assert store.get_by_application_id("a1b2c3d4e5f6") is None
assert store.get_by_application_id("application_1") is not None
# --- put is idempotent (re-put same job_id replaces, not duplicates) ---
def test_put_replaces_existing_job(store):
"""confirm_submit_job may retry; the second put of the same job_id
must replace, not append, so list() doesn't grow on every retry."""
store.put(_job("a", app_id="application_1"))
store.put(_job("a", app_id="application_2")) # same job_id, new app_id
assert len(store.list()) == 1
assert store.get("a").application_id == "application_2"
+44
View File
@@ -1,5 +1,6 @@
# coding=utf-8 # coding=utf-8
from datetime import datetime from datetime import datetime
from pathlib import Path
from unittest.mock import patch from unittest.mock import patch
import pytest import pytest
@@ -11,6 +12,17 @@ from spark_executor.tools import kill
from spark_executor.tools import connections from spark_executor.tools import connections
@pytest.fixture
def fresh_stores(tmp_path: Path):
"""Wire up connection + job stores rooted in tmp_path. Per-test isolation
so file-backed JobStore doesn't leak between cases."""
store = connection_store.ConnectionStore(data_dir=str(tmp_path))
connection_store.store = store
connections.store = store
kill.conn_store = store
kill.store = JobStore(data_dir=str(tmp_path))
def _fresh_stores(): def _fresh_stores():
store = connection_store.ConnectionStore() store = connection_store.ConnectionStore()
connection_store.store = store connection_store.store = store
@@ -65,3 +77,35 @@ def test_kill_job_raises_when_connection_missing():
) )
with pytest.raises(KeyError, match="Connection not found"): with pytest.raises(KeyError, match="Connection not found"):
kill.kill_job("abc") kill.kill_job("abc")
# --- application_id accepted (regression: "agent passed wrong id" bug) ---
def test_kill_job_accepts_application_id(fresh_stores):
kill.conn_store.save(Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088"))
kill.store.put(
Job(
job_id="a1b2c3d4e5f6",
application_id="application_17400000001_0001",
script_path="/tmp/j.py",
queue="default",
submit_time=datetime(2026, 6, 24),
connection="prod",
yarn_rm_url="http://rm:8088",
)
)
with patch("spark_executor.tools.kill.kill_application") as m:
result = kill.kill_job("application_17400000001_0001")
# Underlying YARN call uses the application_id, not the local job_id.
assert m.call_args.args[0] == "application_17400000001_0001"
# And the response still surfaces BOTH ids so the agent can confirm.
assert result["job_id"] == "a1b2c3d4e5f6"
assert result["application_id"] == "application_17400000001_0001"
def test_kill_job_unknown_error_message_mentions_both_ids(fresh_stores):
with pytest.raises(KeyError) as ei:
kill.kill_job("totally-fake")
msg = str(ei.value)
assert "job_id" in msg
assert "application_id" in msg
+42 -7
View File
@@ -1,5 +1,6 @@
# coding=utf-8 # coding=utf-8
from datetime import datetime from datetime import datetime
from pathlib import Path
from unittest.mock import patch from unittest.mock import patch
import pytest import pytest
@@ -11,16 +12,23 @@ from spark_executor.tools import logs
from spark_executor.tools import connections from spark_executor.tools import connections
def _fresh_stores(): @pytest.fixture
def fresh_stores(tmp_path: Path):
"""Wire up connection + job stores rooted in tmp_path. Per-test isolation
so file-backed JobStore doesn't leak between cases."""
store = connection_store.ConnectionStore(data_dir=str(tmp_path))
connection_store.store = store
connections.store = store
logs.conn_store = store
logs.store = JobStore(data_dir=str(tmp_path))
def _seed(job_id="abc", app_id="application_1"):
store = connection_store.ConnectionStore() store = connection_store.ConnectionStore()
connection_store.store = store connection_store.store = store
connections.store = store connections.store = store
logs.conn_store = store logs.conn_store = store
logs.store = JobStore() logs.store = JobStore()
def _seed(job_id="abc", app_id="application_1"):
_fresh_stores()
logs.conn_store.save(Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088")) logs.conn_store.save(Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088"))
logs.store.put( logs.store.put(
Job( Job(
@@ -60,8 +68,7 @@ def test_get_job_logs_raises_for_unknown_job():
logs.get_job_logs("missing") logs.get_job_logs("missing")
def test_get_job_logs_raises_when_connection_missing(): def test_get_job_logs_raises_when_connection_missing(fresh_stores):
_fresh_stores()
logs.store.put( logs.store.put(
Job( Job(
job_id="abc", job_id="abc",
@@ -74,3 +81,31 @@ def test_get_job_logs_raises_when_connection_missing():
) )
with pytest.raises(KeyError, match="Connection not found"): with pytest.raises(KeyError, match="Connection not found"):
logs.get_job_logs("abc") logs.get_job_logs("abc")
# --- application_id accepted (regression: "agent passed wrong id" bug) ---
def test_get_job_logs_accepts_application_id():
"""The agent gets both job_id and application_id back from
confirm_submit_job and routinely passes the wrong one. The tool must
accept EITHER and return the same logs."""
_seed(job_id="a1b2c3d4e5f6", app_id="application_17400000001_0001")
with patch(
"spark_executor.tools.logs.get_application_logs",
return_value="logs here",
) as m:
out = logs.get_job_logs("application_17400000001_0001")
assert out == "logs here"
# And the underlying YARN call used the YARN ID, not the local job_id.
assert m.call_args.args[0] == "application_17400000001_0001"
def test_get_job_logs_unknown_error_message_mentions_both_ids():
"""The "neither matched" error should explicitly call out BOTH
accepted id forms so the agent doesn't guess."""
_seed()
with pytest.raises(KeyError) as ei:
logs.get_job_logs("totally-fake")
msg = str(ei.value)
assert "job_id" in msg
assert "application_id" in msg
+52 -1
View File
@@ -1,5 +1,6 @@
# coding=utf-8 # coding=utf-8
from datetime import datetime from datetime import datetime
from pathlib import Path
from unittest.mock import patch from unittest.mock import patch
import pytest import pytest
@@ -11,6 +12,17 @@ from spark_executor.tools import result
from spark_executor.tools import connections from spark_executor.tools import connections
@pytest.fixture
def fresh_stores(tmp_path: Path):
"""Wire up connection + job stores rooted in tmp_path. Per-test isolation
so file-backed JobStore doesn't leak between cases."""
store = connection_store.ConnectionStore(data_dir=str(tmp_path))
connection_store.store = store
connections.store = store
result.conn_store = store
result.store = JobStore(data_dir=str(tmp_path))
def _fresh_stores(): def _fresh_stores():
store = connection_store.ConnectionStore() store = connection_store.ConnectionStore()
connection_store.store = store connection_store.store = store
@@ -114,8 +126,14 @@ def test_result_handles_missing_optional_fields():
def test_result_raises_keyerror_for_unknown_job(): def test_result_raises_keyerror_for_unknown_job():
_fresh_stores() _fresh_stores()
with pytest.raises(KeyError, match="Unknown job_id"): with pytest.raises(KeyError) as ei:
result.get_job_result("missing") result.get_job_result("missing")
# New error message must still flag "Unknown" so callers / agents can
# recognize the failure, AND mention application_id so the agent
# knows the other form is also accepted.
msg = str(ei.value)
assert "job_id" in msg
assert "application_id" in msg
def test_result_raises_when_connection_missing(): def test_result_raises_when_connection_missing():
@@ -132,3 +150,36 @@ def test_result_raises_when_connection_missing():
) )
with pytest.raises(KeyError, match="Connection not found"): with pytest.raises(KeyError, match="Connection not found"):
result.get_job_result("abc") result.get_job_result("abc")
# --- application_id accepted (regression: "agent passed wrong id" bug) ---
def test_get_job_result_accepts_application_id(fresh_stores):
result.conn_store.save(Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088"))
result.store.put(
Job(
job_id="a1b2c3d4e5f6",
application_id="application_17400000001_0001",
script_path="/tmp/j.py",
queue="default",
submit_time=datetime(2026, 6, 24),
connection="prod",
yarn_rm_url="http://rm:8088",
)
)
raw = {
"app": {
"id": "application_17400000001_0001",
"state": "SUCCEEDED",
"finalStatus": "SUCCEEDED",
}
}
import json
with patch(
"spark_executor.tools.result.get_application_status",
return_value=("FINISHED", json.dumps(raw)),
) as m:
out = result.get_job_result("application_17400000001_0001")
assert out.application_id == "application_17400000001_0001"
assert out.state == "FINISHED"
assert m.call_args.args[0] == "application_17400000001_0001"
+44
View File
@@ -1,5 +1,6 @@
# coding=utf-8 # coding=utf-8
from datetime import datetime from datetime import datetime
from pathlib import Path
from unittest.mock import patch from unittest.mock import patch
import pytest import pytest
@@ -11,6 +12,17 @@ from spark_executor.tools import connections, status
from spark_executor.core.yarn_client import YarnClientConfig from spark_executor.core.yarn_client import YarnClientConfig
@pytest.fixture
def fresh_stores(tmp_path: Path):
"""Wire up connection + job stores rooted in tmp_path. Per-test isolation
so file-backed JobStore doesn't leak between cases."""
store = connection_store.ConnectionStore(data_dir=str(tmp_path))
connection_store.store = store
connections.store = store
status.conn_store = store
status.store = JobStore(data_dir=str(tmp_path))
def _fresh_stores(): def _fresh_stores():
"""Reset job store and connection store singletons for a single test.""" """Reset job store and connection store singletons for a single test."""
store = connection_store.ConnectionStore() store = connection_store.ConnectionStore()
@@ -103,3 +115,35 @@ def test_get_job_status_passes_auth_config():
assert isinstance(args[1], YarnClientConfig) assert isinstance(args[1], YarnClientConfig)
assert args[1].auth_type == "basic" assert args[1].auth_type == "basic"
assert args[1].auth_user == "hdfs" assert args[1].auth_user == "hdfs"
# --- application_id accepted (regression: "agent passed wrong id" bug) ---
def test_get_job_status_accepts_application_id(fresh_stores):
status.conn_store.save(Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088"))
status.store.put(
Job(
job_id="a1b2c3d4e5f6",
application_id="application_17400000001_0001",
script_path="/tmp/j.py",
queue="default",
submit_time=datetime(2026, 6, 24),
connection="prod",
yarn_rm_url="http://rm:8088",
)
)
with patch(
"spark_executor.tools.status.get_application_status",
return_value=("RUNNING", "{}"),
):
out = status.get_job_status("application_17400000001_0001")
assert out.state == "RUNNING"
assert out.application_id == "application_17400000001_0001"
def test_get_job_status_unknown_error_message_mentions_both_ids(fresh_stores):
with pytest.raises(KeyError) as ei:
status.get_job_status("totally-fake")
msg = str(ei.value)
assert "job_id" in msg
assert "application_id" in msg