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"`
+}