diff --git a/services/slackimport/slackimport.go b/services/slackimport/slackimport.go index f9bd9784f9..57227de7e4 100644 --- a/services/slackimport/slackimport.go +++ b/services/slackimport/slackimport.go @@ -6,6 +6,7 @@ package slackimport import ( "archive/zip" "bytes" + "errors" "image" "io" "mime/multipart" @@ -133,33 +134,54 @@ func (si *SlackImporter) SlackImport(fileData multipart.File, fileSize int64, te posts := make(map[string][]slackPost) uploads := make(map[string]*zip.File) for _, file := range zipreader.File { - if file.UncompressedSize64 > slackImportMaxFileSize { - log.WriteString(i18n.T("api.slackimport.slack_import.zip.file_too_large", map[string]interface{}{"Filename": file.Name})) - continue - } - reader, err := file.Open() + fileReader, err := file.Open() if err != nil { log.WriteString(i18n.T("api.slackimport.slack_import.open.app_error", map[string]interface{}{"Filename": file.Name})) return model.NewAppError("SlackImport", "api.slackimport.slack_import.open.app_error", map[string]interface{}{"Filename": file.Name}, err.Error(), http.StatusInternalServerError), log } + reader := utils.NewLimitedReaderWithError(fileReader, slackImportMaxFileSize) if file.Name == "channels.json" { - publicChannels, _ = slackParseChannels(reader, model.ChannelTypeOpen) + publicChannels, err = slackParseChannels(reader, model.ChannelTypeOpen) + if errors.Is(err, utils.SizeLimitExceeded) { + log.WriteString(i18n.T("api.slackimport.slack_import.zip.file_too_large", map[string]interface{}{"Filename": file.Name})) + continue + } channels = append(channels, publicChannels...) } else if file.Name == "dms.json" { - directChannels, _ = slackParseChannels(reader, model.ChannelTypeDirect) + directChannels, err = slackParseChannels(reader, model.ChannelTypeDirect) + if errors.Is(err, utils.SizeLimitExceeded) { + log.WriteString(i18n.T("api.slackimport.slack_import.zip.file_too_large", map[string]interface{}{"Filename": file.Name})) + continue + } channels = append(channels, directChannels...) } else if file.Name == "groups.json" { - privateChannels, _ = slackParseChannels(reader, model.ChannelTypePrivate) + privateChannels, err = slackParseChannels(reader, model.ChannelTypePrivate) + if errors.Is(err, utils.SizeLimitExceeded) { + log.WriteString(i18n.T("api.slackimport.slack_import.zip.file_too_large", map[string]interface{}{"Filename": file.Name})) + continue + } channels = append(channels, privateChannels...) } else if file.Name == "mpims.json" { - groupChannels, _ = slackParseChannels(reader, model.ChannelTypeGroup) + groupChannels, err = slackParseChannels(reader, model.ChannelTypeGroup) + if errors.Is(err, utils.SizeLimitExceeded) { + log.WriteString(i18n.T("api.slackimport.slack_import.zip.file_too_large", map[string]interface{}{"Filename": file.Name})) + continue + } channels = append(channels, groupChannels...) } else if file.Name == "users.json" { - users, _ = slackParseUsers(reader) + users, err = slackParseUsers(reader) + if errors.Is(err, utils.SizeLimitExceeded) { + log.WriteString(i18n.T("api.slackimport.slack_import.zip.file_too_large", map[string]interface{}{"Filename": file.Name})) + continue + } } else { spl := strings.Split(file.Name, "/") if len(spl) == 2 && strings.HasSuffix(spl[1], ".json") { - newposts, _ := slackParsePosts(reader) + newposts, err := slackParsePosts(reader) + if errors.Is(err, utils.SizeLimitExceeded) { + log.WriteString(i18n.T("api.slackimport.slack_import.zip.file_too_large", map[string]interface{}{"Filename": file.Name})) + continue + } channel := spl[0] if _, ok := posts[channel]; !ok { posts[channel] = newposts diff --git a/utils/file.go b/utils/file.go index b4de4bdaf9..175aa9ac21 100644 --- a/utils/file.go +++ b/utils/file.go @@ -4,6 +4,7 @@ package utils import ( + "errors" "fmt" "io" "io/ioutil" @@ -112,3 +113,23 @@ func CopyDir(src string, dst string) (err error) { return } + +var SizeLimitExceeded = errors.New("Size limit exceeded") + +type LimitedReaderWithError struct { + limitedReader *io.LimitedReader +} + +func NewLimitedReaderWithError(reader io.Reader, maxBytes int64) *LimitedReaderWithError { + return &LimitedReaderWithError{ + limitedReader: &io.LimitedReader{R: reader, N: maxBytes + 1}, + } +} + +func (l *LimitedReaderWithError) Read(p []byte) (int, error) { + n, err := l.limitedReader.Read(p) + if l.limitedReader.N <= 0 && err == io.EOF { + return n, SizeLimitExceeded + } + return n, err +} diff --git a/utils/file_test.go b/utils/file_test.go index cdaff587f5..e74c273cb7 100644 --- a/utils/file_test.go +++ b/utils/file_test.go @@ -4,6 +4,9 @@ package utils import ( + "bytes" + "crypto/rand" + "io" "io/ioutil" "os" "path/filepath" @@ -62,3 +65,65 @@ func TestCopyDir(t *testing.T) { err = CopyDir(srcDir, dstDir) assert.Error(t, err) } +func TestLimitedReaderWithError(t *testing.T) { + t.Run("read less than max size", func(t *testing.T) { + maxBytes := 10 + randomBytes := make([]byte, maxBytes) + n, err := rand.Read(randomBytes) + require.NoError(t, err) + require.Equal(t, n, maxBytes) + + lr := NewLimitedReaderWithError(bytes.NewReader(randomBytes), int64(maxBytes)) + smallerBuf := make([]byte, maxBytes-3) + _, err = io.ReadFull(lr, smallerBuf) + require.NoError(t, err) + }) + + t.Run("read equal to max size", func(t *testing.T) { + maxBytes := 10 + randomBytes := make([]byte, maxBytes) + n, err := rand.Read(randomBytes) + require.NoError(t, err) + require.Equal(t, n, maxBytes) + + lr := NewLimitedReaderWithError(bytes.NewReader(randomBytes), int64(maxBytes)) + buf := make([]byte, maxBytes) + _, err = io.ReadFull(lr, buf) + require.Truef(t, err == nil || err == io.EOF, "err must be nil or %v, got %v", io.EOF, err) + }) + + t.Run("single read, larger than max size", func(t *testing.T) { + maxBytes := 5 + moreThanMaxBytes := maxBytes + 10 + randomBytes := make([]byte, moreThanMaxBytes) + n, err := rand.Read(randomBytes) + require.NoError(t, err) + require.Equal(t, moreThanMaxBytes, n) + + lr := NewLimitedReaderWithError(bytes.NewReader(randomBytes), int64(maxBytes)) + buf := make([]byte, moreThanMaxBytes) + _, err = io.ReadFull(lr, buf) + require.Error(t, err) + require.Equal(t, SizeLimitExceeded, err) + }) + + t.Run("multiple small reads, total larger than max size", func(t *testing.T) { + maxBytes := 10 + lessThanMaxBytes := maxBytes - 4 + randomBytesLen := maxBytes * 2 + randomBytes := make([]byte, randomBytesLen) + n, err := rand.Read(randomBytes) + require.NoError(t, err) + require.Equal(t, randomBytesLen, n) + + lr := NewLimitedReaderWithError(bytes.NewReader(randomBytes), int64(maxBytes)) + buf := make([]byte, lessThanMaxBytes) + _, err = io.ReadFull(lr, buf) + require.NoError(t, err) + + // lets do it again + _, err = io.ReadFull(lr, buf) + require.Error(t, err) + require.Equal(t, SizeLimitExceeded, err) + }) +}