diff --git a/pkg/filemanager/fs/fs.go b/pkg/filemanager/fs/fs.go index 501aee31..c27a340e 100644 --- a/pkg/filemanager/fs/fs.go +++ b/pkg/filemanager/fs/fs.go @@ -282,6 +282,10 @@ type ( // uploaded. Used to safely trigger CompleteUpload only after every // chunk has been received when the client uploads chunks concurrently. ChunksReceived map[int]struct{} + // RangesReceived records merged byte intervals [start,end) written for + // arbitrary-range uploads (e.g. WebDAV Content-Range PUTs). Used to + // trigger CompleteUpload only after the whole file is covered. + RangesReceived [][2]int64 } // UploadProps properties of an upload session/request. diff --git a/pkg/filemanager/manager/upload.go b/pkg/filemanager/manager/upload.go index f42d2665..a9b29b9c 100644 --- a/pkg/filemanager/manager/upload.go +++ b/pkg/filemanager/manager/upload.go @@ -48,6 +48,12 @@ type ( // uploads chunks concurrently and the last-indexed chunk arrives before some // earlier chunks are still in flight. MarkChunkUploaded(ctx context.Context, session *fs.UploadSession, chunkIndex int) (allReceived bool, err error) + // MarkRangeUploaded atomically records the given byte range [offset, offset+length) + // as received on the shared upload session and reports whether the whole file + // is covered. Callers must invoke CompleteUpload only when allReceived is true. + // It is used for arbitrary-range uploads (e.g. WebDAV Content-Range PUTs) where + // chunks are addressed by byte offset rather than a fixed chunk index. + MarkRangeUploaded(ctx context.Context, session *fs.UploadSession, offset, length int64) (allReceived bool, err error) } ) @@ -310,6 +316,77 @@ func (m *manager) MarkChunkUploaded(ctx context.Context, session *fs.UploadSessi return false, nil } +// mergeByteRanges inserts [start,end) into the sorted interval list and +// coalesces overlapping or adjacent intervals. +func mergeByteRanges(ranges [][2]int64, start, end int64) [][2]int64 { + res := make([][2]int64, 0, len(ranges)+1) + inserted := false + for _, r := range ranges { + if r[1] < start { + res = append(res, r) + continue + } + if r[0] > end { + if !inserted { + res = append(res, [2]int64{start, end}) + inserted = true + } + res = append(res, r) + continue + } + // Overlapping or adjacent — extend the pending interval. + if r[0] < start { + start = r[0] + } + if r[1] > end { + end = r[1] + } + } + if !inserted { + res = append(res, [2]int64{start, end}) + } + return res +} + +// rangesCoverFull reports whether merged intervals cover [0,total). +func rangesCoverFull(ranges [][2]int64, total int64) bool { + return len(ranges) == 1 && ranges[0][0] <= 0 && ranges[0][1] >= total +} + +func (m *manager) MarkRangeUploaded(ctx context.Context, session *fs.UploadSession, offset, length int64) (bool, error) { + if session == nil || session.Props == nil || length < 0 || offset < 0 { + return false, fmt.Errorf("invalid upload session or range") + } + + sessionID := session.Props.UploadSessionID + mu := lockUploadSession(sessionID) + mu.Lock() + defer mu.Unlock() + + raw, ok := m.kv.Get(UploadSessionCachePrefix + sessionID) + if !ok { + // Session already completed or cancelled elsewhere. + return false, nil + } + fresh, ok := raw.(fs.UploadSession) + if !ok { + return false, fmt.Errorf("unexpected upload session type in KV") + } + + fresh.RangesReceived = mergeByteRanges(fresh.RangesReceived, offset, offset+length) + + ttl := max(1, int(time.Until(fresh.Props.ExpireAt).Seconds())) + if err := m.kv.Set(UploadSessionCachePrefix+sessionID, fresh, ttl); err != nil { + return false, fmt.Errorf("failed to persist upload session progress: %w", err) + } + + if rangesCoverFull(fresh.RangesReceived, fresh.Props.Size) { + session.RangesReceived = fresh.RangesReceived + return true, nil + } + return false, nil +} + func (m *manager) CancelUploadSession(ctx context.Context, path *fs.URI, sessionID string) error { // Get upload session var session *fs.UploadSession diff --git a/pkg/filemanager/manager/upload_test.go b/pkg/filemanager/manager/upload_test.go new file mode 100644 index 00000000..9d2fb040 --- /dev/null +++ b/pkg/filemanager/manager/upload_test.go @@ -0,0 +1,139 @@ +package manager + +import ( + "context" + "testing" + "time" + + "github.com/cloudreve/Cloudreve/v4/pkg/cache" + "github.com/cloudreve/Cloudreve/v4/pkg/filemanager/fs" +) + +func TestMergeByteRanges(t *testing.T) { + tests := []struct { + name string + ranges [][2]int64 + start int64 + end int64 + want [][2]int64 + }{ + {"empty", nil, 0, 10, [][2]int64{{0, 10}}}, + {"append gap", [][2]int64{{0, 10}}, 20, 30, [][2]int64{{0, 10}, {20, 30}}}, + {"prepend gap", [][2]int64{{20, 30}}, 0, 10, [][2]int64{{0, 10}, {20, 30}}}, + {"adjacent merge", [][2]int64{{0, 10}}, 10, 20, [][2]int64{{0, 20}}}, + {"bridge two", [][2]int64{{0, 10}, {20, 30}}, 10, 20, [][2]int64{{0, 30}}}, + {"contained", [][2]int64{{0, 30}}, 5, 15, [][2]int64{{0, 30}}}, + {"extend right", [][2]int64{{0, 10}}, 5, 20, [][2]int64{{0, 20}}}, + {"extend left", [][2]int64{{10, 20}}, 0, 15, [][2]int64{{0, 20}}}, + {"middle gap", [][2]int64{{0, 5}, {30, 40}}, 10, 20, [][2]int64{{0, 5}, {10, 20}, {30, 40}}}, + {"merge all", [][2]int64{{0, 5}, {10, 15}, {20, 25}}, 4, 21, [][2]int64{{0, 25}}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := mergeByteRanges(tt.ranges, tt.start, tt.end) + if len(got) != len(tt.want) { + t.Fatalf("got %v, want %v", got, tt.want) + } + for i := range got { + if got[i] != tt.want[i] { + t.Fatalf("got %v, want %v", got, tt.want) + } + } + }) + } +} + +func TestRangesCoverFull(t *testing.T) { + if !rangesCoverFull([][2]int64{{0, 100}}, 100) { + t.Fatal("expected full coverage") + } + if rangesCoverFull([][2]int64{{0, 99}}, 100) { + t.Fatal("partial coverage reported as full") + } + if rangesCoverFull([][2]int64{{0, 50}, {50, 100}}, 100) { + t.Fatal("unmerged intervals should not report full coverage") + } + if rangesCoverFull(nil, 100) { + t.Fatal("empty intervals reported as full") + } +} + +func newRangedSession(id string, size int64) fs.UploadSession { + return fs.UploadSession{ + Props: &fs.UploadProps{ + UploadSessionID: id, + Size: size, + ExpireAt: time.Now().Add(time.Hour), + }, + } +} + +func TestMarkRangeUploaded(t *testing.T) { + ctx := context.Background() + m := &manager{kv: cache.NewMemoStore("", nil)} + session := newRangedSession("test-range-session", 100) + if err := m.kv.Set(UploadSessionCachePrefix+"test-range-session", session, 60); err != nil { + t.Fatal(err) + } + + // Out-of-order arrival: second half first. + all, err := m.MarkRangeUploaded(ctx, &session, 50, 50) + if err != nil { + t.Fatal(err) + } + if all { + t.Fatal("reported complete after only second half") + } + + all, err = m.MarkRangeUploaded(ctx, &session, 0, 50) + if err != nil { + t.Fatal(err) + } + if !all { + t.Fatal("expected complete after full coverage") + } + + // Session record carries merged coverage. + raw, ok := m.kv.Get(UploadSessionCachePrefix + "test-range-session") + if !ok { + t.Fatal("session missing from KV") + } + stored := raw.(fs.UploadSession) + if len(stored.RangesReceived) != 1 || stored.RangesReceived[0] != [2]int64{0, 100} { + t.Fatalf("unexpected ranges: %v", stored.RangesReceived) + } +} + +func TestMarkRangeUploadedGap(t *testing.T) { + ctx := context.Background() + m := &manager{kv: cache.NewMemoStore("", nil)} + session := newRangedSession("test-gap-session", 100) + if err := m.kv.Set(UploadSessionCachePrefix+"test-gap-session", session, 60); err != nil { + t.Fatal(err) + } + + for _, r := range [][2]int64{{0, 40}, {60, 100}} { + all, err := m.MarkRangeUploaded(ctx, &session, r[0], r[1]-r[0]) + if err != nil { + t.Fatal(err) + } + if all { + t.Fatal("reported complete with a gap in coverage") + } + } +} + +func TestMarkRangeUploadedMissingSession(t *testing.T) { + ctx := context.Background() + m := &manager{kv: cache.NewMemoStore("", nil)} + session := newRangedSession("gone", 100) + + all, err := m.MarkRangeUploaded(ctx, &session, 0, 100) + if err != nil { + t.Fatal(err) + } + if all { + t.Fatal("missing session reported complete") + } +} diff --git a/pkg/webdav/webdav.go b/pkg/webdav/webdav.go index 744a7248..78ec8027 100644 --- a/pkg/webdav/webdav.go +++ b/pkg/webdav/webdav.go @@ -7,11 +7,13 @@ package webdav // import "golang.org/x/net/webdav" import ( "context" + "crypto/sha1" "errors" "fmt" "net/http" "net/url" "path" + "strconv" "strings" "time" @@ -256,6 +258,34 @@ func handlePut(c *gin.Context, user *ent.User, fm manager.FileManager) (status i return http.StatusBadRequest, err } + // A PUT with no length information at all (e.g. chunked transfer encoding) + // cannot be sized — previously this silently created an empty file. + if fileSize == 0 && c.Request.ContentLength < 0 && + c.Request.Header.Get("X-Expected-Entity-Length") == "" { + return http.StatusLengthRequired, nil + } + + // Ranged PUTs ("Content-Range: bytes start-end/total") are used by some + // clients (e.g. Mountain Duck) to upload large files in pieces. A partial + // range goes through the chunked-assembly path; a full-range PUT falls + // through to the regular overwrite path. + contentRange, err := parseContentRange(c.Request.Header.Get("Content-Range")) + if err != nil { + return http.StatusBadRequest, err + } + + m := manager.NewFileManager(dependency.FromContext(ctx), user) + defer m.Recycle() + + if contentRange != nil { + if contentRange.start != 0 || contentRange.end+1 != contentRange.total { + return handleRangedPut(ctx, c, user, m, fm, uri, rc, fileSize, contentRange) + } + if contentRange.total != fileSize { + return http.StatusBadRequest, errInvalidContentRange + } + } + fileData := &fs.UploadRequest{ Props: &fs.UploadProps{ Uri: uri, @@ -266,9 +296,6 @@ func handlePut(c *gin.Context, user *ent.User, fm manager.FileManager) (status i Mode: fs.ModeOverwrite, } - m := manager.NewFileManager(dependency.FromContext(ctx), user) - defer m.Recycle() - // Update file res, err := m.Update(ctx, fileData) if err != nil { @@ -284,6 +311,150 @@ func handlePut(c *gin.Context, user *ent.User, fm manager.FileManager) (status i return http.StatusCreated, nil } +var errInvalidContentRange = errors.New("invalid Content-Range") + +// contentRange describes a "bytes start-end/total" request range, with end +// inclusive per RFC 7233. +type contentRange struct { + start, end, total int64 +} + +// parseContentRange parses a Content-Range header of the form +// "bytes start-end/total". It returns (nil, nil) when the header is absent. +func parseContentRange(h string) (*contentRange, error) { + if h == "" { + return nil, nil + } + + h = strings.TrimSpace(h) + if !strings.HasPrefix(h, "bytes ") { + return nil, errInvalidContentRange + } + + rangePart, totalPart, ok := strings.Cut(h[len("bytes "):], "/") + if !ok || totalPart == "*" || totalPart == "" { + // An unknown total cannot be turned into a sized upload session. + return nil, errInvalidContentRange + } + + startPart, endPart, ok := strings.Cut(rangePart, "-") + if !ok { + return nil, errInvalidContentRange + } + + start, err := strconv.ParseInt(strings.TrimSpace(startPart), 10, 64) + if err != nil { + return nil, errInvalidContentRange + } + end, err := strconv.ParseInt(strings.TrimSpace(endPart), 10, 64) + if err != nil { + return nil, errInvalidContentRange + } + total, err := strconv.ParseInt(strings.TrimSpace(totalPart), 10, 64) + if err != nil { + return nil, errInvalidContentRange + } + + if start < 0 || end < start || total <= 0 || end >= total { + return nil, errInvalidContentRange + } + return &contentRange{start: start, end: end, total: total}, nil +} + +// handleRangedPut assembles a multi-request ranged PUT into a single upload +// session. Byte-range coverage is tracked on the session in KV; the upload is +// completed once [0,total) has been received. Only local storage policies are +// supported, as remote drivers cannot honor arbitrary write offsets. +func handleRangedPut(ctx context.Context, c *gin.Context, user *ent.User, m manager.FileManager, fm manager.FileManager, uri *fs.URI, rc request.LimitReaderCloser, fileSize int64, cr *contentRange) (status int, err error) { + if fileSize != cr.end-cr.start+1 { + return http.StatusBadRequest, errInvalidContentRange + } + + dep := dependency.FromContext(c) + kv := dep.KV() + + // One in-flight ranged upload per user+path, keyed deterministically so + // that subsequent chunk requests resume the same upload session. + sessionKey := fmt.Sprintf("dav-put-%d-%x", user.ID, sha1.Sum([]byte(uri.String()))) + + var session *fs.UploadSession + if raw, ok := kv.Get(manager.UploadSessionCachePrefix + sessionKey); ok { + s, ok := raw.(fs.UploadSession) + if !ok || s.Props == nil { + kv.Delete(manager.UploadSessionCachePrefix, sessionKey) + } else if s.Props.Size == cr.total { + session = &s + } else { + // A different upload to the same path — fail the stale session. + m.OnUploadFailed(ctx, &s) + session = nil + } + } + + if session == nil { + ttl := dep.SettingProvider().UploadSessionTTL(ctx) + if _, err := m.CreateUploadSession(ctx, &fs.UploadRequest{ + Props: &fs.UploadProps{ + Uri: uri, + Size: cr.total, + UploadSessionID: sessionKey, + ExpireAt: time.Now().Add(ttl), + }, + Mode: fs.ModeOverwrite, + }); err != nil { + return purposeStatusCodeFromError(err), err + } + + raw, ok := kv.Get(manager.UploadSessionCachePrefix + sessionKey) + if !ok { + return http.StatusInternalServerError, errors.New("upload session not persisted") + } + s := raw.(fs.UploadSession) + session = &s + + // Only the local driver honors arbitrary write offsets; for other + // policies a ranged PUT cannot be assembled safely. + if session.Policy == nil || session.Policy.Type != types.PolicyTypeLocal { + m.OnUploadFailed(ctx, session) + return http.StatusNotImplemented, errors.New("ranged PUT not supported by this storage policy") + } + } + + chunkReq := &fs.UploadRequest{ + File: rc, + Offset: cr.start, + Props: session.Props.Copy(), + Mode: fs.ModeOverwrite, + } + if err := m.Upload(ctx, chunkReq, session.Policy, session); err != nil { + return purposeStatusCodeFromError(err), err + } + + if lrc, ok := chunkReq.File.(request.LimitReaderCloser); ok && lrc.Count() != fileSize { + return http.StatusInternalServerError, fmt.Errorf("uploaded data(%d) does not match purposed size(%d)", lrc.Count(), fileSize) + } + + allReceived, err := m.MarkRangeUploaded(ctx, session, cr.start, fileSize) + if err != nil { + return http.StatusInternalServerError, err + } + if !allReceived { + return http.StatusCreated, nil + } + + res, err := m.CompleteUpload(ctx, session) + if err != nil { + return purposeStatusCodeFromError(err), err + } + + etag, err := findETag(ctx, fm, res) + if err != nil { + return http.StatusInternalServerError, err + } + c.Writer.Header().Set("ETag", etag) + return http.StatusCreated, nil +} + func handleOptions(c *gin.Context, user *ent.User, fm manager.FileManager) (status int, err error) { allow := []string{"OPTIONS", "LOCK", "PUT", "MKCOL"} diff --git a/pkg/webdav/webdav_test.go b/pkg/webdav/webdav_test.go index 5a644560..8c3682c6 100644 --- a/pkg/webdav/webdav_test.go +++ b/pkg/webdav/webdav_test.go @@ -111,3 +111,48 @@ func mustWebDAVTestURI(t *testing.T, raw string) *fs.URI { } return uri } + +func TestParseContentRange(t *testing.T) { + tests := []struct { + header string + want *contentRange + err bool + }{ + {"", nil, false}, + {"bytes 0-1048575/3145728", &contentRange{0, 1048575, 3145728}, false}, + {"bytes 1048576-2097151/3145728", &contentRange{1048576, 2097151, 3145728}, false}, + {"bytes 0-99/100", &contentRange{0, 99, 100}, false}, + {" bytes 0-9/10 ", &contentRange{0, 9, 10}, false}, + {"bytes 0-9/*", nil, true}, + {"items 0-9/10", nil, true}, + {"bytes 0-9", nil, true}, + {"bytes */10", nil, true}, + {"bytes 9-0/10", nil, true}, + {"bytes 0-10/10", nil, true}, + {"bytes -1-9/10", nil, true}, + {"bytes 0-9/0", nil, true}, + {"bytes a-b/c", nil, true}, + } + + for _, tt := range tests { + got, err := parseContentRange(tt.header) + if tt.err { + if err == nil { + t.Fatalf("header %q: expected error, got %+v", tt.header, got) + } + continue + } + if err != nil { + t.Fatalf("header %q: unexpected error: %v", tt.header, err) + } + if tt.want == nil { + if got != nil { + t.Fatalf("header %q: expected nil, got %+v", tt.header, got) + } + continue + } + if got == nil || *got != *tt.want { + t.Fatalf("header %q: got %+v, want %+v", tt.header, got, tt.want) + } + } +}