Files
spark-mcp/internal/mcp/tools/spark_submit.go
T
tao.chenandClaude c9e54a3f27 tools: 强制 spark_submit 的 script_path 必须来自 upload_file
在 spark_submit handler 中增加 server-minted 校验:script_path 必须通过
本服务器 upload_file 的 UploadStore.Validate,否则立即返回清晰错误,避免
远程 agent 把自身本地路径转发给服务器侧 spark-submit 时出现难以定位的
FileNotFoundException。

校验对 nil UploadStore 是安全的:当 d.UploadStore 为 nil 时直接跳过,保留
不强制依赖 store 的测试或降级场景。当前 testDepsWithDataDir 始终会注入
UploadStore,所以测试走的是真实校验路径。

同步更新了 script_path 的工具描述,明确声明非 upload_file 铸造的路径会被
拒绝。这是对 7920dc9 中“path 永远在最后”规则的收紧——现在不仅位置固定,
而且必须是本机 uploads 目录下的有效上传文件。

测试调整:
- 新增 TestSparkSubmit_RejectsNonMintedPath:使用 t.TempDir() 下未通过
  upload_file 写入的文件作为 script_path,断言返回错误并包含
  "not from upload_file"。
- TestSparkSubmit_MissingRequiredField / TestSparkSubmit_BadSparkConfValue
  原来使用 "/tmp/script.py" 作为占位路径,现在会先通过 UploadStore.Save
  生成一个铸造路径,确保它们继续分别验证 master 缺失和 spark_conf 值类型
  错误,而不是被新的路径守卫拦截。

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-13 19:49:26 +08:00

261 lines
8.5 KiB
Go

