# coding=utf-8 import os 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_thirteen_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", # LLM-driven PySpark generation (1) — Stage 2 "/generate_job_file", ): 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() == [] # --- Stage 2: generate_job_file --- def test_generate_job_file_accepts_code_string_in_body(tmp_path, monkeypatch): """End-to-end: a Pydantic body model lets tools/call pass a code string that FastAPI would reject if it were a query parameter (length limits).""" monkeypatch.chdir(tmp_path) c = TestClient(app) code = "print('from MCP integration test')\n" * 100 # > FastAPI query limit r = c.post("/generate_job_file", json={"code": code}) assert r.status_code == 200, r.text p = r.json()["script_path"] assert os.path.isfile(p) with open(p) as f: assert f.read() == code # --- 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"]