From 77810a7bc1e5ac3569f4a22df303c9bda0dfa49f Mon Sep 17 00:00:00 2001 From: Nagendra Subramanya Date: Mon, 3 May 2021 21:30:23 -0700 Subject: [PATCH] Enable overriding (de)compressors per file Fixes https://github.com/alexmullins/zip/issues/12 Port of https://github.com/golang/go/commit/46300a058dfb078164f29fa1a86a2dbdad55e503: Implement setting the compression level for a zip archive by registering a per-Writer compressor through Writer.RegisterCompressor. If no compressors are registered, fall back to the ones registered at the package level. Also implements per-Reader decompressors. --- example_test.go | 62 +++++++++++++++++++++++++++++++++++++++++++++++++ reader.go | 32 +++++++++++++++++++++---- writer.go | 30 +++++++++++++++++++----- 3 files changed, 113 insertions(+), 11 deletions(-) diff --git a/example_test.go b/example_test.go index d625b7e..2625e7c 100644 --- a/example_test.go +++ b/example_test.go @@ -6,6 +6,7 @@ package zip_test import ( "bytes" + "compress/flate" "fmt" "io" "log" @@ -111,3 +112,64 @@ func ExampleWriter_Encrypt() { // Output: // Hello World } + +func ExampleWriter_RegisterCompressor() { + // Override the default Deflate compressor with a higher compression + // level. + + // Create a buffer to write our archive to. + buf := new(bytes.Buffer) + + // Create a new zip archive. + w := zip.NewWriter(buf) + + var fw *flate.Writer + + // Register the deflator. + w.RegisterCompressor(zip.Deflate, func(out io.Writer) (io.WriteCloser, error) { + var err error + if fw == nil { + // Creating a flate compressor for every file is + // expensive, create one and reuse it. + fw, err = flate.NewWriter(out, flate.BestCompression) + } else { + fw.Reset(out) + } + return fw, err + }) + + // Proceed to add files to w. +} + +func ExampleReader_RegisterDecompressor() { + // Open a zip archive for reading. + r, err := zip.OpenReader("testdata/readme.zip") + if err != nil { + log.Fatal(err) + } + defer r.Close() + + // Override the default Deflate decompressor. + r.RegisterDecompressor(zip.Deflate, func(in io.Reader) io.ReadCloser { + return flate.NewReader(in) + }) + + // Iterate through the files in the archive, + // printing some of their contents. + for _, f := range r.File { + fmt.Printf("Contents of %s:\n", f.Name) + rc, err := f.Open() + if err != nil { + log.Fatal(err) + } + _, err = io.CopyN(os.Stdout, rc, 68) + if err != nil { + log.Fatal(err) + } + rc.Close() + fmt.Println() + } + // Output: + // Contents of README: + // This is the source code repository for the Go programming language. +} diff --git a/reader.go b/reader.go index a9e3f6b..d76f59c 100644 --- a/reader.go +++ b/reader.go @@ -22,9 +22,10 @@ var ( ) type Reader struct { - r io.ReaderAt - File []*File - Comment string + r io.ReaderAt + File []*File + Comment string + decompressors map[uint16]Decompressor } type ReadCloser struct { @@ -34,6 +35,7 @@ type ReadCloser struct { type File struct { FileHeader + zip *Reader zipr io.ReaderAt zipsize int64 headerOffset int64 @@ -95,7 +97,7 @@ func (z *Reader) init(r io.ReaderAt, size int64) error { // a bad one, and then only report a ErrFormat or UnexpectedEOF if // the file count modulo 65536 is incorrect. for { - f := &File{zipr: r, zipsize: size} + f := &File{zip: z, zipr: r, zipsize: size} err = readDirectoryHeader(f, buf) if err == ErrFormat || err == io.ErrUnexpectedEOF { break @@ -113,6 +115,26 @@ func (z *Reader) init(r io.ReaderAt, size int64) error { return nil } +// RegisterDecompressor registers or overrides a custom decompressor for a +// specific method ID. If a decompressor for a given method is not found, +// Reader will default to looking up the decompressor at the package level. +// +// Must not be called concurrently with Open on any Files in the Reader. +func (z *Reader) RegisterDecompressor(method uint16, dcomp Decompressor) { + if z.decompressors == nil { + z.decompressors = make(map[uint16]Decompressor) + } + z.decompressors[method] = dcomp +} + +func (z *Reader) decompressor(method uint16) Decompressor { + dcomp := z.decompressors[method] + if dcomp == nil { + dcomp = decompressor(method) + } + return dcomp +} + // Close closes the Zip file, rendering it unusable for I/O. func (rc *ReadCloser) Close() error { return rc.f.Close() @@ -151,7 +173,7 @@ func (f *File) Open() (rc io.ReadCloser, err error) { } else { r = rr } - dcomp := decompressor(f.Method) + dcomp := f.zip.decompressor(f.Method) if dcomp == nil { err = ErrAlgorithm return diff --git a/writer.go b/writer.go index 27d125e..52da5ab 100644 --- a/writer.go +++ b/writer.go @@ -14,14 +14,14 @@ import ( ) // TODO(adg): support zip file comments -// TODO(adg): support specifying deflate level // Writer implements a zip file writer. type Writer struct { - cw *countWriter - dir []*header - last *fileWriter - closed bool + cw *countWriter + dir []*header + last *fileWriter + closed bool + compressors map[uint16]Compressor } type header struct { @@ -222,7 +222,7 @@ func (w *Writer) CreateHeader(fh *FileHeader) (io.Writer, error) { crc32: crc32.NewIEEE(), } // Get the compressor before possibly changing Method to 99 due to password - comp := compressor(fh.Method) + comp := w.compressor(fh.Method) if comp == nil { return nil, ErrAlgorithm } @@ -284,6 +284,24 @@ func writeHeader(w io.Writer, h *FileHeader) error { return err } +// RegisterCompressor registers or overrides a custom compressor for a specific +// method ID. If a compressor for a given method is not found, Writer will +// default to looking up the compressor at the package level. +func (w *Writer) RegisterCompressor(method uint16, comp Compressor) { + if w.compressors == nil { + w.compressors = make(map[uint16]Compressor) + } + w.compressors[method] = comp +} + +func (w *Writer) compressor(method uint16) Compressor { + comp := w.compressors[method] + if comp == nil { + comp = compressor(method) + } + return comp +} + type fileWriter struct { *header zipw io.Writer