From 64b28bfde806594870f2ad6e15fec41e675cf9bb Mon Sep 17 00:00:00 2001 From: Denis Arh Date: Fri, 18 Oct 2019 08:46:47 +0200 Subject: [PATCH] Refactor role & user repo, cleanup --- system/commands/auth.go | 2 +- system/commands/roles.go | 4 +- system/grpc/roles_service.go | 2 +- system/importer/default.go | 2 +- system/repository/credentials.go | 3 +- system/repository/organisation.go | 3 +- system/repository/repository.go | 79 --------------- system/repository/role.go | 157 +++++++++++++++++++++--------- system/repository/user.go | 17 ++-- system/repository/util.go | 11 --- system/rest/role.go | 154 ++++++++++++++++++++--------- system/service/access_control.go | 4 + system/service/auth.go | 2 +- system/service/role.go | 25 ++--- system/types/role.go | 11 ++- 15 files changed, 260 insertions(+), 216 deletions(-) diff --git a/system/commands/auth.go b/system/commands/auth.go index f8e276cdb..c0846fd28 100644 --- a/system/commands/auth.go +++ b/system/commands/auth.go @@ -98,7 +98,7 @@ func Auth(ctx context.Context, c *cli.Config) *cobra.Command { } if err == nil { - rr, err = roleRepo.FindByMemberID(user.ID) + rr, _, err = roleRepo.Find(types.RoleFilter{MemberID: user.ID}) } if err != nil { diff --git a/system/commands/roles.go b/system/commands/roles.go index f7d64f48d..0b5a42566 100644 --- a/system/commands/roles.go +++ b/system/commands/roles.go @@ -44,12 +44,12 @@ func Roles(ctx context.Context, c *cli.Config) *cobra.Command { c.InitServices(ctx, c) // Try to find role by name and by ID - if rr, err = roleRepo.Find(&types.RoleFilter{Query: roleStr}); err != nil { + if rr, _, err = roleRepo.Find(types.RoleFilter{Query: roleStr}); err != nil { cli.HandleError(err) } else if len(rr) == 1 { role = rr[0] } else if len(rr) > 1 { - cli.HandleError(errors.Errorf("too many roles found with name %q", roleStr)) + cli.HandleError(errors.Errorf("too many roles found with name/handle %q", roleStr)) } else if role == nil { if ID, err = strconv.ParseUint(roleStr, 10, 64); err != nil { // Could not parse ID out of role string diff --git a/system/grpc/roles_service.go b/system/grpc/roles_service.go index 474a7f71a..3b40c2a24 100644 --- a/system/grpc/roles_service.go +++ b/system/grpc/roles_service.go @@ -25,7 +25,7 @@ func (gs roleService) Find(ctx context.Context, req *proto.FindRoleRequest) (rsp rr types.RoleSet ) - if rr, err = gs.roles.Find(&types.RoleFilter{}); err != nil { + if rr, _, err = gs.roles.Find(types.RoleFilter{}); err != nil { return } diff --git a/system/importer/default.go b/system/importer/default.go index 125c66c03..0feb63728 100644 --- a/system/importer/default.go +++ b/system/importer/default.go @@ -19,7 +19,7 @@ func Import(ctx context.Context, ff ...io.Reader) (err error) { aux interface{} ) - roles, err = service.DefaultRole.With(ctx).Find(&types.RoleFilter{}) + roles, _, err = service.DefaultRole.With(ctx).Find(types.RoleFilter{}) if err != nil { return err } diff --git a/system/repository/credentials.go b/system/repository/credentials.go index 669a850c4..83ecaffd4 100644 --- a/system/repository/credentials.go +++ b/system/repository/credentials.go @@ -7,6 +7,7 @@ import ( "github.com/titpetric/factory" + "github.com/cortezaproject/corteza-server/pkg/rh" "github.com/cortezaproject/corteza-server/system/types" ) @@ -58,7 +59,7 @@ func (r *credentials) FindByID(ID uint64) (*types.Credentials, error) { sql := fmt.Sprintf(sqlCredentialsSelect, r.tblname) + " AND id = ?" mod := &types.Credentials{} - return mod, isFound(r.db().Get(mod, sql, ID), mod.ID > 0, ErrCredentialsNotFound) + return mod, rh.IsFound(r.db().Get(mod, sql, ID), mod.ID > 0, ErrCredentialsNotFound) } func (r *credentials) FindByCredentials(kind, credentials string) (cc types.CredentialsSet, err error) { diff --git a/system/repository/organisation.go b/system/repository/organisation.go index e0ccaf90f..d6d43688f 100644 --- a/system/repository/organisation.go +++ b/system/repository/organisation.go @@ -6,6 +6,7 @@ import ( "github.com/titpetric/factory" + "github.com/cortezaproject/corteza-server/pkg/rh" "github.com/cortezaproject/corteza-server/system/types" ) @@ -52,7 +53,7 @@ func (r *organisation) FindByID(id uint64) (*types.Organisation, error) { sql := "SELECT * FROM " + r.organisations + " WHERE id = ? AND " + sqlOrganisationScope mod := &types.Organisation{} - return mod, isFound(r.db().Get(mod, sql, id), mod.ID > 0, ErrOrganisationNotFound) + return mod, rh.IsFound(r.db().Get(mod, sql, id), mod.ID > 0, ErrOrganisationNotFound) } func (r *organisation) Find(filter *types.OrganisationFilter) ([]*types.Organisation, error) { diff --git a/system/repository/repository.go b/system/repository/repository.go index 830842ec8..5383b6bec 100644 --- a/system/repository/repository.go +++ b/system/repository/repository.go @@ -3,11 +3,7 @@ package repository import ( "context" - "github.com/lann/builder" "github.com/titpetric/factory" - "gopkg.in/Masterminds/squirrel.v1" - - "github.com/cortezaproject/corteza-server/pkg/auth" ) type ( @@ -22,11 +18,6 @@ func DB(ctx context.Context) *factory.DB { return factory.Database.MustGet("system").With(ctx) } -// Identity returns the User ID from context -func Identity(ctx context.Context) uint64 { - return auth.GetIdentityFromContext(ctx).Identity() -} - // With updates repository and database contexts func (r *repository) With(ctx context.Context, db *factory.DB) *repository { return &repository{ @@ -47,73 +38,3 @@ func (r *repository) db() *factory.DB { } return DB(r.ctx) } - -func (r repository) fetchOne(one interface{}, q squirrel.SelectBuilder) (err error) { - var ( - sql string - args []interface{} - ) - - if sql, args, err = q.ToSql(); err != nil { - return - } - - if err = r.db().Get(one, sql, args...); err != nil { - return - } - - return -} - -// Fetches single row from table -func (r repository) fetchSet(set interface{}, q squirrel.SelectBuilder) (err error) { - var ( - sql string - args []interface{} - ) - - if sql, args, err = q.ToSql(); err != nil { - return - } - - if err = r.db().Select(set, sql, args...); err != nil { - return - } - - return -} - -// Fetches paged rows -func (r repository) fetchPaged(set interface{}, q squirrel.SelectBuilder, page, perPage uint) error { - if perPage > 0 { - q = q.Limit(uint64(perPage)) - } - - if page > 0 { - q = q.Offset(uint64((page - 1) * perPage)) - } - - if sqlSelect, argsSelect, err := q.ToSql(); err != nil { - return err - } else { - return r.db().Select(set, sqlSelect, argsSelect...) - } -} - -// Counts all rows that match conditions from given query builder -func (r repository) count(q squirrel.SelectBuilder) (uint, error) { - var ( - count uint - cq = builder.Delete(q, "Columns").(squirrel.SelectBuilder).Column("COUNT(*)") - ) - - if sqlSelect, argsSelect, err := cq.ToSql(); err != nil { - return 0, err - } else { - if err := r.db().Get(&count, sqlSelect, argsSelect...); err != nil { - return 0, err - } - } - - return count, nil -} diff --git a/system/repository/role.go b/system/repository/role.go index 0dc45839b..20db1f9df 100644 --- a/system/repository/role.go +++ b/system/repository/role.go @@ -5,7 +5,9 @@ import ( "time" "github.com/titpetric/factory" + "gopkg.in/Masterminds/squirrel.v1" + "github.com/cortezaproject/corteza-server/pkg/rh" "github.com/cortezaproject/corteza-server/system/types" ) @@ -16,8 +18,7 @@ type ( FindByID(id uint64) (*types.Role, error) FindByName(name string) (*types.Role, error) FindByHandle(handle string) (*types.Role, error) - FindByMemberID(userID uint64) (types.RoleSet, error) - Find(filter *types.RoleFilter) (types.RoleSet, error) + Find(filter types.RoleFilter) (types.RoleSet, types.RoleFilter, error) Create(mod *types.Role) (*types.Role, error) Update(mod *types.Role) (*types.Role, error) @@ -45,8 +46,6 @@ type ( ) const ( - sqlRoleScope = "deleted_at IS NULL AND archived_at IS NULL" - ErrRoleNotFound = repositoryError("RoleNotFound") ) @@ -58,82 +57,152 @@ func Role(ctx context.Context, db *factory.DB) RoleRepository { func (r *role) With(ctx context.Context, db *factory.DB) RoleRepository { return &role{ repository: r.repository.With(ctx, db), - roles: "sys_role", - members: "sys_role_member", } } -func (r *role) FindByID(id uint64) (*types.Role, error) { - sql := "SELECT * FROM " + r.roles + " WHERE id = ? AND " + sqlRoleScope - mod := &types.Role{} +func (r role) table() string { + return "sys_role" +} - return mod, isFound(r.db().Get(mod, sql, id), mod.ID > 0, ErrRoleNotFound) +func (r role) tableMember() string { + return "sys_role_member" +} + +func (r role) columns() []string { + return []string{ + "id", + "name", + "handle", + "created_at", + "updated_at", + "archived_at", + "deleted_at", + } +} + +func (r role) query() squirrel.SelectBuilder { + return squirrel. + Select(r.columns()...). + From(r.table() + " AS r"). + Where(squirrel.And{ + squirrel.Eq{"deleted_at": nil}, + squirrel.Eq{"archived_at": nil}, + }) + +} + +func (r role) FindByID(id uint64) (*types.Role, error) { + return r.findOneBy("id", id) } func (r role) FindByHandle(handle string) (*types.Role, error) { - sql := "SELECT * FROM " + r.roles + " WHERE handle = ? AND " + sqlRoleScope - mod := &types.Role{} - - return mod, isFound(r.db().Get(mod, sql, handle), mod.ID > 0, ErrRoleNotFound) + return r.findOneBy("handle", handle) } func (r role) FindByName(name string) (*types.Role, error) { - sql := "SELECT * FROM " + r.roles + " WHERE name = ? AND " + sqlRoleScope - mod := &types.Role{} - - return mod, isFound(r.db().Get(mod, sql, name), mod.ID > 0, ErrRoleNotFound) + return r.findOneBy("name", name) } -func (r *role) FindByMemberID(userID uint64) (types.RoleSet, error) { - sql := "SELECT * FROM " + r.roles + " where id in (select rel_role from " + r.members + " where rel_user=?) and " + sqlRoleScope - rval := make([]*types.Role, 0) - if err := r.db().Select(&rval, sql, userID); err != nil { +func (r role) findOneBy(field string, value interface{}) (*types.Role, error) { + var ( + ro = &types.Role{} + q = r.query(). + Where(squirrel.Eq{field: value}) + + err = rh.FetchOne(r.db(), q, ro) + ) + + if err != nil { return nil, err + } else if ro.ID == 0 { + return nil, ErrRoleNotFound } - return rval, nil + + return ro, nil } -func (r *role) Find(filter *types.RoleFilter) (types.RoleSet, error) { - rval := make([]*types.Role, 0) - params := make([]interface{}, 0) +func (r *role) Find(filter types.RoleFilter) (set types.RoleSet, f types.RoleFilter, err error) { + f = filter - sql := "SELECT * FROM " + r.roles + " WHERE " + sqlRoleScope - - if filter != nil { - if filter.Query != "" { - sql += " AND name LIKE ?" - params = append(params, filter.Query+"%") - } + if f.Sort == "" { + f.Sort = "id" } - sql += " ORDER BY name ASC" + query := r.query() - return rval, r.db().Select(&rval, sql, params...) + if !f.IncDeleted { + query = query.Where(squirrel.Eq{"r.deleted_at": nil}) + } + + if !f.IncArchived { + query = query.Where(squirrel.Eq{"r.archived_at": nil}) + } + + if len(f.RoleID) > 0 { + query = query.Where(squirrel.Eq{"r.ID": f.RoleID}) + } + + if f.MemberID > 0 { + query = query.Where(squirrel.Expr("r.ID IN (SELECT rel_role FROM sys_role_member AS m WHERE m.rel_user = ?)", f.MemberID)) + } + + if f.Query != "" { + qs := f.Query + "%" + query = query.Where(squirrel.Or{ + squirrel.Like{"r.name": qs}, + squirrel.Like{"r.handle": qs}, + }) + } + + if f.Name != "" { + query = query.Where(squirrel.Eq{"r.name": f.Name}) + } + + if f.Handle != "" { + query = query.Where(squirrel.Eq{"r.handle": f.Handle}) + } + + if f.IsReadable != nil { + query = query.Where(f.IsReadable) + } + + var orderBy []string + if orderBy, err = rh.ParseOrder(f.Sort, r.columns()...); err != nil { + return + } else { + query = query.OrderBy(orderBy...) + } + + if f.Count, err = rh.Count(r.db(), query); err != nil || f.Count == 0 { + return + } + + return set, f, rh.FetchPaged(r.db(), query, f.Page, f.PerPage, &set) } func (r *role) Create(mod *types.Role) (*types.Role, error) { mod.ID = factory.Sonyflake.NextID() mod.CreatedAt = time.Now() - return mod, r.db().Insert(r.roles, mod) + return mod, r.db().Insert(r.table(), mod) } func (r *role) Update(mod *types.Role) (*types.Role, error) { mod.UpdatedAt = timeNowPtr() - return mod, r.db().Replace(r.roles, mod) + return mod, r.db().Replace(r.table(), mod) } func (r *role) ArchiveByID(id uint64) error { - return r.updateColumnByID(r.roles, "archived_at", time.Now(), id) + return r.updateColumnByID(r.table(), "archived_at", time.Now(), id) } func (r *role) UnarchiveByID(id uint64) error { - return r.updateColumnByID(r.roles, "archived_at", nil, id) + return r.updateColumnByID(r.table(), "archived_at", nil, id) } func (r *role) DeleteByID(id uint64) error { - return r.updateColumnByID(r.roles, "deleted_at", time.Now(), id) + return r.updateColumnByID(r.table(), "deleted_at", time.Now(), id) } func (r *role) MergeByID(id, targetRoleID uint64) error { @@ -146,13 +215,13 @@ func (r *role) MoveByID(id, targetOrganisationID uint64) error { func (r *role) MembershipsFindByUserID(roleID uint64) (mm []*types.RoleMember, err error) { rval := make([]*types.RoleMember, 0) - sql := "SELECT * FROM " + r.members + " WHERE rel_user = ?" + sql := "SELECT * FROM " + r.tableMember() + " WHERE rel_user = ?" return rval, r.db().Select(&rval, sql, roleID) } func (r *role) MemberFindByRoleID(roleID uint64) (mm []*types.RoleMember, err error) { rval := make([]*types.RoleMember, 0) - sql := "SELECT * FROM " + r.members + " WHERE rel_role = ?" + sql := "SELECT * FROM " + r.tableMember() + " WHERE rel_role = ?" return rval, r.db().Select(&rval, sql, roleID) } @@ -161,7 +230,7 @@ func (r *role) MemberAddByID(roleID, userID uint64) error { RoleID: roleID, UserID: userID, } - return r.db().Replace(r.members, mod) + return r.db().Replace(r.tableMember(), mod) } func (r *role) MemberRemoveByID(roleID, userID uint64) error { @@ -169,5 +238,5 @@ func (r *role) MemberRemoveByID(roleID, userID uint64) error { RoleID: roleID, UserID: userID, } - return r.db().Delete(r.members, mod, "rel_role", "rel_user") + return r.db().Delete(r.tableMember(), mod, "rel_role", "rel_user") } diff --git a/system/repository/user.go b/system/repository/user.go index 15e9f8876..5750a91bd 100644 --- a/system/repository/user.go +++ b/system/repository/user.go @@ -70,13 +70,13 @@ func (r user) columns() []string { } func (r user) query() squirrel.SelectBuilder { - return r.queryNoFilter().Where("u.deleted_at IS NULL AND u.suspended_at IS NULL") -} - -func (r user) queryNoFilter() squirrel.SelectBuilder { return squirrel. Select(r.columns()...). - From(r.table() + " AS u") + From(r.table() + " AS u"). + Where(squirrel.And{ + squirrel.Eq{"deleted_at": nil}, + squirrel.Eq{"suspended_at": nil}, + }) } func (r *user) With(ctx context.Context, db *factory.DB) UserRepository { @@ -127,7 +127,7 @@ func (r user) Find(filter types.UserFilter) (set types.UserSet, f types.UserFilt f.Sort = "id" } - query := r.queryNoFilter() + query := r.query() // Returns user filter (flt) wrapped in IF() function with cnd as condition (when cnd != nil) whereMasked := func(cnd *permissions.ResourceFilter, flt squirrel.Sqlizer) squirrel.Sqlizer { @@ -155,7 +155,7 @@ func (r user) Find(filter types.UserFilter) (set types.UserSet, f types.UserFilt // Due to lack of support for more exotic expressions (slice of values inside subquery) // we'll use set of OR expressions as a workaround for _, roleID := range f.RoleID { - or = append(or, squirrel.Expr("u.ID IN (SELECT rel_user FROM sys_role_member WHERE rel_role IN (?))", roleID)) + or = append(or, squirrel.Expr("u.ID IN (SELECT rel_user FROM sys_role_member WHERE rel_role = ?)", roleID)) } query = query.Where(or) @@ -206,8 +206,7 @@ func (r user) Find(filter types.UserFilter) (set types.UserSet, f types.UserFilt } func (r user) Total() (count uint) { - - count, _ = r.count(r.query()) + count, _ = rh.Count(r.db(), squirrel.Select().From(r.table())) return } diff --git a/system/repository/util.go b/system/repository/util.go index 0b7f4cd8b..2b5f0d1b6 100644 --- a/system/repository/util.go +++ b/system/repository/util.go @@ -18,17 +18,6 @@ func exec(_ interface{}, err error) error { return errors.WithStack(err) } -// Returns err if set otherwise it returns nerr if not valid -func isFound(err error, valid bool, nerr error) error { - if err != nil { - return errors.WithStack(err) - } else if !valid { - return errors.WithStack(nerr) - } - - return nil -} - func timeNowPtr() *time.Time { n := time.Now() return &n diff --git a/system/rest/role.go b/system/rest/role.go index d5fe44930..c65abd8c6 100644 --- a/system/rest/role.go +++ b/system/rest/role.go @@ -16,98 +16,127 @@ var _ = errors.Wrap type ( Role struct { - svc struct { - role service.RoleService - } + role service.RoleService + ac roleAccessController + } + + roleAccessController interface { + CanGrant(context.Context) bool + + CanUpdateRole(context.Context, *types.Role) bool + CanDeleteRole(context.Context, *types.Role) bool + } + + rolePayload struct { + *types.Role + + CanGrant bool `json:"canGrant"` + CanUpdateRole bool `json:"canUpdateRole"` + CanDeleteRole bool `json:"canDeleteRole"` + } + + roleSetPayload struct { + Filter types.RoleFilter `json:"filter"` + Set []*rolePayload `json:"set"` } ) func (Role) New() *Role { - ctrl := &Role{} - ctrl.svc.role = service.DefaultRole - return ctrl -} - -func (ctrl *Role) Read(ctx context.Context, r *request.RoleRead) (interface{}, error) { - return ctrl.svc.role.With(ctx).FindByID(r.RoleID) -} - -func (ctrl *Role) List(ctx context.Context, r *request.RoleList) (interface{}, error) { - return ctrl.svc.role.With(ctx).Find(&types.RoleFilter{Query: r.Query}) -} - -func (ctrl *Role) Create(ctx context.Context, r *request.RoleCreate) (interface{}, error) { - role := &types.Role{ - Name: r.Name, - Handle: r.Handle, + return &Role{ + role: service.DefaultRole, + ac: service.DefaultAccessControl, } +} - role, err := ctrl.svc.role.With(ctx).Create(role) +func (ctrl Role) Read(ctx context.Context, r *request.RoleRead) (interface{}, error) { + role, err := ctrl.role.With(ctx).FindByID(r.RoleID) + return ctrl.makePayload(ctx, role, err) +} + +func (ctrl Role) List(ctx context.Context, r *request.RoleList) (interface{}, error) { + set, filter, err := ctrl.role.With(ctx).Find(types.RoleFilter{Query: r.Query}) + return ctrl.makeFilterPayload(ctx, set, filter, err) +} + +func (ctrl Role) Create(ctx context.Context, r *request.RoleCreate) (interface{}, error) { + var ( + err error + role = &types.Role{ + Name: r.Name, + Handle: r.Handle, + } + ) + + role, err = ctrl.role.With(ctx).Create(role) if err != nil { return nil, err } for _, userID := range payload.ParseUInt64s(r.Members) { - err := ctrl.svc.role.With(ctx).MemberAdd(role.ID, userID) + err := ctrl.role.With(ctx).MemberAdd(role.ID, userID) if err != nil { return nil, err } } - return role, nil + return ctrl.makePayload(ctx, role, err) } -func (ctrl *Role) Update(ctx context.Context, r *request.RoleUpdate) (interface{}, error) { - role := &types.Role{ - ID: r.RoleID, - Name: r.Name, - Handle: r.Handle, - } +func (ctrl Role) Update(ctx context.Context, r *request.RoleUpdate) (interface{}, error) { + var ( + err error + role = &types.Role{ + ID: r.RoleID, + Name: r.Name, + Handle: r.Handle, + } + ) - role, err := ctrl.svc.role.With(ctx).Update(role) + role, err = ctrl.role.With(ctx).Update(role) if err != nil { return nil, err } if len(r.Members) > 0 { - members, err := ctrl.svc.role.With(ctx).MemberList(r.RoleID) + members, err := ctrl.role.With(ctx).MemberList(r.RoleID) if err != nil { return nil, err } for _, member := range members { - err := ctrl.svc.role.With(ctx).MemberRemove(role.ID, member.UserID) + err := ctrl.role.With(ctx).MemberRemove(role.ID, member.UserID) if err != nil { return nil, err } } for _, userID := range payload.ParseUInt64s(r.Members) { - err := ctrl.svc.role.With(ctx).MemberAdd(role.ID, userID) + err := ctrl.role.With(ctx).MemberAdd(role.ID, userID) if err != nil { return nil, err } } } - return role, nil + + return ctrl.makePayload(ctx, role, err) } -func (ctrl *Role) Delete(ctx context.Context, r *request.RoleDelete) (interface{}, error) { - return resputil.OK(), ctrl.svc.role.With(ctx).Delete(r.RoleID) +func (ctrl Role) Delete(ctx context.Context, r *request.RoleDelete) (interface{}, error) { + return resputil.OK(), ctrl.role.With(ctx).Delete(r.RoleID) } -func (ctrl *Role) Archive(ctx context.Context, r *request.RoleArchive) (interface{}, error) { - return resputil.OK(), ctrl.svc.role.With(ctx).Archive(r.RoleID) +func (ctrl Role) Archive(ctx context.Context, r *request.RoleArchive) (interface{}, error) { + return resputil.OK(), ctrl.role.With(ctx).Archive(r.RoleID) } -func (ctrl *Role) Merge(ctx context.Context, r *request.RoleMerge) (interface{}, error) { - return resputil.OK(), ctrl.svc.role.With(ctx).Merge(r.RoleID, r.Destination) +func (ctrl Role) Merge(ctx context.Context, r *request.RoleMerge) (interface{}, error) { + return resputil.OK(), ctrl.role.With(ctx).Merge(r.RoleID, r.Destination) } -func (ctrl *Role) Move(ctx context.Context, r *request.RoleMove) (interface{}, error) { - return resputil.OK(), ctrl.svc.role.With(ctx).Move(r.RoleID, r.OrganisationID) +func (ctrl Role) Move(ctx context.Context, r *request.RoleMove) (interface{}, error) { + return resputil.OK(), ctrl.role.With(ctx).Move(r.RoleID, r.OrganisationID) } -func (ctrl *Role) MemberList(ctx context.Context, r *request.RoleMemberList) (interface{}, error) { - if mm, err := ctrl.svc.role.With(ctx).MemberList(r.RoleID); err != nil { +func (ctrl Role) MemberList(ctx context.Context, r *request.RoleMemberList) (interface{}, error) { + if mm, err := ctrl.role.With(ctx).MemberList(r.RoleID); err != nil { return nil, err } else { rval := make([]string, len(mm)) @@ -118,10 +147,39 @@ func (ctrl *Role) MemberList(ctx context.Context, r *request.RoleMemberList) (in } } -func (ctrl *Role) MemberAdd(ctx context.Context, r *request.RoleMemberAdd) (interface{}, error) { - return resputil.OK(), ctrl.svc.role.With(ctx).MemberAdd(r.RoleID, r.UserID) +func (ctrl Role) MemberAdd(ctx context.Context, r *request.RoleMemberAdd) (interface{}, error) { + return resputil.OK(), ctrl.role.With(ctx).MemberAdd(r.RoleID, r.UserID) } -func (ctrl *Role) MemberRemove(ctx context.Context, r *request.RoleMemberRemove) (interface{}, error) { - return resputil.OK(), ctrl.svc.role.With(ctx).MemberRemove(r.RoleID, r.UserID) +func (ctrl Role) MemberRemove(ctx context.Context, r *request.RoleMemberRemove) (interface{}, error) { + return resputil.OK(), ctrl.role.With(ctx).MemberRemove(r.RoleID, r.UserID) +} + +func (ctrl Role) makePayload(ctx context.Context, m *types.Role, err error) (*rolePayload, error) { + if err != nil || m == nil { + return nil, err + } + + return &rolePayload{ + Role: m, + + CanGrant: ctrl.ac.CanGrant(ctx), + + CanUpdateRole: ctrl.ac.CanUpdateRole(ctx, m), + CanDeleteRole: ctrl.ac.CanDeleteRole(ctx, m), + }, nil +} + +func (ctrl Role) makeFilterPayload(ctx context.Context, nn types.RoleSet, f types.RoleFilter, err error) (*roleSetPayload, error) { + if err != nil { + return nil, err + } + + msp := &roleSetPayload{Filter: f, Set: make([]*rolePayload, len(nn))} + + for i := range nn { + msp.Set[i], _ = ctrl.makePayload(ctx, nn[i], nil) + } + + return msp, nil } diff --git a/system/service/access_control.go b/system/service/access_control.go index 44437e0bf..8a712b8db 100644 --- a/system/service/access_control.go +++ b/system/service/access_control.go @@ -90,6 +90,10 @@ func (svc accessControl) CanReadRole(ctx context.Context, rl *types.Role) bool { return svc.can(ctx, rl, "read", permissions.Allowed) } +func (svc accessControl) FilterReadableRoles(ctx context.Context) *permissions.ResourceFilter { + return svc.permissions.ResourceFilter(ctx, types.RolePermissionResource, "read", permissions.Allow) +} + func (svc accessControl) CanUpdateRole(ctx context.Context, rl *types.Role) bool { return svc.can(ctx, rl, "update") } diff --git a/system/service/auth.go b/system/service/auth.go index d6ef3dc2c..3af1f5d62 100644 --- a/system/service/auth.go +++ b/system/service/auth.go @@ -858,7 +858,7 @@ func (svc auth) autoPromote(u *types.User) (err error) { } func (svc auth) LoadRoleMemberships(u *types.User) error { - rr, err := svc.roles.FindByMemberID(u.ID) + rr, _, err := svc.roles.Find(types.RoleFilter{MemberID: u.ID}) if err != nil { return err } diff --git a/system/service/role.go b/system/service/role.go index 1d2172e24..60c1ef8a3 100644 --- a/system/service/role.go +++ b/system/service/role.go @@ -10,6 +10,7 @@ import ( "github.com/cortezaproject/corteza-server/pkg/handle" "github.com/cortezaproject/corteza-server/pkg/logger" + "github.com/cortezaproject/corteza-server/pkg/permissions" "github.com/cortezaproject/corteza-server/system/repository" "github.com/cortezaproject/corteza-server/system/types" ) @@ -36,6 +37,8 @@ type ( CanUpdateRole(context.Context, *types.Role) bool CanDeleteRole(context.Context, *types.Role) bool CanManageRoleMembers(context.Context, *types.Role) bool + + FilterReadableRoles(ctx context.Context) *permissions.ResourceFilter } RoleService interface { @@ -44,7 +47,7 @@ type ( FindByID(roleID uint64) (*types.Role, error) FindByName(name string) (*types.Role, error) FindByHandle(handle string) (*types.Role, error) - Find(filter *types.RoleFilter) ([]*types.Role, error) + Find(types.RoleFilter) (types.RoleSet, types.RoleFilter, error) Create(role *types.Role) (*types.Role, error) Update(role *types.Role) (*types.Role, error) @@ -105,19 +108,9 @@ func (svc role) findByID(roleID uint64) (*types.Role, error) { return role, nil } -func (svc role) Find(filter *types.RoleFilter) ([]*types.Role, error) { - roles, err := svc.role.Find(filter) - if err != nil { - return nil, err - } - - ret := []*types.Role{} - for _, role := range roles { - if svc.ac.CanReadRole(svc.ctx, role) { - ret = append(ret, role) - } - } - return ret, nil +func (svc role) Find(f types.RoleFilter) (types.RoleSet, types.RoleFilter, error) { + f.IsReadable = svc.ac.FilterReadableRoles(svc.ctx) + return svc.role.Find(f) } func (svc role) FindByName(rolename string) (*types.Role, error) { @@ -186,13 +179,13 @@ func (svc role) Update(mod *types.Role) (t *types.Role, err error) { func (svc role) UniqueCheck(r *types.Role) (err error) { if r.Handle != "" { - if ex, _ := svc.role.FindByHandle(r.Handle); ex.ID > 0 && ex.ID != r.ID { + if ex, _ := svc.role.FindByHandle(r.Handle); ex != nil && ex.ID > 0 && ex.ID != r.ID { return ErrRoleHandleNotUnique } } if r.Name != "" { - if ex, _ := svc.role.FindByName(r.Name); ex.ID > 0 && ex.ID != r.ID { + if ex, _ := svc.role.FindByName(r.Name); ex != nil && ex.ID > 0 && ex.ID != r.ID { return ErrRoleNameNotUnique } } diff --git a/system/types/role.go b/system/types/role.go index 6477441ab..15cf5f7a2 100644 --- a/system/types/role.go +++ b/system/types/role.go @@ -20,7 +20,16 @@ type ( } RoleFilter struct { - Query string + RoleID []uint64 `json:"roleID"` + MemberID uint64 `json:"memberID"` + + Query string `json:"query"` + + Handle string `json:"handle"` + Name string `json:"name"` + + IncDeleted bool `json:"incDeleted"` + IncArchived bool `json:"incArchived"` Sort string `json:"sort"`