338 lines
9.8 KiB
Go
338 lines
9.8 KiB
Go
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)
|
|
}
|
|
}
|