Skip to content
Merged
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
20 changes: 15 additions & 5 deletions packages/orchestrator/pkg/sandbox/block/streaming_chunk.go
Original file line number Diff line number Diff line change
Expand Up @@ -290,13 +290,18 @@ func (c *Chunker) progressiveFetch(ctx context.Context, s *fetchSession, mmapSli
var closeErr error
res.stats, closeErr = reader.Close(context.WithoutCancel(ctx))

// A compressed frame's CRC is only verified once its footer is consumed
// on Close, so a Close error is a verification failure that must fail
// the fetch (don't cache or release a corrupt frame). A read error, if
// any, takes precedence.
if err == nil {
err = closeErr
}

ct := ft.CompressionType()
attrs := storage.OKAttrs(c.objType, source, ct)
switch {
case err != nil:
if err != nil {
attrs = storage.ErrAttrs(c.objType, source, ct, err)
case closeErr != nil:
attrs = storage.ErrAttrs(c.objType, source, ct, closeErr)
}

var readDur time.Duration
Expand All @@ -318,9 +323,14 @@ func (c *Chunker) progressiveFetch(ctx context.Context, s *fetchSession, mmapSli
n, readErr := io.ReadFull(reader, mmapSlice[totalRead:readEnd])
totalRead += int64(n)

if n > 0 {
if n > 0 && !ft.IsCompressed() {
// Dirty marking is deferred to runFetch after the full chunk is fetched.
// With coarse dirty granularity, marking here would expose partially-written data.
//
// Compressed chunks are a single frame whose CRC is only verified
// once the footer is consumed on Close. Releasing waiters here
// would hand out bytes that a later CRC failure proves corrupt, so
// their release is deferred to setDone after verification.
s.advance(totalRead)
}

Expand Down
51 changes: 46 additions & 5 deletions packages/orchestrator/pkg/sandbox/block/streaming_chunk_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -50,10 +50,12 @@ func makeTestData(size int) []byte {
// fakeSeekable implements storage.Seekable backed by in-memory data.
// When ctrl is non-nil, reads are gated through its channels for concurrency tests.
type fakeSeekable struct {
data []byte
failAfter int64 // >0: truncate reads at this offset; 0 = disabled
fetchCount atomic.Int64
ctrl *testControl // nil = ungated immediate reads
data []byte
failAfter int64 // >0: truncate reads at this offset; 0 = disabled
corrupt bool // enable corruptByte injection
corruptByte int64 // absolute offset to XOR 0xFF in served bytes (requires corrupt=true)
fetchCount atomic.Int64
ctrl *testControl // nil = ungated immediate reads
}

var _ storage.Seekable = (*fakeSeekable)(nil)
Expand Down Expand Up @@ -126,7 +128,13 @@ func (s *fakeSeekable) OpenRangeReader(_ context.Context, offsetU int64, length
end = min(end, s.failAfter)
}

r := io.Reader(bytes.NewReader(s.data[fetchOff:end]))
served := s.data[fetchOff:end]
if s.corrupt && s.corruptByte >= fetchOff && s.corruptByte < end {
served = bytes.Clone(served)
served[s.corruptByte-fetchOff] ^= 0xFF
}

r := io.Reader(bytes.NewReader(served))
if frameTable.IsCompressed() {
dec, err := storage.NewDecompressReader(storage.NewRangeReader(io.NopCloser(r)), frameTable.CompressionType(), storage.UnknownSource, storage.UnknownSeekableObjectType)

Expand Down Expand Up @@ -645,3 +653,36 @@ func (r *controlledReader) Close(context.Context) (*storage.ReadStats, error) {

return nil, nil
}

// TestChunker_CorruptCompressedFrameNotServedOrCached verifies the integrity
// gap fix: a compressed frame whose CRC only fails at the footer (content
// decodes to plausible bytes) must not be released to waiters or marked
// cached. Without the fix, an exact-size read never triggers the codec's
// footer CRC, the Close error was ignored, and the frame was advanced to
// waiters and cached. Corrupts the final byte of the first compressed frame.
func TestChunker_CorruptCompressedFrameNotServedOrCached(t *testing.T) {
t.Parallel()

data := makeTestData(testFileSize)
ft, file := makeCompressedTestData(t, data)

// Corrupt the last byte of the first frame's compressed range (footer).
r, err := ft.LocateCompressed(0)
require.NoError(t, err)
file.corrupt = true
file.corruptByte = r.Offset + int64(r.Length) - 1

chunker := newTestChunker(t, int64(len(data)))
defer chunker.Close()

_, err = chunker.Slice(t.Context(), 0, testBlockSize, file, ft)
require.Error(t, err, "corrupt compressed frame must surface an error, not serve unverified bytes")
require.False(t, chunker.IsCached(t.Context(), 0, testBlockSize),
"corrupt compressed frame must not be marked cached")

// A later read of a clean frame still works (chunker remains usable).
lastOff := int64(testFileSize) - testBlockSize
slice, err := chunker.Slice(t.Context(), lastOff, testBlockSize, file, ft)
require.NoError(t, err)
require.Equal(t, data[lastOff:], slice)
}
15 changes: 15 additions & 0 deletions packages/shared/pkg/storage/compress_decode.go
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,17 @@ func (r *decompressReader) Read(p []byte) (int, error) {
}

func (r *decompressReader) Close(ctx context.Context) (*ReadStats, error) {
// Drain any remaining decoded bytes so the codec consumes the frame footer
// and surfaces a CRC/truncation error. A caller reading only the exact
// uncompressed frame size never pulls the footer through, so zstd would
// otherwise report success on a footer-corrupted or truncated frame.
// Bounded: the source is a single frame.
if r.readErr == nil {
if _, err := io.Copy(io.Discard, r.meteredOut); err != nil && !errors.Is(err, io.EOF) {
r.readErr = err
}
}

r.releaseCodec()

stats := &ReadStats{
Expand All @@ -130,5 +141,9 @@ func (r *decompressReader) Close(ctx context.Context) (*ReadStats, error) {

_, innerErr := r.inner.Close(ctx)

if r.readErr != nil {
return stats, r.readErr
}

return stats, innerErr
}
52 changes: 52 additions & 0 deletions packages/shared/pkg/storage/compress_decode_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
package storage

import (
"bytes"
"crypto/sha256"
"io"
"testing"

"github.com/stretchr/testify/require"
)

// TestDecompressReaderVerifiesCRCOnClose ensures a footer-corrupt frame is not
// reported as success when the caller reads exactly the uncompressed size.
// zstd only verifies the frame checksum once the footer is consumed, which an
// exact-size read never triggers, so Close must drain-verify and surface the
// error.
func TestDecompressReaderVerifiesCRCOnClose(t *testing.T) {
t.Parallel()

const frameU = 512 * 1024
data := generateSemiRandomData(frameU)

up := &memPartUploader{}
fullFT, _, err := compressStream(t.Context(), bytes.NewReader(data), defaultCfg(CompressionZstd, 1, frameU), up, 1, nil)
require.NoError(t, err)
require.Equal(t, 1, fullFT.Table().NumFrames())
blob := up.Assemble()

// Corrupt the frame footer (last byte) so only the trailing checksum,
// verified after the last content byte, is wrong.
corrupt := bytes.Clone(blob)
corrupt[len(corrupt)-1] ^= 0xFF

dec, err := NewDecompressReader(bytesRangeReader(corrupt), CompressionZstd, SourceAWS, SeekableObjectType(0))
require.NoError(t, err)

// Read EXACTLY the uncompressed frame size — no read past EOF.
buf := make([]byte, frameU)
_, _ = io.ReadFull(dec, buf)
_, closeErr := dec.Close(t.Context())
require.Error(t, closeErr, "exact-size read of a footer-corrupt frame must fail on Close")

// A clean frame still round-trips and closes without error.
dec2, err := NewDecompressReader(bytesRangeReader(blob), CompressionZstd, SourceAWS, SeekableObjectType(0))
require.NoError(t, err)
got := make([]byte, frameU)
_, err = io.ReadFull(dec2, got)
require.NoError(t, err)
_, closeErr = dec2.Close(t.Context())
require.NoError(t, closeErr)
require.Equal(t, sha256.Sum256(data), sha256.Sum256(got))
}
8 changes: 8 additions & 0 deletions packages/shared/pkg/storage/compress_frame_table.go
Original file line number Diff line number Diff line change
Expand Up @@ -236,6 +236,14 @@ func (ft *FrameTable) LocateCompressed(offset int64) (Range, error) {
return Range{}, err
}

// The compressed range covers a whole frame; a mid-frame uncompressed
// offset would fetch and decode that frame from its start, silently
// returning data for the wrong position. Callers must frame-align via
// LocateUncompressed, so reject anything else rather than mis-serve.
if offset != e.StartU {
return Range{}, fmt.Errorf("offset %d is not frame-aligned (frame starts at %d); align via LocateUncompressed", offset, e.StartU)
Comment thread
dobrac marked this conversation as resolved.
}
Comment thread
dobrac marked this conversation as resolved.

return Range{Offset: e.StartC, Length: int(e.SizeC)}, nil
}

Expand Down
10 changes: 10 additions & 0 deletions packages/shared/pkg/storage/compress_frame_table_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,16 @@ func TestLocate(t *testing.T) {
require.Error(t, err)
})

t.Run("mid-frame offset errors", func(t *testing.T) {
t.Parallel()
// A mid-frame uncompressed offset would fetch and decode the whole
// containing frame from its start, silently returning data for the
// wrong position. It must be rejected so callers frame-align first.
_, err := ft.LocateCompressed((1 << 20) / 2)
require.Error(t, err)
require.Contains(t, err.Error(), "frame-aligned")
})

t.Run("nil table errors", func(t *testing.T) {
t.Parallel()
_, err := (*FrameTable)(nil).LocateCompressed(0)
Expand Down
90 changes: 81 additions & 9 deletions packages/shared/pkg/storage/gcp_multipart.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ import (
"net/http"
"os"
"slices"
"strings"
"sync"
"time"

Expand Down Expand Up @@ -55,6 +56,7 @@ func createRetryableClient(ctx context.Context, config RetryConfig) *retryableht
client.RetryMax = config.MaxAttempts - 1 // go-retryablehttp counts retries, not total attempts
client.RetryWaitMin = config.InitialBackoff
client.RetryWaitMax = config.MaxBackoff
client.CheckRetry = retryOnCompleteError

// Custom backoff function with full jitter to avoid thundering herd
client.Backoff = func(start, maxBackoff time.Duration, attemptNum int, _ *http.Response) time.Duration {
Expand Down Expand Up @@ -95,6 +97,40 @@ func createRetryableClient(ctx context.Context, config RetryConfig) *retryableht
return client
}

// retryOnCompleteError extends the default retry policy so a transient
// CompleteMultipartUpload failure is retried like a 5xx. S3 and the GCS XML
// API can return HTTP 200 with an <Error> body (e.g. InternalError) for a
// commit that didn't happen; the default policy only looks at the status code,
// so without this such a response fails on the first attempt and ignores the
// configured retry budget. Only the complete request (URL query "uploadId=…",
// small XML body) is inspected, so buffering the body to peek + restore it is
// cheap.
func retryOnCompleteError(ctx context.Context, resp *http.Response, err error) (bool, error) {
if retry, rerr := retryablehttp.DefaultRetryPolicy(ctx, resp, err); retry || rerr != nil {
return retry, rerr
}

if resp == nil || resp.StatusCode != http.StatusOK || resp.Request == nil ||
!strings.HasPrefix(resp.Request.URL.RawQuery, "uploadId=") || resp.Body == nil {
return false, nil
}

body, readErr := io.ReadAll(resp.Body)
resp.Body.Close()
// Restore the body so the caller (or a retry) can read it.
resp.Body = io.NopCloser(bytes.NewReader(body))

// Retry when the body couldn't be fully read (transient truncation/reset
// after the headers) or when it carries a transient <Error> (e.g.
// InternalError). Crucially, a read failure must not fall through as
// no-retry: that would hand completeUpload a clean, buffered partial body
// and let it commit on a truncated response.
var apiErr completeMultipartError
retry := readErr != nil || (xml.Unmarshal(body, &apiErr) == nil && apiErr.Code != "")

return retry, nil //nolint:nilerr // decision is carried by the bool; returning readErr would abort the retry loop instead of retrying
}

// zapLogger adapts zap.Logger to retryablehttp.LeveledLogger interface
var _ retryablehttp.LeveledLogger = &leveledLogger{}

Expand Down Expand Up @@ -134,6 +170,15 @@ type Part struct {
ETag string `xml:"ETag"`
}

// completeMultipartError matches the S3/GCS XML error payload that can arrive
// with an HTTP 200 status on CompleteMultipartUpload. Unmarshal only succeeds
// (Code populated) when the response root is <Error>.
type completeMultipartError struct {
XMLName xml.Name `xml:"Error"`
Code string `xml:"Code"`
Message string `xml:"Message"`
}

type MultipartUploader struct {
bucketName string
objectName string
Expand Down Expand Up @@ -257,13 +302,20 @@ func (m *MultipartUploader) uploadPart(ctx context.Context, uploadID string, par
url := fmt.Sprintf("%s/%s?partNumber=%d&uploadId=%s",
m.baseURL, m.objectName, partNumber, uploadID)

req, err := retryablehttp.NewRequestWithContext(ctx, "PUT", url, bytes.NewReader(data))
// A non-nil zero-length body counts as "unknown length" to net/http and
// is sent chunked, which S3-compatible XML backends reject with 411. Pass
// no body for empty parts so Content-Length: 0 is sent instead.
var body any
if len(data) > 0 {
body = bytes.NewReader(data)
}

req, err := retryablehttp.NewRequestWithContext(ctx, "PUT", url, body)
if err != nil {
return "", err
}

req.Header.Set("Authorization", "Bearer "+m.token)
req.Header.Set("Content-Length", fmt.Sprintf("%d", len(data)))
sum := md5.Sum(data) //nolint:gosec // GCS multipart uses Content-MD5 for transport integrity.
req.Header.Set("Content-MD5", base64.StdEncoding.EncodeToString(sum[:]))

Expand Down Expand Up @@ -297,18 +349,26 @@ func (m *MultipartUploader) uploadPartSlices(ctx context.Context, uploadID strin
url := fmt.Sprintf("%s/%s?partNumber=%d&uploadId=%s",
m.baseURL, m.objectName, partNumber, uploadID)

// Use a ReaderFunc so the retryable client can replay the body on retries
bodyFn := func() (io.Reader, error) {
return newMultiSliceReader(slices), nil
// Use a ReaderFunc so the retryable client can replay the body on
// retries. The multiSliceReader's Len makes retryablehttp set the
// request's ContentLength, so parts are sent with an explicit
// Content-Length rather than chunked transfer encoding (which GCS
// tolerates but S3-compatible XML backends reject with 411). Empty parts
// send no body at all for the same reason: a zero-length reader still
// counts as "unknown length" to net/http and would be chunked.
var body any
if totalLen > 0 {
body = retryablehttp.ReaderFunc(func() (io.Reader, error) {
return newMultiSliceReader(slices), nil
})
}

req, err := retryablehttp.NewRequestWithContext(ctx, "PUT", url, retryablehttp.ReaderFunc(bodyFn))
req, err := retryablehttp.NewRequestWithContext(ctx, "PUT", url, body)
if err != nil {
return "", err
}

req.Header.Set("Authorization", "Bearer "+m.token)
req.Header.Set("Content-Length", fmt.Sprintf("%d", totalLen))
h := md5.New() //nolint:gosec // GCS multipart uses Content-MD5 for transport integrity.
for _, s := range slices {
_, _ = h.Write(s)
Expand Down Expand Up @@ -365,12 +425,24 @@ func (m *MultipartUploader) completeUpload(ctx context.Context, uploadID string,
}
defer resp.Body.Close()

if resp.StatusCode != http.StatusOK {
body, _ := io.ReadAll(resp.Body)
body, readErr := io.ReadAll(resp.Body)
if readErr != nil {
return fmt.Errorf("read complete upload response (status %d): %w", resp.StatusCode, readErr)
}

if resp.StatusCode != http.StatusOK {
return fmt.Errorf("failed to complete upload (status %d): %s", resp.StatusCode, string(body))
}

// S3 and the GCS XML API can return HTTP 200 with an <Error> body when the
// commit fails server-side (documented for internal errors/timeouts).
// Treating that as success would record a frame table for an object that
// was never committed, so reject any response whose root is an Error.
var apiErr completeMultipartError
if xml.Unmarshal(body, &apiErr) == nil && apiErr.Code != "" {
return fmt.Errorf("failed to complete upload (status %d, code %s): %s", resp.StatusCode, apiErr.Code, apiErr.Message)
Comment thread
dobrac marked this conversation as resolved.
}

return nil
}

Expand Down
Loading
Loading