diff --git a/drivers/s3/driver.go b/drivers/s3/driver.go index 711f46ab5e..8928035713 100644 --- a/drivers/s3/driver.go +++ b/drivers/s3/driver.go @@ -224,11 +224,25 @@ func (d *S3) GetDirectUploadTools() []string { return []string{"HttpDirect"} } -func (d *S3) GetDirectUploadInfo(ctx context.Context, _ string, dstDir model.Obj, fileName string, _ int64) (any, error) { +func (d *S3) GetDirectUploadInfo(ctx context.Context, _ string, dstDir model.Obj, fileName string, fileSize int64) (any, error) { if !d.EnableDirectUpload { return nil, errs.NotImplement } + maxParts := d.DirectUploadMaxParts + if maxParts == 0 { + maxParts = maxCopyParts + } + return d.getDirectUploadInfo(ctx, dstDir, fileName, fileSize, maxParts, d.DirectUploadMinPartSize) +} + +func (d *S3) getDirectUploadInfo(ctx context.Context, dstDir model.Obj, fileName string, fileSize, maxParts, chunkSize int64) (any, error) { path := getKey(stdpath.Join(dstDir.GetPath(), fileName), false) + if maxParts > 1 && fileSize > minMultipartUploadPartSize { + return d.getMultipartDirectUploadInfo(ctx, path, fileSize, maxParts, chunkSize) + } + if fileSize > maxMultipartUploadPartSize { + return nil, fmt.Errorf("object size %d exceeds direct upload limit", fileSize) + } req, _ := d.directUploadClient.PutObjectRequest(&s3.PutObjectInput{ Bucket: &d.Bucket, Key: &path, @@ -246,6 +260,83 @@ func (d *S3) GetDirectUploadInfo(ctx context.Context, _ string, dstDir model.Obj }, nil } +func (d *S3) getMultipartDirectUploadInfo(ctx context.Context, key string, fileSize, maxParts, chunkSize int64) (*model.S3MultipartDirectUploadInfo, error) { + partSize, err := getMultipartUploadPartSize(fileSize, maxParts, chunkSize) + if err != nil { + return nil, err + } + created, err := d.directUploadClient.CreateMultipartUploadWithContext(ctx, &s3.CreateMultipartUploadInput{ + Bucket: &d.Bucket, + Key: &key, + }) + if err != nil { + return nil, err + } + uploadID := aws.StringValue(created.UploadId) + if uploadID == "" { + return nil, fmt.Errorf("create multipart upload returned an empty upload ID") + } + createdSuccessfully := false + defer func() { + if createdSuccessfully { + return + } + _, _ = d.directUploadClient.AbortMultipartUploadWithContext(context.WithoutCancel(ctx), &s3.AbortMultipartUploadInput{ + Bucket: &d.Bucket, + Key: &key, + UploadId: &uploadID, + }) + }() + + partCount := (fileSize + partSize - 1) / partSize + uploadURLs := make([]string, 0, partCount) + for partNumber := int64(1); partNumber <= partCount; partNumber++ { + req, _ := d.directUploadClient.UploadPartRequest(&s3.UploadPartInput{ + Bucket: &d.Bucket, + Key: &key, + PartNumber: &partNumber, + UploadId: &uploadID, + }) + if req == nil { + return nil, fmt.Errorf("failed to create multipart upload request for part %d", partNumber) + } + url, err := req.Presign(time.Hour * time.Duration(d.SignURLExpire)) + if err != nil { + return nil, err + } + uploadURLs = append(uploadURLs, url) + } + completeReq, _ := d.directUploadClient.CompleteMultipartUploadRequest(&s3.CompleteMultipartUploadInput{ + Bucket: &d.Bucket, + Key: &key, + UploadId: &uploadID, + }) + abortReq, _ := d.directUploadClient.AbortMultipartUploadRequest(&s3.AbortMultipartUploadInput{ + Bucket: &d.Bucket, + Key: &key, + UploadId: &uploadID, + }) + if completeReq == nil || abortReq == nil { + return nil, fmt.Errorf("failed to create multipart completion requests") + } + completeURL, err := completeReq.Presign(time.Hour * time.Duration(d.SignURLExpire)) + if err != nil { + return nil, err + } + abortURL, err := abortReq.Presign(time.Hour * time.Duration(d.SignURLExpire)) + if err != nil { + return nil, err + } + info := &model.S3MultipartDirectUploadInfo{ + ChunkSize: partSize, + UploadURLs: uploadURLs, + CompleteURL: completeURL, + AbortURL: abortURL, + } + createdSuccessfully = true + return info, nil +} + // implements driver.Getter interface func (d *S3) Get(ctx context.Context, path string) (model.Obj, error) { // try to get object as a file using HeadObject diff --git a/drivers/s3/meta.go b/drivers/s3/meta.go index 4243f12d56..7b3a317f88 100644 --- a/drivers/s3/meta.go +++ b/drivers/s3/meta.go @@ -23,6 +23,8 @@ type Addition struct { AddFilenameToDisposition bool `json:"add_filename_to_disposition" help:"Add filename to Content-Disposition header."` EnableDirectUpload bool `json:"enable_direct_upload" default:"false"` DirectUploadHost string `json:"direct_upload_host" required:"false"` + DirectUploadMaxParts int64 `json:"direct_upload_max_parts" type:"number" default:"10000" help:"Maximum number of parts for frontend multipart upload. Set to 1 to disable multipart upload."` + DirectUploadMinPartSize int64 `json:"direct_upload_min_part_size" type:"number" default:"104857600" help:"Minimum part size for frontend multipart upload, in bytes."` UserAgent string `json:"user_agent" required:"false" default:"" help:"Custom User-Agent for S3 requests."` } diff --git a/drivers/s3/util.go b/drivers/s3/util.go index cba8698fae..c4d3a15b1a 100644 --- a/drivers/s3/util.go +++ b/drivers/s3/util.go @@ -25,6 +25,10 @@ const ( defaultCopyPartSize int64 = 100 * 1024 * 1024 maxCopyPartSize int64 = 5 * 1024 * 1024 * 1024 maxCopyParts int64 = 10000 + + minMultipartUploadPartSize int64 = 5 * 1024 * 1024 + defaultMultipartUploadPartSize int64 = 100 * 1024 * 1024 + maxMultipartUploadPartSize int64 = 5 * 1024 * 1024 * 1024 ) // do others that not defined in Driver interface @@ -79,7 +83,9 @@ func (d *S3) getClient(clientType int) *s3.S3 { } if clientType == ClientTypeDirectUpload && d.DirectUploadHost != "" { client.Handlers.Build.PushBack(func(r *request.Request) { - if r.HTTPRequest.Method != http.MethodPut { + switch r.HTTPRequest.Method { + case http.MethodPut, http.MethodPost, http.MethodDelete: + default: return } split := strings.SplitN(d.DirectUploadHost, "://", 2) @@ -102,6 +108,24 @@ func getKey(path string, dir bool) string { return path } +func getMultipartUploadPartSize(size, maxParts, chunkSize int64) (int64, error) { + if maxParts <= 1 { + if size > maxMultipartUploadPartSize { + return 0, fmt.Errorf("object size %d exceeds direct upload limit", size) + } + return size, nil + } + maxParts = min(maxParts, maxCopyParts) + if size > maxMultipartUploadPartSize*maxParts { + return 0, fmt.Errorf("object size %d exceeds multipart upload limit", size) + } + if chunkSize <= 0 { + chunkSize = defaultMultipartUploadPartSize + } + chunkSize = min(chunkSize, maxMultipartUploadPartSize) + return max(chunkSize, (size+maxParts-1)/maxParts, minMultipartUploadPartSize), nil +} + var defaultPlaceholderName = ".openlist" func getPlaceholderName(placeholder string) string { diff --git a/drivers/s3/util_test.go b/drivers/s3/util_test.go index 6c718a2f62..f082bf3073 100644 --- a/drivers/s3/util_test.go +++ b/drivers/s3/util_test.go @@ -1,6 +1,7 @@ package s3 import ( + "bytes" "context" "fmt" "io" @@ -10,6 +11,7 @@ import ( "strings" "testing" + "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/aws/aws-sdk-go/aws" "github.com/aws/aws-sdk-go/aws/credentials" "github.com/aws/aws-sdk-go/aws/session" @@ -186,6 +188,255 @@ func TestGetCopyPartSize(t *testing.T) { } } +func TestGetMultipartUploadPartSize(t *testing.T) { + partSize, err := getMultipartUploadPartSize(defaultMultipartUploadPartSize*maxCopyParts, maxCopyParts, defaultMultipartUploadPartSize) + if err != nil { + t.Fatalf("getMultipartUploadPartSize: %v", err) + } + if partSize != defaultMultipartUploadPartSize { + t.Fatalf("part size = %d, want %d", partSize, defaultMultipartUploadPartSize) + } + + partSize, err = getMultipartUploadPartSize(defaultMultipartUploadPartSize*maxCopyParts+1, maxCopyParts, defaultMultipartUploadPartSize) + if err != nil { + t.Fatalf("getMultipartUploadPartSize: %v", err) + } + if partSize != defaultMultipartUploadPartSize+1 { + t.Fatalf("grown part size = %d, want %d", partSize, defaultMultipartUploadPartSize+1) + } + + partSize, err = getMultipartUploadPartSize(25*1024*1024, 2, 10*1024*1024) + if err != nil { + t.Fatalf("getMultipartUploadPartSize with custom max parts: %v", err) + } + if partSize != 25*1024*1024/2 { + t.Fatalf("custom max parts size = %d, want %d", partSize, 25*1024*1024/2) + } + + partSize, err = getMultipartUploadPartSize(25*1024*1024, maxCopyParts, 20*1024*1024) + if err != nil { + t.Fatalf("getMultipartUploadPartSize with custom chunk size: %v", err) + } + if partSize != 20*1024*1024 { + t.Fatalf("custom chunk size = %d, want %d", partSize, 20*1024*1024) + } + + partSize, err = getMultipartUploadPartSize(25*1024*1024, maxCopyParts, maxMultipartUploadPartSize+1) + if err != nil { + t.Fatalf("getMultipartUploadPartSize with oversized chunk size: %v", err) + } + if partSize != maxMultipartUploadPartSize { + t.Fatalf("oversized chunk size = %d, want %d", partSize, maxMultipartUploadPartSize) + } + + if _, err := getMultipartUploadPartSize(maxMultipartUploadPartSize*2+1, 2, defaultMultipartUploadPartSize); err == nil { + t.Fatal("getMultipartUploadPartSize returned nil error for an oversized object") + } +} + +func TestGetDirectUploadInfoUsesMultipartForLargeFiles(t *testing.T) { + const fileSize = 25 * 1024 * 1024 + created := false + d := newTestS3Driver(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost || !r.URL.Query().Has("uploads") { + t.Errorf("unexpected request: %s %s", r.Method, r.URL.String()) + w.WriteHeader(http.StatusBadRequest) + return + } + created = true + writeTestXML(t, w, `upload-id`) + }) + d.EnableDirectUpload = true + d.DirectUploadMaxParts = 2 + d.DirectUploadMinPartSize = 10 * 1024 * 1024 + + info, err := d.GetDirectUploadInfo(context.Background(), "HttpDirect", &model.Object{Path: "/"}, "large-file", fileSize) + if err != nil { + t.Fatalf("GetDirectUploadInfo: %v", err) + } + multipartInfo, ok := info.(*model.S3MultipartDirectUploadInfo) + if !ok { + t.Fatalf("upload info type = %T, want S3MultipartDirectUploadInfo", info) + } + if multipartInfo.ChunkSize != fileSize/2 { + t.Errorf("chunk size = %d, want %d", multipartInfo.ChunkSize, fileSize/2) + } + if len(multipartInfo.UploadURLs) != 2 { + t.Fatalf("part URL count = %d, want 2", len(multipartInfo.UploadURLs)) + } + if !created { + t.Fatal("multipart upload was not initiated") + } +} + +func TestGetDirectUploadInfoUsesConfiguredDirectUploadMinPartSize(t *testing.T) { + const fileSize = 25 * 1024 * 1024 + d := newTestS3Driver(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost || !r.URL.Query().Has("uploads") { + t.Errorf("unexpected request: %s %s", r.Method, r.URL.String()) + w.WriteHeader(http.StatusBadRequest) + return + } + writeTestXML(t, w, `upload-id`) + }) + d.EnableDirectUpload = true + d.DirectUploadMinPartSize = 20 * 1024 * 1024 + + info, err := d.GetDirectUploadInfo(context.Background(), "HttpDirect", &model.Object{Path: "/"}, "large-file", fileSize) + if err != nil { + t.Fatalf("GetDirectUploadInfo: %v", err) + } + multipartInfo, ok := info.(*model.S3MultipartDirectUploadInfo) + if !ok { + t.Fatalf("upload info type = %T, want S3MultipartDirectUploadInfo", info) + } + if multipartInfo.ChunkSize != d.DirectUploadMinPartSize { + t.Errorf("chunk size = %d, want %d", multipartInfo.ChunkSize, d.DirectUploadMinPartSize) + } + if len(multipartInfo.UploadURLs) != 2 { + t.Fatalf("part URL count = %d, want 2", len(multipartInfo.UploadURLs)) + } +} + +func TestGetDirectUploadInfoRejectsOversizedSinglePut(t *testing.T) { + d := newTestS3Driver(t, func(w http.ResponseWriter, r *http.Request) { + t.Errorf("unexpected request: %s %s", r.Method, r.URL.String()) + w.WriteHeader(http.StatusBadRequest) + }) + d.EnableDirectUpload = true + d.DirectUploadMaxParts = 1 + + if _, err := d.GetDirectUploadInfo(context.Background(), "HttpDirect", &model.Object{Path: "/"}, "large-file", maxMultipartUploadPartSize+1); err == nil { + t.Fatal("GetDirectUploadInfo returned nil error for oversized single PUT") + } +} + +func TestGetDirectUploadInfoUsesSinglePutWhenMultipartMaxPartsIsOne(t *testing.T) { + const fileSize = 25 * 1024 * 1024 + putRequests := 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 + } + putRequests++ + w.WriteHeader(http.StatusOK) + }) + d.EnableDirectUpload = true + d.DirectUploadMaxParts = 1 + + info, err := d.GetDirectUploadInfo(context.Background(), "HttpDirect", &model.Object{Path: "/"}, "large-file", fileSize) + if err != nil { + t.Fatalf("GetDirectUploadInfo: %v", err) + } + httpInfo, ok := info.(*model.HttpDirectUploadInfo) + if !ok { + t.Fatalf("upload info type = %T, want HttpDirectUploadInfo", info) + } + if httpInfo.UploadURL == "" { + t.Fatal("single upload URL is empty") + } + if putRequests != 0 { + t.Fatalf("PutObject was executed while presigning, requests = %d", putRequests) + } +} + +func TestDirectMultipartUploadCompletesWithUploadedPartETags(t *testing.T) { + const fileSize = 25 * 1024 * 1024 + uploadedParts := make(map[string]string) + completed := false + aborted := false + d := newTestS3Driver(t, func(w http.ResponseWriter, r *http.Request) { + switch { + 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": + partNumber := r.URL.Query().Get("partNumber") + body, err := io.ReadAll(r.Body) + if err != nil { + t.Errorf("read part %s body: %v", partNumber, err) + w.WriteHeader(http.StatusBadRequest) + return + } + uploadedParts[partNumber] = string(body) + w.Header().Set("ETag", fmt.Sprintf(`"etag-%s"`, partNumber)) + w.WriteHeader(http.StatusOK) + case r.Method == http.MethodPost && r.URL.Query().Get("uploadId") == "upload-id": + body, err := io.ReadAll(r.Body) + if err != nil { + t.Errorf("read completion body: %v", err) + w.WriteHeader(http.StatusBadRequest) + return + } + for partNumber := 1; partNumber <= 3; partNumber++ { + want := fmt.Sprintf("%d\"etag-%d\"", partNumber, partNumber) + if !strings.Contains(string(body), want) { + t.Errorf("completion body does not contain %q: %s", want, body) + } + } + 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) + } + }) + d.EnableDirectUpload = true + + info, err := d.GetDirectUploadInfo(context.Background(), "HttpDirect", &model.Object{Path: "/"}, "large-file", fileSize) + if err != nil { + t.Fatalf("GetDirectUploadInfo: %v", err) + } + multipartInfo, ok := info.(*model.S3MultipartDirectUploadInfo) + if !ok { + t.Fatalf("upload info type = %T, want S3MultipartDirectUploadInfo", info) + } + for i, uploadURL := range multipartInfo.UploadURLs { + request, err := http.NewRequestWithContext(context.Background(), http.MethodPut, uploadURL, bytes.NewBufferString(fmt.Sprintf("part-%d", i+1))) + if err != nil { + t.Fatalf("create part %d request: %v", i+1, err) + } + response, err := http.DefaultClient.Do(request) + if err != nil { + t.Fatalf("upload part %d: %v", i+1, err) + } + response.Body.Close() + if response.StatusCode != http.StatusOK { + t.Fatalf("part %d status = %d, want %d", i+1, response.StatusCode, http.StatusOK) + } + if got := response.Header.Get("ETag"); got != fmt.Sprintf(`"etag-%d"`, i+1) { + t.Fatalf("part %d ETag = %q", i+1, got) + } + } + + completion := `1"etag-1"2"etag-2"3"etag-3"` + request, err := http.NewRequestWithContext(context.Background(), http.MethodPost, multipartInfo.CompleteURL, strings.NewReader(completion)) + if err != nil { + t.Fatalf("create completion request: %v", err) + } + response, err := http.DefaultClient.Do(request) + if err != nil { + t.Fatalf("complete multipart upload: %v", err) + } + response.Body.Close() + if response.StatusCode != http.StatusOK { + t.Fatalf("completion status = %d, want %d", response.StatusCode, http.StatusOK) + } + if !completed { + t.Fatal("multipart upload was not completed") + } + if aborted { + t.Fatal("successful multipart upload was aborted") + } + if len(uploadedParts) != len(multipartInfo.UploadURLs) { + t.Fatalf("uploaded parts = %d, want %d", len(uploadedParts), len(multipartInfo.UploadURLs)) + } +} + func newTestS3Driver(t *testing.T, handler http.HandlerFunc) *S3 { t.Helper() server := httptest.NewServer(handler) @@ -201,8 +452,9 @@ func newTestS3Driver(t *testing.T, handler http.HandlerFunc) *S3 { t.Fatalf("create AWS session: %v", err) } return &S3{ - Addition: Addition{Bucket: "bucket"}, - client: awss3.New(sess), + Addition: Addition{Bucket: "bucket", SignURLExpire: 4}, + client: awss3.New(sess), + directUploadClient: awss3.New(sess), } } diff --git a/internal/model/direct_upload.go b/internal/model/direct_upload.go index 89bbfeb5d6..84cef3b8b3 100644 --- a/internal/model/direct_upload.go +++ b/internal/model/direct_upload.go @@ -6,3 +6,12 @@ type HttpDirectUploadInfo struct { Headers map[string]string `json:"headers,omitempty"` // Optional headers to include in the upload request Method string `json:"method,omitempty"` // HTTP method, default is PUT } + +// S3MultipartDirectUploadInfo contains presigned URLs for an S3 multipart upload. +// Parts, completion, and cancellation are sent directly to object storage. +type S3MultipartDirectUploadInfo struct { + ChunkSize int64 `json:"chunk_size"` + UploadURLs []string `json:"upload_urls"` + CompleteURL string `json:"complete_url"` + AbortURL string `json:"abort_url"` +}