diff --git a/zstd_stream.go b/zstd_stream.go index 167c764..f24007a 100644 --- a/zstd_stream.go +++ b/zstd_stream.go @@ -67,6 +67,7 @@ import ( "io" "runtime" "sync" + "sync/atomic" "unsafe" ) @@ -74,6 +75,21 @@ var errShortRead = errors.New("short read") var errReaderClosed = errors.New("Reader is closed") var ErrNoParallelSupport = errors.New("No parallel support") +var writerBufferPoolEnabled atomic.Bool + +var writerBufferPool = sync.Pool{ + New: func() interface{} { + buffer := make([]byte, CompressBound(1024)) + return &buffer + }, +} + +// SetWriterBufferPoolEnabled enables or disables destination buffer pooling +// for new Writers. Pooling is disabled by default. +func SetWriterBufferPoolEnabled(enabled bool) { + writerBufferPoolEnabled.Store(enabled) +} + // Writer is an io.WriteCloser that zstd-compresses its input. type Writer struct { CompressionLevel int @@ -81,6 +97,7 @@ type Writer struct { ctx *C.ZSTD_CCtx dict []byte dstBuffer []byte + dstBufferPtr *[]byte firstError error underlyingWriter io.Writer resultBuffer *C.compressStream2_result @@ -119,6 +136,14 @@ func NewWriterLevel(w io.Writer, level int) *Writer { func NewWriterLevelDict(w io.Writer, level int, dict []byte) *Writer { var err error ctx := C.ZSTD_createCStream() + var dstBufferPtr *[]byte + var dstBuffer []byte + if writerBufferPoolEnabled.Load() { + dstBufferPtr = writerBufferPool.Get().(*[]byte) + dstBuffer = *dstBufferPtr + } else { + dstBuffer = make([]byte, CompressBound(1024)) + } // Load dictionnary if any if dict != nil { @@ -137,7 +162,8 @@ func NewWriterLevelDict(w io.Writer, level int, dict []byte) *Writer { CompressionLevel: level, ctx: ctx, dict: dict, - dstBuffer: make([]byte, CompressBound(1024)), + dstBuffer: dstBuffer, + dstBufferPtr: dstBufferPtr, firstError: err, underlyingWriter: w, resultBuffer: new(C.compressStream2_result), @@ -152,6 +178,16 @@ func finalizeWriter(w *Writer) { } } +func (w *Writer) releaseDstBuffer() { + if w.dstBufferPtr == nil { + return + } + *w.dstBufferPtr = w.dstBuffer[:cap(w.dstBuffer)] + writerBufferPool.Put(w.dstBufferPtr) + w.dstBuffer = nil + w.dstBufferPtr = nil +} + // Write writes a compressed form of p to the underlying io.Writer. func (w *Writer) Write(p []byte) (int, error) { if w.firstError != nil { @@ -255,6 +291,8 @@ func (w *Writer) Flush() error { // Close closes the Writer, flushing any unwritten data to the underlying // io.Writer and freeing objects, but does not close the underlying io.Writer. func (w *Writer) Close() error { + defer w.releaseDstBuffer() + if w.ctx == nil { if w.firstError != nil { return w.firstError diff --git a/zstd_stream_test.go b/zstd_stream_test.go index 049dd61..e58629b 100644 --- a/zstd_stream_test.go +++ b/zstd_stream_test.go @@ -102,6 +102,58 @@ func TestZstdReaderLong(t *testing.T) { testCompressionDecompression(t, nil, long.Bytes(), 1) } +func TestWriterBufferPool(t *testing.T) { + SetWriterBufferPoolEnabled(false) + defer SetWriterBufferPoolEnabled(false) + + unpooled := NewWriter(io.Discard) + if unpooled.dstBufferPtr != nil { + t.Fatal("writer buffer pooling is enabled by default") + } + failOnError(t, "failed to close unpooled writer", unpooled.Close()) + + SetWriterBufferPoolEnabled(true) + payload := bytes.Repeat([]byte("payload"), 10000) + pooled := NewWriter(io.Discard) + if pooled.dstBufferPtr == nil { + t.Fatal("writer did not use the enabled buffer pool") + } + buffer := pooled.dstBufferPtr + _, err := pooled.Write(payload) + failOnError(t, "failed to write with pooled writer", err) + failOnError(t, "failed to close pooled writer", pooled.Close()) + + if pooled.dstBuffer != nil || pooled.dstBufferPtr != nil { + t.Fatal("writer kept the pooled buffer after close") + } + if cap(*buffer) < CompressBound(len(payload)) { + t.Fatalf("pooled buffer capacity = %d, want at least %d", cap(*buffer), CompressBound(len(payload))) + } +} + +func BenchmarkStreamWriterBufferPool(b *testing.B) { + payload := bytes.Repeat([]byte("payload"), 10000) + for _, enabled := range []bool{false, true} { + b.Run(fmt.Sprintf("enabled=%t", enabled), func(b *testing.B) { + SetWriterBufferPoolEnabled(enabled) + b.Cleanup(func() { + SetWriterBufferPoolEnabled(false) + }) + b.ReportAllocs() + b.SetBytes(int64(len(payload))) + for i := 0; i < b.N; i++ { + writer := NewWriter(io.Discard) + if _, err := writer.Write(payload); err != nil { + b.Fatal(err) + } + if err := writer.Close(); err != nil { + b.Fatal(err) + } + } + }) + } +} + func doStreamCompressionDecompression() error { payload := []byte("Hello World!") repeat := 10000