diff --git a/inventory/share.go b/inventory/share.go index 0090154b..8ff5e0a5 100644 --- a/inventory/share.go +++ b/inventory/share.go @@ -237,6 +237,13 @@ func IsValidShare(share *ent.Share) error { return ErrOwnerInactive } + // Check the owner's current share permission. + ownerGroup, err := owner.Edges.GroupOrErr() + if err != nil || ownerGroup.Permissions == nil || + !ownerGroup.Permissions.Enabled(int(types.GroupPermissionShare)) { + return ErrSourceFileInvalid + } + // Check source file status file, err := share.Edges.FileOrErr() if err != nil || file.FileChildren == 0 || file.OwnerID != owner.ID { @@ -426,7 +433,8 @@ func withShareEagerLoading(ctx context.Context, q *ent.ShareQuery) *ent.ShareQue } if v, ok := ctx.Value(LoadShareUser{}).(bool); ok && v { q.WithUser(func(q *ent.UserQuery) { - withUserEagerLoading(ctx, q) + userCtx := context.WithValue(ctx, LoadUserGroup{}, true) + withUserEagerLoading(userCtx, q) }) } diff --git a/inventory/share_test.go b/inventory/share_test.go new file mode 100644 index 00000000..ee3415ba --- /dev/null +++ b/inventory/share_test.go @@ -0,0 +1,83 @@ +package inventory + +import ( + "context" + "testing" + + "github.com/cloudreve/Cloudreve/v4/ent" + "github.com/cloudreve/Cloudreve/v4/ent/enttest" + entuser "github.com/cloudreve/Cloudreve/v4/ent/user" + "github.com/cloudreve/Cloudreve/v4/inventory/types" + "github.com/cloudreve/Cloudreve/v4/pkg/boolset" + "github.com/cloudreve/Cloudreve/v4/pkg/conf" + "github.com/stretchr/testify/require" +) + +func TestIsValidShareChecksOwnerAccess(t *testing.T) { + permissions := &boolset.BooleanSet{} + boolset.Set(types.GroupPermissionShare, true, permissions) + allowedGroup := &ent.Group{Permissions: permissions} + + tests := []struct { + name string + status entuser.Status + group *ent.Group + wantErr error + }{ + {name: "active owner with share permission", status: entuser.StatusActive, group: allowedGroup}, + {name: "active owner without share permission", status: entuser.StatusActive, group: &ent.Group{Permissions: &boolset.BooleanSet{}}, wantErr: ErrSourceFileInvalid}, + {name: "missing group", status: entuser.StatusActive, wantErr: ErrSourceFileInvalid}, + {name: "missing permissions", status: entuser.StatusActive, group: &ent.Group{}, wantErr: ErrSourceFileInvalid}, + {name: "manually banned owner", status: entuser.StatusManualBanned, group: allowedGroup, wantErr: ErrOwnerInactive}, + {name: "system banned owner", status: entuser.StatusSysBanned, group: allowedGroup, wantErr: ErrOwnerInactive}, + {name: "inactive owner", status: entuser.StatusInactive, group: allowedGroup, wantErr: ErrOwnerInactive}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + owner := &ent.User{ID: 1, Status: tt.status} + owner.SetGroup(tt.group) + share := &ent.Share{} + share.SetUser(owner) + share.SetFile(&ent.File{OwnerID: owner.ID, FileChildren: 1}) + + require.ErrorIs(t, IsValidShare(share), tt.wantErr) + }) + } +} + +func TestShareClientRevalidatesOwnerGroup(t *testing.T) { + client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared") + t.Cleanup(func() { require.NoError(t, client.Close()) }) + ctx := context.Background() + + permissions := &boolset.BooleanSet{} + boolset.Set(types.GroupPermissionShare, true, permissions) + group := client.Group.Create().SetName("sharing enabled").SetPermissions(permissions).SaveX(ctx) + restrictedGroup := client.Group.Create().SetName("sharing disabled").SetPermissions(&boolset.BooleanSet{}).SaveX(ctx) + owner := client.User.Create().SetEmail("owner@example.com").SetNick("owner").SetGroup(group).SaveX(ctx) + root := client.File.Create().SetName(RootFolderName).SetType(int(types.FileTypeFolder)).SetOwner(owner).SaveX(ctx) + file := client.File.Create().SetName("shared.txt").SetType(int(types.FileTypeFile)).SetOwner(owner).SetParent(root).SaveX(ctx) + share := client.Share.Create().SetUser(owner).SetFile(file).SaveX(ctx) + shareClient := NewShareClient(client, conf.SQLiteDB, nil) + + // Share-info and listing callers only request the owner and file edges. + ctx = context.WithValue(ctx, LoadShareUser{}, true) + ctx = context.WithValue(ctx, LoadShareFile{}, true) + checkShare := func(wantErr error) { + t.Helper() + current, err := shareClient.GetByID(ctx, share.ID) + require.NoError(t, err) + require.ErrorIs(t, IsValidShare(current), wantErr) + } + + checkShare(nil) + client.Group.UpdateOne(group).SetPermissions(&boolset.BooleanSet{}).SaveX(ctx) + checkShare(ErrSourceFileInvalid) + client.Group.UpdateOne(group).SetPermissions(permissions).SaveX(ctx) + checkShare(nil) + client.User.UpdateOne(owner).SetGroup(restrictedGroup).SaveX(ctx) + checkShare(ErrSourceFileInvalid) + client.User.UpdateOne(owner).SetGroup(group).SaveX(ctx) + checkShare(nil) +} diff --git a/pkg/filemanager/fs/dbfs/manage.go b/pkg/filemanager/fs/dbfs/manage.go index 8f28f1b2..48a38c09 100644 --- a/pkg/filemanager/fs/dbfs/manage.go +++ b/pkg/filemanager/fs/dbfs/manage.go @@ -698,6 +698,12 @@ func (f *DBFS) GetFileFromDirectLink(ctx context.Context, dl *ent.DirectLink) (f return nil, fs.ErrDirectLinkInvalid.WithError(fmt.Errorf("file owner is not active")) } + // Check the owner's current direct-link permission. + group, err := owner.Edges.GroupOrErr() + if err != nil || group.Settings == nil || group.Settings.SourceBatchSize <= 0 { + return nil, fs.ErrDirectLinkInvalid + } + file := newFile(nil, fileModel) // Traverse to the root file diff --git a/pkg/filemanager/fs/dbfs/manage_direct_link_test.go b/pkg/filemanager/fs/dbfs/manage_direct_link_test.go new file mode 100644 index 00000000..b6f8d82c --- /dev/null +++ b/pkg/filemanager/fs/dbfs/manage_direct_link_test.go @@ -0,0 +1,87 @@ +package dbfs + +import ( + "context" + "testing" + + "github.com/cloudreve/Cloudreve/v4/ent" + entuser "github.com/cloudreve/Cloudreve/v4/ent/user" + "github.com/cloudreve/Cloudreve/v4/inventory" + "github.com/cloudreve/Cloudreve/v4/inventory/types" + "github.com/cloudreve/Cloudreve/v4/pkg/filemanager/fs" + "github.com/cloudreve/Cloudreve/v4/pkg/serializer" + "github.com/cloudreve/Cloudreve/v4/pkg/setting" + "github.com/stretchr/testify/require" +) + +type directLinkFileClient struct { + inventory.FileClient + root *ent.File +} + +func (c *directLinkFileClient) GetParentFile(_ context.Context, file *ent.File, _ bool) (*ent.File, error) { + if file.FileChildren == c.root.ID { + return c.root, nil + } + return nil, &ent.NotFoundError{} +} + +type directLinkSettingProvider struct { + setting.Provider +} + +func (directLinkSettingProvider) DBFS(context.Context) *setting.DBFS { + return &setting.DBFS{} +} + +func TestGetFileFromDirectLinkChecksOwnerAccess(t *testing.T) { + allowedGroup := &ent.Group{Settings: &types.GroupSetting{SourceBatchSize: 1}} + tests := []struct { + name string + status entuser.Status + group *ent.Group + wantErr bool + }{ + {name: "active owner with direct link permission", status: entuser.StatusActive, group: allowedGroup}, + {name: "active owner without direct link permission", status: entuser.StatusActive, group: &ent.Group{Settings: &types.GroupSetting{}}, wantErr: true}, + {name: "negative batch size", status: entuser.StatusActive, group: &ent.Group{Settings: &types.GroupSetting{SourceBatchSize: -1}}, wantErr: true}, + {name: "missing group", status: entuser.StatusActive, wantErr: true}, + {name: "missing group settings", status: entuser.StatusActive, group: &ent.Group{}, wantErr: true}, + {name: "manually banned owner", status: entuser.StatusManualBanned, group: allowedGroup, wantErr: true}, + {name: "system banned owner", status: entuser.StatusSysBanned, group: allowedGroup, wantErr: true}, + {name: "inactive owner", status: entuser.StatusInactive, group: allowedGroup, wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + owner := &ent.User{ID: 1, Status: tt.status} + owner.SetGroup(tt.group) + root := &ent.File{ID: 1, Name: inventory.RootFolderName, OwnerID: owner.ID} + file := &ent.File{ID: 2, Name: "shared.txt", OwnerID: owner.ID, FileChildren: root.ID} + file.SetOwner(owner) + link := &ent.DirectLink{} + link.SetFile(file) + dbfs := &DBFS{ + user: owner, + fileClient: &directLinkFileClient{root: root}, + settingClient: directLinkSettingProvider{}, + } + + got, err := dbfs.GetFileFromDirectLink(context.Background(), link) + if got != nil { + t.Cleanup(got.(*File).Parent.Recycle) + } + if tt.wantErr { + require.Nil(t, got) + var appErr serializer.AppError + require.ErrorAs(t, err, &appErr) + require.Equal(t, fs.ErrDirectLinkInvalid.Code, appErr.Code) + require.Equal(t, fs.ErrDirectLinkInvalid.Msg, appErr.Msg) + return + } + + require.NoError(t, err) + require.Same(t, file, got.(*File).Model) + }) + } +}