diff --git a/internal/mcp/tools/deps.go b/internal/mcp/tools/deps.go index 718c09c..f9989b9 100644 --- a/internal/mcp/tools/deps.go +++ b/internal/mcp/tools/deps.go @@ -16,6 +16,7 @@ type Deps struct { Logger *slog.Logger AuditRepo *audit.Repo ClusterRepo *storage.ClusterRepo + SubmissionRepo *storage.SubmissionRepo SparkSubmitTimeout time.Duration HTTPClient *httpclient.Client MaxResponseBytes int64 diff --git a/internal/mcp/tools/fetch_upload_test.go b/internal/mcp/tools/fetch_upload_test.go index 683bffc..9d1b46a 100644 --- a/internal/mcp/tools/fetch_upload_test.go +++ b/internal/mcp/tools/fetch_upload_test.go @@ -25,7 +25,7 @@ import ( // testDepsWithDataDir returns dependencies backed by an in-memory DB and a // temporary data directory. -func testDepsWithDataDir(t *testing.T) (*Deps, *storage.ClusterRepo, *audit.Repo) { +func testDepsWithDataDir(t *testing.T) (*Deps, *storage.ClusterRepo, *storage.SubmissionRepo, *audit.Repo) { t.Helper() db, err := storage.Open(":memory:") if err != nil { @@ -46,10 +46,11 @@ func testDepsWithDataDir(t *testing.T) (*Deps, *storage.ClusterRepo, *audit.Repo }), AuditRepo: audit.NewRepo(db), ClusterRepo: db.Clusters(), + SubmissionRepo: db.Submissions(), MaxResponseBytes: 1 << 20, DataDir: dataDir, UploadStore: uploadStore, - }, db.Clusters(), audit.NewRepo(db) + }, db.Clusters(), db.Submissions(), audit.NewRepo(db) } // createCluster creates a cluster in the repository with the given fields. @@ -100,7 +101,7 @@ func TestFetchURL(t *testing.T) { })) defer srv.Close() - deps, repo, _ := testDepsWithDataDir(t) + deps, repo, _, _ := testDepsWithDataDir(t) createCluster(t, repo, &cluster.Cluster{ ID: "cluster-a", Name: "Cluster A", @@ -294,7 +295,7 @@ func TestFetchURL_RedirectPreservesAuth(t *testing.T) { })) defer server1.Close() - deps, repo, _ := testDepsWithDataDir(t) + deps, repo, _, _ := testDepsWithDataDir(t) createCluster(t, repo, &cluster.Cluster{ ID: "cluster-redirect", Name: "Redirect", @@ -323,7 +324,7 @@ func TestFetchURL_RedirectPreservesAuth(t *testing.T) { } func TestUploadFile(t *testing.T) { - deps, _, auditRepo := testDepsWithDataDir(t) + deps, _, _, auditRepo := testDepsWithDataDir(t) req := newToolRequest(UploadFileName, map[string]any{ "filename": "hello.txt", @@ -387,7 +388,7 @@ func TestUploadFile(t *testing.T) { } func TestUploadFile_PathTraversal(t *testing.T) { - deps, _, _ := testDepsWithDataDir(t) + deps, _, _, _ := testDepsWithDataDir(t) cases := []string{ "../../../etc/passwd", @@ -418,7 +419,7 @@ func TestUploadFile_PathTraversal(t *testing.T) { } func TestUploadFile_Base64(t *testing.T) { - deps, _, _ := testDepsWithDataDir(t) + deps, _, _, _ := testDepsWithDataDir(t) req := newToolRequest(UploadFileName, map[string]any{ "filename": "hello.bin", diff --git a/internal/mcp/tools/spark_submit.go b/internal/mcp/tools/spark_submit.go index 4be2493..351a65b 100644 --- a/internal/mcp/tools/spark_submit.go +++ b/internal/mcp/tools/spark_submit.go @@ -6,10 +6,12 @@ import ( "log/slog" "sort" "strconv" + "time" "github.com/mark3labs/mcp-go/mcp" "spark-mcp-go/internal/executor" + "spark-mcp-go/internal/storage" ) const SparkSubmitName = "spark_submit" @@ -98,7 +100,8 @@ func (d *Deps) SparkSubmitHandler(ctx context.Context, req mcp.CallToolRequest) // 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 _, err := d.UploadStore.Validate(scriptPath); err != nil { + fileID, err := d.UploadStore.Validate(scriptPath) + if err != nil { if d.Logger != nil { d.Logger.Warn("spark_submit.unminted_path_rejected", slog.String("path", scriptPath), slog.String("err", err.Error())) } @@ -173,35 +176,58 @@ func (d *Deps) SparkSubmitHandler(ctx context.Context, req mcp.CallToolRequest) }) if err != nil { callLog.WithError(err) - // Even on error, result carries ExitCode/Stderr for the LLM. if result != nil { - callLog.WithResult(map[string]any{ - "binary": binary, - "argv": cmd, - "app_id": result.AppID, - "exit_code": result.ExitCode, - "duration_ms": result.DurationMS, - }) + callLog.WithResult(sparkSubmitResultLog(binary, cmd, result)) 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(), + "app_id": result.AppID, + "tracking_url": result.TrackingURL, + "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{ + if result.AppID != "" && d.SubmissionRepo != nil { + sub := storage.Submission{ + FileID: fileID, + AppID: result.AppID, + TrackingURL: result.TrackingURL, + ClusterID: clusterID, + ExitCode: result.ExitCode, + DurationMS: result.DurationMS, + SubmittedAt: time.Now(), + } + if createErr := d.SubmissionRepo.Create(ctx, sub); createErr != nil { + if d.Logger != nil { + d.Logger.Warn("spark_submit.persist_failed", slog.String("file_id", fileID), slog.String("app_id", result.AppID), slog.String("err", createErr.Error())) + } else { + slog.Default().Warn("spark_submit.persist_failed", slog.String("file_id", fileID), slog.String("app_id", result.AppID), slog.String("err", createErr.Error())) + } + } + } + + callLog.WithResult(sparkSubmitResultLog(binary, cmd, result)) + return textResult(encodeJSON(result)), nil +} + +// sparkSubmitResultLog builds the standard result map used for both logging +// and the admin audit trail. +func sparkSubmitResultLog(binary string, cmd []string, result *executor.Result) map[string]any { + m := map[string]any{ "binary": binary, "argv": cmd, "app_id": result.AppID, "exit_code": result.ExitCode, "duration_ms": result.DurationMS, - }) - return textResult(encodeJSON(result)), nil + } + if result.TrackingURL != "" { + m["tracking_url"] = result.TrackingURL + } + return m } // SparkSubmitCommandOpts holds the structured arguments used to build the diff --git a/internal/mcp/tools/spark_submit_path_test.go b/internal/mcp/tools/spark_submit_path_test.go index 2faf54f..6ce6aa6 100644 --- a/internal/mcp/tools/spark_submit_path_test.go +++ b/internal/mcp/tools/spark_submit_path_test.go @@ -10,11 +10,15 @@ import ( "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, repo, _, _ := testDepsWithDataDir(t) deps.SparkSubmitTimeout = 5 * time.Second store := deps.UploadStore @@ -25,7 +29,7 @@ func TestSparkSubmit_StructuredCommand(t *testing.T) { 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 { + 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) } @@ -97,8 +101,107 @@ func TestSparkSubmit_StructuredCommand(t *testing.T) { } } +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) + deps, _, _, _ := testDepsWithDataDir(t) store := deps.UploadStore _, _, _, _, mintedPath, err := store.Save([]byte("# dummy\n"), "dummy.py") if err != nil { @@ -131,7 +234,7 @@ func TestSparkSubmit_MissingRequiredField(t *testing.T) { } func TestSparkSubmit_BadSparkConfValue(t *testing.T) { - deps, _, _ := testDepsWithDataDir(t) + deps, _, _, _ := testDepsWithDataDir(t) store := deps.UploadStore _, _, _, _, mintedPath, err := store.Save([]byte("# dummy\n"), "dummy.py") if err != nil { @@ -166,7 +269,7 @@ func TestSparkSubmit_BadSparkConfValue(t *testing.T) { } func TestSparkSubmit_RejectsNonMintedPath(t *testing.T) { - deps, _, _ := testDepsWithDataDir(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 { @@ -200,7 +303,7 @@ func TestSparkSubmit_RejectsNonMintedPath(t *testing.T) { } func TestSparkSubmit_EmptyMaster(t *testing.T) { - deps, _, _ := testDepsWithDataDir(t) + deps, _, _, _ := testDepsWithDataDir(t) store := deps.UploadStore _, _, _, _, mintedPath, err := store.Save([]byte("# dummy\n"), "dummy.py") if err != nil {