MM-43716: enforce size limits in slack importer (#20144)

Automatic Merge
Этот коммит содержится в:
Ashish Bhate
2022-05-12 18:39:55 +05:30
коммит произвёл GitHub
родитель 87a719b4b4
Коммит 3c625743e5
3 изменённых файлов: 119 добавлений и 11 удалений

Просмотреть файл

@@ -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

Просмотреть файл

@@ -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
}

Просмотреть файл

@@ -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)
})
}