diff --git a/drivers/s3/driver.go b/drivers/s3/driver.go index 711f46ab5..50a708568 100644 --- a/drivers/s3/driver.go +++ b/drivers/s3/driver.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "fmt" + "io" "net/url" stdpath "path" "strings" @@ -13,6 +14,7 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/internal/stream" + "github.com/OpenListTeam/OpenList/v4/internal/task" "github.com/OpenListTeam/OpenList/v4/pkg/cron" "github.com/OpenListTeam/OpenList/v4/pkg/utils" "github.com/OpenListTeam/OpenList/v4/server/common" @@ -224,11 +226,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 +262,184 @@ 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 +} + +func (d *S3) PutAsTask(ctx context.Context, dstDirPath string, file model.FileStreamer) (task.TaskExtensionInfo, error) { + return nil, errs.NotImplement +} + +func (d *S3) CreateMultipartUpload(ctx context.Context, dstDir model.Obj, fileName string, fileSize int64) (*model.DirectUploadPartOption, error) { + if d.DirectUploadMaxParts == 1 { + return nil, errs.NotSupport + } + maxParts := d.DirectUploadMaxParts + if maxParts == 0 { + maxParts = maxCopyParts + } + if _, err := calculatePartSize(fileSize, maxParts, minMultipartUploadPartSize, maxMultipartUploadPartSize); err != nil { + return nil, err + } + key := getKey(stdpath.Join(dstDir.GetPath(), fileName), false) + created, err := d.client.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") + } + return &model.DirectUploadPartOption{ + Key: key, + UploadId: uploadID, + }, nil +} + +func (d *S3) UploadPart(ctx context.Context, uploadId string, partNumber int, stream model.FileStreamer, options *model.DirectUploadPartOption) (*model.DirectUploadPartInfo, error) { + if d.DirectUploadMaxParts == 1 { + return nil, errs.NotSupport + } + if partNumber < 1 || partNumber > 10000 { + return nil, fmt.Errorf("partNumber must be between 1 and 10000, got %d", partNumber) + } + if uploadId == "" { + return nil, fmt.Errorf("uploadId must not be empty") + } + body, err := io.ReadAll(stream) + if err != nil { + return nil, fmt.Errorf("read part body: %w", err) + } + partNum := int64(partNumber) + key := options.Key + input := &s3.UploadPartInput{ + Bucket: &d.Bucket, + Key: &key, + PartNumber: &partNum, + UploadId: &uploadId, + Body: bytes.NewReader(body), + } + output, err := d.directUploadClient.UploadPartWithContext(ctx, input) + if err != nil { + return nil, err + } + return &model.DirectUploadPartInfo{ + ETag: aws.StringValue(output.ETag), + PartNumber: partNumber, + }, nil +} + +func (d *S3) CompleteMultipartUpload(ctx context.Context, uploadId string, dst string, parts []model.DirectUploadPartInfo) error { + completedParts := make([]*s3.CompletedPart, 0, len(parts)) + for _, p := range parts { + partNumber := int64(p.PartNumber) + etag := p.ETag + completedParts = append(completedParts, &s3.CompletedPart{ + PartNumber: &partNumber, + ETag: &etag, + }) + } + key := getKey(dst, false) + _, err := d.client.CompleteMultipartUploadWithContext(ctx, &s3.CompleteMultipartUploadInput{ + Bucket: &d.Bucket, + Key: &key, + UploadId: &uploadId, + MultipartUpload: &s3.CompletedMultipartUpload{ + Parts: completedParts, + }, + }) + if err != nil { + return fmt.Errorf("failed to complete multipart upload for %s (uploadId=%s): %w", dst, uploadId, err) + } + return nil +} + +func (d *S3) AbortMultipartUpload(ctx context.Context, uploadId string, dst string) error { + key := getKey(dst, false) + _, err := d.client.AbortMultipartUploadWithContext(ctx, &s3.AbortMultipartUploadInput{ + Bucket: &d.Bucket, + Key: &key, + UploadId: &uploadId, + }) + return err +} + // 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 @@ -325,3 +519,5 @@ func (d *S3) Get(ctx context.Context, path string) (model.Obj, error) { var _ driver.Driver = (*S3)(nil) var _ driver.Getter = (*S3)(nil) +var _ driver.MultipartBackend = (*S3)(nil) +var _ driver.PutAsTask = (*S3)(nil) diff --git a/drivers/s3/meta.go b/drivers/s3/meta.go index 4243f12d5..d2a22104c 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 direct multipart upload. Set to 1 to disable multipart and use single-part upload. Valid range: 1-10000."` + 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 cba8698fa..d36dcce21 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 { @@ -343,6 +367,21 @@ func getCopyPartSize(size int64) (int64, error) { return partSize, nil } +func calculatePartSize(fileSize, maxParts, minPartSize, maxPartSize int64) (int64, error) { + if fileSize <= 0 { + return 0, errors.New("file size must be positive") + } + partSize := (fileSize + maxParts - 1) / maxParts + if partSize < minPartSize { + return minPartSize, nil + } + if partSize > maxPartSize { + return 0, fmt.Errorf("file size %d bytes exceeds S3 multipart upload limit (max %d parts x %d bytes/part = %d bytes total)", + fileSize, maxParts, maxPartSize, maxParts*maxPartSize) + } + 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 { diff --git a/drivers/s3/util_test.go b/drivers/s3/util_test.go index 6c718a2f6..0c2092037 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,287 @@ 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 TestCalculatePartSize(t *testing.T) { + tests := []struct { + fileSize int64 + maxParts int64 + minPartSize int64 + maxPartSize int64 + wantErr bool + wantPart int64 + }{ + {100 * 1024 * 1024, 10000, minMultipartUploadPartSize, maxMultipartUploadPartSize, false, minMultipartUploadPartSize}, + {50*1024*1024*1024 + 1, 10000, minMultipartUploadPartSize, maxMultipartUploadPartSize, false, (50*1024*1024*1024+1 + 9999) / 10000}, + {maxMultipartUploadPartSize * 10000, 10000, minMultipartUploadPartSize, maxMultipartUploadPartSize, false, maxMultipartUploadPartSize}, + {maxMultipartUploadPartSize*10000 + 1, 10000, minMultipartUploadPartSize, maxMultipartUploadPartSize, true, 0}, + {25 * 1024 * 1024, 2, 10 * 1024 * 1024, 100 * 1024 * 1024, false, 13107200}, + {0, 10000, minMultipartUploadPartSize, maxMultipartUploadPartSize, true, 0}, + {-1, 10000, minMultipartUploadPartSize, maxMultipartUploadPartSize, true, 0}, + {5 * 1024 * 1024, 1, minMultipartUploadPartSize, maxMultipartUploadPartSize, false, minMultipartUploadPartSize}, + } + for _, tt := range tests { + t.Run(fmt.Sprintf("size=%d_parts=%d", tt.fileSize, tt.maxParts), func(t *testing.T) { + got, err := calculatePartSize(tt.fileSize, tt.maxParts, tt.minPartSize, tt.maxPartSize) + if (err != nil) != tt.wantErr { + t.Errorf("calculatePartSize error = %v, wantErr = %v", err, tt.wantErr) + return + } + if !tt.wantErr && got != tt.wantPart { + t.Errorf("calculatePartSize = %d, want %d", got, tt.wantPart) + } + }) + } +} + func newTestS3Driver(t *testing.T, handler http.HandlerFunc) *S3 { t.Helper() server := httptest.NewServer(handler) @@ -201,8 +484,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/driver/driver.go b/internal/driver/driver.go index 373bb5653..421307f73 100644 --- a/internal/driver/driver.go +++ b/internal/driver/driver.go @@ -4,6 +4,7 @@ import ( "context" "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/OpenListTeam/OpenList/v4/internal/task" ) type Driver interface { @@ -218,3 +219,14 @@ type DirectUploader interface { // return errs.NotImplement if the driver does not support the given direct upload tool GetDirectUploadInfo(ctx context.Context, tool string, dstDir model.Obj, fileName string, fileSize int64) (any, error) } + +type MultipartBackend interface { + CreateMultipartUpload(ctx context.Context, dstDir model.Obj, fileName string, fileSize int64) (*model.DirectUploadPartOption, error) + UploadPart(ctx context.Context, uploadId string, partNumber int, stream model.FileStreamer, options *model.DirectUploadPartOption) (*model.DirectUploadPartInfo, error) + CompleteMultipartUpload(ctx context.Context, uploadId string, dst string, parts []model.DirectUploadPartInfo) error + AbortMultipartUpload(ctx context.Context, uploadId string, dst string) error +} + +type PutAsTask interface { + PutAsTask(ctx context.Context, dstDirPath string, file model.FileStreamer) (task.TaskExtensionInfo, error) +} diff --git a/internal/model/direct_upload.go b/internal/model/direct_upload.go index 89bbfeb5d..6d098955a 100644 --- a/internal/model/direct_upload.go +++ b/internal/model/direct_upload.go @@ -6,3 +6,22 @@ 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"` +} + +type DirectUploadPartOption struct { + Key string `json:"key"` + UploadId string `json:"upload_id"` +} + +type DirectUploadPartInfo struct { + ETag string `json:"etag"` + PartNumber int `json:"part_number"` +}