From 55fa9b6ab2b746fb8ddefdebb8db04941f8acf1e Mon Sep 17 00:00:00 2001 From: Nick Craig-Wood Date: Fri, 28 Aug 2026 14:52:09 +0100 Subject: [PATCH] pool: return an error instead of panicking when an RW is used after Close *pool.RW implements io.Closer, so when one is used directly as an http.Request body the transport closes it after each attempt. A retry which then seeks and re-reads the RW would index the freed pages and panic with "index out of range" in readPage. Read, Write, ReadFrom, WriteTo and Seek now return ErrClosed (which wraps io/fs.ErrClosed) after Close, and Close is safe to call more than once. --- lib/pool/reader_writer.go | 73 ++++++++++++++++++++++++---------- lib/pool/reader_writer_test.go | 43 ++++++++++++++++++++ 2 files changed, 94 insertions(+), 22 deletions(-) diff --git a/lib/pool/reader_writer.go b/lib/pool/reader_writer.go index 3a14af2c0..0806448fe 100644 --- a/lib/pool/reader_writer.go +++ b/lib/pool/reader_writer.go @@ -3,7 +3,9 @@ package pool import ( "context" "errors" + "fmt" "io" + "io/fs" "sync" "time" ) @@ -31,6 +33,7 @@ type RW struct { size int // size written lastOffset int // size in last page written chan struct{} // signalled when a write happens + closed bool // set after Close and the pages are returned to the pool // Read side Variables out int // offset we are reading from @@ -40,6 +43,9 @@ type RW struct { } var ( + // ErrClosed is returned from all operations on an RW after Close. + // It wraps fs.ErrClosed so errors.Is(err, fs.ErrClosed) is true. + ErrClosed = fmt.Errorf("pool.RW: %w", fs.ErrClosed) errInvalidWhence = errors.New("pool.RW Seek: invalid whence") errNegativeSeek = errors.New("pool.RW Seek: negative position") errSeekPastEnd = errors.New("pool.RW Seek: attempt to seek past end of data") @@ -117,10 +123,17 @@ func (rw *RW) DelayAccounting(i int) { // Returns the page and offset of i for reading. // -// Ensure there are pages before calling this. -func (rw *RW) readPage(i int) (page []byte) { +// Returns ErrClosed if the RW has been closed or io.EOF if i is at or +// beyond the end of the data written. +func (rw *RW) readPage(i int) (page []byte, err error) { rw.mu.Lock() defer rw.mu.Unlock() + if rw.closed { + return nil, ErrClosed + } + if i >= rw.size { + return nil, io.EOF + } // Count a read of the data if we read the first page if i == 0 { rw.reads++ @@ -132,7 +145,7 @@ func (rw *RW) readPage(i int) (page []byte) { if pageNumber == len(rw.pages)-1 { page = page[:rw.lastOffset] } - return page[offset:] + return page[offset:], nil } // account for n bytes being read @@ -150,13 +163,6 @@ func (rw *RW) accountRead(n int) error { return nil } -// Returns true if we have read to EOF -func (rw *RW) eof() bool { - rw.mu.Lock() - defer rw.mu.Unlock() - return rw.out >= rw.size -} - // Read reads up to len(p) bytes into p. It returns the number of // bytes read (0 <= n <= len(p)) and any error encountered. If some // data is available but not len(p) bytes, Read returns what is @@ -167,10 +173,10 @@ func (rw *RW) Read(p []byte) (n int, err error) { page []byte ) for len(p) > 0 { - if rw.eof() { - return n, io.EOF + page, err = rw.readPage(rw.out) + if err != nil { + return n, err } - page = rw.readPage(rw.out) nn = copy(p, page) p = p[nn:] n += nn @@ -194,8 +200,13 @@ func (rw *RW) WriteTo(w io.Writer) (n int64, err error) { nn int page []byte ) - for !rw.eof() { - page = rw.readPage(rw.out) + for { + page, err = rw.readPage(rw.out) + if err == io.EOF { + return n, nil + } else if err != nil { + return n, err + } nn, err = w.Write(page) n += int64(nn) rw.out += nn @@ -207,15 +218,19 @@ func (rw *RW) WriteTo(w io.Writer) (n int64, err error) { return n, err } } - return n, nil } // Get the page we are writing to -func (rw *RW) writePage() (page []byte) { +// +// Returns ErrClosed if the RW has been closed. +func (rw *RW) writePage() (page []byte, err error) { rw.mu.Lock() defer rw.mu.Unlock() + if rw.closed { + return nil, ErrClosed + } if len(rw.pages) > 0 && rw.lastOffset < rw.pool.bufferSize { - return rw.pages[len(rw.pages)-1][rw.lastOffset:] + return rw.pages[len(rw.pages)-1][rw.lastOffset:], nil } if len(rw.reserved) > 0 { // Get reserved pages if available @@ -228,7 +243,7 @@ func (rw *RW) writePage() (page []byte) { } rw.pages = append(rw.pages, page) rw.lastOffset = 0 - return page + return page, nil } // Write writes len(p) bytes from p to the underlying data stream. It returns @@ -239,7 +254,10 @@ func (rw *RW) Write(p []byte) (n int, err error) { page []byte ) for len(p) > 0 { - page = rw.writePage() + page, err = rw.writePage() + if err != nil { + return n, err + } nn = copy(page, p) p = p[nn:] n += nn @@ -264,7 +282,10 @@ func (rw *RW) ReadFrom(r io.Reader) (n int64, err error) { page []byte ) for err == nil { - page = rw.writePage() + page, err = rw.writePage() + if err != nil { + return n, err + } nn, err = r.Read(page) n += int64(nn) rw.mu.Lock() @@ -316,7 +337,11 @@ func (rw *RW) Seek(offset int64, whence int) (int64, error) { var abs int64 rw.mu.Lock() size := int64(rw.size) + closed := rw.closed rw.mu.Unlock() + if closed { + return 0, ErrClosed + } switch whence { case io.SeekStart: abs = offset @@ -337,10 +362,14 @@ func (rw *RW) Seek(offset int64, whence int) (int64, error) { return abs, nil } -// Close the buffer returning memory to the pool +// Close the buffer returning memory to the pool. +// +// It is safe to call more than once. All other operations return +// ErrClosed afterwards. func (rw *RW) Close() error { rw.mu.Lock() defer rw.mu.Unlock() + rw.closed = true rw.signalWrite() // signal more data available rw.pool.PutN(rw.pages) clear(rw.pages) diff --git a/lib/pool/reader_writer_test.go b/lib/pool/reader_writer_test.go index 568424c23..c9d3a34bb 100644 --- a/lib/pool/reader_writer_test.go +++ b/lib/pool/reader_writer_test.go @@ -5,6 +5,7 @@ import ( "context" "errors" "io" + "io/fs" "sync" "testing" "time" @@ -607,3 +608,45 @@ func TestRWConcurrency(t *testing.T) { } } + +// Check that using an RW after Close returns an ErrClosed +func TestRWUseAfterClose(t *testing.T) { + rw := NewRW(rwPool) + testData := []byte("Goodness!!") + n, err := rw.Write(testData) + require.NoError(t, err) + assert.Equal(t, len(testData), n) + require.NoError(t, rw.Close()) + + // Rewind and read as an HTTP retry would + _, err = rw.Seek(0, io.SeekStart) + assert.ErrorIs(t, err, ErrClosed) + buf := make([]byte, 16) + _, err = rw.Read(buf) + assert.ErrorIs(t, err, ErrClosed) + _, err = rw.WriteTo(io.Discard) + assert.ErrorIs(t, err, ErrClosed) + _, err = rw.Write(testData) + assert.ErrorIs(t, err, ErrClosed) + _, err = rw.ReadFrom(bytes.NewReader(testData)) + assert.ErrorIs(t, err, ErrClosed) + + // Check it matches the standard library sentinel too + assert.ErrorIs(t, err, fs.ErrClosed) + + // Closing twice is fine + assert.NoError(t, rw.Close()) + + // Check a fully read RW still errors after Close rather than + // looking like an empty body + rw = NewRW(rwPool) + _, err = rw.Write(testData) + require.NoError(t, err) + _, err = rw.WriteTo(io.Discard) + require.NoError(t, err) + require.NoError(t, rw.Close()) + _, err = rw.WriteTo(io.Discard) + assert.ErrorIs(t, err, ErrClosed) + _, err = rw.Read(buf) + assert.ErrorIs(t, err, ErrClosed) +}