Files
spark-mcp/internal/mcp/tools/fetch_upload_test.go
T
tao.chenandClaude 876f1c9edb uploads: dual-write uploads to DB, startup backfill, and audit log
Store now indexes every Save into the upload_files table via SetRepo.

DB failures are logged as warnings and do not fail the upload because

.meta.json remains the source of truth and startup backfill recovers.

Add Store.Backfill to walk Root at startup and insert index rows for

any pre-existing .meta.json sidecars, swallowing duplicate-key races.

The upload_file MCP Tool now writes an audit_log entry on success.

Tests cover dual-write args, repo-error non-failure, and backfill

skipping existing rows.

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-14 11:30:39 +08:00

453 lines
12 KiB
Go

package tools
import (
"context"
"encoding/base64"
"encoding/json"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"regexp"
"strconv"
"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"
)
// testDepsWithDataDir returns dependencies backed by an in-memory DB and a
// temporary data directory.
func testDepsWithDataDir(t *testing.T) (*Deps, *storage.ClusterRepo, *audit.Repo) {
t.Helper()
db, err := storage.Open(":memory:")
if err != nil {
t.Fatalf("open db: %v", err)
}
t.Cleanup(func() { _ = db.Close() })
dataDir := t.TempDir()
uploadStore, err := uploads.New(filepath.Join(dataDir, "uploads"))
if err != nil {
t.Fatalf("new upload store: %v", err)
}
return &Deps{
HTTPClient: httpclient.New(httpclient.Config{
Timeout: 5 * time.Second,
MaxResponseBytes: 1 << 20,
}),
AuditRepo: audit.NewRepo(db),
ClusterRepo: db.Clusters(),
MaxResponseBytes: 1 << 20,
DataDir: dataDir,
UploadStore: uploadStore,
}, db.Clusters(), audit.NewRepo(db)
}
// createCluster creates a cluster in the repository with the given fields.
func createCluster(t *testing.T, repo *storage.ClusterRepo, c *cluster.Cluster) {
t.Helper()
if c.RMURL == "" {
c.RMURL = "http://rm.example.com:8088"
}
if c.SHSURL == "" {
c.SHSURL = "http://shs.example.com:18080"
}
if c.SparkSubmitExecuteBin == "" {
c.SparkSubmitExecuteBin = "/usr/bin/spark-submit"
}
if err := repo.Create(context.Background(), c); err != nil {
t.Fatalf("create cluster: %v", err)
}
}
func resultText(t *testing.T, res *mcp.CallToolResult) string {
t.Helper()
if res.IsError {
t.Fatalf("unexpected error result: %v", res.Content)
}
if len(res.Content) == 0 {
t.Fatal("empty result content")
}
text, ok := mcp.AsTextContent(res.Content[0])
if !ok {
t.Fatalf("content is not text: %T", res.Content[0])
}
return text.Text
}
func resultJSON(t *testing.T, res *mcp.CallToolResult) map[string]any {
t.Helper()
var out map[string]any
if err := json.Unmarshal([]byte(resultText(t, res)), &out); err != nil {
t.Fatalf("unmarshal result: %v", err)
}
return out
}
func TestFetchURL(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.Write([]byte(`{"hello":"world"}`))
}))
defer srv.Close()
deps, repo, _ := testDepsWithDataDir(t)
createCluster(t, repo, &cluster.Cluster{
ID: "cluster-a",
Name: "Cluster A",
IsActive: true,
AuthType: cluster.AuthNone,
URLAllowlist: []string{srv.URL},
})
t.Run("success", func(t *testing.T) {
req := newToolRequest(FetchURLName, map[string]any{
"cluster_id": "cluster-a",
"url": srv.URL + "/foo",
"method": "GET",
})
res, err := deps.FetchURLHandler(context.Background(), req)
if err != nil {
t.Fatalf("handler error: %v", err)
}
payload := resultJSON(t, res)
if payload["status_code"] != float64(200) {
t.Errorf("status_code=%v, want 200", payload["status_code"])
}
if !strings.Contains(payload["body"].(string), "hello") {
t.Errorf("body missing hello: %v", payload["body"])
}
})
t.Run("cluster not found", func(t *testing.T) {
req := newToolRequest(FetchURLName, map[string]any{
"cluster_id": "missing",
"url": srv.URL,
})
res, err := deps.FetchURLHandler(context.Background(), req)
if err != nil {
t.Fatalf("handler error: %v", err)
}
if !res.IsError {
t.Errorf("expected error result, got: %v", res.Content)
}
})
t.Run("private ip without allowlist", func(t *testing.T) {
createCluster(t, repo, &cluster.Cluster{
ID: "cluster-private",
Name: "Private",
IsActive: true,
AuthType: cluster.AuthNone,
})
req := newToolRequest(FetchURLName, map[string]any{
"cluster_id": "cluster-private",
"url": "http://127.0.0.1:12345/",
})
res, err := deps.FetchURLHandler(context.Background(), req)
if err != nil {
t.Fatalf("handler error: %v", err)
}
if !res.IsError {
t.Errorf("expected error result, got: %v", res.Content)
}
})
t.Run("public url without allowlist", func(t *testing.T) {
createCluster(t, repo, &cluster.Cluster{
ID: "cluster-public",
Name: "Public",
IsActive: true,
AuthType: cluster.AuthNone,
})
req := newToolRequest(FetchURLName, map[string]any{
"cluster_id": "cluster-public",
"url": "http://example.com/",
})
res, err := deps.FetchURLHandler(context.Background(), req)
if err != nil {
t.Fatalf("handler error: %v", err)
}
if !res.IsError {
t.Errorf("expected error result, got: %v", res.Content)
}
})
t.Run("invalid method", func(t *testing.T) {
req := newToolRequest(FetchURLName, map[string]any{
"cluster_id": "cluster-a",
"url": srv.URL,
"method": "INVALID",
})
res, err := deps.FetchURLHandler(context.Background(), req)
if err != nil {
t.Fatalf("handler error: %v", err)
}
if !res.IsError {
t.Errorf("expected error result, got: %v", res.Content)
}
})
t.Run("basic auth header sent", func(t *testing.T) {
authSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
auth := r.Header.Get("Authorization")
if auth == "" {
http.Error(w, "missing auth", http.StatusUnauthorized)
return
}
parts := strings.SplitN(auth, " ", 2)
if len(parts) != 2 || parts[0] != "Basic" {
http.Error(w, "bad auth", http.StatusUnauthorized)
return
}
w.Write([]byte(parts[1]))
}))
defer authSrv.Close()
createCluster(t, repo, &cluster.Cluster{
ID: "cluster-auth",
Name: "Auth",
IsActive: true,
AuthType: cluster.AuthBasic,
AuthUsername: "u",
AuthPassword: "p",
URLAllowlist: []string{authSrv.URL},
})
req := newToolRequest(FetchURLName, map[string]any{
"cluster_id": "cluster-auth",
"url": authSrv.URL,
})
res, err := deps.FetchURLHandler(context.Background(), req)
if err != nil {
t.Fatalf("handler error: %v", err)
}
payload := resultJSON(t, res)
got := payload["body"].(string)
want := base64.StdEncoding.EncodeToString([]byte("u:p"))
if got != want {
t.Errorf("auth body=%q, want %q", got, want)
}
})
t.Run("user authorization header preserved", func(t *testing.T) {
authSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Write([]byte(r.Header.Get("Authorization")))
}))
defer authSrv.Close()
createCluster(t, repo, &cluster.Cluster{
ID: "cluster-user-auth",
Name: "UserAuth",
IsActive: true,
AuthType: cluster.AuthBasic,
AuthUsername: "u",
AuthPassword: "p",
URLAllowlist: []string{authSrv.URL},
})
req := newToolRequest(FetchURLName, map[string]any{
"cluster_id": "cluster-user-auth",
"url": authSrv.URL,
"headers": map[string]any{
"Authorization": "Bearer user-token",
},
})
res, err := deps.FetchURLHandler(context.Background(), req)
if err != nil {
t.Fatalf("handler error: %v", err)
}
payload := resultJSON(t, res)
got := payload["body"].(string)
if got != "Bearer user-token" {
t.Errorf("authorization header=%q, want user token preserved", got)
}
})
}
func TestFetchURL_RedirectPreservesAuth(t *testing.T) {
var server2URL string
server2 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
auth := r.Header.Get("Authorization")
parts := strings.SplitN(auth, " ", 2)
if len(parts) == 2 && parts[0] == "Basic" {
w.Write([]byte(parts[1]))
return
}
w.Write([]byte(auth))
}))
defer server2.Close()
server2URL = server2.URL
server1 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Location", server2URL+"/final")
w.WriteHeader(http.StatusTemporaryRedirect)
}))
defer server1.Close()
deps, repo, _ := testDepsWithDataDir(t)
createCluster(t, repo, &cluster.Cluster{
ID: "cluster-redirect",
Name: "Redirect",
IsActive: true,
AuthType: cluster.AuthBasic,
AuthUsername: "u",
AuthPassword: "p",
URLAllowlist: []string{server1.URL, server2.URL},
})
req := newToolRequest(FetchURLName, map[string]any{
"cluster_id": "cluster-redirect",
"url": server1.URL + "/start",
})
res, err := deps.FetchURLHandler(context.Background(), req)
if err != nil {
t.Fatalf("handler error: %v", err)
}
payload := resultJSON(t, res)
got := payload["body"].(string)
want := base64.StdEncoding.EncodeToString([]byte("u:p"))
if got != want {
t.Errorf("redirect auth body=%q, want %q", got, want)
}
}
func TestUploadFile(t *testing.T) {
deps, _, auditRepo := testDepsWithDataDir(t)
req := newToolRequest(UploadFileName, map[string]any{
"filename": "hello.txt",
"content": "hello world",
})
res, err := deps.UploadFileHandler(context.Background(), req)
if err != nil {
t.Fatalf("handler error: %v", err)
}
payload := resultJSON(t, res)
fileID, ok := payload["file_id"].(string)
if !ok || !regexp.MustCompile(`^[0-9a-f]{32}$`).MatchString(fileID) {
t.Errorf("file_id=%v, want 32 lowercase hex chars", payload["file_id"])
}
path, ok := payload["path"].(string)
if !ok || !filepath.IsAbs(path) || !strings.HasSuffix(path, "/"+fileID) {
t.Errorf("path=%v, want absolute path ending in /%s", payload["path"], fileID)
}
if payload["name"] != "hello.txt" {
t.Errorf("name=%v, want hello.txt", payload["name"])
}
if payload["size"] != float64(len("hello world")) {
t.Errorf("size=%v, want %d", payload["size"], len("hello world"))
}
sha, ok := payload["sha256"].(string)
if !ok || !regexp.MustCompile(`^[0-9a-f]{64}$`).MatchString(sha) {
t.Errorf("sha256=%v, want 64-char hex", payload["sha256"])
}
got, err := os.ReadFile(path)
if err != nil {
t.Fatalf("read uploaded file: %v", err)
}
if string(got) != "hello world" {
t.Errorf("content=%q, want %q", got, "hello world")
}
info, err := os.Stat(path)
if err != nil {
t.Fatalf("stat uploaded file: %v", err)
}
if info.Mode().Perm() != 0o640 {
t.Errorf("mode=%o, want %o", info.Mode().Perm(), 0o640)
}
entries, err := auditRepo.List(context.Background(), 100)
if err != nil {
t.Fatalf("list audit: %v", err)
}
found := false
for _, e := range entries {
if e.Action == audit.ActionUploadCreate && e.Actor == "tool:upload_file" && e.ClusterID == fileID {
found = true
break
}
}
if !found {
t.Errorf("missing upload.create audit entry for file_id %s", fileID)
}
}
func TestUploadFile_PathTraversal(t *testing.T) {
deps, _, _ := testDepsWithDataDir(t)
cases := []string{
"../../../etc/passwd",
"/etc/passwd",
"foo/bar",
"foo\\bar",
"",
"hello world",
"中文.txt",
"..",
}
for _, name := range cases {
t.Run(strconv.Quote(name), func(t *testing.T) {
req := newToolRequest(UploadFileName, map[string]any{
"filename": name,
"content": "x",
})
res, err := deps.UploadFileHandler(context.Background(), req)
if err != nil {
t.Fatalf("handler error: %v", err)
}
if !res.IsError {
t.Errorf("expected error for filename %q, got: %v", name, res.Content)
}
})
}
}
func TestUploadFile_Base64(t *testing.T) {
deps, _, _ := testDepsWithDataDir(t)
req := newToolRequest(UploadFileName, map[string]any{
"filename": "hello.bin",
"content": "aGVsbG8=",
"encoding": "base64",
})
res, err := deps.UploadFileHandler(context.Background(), req)
if err != nil {
t.Fatalf("handler error: %v", err)
}
payload := resultJSON(t, res)
path, ok := payload["path"].(string)
if !ok || !filepath.IsAbs(path) {
t.Errorf("path=%v, want absolute path", payload["path"])
}
if payload["name"] != "hello.bin" {
t.Errorf("name=%v, want hello.bin", payload["name"])
}
if payload["size"] != float64(len("hello")) {
t.Errorf("size=%v, want %d", payload["size"], len("hello"))
}
got, err := os.ReadFile(path)
if err != nil {
t.Fatalf("read uploaded file: %v", err)
}
if string(got) != "hello" {
t.Errorf("content=%q, want %q", got, "hello")
}
}