diff --git a/gunzip.go b/gunzip.go index f4e4436..89400b9 100644 --- a/gunzip.go +++ b/gunzip.go @@ -172,6 +172,9 @@ func (z *Reader) Reset(r io.Reader) error { z.digest = crc32.NewIEEE() z.size = 0 z.err = nil + z.lastBlock = false + z.current = nil + z.roff = 0 z.multistream.Store(true) z.readAheadStarted.Store(false) @@ -473,6 +476,11 @@ func (z *Reader) Read(p []byte) (n int, err error) { if len(p) == 0 { return 0, nil } + // Already drained (e.g. via WriteTo/io.Copy). Returning here avoids + // blocking on a full blockPool when current is nil. + if z.lastBlock && len(z.current) == 0 { + return 0, io.EOF + } if z.readAheadStarted.CompareAndSwap(false, true) { z.doReadAhead() @@ -504,8 +512,10 @@ func (z *Reader) Read(p []byte) (n int, err error) { if len(p) >= len(avail) { // If len(p) >= len(current), return all content of current n = copy(p, avail) - z.blockPool <- z.current - z.current = nil + if z.current != nil { + z.blockPool <- z.current + z.current = nil + } if z.lastBlock { err = io.EOF break @@ -555,20 +565,33 @@ func (z *Reader) WriteTo(w io.Writer) (n int64, err error) { } if read.err == io.EOF { z.lastBlock = true - err = nil } } - // Write what we got - n, err := w.Write(read.b) - if n != len(read.b) { - return total, io.ErrShortWrite + // Write what we got (even empty final block). + if len(read.b) > 0 { + n, err := w.Write(read.b) + if n != len(read.b) { + if cap(read.b) > 0 { + z.blockPool <- read.b + } + return total, io.ErrShortWrite + } + total += int64(n) + if err != nil { + if cap(read.b) > 0 { + z.blockPool <- read.b + } + return total, err + } } - total += int64(n) - if err != nil { - return total, err + // Put block back when it came from the pool. + if cap(read.b) > 0 { + z.blockPool <- read.b } - // Put block back - z.blockPool <- read.b + } + // WriterTo must not return io.EOF; io.Copy surfaces it to callers (#38). + if z.err == io.EOF { + return total, nil } return total, z.err } diff --git a/gunzip_test.go b/gunzip_test.go index 70e7258..32eb936 100644 --- a/gunzip_test.go +++ b/gunzip_test.go @@ -847,3 +847,62 @@ func TestWriterTo(t *testing.T) { t.Log("Size", n, "Checksum OK") }) } + +func TestReadAfterWriteToNoDeadlock(t *testing.T) { + // echo hello | gzip -c + gzipData := []byte{ + 0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x03, 0xcb, 0x48, + 0xcd, 0xc9, 0xc9, 0xe7, 0x02, 0x00, 0x20, 0x30, 0x3a, 0x36, 0x06, 0x00, + 0x00, 0x00, + } + rdr, err := NewReader(bytes.NewReader(gzipData)) + if err != nil { + t.Fatal(err) + } + defer rdr.Close() + + n, err := io.Copy(io.Discard, rdr) + if err != nil { + t.Fatalf("WriteTo/Copy: %v", err) + } + if n != 6 { + t.Fatalf("copied %d, want 6", n) + } + + done := make(chan struct{}) + var rn int + var rerr error + go func() { + defer close(done) + var buf [8]byte + rn, rerr = rdr.Read(buf[:]) + }() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("Read after WriteTo deadlocked") + } + if rn != 0 || rerr != io.EOF { + t.Fatalf("Read after drain: n=%d err=%v, want 0, EOF", rn, rerr) + } +} + +func TestWriteToDoesNotReturnEOF(t *testing.T) { + gzipData := []byte{ + 0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x03, 0xcb, 0x48, + 0xcd, 0xc9, 0xc9, 0xe7, 0x02, 0x00, 0x20, 0x30, 0x3a, 0x36, 0x06, 0x00, + 0x00, 0x00, + } + rdr, err := NewReader(bytes.NewReader(gzipData)) + if err != nil { + t.Fatal(err) + } + defer rdr.Close() + n, err := rdr.WriteTo(io.Discard) + if err != nil { + t.Fatalf("WriteTo err=%v, want nil (not EOF)", err) + } + if n != 6 { + t.Fatalf("n=%d want 6", n) + } +}