package tools import ( "context" "os" "path/filepath" "strings" "testing" "time" "github.com/mark3labs/mcp-go/mcp" "spark-mcp-go/internal/audit" "spark-mcp-go/internal/cluster" "spark-mcp-go/internal/httpclient" "spark-mcp-go/internal/storage" "spark-mcp-go/internal/uploads" ) 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\necho 'Submitted application application_test_0001'\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{ "--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) } if lines[0] != "--master" { t.Errorf("line[0]=%q, want \"--master\"", lines[0]) } 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_PersistsSubmission(t *testing.T) { // Open a dedicated DB so we can insert upload_files metadata for the join. db, err := storage.Open(":memory:") if err != nil { t.Fatalf("open db: %v", err) } t.Cleanup(func() { _ = db.Close() }) dataDir := t.TempDir() store, err := uploads.New(filepath.Join(dataDir, "uploads")) if err != nil { t.Fatalf("new upload store: %v", err) } deps := &Deps{ HTTPClient: httpclient.New(httpclient.Config{Timeout: 5 * time.Second, MaxResponseBytes: 1 << 20}), AuditRepo: audit.NewRepo(db), ClusterRepo: db.Clusters(), SubmissionRepo: db.Submissions(), MaxResponseBytes: 1 << 20, DataDir: dataDir, UploadStore: store, SparkSubmitTimeout: 5 * time.Second, } repo := db.Clusters() subRepo := db.Submissions() fileID, _, _, _, absPath, err := store.Save([]byte("print('hi')\n"), "hello.py") if err != nil { t.Fatalf("save upload: %v", err) } if err := db.Uploads().Create(context.Background(), fileID, "hello.py", 14, "dummy-sha", time.Now()); err != nil { t.Fatalf("create upload record: %v", err) } echoScript := filepath.Join(t.TempDir(), "echo-submit.sh") if err := os.WriteFile(echoScript, []byte("#!/bin/sh\necho 'tracking URL: http://rm.example.com:8088/proxy/application_persist_0001/' >&2\necho 'Submitted application application_persist_0001'\n"), 0o755); err != nil { t.Fatalf("write echo script: %v", err) } createCluster(t, repo, &cluster.Cluster{ ID: "cluster-persist", Name: "Persist", IsActive: true, AuthType: cluster.AuthNone, SparkSubmitExecuteBin: echoScript, }) req := newToolRequest(SparkSubmitName, map[string]any{ "cluster_id": "cluster-persist", "master": "yarn", "deploy_mode": "cluster", "script_path": absPath, "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("unexpected error result: %v", res.Content) } payload := resultJSON(t, res) if payload["app_id"] != "application_persist_0001" { t.Errorf("app_id=%v, want application_persist_0001", payload["app_id"]) } if payload["tracking_url"] != "http://rm.example.com:8088/proxy/application_persist_0001/" { t.Errorf("tracking_url=%v, want tracking URL", payload["tracking_url"]) } subs, err := subRepo.List(context.Background(), storage.SubmissionFilter{Limit: 10}) if err != nil { t.Fatalf("list submissions: %v", err) } if len(subs) != 1 { t.Fatalf("got %d submissions, want 1", len(subs)) } s := subs[0] if s.FileID != fileID { t.Errorf("file_id=%q, want %q", s.FileID, fileID) } if s.AppID != "application_persist_0001" { t.Errorf("app_id=%q, want application_persist_0001", s.AppID) } if s.TrackingURL != "http://rm.example.com:8088/proxy/application_persist_0001/" { t.Errorf("tracking_url=%q, want tracking URL", s.TrackingURL) } if s.ClusterID != "cluster-persist" { t.Errorf("cluster_id=%q, want cluster-persist", s.ClusterID) } if s.FileName != "hello.py" { t.Errorf("file_name=%q, want hello.py", s.FileName) } } 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) } } func TestSparkSubmit_EmptyMaster(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": "", "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 is required") && !strings.Contains(text.Text, "master") { t.Errorf("error text=%q, want mention of master", text.Text) } }