# 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.models import Connection 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_sixteen_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 (3) — Stage 2 "/write_job_file", "/read_job_file", "/update_job_file", ): assert path in paths, f"missing MCP tool route: {path}" def test_seventeen_tool_routes_registered(): paths = {r.path for r in app.routes} assert "/update_pending_job" in paths def test_twenty_three_tool_routes_registered(): paths = {r.path for r in app.routes} for path in ( "/get_external_job_logs", "/get_external_job_status", "/get_external_job_result", "/list_applications", "/fetch_url", "/update_connection", ): assert path in paths, f"missing MCP tool route: {path}" # --- operation_id: pin clean MCP tool names (no auto-generated suffixes) --- # # fastapi-mcp uses each route's OpenAPI `operationId` as the MCP tool name # exposed to LLM clients (fastapi_mcp/server.py:619-620). If we omit # `operation_id=`, FastAPI auto-generates a name like # `prepare_submit_job_prepare_submit_job_post` (function_name + path + # method), which is ugly for the agent. This test asserts every MCP tool # route has an explicit, clean operation_id matching its handler. # Set of routes that are NOT MCP tools (no operation_id expected). _NON_MCP_PATHS = {"/health", "/openapi.json", "/docs", "/docs/oauth2-redirect", "/redoc"} def test_every_mcp_route_has_explicit_clean_operation_id(): c = TestClient(app) openapi = c.get("/openapi.json").json() paths = openapi["paths"] seen: list[tuple[str, str, str]] = [] # (path, method, operation_id) for path, methods in paths.items(): if path in _NON_MCP_PATHS: continue for method, op in methods.items(): if method.upper() not in {"GET", "POST", "PUT", "DELETE", "PATCH"}: continue op_id = op.get("operationId") assert op_id, ( f"route {method.upper()} {path} has no operationId in OpenAPI — " f"fastapi-mcp will fall back to an auto-generated name like " f"`{path.strip('/').replace('/', '_')}_{method}_...` and expose " f"that to LLM clients. Add an explicit `operation_id=` to the decorator." ) # Reject FastAPI's auto-generated format: it always embeds the # HTTP method as a `_post`/`_get` suffix, and repeats the path. assert not op_id.endswith(f"_{method}"), ( f"route {method.upper()} {path} has auto-generated operationId " f"{op_id!r} (ends with `_{method}`). Add an explicit " f"`operation_id=` to the decorator." ) seen.append((path, method.upper(), op_id)) # Sanity: we expect at least the 16 MCP tool routes from the project. assert len(seen) >= 16, f"only found {len(seen)} MCP tool routes, expected >=16: {seen}" # Every operation_id must be unique — fastapi-mcp's tool registry keys # by name, so duplicates would silently shadow one another. ids = [op_id for _, _, op_id in seen] dupes = [x for x in set(ids) if ids.count(x) > 1] assert not dupes, f"duplicate operation_ids: {dupes}" # --- 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(tmp_path): c = TestClient(app) c.post("/save_connection", json={"name": "prod", "master": "yarn"}) # The script must exist (Stage 2 cleanup) — write a real file first. script = tmp_path / "demo.py" script.write_text("print('hi')\n") r = c.post( "/prepare_submit_job", json={ "connection": "prod", "script_path": str(script), "queue": "research", "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", }, ) 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_prepare_rejects_unconfirmed_defaults(tmp_path): c = TestClient(app) c.post("/save_connection", json={"name": "prod", "master": "yarn"}) script = tmp_path / "demo.py" script.write_text("print('hi')\n") r = c.post( "/prepare_submit_job", json={ "connection": "prod", "script_path": str(script), "app_name": "test-app", }, ) assert r.status_code == 400 detail = r.json()["detail"] assert "Please confirm default values" in detail assert "queue='default'" in detail assert "executor_memory='2G'" in detail assert "executor_cores=2" in detail assert "num_executors=2" in detail def test_prepare_accepts_confirmed_defaults(tmp_path): c = TestClient(app) c.post("/save_connection", json={"name": "prod", "master": "yarn"}) script = tmp_path / "demo.py" script.write_text("print('hi')\n") r = c.post( "/prepare_submit_job", json={ "connection": "prod", "script_path": str(script), "queue": "default", "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", }, ) assert r.status_code == 200 body = r.json() assert body["status"] == "PENDING" assert body["parameters"]["queue"] == "default" assert body["parameters"]["executor_memory"] == "2G" assert body["parameters"]["executor_cores"] == 2 assert body["parameters"]["num_executors"] == 2 def test_prepare_rejects_nonexistent_script_path_with_400(): """MCP clients should see a clean 400 with a remediation hint, not a 500.""" c = TestClient(app) c.post("/save_connection", json={"name": "prod", "master": "yarn"}) r = c.post( "/prepare_submit_job", json={ "connection": "prod", "script_path": "/nope/does_not_exist.py", "queue": "default", "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", }, ) assert r.status_code == 400 detail = r.json()["detail"] # The message must guide the agent to the right next step assert "does not exist" in detail assert "write_job_file" in detail def test_prepare_rejects_sql_policy_violation_with_400(tmp_path): """SQL guard: even if the file exists, prepare_submit_job must reject code containing DROP/DELETE/etc. (defense in depth — write_job_file also enforces this, but a host-mounted file might bypass that).""" c = TestClient(app) c.post("/save_connection", json={"name": "prod", "master": "yarn"}) bad_script = tmp_path / "evil.py" bad_script.write_text('spark.sql("DROP TABLE users")\n') r = c.post( "/prepare_submit_job", json={ "connection": "prod", "script_path": str(bad_script), "queue": "default", "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", }, ) assert r.status_code == 400 detail = r.json()["detail"] assert "DROP" in detail assert "SELECT" in detail and "INSERT" in detail # the policy is explained def test_write_rejects_sql_policy_violation_with_400(tmp_path): """Same policy at write_job_file — agent gets immediate feedback before the file is even written to disk.""" from common import config config.settings.jobs_dir = str(tmp_path / "jobs") c = TestClient(app) r = c.post( "/write_job_file", json={"code": 'spark.sql("DELETE FROM events")\n'}, ) assert r.status_code == 400 detail = r.json()["detail"] assert "DELETE" in detail # And the file should NOT have been written assert list(tmp_path.glob("*.py")) == [] # --- Stage 2: read_job_file / update_job_file --- def test_read_job_file_returns_content_via_mcp(tmp_path): c = TestClient(app) script = tmp_path / "demo.py" script.write_text("print('hello from read_job_file')\n") r = c.post("/read_job_file", json={"script_path": str(script)}) assert r.status_code == 200 body = r.json() assert body["content"] == "print('hello from read_job_file')\n" assert body["path"] == str(script) def test_read_job_file_404_for_missing_file(tmp_path): c = TestClient(app) r = c.post("/read_job_file", json={"script_path": str(tmp_path / "nope.py")}) assert r.status_code == 400 assert "does not exist" in r.json()["detail"] def test_update_job_file_writes_and_reads_back(tmp_path): from common import config config.settings.jobs_dir = str(tmp_path) c = TestClient(app) script = tmp_path / "edit.py" script.write_text("v1\n") # Update r1 = c.post("/update_job_file", json={"script_path": str(script), "content": "v2\n"}) assert r1.status_code == 200 assert r1.json()["bytes_written"] == 3 # Read back r2 = c.post("/read_job_file", json={"script_path": str(script)}) assert r2.json()["content"] == "v2\n" def test_update_job_file_rejects_paths_outside_jobs_dir(tmp_path): from common import config config.settings.jobs_dir = str(tmp_path / "jobs") c = TestClient(app) other = tmp_path / "elsewhere.py" other.write_text("x\n") r = c.post("/update_job_file", json={"script_path": str(other), "content": "y\n"}) assert r.status_code == 400 assert "must be under" in r.json()["detail"] def test_list_and_get_pending_job_roundtrip_via_body(tmp_path): c = TestClient(app) c.post("/save_connection", json={"name": "prod", "master": "yarn"}) script = tmp_path / "a.py" script.write_text("print('hi')\n") prep = c.post( "/prepare_submit_job", json={ "connection": "prod", "script_path": str(script), "queue": "default", "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", }, ).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"] == str(script) 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: write_job_file --- def test_write_job_file_accepts_code_string_in_body(tmp_path, monkeypatch): """End-to-end: a Pydantic body model lets tools/call pass a code string that FastAPI would reject if it were a query parameter (length limits).""" monkeypatch.chdir(tmp_path) c = TestClient(app) code = "print('from MCP integration test')\n" * 100 # > FastAPI query limit r = c.post("/write_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", "queue": "default", "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", }, ) assert r.status_code == 404 assert "nope" in r.json()["detail"] def test_confirm_non_pending_returns_400(tmp_path): """A CANCELLED pending should refuse confirm; FastAPI should surface ValueError as 400.""" c = TestClient(app) c.post("/save_connection", json={"name": "prod", "master": "yarn"}) script = tmp_path / "x.py" script.write_text("print('hi')\n") prep = c.post( "/prepare_submit_job", json={ "connection": "prod", "script_path": str(script), "queue": "default", "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", }, ).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(tmp_path): """Cancelling a SUBMITTED pending should refuse with ValueError -> 400.""" c = TestClient(app) c.post("/save_connection", json={"name": "prod", "master": "yarn"}) script = tmp_path / "x.py" script.write_text("print('hi')\n") prep = c.post( "/prepare_submit_job", json={ "connection": "prod", "script_path": str(script), "queue": "default", "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", }, ).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"] # --- update_pending_job route --- def test_update_pending_job_route(tmp_path): c = TestClient(app) c.post("/save_connection", json={"name": "prod", "master": "yarn"}) script = tmp_path / "x.py" script.write_text("print('hi')\n") prep = c.post( "/prepare_submit_job", json={ "connection": "prod", "script_path": str(script), "queue": "default", "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", }, ).json() pid = prep["pending_id"] r = c.post( "/update_pending_job", json={"pending_id": pid, "queue": "research"}, ) assert r.status_code == 200, r.text assert r.json()["status"] == "PENDING" assert r.json()["parameters"]["queue"] == "research" got = c.post("/get_pending_job", json={"pending_id": pid}).json() assert got["queue"] == "research" # untouched fields are preserved assert got["executor_memory"] == "2G" assert got["app_name"] == "test-app" def test_update_pending_job_route_rejects_submitted(tmp_path, monkeypatch): c = TestClient(app) c.post("/save_connection", json={"name": "prod", "master": "yarn"}) script = tmp_path / "x.py" script.write_text("print('hi')\n") prep = c.post( "/prepare_submit_job", json={ "connection": "prod", "script_path": str(script), "queue": "default", "executor_memory": "2G", "executor_cores": 2, "num_executors": 2, "app_name": "test-app", }, ).json() pid = prep["pending_id"] fake_proc = type("P", (), { "returncode": 0, "stderr": "tracking URL: http://rm:8088/proxy/application_17400000001/\n", })() monkeypatch.setattr(submit, "run_spark_submit", lambda _cmd: fake_proc) c.post("/confirm_submit_job", json={"pending_id": pid}) r = c.post( "/update_pending_job", json={"pending_id": pid, "queue": "research"}, ) assert r.status_code == 400 assert "only PENDING submissions can be updated" in r.json()["detail"] # --- YarnError -> 502 handler --- def test_yarn_error_returns_502_with_detail(tmp_path, monkeypatch): """When a YARN REST call raises YarnError (host unreachable, YARN returned 4xx/5xx, parse failure), the route must surface HTTP 502 with the YarnError message in the detail — not 500 'Internal Server Error' with no info. Mirrors the same fix for fetch_url. """ from unittest.mock import patch from spark_executor.core.yarn_client import YarnError connection_store.store.save( Connection( name="prod", master="yarn", yarn_rm_url="http://ccam1:8088", ) ) c = TestClient(app) with patch( "spark_executor.tools.status.get_application_status", side_effect=YarnError("YARN connection failed: ConnectError: Connection refused"), ): r = c.post( "/get_job_status", json={"job_id": "a1b2c3d4e5f6"}, ) assert r.status_code == 502 detail = r.json()["detail"] assert "YARN connection failed" in detail assert "Connection refused" in detail def test_external_job_status_returns_502_on_yarn_error(tmp_path): """get_external_job_status uses get_application_status under the hood — same handler, same 502, same detail.""" from unittest.mock import patch from spark_executor.core.yarn_client import YarnError connection_store.store.save( Connection( name="prod", master="yarn", yarn_rm_url="http://ccam1:8088", ) ) c = TestClient(app) with patch( "spark_executor.tools.external_jobs.get_application_status", side_effect=YarnError("YARN application 'application_xxx' not found"), ): r = c.post( "/get_external_job_status", json={"application_id": "application_1740000000001_0001", "connection_name": "prod"}, ) assert r.status_code == 502 assert "not found" in r.json()["detail"] def test_external_job_result_returns_502_on_yarn_error(tmp_path): """get_external_job_result also uses get_application_status (same underlying YARN endpoint), and exercises the same YarnError path.""" from unittest.mock import patch from spark_executor.core.yarn_client import YarnError connection_store.store.save( Connection( name="prod", master="yarn", yarn_rm_url="http://ccam1:8088", ) ) c = TestClient(app) with patch( "spark_executor.tools.external_jobs.get_application_status", side_effect=YarnError("YARN connection failed: ConnectError: Connection refused"), ): r = c.post( "/get_external_job_result", json={"application_id": "application_1740000000001_0001", "connection_name": "prod"}, ) assert r.status_code == 502 assert "YARN connection failed" in r.json()["detail"] def test_external_job_logs_returns_502_on_yarn_error(tmp_path): """get_external_job_logs uses get_application_logs — different function, but raises YarnError the same way. Handler must catch it the same way.""" from unittest.mock import patch from spark_executor.core.yarn_client import YarnError connection_store.store.save( Connection( name="prod", master="yarn", yarn_rm_url="http://ccam1:8088", ) ) c = TestClient(app) with patch( "spark_executor.tools.external_jobs.get_application_logs", side_effect=YarnError("YARN GET logs returned HTTP 500: cluster overloaded"), ): r = c.post( "/get_external_job_logs", json={ "application_id": "application_1740000000001_0001", "connection_name": "prod", "tail_chars": 5000, }, ) assert r.status_code == 502 assert "YARN GET logs returned HTTP 500" in r.json()["detail"] assert "cluster overloaded" in r.json()["detail"] def test_list_applications_returns_502_on_yarn_error(tmp_path): """list_applications uses list_applications_yarn — yet another YARN-touching tool. Same YarnError -> 502 contract.""" from unittest.mock import patch from spark_executor.core.yarn_client import YarnError connection_store.store.save( Connection( name="prod", master="yarn", yarn_rm_url="http://ccam1:8088", ) ) c = TestClient(app) with patch( "spark_executor.tools.external_jobs.list_applications_yarn", side_effect=YarnError("YARN list applications failed: 503 Service Unavailable"), ): r = c.post( "/list_applications", json={"connection_name": "prod"}, ) assert r.status_code == 502 assert "YARN list applications failed" in r.json()["detail"]