diff --git a/packages/orchestrator/cmd/inspect-build/validate.go b/packages/orchestrator/cmd/inspect-build/validate.go index 5ab8011579..9548c04012 100644 --- a/packages/orchestrator/cmd/inspect-build/validate.go +++ b/packages/orchestrator/cmd/inspect-build/validate.go @@ -92,7 +92,7 @@ func validateBuild(ctx context.Context, storagePath, buildID, artifact string, e // fetch still runs and reportValidation shows the checksum as n/a. expected := h.Builds[h.Metadata.BuildId].Checksum - chunker, obj, size, cleanup, err := openChunker(ctx, storagePath, buildID, artifact, h, ft) + chunker, upstream, size, cleanup, err := openChunker(ctx, storagePath, buildID, artifact, h, ft) if err != nil { return err } @@ -119,7 +119,7 @@ func validateBuild(ctx context.Context, storagePath, buildID, artifact string, e } c := chunks[ci] for off := c.lo; off < c.hi; off += blockSize { - if _, err := chunker.Slice(egCtx, off, min(blockSize, c.hi-off), obj, ft); err != nil { + if _, err := chunker.Slice(egCtx, off, min(blockSize, c.hi-off), upstream, ft); err != nil { return fmt.Errorf("fetch block at %d: %w", off, err) } } @@ -133,7 +133,7 @@ func validateBuild(ctx context.Context, storagePath, buildID, artifact string, e // The cache is warm — sweep it one block at a time to hash the image. hasher := sha256.New() for off := int64(0); off < size; off += blockSize { - b, err := chunker.Slice(ctx, off, min(blockSize, size-off), obj, ft) + b, err := chunker.Slice(ctx, off, min(blockSize, size-off), upstream, ft) if err != nil { return fmt.Errorf("read block at %d: %w", off, err) } @@ -146,9 +146,8 @@ func validateBuild(ctx context.Context, storagePath, buildID, artifact string, e } // openChunker wires a production block.Chunker over the build's data file, -// returning the chunker, the upstream storage object, the image's uncompressed -// size, and a cleanup function. -func openChunker(ctx context.Context, storagePath, buildID, artifact string, h *header.Header, ft *storage.FrameTable) (*block.Chunker, storage.Seekable, int64, func(), error) { +// returning the chunker, the image's uncompressed size, and a cleanup function. +func openChunker(ctx context.Context, storagePath, buildID, artifact string, h *header.Header, ft *storage.FrameTable) (*block.Chunker, storage.RangeOpener, int64, func(), error) { if err := cmdutil.SetupStorage(storagePath); err != nil { return nil, nil, 0, nil, err } diff --git a/packages/orchestrator/pkg/sandbox/block/streaming_chunk.go b/packages/orchestrator/pkg/sandbox/block/streaming_chunk.go index c21eedbd19..5c05ff7255 100644 --- a/packages/orchestrator/pkg/sandbox/block/streaming_chunk.go +++ b/packages/orchestrator/pkg/sandbox/block/streaming_chunk.go @@ -241,7 +241,7 @@ func (c *Chunker) progressiveRead(ctx context.Context, s *fetchSession, mmapSlic return 0, fmt.Errorf("failed to open range reader at %d: %w", s.chunkOff, err) } defer func() { - if closeErr := reader.Close(); closeErr != nil && err == nil { + if closeErr := reader.Close(context.WithoutCancel(ctx)); closeErr != nil && err == nil { err = closeErr } }() diff --git a/packages/orchestrator/pkg/sandbox/block/streaming_chunk_test.go b/packages/orchestrator/pkg/sandbox/block/streaming_chunk_test.go index 7b8db7b570..8295daaa46 100644 --- a/packages/orchestrator/pkg/sandbox/block/streaming_chunk_test.go +++ b/packages/orchestrator/pkg/sandbox/block/streaming_chunk_test.go @@ -82,7 +82,7 @@ func (s *fakeSeekable) StoreFile(context.Context, string, ...storage.PutOption) panic("not used") } -func (s *fakeSeekable) OpenRangeReader(_ context.Context, offsetU int64, length int64, frameTable *storage.FrameTable) (io.ReadCloser, error) { +func (s *fakeSeekable) OpenRangeReader(_ context.Context, offsetU int64, length int64, frameTable *storage.FrameTable) (storage.RangeReader, error) { s.fetchCount.Add(1) if s.ctrl != nil { @@ -127,10 +127,10 @@ func (s *fakeSeekable) OpenRangeReader(_ context.Context, offsetU int64, length r := io.Reader(bytes.NewReader(s.data[fetchOff:end])) if frameTable.IsCompressed() { - return storage.NewDecompressingReader(r, frameTable.CompressionType()) + return storage.NewDecompressingReader(storage.NewRangeReader(io.NopCloser(r)), frameTable.CompressionType()) } - return io.NopCloser(r), nil + return storage.NewRangeReader(io.NopCloser(r)), nil } func makeCompressedTestData(tb testing.TB, data []byte) (*storage.FrameTable, *fakeSeekable) { @@ -429,7 +429,7 @@ func (s *panicSeekable) StoreFile(context.Context, string, ...storage.PutOption) panic("not used") } -func (s *panicSeekable) OpenRangeReader(_ context.Context, off int64, length int64, _ *storage.FrameTable) (io.ReadCloser, error) { +func (s *panicSeekable) OpenRangeReader(_ context.Context, off int64, length int64, _ *storage.FrameTable) (storage.RangeReader, error) { end := min(off+length, int64(len(s.data))) return &panicReader{ @@ -460,7 +460,7 @@ func (r *panicReader) Read(p []byte) (int, error) { return n, nil } -func (r *panicReader) Close() error { +func (r *panicReader) Close(context.Context) error { return nil } @@ -587,7 +587,7 @@ func (r *controlledReader) Read(p []byte) (int, error) { return n, nil } -func (r *controlledReader) Close() error { +func (r *controlledReader) Close(context.Context) error { select { case r.closed <- struct{}{}: default: diff --git a/packages/orchestrator/pkg/sandbox/build/storage_diff_test.go b/packages/orchestrator/pkg/sandbox/build/storage_diff_test.go index d6dff8bc1f..e00fe8021a 100644 --- a/packages/orchestrator/pkg/sandbox/build/storage_diff_test.go +++ b/packages/orchestrator/pkg/sandbox/build/storage_diff_test.go @@ -188,10 +188,10 @@ func TestStorageDiff_NoRefreshOnFinalizedHeader(t *testing.T) { uncompressedSeekable := storage.NewMockSeekable(t) uncompressedSeekable.EXPECT(). OpenRangeReader(mock.Anything, mock.Anything, mock.Anything, mock.Anything). - RunAndReturn(func(_ context.Context, off, length int64, _ *storage.FrameTable) (io.ReadCloser, error) { + RunAndReturn(func(_ context.Context, off, length int64, _ *storage.FrameTable) (storage.RangeReader, error) { end := min(off+length, int64(len(payload))) - return io.NopCloser(bytes.NewReader(payload[off:end])), nil + return storage.NewRangeReader(io.NopCloser(bytes.NewReader(payload[off:end]))), nil }) provider.EXPECT(). OpenSeekable(mock.Anything, uncompressedPath, mock.Anything). @@ -235,10 +235,10 @@ func TestStorageDiff_SkipsHeaderRefreshWhenPeerActive(t *testing.T) { Return(int64(payloadSize), nil).Once() uncompressedSeekable.EXPECT(). OpenRangeReader(mock.Anything, mock.Anything, mock.Anything, mock.Anything). - RunAndReturn(func(_ context.Context, off, length int64, _ *storage.FrameTable) (io.ReadCloser, error) { + RunAndReturn(func(_ context.Context, off, length int64, _ *storage.FrameTable) (storage.RangeReader, error) { end := min(off+length, int64(len(payload))) - return io.NopCloser(bytes.NewReader(payload[off:end])), nil + return storage.NewRangeReader(io.NopCloser(bytes.NewReader(payload[off:end]))), nil }) provider.EXPECT(). OpenSeekable(mock.Anything, aPaths.DataFile(storage.MemfileName, storage.CompressionNone), mock.Anything). @@ -299,10 +299,10 @@ func TestStorageDiff_V3AncestorFallsBackToUncompressed(t *testing.T) { Return(int64(payloadSize), nil).Once() rawSeekable.EXPECT(). OpenRangeReader(mock.Anything, mock.Anything, mock.Anything, mock.Anything). - RunAndReturn(func(_ context.Context, off, length int64, _ *storage.FrameTable) (io.ReadCloser, error) { + RunAndReturn(func(_ context.Context, off, length int64, _ *storage.FrameTable) (storage.RangeReader, error) { end := min(off+length, int64(len(payload))) - return io.NopCloser(bytes.NewReader(payload[off:end])), nil + return storage.NewRangeReader(io.NopCloser(bytes.NewReader(payload[off:end]))), nil }) provider.EXPECT(). OpenSeekable(mock.Anything, aPaths.DataFile(storage.MemfileName, storage.CompressionNone), mock.Anything). @@ -348,10 +348,10 @@ func TestStorageDiff_ReloadSourceLatchesV3AsUncompressed(t *testing.T) { rawSeekable := storage.NewMockSeekable(t) rawSeekable.EXPECT(). OpenRangeReader(mock.Anything, mock.Anything, mock.Anything, mock.Anything). - RunAndReturn(func(_ context.Context, off, length int64, _ *storage.FrameTable) (io.ReadCloser, error) { + RunAndReturn(func(_ context.Context, off, length int64, _ *storage.FrameTable) (storage.RangeReader, error) { end := min(off+length, int64(len(payload))) - return io.NopCloser(bytes.NewReader(payload[off:end])), nil + return storage.NewRangeReader(io.NopCloser(bytes.NewReader(payload[off:end]))), nil }) provider.EXPECT(). OpenSeekable(mock.Anything, aPaths.DataFile(storage.MemfileName, storage.CompressionNone), mock.Anything). @@ -434,14 +434,14 @@ func buildHeader(t *testing.T, selfID uuid.UUID, size int64, mapsTo uuid.UUID) * // requested U-offset in the caller's frame table, slices the compressed // payload, and streams it through a decompressor. Mirrors what a real // Seekable does over zstd-compressed data. -func decompressingRangeReader(compressed []byte) func(context.Context, int64, int64, *storage.FrameTable) (io.ReadCloser, error) { - return func(_ context.Context, offsetU, _ int64, ft *storage.FrameTable) (io.ReadCloser, error) { +func decompressingRangeReader(compressed []byte) func(context.Context, int64, int64, *storage.FrameTable) (storage.RangeReader, error) { + return func(_ context.Context, offsetU, _ int64, ft *storage.FrameTable) (storage.RangeReader, error) { r, err := ft.LocateCompressed(offsetU) if err != nil { return nil, err } end := min(r.Offset+int64(r.Length), int64(len(compressed))) - return storage.NewDecompressingReader(bytes.NewReader(compressed[r.Offset:end]), ft.CompressionType()) + return storage.NewDecompressingReader(storage.NewRangeReader(io.NopCloser(bytes.NewReader(compressed[r.Offset:end]))), ft.CompressionType()) } } diff --git a/packages/orchestrator/pkg/sandbox/template/peerclient/blob.go b/packages/orchestrator/pkg/sandbox/template/peerclient/blob.go index 65891b2f70..92067e6011 100644 --- a/packages/orchestrator/pkg/sandbox/template/peerclient/blob.go +++ b/packages/orchestrator/pkg/sandbox/template/peerclient/blob.go @@ -65,7 +65,7 @@ func (b *peerBlob) WriteTo(ctx context.Context, dst io.Writer) (int64, error) { } reader := newPeerStreamReader(recv, cancel) - defer reader.Close() + defer reader.Close(context.WithoutCancel(ctx)) n, err := io.Copy(dst, reader) if err != nil { diff --git a/packages/orchestrator/pkg/sandbox/template/peerclient/seekable.go b/packages/orchestrator/pkg/sandbox/template/peerclient/seekable.go index 792a18c792..ee0c8a972f 100644 --- a/packages/orchestrator/pkg/sandbox/template/peerclient/seekable.go +++ b/packages/orchestrator/pkg/sandbox/template/peerclient/seekable.go @@ -4,7 +4,6 @@ import ( "context" "errors" "fmt" - "io" "sync" "sync/atomic" @@ -98,9 +97,9 @@ func (s *peerSeekable) Size(ctx context.Context) (int64, error) { return 0, &storage.PeerTransitionedError{} } -func (s *peerSeekable) OpenRangeReader(ctx context.Context, off int64, length int64, frameTable *storage.FrameTable) (io.ReadCloser, error) { +func (s *peerSeekable) OpenRangeReader(ctx context.Context, off int64, length int64, frameTable *storage.FrameTable) (storage.RangeReader, error) { res, err := tryPeer(ctx, &s.peerHandle, "peer-seekable-open-range-reader", attrOpRangeReader, - func(ctx context.Context) (peerAttempt[io.ReadCloser], error) { + func(ctx context.Context) (peerAttempt[storage.RangeReader], error) { streamCtx, cancel := context.WithCancel(ctx) recv, err := openPeerSeekableStream(streamCtx, s.client, &orchestrator.ReadAtBuildSeekableRequest{ @@ -113,10 +112,10 @@ func (s *peerSeekable) OpenRangeReader(ctx context.Context, off int64, length in logger.L().Warn(ctx, "failed to open range reader from peer", logger.WithBuildID(s.buildID), zap.Int64("off", off), zap.Int64("length", length), zap.Error(err)) cancel() - return peerAttempt[io.ReadCloser]{}, nil + return peerAttempt[storage.RangeReader]{}, nil } - return peerAttempt[io.ReadCloser]{ + return peerAttempt[storage.RangeReader]{ value: newPeerStreamReader(recv, cancel), hit: true, }, nil diff --git a/packages/orchestrator/pkg/sandbox/template/peerclient/seekable_test.go b/packages/orchestrator/pkg/sandbox/template/peerclient/seekable_test.go index fd8e8a1cfe..5c1f8ea9d6 100644 --- a/packages/orchestrator/pkg/sandbox/template/peerclient/seekable_test.go +++ b/packages/orchestrator/pkg/sandbox/template/peerclient/seekable_test.go @@ -75,7 +75,7 @@ func TestPeerSeekable_OpenRangeReader_PeerSucceeds(t *testing.T) { s := &peerSeekable{peerHandle: peerHandle{client: client, buildID: "build-1", name: storage.MemfileName, uploaded: &atomic.Bool{}}} rc, err := s.OpenRangeReader(t.Context(), 10, int64(len(data)), nil) require.NoError(t, err) - defer rc.Close() + defer rc.Close(t.Context()) got, err := io.ReadAll(rc) require.NoError(t, err) @@ -90,7 +90,7 @@ func TestPeerSeekable_OpenRangeReader_PeerError_FallsBackToBase(t *testing.T) { client.EXPECT().ReadAtBuildSeekable(mock.Anything, mock.Anything).Return(nil, errors.New("peer unavailable")) baseSeekable := storage.NewMockSeekable(t) - baseSeekable.EXPECT().OpenRangeReader(mock.Anything, int64(0), int64(len(baseData)), (*storage.FrameTable)(nil)).Return(io.NopCloser(bytes.NewReader(baseData)), nil) + baseSeekable.EXPECT().OpenRangeReader(mock.Anything, int64(0), int64(len(baseData)), (*storage.FrameTable)(nil)).Return(storage.NewRangeReader(io.NopCloser(bytes.NewReader(baseData))), nil) base := storage.NewMockStorageProvider(t) base.EXPECT().OpenSeekable(mock.Anything, "build-1/memfile", storage.MemfileObjectType).Return(baseSeekable, nil) @@ -107,7 +107,7 @@ func TestPeerSeekable_OpenRangeReader_PeerError_FallsBackToBase(t *testing.T) { } rc, err := s.OpenRangeReader(t.Context(), 0, int64(len(baseData)), nil) require.NoError(t, err) - defer rc.Close() + defer rc.Close(t.Context()) got, err := io.ReadAll(rc) require.NoError(t, err) @@ -185,7 +185,7 @@ func TestPeerStorageProvider_TransitionEmitsError(t *testing.T) { require.NoError(t, err) got, err := io.ReadAll(rc) require.NoError(t, err) - require.NoError(t, rc.Close()) + require.NoError(t, rc.Close(t.Context())) assert.Equal(t, prePeerBytes, got) require.True(t, uploaded.Load(), "uploaded flag should be set after peer EOF with UseStorage") diff --git a/packages/orchestrator/pkg/sandbox/template/peerclient/storage.go b/packages/orchestrator/pkg/sandbox/template/peerclient/storage.go index eb463e53d0..de51c870b8 100644 --- a/packages/orchestrator/pkg/sandbox/template/peerclient/storage.go +++ b/packages/orchestrator/pkg/sandbox/template/peerclient/storage.go @@ -273,9 +273,9 @@ func tryPeer[T any]( return peerAttempt[T]{}, nil } -var _ io.ReadCloser = (*peerStreamReader)(nil) +var _ storage.RangeReader = (*peerStreamReader)(nil) -// peerStreamReader wraps a gRPC streaming recv function as an io.ReadCloser. +// peerStreamReader wraps a gRPC streaming recv function as a storage.RangeReader. // cancel is called on Close to signal the server to terminate the stream. type peerStreamReader struct { recv func() ([]byte, error) @@ -317,7 +317,7 @@ func (r *peerStreamReader) Read(p []byte) (int, error) { } } -func (r *peerStreamReader) Close() error { +func (r *peerStreamReader) Close(context.Context) error { r.cancel() return nil diff --git a/packages/shared/pkg/storage/compress_decode.go b/packages/shared/pkg/storage/compress_decode.go index b8a6d414a7..a0a706bab3 100644 --- a/packages/shared/pkg/storage/compress_decode.go +++ b/packages/shared/pkg/storage/compress_decode.go @@ -1,6 +1,7 @@ package storage import ( + "context" "fmt" "io" "sync" @@ -9,6 +10,8 @@ import ( lz4 "github.com/pierrec/lz4/v4" ) +var _ RangeReader = (*decompressReader)(nil) + var lz4DecoderPool sync.Pool func getLZ4Decoder(r io.Reader) *lz4.Reader { @@ -53,78 +56,47 @@ func putZstdDecoder(dec *zstd.Decoder) { zstdDecoderPool.Put(dec) } -// NewDecompressingReader wraps a reader with the appropriate decompressor. -// Close releases the decompressor back to its pool but does NOT close the -// underlying reader — the caller is responsible for closing it. -func NewDecompressingReader(raw io.Reader, ct CompressionType) (io.ReadCloser, error) { +// decompressReader decompresses inner on Read; Close releases the codec back +// to its pool and closes inner. +type decompressReader struct { + inner RangeReader + dec io.Reader + releaseCodec func() +} + +func NewDecompressingReader(inner RangeReader, ct CompressionType) (RangeReader, error) { + var dec io.Reader + var releaseCodec func() + switch ct { case CompressionLZ4: - dec := getLZ4Decoder(raw) - - return &pooledDecoder{ - Reader: dec, - close: func() { putLZ4Decoder(dec) }, - }, nil + d := getLZ4Decoder(inner) + dec, releaseCodec = d, func() { putLZ4Decoder(d) } case CompressionZstd: - dec, err := getZstdDecoder(raw) + d, err := getZstdDecoder(inner) if err != nil { return nil, fmt.Errorf("failed to create zstd decoder: %w", err) } - - return &pooledDecoder{ - Reader: dec, - close: func() { putZstdDecoder(dec) }, - }, nil + dec, releaseCodec = d, func() { putZstdDecoder(d) } default: return nil, fmt.Errorf("unsupported compression type: %s", ct) } -} - -// pooledDecoder wraps a decompressor from a sync.Pool. -// Close returns the decompressor to the pool. -type pooledDecoder struct { - io.Reader - close func() + return &decompressReader{ + inner: inner, + dec: dec, + releaseCodec: releaseCodec, + }, nil } -func (r *pooledDecoder) Close() error { - r.close() - - return nil +func (r *decompressReader) Read(p []byte) (int, error) { + return r.dec.Read(p) } -// newDecompressingReadCloser wraps raw with the appropriate decompressor and -// takes ownership: Close releases the decompressor back to the pool AND closes raw. -func newDecompressingReadCloser(raw io.ReadCloser, ct CompressionType) (io.ReadCloser, error) { - dec, err := NewDecompressingReader(raw, ct) - if err != nil { - return nil, err - } - - return &decompressingReadCloser{dec: dec, raw: raw}, nil -} - -// decompressingReadCloser reads from the decompressor and closes both the -// decompressor (returning it to the pool) and the underlying raw stream. -type decompressingReadCloser struct { - dec io.ReadCloser // decompressor — reads from raw - raw io.Closer // underlying stream -} - -func (c *decompressingReadCloser) Read(p []byte) (int, error) { - return c.dec.Read(p) -} - -func (c *decompressingReadCloser) Close() error { - decErr := c.dec.Close() - rawErr := c.raw.Close() - - if decErr != nil { - return decErr - } +func (r *decompressReader) Close(ctx context.Context) error { + r.releaseCodec() - return rawErr + return r.inner.Close(ctx) } diff --git a/packages/shared/pkg/storage/io_wrappers.go b/packages/shared/pkg/storage/io_wrappers.go new file mode 100644 index 0000000000..18d5e5ebba --- /dev/null +++ b/packages/shared/pkg/storage/io_wrappers.go @@ -0,0 +1,143 @@ +package storage + +import ( + "bytes" + "context" + "errors" + "io" + "os" + + "go.opentelemetry.io/otel/trace" + + "github.com/e2b-dev/infra/packages/shared/pkg/telemetry" +) + +var ( + _ RangeReader = (*sectionReader)(nil) + _ RangeReader = (*observableReader)(nil) + _ RangeReader = (*rangeReader)(nil) + _ RangeReader = (*captureReader)(nil) +) + +// rangeReader adapts an io.ReadCloser into a RangeReader by ignoring the +// Close context. +type rangeReader struct { + io.ReadCloser +} + +func NewRangeReader(rc io.ReadCloser) RangeReader { return &rangeReader{ReadCloser: rc} } + +func (p *rangeReader) Close(context.Context) error { + return p.ReadCloser.Close() +} + +type sectionReader struct { + *io.SectionReader + + file *os.File +} + +func newSectionReader(f *os.File, off, length int64) *sectionReader { + return §ionReader{ + SectionReader: io.NewSectionReader(f, off, length), + file: f, + } +} + +func (r *sectionReader) Close(context.Context) error { + return r.file.Close() +} + +// captureReader tees every read byte into a buffer and hands the captured +// bytes to onClose on Close. Used by the cache writeback paths. +// +// drainOnClose=true reads inner to EOF on Close even if the caller above hasn't +// consumed everything. Needed when capturing under a decoder that stops short +// of EOF on its source (e.g. lz4.Reader skips the 4-byte EndMark). +type captureReader struct { + inner RangeReader + buf *bytes.Buffer + onClose func(ctx context.Context, captured []byte) + drainOnClose bool +} + +func newCaptureReader(inner RangeReader, capHint int, drainOnClose bool, onClose func(context.Context, []byte)) *captureReader { + return &captureReader{ + inner: inner, + buf: bytes.NewBuffer(make([]byte, 0, capHint)), + onClose: onClose, + drainOnClose: drainOnClose, + } +} + +func (r *captureReader) Read(p []byte) (int, error) { + n, err := r.inner.Read(p) + if n > 0 { + r.buf.Write(p[:n]) + } + + return n, err +} + +func (r *captureReader) Close(ctx context.Context) error { + if r.drainOnClose { + _, _ = io.Copy(io.Discard, r) + } + err := r.inner.Close(ctx) + if err == nil { + r.onClose(ctx, r.buf.Bytes()) + } + + return err +} + +// observableReader layers OTEL observability (legacy per-backend timer + span) +// onto an inner RangeReader, all applied on Close. timer and span are optional; +// pass nil if unused. +type observableReader struct { + inner RangeReader + timer *telemetry.Stopwatch + span trace.Span + + bytes int64 + readErr error +} + +func newObservableReader(inner RangeReader, timer *telemetry.Stopwatch, span trace.Span) *observableReader { + return &observableReader{inner: inner, timer: timer, span: span} +} + +func (r *observableReader) Read(p []byte) (int, error) { + n, err := r.inner.Read(p) + r.bytes += int64(n) + + if err != nil && !errors.Is(err, io.EOF) { + r.readErr = err + } + + return n, err +} + +func (r *observableReader) Close(ctx context.Context) error { + closeErr := r.inner.Close(ctx) + + if r.timer != nil { + if r.readErr != nil || closeErr != nil { + r.timer.Failure(ctx, r.bytes) + } else { + r.timer.Success(ctx, r.bytes) + } + } + + if r.span != nil { + if closeErr != nil { + recordError(r.span, closeErr) + } else if r.readErr != nil { + recordError(r.span, r.readErr) + } + + r.span.End() + } + + return closeErr +} diff --git a/packages/shared/pkg/storage/mock_seekable.go b/packages/shared/pkg/storage/mock_seekable.go index 610bddaca0..6c3e18a4fc 100644 --- a/packages/shared/pkg/storage/mock_seekable.go +++ b/packages/shared/pkg/storage/mock_seekable.go @@ -6,7 +6,6 @@ package storage import ( "context" - "io" mock "github.com/stretchr/testify/mock" ) @@ -39,23 +38,23 @@ func (_m *MockSeekable) EXPECT() *MockSeekable_Expecter { } // OpenRangeReader provides a mock function for the type MockSeekable -func (_mock *MockSeekable) OpenRangeReader(ctx context.Context, offsetU int64, length int64, frameTable *FrameTable) (io.ReadCloser, error) { +func (_mock *MockSeekable) OpenRangeReader(ctx context.Context, offsetU int64, length int64, frameTable *FrameTable) (RangeReader, error) { ret := _mock.Called(ctx, offsetU, length, frameTable) if len(ret) == 0 { panic("no return value specified for OpenRangeReader") } - var r0 io.ReadCloser + var r0 RangeReader var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, int64, int64, *FrameTable) (io.ReadCloser, error)); ok { + if returnFunc, ok := ret.Get(0).(func(context.Context, int64, int64, *FrameTable) (RangeReader, error)); ok { return returnFunc(ctx, offsetU, length, frameTable) } - if returnFunc, ok := ret.Get(0).(func(context.Context, int64, int64, *FrameTable) io.ReadCloser); ok { + if returnFunc, ok := ret.Get(0).(func(context.Context, int64, int64, *FrameTable) RangeReader); ok { r0 = returnFunc(ctx, offsetU, length, frameTable) } else { if ret.Get(0) != nil { - r0 = ret.Get(0).(io.ReadCloser) + r0 = ret.Get(0).(RangeReader) } } if returnFunc, ok := ret.Get(1).(func(context.Context, int64, int64, *FrameTable) error); ok { @@ -108,12 +107,12 @@ func (_c *MockSeekable_OpenRangeReader_Call) Run(run func(ctx context.Context, o return _c } -func (_c *MockSeekable_OpenRangeReader_Call) Return(readCloser io.ReadCloser, err error) *MockSeekable_OpenRangeReader_Call { - _c.Call.Return(readCloser, err) +func (_c *MockSeekable_OpenRangeReader_Call) Return(rangeReader RangeReader, err error) *MockSeekable_OpenRangeReader_Call { + _c.Call.Return(rangeReader, err) return _c } -func (_c *MockSeekable_OpenRangeReader_Call) RunAndReturn(run func(ctx context.Context, offsetU int64, length int64, frameTable *FrameTable) (io.ReadCloser, error)) *MockSeekable_OpenRangeReader_Call { +func (_c *MockSeekable_OpenRangeReader_Call) RunAndReturn(run func(ctx context.Context, offsetU int64, length int64, frameTable *FrameTable) (RangeReader, error)) *MockSeekable_OpenRangeReader_Call { _c.Call.Return(run) return _c } diff --git a/packages/shared/pkg/storage/offset_reader.go b/packages/shared/pkg/storage/offset_reader.go deleted file mode 100644 index 29d9048d6c..0000000000 --- a/packages/shared/pkg/storage/offset_reader.go +++ /dev/null @@ -1,23 +0,0 @@ -package storage - -import ( - "io" -) - -type offsetReader struct { - wrapped io.ReaderAt - offset int64 -} - -var _ io.Reader = (*offsetReader)(nil) - -func (r *offsetReader) Read(p []byte) (n int, err error) { - n, err = r.wrapped.ReadAt(p, r.offset) - r.offset += int64(n) - - return -} - -func newOffsetReader(reader io.ReaderAt, offset int64) *offsetReader { - return &offsetReader{reader, offset} -} diff --git a/packages/shared/pkg/storage/offset_reader_test.go b/packages/shared/pkg/storage/offset_reader_test.go deleted file mode 100644 index 1c59c46cbc..0000000000 --- a/packages/shared/pkg/storage/offset_reader_test.go +++ /dev/null @@ -1,128 +0,0 @@ -package storage - -import ( - "bytes" - "io" - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -func TestOffsetReader_Read(t *testing.T) { - t.Parallel() - - data := []byte("hello world") - readerAt := bytes.NewReader(data) - - tests := []struct { - name string - offset int64 - readSize int - expectedData string - expectedN int - expectedErr error - expectedOffset int64 - }{ - { - name: "read from start", - offset: 0, - readSize: 5, - expectedData: "hello", - expectedN: 5, - expectedErr: nil, - expectedOffset: 5, - }, - { - name: "read from offset", - offset: 6, - readSize: 5, - expectedData: "world", - expectedN: 5, - expectedErr: nil, - expectedOffset: 11, - }, - { - name: "read until EOF", - offset: 0, - readSize: 11, - expectedData: "hello world", - expectedN: 11, - expectedErr: nil, - expectedOffset: 11, - }, - { - name: "read past EOF", - offset: 0, - readSize: 15, - expectedData: "hello world", - expectedN: 11, - expectedErr: io.EOF, - expectedOffset: 11, - }, - { - name: "read exactly at EOF", - offset: 11, - readSize: 5, - expectedData: "", - expectedN: 0, - expectedErr: io.EOF, - expectedOffset: 11, - }, - { - name: "read zero bytes", - offset: 0, - readSize: 0, - expectedData: "", - expectedN: 0, - expectedErr: nil, - expectedOffset: 0, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - - r := newOffsetReader(readerAt, tt.offset) - p := make([]byte, tt.readSize) - n, err := r.Read(p) - - require.ErrorIs(t, err, tt.expectedErr) - assert.Equal(t, tt.expectedN, n) - assert.Equal(t, tt.expectedData, string(p[:n])) - assert.Equal(t, tt.expectedOffset, r.offset) - }) - } -} - -func TestOffsetReader_SequentialReads(t *testing.T) { - t.Parallel() - - data := []byte("hello world") - readerAt := bytes.NewReader(data) - r := newOffsetReader(readerAt, 0) - - // First read - p1 := make([]byte, 6) - n1, err1 := r.Read(p1) - require.NoError(t, err1) - assert.Equal(t, 6, n1) - assert.Equal(t, "hello ", string(p1[:n1])) - assert.Equal(t, int64(6), r.offset) - - // Second read - p2 := make([]byte, 5) - n2, err2 := r.Read(p2) - require.NoError(t, err2) - assert.Equal(t, 5, n2) - assert.Equal(t, "world", string(p2[:n2])) - assert.Equal(t, int64(11), r.offset) - - // Third read (EOF) - p3 := make([]byte, 5) - n3, err3 := r.Read(p3) - require.ErrorIs(t, err3, io.EOF) - assert.Equal(t, 0, n3) - assert.Equal(t, int64(11), r.offset) -} diff --git a/packages/shared/pkg/storage/storage.go b/packages/shared/pkg/storage/storage.go index 9a2c463c3b..43491661ef 100644 --- a/packages/shared/pkg/storage/storage.go +++ b/packages/shared/pkg/storage/storage.go @@ -140,9 +140,14 @@ type Blob interface { Exists(ctx context.Context) (bool, error) } +type RangeReader interface { + io.Reader + Close(ctx context.Context) error +} + // RangeOpener supports progressive reads via a streaming range reader. type RangeOpener interface { - OpenRangeReader(ctx context.Context, offsetU int64, length int64, frameTable *FrameTable) (io.ReadCloser, error) + OpenRangeReader(ctx context.Context, offsetU int64, length int64, frameTable *FrameTable) (RangeReader, error) } type SeekableWriter interface { diff --git a/packages/shared/pkg/storage/storage_aws.go b/packages/shared/pkg/storage/storage_aws.go index 467c01c79f..b608eeb423 100644 --- a/packages/shared/pkg/storage/storage_aws.go +++ b/packages/shared/pkg/storage/storage_aws.go @@ -233,7 +233,7 @@ func (o *awsObject) Put(ctx context.Context, data []byte, opts ...PutOption) err return nil } -func (o *awsObject) OpenRangeReader(ctx context.Context, off, length int64, frameTable *FrameTable) (io.ReadCloser, error) { +func (o *awsObject) OpenRangeReader(ctx context.Context, off, length int64, frameTable *FrameTable) (RangeReader, error) { if frameTable.IsCompressed() { return nil, errors.New("compressed reads are not supported on AWS") } @@ -253,7 +253,7 @@ func (o *awsObject) OpenRangeReader(ctx context.Context, off, length int64, fram return nil, fmt.Errorf("failed to create S3 range reader for %q: %w", o.path, err) } - return resp.Body, nil + return NewRangeReader(resp.Body), nil } func (o *awsObject) Size(ctx context.Context) (int64, error) { diff --git a/packages/shared/pkg/storage/storage_cache.go b/packages/shared/pkg/storage/storage_cache.go index 9f93559662..7d5838758c 100644 --- a/packages/shared/pkg/storage/storage_cache.go +++ b/packages/shared/pkg/storage/storage_cache.go @@ -151,6 +151,10 @@ func ignoreEOF(err error) error { // isCompleteRead reports whether a read of n bytes into a buffer of expected // size represents a valid, cacheable result. A read is complete when either // the full buffer was filled or io.EOF explains a non-empty short read (last chunk). +// +// Writeback callers pass err=nil: a streaming reader always ends in io.EOF +// regardless of whether the upstream was truncated, so the byte count is the +// only reliable signal that the captured bytes are safe to cache. func isCompleteRead(n, expected int, err error) bool { return n == expected || (n > 0 && errors.Is(err, io.EOF)) } diff --git a/packages/shared/pkg/storage/storage_cache_compressed_test.go b/packages/shared/pkg/storage/storage_cache_compressed_test.go index 8314abb0a1..d19b396c15 100644 --- a/packages/shared/pkg/storage/storage_cache_compressed_test.go +++ b/packages/shared/pkg/storage/storage_cache_compressed_test.go @@ -68,19 +68,16 @@ func TestDecompressingCacheReader(t *testing.T) { c := newTestCache(t) framePath := makeFrameFilename(c.path, Range{Offset: 0, Length: len(compressed)}) - rc, err := newDecompressingCacheReader( - io.NopCloser(bytes.NewReader(compressed)), - CompressionLZ4, - len(compressed), - &c, t.Context(), framePath, 0, - ) + capturing := newCaptureReader(bytesRangeReader(compressed), len(compressed), true, + c.compressedFrameWriteback(framePath, 0, len(compressed))) + rc, err := NewDecompressingReader(capturing, CompressionLZ4) require.NoError(t, err) got, err := io.ReadAll(rc) require.NoError(t, err) require.Equal(t, original, got) - require.NoError(t, rc.Close()) + mustClose(t, rc) c.wg.Wait() cached, err := os.ReadFile(framePath) @@ -104,12 +101,9 @@ func TestDecompressingCacheReader(t *testing.T) { compressedProd := lz4CompressProd(t, original) framePath := makeFrameFilename(c.path, Range{Offset: 0, Length: len(compressedProd)}) - rc, err := newDecompressingCacheReader( - io.NopCloser(bytes.NewReader(compressedProd)), - CompressionLZ4, - len(compressedProd), - &c, t.Context(), framePath, 0, - ) + capturing := newCaptureReader(bytesRangeReader(compressedProd), len(compressedProd), true, + c.compressedFrameWriteback(framePath, 0, len(compressedProd))) + rc, err := NewDecompressingReader(capturing, CompressionLZ4) require.NoError(t, err) out := make([]byte, len(original)) @@ -118,7 +112,8 @@ func TestDecompressingCacheReader(t *testing.T) { require.Equal(t, len(original), n) require.Equal(t, original, out) - require.NoError(t, rc.Close(), "writeback failure must not surface as a read error") + closeErr := rc.Close(t.Context()) + require.NoError(t, closeErr, "writeback failure must not surface as a read error") c.wg.Wait() _, err = os.Stat(framePath) @@ -131,19 +126,17 @@ func TestDecompressingCacheReader(t *testing.T) { c := newTestCache(t) framePath := makeFrameFilename(c.path, Range{Offset: 0, Length: len(compressed)}) - rc, err := newDecompressingCacheReader( - io.NopCloser(bytes.NewReader(compressed)), - CompressionLZ4, - len(compressed)+100, // wrong size - &c, t.Context(), framePath, 0, - ) + capturing := newCaptureReader(bytesRangeReader(compressed), len(compressed)+100, true, + c.compressedFrameWriteback(framePath, 0, len(compressed)+100)) // wrong expected size + rc, err := NewDecompressingReader(capturing, CompressionLZ4) require.NoError(t, err) got, err := io.ReadAll(rc) require.NoError(t, err) require.Equal(t, original, got, "decompressed data should be correct regardless") - require.NoError(t, rc.Close(), "writeback failure must not surface as a read error") + closeErr := rc.Close(t.Context()) + require.NoError(t, closeErr, "writeback failure must not surface as a read error") c.wg.Wait() diff --git a/packages/shared/pkg/storage/storage_cache_seekable.go b/packages/shared/pkg/storage/storage_cache_seekable.go index c24bf5fbef..2891047aaf 100644 --- a/packages/shared/pkg/storage/storage_cache_seekable.go +++ b/packages/shared/pkg/storage/storage_cache_seekable.go @@ -1,7 +1,6 @@ package storage import ( - "bytes" "context" "errors" "fmt" @@ -80,7 +79,7 @@ var ( _ RangeOpener = (*cachedSeekable)(nil) ) -func (c *cachedSeekable) OpenRangeReader(ctx context.Context, off int64, length int64, frameTable *FrameTable) (io.ReadCloser, error) { +func (c *cachedSeekable) OpenRangeReader(ctx context.Context, off int64, length int64, frameTable *FrameTable) (RangeReader, error) { compressed := frameTable.IsCompressed() ctx, span := c.tracer.Start(ctx, "read", trace.WithAttributes( @@ -89,27 +88,25 @@ func (c *cachedSeekable) OpenRangeReader(ctx context.Context, off int64, length attribute.Bool("compressed", compressed), )) + var rc RangeReader + var err error if compressed { - rc, err := c.openReaderCompressed(ctx, off, frameTable) - if err != nil { - recordError(span, err) - span.End() - - return nil, err - } - - rc = withSpan(rc, span) - - return rc, nil + rc, err = c.openReaderCompressed(ctx, off, frameTable) + } else if err = c.validateReadParams(length, off); err == nil { + rc, err = c.openReaderUncompressed(ctx, off, length) } - if err := c.validateReadParams(length, off); err != nil { + if err != nil { recordError(span, err) span.End() return nil, err } + return newObservableReader(rc, nil, span), nil +} + +func (c *cachedSeekable) openReaderUncompressed(ctx context.Context, off, length int64) (RangeReader, error) { timer := cacheSlabReadTimerFactory.Begin( attribute.String(nfsCacheOperationAttr, nfsCacheOperationAttrReadAt), attribute.Bool("compressed", false), @@ -122,10 +119,7 @@ func (c *cachedSeekable) OpenRangeReader(ctx context.Context, off int64, length recordCacheRead(ctx, true, length, cacheTypeSeekable, cacheOpOpenRangeReader) timer.Success(ctx, length) - rc := io.ReadCloser(&fsRangeReadCloser{Reader: io.NewSectionReader(fp, 0, length), file: fp}) - rc = withSpan(rc, span) - - return withNFSGauge(ctx, rc), nil + return withNFSGauge(ctx, newSectionReader(fp, 0, length)), nil } if !os.IsNotExist(err) { @@ -136,122 +130,58 @@ func (c *cachedSeekable) OpenRangeReader(ctx context.Context, off int64, length rc, err := c.inner.OpenRangeReader(ctx, off, length, nil) if err != nil { - recordError(span, err) - span.End() - return nil, fmt.Errorf("failed to open inner range reader: %w", err) } recordCacheRead(ctx, false, length, cacheTypeSeekable, cacheOpOpenRangeReader) if !skipCacheWriteback(ctx) { - rc = newCacheWriteThroughReader(rc, c, ctx, off, length, chunkPath) + rc = newCaptureReader(rc, int(length), false, + c.uncompressedChunkWriteback(chunkPath, off, length)) } - rc = withSpan(rc, span) - return rc, nil } -// withSpan wraps a reader with an OTEL span that ends on Close. -func withSpan(rc io.ReadCloser, span trace.Span) io.ReadCloser { - return &spanReadCloser{inner: rc, span: span} -} - -type spanReadCloser struct { - inner io.ReadCloser - span trace.Span -} - -func (r *spanReadCloser) Read(p []byte) (int, error) { - return r.inner.Read(p) -} - -func (r *spanReadCloser) Close() error { - err := r.inner.Close() - recordError(r.span, err) - r.span.End() - - return err -} - // nfsGaugeReadCloser wraps a reader and decrements the NFS concurrent reads // gauge on Close. type nfsGaugeReadCloser struct { - io.ReadCloser - - ctx context.Context //nolint:containedctx // needed for gauge decrement in Close + RangeReader } -func (r *nfsGaugeReadCloser) Close() error { - nfsCacheConcurrentReads.Add(r.ctx, -1) +func (r *nfsGaugeReadCloser) Close(ctx context.Context) error { + nfsCacheConcurrentReads.Add(ctx, -1) - return r.ReadCloser.Close() + return r.RangeReader.Close(ctx) } -func withNFSGauge(ctx context.Context, rc io.ReadCloser) io.ReadCloser { +func withNFSGauge(ctx context.Context, rc RangeReader) RangeReader { nfsCacheConcurrentReads.Add(ctx, 1) - return &nfsGaugeReadCloser{ReadCloser: rc, ctx: ctx} -} - -// newCacheWriteThroughReader wraps a reader, buffering all data read through it. -// On Close, it asynchronously writes the buffered data to the NFS cache only -// if the total bytes read match the expected length (to avoid caching truncated data). -func newCacheWriteThroughReader(inner io.ReadCloser, cache *cachedSeekable, ctx context.Context, off, expectedLen int64, chunkPath string) io.ReadCloser { - return &cacheWriteThroughReader{ - inner: inner, - buf: bytes.NewBuffer(make([]byte, 0, expectedLen)), - cache: cache, - ctx: ctx, - off: off, - expectedLen: expectedLen, - chunkPath: chunkPath, - } -} - -type cacheWriteThroughReader struct { - inner io.ReadCloser - buf *bytes.Buffer - cache *cachedSeekable - ctx context.Context //nolint:containedctx // needed for async cache write-back in Close - off int64 - expectedLen int64 - chunkPath string -} - -func (r *cacheWriteThroughReader) Read(p []byte) (int, error) { - n, err := r.inner.Read(p) - if n > 0 { - r.buf.Write(p[:n]) - } - - return n, err + return &nfsGaugeReadCloser{RangeReader: rc} } -func (r *cacheWriteThroughReader) Close() error { - closeErr := r.inner.Close() - - // Only cache when the total bytes read match the expected length. - // Unlike ReadAt where io.EOF can justify a short read (last chunk), - // a streaming reader always ends with EOF regardless of whether the - // data was truncated, so the byte count is the only reliable check. - if isCompleteRead(r.buf.Len(), int(r.expectedLen), nil) { - data := make([]byte, r.buf.Len()) - copy(data, r.buf.Bytes()) +// uncompressedChunkWriteback returns a captureReader callback that persists +// the captured chunk to the NFS cache in a detached goroutine. Best-effort: +// a short capture (e.g. upstream truncation) is dropped silently — a streaming +// reader always ends in EOF, so byte count is the only reliable signal. +func (c *cachedSeekable) uncompressedChunkWriteback(chunkPath string, off, expectedLen int64) func(context.Context, []byte) { + return func(ctx context.Context, captured []byte) { + if !isCompleteRead(len(captured), int(expectedLen), nil) { + return + } - r.cache.goCtx(r.ctx, func(ctx context.Context) { - ctx, span := r.cache.tracer.Start(ctx, "write range reader chunk back to cache") + c.goCtx(ctx, func(ctx context.Context) { + ctx, span := c.tracer.Start(ctx, "write range reader chunk back to cache") defer span.End() - if err := r.cache.writeToCache(ctx, r.off, r.chunkPath, data); err != nil { + err := c.writeToCache(ctx, off, chunkPath, captured) + if err != nil { recordError(span, err) recordCacheWriteError(ctx, cacheTypeSeekable, cacheOpOpenRangeReader, err) } }) } - - return closeErr } func (c *cachedSeekable) Size(ctx context.Context) (n int64, e error) { @@ -571,9 +501,8 @@ func (c *cachedSeekable) writeChunkFromFile(ctx context.Context, offset int64, i } defer utils.Cleanup(ctx, "failed to close file", output.Close) - offsetReader := newOffsetReader(input, offset) - count, err := io.CopyN(output, offsetReader, c.chunkSize) - if ignoreEOF(err) != nil { + count, err := io.Copy(output, io.NewSectionReader(input, offset, c.chunkSize)) + if err != nil { writeTimer.Failure(ctx, count) safelyRemoveFile(ctx, chunkPath) diff --git a/packages/shared/pkg/storage/storage_cache_seekable_compressed.go b/packages/shared/pkg/storage/storage_cache_seekable_compressed.go index 4886b449d4..2e8077c4bb 100644 --- a/packages/shared/pkg/storage/storage_cache_seekable_compressed.go +++ b/packages/shared/pkg/storage/storage_cache_seekable_compressed.go @@ -1,10 +1,8 @@ package storage import ( - "bytes" "context" "fmt" - "io" "os" "go.opentelemetry.io/otel/attribute" @@ -19,13 +17,14 @@ var compressedCacheReadAttrs = []attribute.KeyValue{ // openReaderCompressed handles the compressed cache path for OpenRangeReader. // NFS stores compressed frames (.frm); on hit we decompress, on miss we fetch // raw compressed bytes and tee them to NFS on Close. -func (c *cachedSeekable) openReaderCompressed(ctx context.Context, offsetU int64, frameTable *FrameTable) (io.ReadCloser, error) { +func (c *cachedSeekable) openReaderCompressed(ctx context.Context, offsetU int64, frameTable *FrameTable) (RangeReader, error) { r, err := frameTable.LocateCompressed(offsetU) if err != nil { return nil, fmt.Errorf("frame lookup for offset %d: %w", offsetU, err) } path := makeFrameFilename(c.path, r) + ct := frameTable.CompressionType() timer := cacheSlabReadTimerFactory.Begin(compressedCacheReadAttrs...) @@ -37,14 +36,14 @@ func (c *cachedSeekable) openReaderCompressed(ctx context.Context, offsetU int64 recordCacheRead(ctx, true, int64(r.Length), cacheTypeSeekable, cacheOpOpenRangeReader) timer.Success(ctx, int64(r.Length)) - decompressed, err := newDecompressingReadCloser(f, frameTable.CompressionType()) + dec, err := NewDecompressingReader(NewRangeReader(f), ct) if err != nil { f.Close() return nil, fmt.Errorf("decompress cached frame: %w", err) } - return withNFSGauge(ctx, decompressed), nil + return withNFSGauge(ctx, dec), nil case statErr == nil: // Confirmed size mismatch: drop the file so the miss path rewrites it. f.Close() @@ -70,111 +69,46 @@ func (c *cachedSeekable) openReaderCompressed(ctx context.Context, offsetU int64 recordCacheRead(ctx, false, int64(r.Length), cacheTypeSeekable, cacheOpOpenRangeReader) - rc, err := newDecompressingCacheReader(raw, frameTable.CompressionType(), r.Length, c, ctx, path, offsetU) - if err != nil { - raw.Close() - - return nil, fmt.Errorf("create decompressor: %w", err) + in := raw + if !skipCacheWriteback(ctx) { + in = newCaptureReader(raw, r.Length, true, + c.compressedFrameWriteback(path, offsetU, r.Length)) } - return rc, nil -} - -// newDecompressingCacheReader creates a reader that decompresses on Read and -// writes the accumulated compressed bytes to the NFS cache on Close. -func newDecompressingCacheReader( - raw io.ReadCloser, - ct CompressionType, - expectedSize int, - cache *cachedSeekable, - ctx context.Context, //nolint:revive // ctx after other params for readability at call site - framePath string, - offset int64, -) (io.ReadCloser, error) { - var compressedBuf bytes.Buffer - compressedBuf.Grow(expectedSize) - - tee := io.TeeReader(raw, &compressedBuf) - - dec, err := NewDecompressingReader(tee, ct) + dec, err := NewDecompressingReader(in, ct) if err != nil { - return nil, err - } + raw.Close(ctx) - return &decompressingCacheReader{ - decompressor: dec, - raw: raw, - compressedBuf: &compressedBuf, - expectedSize: expectedSize, - cache: cache, - ctx: ctx, - framePath: framePath, - offset: offset, - }, nil -} - -type decompressingCacheReader struct { - decompressor io.ReadCloser // decompresses on Read - raw io.ReadCloser // underlying compressed stream (must be closed) - compressedBuf *bytes.Buffer - expectedSize int - cache *cachedSeekable - ctx context.Context //nolint:containedctx // needed for async cache write-back in Close - framePath string - offset int64 -} + return nil, fmt.Errorf("create decompressor: %w", err) + } -func (r *decompressingCacheReader) Read(p []byte) (int, error) { - return r.decompressor.Read(p) + return dec, nil } -func (r *decompressingCacheReader) Close() error { - // Drive the decompressor to EOF before closing it. With io.ReadFull bounded - // by the uncompressed size, an LZ4 frame written with BlockChecksum=true / - // Checksum=false leaves the 4-byte EndMark unread — the next Read on the - // decoder pulls the EndMark (block-size = 0 → io.EOF) from raw through the - // tee, populating compressedBuf with the full encoded frame for cache writeback. - _, _ = io.Copy(io.Discard, r.decompressor) - - decErr := r.decompressor.Close() - rawErr := r.raw.Close() - - if decErr != nil { - return decErr - } - if rawErr != nil { - return rawErr - } - - got := r.compressedBuf.Len() - if skipCacheWriteback(r.ctx) { - return nil - } +// compressedFrameWriteback returns a captureReader callback that +// persists the captured frame to the NFS cache in a detached goroutine. +// Best-effort: a short capture is logged and skipped — the caller already +// got valid decompressed bytes. +func (c *cachedSeekable) compressedFrameWriteback(framePath string, offset int64, expectedSize int) func(context.Context, []byte) { + return func(ctx context.Context, frame []byte) { + if !isCompleteRead(len(frame), expectedSize, nil) { + recordCacheWriteError(ctx, cacheTypeSeekable, cacheOpOpenRangeReader, + fmt.Errorf("compressed frame cache writeback short: got %d bytes, expected %d for %s", len(frame), expectedSize, framePath)) + + return + } - // Cache writeback is best-effort. After draining above, a remaining shortfall - // implies upstream truncation — log/metric and skip writeback rather than - // poison the read (the caller already received valid decompressed bytes). - if !isCompleteRead(got, r.expectedSize, nil) { - recordCacheWriteError(r.ctx, cacheTypeSeekable, cacheOpOpenRangeReader, - fmt.Errorf("compressed frame cache writeback short: got %d bytes, expected %d for %s", got, r.expectedSize, r.framePath)) + c.goCtx(ctx, func(ctx context.Context) { + ctx, span := c.tracer.Start(ctx, "write compressed frame back to cache") + defer span.End() - return nil + err := c.writeToCache(ctx, offset, framePath, frame) + if err != nil { + recordError(span, err) + recordCacheWriteError(ctx, cacheTypeSeekable, cacheOpOpenRangeReader, err) + } + }) } - - data := r.compressedBuf.Bytes() - r.compressedBuf = nil - - r.cache.goCtx(r.ctx, func(ctx context.Context) { - ctx, span := r.cache.tracer.Start(ctx, "write compressed frame back to cache") - defer span.End() - - if err := r.cache.writeToCache(ctx, r.offset, r.framePath, data); err != nil { - recordError(span, err) - recordCacheWriteError(ctx, cacheTypeSeekable, cacheOpOpenRangeReader, err) - } - }) - - return nil } // makeFrameFilename returns the NFS cache path for a compressed frame. diff --git a/packages/shared/pkg/storage/storage_cache_seekable_test.go b/packages/shared/pkg/storage/storage_cache_seekable_test.go index 26bac12f9a..81f7364c0a 100644 --- a/packages/shared/pkg/storage/storage_cache_seekable_test.go +++ b/packages/shared/pkg/storage/storage_cache_seekable_test.go @@ -14,6 +14,17 @@ import ( "github.com/stretchr/testify/require" ) +// mustClose closes a RangeReader and asserts no error. +func mustClose(t *testing.T, rc RangeReader) { + t.Helper() + require.NoError(t, rc.Close(t.Context())) +} + +// bytesRangeReader wraps an in-memory byte slice as a RangeReader for tests. +func bytesRangeReader(b []byte) RangeReader { + return NewRangeReader(io.NopCloser(bytes.NewReader(b))) +} + // testReadAt emulates the removed cachedSeekable.ReadAt via OpenRangeReader. // This preserves the base test structure after ReadAt was removed from the Seekable interface. func testReadAt(ctx context.Context, c *cachedSeekable, buff []byte, off int64) (int, error) { @@ -24,7 +35,7 @@ func testReadAt(ctx context.Context, c *cachedSeekable, buff []byte, off int64) n, err := io.ReadFull(rc, buff) - closeErr := rc.Close() + closeErr := rc.Close(ctx) if errors.Is(err, io.ErrUnexpectedEOF) { err = io.EOF } @@ -181,10 +192,10 @@ func TestCachedFileObjectProvider_WriteTo(t *testing.T) { inner.EXPECT(). OpenRangeReader(mock.Anything, mock.Anything, mock.Anything, (*FrameTable)(nil)). - RunAndReturn(func(_ context.Context, off int64, length int64, _ *FrameTable) (io.ReadCloser, error) { + RunAndReturn(func(_ context.Context, off int64, length int64, _ *FrameTable) (RangeReader, error) { end := min(int(off)+int(length), len(fakeData)) - return io.NopCloser(bytes.NewReader(fakeData[off:end])), nil + return NewRangeReader(io.NopCloser(bytes.NewReader(fakeData[off:end]))), nil }) tempDir := t.TempDir() @@ -316,7 +327,7 @@ func TestCachedSeekableObjectProvider_ReadAt(t *testing.T) { inner := NewMockSeekable(t) inner.EXPECT(). OpenRangeReader(mock.Anything, mock.Anything, mock.Anything, (*FrameTable)(nil)). - Return(io.NopCloser(bytes.NewReader(nil)), nil) + Return(NewRangeReader(io.NopCloser(bytes.NewReader(nil))), nil) c := cachedSeekable{ path: tempDir, @@ -345,7 +356,7 @@ func TestCachedSeekableObjectProvider_ReadAt(t *testing.T) { inner := NewMockSeekable(t) inner.EXPECT(). OpenRangeReader(mock.Anything, mock.Anything, mock.Anything, (*FrameTable)(nil)). - Return(io.NopCloser(bytes.NewReader(data)), nil) + Return(NewRangeReader(io.NopCloser(bytes.NewReader(data))), nil) c := cachedSeekable{ path: tempDir, @@ -407,7 +418,7 @@ func TestCachedSeekable_ReadAt_PreservesEOF(t *testing.T) { inner := NewMockSeekable(t) inner.EXPECT(). OpenRangeReader(mock.Anything, mock.Anything, mock.Anything, (*FrameTable)(nil)). - Return(io.NopCloser(bytes.NewReader([]byte{1, 2, 3})), nil) + Return(NewRangeReader(io.NopCloser(bytes.NewReader([]byte{1, 2, 3}))), nil) c := cachedSeekable{ path: tempDir, @@ -431,7 +442,7 @@ func TestCachedSeekable_ReadAt_PreservesEOF(t *testing.T) { inner := NewMockSeekable(t) inner.EXPECT(). OpenRangeReader(mock.Anything, mock.Anything, mock.Anything, (*FrameTable)(nil)). - Return(io.NopCloser(bytes.NewReader([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10})), nil) + Return(NewRangeReader(io.NopCloser(bytes.NewReader([]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}))), nil) c := cachedSeekable{ path: tempDir, @@ -457,8 +468,8 @@ func TestCachedSeekable_ReadAt_SkipCacheWriteback(t *testing.T) { inner := NewMockSeekable(t) inner.EXPECT(). OpenRangeReader(mock.Anything, mock.Anything, mock.Anything, (*FrameTable)(nil)). - RunAndReturn(func(_ context.Context, _ int64, _ int64, _ *FrameTable) (io.ReadCloser, error) { - return io.NopCloser(bytes.NewReader(data)), nil + RunAndReturn(func(_ context.Context, _ int64, _ int64, _ *FrameTable) (RangeReader, error) { + return NewRangeReader(io.NopCloser(bytes.NewReader(data))), nil }) c := cachedSeekable{ @@ -493,7 +504,7 @@ func TestCachedSeekable_OpenRangeReader(t *testing.T) { inner := NewMockSeekable(t) inner.EXPECT(). OpenRangeReader(mock.Anything, int64(0), int64(len(data)), (*FrameTable)(nil)). - Return(io.NopCloser(bytes.NewReader(data)), nil). + Return(NewRangeReader(io.NopCloser(bytes.NewReader(data))), nil). Once() c := cachedSeekable{ @@ -510,7 +521,7 @@ func TestCachedSeekable_OpenRangeReader(t *testing.T) { got, err := io.ReadAll(rc) require.NoError(t, err) assert.Equal(t, data, got) - require.NoError(t, rc.Close()) + require.NoError(t, rc.Close(t.Context())) c.wg.Wait() @@ -522,7 +533,7 @@ func TestCachedSeekable_OpenRangeReader(t *testing.T) { got2, err := io.ReadAll(rc2) require.NoError(t, err) assert.Equal(t, data, got2) - require.NoError(t, rc2.Close()) + require.NoError(t, rc2.Close(t.Context())) }) t.Run("skip cache writeback returns inner directly", func(t *testing.T) { @@ -534,8 +545,8 @@ func TestCachedSeekable_OpenRangeReader(t *testing.T) { inner := NewMockSeekable(t) inner.EXPECT(). OpenRangeReader(mock.Anything, int64(0), int64(len(data)), (*FrameTable)(nil)). - RunAndReturn(func(_ context.Context, _ int64, _ int64, _ *FrameTable) (io.ReadCloser, error) { - return io.NopCloser(bytes.NewReader(data)), nil + RunAndReturn(func(_ context.Context, _ int64, _ int64, _ *FrameTable) (RangeReader, error) { + return NewRangeReader(io.NopCloser(bytes.NewReader(data))), nil }). Times(2) @@ -554,7 +565,7 @@ func TestCachedSeekable_OpenRangeReader(t *testing.T) { got, err := io.ReadAll(rc) require.NoError(t, err) assert.Equal(t, data, got) - require.NoError(t, rc.Close()) + require.NoError(t, rc.Close(ctx)) c.wg.Wait() @@ -569,7 +580,7 @@ func TestCachedSeekable_OpenRangeReader(t *testing.T) { got2, err := io.ReadAll(rc2) require.NoError(t, err) assert.Equal(t, data, got2) - require.NoError(t, rc2.Close()) + require.NoError(t, rc2.Close(ctx)) }) t.Run("truncated inner read does not populate cache", func(t *testing.T) { @@ -580,7 +591,7 @@ func TestCachedSeekable_OpenRangeReader(t *testing.T) { inner := NewMockSeekable(t) inner.EXPECT(). OpenRangeReader(mock.Anything, int64(0), int64(5), (*FrameTable)(nil)). - Return(io.NopCloser(bytes.NewReader([]byte{0xAA, 0xBB})), nil) + Return(NewRangeReader(io.NopCloser(bytes.NewReader([]byte{0xAA, 0xBB}))), nil) c := cachedSeekable{ path: tempDir, @@ -595,7 +606,7 @@ func TestCachedSeekable_OpenRangeReader(t *testing.T) { got, err := io.ReadAll(rc) require.NoError(t, err) assert.Equal(t, []byte{0xAA, 0xBB}, got) - require.NoError(t, rc.Close()) + require.NoError(t, rc.Close(t.Context())) c.wg.Wait() @@ -681,23 +692,16 @@ func TestCacheWriteThroughReader(t *testing.T) { c := newTestCache(t) data := []byte("hello") - inner := io.NopCloser(bytes.NewReader(data)) - - r := &cacheWriteThroughReader{ - inner: inner, - buf: bytes.NewBuffer(make([]byte, 0, len(data))), - cache: &c, - ctx: t.Context(), - off: 0, - expectedLen: int64(len(data)), - chunkPath: c.makeChunkFilename(0), - } + inner := NewRangeReader(io.NopCloser(bytes.NewReader(data))) + + r := newCaptureReader(inner, len(data), false, + c.uncompressedChunkWriteback(c.makeChunkFilename(0), 0, int64(len(data)))) got, err := io.ReadAll(r) require.NoError(t, err) assert.Equal(t, data, got) - require.NoError(t, r.Close()) + require.NoError(t, r.Close(t.Context())) c.wg.Wait() cached, err := os.ReadFile(c.makeChunkFilename(0)) @@ -712,23 +716,16 @@ func TestCacheWriteThroughReader(t *testing.T) { // Inner has only 2 bytes but expectedLen is 5. The reader is // fully consumed (EOF is reached), yet the total doesn't match // the expected length so it must not be cached. - inner := io.NopCloser(bytes.NewReader([]byte{0xAA, 0xBB})) - - r := &cacheWriteThroughReader{ - inner: inner, - buf: bytes.NewBuffer(make([]byte, 0, 5)), - cache: &c, - ctx: t.Context(), - off: 0, - expectedLen: 5, - chunkPath: c.makeChunkFilename(0), - } + inner := NewRangeReader(io.NopCloser(bytes.NewReader([]byte{0xAA, 0xBB}))) + + r := newCaptureReader(inner, 5, false, + c.uncompressedChunkWriteback(c.makeChunkFilename(0), 0, 5)) got, err := io.ReadAll(r) require.NoError(t, err) assert.Equal(t, []byte{0xAA, 0xBB}, got) - require.NoError(t, r.Close()) + require.NoError(t, r.Close(t.Context())) c.wg.Wait() _, err = os.Stat(c.makeChunkFilename(0)) @@ -740,17 +737,10 @@ func TestCacheWriteThroughReader(t *testing.T) { c := newTestCache(t) data := []byte("hello") - inner := io.NopCloser(bytes.NewReader(data)) - - r := &cacheWriteThroughReader{ - inner: inner, - buf: bytes.NewBuffer(make([]byte, 0, len(data))), - cache: &c, - ctx: t.Context(), - off: 0, - expectedLen: int64(len(data)), - chunkPath: c.makeChunkFilename(0), - } + inner := NewRangeReader(io.NopCloser(bytes.NewReader(data))) + + r := newCaptureReader(inner, len(data), false, + c.uncompressedChunkWriteback(c.makeChunkFilename(0), 0, int64(len(data)))) // Read only 2 of 5 bytes, then close without reaching EOF. buf := make([]byte, 2) @@ -758,7 +748,7 @@ func TestCacheWriteThroughReader(t *testing.T) { require.NoError(t, err) assert.Equal(t, 2, n) - require.NoError(t, r.Close()) + require.NoError(t, r.Close(t.Context())) c.wg.Wait() _, err = os.Stat(c.makeChunkFilename(0)) diff --git a/packages/shared/pkg/storage/storage_fs.go b/packages/shared/pkg/storage/storage_fs.go index ea064ef46a..fcaac07366 100644 --- a/packages/shared/pkg/storage/storage_fs.go +++ b/packages/shared/pkg/storage/storage_fs.go @@ -39,16 +39,6 @@ var ( _ RangeOpener = (*fsObject)(nil) ) -type fsRangeReadCloser struct { - io.Reader - - file *os.File -} - -func (r *fsRangeReadCloser) Close() error { - return r.file.Close() -} - func newFileSystemStorage(cfg StorageConfig) *fsStorage { return &fsStorage{ basePath: cfg.GetLocalBasePath(), @@ -206,16 +196,13 @@ func (o *fsObject) storeFileCompressed(ctx context.Context, localPath string, cf return ft, checksum, nil } -func (o *fsObject) openRangeReader(_ context.Context, off, length int64) (io.ReadCloser, error) { +func (o *fsObject) openRangeReader(_ context.Context, off, length int64) (RangeReader, error) { f, err := o.getHandle(true) if err != nil { return nil, err } - return &fsRangeReadCloser{ - Reader: io.NewSectionReader(f, off, length), - file: f, - }, nil + return newSectionReader(f, off, length), nil } func (o *fsObject) Exists(_ context.Context) (bool, error) { @@ -316,7 +303,7 @@ func (u *fsPartUploader) Complete(_ context.Context) error { return os.WriteFile(u.fullPath, u.Assemble(), 0o644) } -func (o *fsObject) OpenRangeReader(ctx context.Context, offsetU int64, length int64, frameTable *FrameTable) (io.ReadCloser, error) { +func (o *fsObject) OpenRangeReader(ctx context.Context, offsetU int64, length int64, frameTable *FrameTable) (RangeReader, error) { if frameTable.IsCompressed() { r, err := frameTable.LocateCompressed(offsetU) if err != nil { @@ -328,14 +315,14 @@ func (o *fsObject) OpenRangeReader(ctx context.Context, offsetU int64, length in return nil, err } - decompressed, err := newDecompressingReadCloser(raw, frameTable.CompressionType()) + dec, err := NewDecompressingReader(raw, frameTable.CompressionType()) if err != nil { - raw.Close() + raw.Close(ctx) return nil, err } - return decompressed, nil + return dec, nil } return o.openRangeReader(ctx, offsetU, length) diff --git a/packages/shared/pkg/storage/storage_google.go b/packages/shared/pkg/storage/storage_google.go index 9f1b81c1d1..e1fe7c6dc9 100644 --- a/packages/shared/pkg/storage/storage_google.go +++ b/packages/shared/pkg/storage/storage_google.go @@ -264,7 +264,7 @@ func (o *gcpObject) Size(ctx context.Context) (int64, error) { return attrs.Size, nil } -func (o *gcpObject) openRangeReader(ctx context.Context, off, length int64) (io.ReadCloser, error) { +func (o *gcpObject) openRangeReader(ctx context.Context, off, length int64) (RangeReader, error) { readCtx, cancel := context.WithCancel(ctx) openTimer := time.AfterFunc(googleReadTimeout, cancel) @@ -285,6 +285,8 @@ func (o *gcpObject) openRangeReader(ctx context.Context, off, length int64) (io. return nil, fmt.Errorf("failed to create GCS range reader for %q at %d+%d: %w", o.path, off, length, err) } + gcsConcurrentReads.Add(ctx, 1) + return &idleTimeoutReader{ ReadCloser: reader, cancel: cancel, @@ -314,9 +316,10 @@ func (r *idleTimeoutReader) Read(p []byte) (int, error) { return n, err } -func (r *idleTimeoutReader) Close() error { +func (r *idleTimeoutReader) Close(ctx context.Context) error { r.timer.Stop() defer r.cancel() + gcsConcurrentReads.Add(ctx, -1) return r.ReadCloser.Close() } @@ -599,7 +602,7 @@ func parseServiceAccountBase64(serviceAccount string) (*gcpServiceToken, error) return &sa, nil } -func (o *gcpObject) OpenRangeReader(ctx context.Context, offsetU int64, length int64, frameTable *FrameTable) (io.ReadCloser, error) { +func (o *gcpObject) OpenRangeReader(ctx context.Context, offsetU int64, length int64, frameTable *FrameTable) (RangeReader, error) { timer := googleReadTimerFactory.Begin(attribute.String(gcsOperationAttr, gcsOperationAttrReadAt)) if !frameTable.IsCompressed() { @@ -610,9 +613,7 @@ func (o *gcpObject) OpenRangeReader(ctx context.Context, offsetU int64, length i return nil, err } - gcsConcurrentReads.Add(ctx, 1) - - return &timedReadCloser{inner: rc, timer: timer, ctx: ctx}, nil + return newObservableReader(rc, timer, nil), nil } r, err := frameTable.LocateCompressed(offsetU) @@ -629,52 +630,15 @@ func (o *gcpObject) OpenRangeReader(ctx context.Context, offsetU int64, length i return nil, err } - decompressed, err := newDecompressingReadCloser(raw, frameTable.CompressionType()) + dec, err := NewDecompressingReader(raw, frameTable.CompressionType()) if err != nil { - raw.Close() + raw.Close(ctx) timer.Failure(ctx, 0) return nil, err } - gcsConcurrentReads.Add(ctx, 1) - - return &timedReadCloser{inner: decompressed, timer: timer, ctx: ctx}, nil -} - -// timedReadCloser wraps a reader with OTEL timer metrics. -// Close records success (with total bytes read) or failure on the timer. -type timedReadCloser struct { - inner io.ReadCloser - timer *telemetry.Stopwatch - ctx context.Context //nolint:containedctx // needed for timer recording in Close - bytesRead int64 - closeErr error -} - -func (r *timedReadCloser) Read(p []byte) (int, error) { - n, err := r.inner.Read(p) - r.bytesRead += int64(n) - - if err != nil && err != io.EOF { - r.closeErr = err - } - - return n, err -} - -func (r *timedReadCloser) Close() error { - gcsConcurrentReads.Add(r.ctx, -1) - - err := r.inner.Close() - - if r.closeErr != nil || err != nil { - r.timer.Failure(r.ctx, r.bytesRead) - } else { - r.timer.Success(r.ctx, r.bytesRead) - } - - return err + return newObservableReader(dec, timer, nil), nil } func isResourceExhausted(err error) bool { diff --git a/tests/integration/internal/tests/api/sandboxes/sandbox_rapid_pause_resume_test.go b/tests/integration/internal/tests/api/sandboxes/sandbox_rapid_pause_resume_test.go index 474477dcd3..3b36793ff7 100644 --- a/tests/integration/internal/tests/api/sandboxes/sandbox_rapid_pause_resume_test.go +++ b/tests/integration/internal/tests/api/sandboxes/sandbox_rapid_pause_resume_test.go @@ -149,7 +149,7 @@ func verifyChecksum(t *testing.T, ctx context.Context, persistence storage.Stora rc, err := obj.OpenRangeReader(ctx, 0, bd.Size, bd.FrameData) require.NoErrorf(t, err, "%s/%s: open range reader", node.name, fileName) - defer rc.Close() + defer rc.Close(context.WithoutCancel(ctx)) hasher := sha256.New() n, err := io.Copy(hasher, rc)