From c01110ef06f1020c233669acb82ffcf7915e60b2 Mon Sep 17 00:00:00 2001 From: Andrew Nesbitt Date: Sun, 27 Sep 2026 06:45:40 +0100 Subject: [PATCH] Add streaming archive readers with configurable limits --- .github/workflows/ci.yml | 14 +- README.md | 62 +++++ archives.go | 20 +- conda.go | 2 +- gem.go | 2 +- stream.go | 236 +++++++++++++++++++ stream_tar.go | 103 +++++++++ stream_test.go | 481 +++++++++++++++++++++++++++++++++++++++ stream_zip.go | 134 +++++++++++ tar.go | 90 ++++---- zip.go | 8 +- 11 files changed, 1095 insertions(+), 57 deletions(-) create mode 100644 stream.go create mode 100644 stream_tar.go create mode 100644 stream_test.go create mode 100644 stream_zip.go diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b95295b..2566fb8 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -11,6 +11,8 @@ permissions: {} jobs: test: runs-on: ubuntu-latest + permissions: + contents: read steps: - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: @@ -68,14 +70,22 @@ jobs: tar -xJf wasmtime.tar.xz echo "$RUNNER_TEMP/wasmtime-v44.0.1-x86_64-linux" >> "$GITHUB_PATH" + - name: Test Go WebAssembly streaming + run: GOOS=js GOARCH=wasm go test -exec="$(go env GOROOT)/lib/wasm/go_js_wasm_exec" -run '^TestStream' . + - name: Test TinyGo WebAssembly archive reading - run: tinygo test -target=wasm -run '^TestExtractAllUnsupported$' -v . + run: tinygo test -target=wasm -run '^Test(ExtractAllUnsupported|Stream)' -v . - name: Test TinyGo WASI archive reading - run: tinygo test -target=wasip1 -run '^TestExtractAllUnsupported$' -v . + run: tinygo test -target=wasip1 -run '^Test(ExtractAllUnsupported|Stream)' -v . + + - name: Test Go WASI streaming + run: GOOS=wasip1 GOARCH=wasm go test -exec=wasmtime -run '^TestStream' . lint: runs-on: ubuntu-latest + permissions: + contents: read steps: - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: diff --git a/README.md b/README.md index 3514988..6b84d51 100644 --- a/README.md +++ b/README.md @@ -72,6 +72,68 @@ Compressed content is opened as TAR and returns a parser error when it does not contain a TAR archive. `Open` reads at most 512 bytes before rejecting an unsupported stream with no recognised extension. +### Sequential reading + +`OpenStream` reads entries one at a time through `Next` and `Read`, without +retaining expanded file bodies. It supports every format listed below, +including the inner files of gem and conda packages. + +```go +stream, err := archives.OpenStream("package.tgz", f, archives.StreamOptions{ + MaxInputBytes: 64 << 20, + MaxEntryBytes: 8 << 20, + MaxExpandedBytes: 64 << 20, + MaxEntries: 2000, +}) +if err != nil { + return err +} +defer stream.Close() + +for { + entry, err := stream.Next() + if err == io.EOF { + break + } + if err != nil { + return err + } + fmt.Println(entry.Path, entry.Size) + if _, err := io.Copy(io.Discard, stream); err != nil { + return err + } +} +``` + +Import `io` for this example. `Next` discards unread entry data, and skipped +entries still count towards the limits. Entries remain in archive order, +including duplicates. Each entry's `Read` ends at `io.EOF`; iteration errors +are terminal. The caller owns the input, and `Close` releases decoder resources +without closing or draining it. A TAR end marker ends iteration, so trailing +data and compression trailers beyond that marker may remain unchecked. + +ZIP and conda require random access to their compressed input. `OpenStream` +buffers that input within `MaxInputBytes`; `OpenStreamBytes(name, data, options)` +reuses an existing byte slice without copying it. Keep that slice unchanged +until `Close`. TAR and gem can consume a forward-only input directly. + +Zero limits default to 512 MiB for each byte limit and 100,000 entries. +Negative limits are rejected. Entry and expanded-byte limits use logical +header sizes, including sparse files, before exposing a body. Container members +in gem and conda have a separate entry-count budget, with their combined bodies +bounded by `MaxInputBytes`. Decoder workspace and metadata are additional memory; +these limits do not cap total heap usage. For forward-only input, the input +limit covers bytes consumed, with at most one extra byte read to detect overflow. + +For `swhid-go`, regular files can go directly to +`objects.ComputeContentHashReader(stream, entry.Size)`. This preserves its +collision-detecting hash without buffering a file. `StreamEntry` includes the +mode bits, `Linkname`, and `IsHardlink`: TAR symlink targets are in `Linkname`, +while ZIP symlink targets are in the body. A directory-hashing consumer must +retain paths, modes, and content hashes, resolve hard links, and apply its own +duplicate-path policy. The stream cannot rewind to hash the whole artifact; +hash the original byte slice or use a tee on the input and consume it fully. + ### Prefix stripping Some package formats wrap content in a directory (npm uses `package/`). `OpenWithPrefix` strips a path prefix from all entries: diff --git a/archives.go b/archives.go index 00f4550..610e753 100644 --- a/archives.go +++ b/archives.go @@ -36,6 +36,10 @@ const ( formatGem = "gem" formatConda = "conda" contentSniffSize = 512 + compressionGzip = "gzip" + compressionBzip2 = "bzip2" + compressionXZ = "xz" + compressionZstd = "zstd" ) // FileInfo represents metadata about a file in an archive. @@ -128,13 +132,13 @@ func openRaw(format string, raw []byte) (Reader, error) { case formatTAR: return openTar(raw, "") case formatTarGzip, formatTGZ: - return openTar(raw, "gzip") + return openTar(raw, compressionGzip) case formatTarBzip2: - return openTar(raw, "bzip2") + return openTar(raw, compressionBzip2) case formatTarXZ: - return openTar(raw, "xz") + return openTar(raw, compressionXZ) case formatTarZstd: - return openTar(raw, "zstd") + return openTar(raw, compressionZstd) case formatGem: return openGem(raw) case formatConda: @@ -158,13 +162,13 @@ func archiveFormat(detected string) string { return formatZIP case "tar": return formatTAR - case "gzip": + case compressionGzip: return formatTarGzip - case "bzip2": + case compressionBzip2: return formatTarBzip2 - case "xz": + case compressionXZ: return formatTarXZ - case "zstd": + case compressionZstd: return formatTarZstd default: return "" diff --git a/conda.go b/conda.go index 00da255..328e3b1 100644 --- a/conda.go +++ b/conda.go @@ -68,7 +68,7 @@ func readCondaMember(raw []byte, f *zip.File, initialEntryCount int) ([]tarFileE return nil, 0, err } - tr, err := openTarWithInitialEntryCount(data, "zstd", initialEntryCount) + tr, err := openTarWithInitialEntryCount(data, compressionZstd, initialEntryCount) if err != nil { return nil, 0, fmt.Errorf("opening %s: %w", f.Name, err) } diff --git a/gem.go b/gem.go index c5edcc8..f57614d 100644 --- a/gem.go +++ b/gem.go @@ -43,7 +43,7 @@ func openGem(raw []byte) (*gemReader, error) { return nil, fmt.Errorf("%w: data.tar.gz exceeds %d bytes", ErrDecompressLimit, maxDecompressedSize) } - dataReader, err := openTar(dataContent, "gzip") + dataReader, err := openTar(dataContent, compressionGzip) if err != nil { return nil, fmt.Errorf("opening data.tar.gz: %w", err) } diff --git a/stream.go b/stream.go new file mode 100644 index 0000000..0afee00 --- /dev/null +++ b/stream.go @@ -0,0 +1,236 @@ +package archives + +import ( + "bufio" + "bytes" + "errors" + "fmt" + "io" + "io/fs" +) + +var ErrEntrySizeLimit = errors.New("archive entry exceeds size limit") +var ErrInputLimit = errors.New("archive input exceeds size limit") + +// StreamOptions limits entry bodies, including skipped entries. Zero values +// use 512 MiB for byte limits and 100,000 entries; negative values are invalid. +// Decoder buffers and archive metadata require additional memory. +type StreamOptions struct { + MaxInputBytes int64 + MaxEntryBytes int64 + MaxExpandedBytes int64 + MaxEntries int +} + +// StreamEntry includes link metadata needed to hash TAR links. ZIP symlink +// targets are stored in the entry body. Hard links refer to another TAR path. +type StreamEntry struct { + FileInfo + Linkname string + IsHardlink bool +} + +// Stream reads entries in archive order without retaining expanded bodies. +// Paths and duplicate entries are preserved. A Stream is not safe for +// concurrent use. It does not implement the random-access Reader interface. +type Stream struct { + source streamSource + options StreamOptions + expanded int64 + entries int + err error + closed bool +} + +type streamSource interface { + io.ReadCloser + Next() (*StreamEntry, error) +} + +// OpenStream reads supported archives sequentially. ZIP and conda input is +// buffered up to MaxInputBytes for random access; TAR and gem input is streamed. +// The caller owns content and must close it when needed. +func OpenStream(filename string, content io.Reader, options StreamOptions) (*Stream, error) { + options, err := options.defaults() + if err != nil { + return nil, err + } + return openStream(filename, &streamInput{reader: content, remaining: options.MaxInputBytes}, nil, options) +} + +// OpenStreamBytes reuses content without copying it, including ZIP and conda +// input. The caller must not modify the slice until the stream is closed. +func OpenStreamBytes(filename string, content []byte, options StreamOptions) (*Stream, error) { + options, err := options.defaults() + if err != nil { + return nil, err + } + if int64(len(content)) > options.MaxInputBytes { + return nil, ErrInputLimit + } + return openStream(filename, bytes.NewReader(content), content, options) +} + +func openStream(filename string, content io.Reader, raw []byte, options StreamOptions) (*Stream, error) { + format := detectFormat(filename) + if format == "" { + buffered := bufio.NewReaderSize(content, contentSniffSize) + prefix, err := buffered.Peek(contentSniffSize) + if err != nil && err != io.EOF { + return nil, fmt.Errorf("reading archive content: %w", err) + } + format = detectContentPrefixFormat(prefix) + content = buffered + } + var source streamSource + var err error + switch format { + case formatZIP, formatConda: + if raw == nil { + raw, err = io.ReadAll(content) + if err != nil { + return nil, err + } + } + source, err = newZipStream(raw, options, format == formatConda) + case formatGem: + source = &gemStream{outer: newContainerTar(content, options)} + default: + compression, ok := streamCompression(format) + if !ok { + return nil, fmt.Errorf("%w: archive format: %s", errors.ErrUnsupported, filename) + } + source, err = newTarStream(content, compression) + } + if err != nil { + return nil, err + } + return &Stream{source: source, options: options}, nil +} + +func (o StreamOptions) defaults() (StreamOptions, error) { + if o.MaxInputBytes < 0 || o.MaxEntryBytes < 0 || o.MaxExpandedBytes < 0 || o.MaxEntries < 0 { + return o, errors.New("stream limits must not be negative") + } + if o.MaxInputBytes == 0 { + o.MaxInputBytes = maxDecompressedSize + } + if o.MaxEntryBytes == 0 { + o.MaxEntryBytes = maxDecompressedSize + } + if o.MaxExpandedBytes == 0 { + o.MaxExpandedBytes = maxDecompressedSize + } + if o.MaxEntries == 0 { + o.MaxEntries = maxArchiveEntries + } + return o, nil +} + +func streamCompression(format string) (string, bool) { + switch format { + case formatTAR: + return "", true + case formatTarGzip, formatTGZ: + return compressionGzip, true + case formatTarBzip2: + return compressionBzip2, true + case formatTarXZ: + return compressionXZ, true + case formatTarZstd: + return compressionZstd, true + default: + return "", false + } +} + +// Next discards the unread body and returns the next entry, or io.EOF. +// Header sizes, including sparse logical sizes, count against limits before +// the body is exposed. MaxEntries counts visible entries; nested containers +// have a separate MaxEntries budget. Hidden TAR extension headers are excluded. +// Any iteration error is terminal. +func (s *Stream) Next() (*StreamEntry, error) { + if s.closed { + return nil, fs.ErrClosed + } + if s.err != nil { + return nil, s.err + } + entry, err := s.source.Next() + if err == nil { + err = s.checkEntry(entry) + } + if err != nil { + s.err = err + return nil, err + } + s.entries++ + s.expanded += entry.Size + return entry, nil +} + +func (s *Stream) checkEntry(entry *StreamEntry) error { + if s.entries >= s.options.MaxEntries { + return fmt.Errorf("%w: exceeds %d", ErrEntryLimit, s.options.MaxEntries) + } + if entry.Size < 0 || entry.Size > s.options.MaxEntryBytes { + return fmt.Errorf("%w: %s exceeds %d bytes", ErrEntrySizeLimit, entry.Path, s.options.MaxEntryBytes) + } + if entry.Size > s.options.MaxExpandedBytes-s.expanded { + return fmt.Errorf("%w: exceeds %d bytes", ErrDecompressLimit, s.options.MaxExpandedBytes) + } + return nil +} + +// Read reads the current entry and returns io.EOF at its end or before Next. +// Non-EOF errors prevent further iteration. +func (s *Stream) Read(p []byte) (int, error) { + if s.closed { + return 0, fs.ErrClosed + } + if s.err != nil { + return 0, s.err + } + n, err := s.source.Read(p) + if err != nil && err != io.EOF { + s.err = err + } + return n, err +} + +// Close releases resources without draining or closing the input. Unread +// content and compressed trailers after a TAR end marker are not validated. +func (s *Stream) Close() error { + if s.closed { + return nil + } + s.closed = true + err := s.source.Close() + s.source = nil + return err +} + +type streamInput struct { + reader io.Reader + remaining int64 +} + +func (r *streamInput) Read(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + if r.remaining == 0 { + var probe [1]byte + n, err := r.reader.Read(probe[:]) + if n != 0 { + return 0, ErrInputLimit + } + return 0, err + } + if int64(len(p)) > r.remaining { + p = p[:r.remaining] + } + n, err := r.reader.Read(p) + r.remaining -= int64(n) + return n, err +} diff --git a/stream_tar.go b/stream_tar.go new file mode 100644 index 0000000..0f91b06 --- /dev/null +++ b/stream_tar.go @@ -0,0 +1,103 @@ +package archives + +import ( + "archive/tar" + "fmt" + "io" +) + +type tarStream struct { + reader *tar.Reader + closer io.Closer +} + +func newTarStream(content io.Reader, compression string) (*tarStream, error) { + r, closer, err := tarContentReader(content, compression) + if err != nil { + return nil, err + } + return &tarStream{reader: tar.NewReader(r), closer: closer}, nil +} + +func (t *tarStream) Next() (*StreamEntry, error) { + h, err := t.reader.Next() + if err != nil { + return nil, err + } + return &StreamEntry{ + FileInfo: fileInfoFromTar(h), Linkname: h.Linkname, + IsHardlink: h.Typeflag == tar.TypeLink, + }, nil +} + +func (t *tarStream) Read(p []byte) (int, error) { + return t.reader.Read(p) +} + +func (t *tarStream) Close() error { + t.reader = nil + if t.closer != nil { + return t.closer.Close() + } + return nil +} + +func containerOptions(options StreamOptions) StreamOptions { + options.MaxEntryBytes = options.MaxInputBytes + options.MaxExpandedBytes = options.MaxInputBytes + return options +} + +func newContainerTar(content io.Reader, options StreamOptions) *Stream { + return &Stream{ + source: &tarStream{reader: tar.NewReader(content)}, + options: containerOptions(options), + } +} + +type gemStream struct { + outer *Stream + inner *tarStream +} + +func (g *gemStream) Next() (*StreamEntry, error) { + if g.inner == nil { + for { + entry, err := g.outer.Next() + if err == io.EOF { + return nil, fmt.Errorf("data.tar.gz not found in gem") + } + if err != nil { + return nil, err + } + if entry.Path != "data.tar.gz" { + continue + } + g.inner, err = newTarStream(g.outer, compressionGzip) + if err != nil { + return nil, fmt.Errorf("opening data.tar.gz: %w", err) + } + break + } + } + return g.inner.Next() +} + +func (g *gemStream) Read(p []byte) (int, error) { + if g.inner == nil { + return 0, io.EOF + } + return g.inner.Read(p) +} + +func (g *gemStream) Close() error { + var err error + if g.inner != nil { + err = g.inner.Close() + } + outerErr := g.outer.Close() + if err != nil { + return err + } + return outerErr +} diff --git a/stream_test.go b/stream_test.go new file mode 100644 index 0000000..b6b68e0 --- /dev/null +++ b/stream_test.go @@ -0,0 +1,481 @@ +package archives + +import ( + "archive/tar" + "archive/zip" + "bytes" + "errors" + "fmt" + "io" + "io/fs" + "reflect" + "testing" +) + +func TestStreamFormats(t *testing.T) { + for _, tc := range []struct { + name string + data []byte + }{ + {"source.zip", createTestZip()}, + {"source.tar", createTestTar()}, + {"source.tgz", createTestTarGz()}, + {"source.tar.bz2", createTestTarBz2(t)}, + {"source.tar.xz", createTestTarXz(t)}, + {"source.tar.zst", createTestTarZst(t)}, + {"source.gem", createTestGem()}, + {"source.conda", createTestConda(t)}, + {"alpine.apk", createTestTarGz()}, + {"android.apk", createTestZip()}, + {"artifact", createTestTarGz()}, + } { + t.Run(tc.name, func(t *testing.T) { + r, err := OpenBytes(tc.name, tc.data) + if err != nil { + t.Fatal(err) + } + defer func() { _ = r.Close() }() + want, err := r.List() + if err != nil { + t.Fatal(err) + } + for _, buffered := range []bool{false, true} { + var stream *Stream + if buffered { + stream, err = OpenStream(tc.name, bytes.NewBuffer(tc.data), StreamOptions{}) + } else { + stream, err = OpenStreamBytes(tc.name, tc.data, StreamOptions{}) + } + if err != nil { + t.Fatal(err) + } + assertStreamMatches(t, stream, r, want) + } + }) + } +} + +func assertStreamMatches(t *testing.T, stream *Stream, r Reader, want []FileInfo) { + t.Helper() + defer func() { _ = stream.Close() }() + for _, info := range want { + entry, err := stream.Next() + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(entry.FileInfo, info) { + t.Fatalf("metadata = %+v, want %+v", entry.FileInfo, info) + } + got, err := io.ReadAll(stream) + if err != nil { + t.Fatal(err) + } + if info.IsDir { + if len(got) != 0 { + t.Fatal("directory has content") + } + continue + } + body, err := r.Extract(info.Path) + if err != nil { + t.Fatal(err) + } + data, err := io.ReadAll(body) + _ = body.Close() + if err != nil || !bytes.Equal(got, data) { + t.Fatalf("%s: content differs, %v", info.Path, err) + } + } + for range 2 { + if _, err := stream.Next(); err != io.EOF { + t.Fatalf("end: %v", err) + } + } +} + +func TestStreamLimits(t *testing.T) { + raw := tarWithPayloads(t, []byte("aaa"), []byte("bbbb"), nil) + for _, name := range []string{"test.tar", "test.tar.gz", "test.tar.xz", "test.tar.zst", "test.zip", "test.gem", "test.conda"} { + data := streamTestArchive(t, name, raw) + for _, tc := range []struct { + name string + opts StreamOptions + want error + }{ + {"entry", StreamOptions{MaxEntryBytes: 3}, ErrEntrySizeLimit}, + {"total", StreamOptions{MaxExpandedBytes: 6}, ErrDecompressLimit}, + {"count", StreamOptions{MaxEntries: 2}, ErrEntryLimit}, + {"exact", StreamOptions{MaxExpandedBytes: 7, MaxEntryBytes: 4, MaxEntries: 3}, nil}, + {"input", StreamOptions{MaxInputBytes: int64(len(data) - 1)}, ErrInputLimit}, + } { + t.Run(name+"/"+tc.name, func(t *testing.T) { + for _, readBodies := range []bool{false, true} { + s, err := OpenStreamBytes(name, data, tc.opts) + if err == nil { + err = consumeStream(s, readBodies) + _ = s.Close() + } + if !errors.Is(err, tc.want) { + t.Fatalf("read=%v: got %v, want %v", readBodies, err, tc.want) + } + } + }) + } + } +} + +func consumeStream(s *Stream, read bool) error { + for { + _, err := s.Next() + if err == io.EOF { + return nil + } + if err != nil { + return err + } + if read { + if _, err := io.Copy(io.Discard, s); err != nil { + return err + } + } + } +} + +func streamTestArchive(t testing.TB, name string, raw []byte) []byte { + t.Helper() + switch name { + case "test.gem": + var buf bytes.Buffer + tw := tar.NewWriter(&buf) + data := gzipTar(t, raw) + if err := tw.WriteHeader(&tar.Header{Name: "data.tar.gz", Size: int64(len(data)), Mode: 0o644}); err != nil { + t.Fatal(err) + } + if _, err := tw.Write(data); err != nil { + t.Fatal(err) + } + if err := tw.Close(); err != nil { + t.Fatal(err) + } + return buf.Bytes() + case "test.zip", "test.conda": + var buf bytes.Buffer + zw := zip.NewWriter(&buf) + if name == "test.conda" { + w, err := zw.Create("pkg-test.tar.zst") + if err != nil { + t.Fatal(err) + } + if _, err := w.Write(compressTar(t, "test.tar.zst", raw)); err != nil { + t.Fatal(err) + } + } else { + streamTestZipEntries(t, zw, raw) + } + if err := zw.Close(); err != nil { + t.Fatal(err) + } + return buf.Bytes() + default: + return compressTar(t, name, raw) + } +} + +func streamTestZipEntries(t testing.TB, zw *zip.Writer, raw []byte) { + t.Helper() + tr := tar.NewReader(bytes.NewReader(raw)) + for { + h, err := tr.Next() + if err == io.EOF { + return + } + if err != nil { + t.Fatal(err) + } + w, err := zw.Create(h.Name) + if err != nil { + t.Fatal(err) + } + if _, err := io.Copy(w, tr); err != nil { + t.Fatal(err) + } + } +} + +func TestStreamRejectsHeaderBeforeBody(t *testing.T) { + var buf bytes.Buffer + tw := tar.NewWriter(&buf) + if err := tw.WriteHeader(&tar.Header{Name: "huge", Size: 1 << 40, Mode: 0o644, Format: tar.FormatGNU}); err != nil { + t.Fatal(err) + } + input := &countStreamInput{Reader: bytes.NewReader(buf.Bytes())} + s, err := OpenStream("huge.tar", input, StreamOptions{MaxEntryBytes: 8 << 20}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = s.Close() }() + if input.bytes != 0 { + t.Fatal("opened by reading input") + } + for range 2 { + if _, err := s.Next(); !errors.Is(err, ErrEntrySizeLimit) { + t.Fatalf("Next: %v", err) + } + } + if _, err := s.Read(make([]byte, 1)); !errors.Is(err, ErrEntrySizeLimit) { + t.Fatalf("Read after rejection: %v", err) + } + if input.bytes != 512 { + t.Fatalf("read %d input bytes, want header only", input.bytes) + } +} + +type countStreamInput struct { + io.Reader + bytes int + closed bool +} + +func (r *countStreamInput) Read(p []byte) (int, error) { + n, err := r.Reader.Read(p) + r.bytes += n + return n, err +} + +func (r *countStreamInput) Close() error { + r.closed = true + return nil +} + +func TestStreamLifecycle(t *testing.T) { + input := &countStreamInput{Reader: bytes.NewReader(createTestTarGz())} + s, err := OpenStream("source.tgz", input, StreamOptions{}) + if err != nil { + t.Fatal(err) + } + if _, err := s.Read(make([]byte, 1)); err != io.EOF { + t.Fatalf("Read before Next: %v", err) + } + if _, err := s.Next(); err != nil { + t.Fatal(err) + } + for range 2 { + if err := s.Close(); err != nil { + t.Fatal(err) + } + } + if input.closed { + t.Fatal("closed caller's input") + } + if _, err := s.Next(); !errors.Is(err, fs.ErrClosed) { + t.Fatalf("Next after Close: %v", err) + } + if _, err := s.Read(make([]byte, 1)); !errors.Is(err, fs.ErrClosed) { + t.Fatalf("Read after Close: %v", err) + } +} + +func TestStreamLinksAndDuplicatePaths(t *testing.T) { + var buf bytes.Buffer + tw := tar.NewWriter(&buf) + writeTarFile(t, tw, "run", "first", 0o755) + writeTarFile(t, tw, "run", "second", 0o644) + for _, h := range []*tar.Header{ + {Name: "soft", Linkname: "run", Typeflag: tar.TypeSymlink}, + {Name: "hard", Linkname: "run", Typeflag: tar.TypeLink}, + } { + if err := tw.WriteHeader(h); err != nil { + t.Fatal(err) + } + } + if err := tw.Close(); err != nil { + t.Fatal(err) + } + s, err := OpenStreamBytes("links.tar", buf.Bytes(), StreamOptions{}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = s.Close() }() + for i, want := range []string{"first", "second", "", ""} { + e, err := s.Next() + if err != nil { + t.Fatal(err) + } + body, err := io.ReadAll(s) + if err != nil || string(body) != want { + t.Fatalf("body %d: %q, %v", i, body, err) + } + if i == 0 && e.Mode&0o111 == 0 { + t.Fatal("executable mode lost") + } + if i >= 2 && (e.Linkname != "run" || e.IsHardlink != (i == 3)) { + t.Fatalf("link metadata: %+v", e) + } + } +} + +func TestStreamTruncatedSkippedBody(t *testing.T) { + raw := tarWithPayloads(t, make([]byte, 1024))[:1023] + for _, name := range []string{"test.tar", "test.tar.gz", "test.gem", "test.conda"} { + s, err := OpenStreamBytes(name, streamTestArchive(t, name, raw), StreamOptions{}) + if err != nil { + t.Fatal(err) + } + if _, err := s.Next(); err != nil { + t.Fatal(err) + } + if _, err := s.Next(); !errors.Is(err, io.ErrUnexpectedEOF) { + t.Fatalf("%s: %v", name, err) + } + _ = s.Close() + } +} + +func TestStreamInputBounds(t *testing.T) { + data := createTestZip() + for _, delta := range []int{-1, 0, 1} { + input := &countStreamInput{Reader: bytes.NewReader(data)} + s, err := OpenStream("test.zip", input, StreamOptions{MaxInputBytes: int64(len(data) + delta)}) + switch { + case delta < 0: + if !errors.Is(err, ErrInputLimit) { + t.Fatalf("input over limit: %v", err) + } + case err != nil: + t.Fatal(err) + default: + _ = s.Close() + } + if input.bytes > len(data)+delta+1 { + t.Fatal("read beyond input limit probe") + } + } + for _, options := range []StreamOptions{{MaxInputBytes: -1}, {MaxEntryBytes: -1}, {MaxExpandedBytes: -1}, {MaxEntries: -1}} { + if _, err := OpenStreamBytes("test.zip", data, options); err == nil { + t.Fatalf("accepted negative limit: %+v", options) + } + } +} + +func BenchmarkStreamTar(b *testing.B) { + for _, count := range []int{1, 128} { + payloads := make([][]byte, count) + for i := range payloads { + payloads[i] = make([]byte, 64<<10) + } + data := gzipTar(b, tarWithPayloads(b, payloads...)) + b.Run(fmt.Sprint(count), func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + s, err := OpenStreamBytes("source.tgz", data, StreamOptions{}) + if err != nil { + b.Fatal(err) + } + if err := consumeStream(s, true); err != nil { + b.Fatal(err) + } + _ = s.Close() + } + }) + } +} + +func TestStreamPartialRead(t *testing.T) { + raw := tarWithPayloads(t, []byte("first"), []byte("second")) + for _, name := range []string{"test.tar", "test.tar.gz", "test.zip", "test.gem", "test.conda"} { + s, err := OpenStreamBytes(name, streamTestArchive(t, name, raw), StreamOptions{}) + if err != nil { + t.Fatal(err) + } + if _, err := s.Next(); err != nil { + t.Fatal(err) + } + var prefix [2]byte + if _, err := io.ReadFull(s, prefix[:]); err != nil || string(prefix[:]) != "fi" { + t.Fatalf("partial read: %q, %v", prefix, err) + } + if _, err := s.Next(); err != nil { + t.Fatal(err) + } + body, err := io.ReadAll(s) + if err != nil || string(body) != "second" { + t.Fatalf("%s second entry: %q, %v", name, body, err) + } + _ = s.Close() + } +} + +func TestStreamZipChecksumOnSkip(t *testing.T) { + var buf bytes.Buffer + zw := zip.NewWriter(&buf) + w, err := zw.CreateHeader(&zip.FileHeader{Name: "file", Method: zip.Store}) + if err != nil { + t.Fatal(err) + } + if _, err := io.WriteString(w, "payload"); err != nil { + t.Fatal(err) + } + if err := zw.Close(); err != nil { + t.Fatal(err) + } + data := buf.Bytes() + i := bytes.Index(data, []byte("payload")) + if i < 0 { + t.Fatal("missing fixture payload") + } + data[i] ^= 1 + s, err := OpenStreamBytes("bad.zip", data, StreamOptions{}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = s.Close() }() + if _, err := s.Next(); err != nil { + t.Fatal(err) + } + for range 2 { + if _, err := s.Next(); !errors.Is(err, zip.ErrChecksum) { + t.Fatalf("skipped corrupt entry: %v", err) + } + } +} + +func TestStreamCondaCombinedLimits(t *testing.T) { + data := createTestConda(t) + for _, opts := range []StreamOptions{{MaxEntries: 4}, {MaxExpandedBytes: 30}} { + s, err := OpenStreamBytes("test.conda", data, opts) + if err != nil { + t.Fatal(err) + } + err = consumeStream(s, false) + _ = s.Close() + want := ErrDecompressLimit + if opts.MaxEntries != 0 { + want = ErrEntryLimit + } + if !errors.Is(err, want) { + t.Fatalf("combined member budget: %v, want %v", err, want) + } + } +} + +func TestStreamUnsupportedAndMissingMembers(t *testing.T) { + for _, tc := range []struct { + name string + data []byte + }{ + {"unknown", []byte("no archive")}, + {"bad.zip", []byte("no zip")}, + {"bad.tgz", []byte("no gzip")}, + {"bad.gem", createTestTar()}, + {"bad.conda", createTestZip()}, + } { + s, err := OpenStreamBytes(tc.name, tc.data, StreamOptions{}) + if err == nil { + err = consumeStream(s, false) + _ = s.Close() + } + if err == nil { + t.Fatalf("accepted %s", tc.name) + } + } +} diff --git a/stream_zip.go b/stream_zip.go new file mode 100644 index 0000000..5c01504 --- /dev/null +++ b/stream_zip.go @@ -0,0 +1,134 @@ +package archives + +import ( + "archive/zip" + "bytes" + "fmt" + "io" + "strings" +) + +type zipStream struct { + reader *zip.Reader + index int + file *zip.File + current io.ReadCloser +} + +//nolint:ireturn // archive formats share a sequential source +func newZipStream(raw []byte, options StreamOptions, conda bool) (streamSource, error) { + if err := checkZipEntryCountLimit(raw, options.MaxEntries); 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 len(reader.File) > options.MaxEntries { + return nil, ErrEntryLimit + } + reader.RegisterDecompressor(zip.Deflate, newFlateReader) + z := &zipStream{reader: reader} + if conda { + return &condaStream{outer: &Stream{source: z, options: containerOptions(options)}}, nil + } + return z, nil +} + +func (z *zipStream) Next() (*StreamEntry, error) { + if z.file != nil { + if _, err := io.Copy(io.Discard, z); err != nil { + return nil, err + } + if err := z.current.Close(); err != nil { + return nil, err + } + z.current = nil + z.file = nil + } + if z.index == len(z.reader.File) { + return nil, io.EOF + } + z.file = z.reader.File[z.index] + z.index++ + return &StreamEntry{FileInfo: fileInfoFromZip(z.file)}, nil +} + +func (z *zipStream) Read(p []byte) (int, error) { + if z.file == nil { + return 0, io.EOF + } + if z.current == nil { + var err error + z.current, err = z.file.Open() + if err != nil { + return 0, err + } + } + return z.current.Read(p) +} + +func (z *zipStream) Close() error { + z.reader = nil + z.file = nil + if z.current != nil { + return z.current.Close() + } + return nil +} + +type condaStream struct { + outer *Stream + inner *tarStream + found bool +} + +func (c *condaStream) Next() (*StreamEntry, error) { + for { + if c.inner != nil { + entry, err := c.inner.Next() + if err != io.EOF { + return entry, err + } + if err := c.inner.Close(); err != nil { + return nil, err + } + c.inner = nil + } + entry, err := c.outer.Next() + if err == io.EOF && !c.found { + return nil, fmt.Errorf("no pkg-*.tar.zst or info-*.tar.zst member in conda package") + } + if err != nil { + return nil, err + } + if !strings.HasSuffix(entry.Path, ".tar.zst") || + (!strings.HasPrefix(entry.Path, "pkg-") && !strings.HasPrefix(entry.Path, "info-")) { + continue + } + c.found = true + c.inner, err = newTarStream(c.outer, compressionZstd) + if err != nil { + return nil, fmt.Errorf("opening %s: %w", entry.Path, err) + } + } +} + +func (c *condaStream) Read(p []byte) (int, error) { + if c.inner == nil { + return 0, io.EOF + } + return c.inner.Read(p) +} + +func (c *condaStream) Close() error { + var err error + if c.inner != nil { + err = c.inner.Close() + } + outerErr := c.outer.Close() + if err != nil { + return err + } + return outerErr +} diff --git a/tar.go b/tar.go index 94e0598..59bd0bc 100644 --- a/tar.go +++ b/tar.go @@ -41,32 +41,12 @@ func openTar(raw []byte, compression string) (*tarReader, error) { } func openTarWithInitialEntryCount(raw []byte, compression string, initialEntryCount int) (*tarReader, error) { - content := bytes.NewReader(raw) - r := io.Reader(content) - - switch compression { - case "gzip": - gz, err := gzip.NewReader(content) - if err != nil { - return nil, fmt.Errorf("opening gzip: %w", err) - } - defer func() { _ = gz.Close() }() - r = gz - case "bzip2": - r = bzip2.NewReader(content) - case "xz": - xzReader, err := xz.NewReader(content) - if err != nil { - return nil, fmt.Errorf("opening xz: %w", err) - } - r = xzReader - case "zstd": - dec, err := zstd.NewReader(content, zstd.WithDecoderConcurrency(1)) - if err != nil { - return nil, fmt.Errorf("opening zstd: %w", err) - } - defer dec.Close() - r = dec + r, closer, err := tarContentReader(bytes.NewReader(raw), compression) + if err != nil { + return nil, err + } + if closer != nil { + defer func() { _ = closer.Close() }() } tr := tar.NewReader(r) @@ -85,23 +65,7 @@ func openTarWithInitialEntryCount(raw []byte, compression string, initialEntryCo return nil, err } - // FileInfo().Mode() combines header.Mode permission bits with type - // bits derived from Typeflag. It reports hard links as regular - // files, so mark them irregular explicitly since a hard-link entry - // carries no data of its own. - mode := header.FileInfo().Mode() - if header.Typeflag == tar.TypeLink { - mode |= fs.ModeIrregular - } - info := FileInfo{ - Path: header.Name, - Name: extractName(header.Name), - Size: header.Size, - ModTime: header.ModTime, - IsDir: header.Typeflag == tar.TypeDir, - Mode: uint32(mode), - HasMode: true, - } + info := fileInfoFromTar(header) var data []byte if !info.IsDir { @@ -132,6 +96,46 @@ func openTarWithInitialEntryCount(raw []byte, compression string, initialEntryCo return &tarReader{raw: raw, files: files, index: index}, nil } +func fileInfoFromTar(header *tar.Header) FileInfo { + mode := header.FileInfo().Mode() + // Hard links have no body, despite FileInfo reporting a regular file. + if header.Typeflag == tar.TypeLink { + mode |= fs.ModeIrregular + } + return FileInfo{ + Path: header.Name, Name: extractName(header.Name), Size: header.Size, + ModTime: header.ModTime, IsDir: header.Typeflag == tar.TypeDir, + Mode: uint32(mode), HasMode: true, + } +} + +func tarContentReader(content io.Reader, compression string) (io.Reader, io.Closer, error) { + switch compression { + case compressionGzip: + gz, err := gzip.NewReader(content) + if err != nil { + return nil, nil, fmt.Errorf("opening gzip: %w", err) + } + return gz, gz, nil + case compressionBzip2: + return bzip2.NewReader(content), nil, nil + case compressionXZ: + r, err := xz.NewReader(content) + if err != nil { + return nil, nil, fmt.Errorf("opening xz: %w", err) + } + return r, nil, nil + case compressionZstd: + dec, err := zstd.NewReader(content, zstd.WithDecoderConcurrency(1)) + if err != nil { + return nil, nil, fmt.Errorf("opening zstd: %w", err) + } + return dec, dec.IOReadCloser(), nil + default: + return content, nil, nil + } +} + func readTarEntry(r io.Reader, size, remaining int64) ([]byte, error) { limited := io.LimitReader(r, remaining+1) if size <= bytes.MinRead || size > maxInitialEntryBuffer { diff --git a/zip.go b/zip.go index 0c78c85..a4bd766 100644 --- a/zip.go +++ b/zip.go @@ -72,6 +72,10 @@ func openZip(raw []byte) (*zipReader, error) { } func checkZipEntryCount(raw []byte) error { + return checkZipEntryCountLimit(raw, maxArchiveEntries) +} + +func checkZipEntryCountLimit(raw []byte, limit int) error { start, end, ok := zipCentralDirectoryBounds(raw) if !ok { return nil @@ -91,8 +95,8 @@ func checkZipEntryCount(raw []byte) error { } count++ - if err := checkArchiveEntryCount(count); err != nil { - return err + if count > limit { + return fmt.Errorf("%w: count %d exceeds %d", ErrEntryLimit, count, limit) } offset += recordLen }