package tools import ( "bytes" "context" "encoding/json" "log/slog" "os" "path/filepath" "strings" "testing" "time" "spark-mcp-go/internal/cluster" ) func TestSparkSubmit_PathPassthrough(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", "args": []any{absPath, "--class", "Main"}, }) 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{absPath, "--class", "Main"} if len(lines) != len(want) { t.Fatalf("echoed lines=%v, want %v", lines, want) } for i, l := range lines { if l != want[i] { t.Errorf("line[%d]=%q, want %q", i, l, want[i]) } } // Sanity: the canonical path still contains the minted file ID. if !strings.HasSuffix(absPath, "/"+fileID) { t.Errorf("absPath=%q does not end with fileID %q", absPath, fileID) } } func TestTranslateArgs_Logging(t *testing.T) { deps, _ := testDepsWithDataDir(t) t.Run("minted_path_no_warn", func(t *testing.T) { store := deps.UploadStore fileID, _, _, _, absPath, err := store.Save([]byte("print('hi')\n"), "hello.py") if err != nil { t.Fatalf("save upload: %v", err) } logger, buf := capturingLogger(t) out := translateArgs([]string{absPath, "--class", "Main"}, deps.UploadStore, logger) if out[0] != absPath { t.Errorf("out[0]=%q, want %q", out[0], absPath) } for _, rec := range parseLogRecords(t, buf) { if rec["msg"] == "spark_submit.unminted_path" { t.Errorf("unexpected warn for minted path: %v", rec) } } if !strings.HasSuffix(absPath, "/"+fileID) { t.Errorf("absPath=%q does not end with fileID %q", absPath, fileID) } }) t.Run("cluster_local_path_warns", func(t *testing.T) { logger, buf := capturingLogger(t) out := translateArgs([]string{"/opt/spark/examples/pi.py"}, deps.UploadStore, logger) if len(out) != 1 || out[0] != "/opt/spark/examples/pi.py" { t.Errorf("out=%v, want [/opt/spark/examples/pi.py]", out) } var warns int for _, rec := range parseLogRecords(t, buf) { if rec["msg"] == "spark_submit.unminted_path" { warns++ if rec["path"] != "/opt/spark/examples/pi.py" { t.Errorf("warn path=%v, want /opt/spark/examples/pi.py", rec["path"]) } } } if warns != 1 { t.Errorf("warns=%d, want 1; log=%s", warns, buf.String()) } }) t.Run("non_path_args_never_warn", func(t *testing.T) { logger, buf := capturingLogger(t) out := translateArgs([]string{"--class", "Main", "--master", "yarn"}, deps.UploadStore, logger) want := []string{"--class", "Main", "--master", "yarn"} if len(out) != len(want) { t.Fatalf("out=%v, want %v", out, want) } for i := range out { if out[i] != want[i] { t.Errorf("out[%d]=%q, want %q", i, out[i], want[i]) } } for _, rec := range parseLogRecords(t, buf) { if rec["msg"] == "spark_submit.unminted_path" { t.Errorf("unexpected warn for non-path args: %v", rec) } } }) } func capturingLogger(t *testing.T) (*slog.Logger, *bytes.Buffer) { t.Helper() buf := &bytes.Buffer{} return slog.New(slog.NewJSONHandler(buf, &slog.HandlerOptions{Level: slog.LevelDebug})), buf } func parseLogRecords(t *testing.T, buf *bytes.Buffer) []map[string]any { t.Helper() var recs []map[string]any for _, line := range strings.Split(strings.TrimSpace(buf.String()), "\n") { if line == "" { continue } var rec map[string]any if err := json.Unmarshal([]byte(line), &rec); err != nil { t.Fatalf("invalid log JSON: %q (err=%v)", line, err) } recs = append(recs, rec) } return recs } func TestSparkSubmit_ClusterLocalPathPassthrough(t *testing.T) { deps, repo := testDepsWithDataDir(t) deps.SparkSubmitTimeout = 5 * time.Second 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, }) clusterLocalPath := "/usr/bin/env" req := newToolRequest(SparkSubmitName, map[string]any{ "cluster_id": "cluster-echo", "args": []any{clusterLocalPath, "--class", "Main"}, }) 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{clusterLocalPath, "--class", "Main"} if len(lines) != len(want) { t.Fatalf("echoed lines=%v, want %v", lines, want) } for i, l := range lines { if l != want[i] { t.Errorf("line[%d]=%q, want %q", i, l, want[i]) } } }