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
176 changes: 161 additions & 15 deletions packages/shared/pkg/storage/storage_aws.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,9 @@ import (
"fmt"
"io"
"os"
"slices"
"strings"
"sync"
"time"

"github.com/aws/aws-sdk-go-v2/aws"
Expand All @@ -23,9 +25,10 @@ import (
)

const (
awsOperationTimeout = 5 * time.Second
awsWriteTimeout = 30 * time.Second
awsReadTimeout = 15 * time.Second
awsOperationTimeout = 5 * time.Second
awsWriteTimeout = 30 * time.Second
awsReadTimeout = 15 * time.Second
awsMultipartUploadPartSize = 10 * 1024 * 1024
)

type awsStorage struct {
Expand Down Expand Up @@ -188,19 +191,23 @@ func (o *awsObject) WriteTo(ctx context.Context, dst io.Writer) (n int64, err er

func (o *awsObject) StoreFile(ctx context.Context, path string, opts ...PutOption) (*FullFrameTable, [32]byte, error) {
p := ApplyPutOptions(opts)
if CompressConfigFromOpts(p).IsCompressionEnabled() {
return nil, [32]byte{}, errors.New("compressed uploads are not supported on AWS (builds target GCP only)")
}

release, err := o.limiter.AcquireUploadSlot(ctx)
if err != nil {
return nil, [32]byte{}, err
}
defer release()

cfg := CompressConfigFromOpts(p)
if cfg.IsCompressionEnabled() {
return storeFileCompressed(ctx, path, cfg, o.limiter.MaxUploadTasks(ctx), p, func(metadata ObjectMetadata) (partUploader, error) {
Comment thread
dobrac marked this conversation as resolved.
Comment thread
dobrac marked this conversation as resolved.
return &awsPartUploader{client: o.client, bucketName: o.bucketName, objectName: o.path, metadata: metadata}, nil
})
Comment thread
dobrac marked this conversation as resolved.
}

// Inherit the caller's context for the multipart upload. The AWS SDK's
// manager.Uploader reuses the same ctx for CreateMultipartUpload, every
// UploadPart (Concurrency=8, PartSize=10MB), and the final Complete/Abort —
// UploadPart, and the final Complete/Abort —
Comment thread
dobrac marked this conversation as resolved.
Comment thread
dobrac marked this conversation as resolved.
// a tight static timeout here would cancel an in-flight multi-GB snapshot
// upload and surface as "S3: UploadPart ... StatusCode: 0, canceled,
// context deadline exceeded". The caller (pkg/server/sandboxes.go) already
Expand All @@ -215,7 +222,7 @@ func (o *awsObject) StoreFile(ctx context.Context, path string, opts ...PutOptio
uploader := manager.NewUploader(
o.client,
func(u *manager.Uploader) {
u.PartSize = 10 * 1024 * 1024 // 10 MB
u.PartSize = awsMultipartUploadPartSize
u.Concurrency = o.limiter.MaxUploadTasks(ctx)
},
)
Expand Down Expand Up @@ -275,26 +282,51 @@ func (o *awsObject) OpenRangeReader(ctx context.Context, off, length int64, fram
RecordReadOpen(ctx, time.Since(start), objType, SourceAWS, frameTable.CompressionType(), err)
}()

if frameTable.IsCompressed() {
return nil, SourceAWS, errors.New("compressed reads are not supported on AWS")
if !frameTable.IsCompressed() {
rc, err := o.openRangeReader(ctx, off, length)
if err != nil {
return nil, SourceAWS, err
}

return rc, SourceAWS, nil
}

r, err := frameTable.LocateCompressed(off)
if err != nil {
return nil, SourceAWS, fmt.Errorf("get frame for offset %d, S3:%s: %w", off, o.path, err)
}

raw, err := o.openRangeReader(ctx, r.Offset, int64(r.Length))
if err != nil {
return nil, SourceAWS, err
}

dec, err := NewDecompressReader(raw, frameTable.CompressionType(), SourceAWS, objType)
if err != nil {
raw.Close(ctx)

return nil, SourceAWS, err
}

readRange := aws.String(fmt.Sprintf("bytes=%d-%d", off, off+length-1))
return dec, SourceAWS, nil
}

func (o *awsObject) openRangeReader(ctx context.Context, off, length int64) (RangeReader, error) {
resp, err := o.client.GetObject(ctx, &s3.GetObjectInput{
Bucket: aws.String(o.bucketName),
Key: aws.String(o.path),
Range: readRange,
Range: aws.String(fmt.Sprintf("bytes=%d-%d", off, off+length-1)),
})
if err != nil {
var nsk *types.NoSuchKey
if errors.As(err, &nsk) {
return nil, SourceAWS, ErrObjectNotExist
return nil, ErrObjectNotExist
}

return nil, SourceAWS, fmt.Errorf("failed to create S3 range reader for %q: %w", o.path, err)
return nil, fmt.Errorf("failed to create S3 range reader for %q: %w", o.path, err)
}

return NewRangeReader(resp.Body), SourceAWS, nil
return NewRangeReader(resp.Body), nil
}

func (o *awsObject) Size(ctx context.Context) (_ int64, err error) {
Expand All @@ -316,6 +348,10 @@ func (o *awsObject) Size(ctx context.Context) (_ int64, err error) {
return 0, err
}

if size, ok := ObjectMetadata(resp.Metadata).UncompressedSize(); ok {
return size, nil
}

return *resp.ContentLength, nil
}

Expand Down Expand Up @@ -346,3 +382,113 @@ func ignoreNotExists(err error) error {

return err
}

type awsPartUploader struct {
client *s3.Client
bucketName string
objectName string
metadata ObjectMetadata

mu sync.Mutex
uploadID string
parts []types.CompletedPart
// completed needs no lock: compressStream calls Complete and the deferred
// Close sequentially from one goroutine, after all UploadPart calls finish.
completed bool
}

var _ partUploader = (*awsPartUploader)(nil)

func (m *awsPartUploader) Start(ctx context.Context) error {
out, err := m.client.CreateMultipartUpload(ctx, &s3.CreateMultipartUploadInput{
Bucket: aws.String(m.bucketName),
Key: aws.String(m.objectName),
Metadata: m.metadata,
// The SDK's default integrity protections attach CRC32 checksums to
// UploadPart requests; S3 requires the algorithm to be declared at
// initiation and echoed per part in Complete. Declare it explicitly on
// every call so the flow is consistent regardless of SDK/env config
// (manager.Uploader does the same for the uncompressed path).
ChecksumAlgorithm: types.ChecksumAlgorithmCrc32,
})
Comment thread
dobrac marked this conversation as resolved.
if err != nil {
return fmt.Errorf("failed to initiate multipart upload: %w", err)
}

m.uploadID = aws.ToString(out.UploadId)

return nil
}

// UploadPart uploads a single part. Multiple data slices are streamed without
// copying into a contiguous buffer; the section reader's Seek lets the SDK
// compute the payload hash/length and rewind on retries.
func (m *awsPartUploader) UploadPart(ctx context.Context, partIndex int, data ...[]byte) error {
body := newMultiSliceReader(data)
out, err := m.client.UploadPart(ctx, &s3.UploadPartInput{
Bucket: aws.String(m.bucketName),
Key: aws.String(m.objectName),
UploadId: aws.String(m.uploadID),
PartNumber: aws.Int32(int32(partIndex)),
Body: body,
ContentLength: aws.Int64(body.Size()),
ChecksumAlgorithm: types.ChecksumAlgorithmCrc32,
})
if err != nil {
return fmt.Errorf("failed to upload part %d: %w", partIndex, err)
}

m.mu.Lock()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Wondering if it makes sense to make a cloud agnostic abstraction since it looks like we have to implement multipart splitting for both AWS and GCP.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The GCP and AWS multipart protocols are unfortunately different

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'll check if the parts splitting could be refactored to common logic in a follow up work

m.parts = append(m.parts, types.CompletedPart{
ETag: out.ETag,
ChecksumCRC32: out.ChecksumCRC32,
PartNumber: aws.Int32(int32(partIndex)),
})
m.mu.Unlock()

return nil
}

func (m *awsPartUploader) Complete(ctx context.Context) error {
m.mu.Lock()
parts := make([]types.CompletedPart, len(m.parts))
copy(parts, m.parts)
m.mu.Unlock()

slices.SortFunc(parts, func(a, b types.CompletedPart) int {
return int(aws.ToInt32(a.PartNumber) - aws.ToInt32(b.PartNumber))
})

_, err := m.client.CompleteMultipartUpload(ctx, &s3.CompleteMultipartUploadInput{
Bucket: aws.String(m.bucketName),
Key: aws.String(m.objectName),
UploadId: aws.String(m.uploadID),
MultipartUpload: &types.CompletedMultipartUpload{
Parts: parts,
},
})
if err != nil {
return err
}

m.completed = true

return nil
Comment thread
dobrac marked this conversation as resolved.
}

func (m *awsPartUploader) Close() error {
if m.completed || m.uploadID == "" {
return nil
}
Comment thread
dobrac marked this conversation as resolved.

ctx, cancel := context.WithTimeout(context.Background(), awsOperationTimeout)
defer cancel()

_, err := m.client.AbortMultipartUpload(ctx, &s3.AbortMultipartUploadInput{
Bucket: aws.String(m.bucketName),
Key: aws.String(m.objectName),
UploadId: aws.String(m.uploadID),
})

return err
}
Loading
Loading