diff --git a/pkg/webdav/webdav.go b/pkg/webdav/webdav.go index b8c3707d..744a7248 100644 --- a/pkg/webdav/webdav.go +++ b/pkg/webdav/webdav.go @@ -619,12 +619,18 @@ func handleCopyMove(c *gin.Context, user *ent.User, fm manager.FileManager) (sta if err != nil && !ent.IsNotFound(err) { return purposeStatusCodeFromError(err), err } + dstExists := err == nil _, dstFolderUri, err := fm.SharedAddressTranslation(c, dst.DirUri()) if err != nil { return purposeStatusCodeFromError(err), err } + overwrite, err := parseOverwrite(c.Request.Header.Get("Overwrite")) + if err != nil { + return http.StatusBadRequest, err + } + hasher := dependency.FromContext(c).HashIDEncoder() if srcUri.IsSame(dstUri, hashid.EncodeUserID(hasher, user.ID)) { return http.StatusForbidden, errDestinationEqualsSource @@ -655,9 +661,7 @@ func handleCopyMove(c *gin.Context, user *ent.User, fm manager.FileManager) (sta } } - if err := fm.MoveOrCopy(ctx, []*fs.URI{srcUri}, dstFolderUri, true); err != nil { - return purposeStatusCodeFromError(err), err - } + return performCopyMove(ctx, fm, srcUri, dstUri, dstFolderUri, true, overwrite, dstExists) } release, ls, status, err := confirmLock(c, fm, user, srcTarget, dstTarget, srcUri, dstUri) @@ -675,7 +679,42 @@ func handleCopyMove(c *gin.Context, user *ent.User, fm manager.FileManager) (sta return http.StatusBadRequest, errInvalidDepth } } - if err := fm.MoveOrCopy(ctx, []*fs.URI{srcUri}, dstFolderUri, false); err != nil { + return performCopyMove(ctx, fm, srcUri, dstUri, dstFolderUri, false, overwrite, dstExists) +} + +type copyMoveOperations interface { + Delete(ctx context.Context, path []*fs.URI, opts ...fs.Option) error + MoveOrCopy(ctx context.Context, src []*fs.URI, dst *fs.URI, isCopy bool) error + Rename(ctx context.Context, path *fs.URI, newName string) (fs.File, error) +} + +func parseOverwrite(value string) (bool, error) { + switch value { + case "", "T": + return true, nil + case "F": + return false, nil + default: + return false, errInvalidOverwrite + } +} + +func performCopyMove( + ctx context.Context, + fm copyMoveOperations, + srcUri, dstUri, dstFolderUri *fs.URI, + isCopy, overwrite, dstExists bool, +) (int, error) { + if dstExists { + if !overwrite { + return http.StatusPreconditionFailed, errDestinationExists + } + if err := fm.Delete(ctx, []*fs.URI{dstUri}); err != nil { + return purposeStatusCodeFromError(err), err + } + } + + if err := fm.MoveOrCopy(ctx, []*fs.URI{srcUri}, dstFolderUri, isCopy); err != nil { return purposeStatusCodeFromError(err), err } @@ -685,7 +724,10 @@ func handleCopyMove(c *gin.Context, user *ent.User, fm manager.FileManager) (sta } } - return http.StatusNoContent, nil + if dstExists { + return http.StatusNoContent, nil + } + return http.StatusCreated, nil } func handleProppatch(c *gin.Context, user *ent.User, fm manager.FileManager) (status int, err error) { @@ -833,12 +875,14 @@ func StatusText(code int) string { var ( errDestinationEqualsSource = errors.New("webdav: destination equals source") + errDestinationExists = errors.New("webdav: destination exists") errDirectoryNotEmpty = errors.New("webdav: directory not empty") errInvalidDepth = errors.New("webdav: invalid depth") errInvalidDestination = errors.New("webdav: invalid destination") errInvalidIfHeader = errors.New("webdav: invalid If header") errInvalidLockInfo = errors.New("webdav: invalid lock info") errInvalidLockToken = errors.New("webdav: invalid lock token") + errInvalidOverwrite = errors.New("webdav: invalid overwrite") errInvalidPropfind = errors.New("webdav: invalid propfind") errInvalidProppatch = errors.New("webdav: invalid proppatch") errInvalidResponse = errors.New("webdav: invalid response") diff --git a/pkg/webdav/webdav_test.go b/pkg/webdav/webdav_test.go new file mode 100644 index 00000000..5a644560 --- /dev/null +++ b/pkg/webdav/webdav_test.go @@ -0,0 +1,113 @@ +package webdav + +import ( + "context" + "errors" + "net/http" + "testing" + + "github.com/cloudreve/Cloudreve/v4/pkg/filemanager/fs" +) + +type copyMoveOperationsStub struct { + calls []string + deleteErr error + moveErr error + renameErr error + moveIsCopy bool +} + +func (s *copyMoveOperationsStub) Delete(context.Context, []*fs.URI, ...fs.Option) error { + s.calls = append(s.calls, "delete") + return s.deleteErr +} + +func (s *copyMoveOperationsStub) MoveOrCopy(_ context.Context, _ []*fs.URI, _ *fs.URI, isCopy bool) error { + s.calls = append(s.calls, "move") + s.moveIsCopy = isCopy + return s.moveErr +} + +func (s *copyMoveOperationsStub) Rename(context.Context, *fs.URI, string) (fs.File, error) { + s.calls = append(s.calls, "rename") + return nil, s.renameErr +} + +func TestParseOverwrite(t *testing.T) { + tests := []struct { + value string + want bool + err bool + }{ + {value: "", want: true}, + {value: "T", want: true}, + {value: "F", want: false}, + {value: "true", err: true}, + } + + for _, test := range tests { + t.Run(test.value, func(t *testing.T) { + got, err := parseOverwrite(test.value) + if (err != nil) != test.err { + t.Fatalf("unexpected error: %v", err) + } + if got != test.want { + t.Fatalf("unexpected overwrite value: %v", got) + } + }) + } +} + +func TestPerformCopyMove(t *testing.T) { + src := mustWebDAVTestURI(t, "cloudreve://my/source.temp") + dst := mustWebDAVTestURI(t, "cloudreve://my/source.txt") + dstFolder := dst.DirUri() + + t.Run("overwrite disabled", func(t *testing.T) { + operations := ©MoveOperationsStub{} + status, err := performCopyMove(context.Background(), operations, src, dst, dstFolder, false, false, true) + if status != http.StatusPreconditionFailed || !errors.Is(err, errDestinationExists) { + t.Fatalf("unexpected result: status=%d err=%v", status, err) + } + if len(operations.calls) != 0 { + t.Fatalf("unexpected operations: %v", operations.calls) + } + }) + + t.Run("overwrite existing move", func(t *testing.T) { + operations := ©MoveOperationsStub{} + status, err := performCopyMove(context.Background(), operations, src, dst, dstFolder, false, true, true) + if err != nil || status != http.StatusNoContent { + t.Fatalf("unexpected result: status=%d err=%v", status, err) + } + if got := operations.calls; len(got) != 3 || got[0] != "delete" || got[1] != "move" || got[2] != "rename" { + t.Fatalf("unexpected operation order: %v", got) + } + if operations.moveIsCopy { + t.Fatal("move was executed as copy") + } + }) + + t.Run("new copy", func(t *testing.T) { + operations := ©MoveOperationsStub{} + status, err := performCopyMove(context.Background(), operations, src, dst, dstFolder, true, true, false) + if err != nil || status != http.StatusCreated { + t.Fatalf("unexpected result: status=%d err=%v", status, err) + } + if got := operations.calls; len(got) != 2 || got[0] != "move" || got[1] != "rename" { + t.Fatalf("unexpected operation order: %v", got) + } + if !operations.moveIsCopy { + t.Fatal("copy was executed as move") + } + }) +} + +func mustWebDAVTestURI(t *testing.T, raw string) *fs.URI { + t.Helper() + uri, err := fs.NewUriFromString(raw) + if err != nil { + t.Fatal(err) + } + return uri +}