diff --git a/drivers/s3/driver.go b/drivers/s3/driver.go index 46d7f414a..711f46ab5 100644 --- a/drivers/s3/driver.go +++ b/drivers/s3/driver.go @@ -178,7 +178,7 @@ func (d *S3) Move(ctx context.Context, srcObj, dstDir model.Obj) error { } func (d *S3) Rename(ctx context.Context, srcObj model.Obj, newName string) error { - err := d.copy(ctx, srcObj.GetPath(), stdpath.Join(stdpath.Dir(srcObj.GetPath()), newName), srcObj.IsDir()) + err := d.copy(ctx, srcObj.GetPath(), stdpath.Join(stdpath.Dir(srcObj.GetPath()), newName), srcObj.GetSize(), srcObj.IsDir()) if err != nil { return err } @@ -186,7 +186,7 @@ func (d *S3) Rename(ctx context.Context, srcObj model.Obj, newName string) error } func (d *S3) Copy(ctx context.Context, srcObj, dstDir model.Obj) error { - return d.copy(ctx, srcObj.GetPath(), stdpath.Join(dstDir.GetPath(), srcObj.GetName()), srcObj.IsDir()) + return d.copy(ctx, srcObj.GetPath(), stdpath.Join(dstDir.GetPath(), srcObj.GetName()), srcObj.GetSize(), srcObj.IsDir()) } func (d *S3) Remove(ctx context.Context, obj model.Obj) error { diff --git a/drivers/s3/util.go b/drivers/s3/util.go index 8500ed264..cba8698fa 100644 --- a/drivers/s3/util.go +++ b/drivers/s3/util.go @@ -3,6 +3,7 @@ package s3 import ( "context" "errors" + "fmt" "net/http" "net/url" "path" @@ -19,6 +20,13 @@ import ( log "github.com/sirupsen/logrus" ) +const ( + maxCopyObjectSize int64 = 5 * 1000 * 1000 * 1000 + defaultCopyPartSize int64 = 100 * 1024 * 1024 + maxCopyPartSize int64 = 5 * 1024 * 1024 * 1024 + maxCopyParts int64 = 10000 +) + // do others that not defined in Driver interface func (d *S3) initSession() error { @@ -212,17 +220,20 @@ func (d *S3) listV2(dirPath string, args model.ListArgs) ([]model.Obj, error) { return files, nil } -func (d *S3) copy(ctx context.Context, src string, dst string, isDir bool) error { +func (d *S3) copy(ctx context.Context, src string, dst string, size int64, isDir bool) error { if isDir { return d.copyDir(ctx, src, dst) } - return d.copyFile(ctx, src, dst) + return d.copyFile(ctx, src, dst, size) } -func (d *S3) copyFile(ctx context.Context, src string, dst string) error { +func (d *S3) copyFile(ctx context.Context, src string, dst string, size int64) error { srcKey := getKey(src, false) dstKey := getKey(dst, false) encodedKey := strings.ReplaceAll(url.PathEscape(d.Bucket+"/"+srcKey), "+", "%2B") + if size > maxCopyObjectSize { + return d.copyFileMultipart(ctx, srcKey, dstKey, encodedKey, size) + } input := &s3.CopyObjectInput{ Bucket: &d.Bucket, CopySource: aws.String(encodedKey), @@ -232,6 +243,106 @@ func (d *S3) copyFile(ctx context.Context, src string, dst string) error { return err } +func (d *S3) copyFileMultipart(ctx context.Context, srcKey, dstKey, encodedKey string, size int64) (err error) { + head, err := d.client.HeadObjectWithContext(ctx, &s3.HeadObjectInput{ + Bucket: &d.Bucket, + Key: &srcKey, + }) + if err != nil { + return err + } + if head.ContentLength != nil { + size = *head.ContentLength + } + partSize, err := getCopyPartSize(size) + if err != nil { + return err + } + createInput := &s3.CreateMultipartUploadInput{ + Bucket: &d.Bucket, + Key: &dstKey, + CacheControl: head.CacheControl, + ContentDisposition: head.ContentDisposition, + ContentEncoding: head.ContentEncoding, + ContentLanguage: head.ContentLanguage, + ContentType: head.ContentType, + Metadata: head.Metadata, + WebsiteRedirectLocation: head.WebsiteRedirectLocation, + } + if head.Expires != nil { + if expires, parseErr := http.ParseTime(*head.Expires); parseErr == nil { + createInput.Expires = &expires + } + } + created, err := d.client.CreateMultipartUploadWithContext(ctx, createInput) + if err != nil { + return err + } + uploadID := aws.StringValue(created.UploadId) + if uploadID == "" { + return errors.New("create multipart upload returned an empty upload ID") + } + completed := false + defer func() { + if completed { + return + } + _, abortErr := d.client.AbortMultipartUploadWithContext(context.WithoutCancel(ctx), &s3.AbortMultipartUploadInput{ + Bucket: &d.Bucket, + Key: &dstKey, + UploadId: &uploadID, + }) + if abortErr != nil { + err = errors.Join(err, fmt.Errorf("failed to abort multipart copy: %w", abortErr)) + } + }() + + parts := make([]*s3.CompletedPart, 0, (size+partSize-1)/partSize) + for start, partNumber := int64(0), int64(1); start < size; start, partNumber = start+partSize, partNumber+1 { + end := min(start+partSize, size) - 1 + copied, copyErr := d.client.UploadPartCopyWithContext(ctx, &s3.UploadPartCopyInput{ + Bucket: &d.Bucket, + CopySource: &encodedKey, + CopySourceRange: aws.String(fmt.Sprintf("bytes=%d-%d", start, end)), + Key: &dstKey, + PartNumber: &partNumber, + UploadId: &uploadID, + }) + if copyErr != nil { + return copyErr + } + if copied.CopyPartResult == nil || aws.StringValue(copied.CopyPartResult.ETag) == "" { + return fmt.Errorf("multipart copy part %d returned an empty ETag", partNumber) + } + parts = append(parts, &s3.CompletedPart{ + ETag: copied.CopyPartResult.ETag, + PartNumber: &partNumber, + }) + } + + _, err = d.client.CompleteMultipartUploadWithContext(ctx, &s3.CompleteMultipartUploadInput{ + Bucket: &d.Bucket, + Key: &dstKey, + UploadId: &uploadID, + MultipartUpload: &s3.CompletedMultipartUpload{ + Parts: parts, + }, + }) + if err != nil { + return err + } + completed = true + return nil +} + +func getCopyPartSize(size int64) (int64, error) { + partSize := max(defaultCopyPartSize, (size-1)/maxCopyParts+1) + if partSize > maxCopyPartSize { + return 0, fmt.Errorf("object size %d exceeds multipart copy limit", size) + } + return partSize, nil +} + func (d *S3) copyDir(ctx context.Context, src string, dst string) error { objs, err := op.List(ctx, d, src, model.ListArgs{S3ShowPlaceholder: true}) if err != nil { @@ -243,7 +354,7 @@ func (d *S3) copyDir(ctx context.Context, src string, dst string) error { if obj.IsDir() { err = d.copyDir(ctx, cSrc, cDst) } else { - err = d.copyFile(ctx, cSrc, cDst) + err = d.copyFile(ctx, cSrc, cDst, obj.GetSize()) } if err != nil { return err diff --git a/drivers/s3/util_test.go b/drivers/s3/util_test.go new file mode 100644 index 000000000..6c718a2f6 --- /dev/null +++ b/drivers/s3/util_test.go @@ -0,0 +1,215 @@ +package s3 + +import ( + "context" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strconv" + "strings" + "testing" + + "github.com/aws/aws-sdk-go/aws" + "github.com/aws/aws-sdk-go/aws/credentials" + "github.com/aws/aws-sdk-go/aws/session" + awss3 "github.com/aws/aws-sdk-go/service/s3" +) + +func TestCopyFileUsesCopyObjectAtLimit(t *testing.T) { + copyRequests := 0 + d := newTestS3Driver(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPut || r.URL.Query().Get("uploadId") != "" { + t.Errorf("unexpected request: %s %s", r.Method, r.URL.String()) + w.WriteHeader(http.StatusBadRequest) + return + } + copyRequests++ + writeTestXML(t, w, `"copy"`) + }) + + if err := d.copyFile(context.Background(), "source+file", "destination", maxCopyObjectSize); err != nil { + t.Fatalf("copyFile: %v", err) + } + if copyRequests != 1 { + t.Fatalf("copy requests = %d, want 1", copyRequests) + } +} + +func TestCopyFileUsesMultipartCopyAboveLimit(t *testing.T) { + size := maxCopyObjectSize + 1 + wantParts := int((size + defaultCopyPartSize - 1) / defaultCopyPartSize) + ranges := make(map[int]string, wantParts) + completed := false + aborted := false + + d := newTestS3Driver(t, func(w http.ResponseWriter, r *http.Request) { + switch { + case r.Method == http.MethodHead: + w.Header().Set("Content-Length", strconv.FormatInt(size, 10)) + w.Header().Set("Content-Type", "application/octet-stream") + w.Header().Set("Cache-Control", "max-age=60") + w.Header().Set("Content-Disposition", "attachment") + w.Header().Set("Expires", "Wed, 21 Oct 2015 07:28:00 GMT") + w.Header().Set("X-Amz-Meta-Source", "preserved") + w.Header().Set("X-Amz-Website-Redirect-Location", "/redirect") + w.WriteHeader(http.StatusOK) + case r.Method == http.MethodPost && r.URL.Query().Has("uploads"): + if got := r.Header.Get("Cache-Control"); got != "max-age=60" { + t.Errorf("Cache-Control = %q, want %q", got, "max-age=60") + } + if got := r.Header.Get("Content-Disposition"); got != "attachment" { + t.Errorf("Content-Disposition = %q, want %q", got, "attachment") + } + if got := r.Header.Get("Content-Type"); got != "application/octet-stream" { + t.Errorf("Content-Type = %q, want %q", got, "application/octet-stream") + } + if got := r.Header.Get("Expires"); got != "Wed, 21 Oct 2015 07:28:00 GMT" { + t.Errorf("Expires = %q, want an unchanged HTTP date", got) + } + if got := r.Header.Get("X-Amz-Meta-Source"); got != "preserved" { + t.Errorf("metadata = %q, want %q", got, "preserved") + } + if got := r.Header.Get("X-Amz-Website-Redirect-Location"); got != "/redirect" { + t.Errorf("website redirect = %q, want %q", got, "/redirect") + } + writeTestXML(t, w, `upload-id`) + case r.Method == http.MethodPut && r.URL.Query().Get("uploadId") == "upload-id": + partNumber, err := strconv.Atoi(r.URL.Query().Get("partNumber")) + if err != nil { + t.Errorf("invalid part number: %v", err) + w.WriteHeader(http.StatusBadRequest) + return + } + if got := r.Header.Get("X-Amz-Copy-Source"); !strings.Contains(got, "source%2Bfile") { + t.Errorf("copy source = %q, want encoded source key", got) + } + ranges[partNumber] = r.Header.Get("X-Amz-Copy-Source-Range") + writeTestXML(t, w, fmt.Sprintf(`"part-%d"`, partNumber)) + case r.Method == http.MethodPost && r.URL.Query().Get("uploadId") == "upload-id": + body, err := io.ReadAll(r.Body) + if err != nil { + t.Errorf("read complete body: %v", err) + } + if got := strings.Count(string(body), ""); got != wantParts { + t.Errorf("completed parts = %d, want %d", got, wantParts) + } + completed = true + writeTestXML(t, w, `"complete"`) + case r.Method == http.MethodDelete && r.URL.Query().Get("uploadId") == "upload-id": + aborted = true + w.WriteHeader(http.StatusNoContent) + default: + t.Errorf("unexpected request: %s %s", r.Method, r.URL.String()) + w.WriteHeader(http.StatusBadRequest) + } + }) + + if err := d.copyFile(context.Background(), "source+file", "destination", size); err != nil { + t.Fatalf("copyFile: %v", err) + } + if !completed { + t.Fatal("multipart upload was not completed") + } + if aborted { + t.Fatal("successful multipart upload was aborted") + } + if len(ranges) != wantParts { + t.Fatalf("copied parts = %d, want %d", len(ranges), wantParts) + } + if got := ranges[1]; got != fmt.Sprintf("bytes=0-%d", defaultCopyPartSize-1) { + t.Errorf("first range = %q", got) + } + lastStart := int64(wantParts-1) * defaultCopyPartSize + if got := ranges[wantParts]; got != fmt.Sprintf("bytes=%d-%d", lastStart, size-1) { + t.Errorf("last range = %q", got) + } +} + +func TestCopyFileMultipartAbortsOnPartFailure(t *testing.T) { + size := maxCopyObjectSize + 1 + aborted := false + completed := false + + d := newTestS3Driver(t, func(w http.ResponseWriter, r *http.Request) { + switch { + case r.Method == http.MethodHead: + w.Header().Set("Content-Length", strconv.FormatInt(size, 10)) + w.WriteHeader(http.StatusOK) + case r.Method == http.MethodPost && r.URL.Query().Has("uploads"): + writeTestXML(t, w, `upload-id`) + case r.Method == http.MethodPut && r.URL.Query().Get("uploadId") == "upload-id": + w.WriteHeader(http.StatusInternalServerError) + writeTestXML(t, w, `InternalErrorcopy failed`) + case r.Method == http.MethodDelete && r.URL.Query().Get("uploadId") == "upload-id": + aborted = true + w.WriteHeader(http.StatusNoContent) + case r.Method == http.MethodPost && r.URL.Query().Get("uploadId") == "upload-id": + completed = true + w.WriteHeader(http.StatusOK) + default: + t.Errorf("unexpected request: %s %s", r.Method, r.URL.String()) + w.WriteHeader(http.StatusBadRequest) + } + }) + + if err := d.copyFile(context.Background(), "source", "destination", size); err == nil { + t.Fatal("copyFile returned nil error") + } + if !aborted { + t.Fatal("failed multipart upload was not aborted") + } + if completed { + t.Fatal("failed multipart upload was completed") + } +} + +func TestGetCopyPartSize(t *testing.T) { + partSize, err := getCopyPartSize(defaultCopyPartSize * maxCopyParts) + if err != nil { + t.Fatalf("getCopyPartSize: %v", err) + } + if partSize != defaultCopyPartSize { + t.Fatalf("part size = %d, want %d", partSize, defaultCopyPartSize) + } + + partSize, err = getCopyPartSize(defaultCopyPartSize*maxCopyParts + 1) + if err != nil { + t.Fatalf("getCopyPartSize: %v", err) + } + if partSize != defaultCopyPartSize+1 { + t.Fatalf("grown part size = %d, want %d", partSize, defaultCopyPartSize+1) + } + + if _, err := getCopyPartSize(maxCopyPartSize*maxCopyParts + 1); err == nil { + t.Fatal("getCopyPartSize returned nil error for an oversized object") + } +} + +func newTestS3Driver(t *testing.T, handler http.HandlerFunc) *S3 { + t.Helper() + server := httptest.NewServer(handler) + t.Cleanup(server.Close) + sess, err := session.NewSession(&aws.Config{ + Credentials: credentials.NewStaticCredentials("access-key", "secret-key", ""), + Endpoint: aws.String(server.URL), + Region: aws.String("us-east-1"), + S3ForcePathStyle: aws.Bool(true), + MaxRetries: aws.Int(0), + }) + if err != nil { + t.Fatalf("create AWS session: %v", err) + } + return &S3{ + Addition: Addition{Bucket: "bucket"}, + client: awss3.New(sess), + } +} + +func writeTestXML(t *testing.T, w http.ResponseWriter, body string) { + t.Helper() + w.Header().Set("Content-Type", "application/xml") + if _, err := io.WriteString(w, body); err != nil { + t.Errorf("write response: %v", err) + } +}