package tools import ( "context" "encoding/base64" "fmt" "path/filepath" "regexp" "github.com/mark3labs/mcp-go/mcp" "spark-mcp-go/internal/audit" ) const UploadFileName = "upload_file" // filenameRegex restricts upload names to a safe, portable character set. var filenameRegex = regexp.MustCompile(`^[a-zA-Z0-9._-]+$`) // NewUploadFileTool returns the schema for the upload_file MCP Tool. func NewUploadFileTool() mcp.Tool { return mcp.NewTool(UploadFileName, mcp.WithDescription("Upload a script or config file to the server's data/uploads/ directory. Returns {file_id, path, name, size, sha256}. To submit it, pass the returned `path` as the `script_path` field in spark_submit, along with the structured fields (master, deploy_mode, queue, executor_memory, executor_cores, num_executors; optional spark_conf, extra_args)."), mcp.WithString("filename", mcp.Required(), mcp.Description("Plain file name without path separators (1-128 chars, [a-zA-Z0-9._-])"), ), mcp.WithString("content", mcp.Required(), mcp.Description("File contents; text or base64-encoded binary"), ), mcp.WithString("encoding", mcp.Description("Encoding of content"), mcp.Enum("text", "base64"), mcp.DefaultString("text"), ), ) } // UploadFileHandler writes user-provided content to DataDir/uploads/filename. func (d *Deps) UploadFileHandler(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { filename, err := req.RequireString("filename") if err != nil { return errResult("upload_file: " + err.Error()), nil } content, err := req.RequireString("content") if err != nil { return errResult("upload_file: " + err.Error()), nil } args := req.GetArguments() encoding := "text" if v, ok := args["encoding"].(string); ok && v != "" { encoding = v } if encoding != "text" && encoding != "base64" { return errResult(fmt.Sprintf("upload_file: invalid encoding %q", encoding)), nil } if err := validateUploadFilename(filename); err != nil { return errResult("upload_file: " + err.Error()), nil } var data []byte switch encoding { case "base64": decoded, err := base64.StdEncoding.DecodeString(content) if err != nil { return errResult("upload_file: decode base64: " + err.Error()), nil } data = decoded default: data = []byte(content) } fileID, _, size, sha256Hex, absPath, err := d.UploadStore.Save(data, filename) if err != nil { return errResult("upload_file: save upload: " + err.Error()), nil } if d.AuditRepo != nil { details, _ := audit.MarshalDetails(map[string]any{ "file_id": fileID, "name": filename, "size": size, "sha256": sha256Hex, }) _ = d.AuditRepo.Insert(ctx, &audit.Entry{ Actor: "tool:upload_file", Action: audit.ActionUploadCreate, ClusterID: fileID, Details: details, }) } result := map[string]any{ "file_id": fileID, "path": absPath, "name": filename, "size": size, "sha256": sha256Hex, } return textResult(encodeJSON(result)), nil } func validateUploadFilename(filename string) error { if len(filename) == 0 || len(filename) > 128 { return fmt.Errorf("filename length must be 1-128") } if filepath.Base(filename) != filename { return fmt.Errorf("filename must not contain path separators or '..'") } if filename == "." || filename == ".." { return fmt.Errorf("filename must not be '.' or '..'") } if !filenameRegex.MatchString(filename) { return fmt.Errorf("filename contains invalid characters") } return nil }