package tools import ( "context" "os" "path/filepath" "strings" "testing" "time" "github.com/mark3labs/mcp-go/mcp" "spark-mcp-go/internal/cluster" ) func TestSparkSubmit_StructuredCommand(t *testing.T) { deps, repo := testDepsWithDataDir(t) deps.SparkSubmitTimeout = 5 * time.Second store := deps.UploadStore fileID, _, _, _, absPath, err := store.Save([]byte("print('hi')\n"), "hello.py") if err != nil { t.Fatalf("save upload: %v", err) } echoed := filepath.Join(t.TempDir(), "echoed") echoScript := filepath.Join(t.TempDir(), "echo-args.sh") if err := os.WriteFile(echoScript, []byte("#!/bin/sh\nfor arg do\n echo \"$arg\" >> \"$ECHO_FILE\"\ndone\n"), 0o755); err != nil { t.Fatalf("write echo script: %v", err) } t.Setenv("ECHO_FILE", echoed) createCluster(t, repo, &cluster.Cluster{ ID: "cluster-echo", Name: "Echo", IsActive: true, AuthType: cluster.AuthNone, SparkSubmitExecuteBin: echoScript, }) req := newToolRequest(SparkSubmitName, map[string]any{ "cluster_id": "cluster-echo", "master": "yarn", "deploy_mode": "cluster", "script_path": absPath, "queue": "default", "executor_memory": "2G", "executor_cores": 1, "num_executors": 4, "spark_conf": map[string]any{"a": "1", "b": "2"}, "extra_args": map[string]any{"name": "wordcount-job"}, }) res, err := deps.SparkSubmitHandler(context.Background(), req) if err != nil { t.Fatalf("handler error: %v", err) } if res.IsError { t.Fatalf("unexpected error result: %v", res.Content) } got, err := os.ReadFile(echoed) if err != nil { t.Fatalf("read echoed args: %v", err) } lines := strings.Split(strings.TrimSpace(string(got)), "\n") want := []string{ echoScript, "--master", "yarn", "--queue", "default", "--executor-memory", "2G", "--executor-cores", "1", "--num-executors", "4", "--conf", "a=1", "--conf", "b=2", "--name", "wordcount-job", absPath, } if len(lines) != len(want) { t.Fatalf("lines=%v\nwant=%v", lines, want) } for i, l := range lines { if l != want[i] { t.Errorf("line[%d]=%q, want %q", i, l, want[i]) } } if lines[len(lines)-1] != absPath { t.Errorf("last line=%q, want script_path %q", lines[len(lines)-1], absPath) } if !strings.HasSuffix(absPath, "/"+fileID) { t.Errorf("absPath=%q does not end with fileID %q", absPath, fileID) } } func TestSparkSubmit_MissingRequiredField(t *testing.T) { deps, _ := testDepsWithDataDir(t) store := deps.UploadStore _, _, _, _, mintedPath, err := store.Save([]byte("# dummy\n"), "dummy.py") if err != nil { t.Fatalf("save upload: %v", err) } req := newToolRequest(SparkSubmitName, map[string]any{ "cluster_id": "cluster-echo", "deploy_mode": "cluster", "script_path": mintedPath, "queue": "default", "executor_memory": "2G", "executor_cores": 1, "num_executors": 4, }) res, err := deps.SparkSubmitHandler(context.Background(), req) if err != nil { t.Fatalf("handler error: %v", err) } if !res.IsError { t.Fatalf("expected error result, got: %v", res.Content) } text, ok := mcp.AsTextContent(res.Content[0]) if !ok { t.Fatalf("content is not text: %T", res.Content[0]) } if !strings.Contains(text.Text, "master") { t.Errorf("error text=%q, want mention of master", text.Text) } } func TestSparkSubmit_BadSparkConfValue(t *testing.T) { deps, _ := testDepsWithDataDir(t) store := deps.UploadStore _, _, _, _, mintedPath, err := store.Save([]byte("# dummy\n"), "dummy.py") if err != nil { t.Fatalf("save upload: %v", err) } req := newToolRequest(SparkSubmitName, map[string]any{ "cluster_id": "cluster-echo", "master": "yarn", "deploy_mode": "cluster", "script_path": mintedPath, "queue": "default", "executor_memory": "2G", "executor_cores": 1, "num_executors": 4, "spark_conf": map[string]any{"a": 1}, }) res, err := deps.SparkSubmitHandler(context.Background(), req) if err != nil { t.Fatalf("handler error: %v", err) } if !res.IsError { t.Fatalf("expected error result, got: %v", res.Content) } text, ok := mcp.AsTextContent(res.Content[0]) if !ok { t.Fatalf("content is not text: %T", res.Content[0]) } if !strings.Contains(text.Text, "spark_conf") { t.Errorf("error text=%q, want mention of spark_conf", text.Text) } } func TestSparkSubmit_RejectsNonMintedPath(t *testing.T) { deps, _ := testDepsWithDataDir(t) unminted := filepath.Join(t.TempDir(), "unminted.py") if err := os.WriteFile(unminted, []byte("print('not from upload_file')\n"), 0o644); err != nil { t.Fatalf("write unminted file: %v", err) } req := newToolRequest(SparkSubmitName, map[string]any{ "cluster_id": "cluster-echo", "master": "yarn", "deploy_mode": "cluster", "script_path": unminted, "queue": "default", "executor_memory": "2G", "executor_cores": 1, "num_executors": 4, }) res, err := deps.SparkSubmitHandler(context.Background(), req) if err != nil { t.Fatalf("handler error: %v", err) } if !res.IsError { t.Fatalf("expected error result, got: %v", res.Content) } text, ok := mcp.AsTextContent(res.Content[0]) if !ok { t.Fatalf("content is not text: %T", res.Content[0]) } if !strings.Contains(text.Text, "not from upload_file") { t.Errorf("error text=%q, want mention of not from upload_file", text.Text) } }