diff --git a/pkg/compressio/compressio_test.go b/pkg/compressio/compressio_test.go index 3087462ba..7b40a6f1d 100644 --- a/pkg/compressio/compressio_test.go +++ b/pkg/compressio/compressio_test.go @@ -234,7 +234,7 @@ func benchmark(b *testing.B, compression bool, write bool, hash bool, blockSize } } else { opts.NewWriter = func(b *bytes.Buffer) (io.WriteCloser, error) { - return NewSimpleWriter(b, key), nil + return NewSimpleWriter(b, key, blockSize), nil } opts.NewReader = func(b *bytes.Buffer) (io.Reader, error) { return NewSimpleReader(b, key), nil diff --git a/pkg/compressio/nocompressio.go b/pkg/compressio/nocompressio.go index 4647ab766..ae19979a4 100644 --- a/pkg/compressio/nocompressio.go +++ b/pkg/compressio/nocompressio.go @@ -16,7 +16,6 @@ package compressio import ( "bufio" - "bytes" "crypto/hmac" "crypto/sha256" "encoding/binary" @@ -60,8 +59,8 @@ type SimpleReader struct { // current chunk position done uint32 - // scratch is a 4-byte scratch buffer used for 32-bit integers. - scratch [4]byte + // scratch is a scratch buffer used for reading chunk size and hash values. + scratch [2 * sha256.Size]byte } var _ io.Reader = (*SimpleReader)(nil) @@ -94,11 +93,11 @@ func (r *SimpleReader) Read(p []byte) (int, error) { // need next chunk? if r.done >= r.chunkSize { - if _, err := io.ReadFull(r.in, r.scratch[:]); err != nil { + if _, err := io.ReadFull(r.in, r.scratch[:4]); err != nil { return 0, err } - r.chunkSize = binary.BigEndian.Uint32(r.scratch[:]) + r.chunkSize = binary.BigEndian.Uint32(r.scratch[:4]) r.done = 0 r.h.Reset() @@ -125,22 +124,28 @@ func (r *SimpleReader) Read(p []byte) (int, error) { return n, err } + // Add data to hash. _, _ = r.h.Write(p[:n]) r.done += uint32(n) + // Is current chunk done? if r.done >= r.chunkSize { - binary.BigEndian.PutUint32(r.scratch[:], r.chunkSize) - r.h.Write(r.scratch[:]) + // Add data size to hash. + binary.BigEndian.PutUint32(r.scratch[:4], r.chunkSize) + r.h.Write(r.scratch[:4]) - sum := r.h.Sum(nil) - readerSum := make([]byte, len(sum)) - if _, err := io.ReadFull(r.in, readerSum); err != nil { + // Compute the hash into the first 32 bytes of the scratch buffer. Pass a + // 32-byte capacity slice (with 0 length) to avoid allocation. + r.h.Sum(r.scratch[0:0:sha256.Size]) + + // Read the hash value from the stream into the second 32 bytes of scratch. + if _, err := io.ReadFull(r.in, r.scratch[sha256.Size:2*sha256.Size]); err != nil { if err == io.EOF { return n, io.ErrUnexpectedEOF } return n, err } - if !hmac.Equal(readerSum, sum) { + if !hmac.Equal(r.scratch[:sha256.Size], r.scratch[sha256.Size:2*sha256.Size]) { return n, ErrHashMismatch } @@ -156,17 +161,23 @@ type SimpleWriter struct { // base is the underlying writer. base io.Writer - // out is a buffered writer. - out *bufio.Writer + // bufOut is a buffered writer. If nil, SimpleWriter does buffering manually. + bufOut *bufio.Writer - // h is the hash object. + // h is the hash object which will be used to checksum each chunk. h hash.Hash + // chunkSize is the data chunk size. chunkSize is immutable. + chunkSize int + + // done is the current chunk position. + done int + + // buf is used to buffer the output. + buf []byte + // closed indicates whether the file has been closed. closed bool - - // scratch is a 4-byte scratch buffer used for 32-bit integers. - scratch [4]byte } var _ io.Writer = (*SimpleWriter)(nil) @@ -174,16 +185,25 @@ var _ io.Closer = (*SimpleWriter)(nil) // NewSimpleWriter returns a new non-compressing writer. If key is non-nil, // hash values are generated and written out for compressed bytes. See package -// comments for details. -func NewSimpleWriter(out io.Writer, key []byte) *SimpleWriter { - w := &SimpleWriter{ - base: out, - out: bufio.NewWriterSize(out, defaultBufSize), +// comments for details. chunkSize is the buffer size used for buffering. Large +// writes are not buffered and written out directly as a single chunk. +func NewSimpleWriter(out io.Writer, key []byte, chunkSize uint32) *SimpleWriter { + if key == nil { + // Since there is no key, this image doesn't use the data integrity stream + // format mentioned in package comments. We can just use a bufio writer. + return &SimpleWriter{ + base: out, + bufOut: bufio.NewWriterSize(out, defaultBufSize), + } } - if key != nil { - w.h = hmac.New(sha256.New, key) + + return &SimpleWriter{ + base: out, + h: hmac.New(sha256.New, key), + chunkSize: int(chunkSize), + // Allocate space for the data size header and the hash. + buf: make([]byte, 4+chunkSize+sha256.Size), } - return w } // Write implements io.Writer.Write. @@ -193,40 +213,84 @@ func (w *SimpleWriter) Write(p []byte) (int, error) { return 0, io.ErrUnexpectedEOF } - if w.h == nil { - return w.out.Write(p) + if w.bufOut != nil { + return w.bufOut.Write(p) } - l := uint32(len(p)) + total := 0 + for len(p) > 0 { + if len(p) > w.chunkSize && w.done == 0 { + // If the payload is larger than the chunk size and we are not in the + // middle of writing another chunk, we can just write it out as one chunk. + n, err := w.directWrite(p) + return total + n, err + } - // chunk length - binary.BigEndian.PutUint32(w.scratch[:], l) - if _, err := w.out.Write(w.scratch[:]); err != nil { + // Copy to buffer. + n := copy(w.buf[4+w.done:4+w.chunkSize], p) + + // Update state. + w.done += n + p = p[n:] + total += n + + // Flush if necessary. + if w.done >= w.chunkSize { + if err := w.flush(); err != nil { + return total, err + } + } + } + return total, nil +} + +// Precondition: w.done == 0. +func (w *SimpleWriter) directWrite(p []byte) (int, error) { + // Write the data size. + binary.BigEndian.PutUint32(w.buf[:4], uint32(len(p))) + if _, err := w.base.Write(w.buf[:4]); err != nil { return 0, err } - // Write out to the stream. - n, err := w.out.Write(p) + // Write the data. + n, err := w.base.Write(p) if err != nil { return n, err } - // Write out the hash. - - // chunk data - _, _ = w.h.Write(p) - - // chunk length - binary.BigEndian.PutUint32(w.scratch[:], l) - w.h.Write(w.scratch[:]) - - sum := w.h.Sum(nil) + // Write the hash. Compute it by writing the data followed by data size. w.h.Reset() - if _, err := io.CopyN(w.out, bytes.NewReader(sum), int64(len(sum))); err != nil { - return n, err + _, _ = w.h.Write(p) + _, _ = w.h.Write(w.buf[:4]) + // Pass a 32-byte capacity slice (with 0 length) to avoid allocation. + w.h.Sum(w.buf[0:0:sha256.Size]) + _, err = w.base.Write(w.buf[:sha256.Size]) + return n, err +} + +func (w *SimpleWriter) flush() error { + if w.done <= 0 { + return nil } - return n, nil + // Add the data size header at the beginning of the buffer. + binary.BigEndian.PutUint32(w.buf[:4], uint32(w.done)) + + // Compute the hash by writing the data followed by data size. + w.h.Reset() + _, _ = w.h.Write(w.buf[4 : 4+w.done]) + _, _ = w.h.Write(w.buf[:4]) + + // Calculate the hash and write it after the data section in the buffer. + // Pass a 32-byte capacity slice (with 0 length) to avoid allocation. + w.h.Sum(w.buf[4+w.done : 4+w.done : 4+w.done+sha256.Size]) + + // Write out to the stream. + _, err := w.base.Write(w.buf[:4+w.done+sha256.Size]) + + // Reset state. + w.done = 0 + return err } // Close implements io.Closer.Close. @@ -238,9 +302,15 @@ func (w *SimpleWriter) Close() error { } w.closed = true - // Flush buffered writer - if err := w.out.Flush(); err != nil { - return err + // Flush buffers. + if w.bufOut != nil { + if err := w.bufOut.Flush(); err != nil { + return err + } + } else { + if err := w.flush(); err != nil { + return err + } } // Close the underlying writer (if necessary). @@ -248,8 +318,9 @@ func (w *SimpleWriter) Close() error { return closer.Close() } - w.out = nil + w.bufOut = nil w.base = nil + w.buf = nil return nil } diff --git a/pkg/compressio/nocompressio_test.go b/pkg/compressio/nocompressio_test.go index e169ddaef..0e445fbbb 100644 --- a/pkg/compressio/nocompressio_test.go +++ b/pkg/compressio/nocompressio_test.go @@ -54,7 +54,7 @@ func TestNoCompress(t *testing.T) { Name: fmt.Sprintf("len(data)=%d, blockSize=%d, key=%s, corruptData=%v", len(data), blockSize, string(key), corruptData), Data: data, NewWriter: func(b *bytes.Buffer) (io.WriteCloser, error) { - return NewSimpleWriter(b, key), nil + return NewSimpleWriter(b, key, blockSize), nil }, NewReader: func(b *bytes.Buffer) (io.Reader, error) { return NewSimpleReader(b, key), nil @@ -87,7 +87,7 @@ func benchmarkNoCompress1ByteWrite(b *testing.B, key []byte, blockSize uint32) { buf [1]byte out bytes.Buffer ) - w := NewSimpleWriter(&out, key) + w := NewSimpleWriter(&out, key, blockSize) for i := 0; i < b.N; i++ { w.Write(buf[:]) } diff --git a/pkg/state/statefile/statefile.go b/pkg/state/statefile/statefile.go index c0edc6da9..6702d22bc 100644 --- a/pkg/state/statefile/statefile.go +++ b/pkg/state/statefile/statefile.go @@ -62,8 +62,8 @@ import ( // keySize is the AES-256 key length. const keySize = 32 -// compressionChunkSize is the chunk size for compression. -const compressionChunkSize = 1024 * 1024 +// stateFileChunkSize is the chunk size used to read/write the state file. +const stateFileChunkSize = 1024 * 1024 // maxMetadataSize is the size limit of metadata section. const maxMetadataSize = 16 * 1024 * 1024 @@ -230,10 +230,10 @@ func NewWriter(w io.Writer, key []byte, metadata map[string]string) (io.WriteClo // gain in restore latency reduction, while incurring much more CPU usage at // save time. if compression == CompressionLevelFlateBestSpeed { - return compressio.NewWriter(w, key, compressionChunkSize, flate.BestSpeed) + return compressio.NewWriter(w, key, stateFileChunkSize, flate.BestSpeed) } - return compressio.NewSimpleWriter(w, key), nil + return compressio.NewSimpleWriter(w, key, stateFileChunkSize), nil } // MetadataUnsafe reads out the metadata from a state file without verifying any diff --git a/pkg/state/statefile/statefile_test.go b/pkg/state/statefile/statefile_test.go index c45029152..af9049f4c 100644 --- a/pkg/state/statefile/statefile_test.go +++ b/pkg/state/statefile/statefile_test.go @@ -69,8 +69,8 @@ func TestStatefile(t *testing.T) { {"longer than hash", []byte("012356asdjflkasjlk3jlk23j4lkjaso0d789f0aujw3lkjlkxsdf78asdful2kj3ljka78"), nil}, // Make sure we have one longer than the chunk size. - {"chunks", make([]byte, 3*compressionChunkSize), nil}, - {"large", make([]byte, 30*compressionChunkSize), nil}, + {"chunks", make([]byte, 3*stateFileChunkSize), nil}, + {"large", make([]byte, 30*stateFileChunkSize), nil}, // Different metadata. {"one metadata", []byte("data"), map[string]string{"foo": "bar"}},