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 } executorMemory, err := req.RequireString("executor_memory") if err != nil { return errResult("spark_submit: " + err.Error()), nil } queue, err := req.RequireString("queue") if err != nil { return errResult("spark_submit: " + err.Error()), nil } args := req.GetArguments() // 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. // Placed after all required-string parses so error messages surface in // field order (cluster_id, master, ..., script_path, queue, ...). if d.UploadStore != nil { if _, err := d.UploadStore.Validate(scriptPath); err != nil { return errResult(fmt.Sprintf("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", scriptPath)), nil } } // 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 }