package tools
import (
"context"
"fmt"
"sort"
"strconv"
"github.com/mark3labs/mcp-go/mcp"
"spark-mcp-go/internal/executor"
)
const SparkSubmitName = "spark_submit"
// NewSparkSubmitTool returns the schema for the spark_submit MCP Tool.
func NewSparkSubmitTool() mcp.Tool {
return mcp.NewTool(SparkSubmitName,
mcp.WithDescription("Submit a Spark application to a configured cluster. The command is built in this order: --master, --queue, --executor-memory, --executor-cores, --num-executors, then --conf per spark_conf entry, then --flag value per extra_args entry, then script_path last. For script_path, pass the `path` field returned by upload_file (do not construct your own path). Returns {app_id, exit_code, stdout_tail, stderr_tail, duration_ms}. Binary is invoked as a child process — no shell, no command injection."),
mcp.WithString("cluster_id",
mcp.Required(),
mcp.Description("ID of the configured cluster (from list_clusters)"),
),
mcp.WithString("master",
mcp.Required(),
mcp.Description("Spark master URL, e.g. yarn, k8s://https://..."),
),
mcp.WithString("deploy_mode",
mcp.Required(),
mcp.Description("Required for forward compatibility; not currently emitted in the command by the builder (matches Python reference)."),
mcp.Enum("client", "cluster"),
),
mcp.WithString("script_path",
mcp.Required(),
mcp.Description("Absolute path to the script. Pass the `path` field returned by upload_file — do not construct your own path. The script is always the last argv element. Paths not minted by upload_file on this server are rejected."),
),
mcp.WithString("queue",
mcp.Required(),
mcp.Description("YARN queue name"),
),
mcp.WithString("executor_memory",
mcp.Required(),
mcp.Description("e.g. 4G"),
),
mcp.WithNumber("executor_cores",
mcp.Required(),
mcp.Description("Cores per executor (integer)"),
),
mcp.WithNumber("num_executors",
mcp.Required(),
mcp.Description("Static executor count (integer)"),
),
mcp.WithObject("spark_conf",
mcp.Description("Map of key=value pairs, each emitted as --conf key=value"),
mcp.AdditionalProperties(map[string]any{"type": "string"}),
),
mcp.WithObject("extra_args",
mcp.Description("Map of flag -> value, each emitted as --flag value"),
mcp.AdditionalProperties(map[string]any{"type": "string"}),
),
)
}
// SparkSubmitHandler runs spark-submit against the requested cluster.
func (d *Deps) SparkSubmitHandler(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
clusterID, err := req.RequireString("cluster_id")
if err != nil {
return errResult("spark_submit: " + err.Error()), nil
}
master, err := req.RequireString("master")
if err != nil {
return errResult("spark_submit: " + err.Error()), nil
}
deployMode, err := req.RequireString("deploy_mode")
if err != nil {
return errResult("spark_submit: " + err.Error()), nil
}
scriptPath, err := req.RequireString("script_path")
if err != nil {
return errResult("spark_submit: " + err.Error()), nil
}
// Reject paths that didn't come through upload_file on this server.
// Cluster-local paths and arbitrary host paths are unreachable from the MCP
// server's spark-submit, so we fail fast with a clear message rather than
// letting spark-submit produce a confusing FileNotFoundException later.
if d.UploadStore != nil {
if _, err := d.UploadStore.Validate(scriptPath); err != nil {
return errResult(fmt.Errorf("spark_submit: script_path %q is not from upload_file on this server; upload the file first via upload_file and pass back the path field: %w", scriptPath, err).Error()), nil
}
}
queue, err := req.RequireString("queue")
if err != nil {
return errResult("spark_submit: " + err.Error()), nil
}
executorMemory, err := req.RequireString("executor_memory")
if err != nil {
return errResult("spark_submit: " + err.Error()), nil
}
args := req.GetArguments()
// deploy_mode is required for forward compatibility but the builder does not
// emit --deploy-mode, matching the Python reference's actual behavior. If the
// Go side needs --deploy-mode later, add it here and document the divergence.
if deployMode != "client" && deployMode != "cluster" {
return errResult("spark_submit: deploy_mode must be client or cluster"), nil
}
sparkConf, err := parseStringMap(args, "spark_conf")
if err != nil {
return errResult("spark_submit: " + err.Error()), nil
}
extraArgs, err := parseStringMap(args, "extra_args")
if err != nil {
return errResult("spark_submit: " + err.Error()), nil
}
executorCores, err := requireNonNegativeInt(args, "executor_cores")
if err != nil {
return errResult("spark_submit: " + err.Error()), nil
}
numExecutors, err := requireNonNegativeInt(args, "num_executors")
if err != nil {
return errResult("spark_submit: " + err.Error()), nil
}
callLog := startToolCall(ctx, d.Logger, SparkSubmitName, map[string]any{
"cluster_id": clusterID,
"master": master,
"deploy_mode": deployMode,
"script_path": scriptPath,
"queue": queue,
"executor_memory": executorMemory,
"executor_cores": executorCores,
"num_executors": numExecutors,
"spark_conf": sparkConf,
"extra_args": extraArgs,
})
defer callLog.End()
cl, err := d.ClusterRepo.Get(ctx, clusterID)
if err != nil {
callLog.WithError(err)
return errResult("spark_submit: cluster " + clusterID + ": " + err.Error()), nil
}
binary := cl.SparkSubmitExecuteBin
// Build argv inline, mirroring the Python reference's output order exactly.
// DefaultSubmitArgs is no longer prepended; the field is currently unused at
// runtime but kept for backward compatibility with the admin API. Removal is
// a separate follow-up.
cmd := []string{binary}
cmd = append(cmd, "--master", master)
cmd = append(cmd, "--queue", queue)
cmd = append(cmd, "--executor-memory", executorMemory)
cmd = append(cmd, "--executor-cores", strconv.Itoa(executorCores))
cmd = append(cmd, "--num-executors", strconv.Itoa(numExecutors))
for _, k := range sortedStringKeys(sparkConf) {
cmd = append(cmd, "--conf", k+"="+sparkConf[k])
}
for _, k := range sortedStringKeys(extraArgs) {
cmd = append(cmd, "--"+k, extraArgs[k])
}
cmd = append(cmd, scriptPath)
if d.Logger != nil {
d.Logger.Debug("spark_submit.built", "cluster_id", clusterID, "argv", cmd)
}
result, err := executor.Run(ctx, executor.SparkSubmitOpts{
Binary: binary,
Args: cmd,
Timeout: d.SparkSubmitTimeout,
})
if err != nil {
callLog.WithError(err)
// Even on error, result carries ExitCode/Stderr for the LLM.
if result != nil {
callLog.WithResult(map[string]any{
"app_id": result.AppID,
"exit_code": result.ExitCode,
"duration_ms": result.DurationMS,
})
return textResult(encodeJSON(map[string]any{
"app_id": result.AppID,
"exit_code": result.ExitCode,
"stdout_tail": result.StdoutTail,
"stderr_tail": result.StderrTail,
"duration_ms": result.DurationMS,
"error": err.Error(),
})), nil
}
return errResult("spark_submit: " + err.Error()), nil
}
callLog.WithResult(map[string]any{
"app_id": result.AppID,
"exit_code": result.ExitCode,
"duration_ms": result.DurationMS,
})
return textResult(encodeJSON(result)), nil
}
// requireNonNegativeInt extracts an integer parameter from the request and
// rejects negative or non-integer values.
func requireNonNegativeInt(args map[string]any, key string) (int, error) {
v, ok := args[key]
if !ok || v == nil {
return 0, fmt.Errorf("%s is required", key)
}
switch n := v.(type) {
case int:
if n < 0 {
return 0, fmt.Errorf("%s must be a non-negative integer", key)
}
return n, nil
case float64:
ni := int(n)
if n < 0 || float64(ni) != n {
return 0, fmt.Errorf("%s must be a non-negative integer", key)
}
return ni, nil
default:
return 0, fmt.Errorf("%s must be a number", key)
}
}
// parseStringMap extracts an optional object whose values are strings.
func parseStringMap(args map[string]any, key string) (map[string]string, error) {
out := make(map[string]string)
v, ok := args[key]
if !ok || v == nil {
return out, nil
}
m, ok := v.(map[string]any)
if !ok {
return nil, fmt.Errorf("%s must be an object", key)
}
for k, val := range m {
s, ok := val.(string)
if !ok {
return nil, fmt.Errorf("%s[%q] must be a string", key, k)
}
out[k] = s
}
return out, nil
}
// sortedStringKeys returns the keys of m in deterministic order.
func sortedStringKeys(m map[string]string) []string {
keys := make([]string, 0, len(m))
for k := range m {
keys = append(keys, k)
}
sort.Strings(keys)
return keys
}