diff --git a/archives_test.go b/archives_test.go index 524e674..8618263 100644 --- a/archives_test.go +++ b/archives_test.go @@ -6,6 +6,7 @@ import ( "bytes" "compress/gzip" "encoding/base64" + "encoding/binary" "errors" "fmt" "io" @@ -858,6 +859,67 @@ func TestOpenTarRejectsCumulativeOverflow(t *testing.T) { } } +func TestOpenTarRejectsTooManyEntries(t *testing.T) { + oldMax := maxArchiveEntries + maxArchiveEntries = 2 + defer func() { maxArchiveEntries = oldMax }() + + var buf bytes.Buffer + tw := tar.NewWriter(&buf) + for i := 0; i < 3; i++ { + _ = tw.WriteHeader(&tar.Header{Name: fmt.Sprintf("empty-%d", i), Mode: 0644}) + } + _ = tw.Close() + + _, err := openTar(buf.Bytes(), "") + if !errors.Is(err, ErrEntryLimit) { + t.Fatalf("expected ErrEntryLimit, got: %v", err) + } +} + +func TestOpenZipRejectsTooManyEntries(t *testing.T) { + oldMax := maxArchiveEntries + maxArchiveEntries = 2 + defer func() { maxArchiveEntries = oldMax }() + + var buf bytes.Buffer + zw := zip.NewWriter(&buf) + for i := 0; i < 3; i++ { + _, _ = zw.Create(fmt.Sprintf("empty-%d", i)) + } + _ = zw.Close() + raw := buf.Bytes() + directoryEnd := bytes.LastIndex(raw, []byte{'P', 'K', 0x05, 0x06}) + if directoryEnd < 0 { + t.Fatal("zip end-of-directory record not found") + } + binary.LittleEndian.PutUint16(raw[directoryEnd+8:], 1) + binary.LittleEndian.PutUint16(raw[directoryEnd+10:], 1) + + _, err := OpenBytes("too-many.zip", raw) + if !errors.Is(err, ErrEntryLimit) { + t.Fatalf("expected ErrEntryLimit, got: %v", err) + } +} + +func TestOpenGemRejectsTooManyEntries(t *testing.T) { + oldMax := maxArchiveEntries + maxArchiveEntries = 2 + defer func() { maxArchiveEntries = oldMax }() + + var buf bytes.Buffer + tw := tar.NewWriter(&buf) + for i := 0; i < 3; i++ { + _ = tw.WriteHeader(&tar.Header{Name: fmt.Sprintf("empty-%d", i), Mode: 0644}) + } + _ = tw.Close() + + _, err := openGem(buf.Bytes()) + if !errors.Is(err, ErrEntryLimit) { + t.Fatalf("expected ErrEntryLimit, got: %v", err) + } +} + func TestOpenGemRejectsOversizedData(t *testing.T) { oldMax := maxDecompressedSize maxDecompressedSize = 512 diff --git a/conda.go b/conda.go index dba636f..c4786f4 100644 --- a/conda.go +++ b/conda.go @@ -18,10 +18,16 @@ import ( // merged tarReader with raw pointing at the outer .conda bytes so Hash // matches the digest anaconda.org publishes in repodata.json. func openConda(raw []byte) (*tarReader, error) { + if err := checkZipEntryCount(raw); err != nil { + return nil, err + } zr, err := zip.NewReader(bytes.NewReader(raw), int64(len(raw))) if err != nil { return nil, fmt.Errorf("opening conda zip: %w", err) } + if err := checkArchiveEntryCount(len(zr.File)); err != nil { + return nil, err + } var files []tarFileEntry var total int64 @@ -32,7 +38,7 @@ func openConda(raw []byte) (*tarReader, error) { if !strings.HasPrefix(f.Name, "pkg-") && !strings.HasPrefix(f.Name, "info-") { continue } - entries, size, err := readCondaMember(f) + entries, size, err := readCondaMember(f, len(files)) if err != nil { return nil, err } @@ -56,7 +62,7 @@ func openConda(raw []byte) (*tarReader, error) { return &tarReader{raw: raw, files: files, index: index}, nil } -func readCondaMember(f *zip.File) ([]tarFileEntry, int64, error) { +func readCondaMember(f *zip.File, initialEntryCount int) ([]tarFileEntry, int64, error) { rc, err := f.Open() if err != nil { return nil, 0, fmt.Errorf("opening %s: %w", f.Name, err) @@ -71,7 +77,7 @@ func readCondaMember(f *zip.File) ([]tarFileEntry, int64, error) { return nil, 0, fmt.Errorf("%w: %s exceeds %d bytes", ErrDecompressLimit, f.Name, maxDecompressedSize) } - tr, err := openTar(data, "zstd") + tr, err := openTarWithInitialEntryCount(data, "zstd", initialEntryCount) if err != nil { return nil, 0, fmt.Errorf("opening %s: %w", f.Name, err) } diff --git a/conda_test.go b/conda_test.go index 1b09dcf..f284514 100644 --- a/conda_test.go +++ b/conda_test.go @@ -186,6 +186,58 @@ func TestOpenCondaRejectsCumulativeOverflow(t *testing.T) { } } +func TestOpenCondaRejectsTooManyEntriesAcrossMembers(t *testing.T) { + oldMax := maxArchiveEntries + maxArchiveEntries = 3 + defer func() { maxArchiveEntries = oldMax }() + + _, err := openConda(createTestConda(t)) + if !errors.Is(err, ErrEntryLimit) { + t.Fatalf("expected ErrEntryLimit, got: %v", err) + } +} + +func TestOpenCondaStopsAtCombinedEntryLimit(t *testing.T) { + oldEntryMax := maxArchiveEntries + oldSizeMax := maxDecompressedSize + maxArchiveEntries = 2 + maxDecompressedSize = 512 + defer func() { + maxArchiveEntries = oldEntryMax + maxDecompressedSize = oldSizeMax + }() + + members := []struct { + name string + data []byte + }{ + {"pkg-a-1.tar.zst", writeTarZst(t, map[string]string{"first": "", "second": ""})}, + {"info-a-1.tar.zst", writeTarZst(t, map[string]string{"large": strings.Repeat("x", 1024)})}, + } + var buf bytes.Buffer + zw := zip.NewWriter(&buf) + for _, member := range members { + w, err := zw.CreateHeader(&zip.FileHeader{Name: member.name, Method: zip.Store}) + if err != nil { + t.Fatal(err) + } + if _, err := w.Write(member.data); err != nil { + t.Fatal(err) + } + } + if err := zw.Close(); err != nil { + t.Fatal(err) + } + + _, err := OpenBytes("a.conda", buf.Bytes()) + if !errors.Is(err, ErrEntryLimit) { + t.Fatalf("expected ErrEntryLimit before reading the next member, got: %v", err) + } + if !strings.Contains(err.Error(), "count 3 exceeds 2") { + t.Fatalf("expected combined entry count in error, got: %v", err) + } +} + func TestOpenDoesNotInferConda(t *testing.T) { reader, err := OpenBytes("artifact", createTestConda(t)) if err != nil { diff --git a/gem.go b/gem.go index b7706c5..c5edcc8 100644 --- a/gem.go +++ b/gem.go @@ -17,6 +17,7 @@ type gemReader struct { func openGem(raw []byte) (*gemReader, error) { tr := tar.NewReader(bytes.NewReader(raw)) + entryCount := 0 // Find data.tar.gz in the gem for { @@ -27,6 +28,10 @@ func openGem(raw []byte) (*gemReader, error) { if err != nil { return nil, fmt.Errorf("reading gem tar: %w", err) } + entryCount++ + if err := checkArchiveEntryCount(entryCount); err != nil { + return nil, err + } // Look for data.tar.gz if header.Name == "data.tar.gz" { diff --git a/tar.go b/tar.go index e33a710..742aba2 100644 --- a/tar.go +++ b/tar.go @@ -15,9 +15,13 @@ import ( "github.com/ulikunitz/xz" ) -var maxDecompressedSize int64 = 512 << 20 // 512 MiB +var ( + maxDecompressedSize int64 = 512 << 20 // 512 MiB + maxArchiveEntries = 100_000 +) var ErrDecompressLimit = errors.New("decompressed content exceeds size limit") +var ErrEntryLimit = errors.New("archive entry count exceeds limit") type tarReader struct { raw []byte @@ -31,6 +35,10 @@ type tarFileEntry struct { } func openTar(raw []byte, compression string) (*tarReader, error) { + return openTarWithInitialEntryCount(raw, compression, 0) +} + +func openTarWithInitialEntryCount(raw []byte, compression string, initialEntryCount int) (*tarReader, error) { content := bytes.NewReader(raw) r := io.Reader(content) @@ -71,6 +79,9 @@ func openTar(raw []byte, compression string) (*tarReader, error) { if err != nil { return nil, fmt.Errorf("reading tar: %w", err) } + if err := checkArchiveEntryCount(initialEntryCount + len(files) + 1); err != nil { + return nil, err + } // FileInfo().Mode() combines header.Mode permission bits with type // bits derived from Typeflag. It reports hard links as regular @@ -119,6 +130,13 @@ func openTar(raw []byte, compression string) (*tarReader, error) { return &tarReader{raw: raw, files: files, index: index}, nil } +func checkArchiveEntryCount(count int) error { + if count > maxArchiveEntries { + return fmt.Errorf("%w: count %d exceeds %d", ErrEntryLimit, count, maxArchiveEntries) + } + return nil +} + func (t *tarReader) List() ([]FileInfo, error) { files := make([]FileInfo, len(t.files)) for i, f := range t.files { diff --git a/zip.go b/zip.go index 025dd88..838fff9 100644 --- a/zip.go +++ b/zip.go @@ -3,6 +3,7 @@ package archives import ( "archive/zip" "bytes" + "encoding/binary" "fmt" "io" "strings" @@ -15,10 +16,22 @@ import ( // then returns a synthesised 0666/0444 that callers should not treat as a // stored mode. const ( - zipUnixModeShift = 16 - zipCreatorShift = 8 - zipCreatorUnix = 3 - zipCreatorMacOSX = 19 + zipUnixModeShift = 16 + zipCreatorShift = 8 + zipCreatorUnix = 3 + zipCreatorMacOSX = 19 + zipSignatureLen = 4 + zipDirectoryHeaderLen = 46 + zipDirectoryEndLen = 22 + zipDirectory64EndLen = 56 + zipDirectory64LocatorLen = 20 + zipDirectoryHeaderSignature = 0x02014b50 + zipDirectoryEndSignature = 0x06054b50 + zipDirectory64EndSignature = 0x06064b50 + zipDirectory64LocSignature = 0x07064b50 + zipDirectoryEndSearchWindow = 65 * 1024 + zipDirectoryRecords64Marker = 0xffff + zipDirectorySizeOffsetMarker = 0xffffffff ) type zipReader struct { @@ -28,10 +41,16 @@ type zipReader struct { } func openZip(raw []byte) (*zipReader, error) { + if err := checkZipEntryCount(raw); err != nil { + return nil, err + } reader, err := zip.NewReader(bytes.NewReader(raw), int64(len(raw))) if err != nil { return nil, fmt.Errorf("opening zip: %w", err) } + if err := checkArchiveEntryCount(len(reader.File)); err != nil { + return nil, err + } index := make(map[string]*zip.File, len(reader.File)) for _, f := range reader.File { @@ -47,6 +66,88 @@ func openZip(raw []byte) (*zipReader, error) { }, nil } +func checkZipEntryCount(raw []byte) error { + start, end, ok := zipCentralDirectoryBounds(raw) + if !ok { + return nil + } + + count := 0 + for offset := start; offset+zipDirectoryHeaderLen <= end; { + if binary.LittleEndian.Uint32(raw[offset:]) != zipDirectoryHeaderSignature { + break + } + nameLen := int(binary.LittleEndian.Uint16(raw[offset+28:])) + extraLen := int(binary.LittleEndian.Uint16(raw[offset+30:])) + commentLen := int(binary.LittleEndian.Uint16(raw[offset+32:])) + recordLen := zipDirectoryHeaderLen + nameLen + extraLen + commentLen + if recordLen > end-offset { + return nil + } + + count++ + if err := checkArchiveEntryCount(count); err != nil { + return err + } + offset += recordLen + } + + return nil +} + +func zipCentralDirectoryBounds(raw []byte) (int, int, bool) { + searchStart := len(raw) - zipDirectoryEndSearchWindow + if searchStart < 0 { + searchStart = 0 + } + + for offset := len(raw) - zipDirectoryEndLen; offset >= searchStart; offset-- { + if binary.LittleEndian.Uint32(raw[offset:]) != zipDirectoryEndSignature { + continue + } + commentLen := int(binary.LittleEndian.Uint16(raw[offset+20:])) + if offset+zipDirectoryEndLen+commentLen > len(raw) { + continue + } + + directoryEnd := offset + directoryRecords := binary.LittleEndian.Uint16(raw[offset+10:]) + directorySize := uint64(binary.LittleEndian.Uint32(raw[offset+12:])) + directoryOffset := uint64(binary.LittleEndian.Uint32(raw[offset+16:])) + if directoryRecords == zipDirectoryRecords64Marker || + directorySize == zipDirectorySizeOffsetMarker || + directoryOffset == zipDirectorySizeOffsetMarker { + locatorOffset := offset - zipDirectory64LocatorLen + if locatorOffset < 0 || binary.LittleEndian.Uint32(raw[locatorOffset:]) != zipDirectory64LocSignature { + return 0, 0, false + } + zip64Offset := binary.LittleEndian.Uint64(raw[locatorOffset+8:]) + if len(raw) < zipDirectory64EndLen || zip64Offset > uint64(len(raw)-zipDirectory64EndLen) { + return 0, 0, false + } + directoryEnd = int(zip64Offset) + if binary.LittleEndian.Uint32(raw[directoryEnd:]) != zipDirectory64EndSignature { + return 0, 0, false + } + directorySize = binary.LittleEndian.Uint64(raw[directoryEnd+40:]) + directoryOffset = binary.LittleEndian.Uint64(raw[directoryEnd+48:]) + } + + if directorySize > uint64(directoryEnd) { + return 0, 0, false + } + start := directoryEnd - int(directorySize) + if directoryOffset < uint64(directoryEnd)-directorySize && + directoryOffset <= uint64(len(raw)-zipSignatureLen) && + binary.LittleEndian.Uint32(raw[int(directoryOffset):]) == zipDirectoryHeaderSignature { + start = int(directoryOffset) + } + return start, directoryEnd, true + } + + return 0, 0, false +} + func (z *zipReader) List() ([]FileInfo, error) { files := make([]FileInfo, 0, len(z.reader.File))