Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 39 additions & 1 deletion zstd_stream.go
Original file line number Diff line number Diff line change
Expand Up @@ -67,20 +67,37 @@ import (
"io"
"runtime"
"sync"
"sync/atomic"
"unsafe"
)

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

ctx *C.ZSTD_CCtx
dict []byte
dstBuffer []byte
dstBufferPtr *[]byte
firstError error
underlyingWriter io.Writer
resultBuffer *C.compressStream2_result
Expand Down Expand Up @@ -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 {
Expand All @@ -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),
Expand All @@ -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 {
Expand Down Expand Up @@ -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
Expand Down
52 changes: 52 additions & 0 deletions zstd_stream_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down