c9e54a3 引入了 server-minted 校验, 但有两处需要调整:
1. 错误信息用 fmt.Errorf(... %w).Error() 构造, %w 包装语义被丢
掉 (errResult 接收 string, Error() 之后 %w 已经不可达). 改用
fmt.Sprintf 拼接, %q 引用路径, 错误信息保持不变但代码不再误
导.
2. 校验放在 scriptPath 解析之后、queue 解析之前. 错误信息会按字
段出现顺序报 (cluster_id, master, deploy_mode, script_path,
queue, ...), 但当前顺序是 script_path 校验先报, 然后才报 queue
缺失. 把校验挪到所有 RequireString 之后、parseStringMap 之前,
LLM 看错误时字段顺序跟 schema 顺序一致.
nil-safe 校验逻辑不变 (d.UploadStore == nil 时跳过). 现有测试
全部通过:
- TestSparkSubmit_RejectsNonMintedPath 仍通过 (校验位置不影响
行为)
- TestSparkSubmit_StructuredCommand 仍通过 (走 mint 路径)
Co-Authored-By: Claude <noreply@anthropic.com>
264 lines
8.6 KiB
Go
264 lines
8.6 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
|
|
}
|
|
|
|
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
|
|
}
|