diff --git a/internal/filedata/archive_reader.go b/internal/filedata/archive_reader.go index f0527b533..8bb9bab66 100644 --- a/internal/filedata/archive_reader.go +++ b/internal/filedata/archive_reader.go @@ -27,6 +27,8 @@ import ( const ( environmentsChecksumFileName = "checksum.md5" maxDecompressedFileSize = 1024 * 1024 * 200 // arbitrary 200MB limit to avoid decompression bombs + maxDecompressedArchiveSize = 1024 * 1024 * 1024 + maxArchiveEntryCount = 10_000 ) // archiveReader is the low-level implementation of unarchiving a data file and reading the environments. @@ -196,7 +198,13 @@ func readUncompressedArchive(filePath, targetDir string) error { } func readTar(r io.Reader, targetDir string) error { + return readTarWithLimits(r, targetDir, maxDecompressedArchiveSize, maxArchiveEntryCount) +} + +func readTarWithLimits(r io.Reader, targetDir string, maxArchiveSize int64, maxEntryCount int) error { tr := tar.NewReader(r) + var archiveSize int64 + entryCount := 0 for { h, err := tr.Next() if err != nil { @@ -206,10 +214,19 @@ func readTar(r io.Reader, targetDir string) error { return err } + entryCount++ + if entryCount > maxEntryCount { + return errArchiveHasTooManyEntries(maxEntryCount) + } + // In our archive format, there should be no subdirectories, just top-level files if h.Typeflag != tar.TypeReg { continue } + if h.Size > maxArchiveSize-archiveSize { + return errUncompressedArchiveTooBig(maxArchiveSize) + } + archiveSize += h.Size outPath, err := securejoin.SecureJoin(targetDir, h.Name) if err != nil { return err // COVERAGE: can't cause this condition in unit tests diff --git a/internal/filedata/archive_reader_test.go b/internal/filedata/archive_reader_test.go index e24f4ec9c..0331c3dd4 100644 --- a/internal/filedata/archive_reader_test.go +++ b/internal/filedata/archive_reader_test.go @@ -1,6 +1,9 @@ package filedata import ( + "archive/tar" + "bytes" + "fmt" "os" "sort" "testing" @@ -138,6 +141,69 @@ func TestEnvironmentSDKDataItemOfUnknownKindIsIgnored(t *testing.T) { }) } +func TestReadTarEnforcesAggregateLimits(t *testing.T) { + for _, params := range []struct { + name string + fileSizes []int + maxArchiveSize int64 + maxEntryCount int + expectedMessage string + }{ + { + name: "archive size", + fileSizes: []int{4, 4, 4}, + maxArchiveSize: 10, + maxEntryCount: 3, + expectedMessage: "contents exceeded 10 bytes", + }, + { + name: "entry count", + fileSizes: []int{1, 1, 1}, + maxArchiveSize: 3, + maxEntryCount: 2, + expectedMessage: "more than 2 entries", + }, + } { + t.Run(params.name, func(t *testing.T) { + archive := makeTarArchive(t, params.fileSizes...) + + err := readTarWithLimits( + bytes.NewReader(archive), + t.TempDir(), + params.maxArchiveSize, + params.maxEntryCount, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), params.expectedMessage) + }) + } +} + +func TestReadTarAcceptsArchiveAtAggregateLimits(t *testing.T) { + archive := makeTarArchive(t, 2, 3) + + err := readTarWithLimits(bytes.NewReader(archive), t.TempDir(), 5, 2) + + require.NoError(t, err) +} + +func makeTarArchive(t *testing.T, fileSizes ...int) []byte { + var data bytes.Buffer + tw := tar.NewWriter(&data) + for index, size := range fileSizes { + require.NoError(t, tw.WriteHeader(&tar.Header{ + Name: fmt.Sprintf("file-%d", index), + Mode: 0600, + Size: int64(size), + })) + _, err := tw.Write(make([]byte, size)) + require.NoError(t, err) + } + require.NoError(t, tw.Close()) + return data.Bytes() +} + func verifyAllEnvironmentData(t *testing.T, ar *archiveReader) { var expectedEnvIDs []config.EnvironmentID for _, te := range allTestEnvs { diff --git a/internal/filedata/errors_and_messages.go b/internal/filedata/errors_and_messages.go index 5e7c53f84..7a00fb319 100644 --- a/internal/filedata/errors_and_messages.go +++ b/internal/filedata/errors_and_messages.go @@ -43,3 +43,11 @@ func errUncompressedFileTooBig(fileName string, maxSize int64) error { return fmt.Errorf("detected malformed or malicious archive file; it contained a file %q with a size >= %d bytes", fileName, maxSize) } + +func errUncompressedArchiveTooBig(maxSize int64) error { + return fmt.Errorf("detected malformed or malicious archive file; its contents exceeded %d bytes", maxSize) +} + +func errArchiveHasTooManyEntries(maxEntries int) error { + return fmt.Errorf("detected malformed or malicious archive file; it contained more than %d entries", maxEntries) +}