From 23eabf1d8d23acdac29d413ba59b7190985baf13 Mon Sep 17 00:00:00 2001 From: nostalume Date: Thu, 24 Sep 2026 20:05:50 +0800 Subject: [PATCH 1/6] refactor(authz): centralize path authorization - Move user, meta, path, hidden-name, and password policy into internal/authz. - Remove authorization regex state from object merging and migrate protocol callers directly. - Replace structural tests with a compact behavior matrix for the policy owner. Co-authored-by: Codex <267193182+codex@users.noreply.github.com> --- internal/authz/policy.go | 98 +++ internal/authz/policy_test.go | 227 ++++++ internal/fs/list.go | 32 +- internal/fs/list_test.go | 151 ---- internal/model/obj.go | 19 +- server/common/check.go | 86 --- server/common/check_test.go | 1035 ---------------------------- server/ftp/fsmanage.go | 12 +- server/ftp/fsread.go | 8 +- server/ftp/fsup.go | 6 +- server/handles/archive.go | 7 +- server/handles/direct_upload.go | 5 +- server/handles/fsbatch.go | 9 +- server/handles/fsmanage.go | 19 +- server/handles/fsread.go | 21 +- server/handles/offline_download.go | 3 +- server/handles/search.go | 3 +- server/handles/torrent.go | 5 +- server/mcp/fs_get.go | 3 +- server/mcp/fs_link.go | 3 +- server/mcp/fs_list.go | 13 +- server/middlewares/down.go | 3 +- server/middlewares/fsup.go | 5 +- server/webdav/file.go | 8 +- server/webdav/webdav.go | 21 +- 25 files changed, 415 insertions(+), 1387 deletions(-) create mode 100644 internal/authz/policy.go create mode 100644 internal/authz/policy_test.go delete mode 100644 internal/fs/list_test.go delete mode 100644 server/common/check_test.go diff --git a/internal/authz/policy.go b/internal/authz/policy.go new file mode 100644 index 0000000000..ce7367365a --- /dev/null +++ b/internal/authz/policy.go @@ -0,0 +1,98 @@ +package authz + +import ( + "path" + "slices" + "strings" + + "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/OpenListTeam/OpenList/v4/pkg/utils" + "github.com/dlclark/regexp2" +) + +func CanRead(user *model.User, meta *model.Meta, reqPath string) bool { + if user == nil { + return true + } + return meta == nil || len(meta.ReadUsers) == 0 || slices.Contains(meta.ReadUsers, user.ID) || + !MetaCoversPath(meta.Path, reqPath, meta.ReadUsersSub) +} + +func CanWrite(user *model.User, meta *model.Meta, reqPath string) bool { + if user == nil { + return true + } + return meta == nil || len(meta.WriteUsers) == 0 || slices.Contains(meta.WriteUsers, user.ID) || + !MetaCoversPath(meta.Path, reqPath, meta.WriteUsersSub) +} + +func CanWriteContentBypassUserPerms(meta *model.Meta, reqPath string) bool { + if meta == nil || !meta.Write { + return false + } + return utils.PathEqual(meta.Path, reqPath) || meta.WSub && utils.IsSubPath(meta.Path, reqPath) +} + +func IsHidden(user *model.User, meta *model.Meta, reqPath string) bool { + return matchesHidden(hidePatterns(user, meta, path.Dir(reqPath)), path.Base(reqPath)) +} + +func FilterHidden(user *model.User, meta *model.Meta, parentPath string, objs []model.Obj) []model.Obj { + patterns := hidePatterns(user, meta, parentPath) + if len(patterns) == 0 { + return objs + } + return slices.DeleteFunc(objs, func(obj model.Obj) bool { + return matchesHidden(patterns, obj.GetName()) + }) +} + +func hidePatterns(user *model.User, meta *model.Meta, parentPath string) []*regexp2.Regexp { + if user == nil || user.CanSeeHides() || meta == nil || meta.Hide == "" || + !MetaCoversPath(meta.Path, parentPath, meta.HSub) { + return nil + } + patterns := make([]*regexp2.Regexp, 0, strings.Count(meta.Hide, "\n")+1) + for hide := range strings.SplitSeq(meta.Hide, "\n") { + patterns = append(patterns, regexp2.MustCompile(hide, regexp2.None)) + } + return patterns +} + +func matchesHidden(patterns []*regexp2.Regexp, name string) bool { + for _, pattern := range patterns { + matched, _ := pattern.MatchString(name) + if matched { + return true + } + } + return false +} + +func CanAccess(user *model.User, meta *model.Meta, reqPath, password string) bool { + if IsHidden(user, meta, reqPath) || !CanRead(user, meta, reqPath) { + return false + } + if user.CanAccessWithoutPassword() || meta == nil || meta.Password == "" { + return true + } + return !MetaCoversPath(meta.Path, reqPath, meta.PSub) || meta.Password == password +} + +func MetaCoversPath(metaPath, reqPath string, applyToSubFolder bool) bool { + metaPath = utils.FixAndCleanPath(metaPath) + reqPath = utils.FixAndCleanPath(reqPath) + if strings.EqualFold(metaPath, reqPath) { + return true + } + if !applyToSubFolder { + return false + } + for reqPath != "/" { + reqPath = path.Dir(reqPath) + if strings.EqualFold(metaPath, reqPath) { + return true + } + } + return false +} diff --git a/internal/authz/policy_test.go b/internal/authz/policy_test.go new file mode 100644 index 0000000000..9ef1da6c9c --- /dev/null +++ b/internal/authz/policy_test.go @@ -0,0 +1,227 @@ +package authz + +import ( + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/model" +) + +func TestMetaCoversPath(t *testing.T) { + tests := []struct { + name, metaPath, reqPath string + sub, want bool + }{ + {"exact", "/folder", "/folder", false, true}, + {"exact case insensitive", "/Folder", "/folder", false, true}, + {"child enabled", "/folder", "/folder/child", true, true}, + {"deep child case insensitive", "/Folder", "/folder/child/file", true, true}, + {"child disabled", "/folder", "/folder/child", false, false}, + {"sibling", "/folder", "/other", true, false}, + {"prefix sibling", "/folder", "/folder-name", true, false}, + {"root subtree", "/", "/any/deep/path", true, true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := MetaCoversPath(tt.metaPath, tt.reqPath, tt.sub); got != tt.want { + t.Fatalf("MetaCoversPath(%q, %q, %v) = %v, want %v", tt.metaPath, tt.reqPath, tt.sub, got, tt.want) + } + }) + } +} + +func TestReadAndWriteRestrictions(t *testing.T) { + tests := []struct { + name string + user *model.User + meta *model.Meta + path string + want bool + }{ + {"system context", nil, restrictedMeta(false), "/folder", true}, + {"no meta", &model.User{ID: 5}, nil, "/folder", true}, + {"unrestricted", &model.User{ID: 5}, &model.Meta{Path: "/folder"}, "/folder", true}, + {"listed exact", &model.User{ID: 1}, restrictedMeta(false), "/folder", true}, + {"unlisted exact", &model.User{ID: 5}, restrictedMeta(false), "/folder", false}, + {"unlisted child not inherited", &model.User{ID: 5}, restrictedMeta(false), "/folder/child", true}, + {"unlisted child inherited", &model.User{ID: 5}, restrictedMeta(true), "/folder/child", false}, + {"listed deep child", &model.User{ID: 1}, restrictedMeta(true), "/folder/child/file", true}, + {"restriction path is case insensitive", &model.User{ID: 5}, restrictedMeta(true), "/Folder/child", false}, + {"unrelated path", &model.User{ID: 5}, restrictedMeta(true), "/other", true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := CanRead(tt.user, tt.meta, tt.path); got != tt.want { + t.Errorf("CanRead() = %v, want %v", got, tt.want) + } + if got := CanWrite(tt.user, tt.meta, tt.path); got != tt.want { + t.Errorf("CanWrite() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestCanWriteContentBypassUserPerms(t *testing.T) { + tests := []struct { + name string + meta *model.Meta + path string + want bool + }{ + {"no meta", nil, "/folder", false}, + {"disabled", &model.Meta{Path: "/folder"}, "/folder", false}, + {"exact", &model.Meta{Path: "/folder", Write: true}, "/folder", true}, + {"child enabled", &model.Meta{Path: "/folder", Write: true, WSub: true}, "/folder/child", true}, + {"child disabled", &model.Meta{Path: "/folder", Write: true}, "/folder/child", false}, + {"unrelated", &model.Meta{Path: "/folder", Write: true, WSub: true}, "/other", false}, + {"case mismatch does not grant", &model.Meta{Path: "/Folder", Write: true, WSub: true}, "/folder/child", false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := CanWriteContentBypassUserPerms(tt.meta, tt.path); got != tt.want { + t.Fatalf("CanWriteContentBypassUserPerms() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestIsHidden(t *testing.T) { + hideMeta := func(sub bool) *model.Meta { + return &model.Meta{Path: "/folder", Hide: "^secret$\n\\.private$", HSub: sub} + } + tests := []struct { + name string + user *model.User + meta *model.Meta + path string + want bool + }{ + {"system context", nil, hideMeta(true), "/folder/secret", false}, + {"hide-capable user", &model.User{Permission: 1}, hideMeta(true), "/folder/secret", false}, + {"no meta", &model.User{}, nil, "/folder/secret", false}, + {"empty patterns", &model.User{}, &model.Meta{Path: "/folder", HSub: true}, "/folder/secret", false}, + {"matching child", &model.User{}, hideMeta(false), "/folder/secret", true}, + {"nonmatching child", &model.User{}, hideMeta(false), "/folder/public", false}, + {"second pattern", &model.User{}, hideMeta(false), "/folder/file.private", true}, + {"nested enabled", &model.User{}, hideMeta(true), "/folder/child/secret", true}, + {"nested disabled", &model.User{}, hideMeta(false), "/folder/child/secret", false}, + {"unrelated parent", &model.User{}, hideMeta(true), "/other/secret", false}, + {"coverage is case insensitive", &model.User{}, hideMeta(true), "/Folder/secret", true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := IsHidden(tt.user, tt.meta, tt.path); got != tt.want { + t.Fatalf("IsHidden() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestFilterHidden(t *testing.T) { + objs := []model.Obj{ + &model.Object{Name: "public"}, + &model.Object{Name: "secret"}, + &model.Object{Name: "file.private"}, + } + meta := &model.Meta{Path: "/folder", Hide: "^secret$\n\\.private$"} + + got := FilterHidden(&model.User{}, meta, "/folder", objs) + if len(got) != 1 || got[0].GetName() != "public" { + t.Fatalf("FilterHidden() returned %v, want only public", objectNames(got)) + } +} + +func TestCanAccess(t *testing.T) { + user := func(id uint, permission int32) *model.User { + return &model.User{ID: id, Permission: permission} + } + tests := []struct { + name string + user *model.User + meta *model.Meta + path string + password string + want bool + }{ + {"no policy", user(1, 0), nil, "/file", "", true}, + {"hidden name", user(1, 0), &model.Meta{Path: "/", Hide: "^secret$", HSub: true}, "/secret", "", false}, + {"hidden permission bypass", user(1, 1), &model.Meta{Path: "/", Hide: "^secret$", HSub: true}, "/secret", "", true}, + {"read restriction precedes password", user(5, 0), accessMeta(), "/folder/file", "secret", false}, + {"correct password", user(1, 0), accessMeta(), "/folder/file", "secret", true}, + {"wrong password", user(1, 0), accessMeta(), "/folder/file", "wrong", false}, + {"password permission bypass", user(1, 2), accessMeta(), "/folder/file", "wrong", true}, + {"password not inherited", user(1, 0), &model.Meta{Path: "/folder", Password: "secret"}, "/folder/file", "wrong", true}, + {"case-insensitive exact password coverage", user(1, 0), &model.Meta{Path: "/Folder", Password: "secret"}, "/folder", "wrong", false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := CanAccess(tt.user, tt.meta, tt.path, tt.password); got != tt.want { + t.Fatalf("CanAccess() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestWritePermissionCombination(t *testing.T) { + tests := []struct { + name string + user *model.User + meta *model.Meta + path string + want bool + }{ + {"user permission and listed", &model.User{ID: 1, Permission: 1 << 3}, writeMeta(false, []uint{1}, false), "/folder", true}, + {"user permission but unlisted", &model.User{ID: 5, Permission: 1 << 3}, writeMeta(false, []uint{1}, false), "/folder", false}, + {"meta bypass and listed", &model.User{ID: 1}, writeMeta(true, []uint{1}, false), "/folder", true}, + {"meta bypass but unlisted", &model.User{ID: 5}, writeMeta(true, []uint{1}, false), "/folder", false}, + {"no permission or bypass", &model.User{ID: 1}, writeMeta(false, []uint{1}, false), "/folder", false}, + {"inherited bypass and whitelist", &model.User{ID: 1}, writeMeta(true, []uint{1}, true), "/folder/child", true}, + {"non-inherited bypass", &model.User{ID: 1}, writeMeta(true, []uint{1}, false), "/folder/child", false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := (tt.user.CanWriteContent() || CanWriteContentBypassUserPerms(tt.meta, tt.path)) && + CanWrite(tt.user, tt.meta, tt.path) + if got != tt.want { + t.Fatalf("combined write permission = %v, want %v", got, tt.want) + } + }) + } +} + +func restrictedMeta(sub bool) *model.Meta { + return &model.Meta{ + Path: "/folder", + ReadUsers: []uint{1, 2}, + ReadUsersSub: sub, + WriteUsers: []uint{1, 2}, + WriteUsersSub: sub, + } +} + +func accessMeta() *model.Meta { + return &model.Meta{ + Path: "/folder", + ReadUsers: []uint{1, 2}, + ReadUsersSub: true, + Password: "secret", + PSub: true, + } +} + +func writeMeta(write bool, users []uint, sub bool) *model.Meta { + return &model.Meta{ + Path: "/folder", + Write: write, + WSub: sub, + WriteUsers: users, + WriteUsersSub: sub, + } +} + +func objectNames(objs []model.Obj) []string { + names := make([]string, len(objs)) + for i, obj := range objs { + names[i] = obj.GetName() + } + return names +} diff --git a/internal/fs/list.go b/internal/fs/list.go index 113ba8231a..1d70c96a00 100644 --- a/internal/fs/list.go +++ b/internal/fs/list.go @@ -2,14 +2,15 @@ package fs import ( "context" + "path" + + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/internal/op" - "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/pkg/errors" log "github.com/sirupsen/logrus" - "path" ) // List files @@ -40,10 +41,8 @@ func list(ctx context.Context, path string, args *ListArgs) ([]model.Obj, error) } om := model.NewObjMerge() - if whetherHide(user, meta, path) { - om.InitHideReg(meta.Hide) - } objs := om.Merge(_objs, virtualFiles...) + objs = authz.FilterHidden(user, meta, path, objs) objs, err = filterReadableObjs(objs, user, path, meta) return objs, err } @@ -62,30 +61,9 @@ func filterReadableObjs(objs []model.Obj, user *model.User, reqPath string, pare } else { meta = parentMeta } - if common.CanRead(user, meta, objPath) { + if authz.CanRead(user, meta, objPath) { result = append(result, obj) } } return result, nil } - -func whetherHide(user *model.User, meta *model.Meta, path string) bool { - // if is admin, don't hide - if user == nil || user.CanSeeHides() { - return false - } - // if meta is nil, don't hide - if meta == nil { - return false - } - // if meta.Hide is empty, don't hide - if meta.Hide == "" { - return false - } - // if meta doesn't apply to sub_folder, don't hide - if !common.MetaCoversPath(meta.Path, path, meta.HSub) { - return false - } - // if is guest, hide - return true -} diff --git a/internal/fs/list_test.go b/internal/fs/list_test.go deleted file mode 100644 index d8c8e47fef..0000000000 --- a/internal/fs/list_test.go +++ /dev/null @@ -1,151 +0,0 @@ -package fs - -import ( - "testing" - - "github.com/OpenListTeam/OpenList/v4/internal/model" -) - -func TestWhetherHide(t *testing.T) { - tests := []struct { - name string - user *model.User - meta *model.Meta - path string - want bool - reason string - }{ - { - name: "nil user", - user: nil, - meta: &model.Meta{ - Path: "/folder", - Hide: "secret", - HSub: true, - }, - path: "/folder", - want: false, - reason: "nil user (treated as admin) should not hide", - }, - { - name: "user with can_see_hides permission", - user: &model.User{ - Role: model.GENERAL, - Permission: 1, // bit 0 set = can see hides - }, - meta: &model.Meta{ - Path: "/folder", - Hide: "secret", - HSub: true, - }, - path: "/folder", - want: false, - reason: "user with can_see_hides permission should not hide", - }, - { - name: "nil meta", - user: &model.User{ - Role: model.GUEST, - }, - meta: nil, - path: "/folder", - want: false, - reason: "nil meta should not hide", - }, - { - name: "empty hide string", - user: &model.User{ - Role: model.GUEST, - }, - meta: &model.Meta{ - Path: "/folder", - Hide: "", - HSub: true, - }, - path: "/folder", - want: false, - reason: "empty hide string should not hide", - }, - { - name: "exact path match with HSub=false", - user: &model.User{ - Role: model.GUEST, - }, - meta: &model.Meta{ - Path: "/folder", - Hide: "secret", - HSub: false, - }, - path: "/folder", - want: true, - reason: "exact path match should hide for guest", - }, - { - name: "sub path with HSub=true", - user: &model.User{ - Role: model.GUEST, - }, - meta: &model.Meta{ - Path: "/folder", - Hide: "secret", - HSub: true, - }, - path: "/folder/subfolder", - want: true, - reason: "sub path with HSub=true should hide for guest", - }, - { - name: "sub path with HSub=false", - user: &model.User{ - Role: model.GUEST, - }, - meta: &model.Meta{ - Path: "/folder", - Hide: "secret", - HSub: false, - }, - path: "/folder/subfolder", - want: false, - reason: "sub path with HSub=false should not hide", - }, - { - name: "non-sub path with HSub=true", - user: &model.User{ - Role: model.GUEST, - }, - meta: &model.Meta{ - Path: "/folder", - Hide: "secret", - HSub: true, - }, - path: "/other", - want: false, - reason: "non-sub path should not hide even with HSub=true", - }, - { - name: "user without can_see_hides permission", - user: &model.User{ - Role: model.GENERAL, - Permission: 0, // bit 0 not set = cannot see hides - }, - meta: &model.Meta{ - Path: "/folder", - Hide: "secret", - HSub: true, - }, - path: "/folder", - want: true, - reason: "user without can_see_hides permission should hide", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := whetherHide(tt.user, tt.meta, tt.path) - if got != tt.want { - t.Errorf("whetherHide() = %v, want %v\nReason: %s", - got, tt.want, tt.reason) - } - }) - } -} diff --git a/internal/model/obj.go b/internal/model/obj.go index 1269b5b797..e3cfb422fd 100644 --- a/internal/model/obj.go +++ b/internal/model/obj.go @@ -3,13 +3,10 @@ package model import ( "io" "sort" - "strings" "time" "github.com/OpenListTeam/OpenList/v4/pkg/http_range" "github.com/OpenListTeam/OpenList/v4/pkg/utils" - "github.com/dlclark/regexp2" - mapset "github.com/deckarep/golang-set/v2" "github.com/maruel/natural" @@ -204,8 +201,7 @@ func NewObjMerge() *ObjMerge { } type ObjMerge struct { - regs []*regexp2.Regexp - set mapset.Set[string] + set mapset.Set[string] } func (om *ObjMerge) Merge(objs []Obj, objs_ ...Obj) []Obj { @@ -224,22 +220,9 @@ func (om *ObjMerge) insertObjs(objs []Obj, objs_ ...Obj) []Obj { } func (om *ObjMerge) clickObj(obj Obj) bool { - for _, reg := range om.regs { - if isMatch, _ := reg.MatchString(obj.GetName()); isMatch { - return false - } - } return om.set.Add(obj.GetName()) } -func (om *ObjMerge) InitHideReg(hides string) { - rs := strings.Split(hides, "\n") - om.regs = make([]*regexp2.Regexp, 0, len(rs)) - for _, r := range rs { - om.regs = append(om.regs, regexp2.MustCompile(r, regexp2.None)) - } -} - func (om *ObjMerge) Reset() { om.set.Clear() } diff --git a/server/common/check.go b/server/common/check.go index 728897013a..7f8610e225 100644 --- a/server/common/check.go +++ b/server/common/check.go @@ -1,16 +1,10 @@ package common import ( - "path" - "slices" - "strings" - "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/driver" - "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/internal/op" "github.com/OpenListTeam/OpenList/v4/pkg/utils" - "github.com/dlclark/regexp2" ) func IsStorageSignEnabled(rawPath string) bool { @@ -18,86 +12,6 @@ func IsStorageSignEnabled(rawPath string) bool { return storage != nil && storage.GetStorage().EnableSign } -func CanRead(user *model.User, meta *model.Meta, path string) bool { - // nil user is treated as internal/system context and bypasses per-user read restrictions - if user == nil { - return true - } - if meta != nil && len(meta.ReadUsers) > 0 && !slices.Contains(meta.ReadUsers, user.ID) && MetaCoversPath(meta.Path, path, meta.ReadUsersSub) { - return false - } - return true -} - -func CanWrite(user *model.User, meta *model.Meta, path string) bool { - // nil user is treated as internal/system context and bypasses per-user write restrictions - if user == nil { - return true - } - if meta != nil && len(meta.WriteUsers) > 0 && !slices.Contains(meta.WriteUsers, user.ID) && MetaCoversPath(meta.Path, path, meta.WriteUsersSub) { - return false - } - return true -} - -func CanWriteContentBypassUserPerms(meta *model.Meta, path string) bool { - if meta == nil || !meta.Write { - return false - } - if utils.PathEqual(meta.Path, path) { - return true - } - return utils.IsSubPath(meta.Path, path) && meta.WSub -} - -func CanAccess(user *model.User, meta *model.Meta, reqPath string, password string) bool { - // if the reqPath is in hide (only can check the nearest meta) and user can't see hides, can't access - if meta != nil && !user.CanSeeHides() && meta.Hide != "" && - MetaCoversPath(meta.Path, path.Dir(reqPath), meta.HSub) { // the meta should apply to the parent of current path - for hide := range strings.SplitSeq(meta.Hide, "\n") { - re := regexp2.MustCompile(hide, regexp2.None) - if isMatch, _ := re.MatchString(path.Base(reqPath)); isMatch { - return false - } - } - } - if !CanRead(user, meta, reqPath) { - return false - } - // if is not guest and can access without password - if user.CanAccessWithoutPassword() { - return true - } - // if meta is nil or password is empty, can access - if meta == nil || meta.Password == "" { - return true - } - // if meta doesn't apply to sub_folder, can access - if !MetaCoversPath(meta.Path, reqPath, meta.PSub) { - return true - } - // validate password - return meta.Password == password -} - -func MetaCoversPath(metaPath, reqPath string, applyToSubFolder bool) bool { - metaPath = utils.FixAndCleanPath(metaPath) - reqPath = utils.FixAndCleanPath(reqPath) - if strings.EqualFold(metaPath, reqPath) { - return true - } - if !applyToSubFolder { - return false - } - for reqPath != "/" { - reqPath = path.Dir(reqPath) - if strings.EqualFold(metaPath, reqPath) { - return true - } - } - return false -} - // ShouldProxy TODO need optimize // when should be proxy? // 1. config.MustProxy() diff --git a/server/common/check_test.go b/server/common/check_test.go deleted file mode 100644 index 631cf26671..0000000000 --- a/server/common/check_test.go +++ /dev/null @@ -1,1035 +0,0 @@ -package common - -import ( - "testing" - - "github.com/OpenListTeam/OpenList/v4/internal/model" -) - -func TestCoversPath(t *testing.T) { - tests := []struct { - name string - metaPath string - reqPath string - applySub bool - want bool - }{ - { - name: "exact path match with applySub=false", - metaPath: "/folder", - reqPath: "/folder", - applySub: false, - want: true, - }, - { - name: "exact path match with applySub=true", - metaPath: "/folder", - reqPath: "/folder", - applySub: true, - want: true, - }, - { - name: "sub path with applySub=true", - metaPath: "/folder", - reqPath: "/folder/subfolder", - applySub: true, - want: true, - }, - { - name: "sub path with applySub=false", - metaPath: "/folder", - reqPath: "/folder/subfolder", - applySub: false, - want: false, - }, - { - name: "non-sub path with applySub=true", - metaPath: "/folder", - reqPath: "/other", - applySub: true, - want: false, - }, - { - name: "non-sub path with applySub=false", - metaPath: "/folder", - reqPath: "/other", - applySub: false, - want: false, - }, - { - name: "root path covers all with applySub=true", - metaPath: "/", - reqPath: "/any/deep/path", - applySub: true, - want: true, - }, - { - name: "root path exact match", - metaPath: "/", - reqPath: "/", - applySub: false, - want: true, - }, - { - name: "deep sub path with applySub=true", - metaPath: "/folder", - reqPath: "/folder/sub1/sub2/file.txt", - applySub: true, - want: true, - }, - { - name: "sibling paths with applySub=true", - metaPath: "/folder1", - reqPath: "/folder2", - applySub: true, - want: false, - }, - { - name: "case-insensitive exact path match", - metaPath: "/Folder", - reqPath: "/folder", - applySub: false, - want: true, - }, - { - name: "case-insensitive sub path match", - metaPath: "/Folder", - reqPath: "/folder/Subfolder", - applySub: true, - want: true, - }, - { - name: "case-insensitive sibling prefix does not match", - metaPath: "/Folder", - reqPath: "/folder-name", - applySub: true, - want: false, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := MetaCoversPath(tt.metaPath, tt.reqPath, tt.applySub) - if got != tt.want { - t.Errorf("MetaCoversPath(%q, %q, %v) = %v, want %v", - tt.metaPath, tt.reqPath, tt.applySub, got, tt.want) - } - }) - } -} - -func TestCanWriteContentIgnoringUserPerms(t *testing.T) { - tests := []struct { - name string - meta *model.Meta - path string - want bool - reason string - }{ - { - name: "nil meta", - meta: nil, - path: "/any", - want: false, - reason: "nil meta should deny write", - }, - { - name: "meta.Write=false", - meta: &model.Meta{ - Path: "/folder", - Write: false, - }, - path: "/folder", - want: false, - reason: "Write=false should deny write", - }, - { - name: "exact path match with WSub=false", - meta: &model.Meta{ - Path: "/folder", - Write: true, - WSub: false, - }, - path: "/folder", - want: true, - reason: "exact path match should allow write", - }, - { - name: "sub path with WSub=true", - meta: &model.Meta{ - Path: "/folder", - Write: true, - WSub: true, - }, - path: "/folder/subfolder", - want: true, - reason: "sub path with WSub=true should allow write", - }, - { - name: "sub path with WSub=false (BEHAVIOR CHANGE)", - meta: &model.Meta{ - Path: "/folder", - Write: true, - WSub: false, - }, - path: "/folder/subfolder", - want: false, - reason: "sub path with WSub=false should deny write (fixed bug)", - }, - { - name: "non-sub path with WSub=true", - meta: &model.Meta{ - Path: "/folder", - Write: true, - WSub: true, - }, - path: "/other", - want: false, - reason: "non-sub path should deny write even with WSub=true", - }, - { - name: "case-insensitive match does not broaden write bypass", - meta: &model.Meta{ - Path: "/Folder", - Write: true, - WSub: true, - }, - path: "/folder/subfolder", - want: false, - reason: "case-insensitive matching should only enforce restrictions, not grant write bypass", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := CanWriteContentBypassUserPerms(tt.meta, tt.path) - if got != tt.want { - t.Errorf("CanWriteContentBypassUserPerms() = %v, want %v\nReason: %s", - got, tt.want, tt.reason) - } - }) - } -} - -func TestCanRead(t *testing.T) { - tests := []struct { - name string - user *model.User - meta *model.Meta - path string - want bool - reason string - }{ - { - name: "nil user should allow access", - user: nil, - meta: nil, - path: "/any", - want: true, - reason: "nil user represents internal/system context and bypasses per-user read restrictions", - }, - { - name: "nil meta should allow access", - user: &model.User{ - ID: 1, - }, - meta: nil, - path: "/any", - want: true, - reason: "nil meta means no restrictions", - }, - { - name: "empty ReadUsers list should allow access", - user: &model.User{ - ID: 1, - }, - meta: &model.Meta{ - Path: "/folder", - ReadUsers: []uint{}, - }, - path: "/folder", - want: true, - reason: "empty ReadUsers means no user-level restrictions", - }, - { - name: "user in ReadUsers list with exact path match", - user: &model.User{ - ID: 1, - }, - meta: &model.Meta{ - Path: "/folder", - ReadUsers: []uint{1, 2, 3}, - ReadUsersSub: false, - }, - path: "/folder", - want: true, - reason: "user ID 1 is in ReadUsers list", - }, - { - name: "user not in ReadUsers list with exact path match", - user: &model.User{ - ID: 5, - }, - meta: &model.Meta{ - Path: "/folder", - ReadUsers: []uint{1, 2, 3}, - ReadUsersSub: false, - }, - path: "/folder", - want: false, - reason: "user ID 5 is not in ReadUsers list and path matches", - }, - { - name: "user not in ReadUsers list with ReadUsersSub=true for sub path", - user: &model.User{ - ID: 5, - }, - meta: &model.Meta{ - Path: "/folder", - ReadUsers: []uint{1, 2, 3}, - ReadUsersSub: true, - }, - path: "/folder/subfolder", - want: false, - reason: "user ID 5 is not in ReadUsers list and ReadUsersSub applies to sub paths", - }, - { - name: "user not in ReadUsers list with ReadUsersSub=false for sub path", - user: &model.User{ - ID: 5, - }, - meta: &model.Meta{ - Path: "/folder", - ReadUsers: []uint{1, 2, 3}, - ReadUsersSub: false, - }, - path: "/folder/subfolder", - want: true, - reason: "ReadUsersSub=false means restriction doesn't apply to sub paths", - }, - { - name: "user in ReadUsers list with ReadUsersSub=true for sub path", - user: &model.User{ - ID: 2, - }, - meta: &model.Meta{ - Path: "/folder", - ReadUsers: []uint{1, 2, 3}, - ReadUsersSub: true, - }, - path: "/folder/subfolder/deep", - want: true, - reason: "user ID 2 is in ReadUsers list so can access sub paths", - }, - { - name: "user not in ReadUsers list for different path", - user: &model.User{ - ID: 5, - }, - meta: &model.Meta{ - Path: "/folder", - ReadUsers: []uint{1, 2, 3}, - ReadUsersSub: false, - }, - path: "/other", - want: true, - reason: "meta path doesn't match request path, so restriction doesn't apply", - }, - { - name: "root level restriction with ReadUsersSub=true", - user: &model.User{ - ID: 5, - }, - meta: &model.Meta{ - Path: "/", - ReadUsers: []uint{1, 2, 3}, - ReadUsersSub: true, - }, - path: "/any/deep/path", - want: false, - reason: "root level restriction with ReadUsersSub affects all paths", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := CanRead(tt.user, tt.meta, tt.path) - if got != tt.want { - t.Errorf("CanRead() = %v, want %v\nReason: %s\nUser ID: %v, Meta: %+v, Path: %s", - got, tt.want, tt.reason, getUserID(tt.user), tt.meta, tt.path) - } - }) - } -} - -func TestCanWrite(t *testing.T) { - tests := []struct { - name string - user *model.User - meta *model.Meta - path string - want bool - reason string - }{ - { - name: "nil user should allow access", - user: nil, - meta: nil, - path: "/any", - want: true, - reason: "nil user represents internal/system context and bypasses per-user write restrictions", - }, - { - name: "nil meta should allow access", - user: &model.User{ - ID: 1, - }, - meta: nil, - path: "/any", - want: true, - reason: "nil meta means no restrictions", - }, - { - name: "empty WriteUsers list should allow access", - user: &model.User{ - ID: 1, - }, - meta: &model.Meta{ - Path: "/folder", - WriteUsers: []uint{}, - }, - path: "/folder", - want: true, - reason: "empty WriteUsers means no user-level restrictions", - }, - { - name: "user in WriteUsers list with exact path match", - user: &model.User{ - ID: 1, - }, - meta: &model.Meta{ - Path: "/folder", - WriteUsers: []uint{1, 2, 3}, - WriteUsersSub: false, - }, - path: "/folder", - want: true, - reason: "user ID 1 is in WriteUsers list", - }, - { - name: "user not in WriteUsers list with exact path match", - user: &model.User{ - ID: 5, - }, - meta: &model.Meta{ - Path: "/folder", - WriteUsers: []uint{1, 2, 3}, - WriteUsersSub: false, - }, - path: "/folder", - want: false, - reason: "user ID 5 is not in WriteUsers list and path matches", - }, - { - name: "user not in WriteUsers list with WriteUsersSub=true for sub path", - user: &model.User{ - ID: 5, - }, - meta: &model.Meta{ - Path: "/folder", - WriteUsers: []uint{1, 2, 3}, - WriteUsersSub: true, - }, - path: "/folder/subfolder", - want: false, - reason: "user ID 5 is not in WriteUsers list and WriteUsersSub applies to sub paths", - }, - { - name: "user not in WriteUsers list with WriteUsersSub=false for sub path", - user: &model.User{ - ID: 5, - }, - meta: &model.Meta{ - Path: "/folder", - WriteUsers: []uint{1, 2, 3}, - WriteUsersSub: false, - }, - path: "/folder/subfolder", - want: true, - reason: "WriteUsersSub=false means restriction doesn't apply to sub paths", - }, - { - name: "user in WriteUsers list with WriteUsersSub=true for sub path", - user: &model.User{ - ID: 2, - }, - meta: &model.Meta{ - Path: "/folder", - WriteUsers: []uint{1, 2, 3}, - WriteUsersSub: true, - }, - path: "/folder/subfolder/deep", - want: true, - reason: "user ID 2 is in WriteUsers list so can write to sub paths", - }, - { - name: "user not in WriteUsers list for different path", - user: &model.User{ - ID: 5, - }, - meta: &model.Meta{ - Path: "/folder", - WriteUsers: []uint{1, 2, 3}, - WriteUsersSub: false, - }, - path: "/other", - want: true, - reason: "meta path doesn't match request path, so restriction doesn't apply", - }, - { - name: "multiple users with mixed permissions", - user: &model.User{ - ID: 10, - }, - meta: &model.Meta{ - Path: "/folder", - WriteUsers: []uint{1, 5, 10, 15}, - WriteUsersSub: true, - }, - path: "/folder/file.txt", - want: true, - reason: "user ID 10 is in WriteUsers list", - }, - { - name: "write restriction at root level", - user: &model.User{ - ID: 5, - }, - meta: &model.Meta{ - Path: "/", - WriteUsers: []uint{1}, - WriteUsersSub: true, - }, - path: "/any/path", - want: false, - reason: "only user ID 1 can write when root has WriteUsers restriction", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := CanWrite(tt.user, tt.meta, tt.path) - if got != tt.want { - t.Errorf("CanWrite() = %v, want %v\nReason: %s\nUser ID: %v, Meta: %+v, Path: %s", - got, tt.want, tt.reason, getUserID(tt.user), tt.meta, tt.path) - } - }) - } -} - -func TestCanAccessWithReadPermissions(t *testing.T) { - tests := []struct { - name string - user *model.User - meta *model.Meta - reqPath string - password string - want bool - reason string - }{ - { - name: "user with read permission and correct password", - user: &model.User{ - ID: 1, - Role: model.GENERAL, - Permission: 0, - }, - meta: &model.Meta{ - Path: "/folder", - ReadUsers: []uint{1, 2}, - ReadUsersSub: true, - Password: "secret", - PSub: true, - }, - reqPath: "/folder/file.txt", - password: "secret", - want: true, - reason: "user in ReadUsers list with correct password", - }, - { - name: "user without read permission even with correct password", - user: &model.User{ - ID: 5, - Role: model.GENERAL, - Permission: 0, - }, - meta: &model.Meta{ - Path: "/folder", - ReadUsers: []uint{1, 2}, - ReadUsersSub: true, - Password: "secret", - PSub: true, - }, - reqPath: "/folder/file.txt", - password: "secret", - want: false, - reason: "user not in ReadUsers list, should be denied before password check", - }, - { - name: "user with read permission but wrong password", - user: &model.User{ - ID: 1, - Role: model.GENERAL, - Permission: 0, - }, - meta: &model.Meta{ - Path: "/folder", - ReadUsers: []uint{1, 2}, - ReadUsersSub: true, - Password: "secret", - PSub: true, - }, - reqPath: "/folder/file.txt", - password: "wrong", - want: false, - reason: "user in ReadUsers list but wrong password", - }, - { - name: "user without read permission and no password", - user: &model.User{ - ID: 5, - Role: model.GENERAL, - Permission: 0, - }, - meta: &model.Meta{ - Path: "/folder", - ReadUsers: []uint{1, 2}, - ReadUsersSub: true, - }, - reqPath: "/folder/file.txt", - password: "", - want: false, - reason: "user not in ReadUsers list should be denied", - }, - { - name: "case-insensitive exact path still requires password", - user: &model.User{ - ID: 1, - Role: model.GENERAL, - Permission: 0, - }, - meta: &model.Meta{ - Path: "/Folder", - Password: "secret", - PSub: false, - }, - reqPath: "/folder", - password: "wrong", - want: false, - reason: "changing only path casing must not bypass an exact-path password", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := CanAccess(tt.user, tt.meta, tt.reqPath, tt.password) - if got != tt.want { - t.Errorf("CanAccess() = %v, want %v\nReason: %s", - got, tt.want, tt.reason) - } - }) - } -} - -// Helper function to safely get user ID -func getUserID(user *model.User) uint { - if user == nil { - return 0 - } - return user.ID -} - -// TestWritePermissionCombinations tests the combined permission check logic -// that is actually used in the codebase: -// -// if !user.CanWriteContent() && !CanWriteContentBypassUserPerms(meta, path) { -// deny -// } -// if !CanWrite(user, meta, path) { -// deny -// } -// -// This ensures the three-layer permission system works correctly: -// 1. User-level global write permission (CanWriteContent) -// 2. Meta-level global write permission (CanWriteContentBypassUserPerms) -// 3. Meta-level user whitelist (CanWrite) -func TestWritePermissionCombinations(t *testing.T) { - tests := []struct { - name string - user *model.User - meta *model.Meta - path string - want bool - reason string - checkFirstLayer bool // whether first layer should pass - checkSecondLayer bool // whether second layer should pass - expectedDenyReason string - }{ - // === Scenario 1: User has global write permission === - { - name: "user has CanWriteContent + in WriteUsers whitelist", - user: &model.User{ - ID: 1, - Permission: 1 << 3, // CanWriteContent = true - }, - meta: &model.Meta{ - Path: "/folder", - Write: false, - WriteUsers: []uint{1}, - WriteUsersSub: false, - }, - path: "/folder", - want: true, - reason: "user has global write permission AND is in whitelist", - checkFirstLayer: true, - checkSecondLayer: true, - expectedDenyReason: "", - }, - { - name: "user has CanWriteContent but NOT in WriteUsers whitelist", - user: &model.User{ - ID: 1, - Permission: 1 << 3, // CanWriteContent = true - }, - meta: &model.Meta{ - Path: "/folder", - Write: false, - WriteUsers: []uint{2, 3}, // user 1 not in list - WriteUsersSub: false, - }, - path: "/folder", - want: false, - reason: "even with global write permission, must pass whitelist check", - checkFirstLayer: true, - checkSecondLayer: false, - expectedDenyReason: "whitelist check failed", - }, - - // === Scenario 2: User lacks global permission but meta.Write=true === - { - name: "no CanWriteContent + meta.Write=true + in WriteUsers", - user: &model.User{ - ID: 1, - Permission: 0, // CanWriteContent = false - }, - meta: &model.Meta{ - Path: "/folder", - Write: true, // bypass enabled - WSub: false, - WriteUsers: []uint{1}, - WriteUsersSub: false, - }, - path: "/folder", - want: true, - reason: "meta.Write bypasses user permission check, and user is in whitelist", - checkFirstLayer: true, - checkSecondLayer: true, - expectedDenyReason: "", - }, - { - name: "no CanWriteContent + meta.Write=true + NOT in WriteUsers (KEY TEST)", - user: &model.User{ - ID: 5, - Permission: 0, // CanWriteContent = false - }, - meta: &model.Meta{ - Path: "/folder", - Write: true, // bypass enabled - WSub: false, - WriteUsers: []uint{1, 2, 3}, // user 5 not in list - WriteUsersSub: false, - }, - path: "/folder", - want: false, - reason: "CRITICAL: meta.Write cannot bypass whitelist check (new behavior)", - checkFirstLayer: true, - checkSecondLayer: false, - expectedDenyReason: "whitelist check failed even with meta.Write=true", - }, - - // === Scenario 3: Both checks fail === - { - name: "no CanWriteContent + meta.Write=false", - user: &model.User{ - ID: 1, - Permission: 0, // CanWriteContent = false - }, - meta: &model.Meta{ - Path: "/folder", - Write: false, // no bypass - WriteUsers: []uint{1}, - WriteUsersSub: false, - }, - path: "/folder", - want: false, - reason: "denied at first layer: no global permission and no bypass", - checkFirstLayer: false, - checkSecondLayer: false, - expectedDenyReason: "first layer check failed", - }, - - // === Scenario 4: Empty WriteUsers (no whitelist restriction) === - { - name: "user has CanWriteContent + empty WriteUsers", - user: &model.User{ - ID: 1, - Permission: 1 << 3, // CanWriteContent = true - }, - meta: &model.Meta{ - Path: "/folder", - Write: false, - WriteUsers: []uint{}, // empty = no restriction - WriteUsersSub: false, - }, - path: "/folder", - want: true, - reason: "empty WriteUsers means no whitelist restriction", - checkFirstLayer: true, - checkSecondLayer: true, - expectedDenyReason: "", - }, - { - name: "no CanWriteContent + meta.Write=true + empty WriteUsers", - user: &model.User{ - ID: 1, - Permission: 0, - }, - meta: &model.Meta{ - Path: "/folder", - Write: true, - WSub: false, - WriteUsers: []uint{}, // empty = no restriction - WriteUsersSub: false, - }, - path: "/folder", - want: true, - reason: "meta.Write bypasses first check, empty whitelist passes second", - checkFirstLayer: true, - checkSecondLayer: true, - expectedDenyReason: "", - }, - - // === Scenario 5: Nil meta (no restrictions) === - { - name: "user has CanWriteContent + nil meta", - user: &model.User{ - ID: 1, - Permission: 1 << 3, - }, - meta: nil, - path: "/folder", - want: true, - reason: "nil meta means no restrictions", - checkFirstLayer: true, - checkSecondLayer: true, - expectedDenyReason: "", - }, - { - name: "no CanWriteContent + nil meta", - user: &model.User{ - ID: 1, - Permission: 0, - }, - meta: nil, - path: "/folder", - want: false, - reason: "nil meta cannot bypass lack of user permission", - checkFirstLayer: false, - checkSecondLayer: true, // would pass if first layer passed - expectedDenyReason: "first layer check failed", - }, - - // === Scenario 6: Sub-directory inheritance === - { - name: "meta.Write with WSub=true for subdirectory", - user: &model.User{ - ID: 1, - Permission: 0, - }, - meta: &model.Meta{ - Path: "/folder", - Write: true, - WSub: true, // applies to subdirectories - WriteUsers: []uint{1}, - WriteUsersSub: true, - }, - path: "/folder/subfolder", - want: true, - reason: "WSub=true applies meta.Write to subdirectories", - checkFirstLayer: true, - checkSecondLayer: true, - expectedDenyReason: "", - }, - { - name: "meta.Write with WSub=false for subdirectory", - user: &model.User{ - ID: 1, - Permission: 0, - }, - meta: &model.Meta{ - Path: "/folder", - Write: true, - WSub: false, // does NOT apply to subdirectories - WriteUsers: []uint{1}, - WriteUsersSub: false, - }, - path: "/folder/subfolder", - want: false, - reason: "WSub=false means meta.Write doesn't apply to subdirectories", - checkFirstLayer: false, - checkSecondLayer: true, - expectedDenyReason: "first layer check failed (WSub=false)", - }, - { - name: "WriteUsersSub=false for subdirectory bypasses whitelist", - user: &model.User{ - ID: 5, // not in WriteUsers - Permission: 1 << 3, - }, - meta: &model.Meta{ - Path: "/folder", - Write: false, - WriteUsers: []uint{1, 2}, - WriteUsersSub: false, // whitelist does NOT apply to subdirectories - }, - path: "/folder/subfolder", - want: true, - reason: "WriteUsersSub=false means whitelist doesn't apply to subdirectories", - checkFirstLayer: true, - checkSecondLayer: true, // passes because restriction doesn't apply - expectedDenyReason: "", - }, - - // === Scenario 7: Root level restriction === - { - name: "root level meta.Write with user in whitelist", - user: &model.User{ - ID: 1, - Permission: 0, - }, - meta: &model.Meta{ - Path: "/", - Write: true, - WSub: true, - WriteUsers: []uint{1}, - WriteUsersSub: true, - }, - path: "/any/deep/path", - want: true, - reason: "root level permissions apply to all paths", - checkFirstLayer: true, - checkSecondLayer: true, - expectedDenyReason: "", - }, - { - name: "root level restriction denies non-whitelisted user", - user: &model.User{ - ID: 5, - Permission: 1 << 3, // has global permission - }, - meta: &model.Meta{ - Path: "/", - Write: false, - WriteUsers: []uint{1, 2}, - WriteUsersSub: true, - }, - path: "/any/path", - want: false, - reason: "root level whitelist restricts all paths", - checkFirstLayer: true, - checkSecondLayer: false, - expectedDenyReason: "not in root level whitelist", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - // Simulate the actual permission check logic - firstLayerPass := tt.user.CanWriteContent() || CanWriteContentBypassUserPerms(tt.meta, tt.path) - secondLayerPass := CanWrite(tt.user, tt.meta, tt.path) - - // Verify our understanding of each layer - if firstLayerPass != tt.checkFirstLayer { - t.Errorf("First layer check mismatch: got %v, expected %v\n"+ - "CanWriteContent()=%v, CanWriteContentBypassUserPerms()=%v", - firstLayerPass, tt.checkFirstLayer, - tt.user.CanWriteContent(), CanWriteContentBypassUserPerms(tt.meta, tt.path)) - } - - if firstLayerPass && secondLayerPass != tt.checkSecondLayer { - t.Errorf("Second layer check mismatch: got %v, expected %v\n"+ - "CanWrite()=%v", - secondLayerPass, tt.checkSecondLayer, - CanWrite(tt.user, tt.meta, tt.path)) - } - - // Final result - got := firstLayerPass && secondLayerPass - - if got != tt.want { - t.Errorf("Permission check failed:\n"+ - " Result: %v, want %v\n"+ - " Reason: %s\n"+ - " First layer (CanWriteContent || CanWriteContentBypassUserPerms): %v\n"+ - " Second layer (CanWrite): %v\n"+ - " User: ID=%d, Permission=%d, CanWriteContent=%v\n"+ - " Meta: Path=%s, Write=%v, WSub=%v, WriteUsers=%v, WriteUsersSub=%v\n"+ - " Check Path: %s", - got, tt.want, - tt.reason, - firstLayerPass, - secondLayerPass, - tt.user.ID, tt.user.Permission, tt.user.CanWriteContent(), - getMetaPath(tt.meta), getMetaWrite(tt.meta), getMetaWSub(tt.meta), - getMetaWriteUsers(tt.meta), getMetaWriteUsersSub(tt.meta), - tt.path) - } - }) - } -} - -// Helper functions to safely extract meta fields -func getMetaPath(meta *model.Meta) string { - if meta == nil { - return "nil" - } - return meta.Path -} - -func getMetaWrite(meta *model.Meta) bool { - if meta == nil { - return false - } - return meta.Write -} - -func getMetaWSub(meta *model.Meta) bool { - if meta == nil { - return false - } - return meta.WSub -} - -func getMetaWriteUsers(meta *model.Meta) []uint { - if meta == nil { - return nil - } - return meta.WriteUsers -} - -func getMetaWriteUsersSub(meta *model.Meta) bool { - if meta == nil { - return false - } - return meta.WriteUsersSub -} diff --git a/server/ftp/fsmanage.go b/server/ftp/fsmanage.go index 3e98d6d14d..4993f69ab0 100644 --- a/server/ftp/fsmanage.go +++ b/server/ftp/fsmanage.go @@ -4,12 +4,12 @@ import ( "context" stdpath "path" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/fs" "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/internal/op" - "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/pkg/errors" ) @@ -27,10 +27,10 @@ func Mkdir(ctx context.Context, path string) error { if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return err } - if !user.CanWriteContent() && !common.CanWriteContentBypassUserPerms(parentMeta, parentPath) { + if !user.CanWriteContent() && !authz.CanWriteContentBypassUserPerms(parentMeta, parentPath) { return errs.PermissionDenied } - if !common.CanWrite(user, parentMeta, parentPath) { + if !authz.CanWrite(user, parentMeta, parentPath) { return errs.PermissionDenied } return fs.MakeDir(ctx, reqPath) @@ -49,7 +49,7 @@ func Remove(ctx context.Context, path string) error { if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return err } - if !common.CanWrite(user, meta, reqPath) { + if !authz.CanWrite(user, meta, reqPath) { return errs.PermissionDenied } if err = RemoveStage(reqPath); !errors.Is(err, errs.ObjectNotFound) { @@ -75,7 +75,7 @@ func Rename(ctx context.Context, oldPath, newPath string) error { return err } if srcDir == dstDir { - if !user.CanRename() || !user.CanFTPManage() || !common.CanWrite(user, dstMeta, dstDir) { + if !user.CanRename() || !user.CanFTPManage() || !authz.CanWrite(user, dstMeta, dstDir) { return errs.PermissionDenied } if err = MoveStage(srcPath, dstPath); !errors.Is(err, errs.ObjectNotFound) { @@ -87,7 +87,7 @@ func Rename(ctx context.Context, oldPath, newPath string) error { if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return err } - if !user.CanMove() || !user.CanFTPManage() || (srcBase != dstBase && !user.CanRename()) || !common.CanWrite(user, srcMeta, srcDir) || !common.CanWrite(user, dstMeta, dstDir) { + if !user.CanMove() || !user.CanFTPManage() || (srcBase != dstBase && !user.CanRename()) || !authz.CanWrite(user, srcMeta, srcDir) || !authz.CanWrite(user, dstMeta, dstDir) { return errs.PermissionDenied } if err = MoveStage(srcPath, dstPath); !errors.Is(err, errs.ObjectNotFound) { diff --git a/server/ftp/fsread.go b/server/ftp/fsread.go index 54a3de8f2c..31825e350b 100644 --- a/server/ftp/fsread.go +++ b/server/ftp/fsread.go @@ -8,13 +8,13 @@ import ( "os" "time" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/fs" "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/internal/op" "github.com/OpenListTeam/OpenList/v4/internal/stream" - "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/pkg/errors" ) @@ -31,7 +31,7 @@ func OpenDownload(ctx context.Context, reqPath string, offset int64) (*FileDownl return nil, err } ctx = context.WithValue(ctx, conf.MetaKey, meta) - if !common.CanAccess(user, meta, reqPath, ctx.Value(conf.MetaPassKey).(string)) { + if !authz.CanAccess(user, meta, reqPath, ctx.Value(conf.MetaPassKey).(string)) { return nil, errs.PermissionDenied } @@ -123,7 +123,7 @@ func Stat(ctx context.Context, path string) (os.FileInfo, error) { return nil, err } ctx = context.WithValue(ctx, conf.MetaKey, meta) - if !common.CanAccess(user, meta, reqPath, ctx.Value(conf.MetaPassKey).(string)) { + if !authz.CanAccess(user, meta, reqPath, ctx.Value(conf.MetaPassKey).(string)) { return nil, errs.PermissionDenied } if ret, err := StatStage(reqPath); !errors.Is(err, errs.ObjectNotFound) { @@ -147,7 +147,7 @@ func List(ctx context.Context, path string) ([]os.FileInfo, error) { return nil, err } ctx = context.WithValue(ctx, conf.MetaKey, meta) - if !common.CanAccess(user, meta, reqPath, ctx.Value(conf.MetaPassKey).(string)) { + if !authz.CanAccess(user, meta, reqPath, ctx.Value(conf.MetaPassKey).(string)) { return nil, errs.PermissionDenied } objs, err := fs.List(ctx, reqPath, &fs.ListArgs{}) diff --git a/server/ftp/fsup.go b/server/ftp/fsup.go index 1894fc0ee0..b70cafa1a1 100644 --- a/server/ftp/fsup.go +++ b/server/ftp/fsup.go @@ -10,6 +10,7 @@ import ( stdpath "path" "time" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/fs" @@ -18,7 +19,6 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/setting" "github.com/OpenListTeam/OpenList/v4/internal/stream" "github.com/OpenListTeam/OpenList/v4/pkg/utils" - "github.com/OpenListTeam/OpenList/v4/server/common" ftpserver "github.com/fclairamb/ftpserverlib" "github.com/pkg/errors" ) @@ -41,10 +41,10 @@ func uploadAuth(ctx context.Context, path string) error { if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return err } - if !user.CanWriteContent() && !common.CanWriteContentBypassUserPerms(parentMeta, parentPath) { + if !user.CanWriteContent() && !authz.CanWriteContentBypassUserPerms(parentMeta, parentPath) { return errs.PermissionDenied } - if !common.CanWrite(user, parentMeta, parentPath) { + if !authz.CanWrite(user, parentMeta, parentPath) { return errs.PermissionDenied } return nil diff --git a/server/handles/archive.go b/server/handles/archive.go index 364e93edc5..a60ca71468 100644 --- a/server/handles/archive.go +++ b/server/handles/archive.go @@ -7,6 +7,7 @@ import ( "strings" "github.com/OpenListTeam/OpenList/v4/internal/archive/tool" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/fs" @@ -106,7 +107,7 @@ func FsArchiveMeta(c *gin.Context, req *ArchiveMetaReq, user *model.User) { return } common.GinAppendValues(c, conf.MetaKey, meta) - if !common.CanAccess(user, meta, reqPath, req.Password) { + if !authz.CanAccess(user, meta, reqPath, req.Password) { common.ErrorStrResp(c, "password is incorrect or you have no permission", 403) return } @@ -189,7 +190,7 @@ func FsArchiveList(c *gin.Context, req *ArchiveListReq, user *model.User) { return } common.GinAppendValues(c, conf.MetaKey, meta) - if !common.CanAccess(user, meta, reqPath, req.Password) { + if !authz.CanAccess(user, meta, reqPath, req.Password) { common.ErrorStrResp(c, "password is incorrect or you have no permission", 403) return } @@ -265,7 +266,7 @@ func FsArchiveDecompress(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, dstMeta, dstDir) { + if !authz.CanWrite(user, dstMeta, dstDir) { common.ErrorResp(c, errs.PermissionDenied, 403) return } diff --git a/server/handles/direct_upload.go b/server/handles/direct_upload.go index b016f028c2..47b6c29d24 100644 --- a/server/handles/direct_upload.go +++ b/server/handles/direct_upload.go @@ -4,6 +4,7 @@ import ( "net/url" stdpath "path" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/fs" @@ -56,11 +57,11 @@ func FsGetDirectUploadInfo(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !user.CanWriteContent() && !common.CanWriteContentBypassUserPerms(parentMeta, path) { + if !user.CanWriteContent() && !authz.CanWriteContentBypassUserPerms(parentMeta, path) { common.ErrorResp(c, errs.PermissionDenied, 403) return } - if !common.CanWrite(user, parentMeta, path) { + if !authz.CanWrite(user, parentMeta, path) { common.ErrorResp(c, errs.PermissionDenied, 403) return } diff --git a/server/handles/fsbatch.go b/server/handles/fsbatch.go index e0f98284a3..377863900f 100644 --- a/server/handles/fsbatch.go +++ b/server/handles/fsbatch.go @@ -5,6 +5,7 @@ import ( "regexp" "slices" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/fs" @@ -45,7 +46,7 @@ func FsRecursiveMove(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, srcMeta, srcDir) { + if !authz.CanWrite(user, srcMeta, srcDir) { common.ErrorResp(c, errs.PermissionDenied, 403) return } @@ -61,7 +62,7 @@ func FsRecursiveMove(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, dstMeta, dstDir) { + if !authz.CanWrite(user, dstMeta, dstDir) { common.ErrorResp(c, errs.PermissionDenied, 403) return } @@ -179,7 +180,7 @@ func FsBatchRename(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, meta, reqPath) { + if !authz.CanWrite(user, meta, reqPath) { common.ErrorResp(c, errs.PermissionDenied, 403) return } @@ -237,7 +238,7 @@ func FsRegexRename(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, meta, reqPath) { + if !authz.CanWrite(user, meta, reqPath) { common.ErrorResp(c, errs.PermissionDenied, 403) return } diff --git a/server/handles/fsmanage.go b/server/handles/fsmanage.go index d97d36d24a..c19444cf40 100644 --- a/server/handles/fsmanage.go +++ b/server/handles/fsmanage.go @@ -5,6 +5,7 @@ import ( stdpath "path" "strings" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/fs" @@ -41,11 +42,11 @@ func FsMkdir(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !user.CanWriteContent() && !common.CanWriteContentBypassUserPerms(parentMeta, parentPath) { + if !user.CanWriteContent() && !authz.CanWriteContentBypassUserPerms(parentMeta, parentPath) { common.ErrorResp(c, errs.PermissionDenied, 403) return } - if !common.CanWrite(user, parentMeta, parentPath) { + if !authz.CanWrite(user, parentMeta, parentPath) { common.ErrorResp(c, errs.PermissionDenied, 403) return } @@ -91,7 +92,7 @@ func FsMove(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, srcMeta, srcDir) { + if !authz.CanWrite(user, srcMeta, srcDir) { common.ErrorResp(c, errs.PermissionDenied, 403) return } @@ -105,7 +106,7 @@ func FsMove(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, dstMeta, dstDir) { + if !authz.CanWrite(user, dstMeta, dstDir) { common.ErrorResp(c, errs.PermissionDenied, 403) return } @@ -193,7 +194,7 @@ func FsCopy(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanRead(user, srcMeta, srcDir) { + if !authz.CanRead(user, srcMeta, srcDir) { common.ErrorResp(c, errs.PermissionDenied, 403) return } @@ -207,7 +208,7 @@ func FsCopy(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, dstMeta, dstDir) { + if !authz.CanWrite(user, dstMeta, dstDir) { common.ErrorResp(c, errs.PermissionDenied, 403) return } @@ -307,7 +308,7 @@ func FsRename(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, parentMeta, parentPath) { + if !authz.CanWrite(user, parentMeta, parentPath) { common.ErrorResp(c, errs.PermissionDenied, 403) return } @@ -365,7 +366,7 @@ func FsRemove(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, meta, reqPath) { + if !authz.CanWrite(user, meta, reqPath) { common.ErrorResp(c, errs.PermissionDenied, 403) return } @@ -422,7 +423,7 @@ func FsRemoveEmptyDirectory(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, meta, srcDir) { + if !authz.CanWrite(user, meta, srcDir) { common.ErrorResp(c, errs.PermissionDenied, 403) return } diff --git a/server/handles/fsread.go b/server/handles/fsread.go index 841cb5afb9..ef01b86f3c 100644 --- a/server/handles/fsread.go +++ b/server/handles/fsread.go @@ -6,6 +6,7 @@ import ( "strings" "time" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/fs" @@ -89,11 +90,11 @@ func FsList(c *gin.Context, req *ListReq, user *model.User) { return } common.GinAppendValues(c, conf.MetaKey, meta) - if !common.CanAccess(user, meta, reqPath, req.Password) { + if !authz.CanAccess(user, meta, reqPath, req.Password) { common.ErrorStrResp(c, "password is incorrect or you have no permission", 403) return } - canWriteContentAtPath := common.CanWrite(user, meta, reqPath) && (user.CanWriteContent() || common.CanWriteContentBypassUserPerms(meta, reqPath)) + canWriteContentAtPath := authz.CanWrite(user, meta, reqPath) && (user.CanWriteContent() || authz.CanWriteContentBypassUserPerms(meta, reqPath)) if req.Refresh && !canWriteContentAtPath { common.ErrorStrResp(c, "Refresh without permission", 403) return @@ -119,8 +120,8 @@ func FsList(c *gin.Context, req *ListReq, user *model.User) { Total: int64(total), Readme: getReadme(meta, reqPath), Header: getHeader(meta, reqPath), - Write: common.CanWrite(user, meta, reqPath), - WriteContentBypass: common.CanWriteContentBypassUserPerms(meta, reqPath), + Write: authz.CanWrite(user, meta, reqPath), + WriteContentBypass: authz.CanWriteContentBypassUserPerms(meta, reqPath), Provider: provider, DirectUploadTools: directUploadTools, }) @@ -153,7 +154,7 @@ func FsDirs(c *gin.Context) { return } common.GinAppendValues(c, conf.MetaKey, meta) - if !common.CanAccess(user, meta, reqPath, req.Password) { + if !authz.CanAccess(user, meta, reqPath, req.Password) { common.ErrorStrResp(c, "password is incorrect or you have no permission", 403) return } @@ -185,14 +186,14 @@ func filterDirs(objs []model.Obj) []DirResp { } func getReadme(meta *model.Meta, path string) string { - if meta != nil && common.MetaCoversPath(meta.Path, path, meta.RSub) { + if meta != nil && authz.MetaCoversPath(meta.Path, path, meta.RSub) { return meta.Readme } return "" } func getHeader(meta *model.Meta, path string) string { - if meta != nil && common.MetaCoversPath(meta.Path, path, meta.HeaderSub) { + if meta != nil && authz.MetaCoversPath(meta.Path, path, meta.HeaderSub) { return meta.Header } return "" @@ -205,7 +206,7 @@ func isEncrypt(meta *model.Meta, path string) bool { if meta == nil || meta.Password == "" { return false } - if !common.MetaCoversPath(meta.Path, path, meta.PSub) { + if !authz.MetaCoversPath(meta.Path, path, meta.PSub) { return false } return true @@ -292,7 +293,7 @@ func FsGet(c *gin.Context, req *FsGetReq, user *model.User) { return } common.GinAppendValues(c, conf.MetaKey, meta) - if !common.CanAccess(user, meta, reqPath, req.Password) { + if !authz.CanAccess(user, meta, reqPath, req.Password) { common.ErrorStrResp(c, "password is incorrect or you have no permission", 403) return } @@ -418,7 +419,7 @@ func FsOther(c *gin.Context) { return } common.GinAppendValues(c, conf.MetaKey, meta) - if !common.CanAccess(user, meta, req.Path, req.Password) { + if !authz.CanAccess(user, meta, req.Path, req.Password) { common.ErrorStrResp(c, "password is incorrect or you have no permission", 403) return } diff --git a/server/handles/offline_download.go b/server/handles/offline_download.go index 53f30662e4..a14815680d 100644 --- a/server/handles/offline_download.go +++ b/server/handles/offline_download.go @@ -3,6 +3,7 @@ package handles import ( "strings" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/model" @@ -350,7 +351,7 @@ func AddOfflineDownload(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, meta, reqPath) { + if !authz.CanWrite(user, meta, reqPath) { common.ErrorResp(c, errs.PermissionDenied, 403) return } diff --git a/server/handles/search.go b/server/handles/search.go index bbc18cae03..1696576f86 100644 --- a/server/handles/search.go +++ b/server/handles/search.go @@ -3,6 +3,7 @@ package handles import ( "path" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/model" @@ -51,7 +52,7 @@ func Search(c *gin.Context) { if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return false } - return common.CanAccess(user, meta, path.Join(node.Parent, node.Name), req.Password) + return authz.CanAccess(user, meta, path.Join(node.Parent, node.Name), req.Password) }) if err != nil { common.ErrorResp(c, err, 500) diff --git a/server/handles/torrent.go b/server/handles/torrent.go index 8b6ee1b6bf..d05e65b1be 100644 --- a/server/handles/torrent.go +++ b/server/handles/torrent.go @@ -7,6 +7,7 @@ import ( "strings" _189pc "github.com/OpenListTeam/OpenList/v4/drivers/189pc" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/model" @@ -172,7 +173,7 @@ func TorrentRapidUpload(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanWrite(user, meta, reqPath) { + if !authz.CanWrite(user, meta, reqPath) { common.ErrorResp(c, errs.PermissionDenied, 403) return } @@ -346,7 +347,7 @@ func GenerateTorrentForPath(c *gin.Context) { common.ErrorResp(c, err, 500, true) return } - if !common.CanRead(user, meta, reqPath) { + if !authz.CanRead(user, meta, reqPath) { common.ErrorResp(c, errs.PermissionDenied, 403) return } diff --git a/server/mcp/fs_get.go b/server/mcp/fs_get.go index c975695fde..0bce0835a8 100644 --- a/server/mcp/fs_get.go +++ b/server/mcp/fs_get.go @@ -7,6 +7,7 @@ import ( stdpath "path" "strings" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/fs" @@ -49,7 +50,7 @@ func (s *Server) callFSGet(c *gin.Context, raw json.RawMessage) (any, *rpcError) if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return nil, &rpcError{Code: -32603, Message: err.Error()} } - if !common.CanAccess(user, meta, reqPath, args.Password) { + if !authz.CanAccess(user, meta, reqPath, args.Password) { return nil, &rpcError{Code: -32003, Message: "password is incorrect or you have no permission"} } diff --git a/server/mcp/fs_link.go b/server/mcp/fs_link.go index 43ddad8522..f8d18f9a29 100644 --- a/server/mcp/fs_link.go +++ b/server/mcp/fs_link.go @@ -9,6 +9,7 @@ import ( stdpath "path" "time" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/driver" "github.com/OpenListTeam/OpenList/v4/internal/errs" @@ -70,7 +71,7 @@ func (s *Server) callFSLink(c *gin.Context, raw json.RawMessage) (any, *rpcError if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return nil, &rpcError{Code: -32603, Message: err.Error()} } - if !common.CanAccess(user, meta, reqPath, args.Password) { + if !authz.CanAccess(user, meta, reqPath, args.Password) { return nil, &rpcError{Code: -32003, Message: "password is incorrect or you have no permission"} } diff --git a/server/mcp/fs_list.go b/server/mcp/fs_list.go index ac28cf7561..c966e879cd 100644 --- a/server/mcp/fs_list.go +++ b/server/mcp/fs_list.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/fs" @@ -48,12 +49,12 @@ func (s *Server) callFSList(c *gin.Context, raw json.RawMessage) (any, *rpcError if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return nil, &rpcError{Code: -32603, Message: err.Error()} } - if !common.CanAccess(user, meta, reqPath, args.Password) { + if !authz.CanAccess(user, meta, reqPath, args.Password) { return nil, &rpcError{Code: -32003, Message: "password is incorrect or you have no permission"} } - write := common.CanWrite(user, meta, reqPath) - writeContentBypass := common.CanWriteContentBypassUserPerms(meta, reqPath) + write := authz.CanWrite(user, meta, reqPath) + writeContentBypass := authz.CanWriteContentBypassUserPerms(meta, reqPath) canWriteContentAtPath := write && (user.CanWriteContent() || writeContentBypass) if args.Refresh && !canWriteContentAtPath { return nil, &rpcError{Code: -32003, Message: "refresh without permission"} @@ -158,14 +159,14 @@ func toObjResp(objs []model.Obj, parent string, encrypt bool) []handles.ObjResp } func getReadme(meta *model.Meta, path string) string { - if meta != nil && common.MetaCoversPath(meta.Path, path, meta.RSub) { + if meta != nil && authz.MetaCoversPath(meta.Path, path, meta.RSub) { return meta.Readme } return "" } func getHeader(meta *model.Meta, path string) string { - if meta != nil && common.MetaCoversPath(meta.Path, path, meta.HeaderSub) { + if meta != nil && authz.MetaCoversPath(meta.Path, path, meta.HeaderSub) { return meta.Header } return "" @@ -178,5 +179,5 @@ func isEncrypt(meta *model.Meta, path string) bool { if meta == nil || meta.Password == "" { return false } - return common.MetaCoversPath(meta.Path, path, meta.PSub) + return authz.MetaCoversPath(meta.Path, path, meta.PSub) } diff --git a/server/middlewares/down.go b/server/middlewares/down.go index 787936d665..0a1f350c75 100644 --- a/server/middlewares/down.go +++ b/server/middlewares/down.go @@ -3,6 +3,7 @@ package middlewares import ( "strings" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/setting" @@ -60,7 +61,7 @@ func needSign(meta *model.Meta, path string) bool { if meta == nil || meta.Password == "" { return false } - if !meta.PSub && !common.MetaCoversPath(meta.Path, path, false) { + if !meta.PSub && !authz.MetaCoversPath(meta.Path, path, false) { return false } return true diff --git a/server/middlewares/fsup.go b/server/middlewares/fsup.go index d99e62aea7..bea023b0c2 100644 --- a/server/middlewares/fsup.go +++ b/server/middlewares/fsup.go @@ -4,6 +4,7 @@ import ( "net/url" stdpath "path" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/model" @@ -34,12 +35,12 @@ func FsUp(c *gin.Context) { c.Abort() return } - if !user.CanWriteContent() && !common.CanWriteContentBypassUserPerms(parentMeta, parentPath) { + if !user.CanWriteContent() && !authz.CanWriteContentBypassUserPerms(parentMeta, parentPath) { common.ErrorResp(c, errs.PermissionDenied, 403) c.Abort() return } - if !common.CanWrite(user, parentMeta, parentPath) { + if !authz.CanWrite(user, parentMeta, parentPath) { common.ErrorResp(c, errs.PermissionDenied, 403) c.Abort() return diff --git a/server/webdav/file.go b/server/webdav/file.go index ea60997359..4d29ff1b45 100644 --- a/server/webdav/file.go +++ b/server/webdav/file.go @@ -10,12 +10,12 @@ import ( "path" "path/filepath" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/fs" "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/internal/op" - "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/pkg/errors" ) @@ -52,7 +52,7 @@ func moveFiles(ctx context.Context, src, dst string, overwrite bool) (status int if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return http.StatusInternalServerError, err } - if !common.CanWrite(user, srcMeta, srcDir) || !common.CanWrite(user, dstMeta, dstDir) { + if !authz.CanWrite(user, srcMeta, srcDir) || !authz.CanWrite(user, dstMeta, dstDir) { return http.StatusForbidden, nil } if srcDir == dstDir { @@ -88,14 +88,14 @@ func copyFiles(ctx context.Context, src, dst string, overwrite bool) (status int if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return http.StatusInternalServerError, err } - if !common.CanRead(user, srcMeta, srcDir) { + if !authz.CanRead(user, srcMeta, srcDir) { return http.StatusForbidden, nil } dstMeta, err := op.GetNearestMeta(dstDir) if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return http.StatusInternalServerError, err } - if !common.CanWrite(user, dstMeta, dstDir) { + if !authz.CanWrite(user, dstMeta, dstDir) { return http.StatusForbidden, nil } _, err = fs.Copy(context.WithValue(ctx, conf.NoTaskKey, struct{}{}), src, dstDir) diff --git a/server/webdav/webdav.go b/server/webdav/webdav.go index 06d1431ac3..a27cb08419 100644 --- a/server/webdav/webdav.go +++ b/server/webdav/webdav.go @@ -17,6 +17,7 @@ import ( "strings" "time" + "github.com/OpenListTeam/OpenList/v4/internal/authz" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/net" "github.com/OpenListTeam/OpenList/v4/internal/op" @@ -236,7 +237,7 @@ func (h *Handler) handleGetHeadPost(w http.ResponseWriter, r *http.Request) (sta if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return http.StatusInternalServerError, err } - if !common.CanAccess(user, meta, reqPath, password) { + if !authz.CanAccess(user, meta, reqPath, password) { return http.StatusForbidden, errs.PermissionDenied } fi, err := fs.Get(ctx, reqPath, &fs.GetArgs{}) @@ -326,7 +327,7 @@ func (h *Handler) handleDelete(w http.ResponseWriter, r *http.Request) (status i if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return http.StatusInternalServerError, err } - if !common.CanWrite(user, parentMeta, parentPath) { + if !authz.CanWrite(user, parentMeta, parentPath) { return http.StatusForbidden, errs.PermissionDenied } if err := fs.Remove(ctx, reqPath); err != nil { @@ -388,10 +389,10 @@ func (h *Handler) handlePut(w http.ResponseWriter, r *http.Request) (status int, if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return http.StatusInternalServerError, err } - if !user.CanWriteContent() && !common.CanWriteContentBypassUserPerms(parentMeta, parentPath) { + if !user.CanWriteContent() && !authz.CanWriteContentBypassUserPerms(parentMeta, parentPath) { return http.StatusForbidden, errs.PermissionDenied } - if !common.CanWrite(user, parentMeta, parentPath) { + if !authz.CanWrite(user, parentMeta, parentPath) { return http.StatusForbidden, errs.PermissionDenied } fsStream := &stream.FileStream{ @@ -463,10 +464,10 @@ func (h *Handler) handleMkcol(w http.ResponseWriter, r *http.Request) (status in if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return http.StatusInternalServerError, err } - if !user.CanWriteContent() && !common.CanWriteContentBypassUserPerms(parentMeta, parentPath) { + if !user.CanWriteContent() && !authz.CanWriteContentBypassUserPerms(parentMeta, parentPath) { return http.StatusForbidden, errs.PermissionDenied } - if !common.CanWrite(user, parentMeta, parentPath) { + if !authz.CanWrite(user, parentMeta, parentPath) { return http.StatusForbidden, errs.PermissionDenied } if err := fs.MakeDir(ctx, reqPath); err != nil { @@ -619,7 +620,7 @@ func (h *Handler) handleLock(w http.ResponseWriter, r *http.Request) (retStatus if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return http.StatusInternalServerError, err } - if !common.CanWrite(user, meta, reqPath) { + if !authz.CanWrite(user, meta, reqPath) { return http.StatusForbidden, errs.PermissionDenied } ld = LockDetails{ @@ -692,7 +693,7 @@ func (h *Handler) handleUnlock(w http.ResponseWriter, r *http.Request) (status i if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return http.StatusInternalServerError, err } - if !common.CanWrite(user, meta, reqPath) { + if !authz.CanWrite(user, meta, reqPath) { return http.StatusForbidden, errs.PermissionDenied } @@ -728,7 +729,7 @@ func (h *Handler) handlePropfind(w http.ResponseWriter, r *http.Request) (status if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return http.StatusInternalServerError, err } - if !common.CanAccess(user, meta, reqPath, password) { + if !authz.CanAccess(user, meta, reqPath, password) { return http.StatusForbidden, errs.PermissionDenied } fi, err := fs.Get(ctx, reqPath, &fs.GetArgs{}) @@ -814,7 +815,7 @@ func (h *Handler) handleProppatch(w http.ResponseWriter, r *http.Request) (statu if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { return http.StatusInternalServerError, err } - if !common.CanWrite(user, meta, reqPath) { + if !authz.CanWrite(user, meta, reqPath) { return http.StatusForbidden, errs.PermissionDenied } if _, err := fs.Get(ctx, reqPath, &fs.GetArgs{}); err != nil { From 6865b330473052c3e36decb178b6655f96c7f771 Mon Sep 17 00:00:00 2001 From: nostalume Date: Thu, 24 Sep 2026 23:42:13 +0800 Subject: [PATCH 2/6] fix(fs): preserve mount operation failures - Propagate directory-listing failures through mount traversal and protocol callers. - Keep list diagnostics at their owning outcome and remove an unused argument placeholder. Co-authored-by: Codex <267193182+codex@users.noreply.github.com> --- drivers/alias/driver.go | 4 ++-- drivers/alias/util.go | 12 ++++++------ drivers/chunk/driver.go | 2 +- internal/fs/fs.go | 5 +---- internal/fs/list.go | 6 +++--- internal/fs/walk.go | 2 +- internal/fs/walk_test.go | 30 ++++++++++++++++++++++++++++++ server/handles/archive.go | 4 ++-- server/handles/down.go | 4 ++-- server/handles/fsmanage.go | 2 +- server/handles/fsread.go | 4 ++-- server/handles/fsup.go | 2 +- server/mcp/fs_get.go | 2 +- server/mcp/fs_link.go | 2 +- server/s3/redirect.go | 2 +- server/webdav/webdav.go | 2 +- 16 files changed, 56 insertions(+), 29 deletions(-) create mode 100644 internal/fs/walk_test.go diff --git a/drivers/alias/driver.go b/drivers/alias/driver.go index d69d6cf502..5d43a879b6 100644 --- a/drivers/alias/driver.go +++ b/drivers/alias/driver.go @@ -138,7 +138,7 @@ func (d *Alias) Get(ctx context.Context, path string) (model.Obj, error) { } obj = &ret if d.ProviderPassThrough && !obj.IsDir() { - if storage, err := fs.GetStorage(rawPath, &fs.GetStoragesArgs{}); err == nil { + if storage, err := fs.GetStorage(rawPath); err == nil { obj = &model.ObjectProvider{ Object: ret, Provider: model.Provider{ @@ -536,7 +536,7 @@ func (d *Alias) GetDetails(ctx context.Context) (*model.StorageDetails, error) { backends := d.pathMap[d.rootOrder[0]] var storage driver.Driver for _, backend := range backends { - s, err := fs.GetStorage(backend, &fs.GetStoragesArgs{}) + s, err := fs.GetStorage(backend) if err != nil { return nil, errs.NotImplement } diff --git a/drivers/alias/util.go b/drivers/alias/util.go index 8e5eb8a843..98f8ffe886 100644 --- a/drivers/alias/util.go +++ b/drivers/alias/util.go @@ -262,7 +262,7 @@ func getRandomObjByQuotaBalanced(ctx context.Context, reqPath BalancedObjs, stri detailsChan := make(chan detailWithIndex, len(reqPath)) workerCount := 0 for i, p := range reqPath { - s, err := fs.GetStorage(p.GetPath(), &fs.GetStoragesArgs{}) + s, err := fs.GetStorage(p.GetPath()) if err != nil { continue } @@ -347,7 +347,7 @@ func (d *Alias) getCopyObjs(ctx context.Context, srcObj, dstDir model.Obj) (Bala dstStorageMap := make(map[string][]model.Obj) allocatingDst := make(map[model.Obj]struct{}) for _, o := range dstObjs { - storage, e := fs.GetStorage(o.GetPath(), &fs.GetStoragesArgs{}) + storage, e := fs.GetStorage(o.GetPath()) if e != nil { return nil, nil, errors.WithMessagef(e, "cannot copy to virtual path [%s]", o.GetPath()) } @@ -361,7 +361,7 @@ func (d *Alias) getCopyObjs(ctx context.Context, srcObj, dstDir model.Obj) (Bala } srcObjs := make(BalancedObjs, 0, len(dstObjs)) for _, src := range tmpSrcObjs { - storage, e := fs.GetStorage(src.GetPath(), &fs.GetStoragesArgs{}) + storage, e := fs.GetStorage(src.GetPath()) if e != nil { continue } @@ -405,7 +405,7 @@ func (d *Alias) getMoveObjs(ctx context.Context, srcObj, dstDir model.Obj) (Bala dstStorageMap := make(map[string][]model.Obj) allocatingDst := make(map[model.Obj]struct{}) for _, o := range dstObjs { - storage, e := fs.GetStorage(o.GetPath(), &fs.GetStoragesArgs{}) + storage, e := fs.GetStorage(o.GetPath()) if e != nil { return nil, nil, errors.WithMessagef(e, "cannot move to virtual path [%s]", o.GetPath()) } @@ -416,7 +416,7 @@ func (d *Alias) getMoveObjs(ctx context.Context, srcObj, dstDir model.Obj) (Bala srcObjs := make(BalancedObjs, 0, len(tmpSrcObjs)) restSrcObjs := make(BalancedObjs, 0, len(tmpSrcObjs)-len(dstObjs)) for _, src := range tmpSrcObjs { - storage, e := fs.GetStorage(src.GetPath(), &fs.GetStoragesArgs{}) + storage, e := fs.GetStorage(src.GetPath()) if e != nil { continue } @@ -499,7 +499,7 @@ func getAllSort(dirs []model.Obj) model.Sort { if dir == nil { continue } - storage, err := fs.GetStorage(dir.GetPath(), &fs.GetStoragesArgs{}) + storage, err := fs.GetStorage(dir.GetPath()) if err != nil { continue } diff --git a/drivers/chunk/driver.go b/drivers/chunk/driver.go index 1ebc8aef6c..02b5cf2b44 100644 --- a/drivers/chunk/driver.go +++ b/drivers/chunk/driver.go @@ -529,7 +529,7 @@ func (d *Chunk) getPartName(part int) string { } func (d *Chunk) GetDetails(ctx context.Context) (*model.StorageDetails, error) { - remoteStorage, err := fs.GetStorage(d.RemotePath, &fs.GetStoragesArgs{}) + remoteStorage, err := fs.GetStorage(d.RemotePath) if err != nil { return nil, errs.NotImplement } diff --git a/internal/fs/fs.go b/internal/fs/fs.go index 67a1ac065e..e4cb6ef389 100644 --- a/internal/fs/fs.go +++ b/internal/fs/fs.go @@ -164,10 +164,7 @@ func ArchiveInternalExtract(ctx context.Context, path string, args model.Archive return l, obj, err } -type GetStoragesArgs struct { -} - -func GetStorage(path string, args *GetStoragesArgs) (driver.Driver, error) { +func GetStorage(path string) (driver.Driver, error) { storageDriver, _, err := op.GetStorageAndActualPath(path) if err != nil { return nil, err diff --git a/internal/fs/list.go b/internal/fs/list.go index 1d70c96a00..5a9025b034 100644 --- a/internal/fs/list.go +++ b/internal/fs/list.go @@ -31,12 +31,12 @@ func list(ctx context.Context, path string, args *ListArgs) ([]model.Obj, error) WithStorageDetails: args.WithStorageDetails, }) if err != nil { - if !args.NoLog { - log.Errorf("fs/list: %+v", err) - } if len(virtualFiles) == 0 { return nil, errors.WithMessage(err, "failed get objs") } + if !args.NoLog { + log.Errorf("fs/list: %+v", err) + } } } diff --git a/internal/fs/walk.go b/internal/fs/walk.go index a534dc4ddd..5e669ff756 100644 --- a/internal/fs/walk.go +++ b/internal/fs/walk.go @@ -31,7 +31,7 @@ func WalkFS(ctx context.Context, depth int, name string, info model.Obj, walkFn // Read directory names. objs, err := List(context.WithValue(ctx, conf.MetaKey, meta), name, &ListArgs{}) if err != nil { - return walkFnErr + return err } for _, fileInfo := range objs { filename := path.Join(name, fileInfo.GetName()) diff --git a/internal/fs/walk_test.go b/internal/fs/walk_test.go new file mode 100644 index 0000000000..ccdc400d4a --- /dev/null +++ b/internal/fs/walk_test.go @@ -0,0 +1,30 @@ +package fs + +import ( + "context" + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/db" + "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/glebarez/sqlite" + "gorm.io/gorm" +) + +func init() { + conf.Conf = conf.DefaultConfig("data") + database, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) + if err != nil { + panic(err) + } + db.Init(database) +} + +func TestWalkFSPropagatesListFailure(t *testing.T) { + err := WalkFS(context.Background(), 1, "/missing", &model.Object{IsFolder: true}, func(string, model.Obj) error { + return nil + }) + if err == nil { + t.Fatal("WalkFS returned nil after listing failed") + } +} diff --git a/server/handles/archive.go b/server/handles/archive.go index a60ca71468..0438acbfd6 100644 --- a/server/handles/archive.go +++ b/server/handles/archive.go @@ -309,7 +309,7 @@ func ArchiveDown(c *gin.Context) { innerPath := utils.FixAndCleanPath(c.Query("inner")) password := c.Query("pass") filename := stdpath.Base(innerPath) - storage, err := fs.GetStorage(archiveRawPath, &fs.GetStoragesArgs{}) + storage, err := fs.GetStorage(archiveRawPath) if err != nil { common.ErrorPage(c, err, 500) return @@ -343,7 +343,7 @@ func ArchiveProxy(c *gin.Context) { innerPath := utils.FixAndCleanPath(c.Query("inner")) password := c.Query("pass") filename := stdpath.Base(innerPath) - storage, err := fs.GetStorage(archiveRawPath, &fs.GetStoragesArgs{}) + storage, err := fs.GetStorage(archiveRawPath) if err != nil { common.ErrorPage(c, err, 500) return diff --git a/server/handles/down.go b/server/handles/down.go index 50025a0b69..ec97247b22 100644 --- a/server/handles/down.go +++ b/server/handles/down.go @@ -20,7 +20,7 @@ import ( func Down(c *gin.Context) { rawPath := c.Request.Context().Value(conf.PathKey).(string) filename := stdpath.Base(rawPath) - storage, err := fs.GetStorage(rawPath, &fs.GetStoragesArgs{}) + storage, err := fs.GetStorage(rawPath) if err != nil { common.ErrorPage(c, err, 500) return @@ -46,7 +46,7 @@ func Down(c *gin.Context) { func Proxy(c *gin.Context) { rawPath := c.Request.Context().Value(conf.PathKey).(string) filename := stdpath.Base(rawPath) - storage, err := fs.GetStorage(rawPath, &fs.GetStoragesArgs{}) + storage, err := fs.GetStorage(rawPath) if err != nil { common.ErrorPage(c, err, 500) return diff --git a/server/handles/fsmanage.go b/server/handles/fsmanage.go index c19444cf40..ba512423d5 100644 --- a/server/handles/fsmanage.go +++ b/server/handles/fsmanage.go @@ -508,7 +508,7 @@ func Link(c *gin.Context) { //rawPath := stdpath.Join(user.BasePath, req.Path) // why need not join base_path? because it's always the full path rawPath := req.Path - storage, err := fs.GetStorage(rawPath, &fs.GetStoragesArgs{}) + storage, err := fs.GetStorage(rawPath) if err != nil { common.ErrorResp(c, err, 500) return diff --git a/server/handles/fsread.go b/server/handles/fsread.go index ef01b86f3c..4e5b2a0ec0 100644 --- a/server/handles/fsread.go +++ b/server/handles/fsread.go @@ -111,7 +111,7 @@ func FsList(c *gin.Context, req *ListReq, user *model.User) { provider := "unknown" var directUploadTools []string if canWriteContentAtPath { - if storage, err := fs.GetStorage(reqPath, &fs.GetStoragesArgs{}); err == nil { + if storage, err := fs.GetStorage(reqPath); err == nil { directUploadTools = op.GetDirectUploadTools(storage) } } @@ -306,7 +306,7 @@ func FsGet(c *gin.Context, req *FsGetReq, user *model.User) { } var rawURL string - storage, err := fs.GetStorage(reqPath, &fs.GetStoragesArgs{}) + storage, err := fs.GetStorage(reqPath) provider, ok := model.GetProvider(obj) if !ok && err == nil { provider = storage.Config().Name diff --git a/server/handles/fsup.go b/server/handles/fsup.go index 0f46398cdf..7d5af7b65b 100644 --- a/server/handles/fsup.go +++ b/server/handles/fsup.go @@ -154,7 +154,7 @@ func FsForm(c *gin.Context) { return } } - storage, err := fs.GetStorage(path, &fs.GetStoragesArgs{}) + storage, err := fs.GetStorage(path) if err != nil { common.ErrorResp(c, err, 400) return diff --git a/server/mcp/fs_get.go b/server/mcp/fs_get.go index 0bce0835a8..7bdc402310 100644 --- a/server/mcp/fs_get.go +++ b/server/mcp/fs_get.go @@ -117,7 +117,7 @@ func parseFSGetArgs(raw json.RawMessage) (*fsGetArgs, *rpcError) { } func buildFSGetRawURL(ctx context.Context, c *gin.Context, reqPath string, obj model.Obj, meta *model.Meta) (string, string, error) { - storage, storageErr := fs.GetStorage(reqPath, &fs.GetStoragesArgs{}) + storage, storageErr := fs.GetStorage(reqPath) provider, ok := model.GetProvider(obj) if !ok && storageErr == nil { provider = storage.Config().Name diff --git a/server/mcp/fs_link.go b/server/mcp/fs_link.go index f8d18f9a29..aef2fb1fdd 100644 --- a/server/mcp/fs_link.go +++ b/server/mcp/fs_link.go @@ -86,7 +86,7 @@ func (s *Server) callFSLink(c *gin.Context, raw json.RawMessage) (any, *rpcError return nil, &rpcError{Code: -32003, Message: "path is a directory"} } - storage, err := fs.GetStorage(reqPath, &fs.GetStoragesArgs{}) + storage, err := fs.GetStorage(reqPath) if err != nil { return nil, &rpcError{Code: -32603, Message: err.Error()} } diff --git a/server/s3/redirect.go b/server/s3/redirect.go index 6d8d430bc3..7dfeede475 100644 --- a/server/s3/redirect.go +++ b/server/s3/redirect.go @@ -55,7 +55,7 @@ func directObjectURL(r *http.Request, authPairs map[string]string) (string, bool reqPath := path.Join(bucket.Path, objectName) meta, _ := op.GetNearestMeta(reqPath) ctx := context.WithValue(r.Context(), conf.MetaKey, meta) - storage, err := fs.GetStorage(reqPath, &fs.GetStoragesArgs{}) + storage, err := fs.GetStorage(reqPath) if err != nil || common.ShouldProxy(storage, path.Base(reqPath)) { return "", false } diff --git a/server/webdav/webdav.go b/server/webdav/webdav.go index a27cb08419..f4dc6d6b49 100644 --- a/server/webdav/webdav.go +++ b/server/webdav/webdav.go @@ -253,7 +253,7 @@ func (h *Handler) handleGetHeadPost(w http.ResponseWriter, r *http.Request) (sta return http.StatusMethodNotAllowed, nil } // Let ServeContent determine the Content-Type header. - storage, _ := fs.GetStorage(reqPath, &fs.GetStoragesArgs{}) + storage, _ := fs.GetStorage(reqPath) if storage.GetStorage().Webdav302() { link, _, err := fs.Link(ctx, reqPath, model.LinkArgs{IP: utils.ClientIP(r), Header: r.Header, Redirect: true}) if err != nil { From 8280d016acfddf7c24b78f5011c570c8048a377a Mon Sep 17 00:00:00 2001 From: nostalume Date: Thu, 24 Sep 2026 23:42:26 +0800 Subject: [PATCH 3/6] refactor(mq): add bounded latest-value processor - Admit keyed updates through a bounded processor with explicit shutdown and replacement semantics. - Retain behavior tests for overwrite, drain, cancellation, and refusal paths. Co-authored-by: Codex <267193182+codex@users.noreply.github.com> --- pkg/mq/latest.go | 182 ++++++++++++++++++++++++++++++++++++++++++ pkg/mq/latest_test.go | 126 +++++++++++++++++++++++++++++ 2 files changed, 308 insertions(+) create mode 100644 pkg/mq/latest.go create mode 100644 pkg/mq/latest_test.go diff --git a/pkg/mq/latest.go b/pkg/mq/latest.go new file mode 100644 index 0000000000..86f54a415f --- /dev/null +++ b/pkg/mq/latest.go @@ -0,0 +1,182 @@ +package mq + +import ( + "context" + "sync" +) + +type OfferResult uint8 + +const ( + OfferAdded OfferResult = iota + OfferReplaced + OfferRejectedCapacity + OfferRejectedClosed +) + +type MailboxStats struct { + Replaced, Rejected uint64 + Pending, InFlight, Weight int +} + +type latestEntry[V any] struct { + value V + weight int +} + +// LatestProcessor runs fixed workers over a bounded latest-value-per-key queue. +type LatestProcessor[K comparable, V any] struct { + mu sync.Mutex + entries map[K]latestEntry[V] + inFlight map[K]int + order []K + ready, exited chan struct{} + maxKeys, maxWeight int + weight int + stopped bool + stats MailboxStats + cancel context.CancelFunc + wg sync.WaitGroup + process func(context.Context, K, V) +} + +func NewLatestProcessor[K comparable, V any](maxKeys, maxWeight, workers int, process func(context.Context, K, V)) *LatestProcessor[K, V] { + if maxKeys < 1 || maxWeight < 0 || workers < 1 || process == nil { + panic("mq: invalid latest processor") + } + ctx, cancel := context.WithCancel(context.Background()) + p := &LatestProcessor[K, V]{ + entries: make(map[K]latestEntry[V]), inFlight: make(map[K]int), + ready: make(chan struct{}, 1), exited: make(chan struct{}), + maxKeys: maxKeys, maxWeight: maxWeight, cancel: cancel, process: process, + } + for range workers { + p.wg.Add(1) + go p.run(ctx) + } + go func() { p.wg.Wait(); close(p.exited) }() + return p +} + +func (p *LatestProcessor[K, V]) Offer(key K, value V, weight int) OfferResult { + if weight < 0 { + panic("mq: negative latest processor weight") + } + p.mu.Lock() + defer p.mu.Unlock() + if p.stopped { + p.stats.Rejected++ + return OfferRejectedClosed + } + if entry, ok := p.entries[key]; ok { + if next := p.weight - entry.weight + weight; next <= p.maxWeight { + p.entries[key], p.weight = latestEntry[V]{value, weight}, next + p.stats.Replaced++ + return OfferReplaced + } + p.stats.Rejected++ + return OfferRejectedCapacity + } + _, busy := p.inFlight[key] + if (!busy && len(p.entries)+len(p.inFlight) == p.maxKeys) || p.weight+weight > p.maxWeight { + p.stats.Rejected++ + return OfferRejectedCapacity + } + p.entries[key] = latestEntry[V]{value, weight} + p.order, p.weight = append(p.order, key), p.weight+weight + p.signal() + return OfferAdded +} + +func (p *LatestProcessor[K, V]) Stats() MailboxStats { + p.mu.Lock() + defer p.mu.Unlock() + stats := p.stats + stats.Pending, stats.InFlight, stats.Weight = len(p.entries), len(p.inFlight), p.weight + return stats +} + +func (p *LatestProcessor[K, V]) Stop(ctx context.Context) error { + p.mu.Lock() + if !p.stopped { + p.stopped = true + p.signal() + } + p.mu.Unlock() + select { + case <-p.exited: + p.cancel() + return nil + case <-ctx.Done(): + p.cancel() + return ctx.Err() + } +} + +func (p *LatestProcessor[K, V]) run(ctx context.Context) { + defer p.wg.Done() + for { + key, value, ok := p.take(ctx) + if !ok { + return + } + func() { defer p.done(key); p.process(ctx, key, value) }() + } +} + +func (p *LatestProcessor[K, V]) take(ctx context.Context) (K, V, bool) { + for { + p.mu.Lock() + for i, key := range p.order { + if _, busy := p.inFlight[key]; busy { + continue + } + entry := p.entries[key] + p.order = append(p.order[:i], p.order[i+1:]...) + delete(p.entries, key) + p.inFlight[key] = entry.weight + if len(p.order) > 0 { + p.signal() + } + p.mu.Unlock() + return key, entry.value, true + } + if p.stopped && len(p.entries) == 0 { + p.signal() + p.mu.Unlock() + return emptyLatest[K, V]() + } + p.mu.Unlock() + select { + case <-ctx.Done(): + return emptyLatest[K, V]() + case <-p.ready: + } + } +} + +func emptyLatest[K comparable, V any]() (K, V, bool) { + var key K + var value V + return key, value, false +} + +func (p *LatestProcessor[K, V]) done(key K) { + p.mu.Lock() + weight, ok := p.inFlight[key] + if ok { + delete(p.inFlight, key) + p.weight -= weight + } + p.mu.Unlock() + if ok { + p.signal() + } +} + +func (p *LatestProcessor[K, V]) signal() { + select { + case p.ready <- struct{}{}: + default: + } +} diff --git a/pkg/mq/latest_test.go b/pkg/mq/latest_test.go new file mode 100644 index 0000000000..7fa305d9ec --- /dev/null +++ b/pkg/mq/latest_test.go @@ -0,0 +1,126 @@ +package mq + +import ( + "context" + "testing" + "time" +) + +func newTestLatestProcessor[K comparable, V any](maxKeys, maxWeight int) *LatestProcessor[K, V] { + return &LatestProcessor[K, V]{ + entries: make(map[K]latestEntry[V]), inFlight: make(map[K]int), + ready: make(chan struct{}, 1), + maxKeys: maxKeys, maxWeight: maxWeight, + } +} + +func TestLatestProcessorCoalescesByKey(t *testing.T) { + mailbox := newTestLatestProcessor[string, string](2, 4) + if got := mailbox.Offer("/a", "old", 2); got != OfferAdded { + t.Fatalf("first offer = %v, want added", got) + } + if got := mailbox.Offer("/a", "new", 3); got != OfferReplaced { + t.Fatalf("replacement = %v, want replaced", got) + } + if got := mailbox.Offer("/b", "other", 1); got != OfferAdded { + t.Fatalf("second key = %v, want added", got) + } + + key, value, ok := mailbox.take(context.Background()) + if !ok || key != "/a" || value != "new" { + t.Fatalf("first take = (%q, %q, %t), want latest /a", key, value, ok) + } + mailbox.done(key) + stats := mailbox.Stats() + if stats.Replaced != 1 || stats.Pending != 1 || stats.Weight != 1 { + t.Fatalf("stats after coalescing = %+v", stats) + } +} + +func TestLatestProcessorRejectsCapacityWithoutDiscardingPendingValue(t *testing.T) { + mailbox := newTestLatestProcessor[string, string](1, 2) + mailbox.Offer("/a", "kept", 2) + if got := mailbox.Offer("/b", "other", 1); got != OfferRejectedCapacity { + t.Fatalf("new key over capacity = %v, want rejected", got) + } + if got := mailbox.Offer("/a", "too-large", 3); got != OfferRejectedCapacity { + t.Fatalf("replacement over capacity = %v, want rejected", got) + } + key, value, ok := mailbox.take(context.Background()) + if !ok || key != "/a" || value != "kept" { + t.Fatalf("take after rejection = (%q, %q, %t), want retained value", key, value, ok) + } + mailbox.done(key) + if stats := mailbox.Stats(); stats.Rejected != 2 || stats.Pending != 0 || stats.Weight != 0 { + t.Fatalf("stats after rejection = %+v", stats) + } +} + +func TestLatestProcessorCloseRejectsNewWorkAndDrainsPending(t *testing.T) { + mailbox := newTestLatestProcessor[string, string](1, 1) + mailbox.Offer("/a", "pending", 1) + mailbox.stopped = true + if got := mailbox.Offer("/b", "rejected", 1); got != OfferRejectedClosed { + t.Fatalf("offer after close = %v, want closed rejection", got) + } + key, value, ok := mailbox.take(context.Background()) + if !ok || value != "pending" { + t.Fatalf("pending value was not drained: (%q, %t)", value, ok) + } + mailbox.done(key) + if _, _, ok := mailbox.take(context.Background()); ok { + t.Fatal("closed drained mailbox returned another value") + } +} + +func TestLatestProcessorSerializesOneKeyAcrossWorkers(t *testing.T) { + mailbox := newTestLatestProcessor[string, string](2, 4) + mailbox.Offer("/a", "first", 1) + key, _, ok := mailbox.take(context.Background()) + if !ok { + t.Fatal("first take failed") + } + mailbox.Offer("/a", "second", 1) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if _, _, ok := mailbox.take(ctx); ok { + t.Fatal("same key was taken while its first value was in flight") + } + mailbox.done(key) + if _, value, ok := mailbox.take(context.Background()); !ok || value != "second" { + t.Fatalf("latest same-key value after completion = (%q, %t)", value, ok) + } +} + +func TestLatestProcessorStopDrainsAcceptedWork(t *testing.T) { + processed := make(chan string, 1) + processor := NewLatestProcessor(1, 1, 1, func(_ context.Context, _ string, value string) { + processed <- value + }) + processor.Offer("/a", "accepted", 1) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := processor.Stop(ctx); err != nil { + t.Fatal(err) + } + if got := <-processed; got != "accepted" { + t.Fatalf("processed %q, want accepted", got) + } +} + +func BenchmarkLatestProcessorReplace(b *testing.B) { + mailbox := newTestLatestProcessor[string, int](1, 1) + for i := 0; b.Loop(); i++ { + mailbox.Offer("parent", i, 1) + } +} + +func BenchmarkLatestProcessorAdmissionRoundTrip(b *testing.B) { + mailbox := newTestLatestProcessor[int, int](1, 1) + for i := 0; b.Loop(); i++ { + mailbox.Offer(i, i, 1) + key, _, _ := mailbox.take(context.Background()) + mailbox.done(key) + } +} From 7dadf5e256019f32c528cffc535bff069a2c77ba Mon Sep 17 00:00:00 2001 From: nostalume Date: Thu, 24 Sep 2026 23:42:42 +0800 Subject: [PATCH 4/6] refactor(namespace): converge mutation projections - Route accepted namespace changes through one bounded latest-value projection owner. - Remove generic hooks, duplicate search queues and locks, skip flags, and old task payload paths. - Preserve mount, search, STRM, and cache behavior with identity and task-group regressions. Co-authored-by: Codex <267193182+codex@users.noreply.github.com> --- drivers/chunk/driver.go | 6 +- drivers/strm/hook.go | 54 ++- internal/bootstrap/projection.go | 53 +++ internal/bootstrap/run.go | 2 + internal/conf/const.go | 2 - internal/fs/archive.go | 47 +-- internal/fs/copy_move.go | 31 +- internal/fs/fs.go | 34 +- internal/fs/other.go | 7 +- internal/fs/put.go | 13 +- internal/model/args.go | 1 - internal/offline_download/tool/download.go | 2 +- internal/offline_download/tool/transfer.go | 18 +- internal/op/archive.go | 39 +- internal/op/cache.go | 81 +++- internal/op/fs.go | 428 ++++++++++----------- internal/op/hook.go | 28 +- internal/op/list_identity_test.go | 257 +++++++++++++ internal/op/reconcile.go | 127 ++++++ internal/op/recursive_list.go | 15 +- internal/search/build.go | 115 +++--- internal/search/build_test.go | 58 --- internal/search/meilisearch/init.go | 4 - internal/search/meilisearch/search.go | 57 +-- internal/search/meilisearch/task_queue.go | 265 ------------- internal/search/meilisearch/update.go | 63 +++ internal/search/searcher/snapshot.go | 23 ++ internal/search/update_lock.go | 38 -- internal/task_group/group.go | 92 ++--- internal/task_group/group_test.go | 26 ++ internal/task_group/transfer.go | 62 +-- pkg/mq/mq.go | 63 --- server/ftp/fsmanage.go | 2 +- server/ftp/fsup.go | 4 +- server/handles/fsbatch.go | 4 +- server/handles/fsmanage.go | 10 +- server/webdav/file.go | 4 +- 37 files changed, 1089 insertions(+), 1046 deletions(-) create mode 100644 internal/bootstrap/projection.go create mode 100644 internal/op/list_identity_test.go create mode 100644 internal/op/reconcile.go delete mode 100644 internal/search/build_test.go delete mode 100644 internal/search/meilisearch/task_queue.go create mode 100644 internal/search/meilisearch/update.go create mode 100644 internal/search/searcher/snapshot.go delete mode 100644 internal/search/update_lock.go create mode 100644 internal/task_group/group_test.go delete mode 100644 pkg/mq/mq.go diff --git a/drivers/chunk/driver.go b/drivers/chunk/driver.go index 02b5cf2b44..ba2a2c9069 100644 --- a/drivers/chunk/driver.go +++ b/drivers/chunk/driver.go @@ -10,7 +10,6 @@ import ( "strconv" "strings" - "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/driver" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/fs" @@ -472,10 +471,9 @@ func (d *Chunk) Put(ctx context.Context, dstDir model.Obj, file model.FileStream UpdateProgress: up, } dst := stdpath.Join(remoteActualPath, dstDir.GetPath(), d.ChunkPrefix+file.GetName()) - skipHookCtx := context.WithValue(ctx, conf.SkipHookKey, struct{}{}) if d.StoreHash { for ht, value := range file.GetHash().All() { - _ = op.Put(skipHookCtx, remoteStorage, dst, &stream.FileStream{ + _ = op.Put(ctx, remoteStorage, dst, &stream.FileStream{ Obj: &model.Object{ Name: fmt.Sprintf("hash_%s_%s%s", ht.Name, value, d.CustomExt), Size: 1, @@ -494,7 +492,7 @@ func (d *Chunk) Put(ctx context.Context, dstDir model.Obj, file model.FileStream } partIndex := 0 for partIndex < fullPartCount { - err = op.Put(skipHookCtx, remoteStorage, dst, &stream.FileStream{ + err = op.Put(ctx, remoteStorage, dst, &stream.FileStream{ Obj: &model.Object{ Name: d.getPartName(partIndex), Size: d.PartSize, diff --git a/drivers/strm/hook.go b/drivers/strm/hook.go index 24b31ee3e9..1971f665c1 100644 --- a/drivers/strm/hook.go +++ b/drivers/strm/hook.go @@ -10,9 +10,9 @@ import ( stdpath "path" "path/filepath" "strings" + "sync" "github.com/OpenListTeam/OpenList/v4/internal/model" - "github.com/OpenListTeam/OpenList/v4/internal/op" "github.com/OpenListTeam/OpenList/v4/internal/stream" "github.com/OpenListTeam/OpenList/v4/pkg/http_range" "github.com/OpenListTeam/OpenList/v4/pkg/utils" @@ -20,10 +20,35 @@ import ( "github.com/tchap/go-patricia/v2/patricia" ) -var strmTrie = patricia.NewTrie() +var ( + strmTrie = patricia.NewTrie() + strmTrieMu sync.RWMutex +) func UpdateLocalStrm(ctx context.Context, path string, objs []model.Obj) { path = utils.FixAndCleanPath(path) + type target struct { + driver *Strm + basePath string + } + var targets []target + strmTrieMu.RLock() + _ = strmTrie.VisitPrefixes(patricia.Prefix(path), func(needPathPrefix patricia.Prefix, item patricia.Item) error { + needPath := string(needPathPrefix) + restPath := strings.TrimPrefix(path, needPath) + if len(restPath) > 0 && restPath[0] != '/' { + return nil + } + for _, strmDriver := range item.([]*Strm) { + targets = append(targets, target{ + driver: strmDriver, + basePath: stdpath.Join(stdpath.Base(needPath), restPath), + }) + } + return nil + }) + strmTrieMu.RUnlock() + updateLocal := func(driver *Strm, basePath string, objs []model.Obj) { relParent := strings.TrimPrefix(basePath, utils.GetActualMountPath(driver.MountPath)) localParentPath := stdpath.Join(driver.SaveStrmLocalPath, relParent) @@ -38,22 +63,15 @@ func UpdateLocalStrm(ctx context.Context, path string, objs []model.Obj) { deleteExtraFiles(driver, localParentPath, objs) } - _ = strmTrie.VisitPrefixes(patricia.Prefix(path), func(needPathPrefix patricia.Prefix, item patricia.Item) error { - strmDrivers := item.([]*Strm) - needPath := string(needPathPrefix) - restPath := strings.TrimPrefix(path, needPath) - if len(restPath) > 0 && restPath[0] != '/' { - return nil - } - for _, strmDriver := range strmDrivers { - strmObjs := strmDriver.convert2strmObjs(ctx, path, objs) - updateLocal(strmDriver, stdpath.Join(stdpath.Base(needPath), restPath), strmObjs) - } - return nil - }) + for _, target := range targets { + strmObjs := target.driver.convert2strmObjs(ctx, path, objs) + updateLocal(target.driver, target.basePath, strmObjs) + } } func InsertStrm(dstPath string, d *Strm) error { + strmTrieMu.Lock() + defer strmTrieMu.Unlock() prefix := patricia.Prefix(strings.TrimRight(dstPath, "/")) existing := strmTrie.Get(prefix) @@ -73,6 +91,8 @@ func InsertStrm(dstPath string, d *Strm) error { } func RemoveStrm(dstPath string, d *Strm) { + strmTrieMu.Lock() + defer strmTrieMu.Unlock() prefix := patricia.Prefix(strings.TrimRight(dstPath, "/")) existing := strmTrie.Get(prefix) if existing == nil { @@ -249,7 +269,3 @@ func getLocalDirsAndFiles(localPath string) ([]string, []string, error) { } return files, dirs, nil } - -func init() { - op.RegisterObjsUpdateHook(UpdateLocalStrm) -} diff --git a/internal/bootstrap/projection.go b/internal/bootstrap/projection.go new file mode 100644 index 0000000000..728f78c35d --- /dev/null +++ b/internal/bootstrap/projection.go @@ -0,0 +1,53 @@ +package bootstrap + +import ( + "context" + "time" + + "github.com/OpenListTeam/OpenList/v4/drivers/strm" + "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/OpenListTeam/OpenList/v4/internal/op" + "github.com/OpenListTeam/OpenList/v4/internal/search" + "github.com/OpenListTeam/OpenList/v4/pkg/mq" + log "github.com/sirupsen/logrus" +) + +var ( + searchProjection *mq.LatestProcessor[string, []model.Obj] + strmProjection *mq.LatestProcessor[string, []model.Obj] +) + +func InitSnapshotProjection() { + searchProjection = mq.NewLatestProcessor(256, 65536, 4, search.UpdateSnapshot) + strmProjection = mq.NewLatestProcessor(128, 32768, 1, strm.UpdateLocalStrm) + op.SetSnapshotProjector(func(_ context.Context, parent string, objs []model.Obj) { + offerSnapshot("search", searchProjection, parent, objs) + offerSnapshot("strm", strmProjection, parent, objs) + }) + op.StartSnapshotReconciliation() +} + +func offerSnapshot(name string, processor *mq.LatestProcessor[string, []model.Obj], parent string, objs []model.Obj) { + result := processor.Offer(parent, objs, len(objs)) + if result == mq.OfferRejectedCapacity || result == mq.OfferRejectedClosed { + count := processor.Stats().Rejected + if count == 1 || count%100 == 0 { + log.Warnf("%s snapshot rejected: parent=%s result=%d rejected=%d", name, parent, result, count) + } + } +} + +func ReleaseSnapshotProjection() { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := op.StopSnapshotReconciliation(ctx); err != nil { + log.Warnf("stop snapshot reconciliation: %v", err) + } + op.SetSnapshotProjector(nil) + if err := searchProjection.Stop(ctx); err != nil { + log.Warnf("stop search snapshot projection: %v", err) + } + if err := strmProjection.Stop(ctx); err != nil { + log.Warnf("stop strm snapshot projection: %v", err) + } +} diff --git a/internal/bootstrap/run.go b/internal/bootstrap/run.go index 6740dba657..70325786a7 100644 --- a/internal/bootstrap/run.go +++ b/internal/bootstrap/run.go @@ -36,10 +36,12 @@ func Init() { data.InitData() InitStreamLimit() InitIndex() + InitSnapshotProjection() InitUpgradePatch() } func Release() { + ReleaseSnapshotProjection() db.Close() } diff --git a/internal/conf/const.go b/internal/conf/const.go index cc8a51d416..8c76e7b7e4 100644 --- a/internal/conf/const.go +++ b/internal/conf/const.go @@ -184,7 +184,6 @@ type ContextKey int8 const ( _ ContextKey = iota - NoTaskKey ApiUrlKey UserKey MetaKey @@ -195,5 +194,4 @@ const ( UserAgentKey PathKey SharingIDKey - SkipHookKey ) diff --git a/internal/fs/archive.go b/internal/fs/archive.go index 0219835628..c5ee0e26b6 100644 --- a/internal/fs/archive.go +++ b/internal/fs/archive.go @@ -57,7 +57,7 @@ func (t *ArchiveDownloadTask) Run() error { return err } uploadTask.groupID = stdpath.Join(uploadTask.DstStorageMp, uploadTask.DstActualPath) - task_group.TransferCoordinator.AddTask(uploadTask.groupID, nil) + task_group.TransferCoordinator.AddTask(uploadTask.groupID) ArchiveContentUploadTaskManager.Add(uploadTask) return nil } @@ -155,7 +155,7 @@ func (t *ArchiveContentUploadTask) Run() error { t.SetStartTime(time.Now()) defer func() { t.SetEndTime(time.Now()) }() return t.RunWithNextTaskCallback(func(nextTsk *ArchiveContentUploadTask) error { - task_group.TransferCoordinator.AddTask(t.groupID, nil) + task_group.TransferCoordinator.AddTask(t.groupID) ArchiveContentUploadTaskManager.Add(nextTsk) return nil }) @@ -175,7 +175,7 @@ func (t *ArchiveContentUploadTask) SetRetry(retry int, maxRetry int) { (len(t.groupID) == 0 || // 重启恢复 (t.GetErr() == nil && t.GetState() != tache.StatePending)) { // 手动重试 t.groupID = stdpath.Join(t.DstStorageMp, t.DstActualPath) - task_group.TransferCoordinator.AddTask(t.groupID, nil) + task_group.TransferCoordinator.AddTask(t.groupID) } } @@ -199,7 +199,6 @@ func (t *ArchiveContentUploadTask) RunWithNextTaskCallback(f func(nextTask *Arch return err } if !t.InPlace { - task_group.TransferCoordinator.AppendPayload(t.groupID, task_group.DstPathToHook(nextDstActualPath)) } var es error for _, entry := range entries { @@ -258,7 +257,7 @@ func (t *ArchiveContentUploadTask) RunWithNextTaskCallback(f func(nextTask *Arch } fs.Closers.Add(file) t.status = "uploading" - err = op.Put(context.WithValue(t.Ctx(), conf.SkipHookKey, struct{}{}), t.dstStorage, t.DstActualPath, fs, t.SetProgress) + err = op.Put(t.Ctx(), t.dstStorage, t.DstActualPath, fs, t.SetProgress) if err != nil { return err } @@ -362,7 +361,7 @@ func archiveList(ctx context.Context, path string, args model.ArchiveListArgs) ( return op.ListArchive(ctx, storage, actualPath, args) } -func archiveDecompress(ctx context.Context, srcObjPath, dstDirPath string, args model.ArchiveDecompressArgs, lazyCache ...bool) (task.TaskExtensionInfo, error) { +func archiveDecompress(ctx context.Context, srcObjPath, dstDirPath string, args model.ArchiveDecompressArgs) (task.TaskExtensionInfo, error) { srcStorage, srcObjActualPath, err := op.GetStorageAndActualPath(srcObjPath) if err != nil { return nil, errors.WithMessage(err, "failed get src storage") @@ -372,7 +371,7 @@ func archiveDecompress(ctx context.Context, srcObjPath, dstDirPath string, args return nil, errors.WithMessage(err, "failed get dst storage") } if srcStorage.GetStorage() == dstStorage.GetStorage() { - err = op.ArchiveDecompress(ctx, srcStorage, srcObjActualPath, dstDirActualPath, args, lazyCache...) + err = op.ArchiveDecompress(ctx, srcStorage, srcObjActualPath, dstDirActualPath, args) if !errors.Is(err, errs.NotImplement) { return nil, err } @@ -388,36 +387,10 @@ func archiveDecompress(ctx context.Context, srcObjPath, dstDirPath string, args }, ArchiveDecompressArgs: args, } - if ctx.Value(conf.NoTaskKey) != nil { - tsk.Base.SetCtx(ctx) - uploadTask, err := tsk.RunWithoutPushUploadTask() - if err != nil { - return nil, errors.WithMessagef(err, "failed download [%s]", srcObjPath) - } - defer uploadTask.deleteSrcFile() - var callback func(t *ArchiveContentUploadTask) error - var hasSuccess bool - callback = func(t *ArchiveContentUploadTask) error { - t.Base.SetCtx(ctx) - e := t.RunWithNextTaskCallback(callback) - if e == nil { - hasSuccess = true - } - t.deleteSrcFile() - return e - } - uploadTask.Base.SetCtx(ctx) - uploadTask.groupID = stdpath.Join(uploadTask.DstStorageMp, uploadTask.DstActualPath) - task_group.TransferCoordinator.AddTask(uploadTask.groupID, nil) - err = uploadTask.RunWithNextTaskCallback(callback) - task_group.TransferCoordinator.Done(context.WithoutCancel(ctx), uploadTask.groupID, hasSuccess) - return nil, err - } else { - tsk.Creator, _ = ctx.Value(conf.UserKey).(*model.User) - tsk.ApiUrl = conf.GetApiUrl(ctx) - ArchiveDownloadTaskManager.Add(tsk) - return tsk, nil - } + tsk.Creator, _ = ctx.Value(conf.UserKey).(*model.User) + tsk.ApiUrl = conf.GetApiUrl(ctx) + ArchiveDownloadTaskManager.Add(tsk) + return tsk, nil } func archiveDriverExtract(ctx context.Context, path string, args model.ArchiveInnerArgs) (*model.Link, model.Obj, error) { diff --git a/internal/fs/copy_move.go b/internal/fs/copy_move.go index 1d171a9b9c..f58a25df06 100644 --- a/internal/fs/copy_move.go +++ b/internal/fs/copy_move.go @@ -13,7 +13,6 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/stream" "github.com/OpenListTeam/OpenList/v4/internal/task" "github.com/OpenListTeam/OpenList/v4/internal/task_group" - "github.com/OpenListTeam/OpenList/v4/pkg/utils" "github.com/OpenListTeam/tache" "github.com/pkg/errors" ) @@ -39,6 +38,13 @@ const ( merge ) +type transferMode uint8 + +const ( + scheduledTransfer transferMode = iota + synchronousTransfer +) + type FileTransferTask struct { TaskData TaskType taskType @@ -67,7 +73,7 @@ func (t *FileTransferTask) Run() error { t.SetStartTime(time.Now()) defer func() { t.SetEndTime(time.Now()) }() return t.RunWithNextTaskCallback(func(nextTask *FileTransferTask) error { - task_group.TransferCoordinator.AddTask(t.groupID, nil) + task_group.TransferCoordinator.AddTask(t.groupID) if t.TaskType == copy || t.TaskType == merge { CopyTaskManager.Add(nextTask) } else { @@ -91,15 +97,14 @@ func (t *FileTransferTask) SetRetry(retry int, maxRetry int) { (len(t.groupID) == 0 || // 重启恢复 (t.GetErr() == nil && t.GetState() != tache.StatePending)) { // 手动重试 t.groupID = stdpath.Join(t.DstStorageMp, t.DstActualPath) - var payload any + task_group.TransferCoordinator.AddTask(t.groupID) if t.TaskType == move { - payload = task_group.SrcPathToRemove(stdpath.Join(t.SrcStorageMp, t.SrcActualPath)) + task_group.TransferCoordinator.RemoveSource(t.groupID, stdpath.Join(t.SrcStorageMp, t.SrcActualPath)) } - task_group.TransferCoordinator.AddTask(t.groupID, payload) } } -func transfer(ctx context.Context, taskType taskType, srcObjPath, dstDirPath string, skipHook ...bool) (task.TaskExtensionInfo, error) { +func transfer(ctx context.Context, taskType taskType, mode transferMode, srcObjPath, dstDirPath string) (task.TaskExtensionInfo, error) { srcStorage, srcObjActualPath, err := op.GetStorageAndActualPath(srcObjPath) if err != nil { return nil, errors.WithMessage(err, "failed get src storage") @@ -110,9 +115,6 @@ func transfer(ctx context.Context, taskType taskType, srcObjPath, dstDirPath str } if srcStorage.GetStorage() == dstStorage.GetStorage() { - if utils.IsBool(skipHook...) { - ctx = context.WithValue(ctx, conf.SkipHookKey, struct{}{}) - } if taskType == copy || taskType == merge { err = op.Copy(ctx, srcStorage, srcObjActualPath, dstDirActualPath) if !errors.Is(err, errs.NotImplement) && !errors.Is(err, errs.NotSupport) { @@ -140,8 +142,8 @@ func transfer(ctx context.Context, taskType taskType, srcObjPath, dstDirPath str } t.groupID = stdpath.Join(t.DstStorageMp, t.DstActualPath) - task_group.TransferCoordinator.AddTask(t.groupID, nil) - if ctx.Value(conf.NoTaskKey) != nil { + task_group.TransferCoordinator.AddTask(t.groupID) + if mode == synchronousTransfer { var callback func(nextTask *FileTransferTask) error hasSuccess := false callback = func(nextTask *FileTransferTask) error { @@ -158,7 +160,7 @@ func transfer(ctx context.Context, taskType taskType, srcObjPath, dstDirPath str hasSuccess = true } if taskType == move { - task_group.TransferCoordinator.AppendPayload(t.groupID, task_group.SrcPathToRemove(srcObjPath)) + task_group.TransferCoordinator.RemoveSource(t.groupID, srcObjPath) } task_group.TransferCoordinator.Done(context.WithoutCancel(ctx), t.groupID, hasSuccess) return nil, err @@ -169,7 +171,7 @@ func transfer(ctx context.Context, taskType taskType, srcObjPath, dstDirPath str if taskType == copy || taskType == merge { CopyTaskManager.Add(t) } else { - task_group.TransferCoordinator.AppendPayload(t.groupID, task_group.SrcPathToRemove(srcObjPath)) + task_group.TransferCoordinator.RemoveSource(t.groupID, srcObjPath) MoveTaskManager.Add(t) } return t, nil @@ -189,7 +191,6 @@ func (t *FileTransferTask) RunWithNextTaskCallback(f func(nextTask *FileTransfer return errors.WithMessagef(err, "failed list src [%s] objs", t.SrcActualPath) } dstActualPath := stdpath.Join(t.DstActualPath, srcObj.GetName()) - task_group.TransferCoordinator.AppendPayload(t.groupID, task_group.DstPathToHook(dstActualPath)) existedObjs := make(map[string]bool) if t.TaskType == merge { @@ -259,7 +260,7 @@ func (t *FileTransferTask) RunWithNextTaskCallback(f func(nextTask *FileTransfer } t.SetTotalBytes(ss.GetSize()) t.Status = "uploading" - return op.Put(context.WithValue(t.Ctx(), conf.SkipHookKey, struct{}{}), t.DstStorage, t.DstActualPath, ss, t.SetProgress) + return op.Put(t.Ctx(), t.DstStorage, t.DstActualPath, ss, t.SetProgress) } var ( diff --git a/internal/fs/fs.go b/internal/fs/fs.go index e4cb6ef389..987836bfab 100644 --- a/internal/fs/fs.go +++ b/internal/fs/fs.go @@ -68,32 +68,42 @@ func MakeDir(ctx context.Context, path string) error { return err } -func Move(ctx context.Context, srcPath, dstDirPath string, skipHook ...bool) (task.TaskExtensionInfo, error) { - req, err := transfer(ctx, move, srcPath, dstDirPath, skipHook...) +func Move(ctx context.Context, srcPath, dstDirPath string) (task.TaskExtensionInfo, error) { + req, err := transfer(ctx, move, scheduledTransfer, srcPath, dstDirPath) if err != nil { log.Errorf("failed move %s to %s: %+v", srcPath, dstDirPath, err) } return req, err } -func Copy(ctx context.Context, srcObjPath, dstDirPath string, skipHook ...bool) (task.TaskExtensionInfo, error) { - res, err := transfer(ctx, copy, srcObjPath, dstDirPath, skipHook...) +func MoveDirectly(ctx context.Context, srcPath, dstDirPath string) error { + _, err := transfer(ctx, move, synchronousTransfer, srcPath, dstDirPath) + return err +} + +func Copy(ctx context.Context, srcObjPath, dstDirPath string) (task.TaskExtensionInfo, error) { + res, err := transfer(ctx, copy, scheduledTransfer, srcObjPath, dstDirPath) if err != nil { log.Errorf("failed copy %s to %s: %+v", srcObjPath, dstDirPath, err) } return res, err } -func Merge(ctx context.Context, srcObjPath, dstDirPath string, skipHook ...bool) (task.TaskExtensionInfo, error) { - res, err := transfer(ctx, merge, srcObjPath, dstDirPath, skipHook...) +func CopyDirectly(ctx context.Context, srcObjPath, dstDirPath string) error { + _, err := transfer(ctx, copy, synchronousTransfer, srcObjPath, dstDirPath) + return err +} + +func Merge(ctx context.Context, srcObjPath, dstDirPath string) (task.TaskExtensionInfo, error) { + res, err := transfer(ctx, merge, scheduledTransfer, srcObjPath, dstDirPath) if err != nil { log.Errorf("failed merge %s to %s: %+v", srcObjPath, dstDirPath, err) } return res, err } -func Rename(ctx context.Context, srcPath, dstName string, skipHook ...bool) error { - err := rename(ctx, srcPath, dstName, skipHook...) +func Rename(ctx context.Context, srcPath, dstName string) error { + err := rename(ctx, srcPath, dstName) if err != nil { log.Errorf("failed rename %s to %s: %+v", srcPath, dstName, err) } @@ -108,8 +118,8 @@ func Remove(ctx context.Context, path string) error { return err } -func PutDirectly(ctx context.Context, dstDirPath string, file model.FileStreamer, skipHook ...bool) error { - err := putDirectly(ctx, dstDirPath, file, skipHook...) +func PutDirectly(ctx context.Context, dstDirPath string, file model.FileStreamer) error { + err := putDirectly(ctx, dstDirPath, file) if err != nil { log.Errorf("failed put %s: %+v", dstDirPath, err) } @@ -140,8 +150,8 @@ func ArchiveList(ctx context.Context, path string, args model.ArchiveListArgs) ( return objs, err } -func ArchiveDecompress(ctx context.Context, srcObjPath, dstDirPath string, args model.ArchiveDecompressArgs, lazyCache ...bool) (task.TaskExtensionInfo, error) { - t, err := archiveDecompress(ctx, srcObjPath, dstDirPath, args, lazyCache...) +func ArchiveDecompress(ctx context.Context, srcObjPath, dstDirPath string, args model.ArchiveDecompressArgs) (task.TaskExtensionInfo, error) { + t, err := archiveDecompress(ctx, srcObjPath, dstDirPath, args) if err != nil { log.Errorf("failed decompress [%s]%s: %+v", srcObjPath, args.InnerPath, err) } diff --git a/internal/fs/other.go b/internal/fs/other.go index a23beb73bc..2123c0f824 100644 --- a/internal/fs/other.go +++ b/internal/fs/other.go @@ -3,12 +3,10 @@ package fs import ( "context" - "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/driver" "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/internal/op" "github.com/OpenListTeam/OpenList/v4/internal/task" - "github.com/OpenListTeam/OpenList/v4/pkg/utils" "github.com/pkg/errors" ) @@ -20,14 +18,11 @@ func makeDir(ctx context.Context, path string) error { return op.MakeDir(ctx, storage, actualPath) } -func rename(ctx context.Context, srcPath, dstName string, skipHook ...bool) error { +func rename(ctx context.Context, srcPath, dstName string) error { storage, srcActualPath, err := op.GetStorageAndActualPath(srcPath) if err != nil { return errors.WithMessage(err, "failed get storage") } - if utils.IsBool(skipHook...) { - ctx = context.WithValue(ctx, conf.SkipHookKey, struct{}{}) - } return op.Rename(ctx, storage, srcActualPath, dstName) } diff --git a/internal/fs/put.go b/internal/fs/put.go index 9042ae88a0..e8d1c2ee1e 100644 --- a/internal/fs/put.go +++ b/internal/fs/put.go @@ -6,8 +6,6 @@ import ( stdpath "path" "time" - "github.com/OpenListTeam/OpenList/v4/pkg/utils" - "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/driver" "github.com/OpenListTeam/OpenList/v4/internal/errs" @@ -38,7 +36,7 @@ func (t *UploadTask) Run() error { t.ClearEndTime() t.SetStartTime(time.Now()) defer func() { t.SetEndTime(time.Now()) }() - return op.Put(context.WithValue(t.Ctx(), conf.SkipHookKey, struct{}{}), t.storage, t.dstDirActualPath, t.file, t.SetProgress) + return op.Put(t.Ctx(), t.storage, t.dstDirActualPath, t.file, t.SetProgress) } func (t *UploadTask) OnSucceeded() { @@ -53,7 +51,7 @@ func (t *UploadTask) SetRetry(retry int, maxRetry int) { t.TaskExtension.SetRetry(retry, maxRetry) if retry == 0 && (t.GetErr() == nil && t.GetState() != tache.StatePending) { // 手动重试 - task_group.TransferCoordinator.AddTask(stdpath.Join(t.storage.GetStorage().MountPath, t.dstDirActualPath), nil) + task_group.TransferCoordinator.AddTask(stdpath.Join(t.storage.GetStorage().MountPath, t.dstDirActualPath)) } } @@ -87,13 +85,13 @@ func putAsTask(ctx context.Context, dstDirPath string, file model.FileStreamer) file: file, } t.SetTotalBytes(file.GetSize()) - task_group.TransferCoordinator.AddTask(stdpath.Join(storage.GetStorage().MountPath, dstDirActualPath), nil) + task_group.TransferCoordinator.AddTask(stdpath.Join(storage.GetStorage().MountPath, dstDirActualPath)) UploadTaskManager.Add(t) return t, nil } // putDirect put the file and return after finish -func putDirectly(ctx context.Context, dstDirPath string, file model.FileStreamer, skipHook ...bool) error { +func putDirectly(ctx context.Context, dstDirPath string, file model.FileStreamer) error { storage, dstDirActualPath, err := op.GetStorageAndActualPath(dstDirPath) if err != nil { _ = file.Close() @@ -103,9 +101,6 @@ func putDirectly(ctx context.Context, dstDirPath string, file model.FileStreamer _ = file.Close() return errors.WithStack(errs.UploadNotSupported) } - if utils.IsBool(skipHook...) { - ctx = context.WithValue(ctx, conf.SkipHookKey, struct{}{}) - } return op.Put(ctx, storage, dstDirActualPath, file, nil) } diff --git a/internal/model/args.go b/internal/model/args.go index 16a5c1722f..87cedab082 100644 --- a/internal/model/args.go +++ b/internal/model/args.go @@ -15,7 +15,6 @@ type ListArgs struct { S3ShowPlaceholder bool Refresh bool WithStorageDetails bool - SkipHook bool } type LinkArgs struct { diff --git a/internal/offline_download/tool/download.go b/internal/offline_download/tool/download.go index ca87768397..39d6f8dc68 100644 --- a/internal/offline_download/tool/download.go +++ b/internal/offline_download/tool/download.go @@ -208,7 +208,7 @@ func (t *DownloadTask) Transfer() error { } tsk.SetTotalBytes(t.GetTotalBytes()) tsk.groupID = path.Join(tsk.DstStorageMp, tsk.DstActualPath) - task_group.TransferCoordinator.AddTask(tsk.groupID, nil) + task_group.TransferCoordinator.AddTask(tsk.groupID) TransferTaskManager.Add(tsk) return nil } diff --git a/internal/offline_download/tool/transfer.go b/internal/offline_download/tool/transfer.go index 7c5bd164a1..9a09181c9a 100644 --- a/internal/offline_download/tool/transfer.go +++ b/internal/offline_download/tool/transfer.go @@ -74,7 +74,7 @@ func (t *TransferTask) Run() error { Mimetype: mimetype, Closers: utils.NewClosers(r), } - return op.Put(context.WithValue(t.Ctx(), conf.SkipHookKey, struct{}{}), t.DstStorage, t.DstActualPath, s, t.SetProgress) + return op.Put(t.Ctx(), t.DstStorage, t.DstActualPath, s, t.SetProgress) } return transferStdPath(t) } @@ -115,7 +115,7 @@ func (t *TransferTask) SetRetry(retry int, maxRetry int) { (len(t.groupID) == 0 || // 重启恢复 (t.GetErr() == nil && t.GetState() != tache.StatePending)) { // 手动重试 t.groupID = stdpath.Join(t.DstStorageMp, t.DstActualPath) - task_group.TransferCoordinator.AddTask(t.groupID, nil) + task_group.TransferCoordinator.AddTask(t.groupID) } t.TaskData.SetRetry(retry, maxRetry) } @@ -149,7 +149,7 @@ func transferStd(ctx context.Context, tempDir, dstDirPath string, deletePolicy D DeletePolicy: deletePolicy, } t.groupID = path.Join(t.DstStorageMp, t.DstActualPath) - task_group.TransferCoordinator.AddTask(t.groupID, nil) + task_group.TransferCoordinator.AddTask(t.groupID) TransferTaskManager.Add(t) } return nil @@ -168,7 +168,6 @@ func transferStdPath(t *TransferTask) error { return err } dstDirActualPath := stdpath.Join(t.DstActualPath, info.Name()) - task_group.TransferCoordinator.AppendPayload(t.groupID, task_group.DstPathToHook(dstDirActualPath)) for _, entry := range entries { srcRawPath := stdpath.Join(t.SrcActualPath, entry.Name()) task := &TransferTask{ @@ -186,7 +185,7 @@ func transferStdPath(t *TransferTask) error { groupID: t.groupID, DeletePolicy: t.DeletePolicy, } - task_group.TransferCoordinator.AddTask(t.groupID, nil) + task_group.TransferCoordinator.AddTask(t.groupID) TransferTaskManager.Add(task) } t.Status = "src object is dir, added all transfer tasks of files" @@ -238,7 +237,7 @@ func transferStdFile(t *TransferTask) error { Closers: utils.NewClosers(rc), } t.SetTotalBytes(info.Size()) - err = op.Put(context.WithValue(t.Ctx(), conf.SkipHookKey, struct{}{}), t.DstStorage, t.DstActualPath, s, t.SetProgress) + err = op.Put(t.Ctx(), t.DstStorage, t.DstActualPath, s, t.SetProgress) if err != nil { return err } @@ -287,7 +286,7 @@ func transferObj(ctx context.Context, tempDir, dstDirPath string, deletePolicy D DeletePolicy: deletePolicy, } t.groupID = path.Join(t.DstStorageMp, t.DstActualPath) - task_group.TransferCoordinator.AddTask(t.groupID, nil) + task_group.TransferCoordinator.AddTask(t.groupID) TransferTaskManager.Add(t) } return nil @@ -306,13 +305,12 @@ func transferObjPath(t *TransferTask) error { return errors.WithMessagef(err, "failed list src [%s] objs", t.SrcActualPath) } dstDirActualPath := stdpath.Join(t.DstActualPath, srcObj.GetName()) - task_group.TransferCoordinator.AppendPayload(t.groupID, task_group.DstPathToHook(dstDirActualPath)) for _, obj := range objs { if utils.IsCanceled(t.Ctx()) { return nil } srcObjPath := stdpath.Join(t.SrcActualPath, obj.GetName()) - task_group.TransferCoordinator.AddTask(t.groupID, nil) + task_group.TransferCoordinator.AddTask(t.groupID) TransferTaskManager.Add(&TransferTask{ TaskData: fs.TaskData{ TaskExtension: task.TaskExtension{ @@ -355,7 +353,7 @@ func transferObjFile(t *TransferTask) error { return errors.WithMessagef(err, "failed get [%s] stream", t.SrcActualPath) } t.SetTotalBytes(ss.GetSize()) - return op.Put(context.WithValue(t.Ctx(), conf.SkipHookKey, struct{}{}), t.DstStorage, t.DstActualPath, ss, t.SetProgress) + return op.Put(t.Ctx(), t.DstStorage, t.DstActualPath, ss, t.SetProgress) } func removeObjTemp(t *TransferTask) { diff --git a/internal/op/archive.go b/internal/op/archive.go index bb3de11a10..4918f5358b 100644 --- a/internal/op/archive.go +++ b/internal/op/archive.go @@ -11,7 +11,6 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/archive/tool" "github.com/OpenListTeam/OpenList/v4/internal/cache" - "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/driver" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/model" @@ -21,7 +20,6 @@ import ( gocache "github.com/OpenListTeam/go-cache" "github.com/pkg/errors" log "github.com/sirupsen/logrus" - "golang.org/x/time/rate" ) var ( @@ -494,7 +492,7 @@ func InternalExtract(ctx context.Context, storage driver.Driver, path string, ar return &streamWithParent{rc: rc, parents: ss}, size, nil } -func ArchiveDecompress(ctx context.Context, storage driver.Driver, srcPath, dstDirPath string, args model.ArchiveDecompressArgs, lazyCache ...bool) error { +func ArchiveDecompress(ctx context.Context, storage driver.Driver, srcPath, dstDirPath string, args model.ArchiveDecompressArgs) error { if storage.Config().CheckStatus && storage.GetStorage().Status != WORK { return errors.WithMessagef(errs.StorageNotInit, "storage status: %s", storage.GetStorage().Status) } @@ -509,32 +507,39 @@ func ArchiveDecompress(ctx context.Context, storage driver.Driver, srcPath, dstD return errors.WithMessage(err, "failed to get dst dir") } + mutationCtx, finishMutation := enterMutationFrame(ctx, storage, reconciliationTarget{path: dstDirPath, recursive: true}) + defer finishMutation() var newObjs []model.Obj + var mutateCache func() switch s := storage.(type) { case driver.ArchiveDecompressResult: - newObjs, err = s.ArchiveDecompress(ctx, srcObj, dstDir, args) + newObjs, err = s.ArchiveDecompress(mutationCtx, srcObj, dstDir, args) if err == nil { - if len(newObjs) > 0 { - if !storage.Config().NoCache { + mutateCache = func() { + if len(newObjs) > 0 && !storage.Config().NoCache { if cache, exist := Cache.dirCache.Get(Key(storage, dstDirPath)); exist { for _, newObj := range newObjs { cache.UpdateObject(newObj.GetName(), newObj) } } + } else if len(newObjs) == 0 && !storage.Config().NoCache { + Cache.dirCache.Delete(Key(storage, dstDirPath)) } - } else if !utils.IsBool(lazyCache...) { - Cache.DeleteDirectory(storage, dstDirPath) } } case driver.ArchiveDecompress: - err = s.ArchiveDecompress(ctx, srcObj, dstDir, args) - if err == nil && !utils.IsBool(lazyCache...) { - Cache.DeleteDirectory(storage, dstDirPath) + err = s.ArchiveDecompress(mutationCtx, srcObj, dstDir, args) + if err == nil { + mutateCache = func() { + if !storage.Config().NoCache { + Cache.dirCache.Delete(Key(storage, dstDirPath)) + } + } } default: return errs.NotImplement } - if !utils.IsBool(lazyCache...) && err == nil && needHandleObjsUpdateHook() { + if err == nil { onlyList := false targetPath := dstDirPath if len(newObjs) == 1 && newObjs[0].IsDir() { @@ -549,15 +554,9 @@ func ArchiveDecompress(ctx context.Context, storage driver.Driver, srcPath, dstD onlyList = e != nil || !dstObj.IsDir() } if onlyList { - go List(context.Background(), storage, dstDirPath, model.ListArgs{Refresh: true}) + commitNamespaceMutation(mutationCtx, storage, []string{Key(storage, dstDirPath)}, mutateCache, reconciliationTarget{path: dstDirPath}) } else { - var limiter *rate.Limiter - if l, _ := GetSettingItemByKey(conf.HandleHookRateLimit); l != nil { - if f, e := strconv.ParseFloat(l.Value, 64); e == nil && f > .0 { - limiter = rate.NewLimiter(rate.Limit(f), 1) - } - } - go RecursivelyListStorage(context.Background(), storage, targetPath, limiter, nil) + commitNamespaceMutation(mutationCtx, storage, []string{Key(storage, dstDirPath)}, mutateCache, reconciliationTarget{path: targetPath, recursive: true}) } } return errors.WithStack(err) diff --git a/internal/op/cache.go b/internal/op/cache.go index d8d32a74b6..22ed35ecc7 100644 --- a/internal/op/cache.go +++ b/internal/op/cache.go @@ -17,6 +17,8 @@ type CacheManager struct { userCache *cache.KeyedCache[*model.User] // Cache for user data settingCache *cache.KeyedCache[any] // Cache for settings detailCache *cache.KeyedCache[*model.StorageDetails] // Cache for storage details + loadMu sync.Mutex + loads map[string]*directoryLoadState } func NewCacheManager() *CacheManager { @@ -26,9 +28,84 @@ func NewCacheManager() *CacheManager { userCache: cache.NewKeyedCache[*model.User](time.Hour), settingCache: cache.NewKeyedCache[any](time.Hour), detailCache: cache.NewKeyedCache[*model.StorageDetails](time.Minute * 30), + loads: make(map[string]*directoryLoadState), } } +type directoryLoadState struct { + revision uint64 + active int + refreshing int +} + +type directoryLoadToken struct { + cache *CacheManager + key string + revision uint64 + refresh bool + commit bool +} + +func (cm *CacheManager) beginDirectoryLoad(key string, refresh bool) directoryLoadToken { + cm.loadMu.Lock() + defer cm.loadMu.Unlock() + state := cm.loads[key] + if state == nil { + state = &directoryLoadState{} + cm.loads[key] = state + } + if refresh { + state.revision++ + state.refreshing++ + } + state.active++ + return directoryLoadToken{ + cache: cm, + key: key, + revision: state.revision, + refresh: refresh, + commit: refresh || state.refreshing == 0, + } +} + +func (token directoryLoadToken) commitIfCurrent(commit func()) bool { + token.cache.loadMu.Lock() + defer token.cache.loadMu.Unlock() + state := token.cache.loads[token.key] + if !token.commit || state == nil || state.revision != token.revision { + return false + } + commit() + return true +} + +func (token directoryLoadToken) done() { + token.cache.loadMu.Lock() + defer token.cache.loadMu.Unlock() + state := token.cache.loads[token.key] + if state == nil { + return + } + state.active-- + if token.refresh { + state.refreshing-- + } + if state.active == 0 { + delete(token.cache.loads, token.key) + } +} + +func (cm *CacheManager) mutateDirectories(keys []string, mutate func()) { + cm.loadMu.Lock() + defer cm.loadMu.Unlock() + for _, key := range keys { + if state := cm.loads[key]; state != nil { + state.revision++ + } + } + mutate() +} + // global instance var Cache = NewCacheManager() @@ -167,10 +244,12 @@ const ( ) func newDirectoryCache(objs []model.Obj) *directoryCache { + owned := make([]model.Obj, len(objs)) + copy(owned, objs) sorted := make([]model.Obj, len(objs)) copy(sorted, objs) return &directoryCache{ - objs: objs, + objs: owned, sorted: sorted, } } diff --git a/internal/op/fs.go b/internal/op/fs.go index e6278ca0aa..c299091dbc 100644 --- a/internal/op/fs.go +++ b/internal/op/fs.go @@ -7,7 +7,6 @@ import ( "strings" "time" - "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/driver" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/model" @@ -17,7 +16,6 @@ import ( "github.com/bmatcuk/doublestar/v4" "github.com/pkg/errors" log "github.com/sirupsen/logrus" - "golang.org/x/time/rate" ) var listG singleflight.Group[[]model.Obj] @@ -32,9 +30,15 @@ func list(ctx context.Context, storage driver.Driver, path string, args model.Li return nil, errors.WithMessagef(errs.StorageNotInit, "storage status: %s", storage.GetStorage().Status) } path = utils.FixAndCleanPath(path) + ctx, ownsProjection := enterSnapshotFrame(ctx) log.Debugf("op.List %s", path) key := Key(storage, path) - if !args.Refresh { + canonicalReqPath := utils.GetFullPath(storage.GetStorage().MountPath, path) + if args.ReqPath == "" { + args.ReqPath = canonicalReqPath + } + cacheable := args.ReqPath == canonicalReqPath && !args.S3ShowPlaceholder && !args.WithStorageDetails + if cacheable && !args.Refresh { if dirCache, exists := Cache.dirCache.Get(key); exists { log.Debugf("use cache when list %s", path) objs := dirCache.GetSortedObjects(storage) @@ -48,7 +52,22 @@ func list(ctx context.Context, storage driver.Driver, path string, args model.Li } } - objs, err, _ := listG.Do(key, func() ([]model.Obj, error) { + flightKey := key + "\x00" + args.ReqPath + if args.S3ShowPlaceholder { + flightKey += "\x00placeholder" + } + if args.WithStorageDetails { + flightKey += "\x00details" + } + if args.Refresh { + flightKey += "\x00fresh" + } + objs, err, _ := listG.Do(flightKey, func() ([]model.Obj, error) { + var load directoryLoadToken + if cacheable { + load = Cache.beginDirectoryLoad(key, args.Refresh) + defer load.done() + } dir, err := GetUnwrap(ctx, storage, path) if err != nil { return nil, errors.WithMessage(err, "failed get dir") @@ -69,46 +88,47 @@ func list(ctx context.Context, storage driver.Driver, path string, args model.Li } model.ExtractFolder(files, storage.GetStorage().ExtractFolder) - if !args.SkipHook { - // call hooks - go func(reqPath string, files []model.Obj) { - HandleObjsUpdateHook(context.WithoutCancel(ctx), reqPath, files) - }(utils.GetFullPath(storage.GetStorage().MountPath, path), files) - } - - if !storage.Config().NoCache { - if len(files) > 0 { - log.Debugf("set cache: %s => %+v", key, files) - - ttl := storage.GetStorage().CacheExpiration - - customCachePolicies := storage.GetStorage().CustomCachePolicies - if len(customCachePolicies) > 0 { - for configPolicy := range strings.SplitSeq(customCachePolicies, "\n") { - pattern, ttlstr, ok := strings.Cut(strings.TrimSpace(configPolicy), ":") - if !ok { - log.Warnf("Malformed custom cache policy entry: %s in storage %s for path %s. Expected format: pattern:ttl", configPolicy, storage.GetStorage().MountPath, path) - continue - } - if match, err1 := doublestar.Match(pattern, path); err1 != nil { - log.Warnf("Invalid glob pattern in custom cache policy: %s, error: %v", pattern, err1) - continue - } else if !match { - continue - } - - if configTtl, err1 := strconv.ParseInt(ttlstr, 10, 64); err1 == nil { - ttl = int(configTtl) - break + if cacheable { + accepted := load.commitIfCurrent(func() { + if !storage.Config().NoCache && len(files) > 0 { + log.Debugf("set cache: %s => %+v", key, files) + + ttl := storage.GetStorage().CacheExpiration + + customCachePolicies := storage.GetStorage().CustomCachePolicies + if len(customCachePolicies) > 0 { + for configPolicy := range strings.SplitSeq(customCachePolicies, "\n") { + pattern, ttlstr, ok := strings.Cut(strings.TrimSpace(configPolicy), ":") + if !ok { + log.Warnf("Malformed custom cache policy entry: %s in storage %s for path %s. Expected format: pattern:ttl", configPolicy, storage.GetStorage().MountPath, path) + continue + } + if match, err1 := doublestar.Match(pattern, path); err1 != nil { + log.Warnf("Invalid glob pattern in custom cache policy: %s, error: %v", pattern, err1) + continue + } else if !match { + continue + } + + if configTtl, err1 := strconv.ParseInt(ttlstr, 10, 64); err1 == nil { + ttl = int(configTtl) + break + } } } - } - duration := time.Minute * time.Duration(ttl) - Cache.dirCache.SetWithTTL(key, newDirectoryCache(files), duration) - } else { - log.Debugf("del cache: %s", key) - Cache.deleteDirectoryTree(key) + duration := time.Minute * time.Duration(ttl) + Cache.dirCache.SetWithTTL(key, newDirectoryCache(files), duration) + } else if !storage.Config().NoCache { + log.Debugf("del cache: %s", key) + Cache.deleteDirectoryTree(key) + } + if ownsProjection { + ProjectSnapshot(context.WithoutCancel(ctx), canonicalReqPath, files) + } + }) + if !accepted { + log.Debugf("discard stale list snapshot: %s", path) } } return files, nil @@ -342,33 +362,37 @@ func MakeDir(ctx context.Context, storage driver.Driver, path string) error { } var newObj model.Obj + mutationCtx, finishMutation := enterMutationFrame(ctx, storage, reconciliationTarget{path: parentPath}) + defer finishMutation() switch s := storage.(type) { case driver.MkdirResult: - newObj, err = s.MakeDir(ctx, parentDir, dirName) + newObj, err = s.MakeDir(mutationCtx, parentDir, dirName) case driver.Mkdir: - err = s.MakeDir(ctx, parentDir, dirName) + err = s.MakeDir(mutationCtx, parentDir, dirName) default: return nil, errs.NotImplement } if err != nil && !errs.IsObjectAlreadyExists(err) { return nil, errors.WithStack(err) } - if storage.Config().NoCache { - return nil, nil - } - if dirCache, exist := Cache.dirCache.Get(Key(storage, parentPath)); exist { - if newObj == nil { - t := time.Now() - newObj = &model.Object{ - Name: dirName, - IsFolder: true, - Modified: t, - Ctime: t, - Mask: model.Temp, + parentKey := Key(storage, parentPath) + commitNamespaceMutation(mutationCtx, storage, []string{parentKey}, func() { + if !storage.Config().NoCache { + if dirCache, exist := Cache.dirCache.Get(parentKey); exist { + if newObj == nil { + t := time.Now() + newObj = &model.Object{ + Name: dirName, + IsFolder: true, + Modified: t, + Ctime: t, + Mask: model.Temp, + } + } + dirCache.UpdateObject("", wrapObjName(storage, newObj)) } } - dirCache.UpdateObject("", wrapObjName(storage, newObj)) - } + }, reconciliationTarget{path: parentPath}) return nil, nil }) return err @@ -403,12 +427,20 @@ func Move(ctx context.Context, storage driver.Driver, srcPath, dstDirPath string return errors.WithStack(errs.PermissionDenied) } + targets := []reconciliationTarget{{path: srcDirPath}} + if srcObj.IsDir() { + targets = append(targets, reconciliationTarget{path: stdpath.Join(dstDirPath, srcObj.GetName()), recursive: true}) + } else { + targets = append(targets, reconciliationTarget{path: dstDirPath}) + } + mutationCtx, finishMutation := enterMutationFrame(ctx, storage, targets...) + defer finishMutation() var newObj model.Obj switch s := storage.(type) { case driver.MoveResult: - newObj, err = s.Move(ctx, srcObj, dstDir) + newObj, err = s.Move(mutationCtx, srcObj, dstDir) case driver.Move: - err = s.Move(ctx, srcObj, dstDir) + err = s.Move(mutationCtx, srcObj, dstDir) default: err = errs.NotImplement } @@ -418,35 +450,28 @@ func Move(ctx context.Context, storage driver.Driver, srcPath, dstDirPath string srcKey := Key(storage, srcDirPath) dstKey := Key(storage, dstDirPath) - if !srcRawObj.IsDir() { - Cache.linkCache.DeleteKey(stdpath.Join(srcKey, srcRawObj.GetName())) - Cache.linkCache.DeleteKey(stdpath.Join(dstKey, srcRawObj.GetName())) - } - if !storage.Config().NoCache { - if cache, exist := Cache.dirCache.Get(srcKey); exist { - if srcRawObj.IsDir() { - Cache.deleteDirectoryTree(stdpath.Join(srcKey, srcRawObj.GetName())) - } - cache.RemoveObject(srcRawObj.GetName()) + commitNamespaceMutation(mutationCtx, storage, []string{srcKey, dstKey}, func() { + if !srcRawObj.IsDir() { + Cache.linkCache.DeleteKey(stdpath.Join(srcKey, srcRawObj.GetName())) + Cache.linkCache.DeleteKey(stdpath.Join(dstKey, srcRawObj.GetName())) } - if cache, exist := Cache.dirCache.Get(dstKey); exist { - if newObj == nil { - newObj = &model.ObjWrapMask{Obj: srcRawObj, Mask: model.Temp} - } else { - newObj = wrapObjName(storage, newObj) + if !storage.Config().NoCache { + if cache, exist := Cache.dirCache.Get(srcKey); exist { + if srcRawObj.IsDir() { + Cache.deleteDirectoryTree(stdpath.Join(srcKey, srcRawObj.GetName())) + } + cache.RemoveObject(srcRawObj.GetName()) + } + if cache, exist := Cache.dirCache.Get(dstKey); exist { + if newObj == nil { + newObj = &model.ObjWrapMask{Obj: srcRawObj, Mask: model.Temp} + } else { + newObj = wrapObjName(storage, newObj) + } + cache.UpdateObject(srcRawObj.GetName(), newObj) } - cache.UpdateObject(srcRawObj.GetName(), newObj) } - } - - if ctx.Value(conf.SkipHookKey) != nil || !needHandleObjsUpdateHook() { - return nil - } - if !srcObj.IsDir() { - go objsUpdateHook(context.WithoutCancel(ctx), storage, dstDirPath, false) - } else { - go objsUpdateHook(context.WithoutCancel(ctx), storage, stdpath.Join(dstDirPath, srcObj.GetName()), true) - } + }, targets...) return nil } @@ -468,12 +493,19 @@ func Rename(ctx context.Context, storage driver.Driver, srcPath, dstName string) oldName := srcRawObj.GetName() srcObj := model.UnwrapObjName(srcRawObj) + dirPath := stdpath.Dir(srcPath) + targets := []reconciliationTarget{{path: dirPath}} + if srcObj.IsDir() { + targets = append(targets, reconciliationTarget{path: stdpath.Join(dirPath, dstName), recursive: true}) + } + mutationCtx, finishMutation := enterMutationFrame(ctx, storage, targets...) + defer finishMutation() var newObj model.Obj switch s := storage.(type) { case driver.RenameResult: - newObj, err = s.Rename(ctx, srcObj, dstName) + newObj, err = s.Rename(mutationCtx, srcObj, dstName) case driver.Rename: - err = s.Rename(ctx, srcObj, dstName) + err = s.Rename(mutationCtx, srcObj, dstName) default: return errs.NotImplement } @@ -481,33 +513,25 @@ func Rename(ctx context.Context, storage driver.Driver, srcPath, dstName string) return errors.WithStack(err) } - dirKey := Key(storage, stdpath.Dir(srcPath)) - if !srcRawObj.IsDir() { - Cache.linkCache.DeleteKey(stdpath.Join(dirKey, oldName)) - Cache.linkCache.DeleteKey(stdpath.Join(dirKey, dstName)) - } - if !storage.Config().NoCache { - if cache, exist := Cache.dirCache.Get(dirKey); exist { - if srcRawObj.IsDir() { - Cache.deleteDirectoryTree(stdpath.Join(dirKey, oldName)) - } - if newObj == nil { - newObj = &model.ObjWrapMask{Obj: &model.ObjWrapName{Name: dstName, Obj: srcObj}, Mask: model.Temp} + dirKey := Key(storage, dirPath) + commitNamespaceMutation(mutationCtx, storage, []string{dirKey}, func() { + if !srcRawObj.IsDir() { + Cache.linkCache.DeleteKey(stdpath.Join(dirKey, oldName)) + Cache.linkCache.DeleteKey(stdpath.Join(dirKey, dstName)) + } + if !storage.Config().NoCache { + if cache, exist := Cache.dirCache.Get(dirKey); exist { + if srcRawObj.IsDir() { + Cache.deleteDirectoryTree(stdpath.Join(dirKey, oldName)) + } + if newObj == nil { + newObj = &model.ObjWrapMask{Obj: &model.ObjWrapName{Name: dstName, Obj: srcObj}, Mask: model.Temp} + } + newObj = wrapObjName(storage, newObj) + cache.UpdateObject(oldName, newObj) } - newObj = wrapObjName(storage, newObj) - cache.UpdateObject(oldName, newObj) } - } - - if ctx.Value(conf.SkipHookKey) != nil || !needHandleObjsUpdateHook() { - return nil - } - dstDirPath := stdpath.Dir(srcPath) - if !srcObj.IsDir() { - go objsUpdateHook(context.WithoutCancel(ctx), storage, dstDirPath, false) - } else { - go objsUpdateHook(context.WithoutCancel(ctx), storage, stdpath.Join(dstDirPath, srcObj.GetName()), true) - } + }, targets...) return nil } @@ -537,12 +561,18 @@ func Copy(ctx context.Context, storage driver.Driver, srcPath, dstDirPath string return errors.WithStack(errs.PermissionDenied) } + target := reconciliationTarget{path: dstDirPath} + if srcObj.IsDir() { + target = reconciliationTarget{path: stdpath.Join(dstDirPath, srcObj.GetName()), recursive: true} + } + mutationCtx, finishMutation := enterMutationFrame(ctx, storage, target) + defer finishMutation() var newObj model.Obj switch s := storage.(type) { case driver.CopyResult: - newObj, err = s.Copy(ctx, srcObj, dstDir) + newObj, err = s.Copy(mutationCtx, srcObj, dstDir) case driver.Copy: - err = s.Copy(ctx, srcObj, dstDir) + err = s.Copy(mutationCtx, srcObj, dstDir) default: err = errs.NotImplement } @@ -551,28 +581,21 @@ func Copy(ctx context.Context, storage driver.Driver, srcPath, dstDirPath string } dstKey := Key(storage, dstDirPath) - if !srcRawObj.IsDir() { - Cache.linkCache.DeleteKey(stdpath.Join(dstKey, srcRawObj.GetName())) - } - if !storage.Config().NoCache { - if cache, exist := Cache.dirCache.Get(dstKey); exist { - if newObj == nil { - newObj = &model.ObjWrapMask{Obj: srcRawObj, Mask: model.Temp} - } else { - newObj = wrapObjName(storage, newObj) + commitNamespaceMutation(mutationCtx, storage, []string{dstKey}, func() { + if !srcRawObj.IsDir() { + Cache.linkCache.DeleteKey(stdpath.Join(dstKey, srcRawObj.GetName())) + } + if !storage.Config().NoCache { + if cache, exist := Cache.dirCache.Get(dstKey); exist { + if newObj == nil { + newObj = &model.ObjWrapMask{Obj: srcRawObj, Mask: model.Temp} + } else { + newObj = wrapObjName(storage, newObj) + } + cache.UpdateObject(srcRawObj.GetName(), newObj) } - cache.UpdateObject(srcRawObj.GetName(), newObj) } - } - - if ctx.Value(conf.SkipHookKey) != nil || !needHandleObjsUpdateHook() { - return nil - } - if !srcObj.IsDir() { - go objsUpdateHook(context.WithoutCancel(ctx), storage, dstDirPath, false) - } else { - go objsUpdateHook(context.WithoutCancel(ctx), storage, stdpath.Join(dstDirPath, srcObj.GetName()), true) - } + }, target) return nil } @@ -597,12 +620,16 @@ func Remove(ctx context.Context, storage driver.Driver, path string) error { return errors.WithStack(errs.PermissionDenied) } dirPath := stdpath.Dir(path) + mutationCtx, finishMutation := enterMutationFrame(ctx, storage, reconciliationTarget{path: dirPath}) + defer finishMutation() switch s := storage.(type) { case driver.Remove: - err = s.Remove(ctx, model.UnwrapObjName(rawObj)) + err = s.Remove(mutationCtx, model.UnwrapObjName(rawObj)) if err == nil { - Cache.removeDirectoryObject(storage, dirPath, rawObj) + commitNamespaceMutation(mutationCtx, storage, []string{Key(storage, dirPath)}, func() { + Cache.removeDirectoryObject(storage, dirPath, rawObj) + }, reconciliationTarget{path: dirPath}) } default: return errs.NotImplement @@ -670,48 +697,49 @@ func Put(ctx context.Context, storage driver.Driver, dstDirPath string, file mod file.CacheFullAndWriter(nil, nil) } + mutationCtx, finishMutation := enterMutationFrame(ctx, storage, reconciliationTarget{path: dstDirPath}) + defer finishMutation() var newObj model.Obj switch s := storage.(type) { case driver.PutResult: - newObj, err = s.Put(ctx, parentDir, file, up) + newObj, err = s.Put(mutationCtx, parentDir, file, up) case driver.Put: - err = s.Put(ctx, parentDir, file, up) + err = s.Put(mutationCtx, parentDir, file, up) default: return errs.NotImplement } if err == nil { - Cache.linkCache.DeleteKey(Key(storage, dstPath)) - if !storage.Config().NoCache { - if cache, exist := Cache.dirCache.Get(Key(storage, dstDirPath)); exist { - if newObj == nil { - newObj = &model.Object{ - Name: file.GetName(), - Size: file.GetSize(), - Modified: file.ModTime(), - Ctime: file.CreateTime(), - Mask: model.Temp, + dstKey := Key(storage, dstDirPath) + commitNamespaceMutation(mutationCtx, storage, []string{dstKey}, func() { + Cache.linkCache.DeleteKey(Key(storage, dstPath)) + if !storage.Config().NoCache { + if cache, exist := Cache.dirCache.Get(dstKey); exist { + if newObj == nil { + newObj = &model.Object{ + Name: file.GetName(), + Size: file.GetSize(), + Modified: file.ModTime(), + Ctime: file.CreateTime(), + Mask: model.Temp, + } } + newObj = wrapObjName(storage, newObj) + cache.UpdateObject(newObj.GetName(), newObj) } - newObj = wrapObjName(storage, newObj) - cache.UpdateObject(newObj.GetName(), newObj) } - } - - if ctx.Value(conf.SkipHookKey) == nil && needHandleObjsUpdateHook() { - go objsUpdateHook(context.WithoutCancel(ctx), storage, dstDirPath, false) - } + }, reconciliationTarget{path: dstDirPath}) } log.Debugf("put file [%s] done", file.GetName()) if storage.Config().NoOverwriteUpload && fi != nil && fi.GetSize() > 0 { if err != nil { // upload failed, recover old obj - err := Rename(ctx, storage, tempPath, file.GetName()) + err := Rename(mutationCtx, storage, tempPath, file.GetName()) if err != nil { log.Errorf("failed recover old obj: %+v", err) } } else { // upload success, remove old obj - err = Remove(ctx, storage, tempPath) + err = Remove(mutationCtx, storage, tempPath) } } return errors.WithStack(err) @@ -738,36 +766,37 @@ func PutURL(ctx context.Context, storage driver.Driver, dstDirPath, dstName, url if model.ObjHasMask(dstDir, model.NoWrite) { return errors.WithStack(errs.PermissionDenied) } + mutationCtx, finishMutation := enterMutationFrame(ctx, storage, reconciliationTarget{path: dstDirPath}) + defer finishMutation() var newObj model.Obj switch s := storage.(type) { case driver.PutURLResult: - newObj, err = s.PutURL(ctx, dstDir, dstName, url) + newObj, err = s.PutURL(mutationCtx, dstDir, dstName, url) case driver.PutURL: - err = s.PutURL(ctx, dstDir, dstName, url) + err = s.PutURL(mutationCtx, dstDir, dstName, url) default: return errors.WithStack(errs.NotImplement) } if err == nil { - Cache.linkCache.DeleteKey(Key(storage, dstPath)) - if !storage.Config().NoCache { - if cache, exist := Cache.dirCache.Get(Key(storage, dstDirPath)); exist { - if newObj == nil { - t := time.Now() - newObj = &model.Object{ - Name: dstName, - Modified: t, - Ctime: t, - Mask: model.Temp, + dstKey := Key(storage, dstDirPath) + commitNamespaceMutation(mutationCtx, storage, []string{dstKey}, func() { + Cache.linkCache.DeleteKey(Key(storage, dstPath)) + if !storage.Config().NoCache { + if cache, exist := Cache.dirCache.Get(dstKey); exist { + if newObj == nil { + t := time.Now() + newObj = &model.Object{ + Name: dstName, + Modified: t, + Ctime: t, + Mask: model.Temp, + } } + newObj = wrapObjName(storage, newObj) + cache.UpdateObject(newObj.GetName(), newObj) } - newObj = wrapObjName(storage, newObj) - cache.UpdateObject(newObj.GetName(), newObj) - } - - if ctx.Value(conf.SkipHookKey) == nil && needHandleObjsUpdateHook() { - go objsUpdateHook(context.WithoutCancel(ctx), storage, dstDirPath, false) } - } + }, reconciliationTarget{path: dstDirPath}) } log.Debugf("put url [%s](%s) done", dstName, url) return errors.WithStack(err) @@ -819,53 +848,6 @@ func GetDirectUploadInfo(ctx context.Context, tool string, storage driver.Driver return info, nil } -func objsUpdateHook(ctx context.Context, storage driver.Driver, dirPath string, recursive bool) { - files, err := List(ctx, storage, dirPath, model.ListArgs{SkipHook: true}) - if err != nil { - return - } - if !recursive { - HandleObjsUpdateHook(ctx, utils.GetFullPath(storage.GetStorage().MountPath, dirPath), files) - return - } - var limiter *rate.Limiter - if l, _ := GetSettingItemByKey(conf.HandleHookRateLimit); l != nil { - if f, e := strconv.ParseFloat(l.Value, 64); e == nil && f > .0 { - limiter = rate.NewLimiter(rate.Limit(f), 1) - } - } - recursivelyObjsUpdateHook(ctx, storage, dirPath, files, limiter) -} -func recursivelyObjsUpdateHook(ctx context.Context, storage driver.Driver, dirPath string, files []model.Obj, limiter *rate.Limiter) { - HandleObjsUpdateHook(ctx, utils.GetFullPath(storage.GetStorage().MountPath, dirPath), files) - for _, f := range files { - if utils.IsCanceled(ctx) { - return - } - if !f.IsDir() { - continue - } - dstPath := stdpath.Join(dirPath, f.GetName()) - if limiter != nil { - if err := limiter.Wait(ctx); err != nil { - return - } - } - files, err := List(ctx, storage, dstPath, model.ListArgs{SkipHook: true}) - if err == nil { - recursivelyObjsUpdateHook(ctx, storage, dstPath, files, limiter) - } - } -} - -func needHandleObjsUpdateHook() bool { - if len(objsUpdateHooks) < 1 { - return false - } - needHandle, _ := GetSettingItemByKey(conf.HandleHookAfterWriting) - return needHandle != nil && (needHandle.Value == "true" || needHandle.Value == "1") -} - func wrapObjsName(storage driver.Driver, objs []model.Obj) { if _, ok := storage.(driver.Getter); !ok { model.WrapObjsName(objs) diff --git a/internal/op/hook.go b/internal/op/hook.go index 3d8530f933..13cb059653 100644 --- a/internal/op/hook.go +++ b/internal/op/hook.go @@ -6,6 +6,7 @@ import ( "regexp" "strconv" "strings" + "sync" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/driver" @@ -15,23 +16,34 @@ import ( log "github.com/sirupsen/logrus" ) -// Obj -type ObjsUpdateHook = func(ctx context.Context, parent string, objs []model.Obj) +type SnapshotProjector func(context.Context, string, []model.Obj) var ( - objsUpdateHooks = make([]ObjsUpdateHook, 0) + snapshotProjectorMu sync.RWMutex + snapshotProjector SnapshotProjector ) -func RegisterObjsUpdateHook(hook ObjsUpdateHook) { - objsUpdateHooks = append(objsUpdateHooks, hook) +func SetSnapshotProjector(projector SnapshotProjector) { + snapshotProjectorMu.Lock() + snapshotProjector = projector + snapshotProjectorMu.Unlock() } -func HandleObjsUpdateHook(ctx context.Context, parent string, objs []model.Obj) { - for _, hook := range objsUpdateHooks { - hook(ctx, parent, objs) +func ProjectSnapshot(ctx context.Context, parent string, objs []model.Obj) { + snapshotProjectorMu.RLock() + projector := snapshotProjector + snapshotProjectorMu.RUnlock() + if projector != nil { + projector(ctx, parent, append([]model.Obj(nil), objs...)) } } +func hasSnapshotProjector() bool { + snapshotProjectorMu.RLock() + defer snapshotProjectorMu.RUnlock() + return snapshotProjector != nil +} + // Setting type SettingItemHook func(item *model.SettingItem) error diff --git a/internal/op/list_identity_test.go b/internal/op/list_identity_test.go new file mode 100644 index 0000000000..2fb869a5be --- /dev/null +++ b/internal/op/list_identity_test.go @@ -0,0 +1,257 @@ +package op + +import ( + "context" + "fmt" + "sync/atomic" + "testing" + "time" + + "github.com/OpenListTeam/OpenList/v4/internal/driver" + "github.com/OpenListTeam/OpenList/v4/internal/errs" + "github.com/OpenListTeam/OpenList/v4/internal/model" +) + +type listIdentityDriver struct { + model.Storage + list func(model.ListArgs) []model.Obj +} + +func newListIdentityDriver(mount string, list func(model.ListArgs) []model.Obj) *listIdentityDriver { + return &listIdentityDriver{ + Storage: model.Storage{MountPath: mount, CacheExpiration: 10}, + list: list, + } +} + +func (d *listIdentityDriver) Config() driver.Config { return driver.Config{} } +func (d *listIdentityDriver) GetAddition() driver.Additional { + return nil +} +func (d *listIdentityDriver) Init(context.Context) error { return nil } +func (d *listIdentityDriver) Drop(context.Context) error { return nil } +func (d *listIdentityDriver) GetRoot(context.Context) (model.Obj, error) { + return &model.Object{Name: "root", Path: "/", IsFolder: true}, nil +} +func (d *listIdentityDriver) List(_ context.Context, _ model.Obj, args model.ListArgs) ([]model.Obj, error) { + return d.list(args), nil +} +func (d *listIdentityDriver) Link(context.Context, model.Obj, model.LinkArgs) (*model.Link, error) { + return nil, errs.NotImplement +} + +func listedName(t *testing.T, objs []model.Obj, err error) string { + t.Helper() + if err != nil { + t.Fatalf("list failed: %v", err) + } + if len(objs) != 1 { + t.Fatalf("expected one object, got %d", len(objs)) + } + return objs[0].GetName() +} + +func listName(t *testing.T, d driver.Driver, args model.ListArgs) string { + t.Helper() + objs, err := List(context.Background(), d, "/", args) + return listedName(t, objs, err) +} + +func TestDirectoryCacheOwnsSnapshotViews(t *testing.T) { + original := []model.Obj{&model.Object{Name: "before"}} + cached := newDirectoryCache(original) + original[0] = &model.Object{Name: "after"} + if got := cached.GetSortedObjects(&listIdentityDriver{}); got[0].GetName() != "before" { + t.Fatalf("cached snapshot changed through caller slice: %q", got[0].GetName()) + } +} + +func TestListNormalizesCanonicalRequestPath(t *testing.T) { + Cache = NewCacheManager() + var calls atomic.Int32 + d := newListIdentityDriver("/canonical", func(args model.ListArgs) []model.Obj { + calls.Add(1) + return []model.Obj{&model.Object{Name: args.ReqPath}} + }) + + if got := listName(t, d, model.ListArgs{}); got != "/canonical" { + t.Fatalf("canonical request path = %q, want /canonical", got) + } + if got := listName(t, d, model.ListArgs{}); got != "/canonical" { + t.Fatalf("cached canonical request path = %q, want /canonical", got) + } + if got := calls.Load(); got != 1 { + t.Fatalf("canonical driver calls = %d, want 1", got) + } +} + +func TestListDoesNotReuseSpecialVariants(t *testing.T) { + tests := []struct { + name string + mount string + first model.ListArgs + second model.ListArgs + label func(model.ListArgs) string + }{ + { + name: "request path", + mount: "/request-path", + first: model.ListArgs{ReqPath: "/alias-a"}, + second: model.ListArgs{ReqPath: "/alias-b"}, + label: func(args model.ListArgs) string { return args.ReqPath }, + }, + { + name: "placeholder visibility", + mount: "/placeholder", + first: model.ListArgs{ReqPath: "/placeholder"}, + second: model.ListArgs{ReqPath: "/placeholder", S3ShowPlaceholder: true}, + label: func(args model.ListArgs) string { return fmt.Sprintf("placeholder=%t", args.S3ShowPlaceholder) }, + }, + { + name: "storage details", + mount: "/details", + first: model.ListArgs{ReqPath: "/details"}, + second: model.ListArgs{ReqPath: "/details", WithStorageDetails: true}, + label: func(args model.ListArgs) string { return fmt.Sprintf("details=%t", args.WithStorageDetails) }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + Cache = NewCacheManager() + var calls atomic.Int32 + d := newListIdentityDriver(tt.mount, func(args model.ListArgs) []model.Obj { + calls.Add(1) + return []model.Obj{&model.Object{Name: tt.label(args)}} + }) + + first := listName(t, d, tt.first) + second := listName(t, d, tt.second) + if first == second { + t.Fatalf("different list variants shared %q", first) + } + if got := calls.Load(); got != 2 { + t.Fatalf("driver calls = %d, want 2", got) + } + }) + } +} + +func TestRefreshDoesNotJoinOrLoseToOrdinaryLoad(t *testing.T) { + Cache = NewCacheManager() + normalStarted := make(chan struct{}) + refreshStarted := make(chan struct{}) + releaseNormal := make(chan struct{}) + d := newListIdentityDriver("/refresh", func(args model.ListArgs) []model.Obj { + if args.Refresh { + close(refreshStarted) + return []model.Obj{&model.Object{Name: "fresh"}} + } + close(normalStarted) + <-releaseNormal + return []model.Obj{&model.Object{Name: "ordinary"}} + }) + + type result struct { + objs []model.Obj + err error + } + normalResult := make(chan result, 1) + go func() { + objs, err := List(context.Background(), d, "/", model.ListArgs{}) + normalResult <- result{objs: objs, err: err} + }() + <-normalStarted + + refreshResult := make(chan result, 1) + go func() { + objs, err := List(context.Background(), d, "/", model.ListArgs{Refresh: true}) + refreshResult <- result{objs: objs, err: err} + }() + + select { + case <-refreshStarted: + case <-time.After(100 * time.Millisecond): + close(releaseNormal) + <-normalResult + <-refreshResult + t.Fatal("refresh joined an ordinary in-flight load") + } + refresh := <-refreshResult + if got := listedName(t, refresh.objs, refresh.err); got != "fresh" { + t.Fatalf("refresh result = %q, want fresh", got) + } + close(releaseNormal) + normal := <-normalResult + if got := listedName(t, normal.objs, normal.err); got != "ordinary" { + t.Fatalf("ordinary result = %q, want ordinary", got) + } + + if got := listName(t, d, model.ListArgs{}); got != "fresh" { + t.Fatalf("cached result after refresh = %q, want fresh", got) + } +} + +func TestMutationInvalidatesInFlightListCommit(t *testing.T) { + Cache = NewCacheManager() + started := make(chan struct{}) + release := make(chan struct{}) + var calls atomic.Int32 + d := newListIdentityDriver("/mutation", func(model.ListArgs) []model.Obj { + if calls.Add(1) == 1 { + close(started) + <-release + return []model.Obj{&model.Object{Name: "stale"}} + } + return []model.Obj{&model.Object{Name: "fresh"}} + }) + + done := make(chan struct{}) + go func() { + _, _ = List(context.Background(), d, "/", model.ListArgs{}) + close(done) + }() + <-started + Cache.mutateDirectories([]string{Key(d, "/")}, func() {}) + close(release) + <-done + + if got := listName(t, d, model.ListArgs{}); got != "fresh" { + t.Fatalf("cached stale in-flight result: got %q", got) + } +} + +type compositeListDriver struct { + *listIdentityDriver + backing driver.Driver +} + +func (d *compositeListDriver) List(ctx context.Context, _ model.Obj, _ model.ListArgs) ([]model.Obj, error) { + return List(ctx, d.backing, "/", model.ListArgs{Refresh: true}) +} + +func TestNestedListProjectsOnlyPublicSnapshot(t *testing.T) { + Cache = NewCacheManager() + backing := newListIdentityDriver("/backing", func(model.ListArgs) []model.Obj { + return []model.Obj{&model.Object{Name: "file"}} + }) + outer := &compositeListDriver{ + listIdentityDriver: newListIdentityDriver("/public", nil), + backing: backing, + } + var projected []string + SetSnapshotProjector(func(_ context.Context, parent string, _ []model.Obj) { + projected = append(projected, parent) + }) + t.Cleanup(func() { SetSnapshotProjector(nil) }) + + if got := listName(t, outer, model.ListArgs{Refresh: true}); got != "file" { + t.Fatalf("outer result = %q, want file", got) + } + if len(projected) != 1 || projected[0] != "/public" { + t.Fatalf("projected parents = %v, want [/public]", projected) + } +} + +var _ driver.Driver = (*listIdentityDriver)(nil) +var _ driver.GetRooter = (*listIdentityDriver)(nil) diff --git a/internal/op/reconcile.go b/internal/op/reconcile.go new file mode 100644 index 0000000000..724728d4db --- /dev/null +++ b/internal/op/reconcile.go @@ -0,0 +1,127 @@ +package op + +import ( + "context" + "strconv" + "sync/atomic" + + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/driver" + "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/OpenListTeam/OpenList/v4/pkg/mq" + "github.com/OpenListTeam/OpenList/v4/pkg/utils" + log "github.com/sirupsen/logrus" + "golang.org/x/time/rate" +) + +type reconciliationTarget struct { + storage driver.Driver + path string + recursive bool +} + +type snapshotFrameKey struct{} +type mutationFrameKey struct{} + +type mutationFrame struct { + dirty atomic.Bool + storage driver.Driver + targets []reconciliationTarget +} + +var reconciliation atomic.Pointer[mq.LatestProcessor[string, reconciliationTarget]] + +func StartSnapshotReconciliation() { + if reconciliation.Load() != nil { + return + } + reconciliation.Store(mq.NewLatestProcessor( + 256, + 256, + 2, + func(ctx context.Context, _ string, request reconciliationTarget) { + reconcileSnapshot(ctx, request.storage, request.path, request.recursive) + }, + )) +} + +func StopSnapshotReconciliation(ctx context.Context) error { + processor := reconciliation.Swap(nil) + if processor == nil { + return nil + } + return processor.Stop(ctx) +} + +func ScheduleSnapshotReconciliation(storage driver.Driver, path string, recursive bool) { + if !hasSnapshotProjector() { + return + } + setting, _ := GetSettingItemByKey(conf.HandleHookAfterWriting) + if setting == nil || (setting.Value != "true" && setting.Value != "1") { + return + } + path = utils.FixAndCleanPath(path) + key := Key(storage, path) + if recursive { + key += "\x00recursive" + } + processor := reconciliation.Load() + if processor != nil { + result := processor.Offer(key, reconciliationTarget{storage: storage, path: path, recursive: recursive}, 1) + if result == mq.OfferRejectedCapacity { + if count := processor.Stats().Rejected; count == 1 || count%100 == 0 { + log.Warnf("snapshot reconciliation capacity reached for %s (rejected=%d)", key, count) + } + } + } +} + +func enterMutationFrame(ctx context.Context, storage driver.Driver, targets ...reconciliationTarget) (context.Context, func()) { + if ctx.Value(mutationFrameKey{}) != nil { + return ctx, func() {} + } + frame := &mutationFrame{storage: storage, targets: targets} + framed := context.WithValue(ctx, mutationFrameKey{}, frame) + return framed, func() { + if frame.dirty.Load() { + for _, target := range frame.targets { + ScheduleSnapshotReconciliation(frame.storage, target.path, target.recursive) + } + } + } +} + +func enterSnapshotFrame(ctx context.Context) (context.Context, bool) { + if ctx.Value(snapshotFrameKey{}) != nil || ctx.Value(mutationFrameKey{}) != nil { + return ctx, false + } + return context.WithValue(ctx, snapshotFrameKey{}, struct{}{}), true +} + +func commitNamespaceMutation(ctx context.Context, storage driver.Driver, cacheKeys []string, mutateCache func(), targets ...reconciliationTarget) { + Cache.mutateDirectories(cacheKeys, mutateCache) + if frame, _ := ctx.Value(mutationFrameKey{}).(*mutationFrame); frame != nil { + frame.dirty.Store(true) + return + } + for _, target := range targets { + ScheduleSnapshotReconciliation(storage, target.path, target.recursive) + } +} + +func reconcileSnapshot(ctx context.Context, storage driver.Driver, path string, recursive bool) { + if recursive { + var limiter *rate.Limiter + if item, _ := GetSettingItemByKey(conf.HandleHookRateLimit); item != nil { + if limit, err := strconv.ParseFloat(item.Value, 64); err == nil && limit > 0 { + limiter = rate.NewLimiter(rate.Limit(limit), 1) + } + } + RecursivelyListStorage(ctx, storage, path, limiter, nil) + return + } + if _, err := List(ctx, storage, path, model.ListArgs{Refresh: true}); err != nil && ctx.Err() == nil { + log.Warnf("reconcile snapshot %s: %v", Key(storage, path), err) + } +} diff --git a/internal/op/recursive_list.go b/internal/op/recursive_list.go index de8a27b2b1..ac5403144e 100644 --- a/internal/op/recursive_list.go +++ b/internal/op/recursive_list.go @@ -3,7 +3,6 @@ package op import ( "context" stdpath "path" - "sync" "sync/atomic" "github.com/OpenListTeam/OpenList/v4/internal/driver" @@ -60,14 +59,12 @@ func RecursivelyList(ctx context.Context, rawPath string, limit rate.Limit, coun } RecursivelyListStorage(ctx, storage, actualPath, limiter, counter) } else { - var wg sync.WaitGroup - recursivelyListVirtual(ctx, rawPath, limit, counter, &wg) - wg.Wait() + recursivelyListVirtual(ctx, rawPath, limit, counter) } return nil } -func recursivelyListVirtual(ctx context.Context, rawPath string, limit rate.Limit, counter *atomic.Uint64, wg *sync.WaitGroup) { +func recursivelyListVirtual(ctx context.Context, rawPath string, limit rate.Limit, counter *atomic.Uint64) { objs := GetStorageVirtualFilesByPath(rawPath) if counter != nil { counter.Add(uint64(len(objs))) @@ -85,13 +82,9 @@ func recursivelyListVirtual(ctx context.Context, rawPath string, limit rate.Limi if limit > .0 { limiter = rate.NewLimiter(limit, 1) } - wg.Add(1) - go func() { - defer wg.Done() - RecursivelyListStorage(ctx, storage, actualPath, limiter, counter) - }() + RecursivelyListStorage(ctx, storage, actualPath, limiter, counter) } else { - recursivelyListVirtual(ctx, nextPath, limit, counter, wg) + recursivelyListVirtual(ctx, nextPath, limit, counter) } } } diff --git a/internal/search/build.go b/internal/search/build.go index 8dbe7d498e..3385ca4aef 100644 --- a/internal/search/build.go +++ b/internal/search/build.go @@ -16,9 +16,6 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/op" "github.com/OpenListTeam/OpenList/v4/internal/search/searcher" "github.com/OpenListTeam/OpenList/v4/internal/setting" - "github.com/OpenListTeam/OpenList/v4/pkg/mq" - "github.com/OpenListTeam/OpenList/v4/pkg/utils" - mapset "github.com/deckarep/golang-set/v2" log "github.com/sirupsen/logrus" ) @@ -26,6 +23,31 @@ var ( Quit = atomic.Pointer[chan struct{}]{} ) +type indexBatch struct { + sync.Mutex + objs []ObjWithParent +} + +func (batch *indexBatch) add(obj ObjWithParent) { + batch.Lock() + batch.objs = append(batch.objs, obj) + batch.Unlock() +} + +func (batch *indexBatch) take() []ObjWithParent { + batch.Lock() + objs := batch.objs + batch.objs = nil + batch.Unlock() + return objs +} + +func (batch *indexBatch) len() int { + batch.Lock() + defer batch.Unlock() + return len(batch.objs) +} + func Running() bool { return Quit.Load() != nil } @@ -44,7 +66,7 @@ func BuildIndex(ctx context.Context, indexPaths, ignorePaths []string, maxDepth return errs.BuildIndexIsRunning } var ( - indexMQ = mq.NewInMemoryMQ[ObjWithParent]() + indexMQ = &indexBatch{} running = atomic.Bool{} // current goroutine running wg = &sync.WaitGroup{} ) @@ -64,20 +86,17 @@ func BuildIndex(ctx context.Context, indexPaths, ignorePaths []string, maxDepth select { case <-ticker.C: tickCount += 1 - if indexMQ.Len() < 1000 && tickCount != 5 { + if indexMQ.len() < 1000 && tickCount != 5 { continue } else if tickCount >= 5 { tickCount = 0 } log.Infof("index obj count: %d", objCount) - indexMQ.ConsumeAll(func(messages []mq.Message[ObjWithParent]) { + func(messages []ObjWithParent) { if len(messages) != 0 { - log.Debugf("current index: %s", messages[len(messages)-1].Content.Parent) + log.Debugf("current index: %s", messages[len(messages)-1].Parent) } - if err = BatchIndex(ctx, utils.MustSliceConvert(messages, - func(src mq.Message[ObjWithParent]) ObjWithParent { - return src.Content - })); err != nil { + if err = BatchIndex(ctx, messages); err != nil { log.Errorf("build index in batch error: %+v", err) } else { objCount = objCount + uint64(len(messages)) @@ -89,18 +108,15 @@ func BuildIndex(ctx context.Context, indexPaths, ignorePaths []string, maxDepth LastDoneTime: nil, }) } - }) + }(indexMQ.take()) case <-quit: log.Debugf("build index for %+v received quit", indexPaths) eMsg := "" now := time.Now() originErr := err - indexMQ.ConsumeAll(func(messages []mq.Message[ObjWithParent]) { - if err = BatchIndex(ctx, utils.MustSliceConvert(messages, - func(src mq.Message[ObjWithParent]) ObjWithParent { - return src.Content - })); err != nil { + func(messages []ObjWithParent) { + if err = BatchIndex(ctx, messages); err != nil { log.Errorf("build index in batch error: %+v", err) } else { objCount = objCount + uint64(len(messages)) @@ -119,7 +135,7 @@ func BuildIndex(ctx context.Context, indexPaths, ignorePaths []string, maxDepth Error: eMsg, }) } - }) + }(indexMQ.take()) log.Debugf("build index for %+v quit success", indexPaths) return } @@ -166,12 +182,7 @@ func BuildIndex(ctx context.Context, indexPaths, ignorePaths []string, maxDepth if indexPath == "/" { return nil } - indexMQ.Publish(mq.Message[ObjWithParent]{ - Content: ObjWithParent{ - Obj: info, - Parent: path.Dir(indexPath), - }, - }) + indexMQ.add(ObjWithParent{Obj: info, Parent: path.Dir(indexPath)}) return nil } fi, err = fs.Get(ctx, indexPath, &fs.GetArgs{}) @@ -199,7 +210,11 @@ func Config(ctx context.Context) searcher.Config { return instance.Config() } -func Update(ctx context.Context, parent string, objs []model.Obj) { +type snapshotUpdater interface { + UpdateSnapshot(context.Context, string, []model.Obj) error +} + +func UpdateSnapshot(ctx context.Context, parent string, objs []model.Obj) { if instance == nil || !instance.Config().AutoUpdate || !setting.GetBool(conf.AutoUpdateIndex) || Running() { return } @@ -216,38 +231,23 @@ func Update(ctx context.Context, parent string, objs []model.Obj) { return } - // Use task queue for Meilisearch to avoid race conditions with async indexing - if msInstance, ok := instance.(interface { - EnqueueUpdate(parent string, objs []model.Obj) - }); ok { - // Enqueue task for async processing (diff calculation happens at consumption time) - msInstance.EnqueueUpdate(parent, objs) + if updater, ok := instance.(snapshotUpdater); ok { + if err := updater.UpdateSnapshot(ctx, parent, objs); err != nil { + log.Errorf("update search index error for %s: %+v", parent, err) + } return } - unlock := lockUpdate(parent) - defer unlock() - nodes, err := instance.Get(ctx, parent) if err != nil { log.Errorf("update search index error while get nodes: %+v", err) return } - now := mapset.NewSet[string]() - for i := range objs { - now.Add(objs[i].GetName()) - } - old := mapset.NewSet[string]() - for i := range nodes { - old.Add(nodes[i].Name) - } - // delete data that no longer exists - toDelete := old.Difference(now) - toAdd := now.Difference(old) - for i := range nodes { - if toDelete.Contains(nodes[i].Name) && !op.HasStorage(path.Join(parent, nodes[i].Name)) { - log.Debugf("delete index: %s", path.Join(parent, nodes[i].Name)) - err = instance.Del(ctx, path.Join(parent, nodes[i].Name)) + removed, added := searcher.SnapshotDiff(objs, nodes) + for _, node := range removed { + if !op.HasStorage(path.Join(parent, node.Name)) { + log.Debugf("delete index: %s", path.Join(parent, node.Name)) + err = instance.Del(ctx, path.Join(parent, node.Name)) if err != nil { log.Errorf("update search index error while del old node: %+v", err) return @@ -255,15 +255,10 @@ func Update(ctx context.Context, parent string, objs []model.Obj) { } } // collect files and folders to add in batch - var toAddObjs []ObjWithParent - for i := range objs { - if toAdd.Contains(objs[i].GetName()) { - log.Debugf("add index: %s", path.Join(parent, objs[i].GetName())) - toAddObjs = append(toAddObjs, ObjWithParent{ - Parent: parent, - Obj: objs[i], - }) - } + toAddObjs := make([]ObjWithParent, 0, len(added)) + for _, obj := range added { + log.Debugf("add index: %s", path.Join(parent, obj.GetName())) + toAddObjs = append(toAddObjs, ObjWithParent{Parent: parent, Obj: obj}) } // batch index all files and folders at once if len(toAddObjs) > 0 { @@ -274,7 +269,3 @@ func Update(ctx context.Context, parent string, objs []model.Obj) { } } } - -func init() { - op.RegisterObjsUpdateHook(Update) -} diff --git a/internal/search/build_test.go b/internal/search/build_test.go deleted file mode 100644 index fed2e155e5..0000000000 --- a/internal/search/build_test.go +++ /dev/null @@ -1,58 +0,0 @@ -package search - -import ( - "testing" - "time" -) - -func TestLockUpdateSerializesSameParent(t *testing.T) { - unlockFirst := lockUpdate("/same-parent") - secondStarted := make(chan struct{}) - secondAcquired := make(chan struct{}) - secondReleased := make(chan struct{}) - go func() { - close(secondStarted) - unlockSecond := lockUpdate("/same-parent") - close(secondAcquired) - unlockSecond() - close(secondReleased) - }() - <-secondStarted - - select { - case <-secondAcquired: - t.Fatal("second update acquired the same parent lock") - case <-time.After(20 * time.Millisecond): - } - - unlockFirst() - select { - case <-secondReleased: - case <-time.After(time.Second): - t.Fatal("second update did not acquire the released parent lock") - } - - updateLocksMu.Lock() - defer updateLocksMu.Unlock() - if len(updateLocks) != 0 { - t.Fatalf("update locks were not cleaned up: %d", len(updateLocks)) - } -} - -func TestLockUpdateAllowsDifferentParents(t *testing.T) { - unlockFirst := lockUpdate("/first-parent") - defer unlockFirst() - - secondAcquired := make(chan struct{}) - go func() { - unlockSecond := lockUpdate("/second-parent") - unlockSecond() - close(secondAcquired) - }() - - select { - case <-secondAcquired: - case <-time.After(time.Second): - t.Fatal("update for a different parent was blocked") - } -} diff --git a/internal/search/meilisearch/init.go b/internal/search/meilisearch/init.go index 3c379b7907..d52e3fc675 100644 --- a/internal/search/meilisearch/init.go +++ b/internal/search/meilisearch/init.go @@ -92,10 +92,6 @@ func init() { } } - // Initialize and start task queue manager - m.taskQueue = NewTaskQueueManager(&m) - m.taskQueue.Start() - return &m, nil }) } diff --git a/internal/search/meilisearch/search.go b/internal/search/meilisearch/search.go index 75283228c3..5795d16254 100644 --- a/internal/search/meilisearch/search.go +++ b/internal/search/meilisearch/search.go @@ -33,7 +33,6 @@ type Meilisearch struct { IndexUid string FilterableAttributes []string SearchableAttributes []string - taskQueue *TaskQueueManager } func (m *Meilisearch) Config() searcher.Config { @@ -83,48 +82,8 @@ func (m *Meilisearch) Index(ctx context.Context, node model.SearchNode) error { } func (m *Meilisearch) BatchIndex(ctx context.Context, nodes []model.SearchNode) error { - documents, err := utils.SliceConvert(nodes, func(src model.SearchNode) (*searchDocument, error) { - parentHash := hashPath(src.Parent) - nodePath := path.Join(src.Parent, src.Name) - nodePathHash := hashPath(nodePath) - parentPaths := utils.GetPathHierarchy(src.Parent) - parentPathHashes, err := utils.SliceConvert(parentPaths, func(parentPath string) (string, error) { - return hashPath(parentPath), nil - }) - if err != nil { - return nil, err - } - - return &searchDocument{ - ID: nodePathHash, - ParentHash: parentHash, - ParentPathHashes: parentPathHashes, - SearchNode: src, - }, nil - }) - if err != nil { - return err - } - - // max up to 10,000 documents per batch to reduce error rate while uploading over the Internet - _, err = m.Client.Index(m.IndexUid).AddDocumentsInBatchesWithContext(ctx, documents, 10000) - if err != nil { - return err - } - - // documents were uploaded and enqueued for indexing, just return early - //// Wait for the task to complete and check - //forTask, err := m.Client.WaitForTask(task.TaskUID, meilisearch.WaitParams{ - // Context: ctx, - // Interval: time.Millisecond * 50, - //}) - //if err != nil { - // return err - //} - //if forTask.Status != meilisearch.TaskStatusSucceeded { - // return fmt.Errorf("BatchIndex failed, task status is %s", forTask.Status) - //} - return nil + _, err := m.batchIndexWithTaskUID(ctx, nodes) + return err } func (m *Meilisearch) getDocumentsByParent(ctx context.Context, parent string) ([]*searchDocument, error) { @@ -210,9 +169,6 @@ func (m *Meilisearch) Del(ctx context.Context, prefix string) error { } func (m *Meilisearch) Release(ctx context.Context) error { - if m.taskQueue != nil { - m.taskQueue.Stop() - } return nil } @@ -230,15 +186,6 @@ func (m *Meilisearch) getTaskStatus(ctx context.Context, taskUID int64) (meilise return forTask.Status, nil } -// EnqueueUpdate enqueues an update task to the task queue -func (m *Meilisearch) EnqueueUpdate(parent string, objs []model.Obj) { - if m.taskQueue == nil { - return - } - - m.taskQueue.Enqueue(parent, objs) -} - // batchIndexWithTaskUID indexes documents and returns all taskUIDs func (m *Meilisearch) batchIndexWithTaskUID(ctx context.Context, nodes []model.SearchNode) ([]int64, error) { if len(nodes) == 0 { diff --git a/internal/search/meilisearch/task_queue.go b/internal/search/meilisearch/task_queue.go deleted file mode 100644 index c5384e6e5a..0000000000 --- a/internal/search/meilisearch/task_queue.go +++ /dev/null @@ -1,265 +0,0 @@ -package meilisearch - -import ( - "context" - "path" - "sort" - "strings" - "sync" - "sync/atomic" - "time" - - "github.com/OpenListTeam/OpenList/v4/internal/model" - "github.com/OpenListTeam/OpenList/v4/internal/op" - mapset "github.com/deckarep/golang-set/v2" - log "github.com/sirupsen/logrus" -) - -// QueuedTask represents a task in the queue -type QueuedTask struct { - Parent string - Objs []model.Obj // current file system state - Depth int // path depth for sorting - EnqueueAt time.Time // enqueue time -} - -// TaskQueueManager manages the task queue for async index operations -type TaskQueueManager struct { - queue map[string]*QueuedTask // parent -> task - pendingTasks map[string][]int64 // parent -> all submitted taskUIDs - mu sync.RWMutex - ticker *time.Ticker - stopCh chan struct{} - m *Meilisearch - consuming atomic.Bool // flag to prevent concurrent consumption -} - -// NewTaskQueueManager creates a new task queue manager -func NewTaskQueueManager(m *Meilisearch) *TaskQueueManager { - return &TaskQueueManager{ - queue: make(map[string]*QueuedTask), - pendingTasks: make(map[string][]int64), - stopCh: make(chan struct{}), - m: m, - } -} - -// calculateDepth calculates the depth of a path -func calculateDepth(path string) int { - if path == "/" { - return 0 - } - return strings.Count(strings.Trim(path, "/"), "/") + 1 -} - -// Enqueue enqueues a task with current file system state -func (tqm *TaskQueueManager) Enqueue(parent string, objs []model.Obj) { - tqm.mu.Lock() - defer tqm.mu.Unlock() - - // deduplicate: overwrite existing task with the same parent - tqm.queue[parent] = &QueuedTask{ - Parent: parent, - Objs: objs, - Depth: calculateDepth(parent), - EnqueueAt: time.Now(), - } - log.Debugf("enqueued update task for parent: %s, depth: %d, objs: %d", parent, calculateDepth(parent), len(objs)) -} - -// Start starts the task queue consumer -func (tqm *TaskQueueManager) Start() { - tqm.ticker = time.NewTicker(30 * time.Second) - go func() { - for { - select { - case <-tqm.ticker.C: - tqm.consume() - case <-tqm.stopCh: - log.Info("task queue manager stopped") - return - } - } - }() - log.Info("task queue manager started, will consume every 30 seconds") -} - -// Stop stops the task queue consumer -func (tqm *TaskQueueManager) Stop() { - if tqm.ticker != nil { - tqm.ticker.Stop() - } - close(tqm.stopCh) -} - -// consume processes all tasks in the queue -func (tqm *TaskQueueManager) consume() { - // Prevent concurrent consumption - if !tqm.consuming.CompareAndSwap(false, true) { - log.Warn("previous consume still running, skip this round") - return - } - defer tqm.consuming.Store(false) - - tqm.mu.Lock() - - // extract all tasks - tasks := make([]*QueuedTask, 0, len(tqm.queue)) - for _, task := range tqm.queue { - tasks = append(tasks, task) - } - - // clear queue - tqm.queue = make(map[string]*QueuedTask) - - tqm.mu.Unlock() - - if len(tasks) == 0 { - return - } - - log.Infof("consuming task queue: %d tasks", len(tasks)) - - // sort tasks: shallow paths first, then by enqueue time - sort.Slice(tasks, func(i, j int) bool { - if tasks[i].Depth != tasks[j].Depth { - return tasks[i].Depth < tasks[j].Depth - } - return tasks[i].EnqueueAt.Before(tasks[j].EnqueueAt) - }) - - ctx := context.Background() - - // execute tasks in order - for _, task := range tasks { - // Check if there are pending tasks for this parent - tqm.mu.RLock() - pendingTaskUIDs, hasPending := tqm.pendingTasks[task.Parent] - tqm.mu.RUnlock() - - if hasPending && len(pendingTaskUIDs) > 0 { - // Check all pending task statuses - allCompleted := true - for _, taskUID := range pendingTaskUIDs { - taskStatus, err := tqm.m.getTaskStatus(ctx, taskUID) - if err != nil { - log.Errorf("failed to get task status for parent %s (taskUID: %d): %v", task.Parent, taskUID, err) - // If we can't get status, assume it's done and continue checking - continue - } - - // Check if task is still running - if taskStatus == "enqueued" || taskStatus == "processing" { - log.Warnf("skipping task for parent %s: previous task %d still %s", task.Parent, taskUID, taskStatus) - allCompleted = false - break // No need to check remaining tasks - } - } - - if !allCompleted { - // Re-enqueue the task if not already in queue (avoid overwriting newer snapshots) - tqm.mu.Lock() - if _, exists := tqm.queue[task.Parent]; !exists { - tqm.queue[task.Parent] = task - log.Debugf("re-enqueued skipped task for parent %s due to pending tasks", task.Parent) - } else { - log.Debugf("skipped task for parent %s not re-enqueued (newer task already in queue)", task.Parent) - } - tqm.mu.Unlock() - continue // Skip this task, some previous tasks are still running - } - - // All tasks are in terminal state, remove from pending - log.Debugf("all previous tasks for parent %s are completed, proceeding with new task", task.Parent) - tqm.mu.Lock() - delete(tqm.pendingTasks, task.Parent) - tqm.mu.Unlock() - } - - // Execute the task - tqm.executeTask(ctx, task) - } - - log.Infof("task queue consumption completed") -} - -// executeTask executes a single task -func (tqm *TaskQueueManager) executeTask(ctx context.Context, task *QueuedTask) { - parent := task.Parent - currentObjs := task.Objs - - // Query index to get old state - nodes, err := tqm.m.Get(ctx, parent) - if err != nil { - log.Errorf("failed to get indexed nodes for parent %s: %v", parent, err) - return - } - - // Calculate diff based on current index state - now := mapset.NewSet[string]() - for i := range currentObjs { - now.Add(currentObjs[i].GetName()) - } - old := mapset.NewSet[string]() - for i := range nodes { - old.Add(nodes[i].Name) - } - - toDelete := old.Difference(now) - toAdd := now.Difference(old) - - // Collect paths to delete - var pathsToDelete []string - for i := range nodes { - if toDelete.Contains(nodes[i].Name) && !op.HasStorage(path.Join(parent, nodes[i].Name)) { - pathsToDelete = append(pathsToDelete, path.Join(parent, nodes[i].Name)) - } - } - - var allTaskUIDs []int64 - - // Execute delete first - if len(pathsToDelete) > 0 { - log.Debugf("executing delete for parent %s: %d paths", parent, len(pathsToDelete)) - taskUIDs, err := tqm.m.batchDeleteWithTaskUID(ctx, pathsToDelete) - if err != nil { - log.Errorf("failed to batch delete for parent %s: %v", parent, err) - // Continue to add even if delete fails - } else { - allTaskUIDs = append(allTaskUIDs, taskUIDs...) - } - } - - // Collect objects to add - var nodesToAdd []model.SearchNode - for i := range currentObjs { - if toAdd.Contains(currentObjs[i].GetName()) { - log.Debugf("will add index: %s", path.Join(parent, currentObjs[i].GetName())) - nodesToAdd = append(nodesToAdd, model.SearchNode{ - Parent: parent, - Name: currentObjs[i].GetName(), - IsDir: currentObjs[i].IsDir(), - Size: currentObjs[i].GetSize(), - }) - } - } - - // Execute add - if len(nodesToAdd) > 0 { - log.Debugf("executing add for parent %s: %d nodes", parent, len(nodesToAdd)) - taskUIDs, err := tqm.m.batchIndexWithTaskUID(ctx, nodesToAdd) - if err != nil { - log.Errorf("failed to batch index for parent %s: %v", parent, err) - } else { - allTaskUIDs = append(allTaskUIDs, taskUIDs...) - } - } - - // Record all task UIDs for this parent - if len(allTaskUIDs) > 0 { - tqm.mu.Lock() - tqm.pendingTasks[parent] = allTaskUIDs - tqm.mu.Unlock() - log.Debugf("recorded %d taskUIDs for parent %s", len(allTaskUIDs), parent) - } -} diff --git a/internal/search/meilisearch/update.go b/internal/search/meilisearch/update.go new file mode 100644 index 0000000000..359748f9f1 --- /dev/null +++ b/internal/search/meilisearch/update.go @@ -0,0 +1,63 @@ +package meilisearch + +import ( + "context" + "errors" + "fmt" + "path" + + "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/OpenListTeam/OpenList/v4/internal/op" + "github.com/OpenListTeam/OpenList/v4/internal/search/searcher" + "github.com/meilisearch/meilisearch-go" +) + +func (m *Meilisearch) UpdateSnapshot(ctx context.Context, parent string, objs []model.Obj) error { + nodes, err := m.Get(ctx, parent) + if err != nil { + return fmt.Errorf("get indexed nodes: %w", err) + } + + removed, added := searcher.SnapshotDiff(objs, nodes) + pathsToDelete := make([]string, 0, len(removed)) + for _, node := range removed { + objPath := path.Join(parent, node.Name) + if !op.HasStorage(objPath) { + pathsToDelete = append(pathsToDelete, objPath) + } + } + + var ( + taskUIDs []int64 + updateErrs []error + ) + deleteUIDs, err := m.batchDeleteWithTaskUID(ctx, pathsToDelete) + if err != nil { + updateErrs = append(updateErrs, fmt.Errorf("delete stale nodes: %w", err)) + } else { + taskUIDs = append(taskUIDs, deleteUIDs...) + } + + nodesToAdd := make([]model.SearchNode, 0, len(added)) + for _, obj := range added { + nodesToAdd = append(nodesToAdd, model.SearchNode{Parent: parent, Name: obj.GetName(), IsDir: obj.IsDir(), Size: obj.GetSize()}) + } + addUIDs, err := m.batchIndexWithTaskUID(ctx, nodesToAdd) + if err != nil { + updateErrs = append(updateErrs, fmt.Errorf("index new nodes: %w", err)) + } else { + taskUIDs = append(taskUIDs, addUIDs...) + } + + for _, taskUID := range taskUIDs { + status, err := m.getTaskStatus(ctx, taskUID) + if err != nil { + updateErrs = append(updateErrs, fmt.Errorf("wait for task %d: %w", taskUID, err)) + continue + } + if status != meilisearch.TaskStatusSucceeded { + updateErrs = append(updateErrs, fmt.Errorf("task %d completed with status %s", taskUID, status)) + } + } + return errors.Join(updateErrs...) +} diff --git a/internal/search/searcher/snapshot.go b/internal/search/searcher/snapshot.go new file mode 100644 index 0000000000..9f85864b16 --- /dev/null +++ b/internal/search/searcher/snapshot.go @@ -0,0 +1,23 @@ +package searcher + +import "github.com/OpenListTeam/OpenList/v4/internal/model" + +func SnapshotDiff(objs []model.Obj, nodes []model.SearchNode) (removed []model.SearchNode, added []model.Obj) { + current := make(map[string]struct{}, len(objs)) + for _, obj := range objs { + current[obj.GetName()] = struct{}{} + } + indexed := make(map[string]struct{}, len(nodes)) + for _, node := range nodes { + indexed[node.Name] = struct{}{} + if _, ok := current[node.Name]; !ok { + removed = append(removed, node) + } + } + for _, obj := range objs { + if _, ok := indexed[obj.GetName()]; !ok { + added = append(added, obj) + } + } + return +} diff --git a/internal/search/update_lock.go b/internal/search/update_lock.go deleted file mode 100644 index 7ca8d28ddb..0000000000 --- a/internal/search/update_lock.go +++ /dev/null @@ -1,38 +0,0 @@ -package search - -import "sync" - -var ( - updateLocksMu sync.Mutex - updateLocks = make(map[string]*updateLock) -) - -type updateLock struct { - mu sync.Mutex - refs uint -} - -// lockUpdate serializes index updates for the same parent while allowing -// unrelated directories to update concurrently. -func lockUpdate(parent string) func() { - updateLocksMu.Lock() - lock, ok := updateLocks[parent] - if !ok { - lock = &updateLock{} - updateLocks[parent] = lock - } - lock.refs++ - updateLocksMu.Unlock() - - lock.mu.Lock() - return func() { - lock.mu.Unlock() - - updateLocksMu.Lock() - lock.refs-- - if lock.refs == 0 { - delete(updateLocks, parent) - } - updateLocksMu.Unlock() - } -} diff --git a/internal/task_group/group.go b/internal/task_group/group.go index edd51fe746..68aef4f034 100644 --- a/internal/task_group/group.go +++ b/internal/task_group/group.go @@ -3,80 +3,54 @@ package task_group import ( "context" "sync" - - "github.com/sirupsen/logrus" ) -type OnCompletionFunc func(ctx context.Context, groupID string, payloads ...any) -type TaskGroupCoordinator struct { - name string - mu sync.Mutex - - groupPayloads map[string][]any - groupStates map[string]groupState - onCompletion OnCompletionFunc +type transferGroup struct { + pending int + hasSuccess bool + removeSource map[string]struct{} } -type groupState struct { - pending int - hasSuccess bool +type TransferGroupCoordinator struct { + mu sync.Mutex + groups map[string]*transferGroup } -func NewTaskGroupCoordinator(name string, f OnCompletionFunc) *TaskGroupCoordinator { - return &TaskGroupCoordinator{ - name: name, - groupPayloads: map[string][]any{}, - groupStates: map[string]groupState{}, - onCompletion: f, +func (coordinator *TransferGroupCoordinator) AddTask(groupID string) { + coordinator.mu.Lock() + group := coordinator.groups[groupID] + if group == nil { + group = &transferGroup{removeSource: make(map[string]struct{})} + coordinator.groups[groupID] = group } + group.pending++ + coordinator.mu.Unlock() } -// payload可为nil -func (tgc *TaskGroupCoordinator) AddTask(groupID string, payload any) { - tgc.mu.Lock() - defer tgc.mu.Unlock() - state := tgc.groupStates[groupID] - state.pending++ - tgc.groupStates[groupID] = state - logrus.Debugf("AddTask:%s ,count=%+v", groupID, state) - if payload == nil { - return +func (coordinator *TransferGroupCoordinator) RemoveSource(groupID, path string) { + coordinator.mu.Lock() + if group := coordinator.groups[groupID]; group != nil { + group.removeSource[path] = struct{}{} } - tgc.groupPayloads[groupID] = append(tgc.groupPayloads[groupID], payload) + coordinator.mu.Unlock() } -func (tgc *TaskGroupCoordinator) AppendPayload(groupID string, payload any) { - if payload == nil { +func (coordinator *TransferGroupCoordinator) Done(ctx context.Context, groupID string, success bool) { + coordinator.mu.Lock() + group := coordinator.groups[groupID] + if group == nil || group.pending == 0 { + coordinator.mu.Unlock() return } - tgc.mu.Lock() - defer tgc.mu.Unlock() - tgc.groupPayloads[groupID] = append(tgc.groupPayloads[groupID], payload) -} - -func (tgc *TaskGroupCoordinator) Done(ctx context.Context, groupID string, success bool) { - tgc.mu.Lock() - defer tgc.mu.Unlock() - state, ok := tgc.groupStates[groupID] - if !ok || state.pending == 0 { + group.hasSuccess = group.hasSuccess || success + group.pending-- + if group.pending != 0 { + coordinator.mu.Unlock() return } - if success { - state.hasSuccess = true - } - logrus.Debugf("Done:%s ,state=%+v", groupID, state) - if state.pending == 1 { - payloads := tgc.groupPayloads[groupID] - delete(tgc.groupStates, groupID) - delete(tgc.groupPayloads, groupID) - if tgc.onCompletion != nil && state.hasSuccess { - logrus.Debugf("OnCompletion:%s", groupID) - tgc.mu.Unlock() - tgc.onCompletion(ctx, groupID, payloads...) - tgc.mu.Lock() - } - return + delete(coordinator.groups, groupID) + coordinator.mu.Unlock() + if group.hasSuccess { + finalizeTransferGroup(ctx, groupID, group) } - state.pending-- - tgc.groupStates[groupID] = state } diff --git a/internal/task_group/group_test.go b/internal/task_group/group_test.go new file mode 100644 index 0000000000..7216caa481 --- /dev/null +++ b/internal/task_group/group_test.go @@ -0,0 +1,26 @@ +package task_group + +import ( + "context" + "testing" +) + +func TestTransferGroupKeepsAnySuccessUntilTerminalCompletion(t *testing.T) { + coordinator := &TransferGroupCoordinator{groups: map[string]*transferGroup{}} + const groupID = "/missing-destination" + coordinator.AddTask(groupID) + coordinator.AddTask(groupID) + coordinator.RemoveSource(groupID, "/source") + coordinator.RemoveSource(groupID, "/source") + + coordinator.Done(context.Background(), groupID, true) + group := coordinator.groups[groupID] + if group == nil || group.pending != 1 || !group.hasSuccess || len(group.removeSource) != 1 { + t.Fatalf("intermediate group state = %+v", group) + } + + coordinator.Done(context.Background(), groupID, false) + if coordinator.groups[groupID] != nil { + t.Fatal("terminal group state was not released") + } +} diff --git a/internal/task_group/transfer.go b/internal/task_group/transfer.go index 0e9661b45c..1fa55e30c2 100644 --- a/internal/task_group/transfer.go +++ b/internal/task_group/transfer.go @@ -5,69 +5,29 @@ import ( "fmt" "path" - "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/driver" "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/internal/op" - "github.com/OpenListTeam/OpenList/v4/internal/setting" - "github.com/OpenListTeam/OpenList/v4/pkg/utils" "github.com/pkg/errors" log "github.com/sirupsen/logrus" - "golang.org/x/time/rate" ) -type SrcPathToRemove string - -// ActualPath -type DstPathToHook string - -func HookAndRemove(ctx context.Context, dstPath string, payloads ...any) { +func finalizeTransferGroup(ctx context.Context, dstPath string, group *transferGroup) { dstStorage, dstActualPath, err := op.GetStorageAndActualPath(dstPath) if err != nil { log.Error(errors.WithMessage(err, "failed get dst storage")) return } - dstNeedHandleHook := setting.GetBool(conf.HandleHookAfterWriting) - dstHandleHookLimit := setting.GetFloat(conf.HandleHookRateLimit, .0) - var listLimiter *rate.Limiter - if dstNeedHandleHook && dstHandleHookLimit > .0 { - listLimiter = rate.NewLimiter(rate.Limit(dstHandleHookLimit), 1) - } - hookedPaths := make(map[string]struct{}) - handleHook := func(actualPath string) { - if _, ok := hookedPaths[actualPath]; ok { - return - } - if listLimiter != nil { - _ = listLimiter.Wait(ctx) - } - files, e := op.List(ctx, dstStorage, actualPath, model.ListArgs{SkipHook: true}) - if e != nil { - log.Errorf("failed handle objs update hook: %v", e) - } else { - op.HandleObjsUpdateHook(ctx, utils.GetFullPath(dstStorage.GetStorage().MountPath, actualPath), files) - hookedPaths[actualPath] = struct{}{} + op.ScheduleSnapshotReconciliation(dstStorage, dstActualPath, false) + for path := range group.removeSource { + srcStorage, srcActualPath, err := op.GetStorageAndActualPath(path) + if err != nil { + log.Error(errors.WithMessage(err, "failed get src storage")) + continue } - } - if dstNeedHandleHook { - handleHook(dstActualPath) - } - for _, payload := range payloads { - switch p := payload.(type) { - case DstPathToHook: - if dstNeedHandleHook { - handleHook(string(p)) - } - case SrcPathToRemove: - srcStorage, srcActualPath, err := op.GetStorageAndActualPath(string(p)) - if err != nil { - log.Error(errors.WithMessage(err, "failed get src storage")) - continue - } - err = verifyAndRemove(ctx, srcStorage, dstStorage, srcActualPath, dstActualPath) - if err != nil { - log.Error(err) - } + err = verifyAndRemove(ctx, srcStorage, dstStorage, srcActualPath, dstActualPath) + if err != nil { + log.Error(err) } } } @@ -117,4 +77,4 @@ func verifyAndRemove(ctx context.Context, srcStorage, dstStorage driver.Driver, return nil } -var TransferCoordinator *TaskGroupCoordinator = NewTaskGroupCoordinator("HookAndRemove", HookAndRemove) +var TransferCoordinator = &TransferGroupCoordinator{groups: map[string]*transferGroup{}} diff --git a/pkg/mq/mq.go b/pkg/mq/mq.go deleted file mode 100644 index e442a922c1..0000000000 --- a/pkg/mq/mq.go +++ /dev/null @@ -1,63 +0,0 @@ -package mq - -import ( - "sync" - - "github.com/OpenListTeam/OpenList/v4/pkg/generic" -) - -type Message[T any] struct { - Content T -} - -type BasicConsumer[T any] func(Message[T]) -type AllConsumer[T any] func([]Message[T]) - -type MQ[T any] interface { - Publish(Message[T]) - Consume(BasicConsumer[T]) - ConsumeAll(AllConsumer[T]) - Clear() - Len() int -} - -type inMemoryMQ[T any] struct { - queue generic.Queue[Message[T]] - sync.Mutex -} - -func NewInMemoryMQ[T any]() MQ[T] { - return &inMemoryMQ[T]{queue: *generic.NewQueue[Message[T]]()} -} - -func (mq *inMemoryMQ[T]) Publish(msg Message[T]) { - mq.Lock() - defer mq.Unlock() - mq.queue.Push(msg) -} - -func (mq *inMemoryMQ[T]) Consume(consumer BasicConsumer[T]) { - mq.Lock() - defer mq.Unlock() - for !mq.queue.IsEmpty() { - consumer(mq.queue.Pop()) - } -} - -func (mq *inMemoryMQ[T]) ConsumeAll(consumer AllConsumer[T]) { - mq.Lock() - defer mq.Unlock() - consumer(mq.queue.PopAll()) -} - -func (mq *inMemoryMQ[T]) Clear() { - mq.Lock() - defer mq.Unlock() - mq.queue.Clear() -} - -func (mq *inMemoryMQ[T]) Len() int { - mq.Lock() - defer mq.Unlock() - return mq.queue.Len() -} diff --git a/server/ftp/fsmanage.go b/server/ftp/fsmanage.go index 4993f69ab0..86eb9e5c6d 100644 --- a/server/ftp/fsmanage.go +++ b/server/ftp/fsmanage.go @@ -94,7 +94,7 @@ func Rename(ctx context.Context, oldPath, newPath string) error { return err } if srcBase != dstBase { - err = fs.Rename(ctx, srcPath, dstBase, true) + err = fs.Rename(ctx, srcPath, dstBase) if err != nil { return err } diff --git a/server/ftp/fsup.go b/server/ftp/fsup.go index b70cafa1a1..1fe7b79c70 100644 --- a/server/ftp/fsup.go +++ b/server/ftp/fsup.go @@ -110,7 +110,7 @@ func (f *FileUploadProxy) Close() error { _ = fs.Rename(ctx, f.path, dstBase) } else { if name != dstBase { - e := fs.Rename(ctx, f.path, dstBase, true) + e := fs.Rename(ctx, f.path, dstBase) if e != nil { return } @@ -206,7 +206,7 @@ func (f *FileUploadWithLengthProxy) write(p []byte) (n int, err error) { Reader: reader, } go func() { - e := fs.PutDirectly(f.ctx, dir, s, true) + e := fs.PutDirectly(f.ctx, dir, s) f.errChan <- e close(f.errChan) }() diff --git a/server/handles/fsbatch.go b/server/handles/fsbatch.go index 377863900f..ba0173fd35 100644 --- a/server/handles/fsbatch.go +++ b/server/handles/fsbatch.go @@ -135,9 +135,9 @@ func FsRecursiveMove(c *gin.Context) { } var count = 0 - for i, fileName := range movingFileNames { + for _, fileName := range movingFileNames { // move - _, err := fs.Move(c.Request.Context(), fileName, dstDir, len(movingFileNames) > i+1) + _, err := fs.Move(c.Request.Context(), fileName, dstDir) if err != nil { common.ErrorResp(c, err, 500) return diff --git a/server/handles/fsmanage.go b/server/handles/fsmanage.go index ba512423d5..265d43ecae 100644 --- a/server/handles/fsmanage.go +++ b/server/handles/fsmanage.go @@ -141,11 +141,11 @@ func FsMove(c *gin.Context) { // Create all tasks immediately without any synchronous validation // All validation will be done asynchronously in the background var addedTasks []task.TaskExtensionInfo - for i, p := range req.Names { + for _, p := range req.Names { if p == "" { continue } - t, err := fs.Move(c.Request.Context(), p, dstDir, len(req.Names) > i+1) + t, err := fs.Move(c.Request.Context(), p, dstDir) if t != nil { addedTasks = append(addedTasks, t) } @@ -245,15 +245,15 @@ func FsCopy(c *gin.Context) { // Create all tasks immediately without any synchronous validation // All validation will be done asynchronously in the background var addedTasks []task.TaskExtensionInfo - for i, p := range req.Names { + for _, p := range req.Names { if p == "" { continue } var t task.TaskExtensionInfo if req.Merge { - t, err = fs.Merge(c.Request.Context(), p, dstDir, len(req.Names) > i+1) + t, err = fs.Merge(c.Request.Context(), p, dstDir) } else { - t, err = fs.Copy(c.Request.Context(), p, dstDir, len(req.Names) > i+1) + t, err = fs.Copy(c.Request.Context(), p, dstDir) } if t != nil { addedTasks = append(addedTasks, t) diff --git a/server/webdav/file.go b/server/webdav/file.go index 4d29ff1b45..77fde5c6b0 100644 --- a/server/webdav/file.go +++ b/server/webdav/file.go @@ -58,7 +58,7 @@ func moveFiles(ctx context.Context, src, dst string, overwrite bool) (status int if srcDir == dstDir { err = fs.Rename(ctx, src, dstName) } else { - _, err = fs.Move(context.WithValue(ctx, conf.NoTaskKey, struct{}{}), src, dstDir) + err = fs.MoveDirectly(ctx, src, dstDir) if err != nil { return http.StatusInternalServerError, err } @@ -98,7 +98,7 @@ func copyFiles(ctx context.Context, src, dst string, overwrite bool) (status int if !authz.CanWrite(user, dstMeta, dstDir) { return http.StatusForbidden, nil } - _, err = fs.Copy(context.WithValue(ctx, conf.NoTaskKey, struct{}{}), src, dstDir) + err = fs.CopyDirectly(ctx, src, dstDir) if err != nil { return http.StatusInternalServerError, err } From 4ba846ae3ae8dfe37cedb2228ad24f357121a439 Mon Sep 17 00:00:00 2001 From: nostalume Date: Thu, 24 Sep 2026 23:52:32 +0800 Subject: [PATCH 5/6] fix(search): normalize deleted node parent path - Remove the exact search node and its descendants when a renamed path is deleted. - Cover sibling preservation and close the SQLite fixture before Windows temp cleanup. Co-authored-by: Codex <267193182+codex@users.noreply.github.com> --- internal/db/searchnode.go | 2 +- internal/db/searchnode_test.go | 56 ++++++++++++++++++++++++++++++++++ 2 files changed, 57 insertions(+), 1 deletion(-) create mode 100644 internal/db/searchnode_test.go diff --git a/internal/db/searchnode.go b/internal/db/searchnode.go index 18aa93053c..3f8d0b1d44 100644 --- a/internal/db/searchnode.go +++ b/internal/db/searchnode.go @@ -38,7 +38,7 @@ func DeleteSearchNodesByParent(path string) error { dir, name := stdpath.Dir(path), stdpath.Base(path) return db.Where(fmt.Sprintf("%s = ? AND %s = ?", columnName("parent"), columnName("name")), - dir, name).Delete(&model.SearchNode{}).Error + utils.FixAndCleanPath(dir), name).Delete(&model.SearchNode{}).Error } func ClearSearchNodes() error { diff --git a/internal/db/searchnode_test.go b/internal/db/searchnode_test.go new file mode 100644 index 0000000000..3031238ec4 --- /dev/null +++ b/internal/db/searchnode_test.go @@ -0,0 +1,56 @@ +package db + +import ( + "path/filepath" + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/glebarez/sqlite" + "gorm.io/gorm" +) + +func TestDeleteSearchNodesByParentRemovesNodeAndDescendants(t *testing.T) { + dataDir := t.TempDir() + oldConf, oldDB := conf.Conf, db + t.Cleanup(func() { conf.Conf, db = oldConf, oldDB }) + conf.Conf = conf.DefaultConfig(dataDir) + database, err := gorm.Open(sqlite.Open(filepath.Join(dataDir, "search.db")), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + sqlDB, err := database.DB() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = sqlDB.Close() }) + Init(database) + + nodes := []model.SearchNode{ + {Parent: "/source", Name: "old.srt"}, + {Parent: "/source/old.srt", Name: "child"}, + {Parent: "/source", Name: "keep.srt"}, + } + if err := BatchCreateSearchNodes(&nodes); err != nil { + t.Fatal(err) + } + + if err := DeleteSearchNodesByParent("/source/old.srt"); err != nil { + t.Fatal(err) + } + + remaining, err := GetSearchNodesByParent("/source") + if err != nil { + t.Fatal(err) + } + if len(remaining) != 1 || remaining[0].Name != "keep.srt" { + t.Fatalf("remaining nodes = %#v, want only keep.srt", remaining) + } + descendants, err := GetSearchNodesByParent("/source/old.srt") + if err != nil { + t.Fatal(err) + } + if len(descendants) != 0 { + t.Fatalf("remaining descendants = %#v, want none", descendants) + } +} From b25ba197881667e83a0dc10e31673534570a041f Mon Sep 17 00:00:00 2001 From: nostalume Date: Thu, 24 Sep 2026 23:52:33 +0800 Subject: [PATCH 6/6] test(bootstrap): retain configured projection benchmark - Measure end-to-end namespace mutation through configured search and STRM sinks. - Verify completion, bounded admission, and final projected state for local or loopback WebDAV sources. Co-authored-by: Codex <267193182+codex@users.noreply.github.com> --- .../bootstrap/projection_benchmark_test.go | 272 ++++++++++++++++++ 1 file changed, 272 insertions(+) create mode 100644 internal/bootstrap/projection_benchmark_test.go diff --git a/internal/bootstrap/projection_benchmark_test.go b/internal/bootstrap/projection_benchmark_test.go new file mode 100644 index 0000000000..07dc1116af --- /dev/null +++ b/internal/bootstrap/projection_benchmark_test.go @@ -0,0 +1,272 @@ +//go:build benchmark + +package bootstrap + +import ( + "context" + "encoding/json" + "fmt" + "net/http/httptest" + "os" + "path/filepath" + "sync" + "testing" + "time" + + _ "github.com/OpenListTeam/OpenList/v4/drivers/local" + "github.com/OpenListTeam/OpenList/v4/drivers/strm" + _ "github.com/OpenListTeam/OpenList/v4/drivers/webdav" + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/db" + "github.com/OpenListTeam/OpenList/v4/internal/fs" + "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/OpenListTeam/OpenList/v4/internal/op" + "github.com/OpenListTeam/OpenList/v4/internal/search" + "github.com/OpenListTeam/OpenList/v4/pkg/mq" + "github.com/glebarez/sqlite" + xwebdav "golang.org/x/net/webdav" + "gorm.io/gorm" + "gorm.io/gorm/logger" +) + +// BenchmarkConfiguredProjectionMutation measures one changed 100-object snapshot +// through the configured search and STRM sinks. Set OPENLIST_BENCH_REMOTE=webdav +// to exercise the real WebDav Driver, and set OPENLIST_BENCH_MEILI_URL to select +// a disposable Meilisearch instance instead of SQLite search. +func BenchmarkConfiguredProjectionMutation(b *testing.B) { + root := b.TempDir() + sourceRoot := filepath.Join(root, "source") + projectionRoot := filepath.Join(root, "projection") + if err := os.MkdirAll(sourceRoot, 0o755); err != nil { + b.Fatal(err) + } + for i := range 99 { + name := filepath.Join(sourceRoot, fmt.Sprintf("subtitle-%02d.srt", i)) + if err := os.WriteFile(name, []byte("subtitle"), 0o644); err != nil { + b.Fatal(err) + } + } + toggleA := filepath.Join(sourceRoot, "toggle-a.srt") + toggleB := filepath.Join(sourceRoot, "toggle-b.srt") + if err := os.WriteFile(toggleA, make([]byte, 1<<20), 0o644); err != nil { + b.Fatal(err) + } + + conf.Conf = conf.DefaultConfig(root) + database, err := gorm.Open(sqlite.Open(filepath.Join(root, "benchmark.db")), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) + if err != nil { + b.Fatal(err) + } + sqlDB, err := database.DB() + if err != nil { + b.Fatal(err) + } + b.Cleanup(func() { _ = sqlDB.Close() }) + db.Init(database) + op.Cache.ClearAll() + + saveBenchmarkSetting(b, conf.AutoUpdateIndex, "true") + progress, err := json.Marshal(model.IndexProgress{IsDone: true}) + if err != nil { + b.Fatal(err) + } + saveBenchmarkSetting(b, conf.IndexProgress, string(progress)) + searchMode := "database" + if host := os.Getenv("OPENLIST_BENCH_MEILI_URL"); host != "" { + searchMode = "meilisearch" + conf.Conf.Meilisearch.Host = host + conf.Conf.Meilisearch.Index = fmt.Sprintf("openlist-benchmark-%d", time.Now().UnixNano()) + } + if err := search.Init(searchMode); err != nil { + b.Fatal(err) + } + b.Cleanup(func() { _ = search.Init("none") }) + + storageDriver := "Local" + additionConfig := map[string]any{"root_folder_path": sourceRoot, "thumbnail": false} + if os.Getenv("OPENLIST_BENCH_REMOTE") == "webdav" { + server := httptest.NewServer(&xwebdav.Handler{ + Prefix: "/", FileSystem: xwebdav.Dir(sourceRoot), LockSystem: xwebdav.NewMemLS(), + }) + b.Cleanup(server.Close) + storageDriver = "WebDav" + additionConfig = map[string]any{ + "vendor": "other", "address": server.URL, + "username": "benchmark", "password": "benchmark", "root_folder_path": "/", + } + } + addition, err := json.Marshal(additionConfig) + if err != nil { + b.Fatal(err) + } + storageID, err := op.CreateStorage(context.Background(), model.Storage{ + Driver: storageDriver, MountPath: "/source", Addition: string(addition), CacheExpiration: 5, + }) + if err != nil { + b.Fatal(err) + } + b.Cleanup(func() { _ = op.DeleteStorageById(context.Background(), storageID) }) + + strmDriver := &strm.Strm{ + Storage: model.Storage{MountPath: "/strm"}, + Addition: strm.Addition{ + Paths: "/source", DownloadFileTypes: "srt", SaveStrmToLocal: true, + SaveStrmLocalPath: projectionRoot, SaveLocalMode: strm.SaveLocalSyncMode, Version: 5, + }, + } + if err := strmDriver.Init(context.Background()); err != nil { + b.Fatal(err) + } + b.Cleanup(func() { _ = strmDriver.Drop(context.Background()) }) + + searchTiming, strmTiming := newBenchmarkSinkTiming(), newBenchmarkSinkTiming() + searchProjection = mq.NewLatestProcessor(256, 65536, 4, func(ctx context.Context, parent string, objs []model.Obj) { + search.UpdateSnapshot(ctx, parent, objs) + searchTiming.complete(parent) + }) + strmProjection = mq.NewLatestProcessor(128, 32768, 1, func(ctx context.Context, parent string, objs []model.Obj) { + strm.UpdateLocalStrm(ctx, parent, objs) + strmTiming.complete(parent) + }) + op.SetSnapshotProjector(func(_ context.Context, parent string, objs []model.Obj) { + searchTiming.offer(parent) + offerSnapshot("search", searchProjection, parent, objs) + strmTiming.offer(parent) + offerSnapshot("strm", strmProjection, parent, objs) + }) + b.Cleanup(func() { + op.SetSnapshotProjector(nil) + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + _ = searchProjection.Stop(ctx) + _ = strmProjection.Stop(ctx) + }) + + if _, err := fs.List(context.Background(), "/source", &fs.ListArgs{Refresh: true}); err != nil { + b.Fatal(err) + } + waitBenchmarkProjectionDrain(b) + searchTiming.reset() + strmTiming.reset() + + b.ReportAllocs() + b.ResetTimer() + current, next := toggleA, toggleB + for b.Loop() { + if err := os.Rename(current, next); err != nil { + b.Fatal(err) + } + current, next = next, current + if _, err := fs.List(context.Background(), "/source", &fs.ListArgs{Refresh: true}); err != nil { + b.Fatal(err) + } + waitBenchmarkProjectionDrain(b) + } + b.StopTimer() + + searchTotal, searchCompleted := searchTiming.result() + strmTotal, strmCompleted := strmTiming.result() + if searchCompleted != int64(b.N) || strmCompleted != int64(b.N) { + b.Fatalf("completed search=%d strm=%d, want %d", searchCompleted, strmCompleted, b.N) + } + b.ReportMetric(float64(searchTotal.Nanoseconds())/float64(b.N), "search-ns/op") + b.ReportMetric(float64(strmTotal.Nanoseconds())/float64(b.N), "strm-ns/op") + if searchProjection.Stats().Rejected != 0 || strmProjection.Stats().Rejected != 0 { + b.Fatalf("projection rejected work: search=%+v strm=%+v", searchProjection.Stats(), strmProjection.Stats()) + } + validateBenchmarkProjection(b, projectionRoot, filepath.Base(current), filepath.Base(next)) +} + +type benchmarkSinkTiming struct { + mu sync.Mutex + offered map[string]time.Time + total time.Duration + completed int64 +} + +func newBenchmarkSinkTiming() *benchmarkSinkTiming { + return &benchmarkSinkTiming{offered: make(map[string]time.Time)} +} + +func (s *benchmarkSinkTiming) offer(parent string) { + s.mu.Lock() + s.offered[parent] = time.Now() + s.mu.Unlock() +} + +func (s *benchmarkSinkTiming) complete(parent string) { + s.mu.Lock() + if start, ok := s.offered[parent]; ok { + s.total += time.Since(start) + s.completed++ + delete(s.offered, parent) + } + s.mu.Unlock() +} + +func (s *benchmarkSinkTiming) reset() { + s.mu.Lock() + clear(s.offered) + s.total = 0 + s.completed = 0 + s.mu.Unlock() +} + +func (s *benchmarkSinkTiming) result() (time.Duration, int64) { + s.mu.Lock() + defer s.mu.Unlock() + return s.total, s.completed +} + +func saveBenchmarkSetting(b *testing.B, key, value string) { + b.Helper() + if err := op.SaveSettingItem(&model.SettingItem{Key: key, Value: value}); err != nil { + b.Fatal(err) + } +} + +func validateBenchmarkProjection(b *testing.B, projectionRoot, current, absent string) { + b.Helper() + nodes, _, err := search.Search(context.Background(), model.SearchReq{ + Parent: "/source", PageReq: model.PageReq{Page: 1, PerPage: 1000}, + }) + if err != nil { + b.Fatal(err) + } + count, foundCurrent, foundAbsent := 0, false, false + for _, node := range nodes { + if node.Parent != "/source" { + continue + } + count++ + foundCurrent = foundCurrent || node.Name == current + foundAbsent = foundAbsent || node.Name == absent + } + if count != 100 || !foundCurrent || foundAbsent { + b.Fatalf("search projection count=%d current=%t absent=%t", count, foundCurrent, foundAbsent) + } + if _, err := os.Stat(filepath.Join(projectionRoot, "source", current)); err != nil { + b.Fatalf("current STRM projection: %v", err) + } + if _, err := os.Stat(filepath.Join(projectionRoot, "source", absent)); !os.IsNotExist(err) { + b.Fatalf("stale STRM projection exists or cannot be checked: %v", err) + } +} + +func waitBenchmarkProjectionDrain(b *testing.B) { + b.Helper() + deadline := time.Now().Add(time.Minute) + for { + searchStats := searchProjection.Stats() + strmStats := strmProjection.Stats() + if searchStats.Pending+searchStats.InFlight+strmStats.Pending+strmStats.InFlight == 0 { + return + } + if time.Now().After(deadline) { + b.Fatalf("projection did not drain: search=%+v strm=%+v", searchStats, strmStats) + } + time.Sleep(time.Millisecond) + } +}