package tools import ( "context" "encoding/base64" "fmt" "os" "path/filepath" "regexp" "github.com/mark3labs/mcp-go/mcp" ) 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 ./data/uploads/ so it can be referenced by spark_submit later."), 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) } final := filepath.Join(d.DataDir, "uploads", filename) if err := os.MkdirAll(filepath.Dir(final), 0o750); err != nil { return errResult("upload_file: create uploads dir: " + err.Error()), nil } if err := os.WriteFile(final, data, 0o640); err != nil { return errResult("upload_file: write file: " + err.Error()), nil } result := map[string]any{ "path": fmt.Sprintf("uploads/%s", filename), "size": len(data), } 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 }