diff --git a/spark_executor/server.py b/spark_executor/server.py index fe18745..578c733 100644 --- a/spark_executor/server.py +++ b/spark_executor/server.py @@ -13,6 +13,15 @@ from spark_executor.tools.connections import ( ) 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, @@ -31,85 +40,72 @@ def health_check(): # 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(connection: str, script_path: str, queue: str = "default", - executor_memory: str = "4G", executor_cores: int = 2, - num_executors: int = 2): - return prepare_submit_job( - connection=connection, - script_path=script_path, - queue=queue, - executor_memory=executor_memory, - executor_cores=executor_cores, - num_executors=num_executors, - ) +def _prepare_submit_job(req: PrepareSubmitJobRequest): + return prepare_submit_job(**req.model_dump()) @app.post("/confirm_submit_job") -def _confirm_submit_job(pending_id: str): - return confirm_submit_job(pending_id=pending_id) +def _confirm_submit_job(req: PendingIdRequest): + return confirm_submit_job(pending_id=req.pending_id) @app.post("/list_pending_jobs") -def _list_pending_jobs(): +def _list_pending_jobs(_req: EmptyRequest = EmptyRequest()): return list_pending_jobs() @app.post("/get_pending_job") -def _get_pending_job(pending_id: str): - return get_pending_job(pending_id) +def _get_pending_job(req: PendingIdRequest): + return get_pending_job(req.pending_id) @app.post("/cancel_pending_job") -def _cancel_pending_job(pending_id: str): - return cancel_pending_job(pending_id) +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(job_id: str): - return get_job_status(job_id) +def _get_job_status(req: JobIdRequest): + return get_job_status(req.job_id) @app.post("/get_job_logs") -def _get_job_logs(job_id: str, tail_chars: int = 5000): - return get_job_logs(job_id, tail_chars=tail_chars) +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(job_id: str): - return kill_job(job_id) +def _kill_job(req: JobIdRequest): + return kill_job(req.job_id) # --- Connection management tools --- @app.post("/save_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): - return save_connection( - name=name, - master=master, - deploy_mode=deploy_mode, - yarn_rm_url=yarn_rm_url, - spark_conf=spark_conf, - ) +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(): +def _list_connections(_req: EmptyRequest = EmptyRequest()): return list_connections() @app.post("/get_connection") -def _get_connection(name: str): - return get_connection(name) +def _get_connection(req: ConnectionNameRequest): + return get_connection(req.name) @app.post("/delete_connection") -def _delete_connection(name: str): - return delete_connection(name) +def _delete_connection(req: ConnectionNameRequest): + return delete_connection(req.name) diff --git a/spark_executor/tools/requests.py b/spark_executor/tools/requests.py new file mode 100644 index 0000000..d7c0d55 --- /dev/null +++ b/spark_executor/tools/requests.py @@ -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 diff --git a/tests/integration/test_mcp_routes.py b/tests/integration/test_mcp_routes.py index dd1ded9..4b85cf8 100644 --- a/tests/integration/test_mcp_routes.py +++ b/tests/integration/test_mcp_routes.py @@ -1,7 +1,28 @@ # 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(): @@ -31,3 +52,61 @@ def test_twelve_tool_routes_registered(): "/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() == []