fix: switch FastAPI routes to Pydantic body models for MCP tools/call
The plan's route signatures used query parameters, which worked for direct HTTP callers and unit tests, but fastapi-mcp's HTTP transport passes tools/call arguments as a JSON body. dict-typed parameters like spark_conf arrived as a string and the route returned 422. Refactor each route to take a single Pydantic body model (saved in spark_executor/tools/requests.py). Underlying tool functions unchanged. Integration tests in tests/integration/test_mcp_routes.py now exercise the full body-based roundtrip (save_connection with spark_conf, prepare → list → get, list_connections with empty body).
This commit is contained in:
@@ -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() == []
|
||||
|
||||
Reference in New Issue
Block a user