From 608c4fe91d86e6b63e214ab6a5f14427c1a8a01b Mon Sep 17 00:00:00 2001 From: Denis Arh Date: Fri, 18 Oct 2019 18:02:00 +0200 Subject: [PATCH] Refactor channel repo --- messaging/importer/default.go | 2 +- messaging/provision-config.go | 4 +- messaging/repository/channel.go | 208 ++++++++++-------- messaging/repository/channel_member.go | 16 +- messaging/rest/channel.go | 4 +- messaging/service/channel.go | 21 +- messaging/types/channel.go | 6 + messaging/websocket/session.go | 2 +- .../websocket/session_incoming_channel.go | 2 +- 9 files changed, 155 insertions(+), 110 deletions(-) diff --git a/messaging/importer/default.go b/messaging/importer/default.go index f67aa9cde..5834387e3 100644 --- a/messaging/importer/default.go +++ b/messaging/importer/default.go @@ -20,7 +20,7 @@ func Import(ctx context.Context, ff ...io.Reader) (err error) { aux interface{} ) - cc, err = service.DefaultChannel.With(ctx).Find(&types.ChannelFilter{}) + cc, _, err = service.DefaultChannel.With(ctx).Find(types.ChannelFilter{}) if err != nil { return err } diff --git a/messaging/provision-config.go b/messaging/provision-config.go index 6626e9dc8..f7ec57a64 100644 --- a/messaging/provision-config.go +++ b/messaging/provision-config.go @@ -42,6 +42,6 @@ func provisionConfig(ctx context.Context, cmd *cobra.Command, c *cli.Config) err // Provision ONLY when there are no channels (even if we find delete channels we abort provisioning func isProvisioned(ctx context.Context) (bool, error) { - cc, err := service.DefaultChannel.With(ctx).Find(&types.ChannelFilter{IncludeDeleted: true}) - return len(cc) > 0, err + _, f, err := service.DefaultChannel.With(ctx).Find(types.ChannelFilter{IncludeDeleted: true}) + return f.Count > 0, err } diff --git a/messaging/repository/channel.go b/messaging/repository/channel.go index 9a43298a9..a3927d763 100644 --- a/messaging/repository/channel.go +++ b/messaging/repository/channel.go @@ -4,9 +4,11 @@ import ( "context" "sort" "strconv" + "strings" "time" "github.com/titpetric/factory" + "gopkg.in/Masterminds/squirrel.v1" "github.com/cortezaproject/corteza-server/messaging/types" "github.com/cortezaproject/corteza-server/pkg/rh" @@ -18,7 +20,7 @@ type ( FindByID(id uint64) (*types.Channel, error) FindByMemberSet(memberID ...uint64) (*types.Channel, error) - Find(filter *types.ChannelFilter) (types.ChannelSet, error) + Find(types.ChannelFilter) (types.ChannelSet, types.ChannelFilter, error) Create(mod *types.Channel) (*types.Channel, error) Update(mod *types.Channel) (*types.Channel, error) @@ -38,45 +40,6 @@ type ( ) const ( - sqlChannelColumns = " id," + - "name, " + - "meta, " + - "membership_policy, " + - "created_at, " + - "updated_at, " + - "archived_at, " + - "deleted_at, " + - "rel_organisation, " + - "rel_creator, " + - "type , " + - "rel_last_message, " + - "topic" - - sqlChannelSelect = `SELECT ` + sqlChannelColumns + ` - FROM messaging_channel AS c - WHERE true ` - - sqlChannelGroupByMemberSet = sqlChannelSelect + ` AND c.type = ? AND c.id IN ( - SELECT rel_channel - FROM messaging_channel_member - GROUP BY rel_channel - HAVING COUNT(*) = ? - AND CONCAT(GROUP_CONCAT(rel_user ORDER BY 1 ASC SEPARATOR ','),',') = ? - )` - - // subquery that filters out all channels that current user has access to as a member - // or via channel type (public channels) - sqlChannelAccess = ` ( - SELECT id - FROM messaging_channel c - LEFT OUTER JOIN messaging_channel_member AS m ON (c.id = m.rel_channel) - WHERE rel_user = ? - UNION - SELECT id - FROM messaging_channel c - WHERE c.type = ? - )` - ErrChannelNotFound = repositoryError("ChannelNotFound") ) @@ -84,66 +47,139 @@ func Channel(ctx context.Context, db *factory.DB) ChannelRepository { return (&channel{}).With(ctx, db) } -func (r *channel) With(ctx context.Context, db *factory.DB) ChannelRepository { +func (r channel) With(ctx context.Context, db *factory.DB) ChannelRepository { return &channel{ repository: r.repository.With(ctx, db), } } -func (r *channel) FindByID(id uint64) (*types.Channel, error) { - mod := &types.Channel{} - sql := sqlChannelSelect + " AND id = ?" - - return mod, isFound(r.db().Get(mod, sql, id), mod.ID > 0, ErrChannelNotFound) +func (r channel) table() string { + return "messaging_channel" } -// FindChannelByMemberSet searches for channel (group!) with exactly the same membership structure -func (r *channel) FindByMemberSet(memberIDs ...uint64) (*types.Channel, error) { - mod := &types.Channel{} +func (r channel) tableMember() string { + return "messaging_channel_member" +} +func (r channel) columns() []string { + return []string{ + "c.id", + "c.name", + "c.meta", + "c.membership_policy", + "c.created_at", + "c.updated_at", + "c.archived_at", + "c.deleted_at", + "c.rel_organisation", + "c.rel_creator", + "c.type", + "c.rel_last_message", + "c.topic", + } +} + +func (r channel) query() squirrel.SelectBuilder { + return squirrel. + Select(r.columns()...). + From(r.table() + " AS c") +} + +func (r channel) FindByID(ID uint64) (*types.Channel, error) { + return r.findOneBy(squirrel.Eq{"c.id": ID}) +} + +// FindByMemberSet searches for channel (group!) with exactly the same membership structure +func (r channel) FindByMemberSet(memberIDs ...uint64) (*types.Channel, error) { + // Make sure members are sorted sort.Slice(memberIDs, func(i, j int) bool { return memberIDs[i] < memberIDs[j] }) + // Concatentating members fore membersConcat := "" for i := range memberIDs { // Don't panic, we're adding , in the SQL as well membersConcat += strconv.FormatUint(memberIDs[i], 10) + "," } - return mod, isFound(r.db().Get(mod, sqlChannelGroupByMemberSet, types.ChannelTypeGroup, len(memberIDs), membersConcat), mod.ID > 0, ErrChannelNotFound) + return r.findOneBy( + squirrel.And{ + squirrel.Eq{"type": types.ChannelTypeGroup}, + squirrel. + Select("rel_channel"). + From(r.tableMember()). + GroupBy("rel_channel"). + Having(squirrel.Eq{ + "COUNT(*)": len(memberIDs), + "CONCAT(GROUP_CONCAT(rel_user ORDER BY 1 ASC SEPARATOR ','),',')": membersConcat, + }), + }) } -func (r *channel) Find(filter *types.ChannelFilter) (types.ChannelSet, error) { - // @todo: actual searching (filter.Query) not just a full select +func (r channel) findOneBy(cnd squirrel.Sqlizer) (*types.Channel, error) { + var ( + app = &types.Channel{} - params := make([]interface{}, 0) - rval := types.ChannelSet{} + q = r.query(). + Where(cnd) - sql := sqlChannelSelect + err = rh.FetchOne(r.db(), q, app) + ) - if filter != nil { - if filter.Query != "" { - sql += " AND c.name LIKE ?" - params = append(params, filter.Query+"%") - } - - if filter.CurrentUserID > 0 { - sql += " AND c.id IN " + sqlChannelAccess - params = append(params, filter.CurrentUserID, types.ChannelTypePublic) - } - - if !filter.IncludeDeleted { - sql += " AND deleted_at IS NULL" - } + if err != nil { + return nil, err + } else if app.ID == 0 { + return nil, ErrChannelNotFound } - sql += " ORDER BY c.name ASC" - - return rval, r.db().Select(&rval, sql, params...) + return app, nil } -func (r *channel) Create(mod *types.Channel) (*types.Channel, error) { +func (r channel) Find(filter types.ChannelFilter) (set types.ChannelSet, f types.ChannelFilter, err error) { + f = filter + + if f.Sort == "" { + f.Sort = "c.name ASC" + } + + query := r.query() + + if !f.IncludeDeleted { + query = query.Where(squirrel.Eq{"c.deleted_at": nil}) + } + + if f.Query != "" { + q := "%" + strings.ToLower(f.Query) + "%" + query = query.Where(squirrel.Like{"LOWER(name)": q}) + } + + if f.CurrentUserID > 0 { + query = query.Where(squirrel.Or{ + squirrel.Eq{"c.type": types.ChannelTypePublic}, + squirrel.ConcatExpr("c.id IN (", squirrel. + Select("rel_channel"). + From(r.tableMember()). + Where(squirrel.Eq{"rel_user": f.CurrentUserID}), ")"), + }) + } + + 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 channel) Create(mod *types.Channel) (*types.Channel, error) { mod.ID = factory.Sonyflake.NextID() rh.SetCurrentTimeRounded(&mod.CreatedAt) @@ -156,7 +192,7 @@ func (r *channel) Create(mod *types.Channel) (*types.Channel, error) { return mod, r.db().Insert("messaging_channel", mod) } -func (r *channel) Update(mod *types.Channel) (*types.Channel, error) { +func (r channel) Update(mod *types.Channel) (*types.Channel, error) { rh.SetCurrentTimeRounded(&mod.UpdatedAt) if mod.Type == "" { @@ -168,27 +204,27 @@ func (r *channel) Update(mod *types.Channel) (*types.Channel, error) { return mod, r.db().UpdatePartial("messaging_channel", mod, whitelist, "id") } -func (r *channel) ArchiveByID(id uint64) error { - return r.updateColumnByID("messaging_channel", "archived_at", time.Now(), id) +func (r channel) ArchiveByID(id uint64) error { + return r.updateColumnByID(r.table(), "archived_at", time.Now(), id) } -func (r *channel) UnarchiveByID(id uint64) error { - return r.updateColumnByID("messaging_channel", "archived_at", nil, id) +func (r channel) UnarchiveByID(id uint64) error { + return r.updateColumnByID(r.table(), "archived_at", nil, id) } -func (r *channel) DeleteByID(id uint64) error { - return r.updateColumnByID("messaging_channel", "deleted_at", time.Now(), id) +func (r channel) DeleteByID(id uint64) error { + return r.updateColumnByID(r.table(), "deleted_at", time.Now(), id) } -func (r *channel) UndeleteByID(id uint64) error { - return r.updateColumnByID("messaging_channel", "deleted_at", nil, id) +func (r channel) UndeleteByID(id uint64) error { + return r.updateColumnByID(r.table(), "deleted_at", nil, id) } -func (r *channel) CountCreated(userID uint64) (c int, err error) { - return c, r.db().Get(&c, "SELECT COUNT(*) FROM messaging_channel WHERE rel_creator = ?", userID) +func (r channel) CountCreated(userID uint64) (c int, err error) { + return c, r.db().Get(&c, "SELECT COUNT(*) FROM "+r.table()+" WHERE rel_creator = ?", userID) } -func (r *channel) ChangeCreator(userID, target uint64) error { - _, err := r.db().Exec("UPDATE messaging_channel SET rel_creator = ? WHERE rel_creator = ?", target, userID) +func (r channel) ChangeCreator(userID, target uint64) error { + _, err := r.db().Exec("UPDATE "+r.table()+" SET rel_creator = ? WHERE rel_creator = ?", target, userID) return err } diff --git a/messaging/repository/channel_member.go b/messaging/repository/channel_member.go index ea55fe909..b74be3ab0 100644 --- a/messaging/repository/channel_member.go +++ b/messaging/repository/channel_member.go @@ -30,8 +30,18 @@ type ( ) const ( - // Copy definitions to make it more obvious that we're reusing channel-scope sql - sqlChannelMemberChannelAccess = sqlChannelAccess + // subquery that filters out all channels that current user has access to as a member + // or via channel type (public channels) + sqlChannelAccess = ` ( + SELECT id + FROM messaging_channel c + LEFT OUTER JOIN messaging_channel_member AS m ON (c.id = m.rel_channel) + WHERE rel_user = ? + UNION + SELECT id + FROM messaging_channel c + WHERE c.type = ? + )` // Fetching channel members of all channels a specific user has access to sqlChannelMemberSelect = `SELECT m.* @@ -66,7 +76,7 @@ func (r *channelMember) Find(filter *types.ChannelMemberFilter) (types.ChannelMe if filter != nil { if filter.ComembersOf > 0 { // scope: only channel we have access to - sql += " AND m.rel_channel IN " + sqlChannelMemberChannelAccess + sql += " AND m.rel_channel IN " + sqlChannelAccess params = append(params, filter.ComembersOf, types.ChannelTypePublic) } diff --git a/messaging/rest/channel.go b/messaging/rest/channel.go index 472dd1246..d9b5bd66a 100644 --- a/messaging/rest/channel.go +++ b/messaging/rest/channel.go @@ -95,7 +95,7 @@ func (ctrl *Channel) Read(ctx context.Context, r *request.ChannelRead) (interfac } func (ctrl *Channel) List(ctx context.Context, r *request.ChannelList) (interface{}, error) { - return ctrl.wrapSet(ctrl.svc.ch.With(ctx).Find(&types.ChannelFilter{Query: r.Query})) + return ctrl.wrapSet(ctrl.svc.ch.With(ctx).Find(types.ChannelFilter{Query: r.Query})) } func (ctrl *Channel) Members(ctx context.Context, r *request.ChannelMembers) (interface{}, error) { @@ -148,7 +148,7 @@ func (ctrl *Channel) wrap(channel *types.Channel, err error) (*outgoing.Channel, } } -func (ctrl *Channel) wrapSet(cc types.ChannelSet, err error) (*outgoing.ChannelSet, error) { +func (ctrl *Channel) wrapSet(cc types.ChannelSet, f types.ChannelFilter, err error) (*outgoing.ChannelSet, error) { if err != nil { return nil, err } else { diff --git a/messaging/service/channel.go b/messaging/service/channel.go index 53edbed5d..f3a14fa37 100644 --- a/messaging/service/channel.go +++ b/messaging/service/channel.go @@ -54,7 +54,7 @@ type ( With(ctx context.Context) ChannelService FindByID(channelID uint64) (*types.Channel, error) - Find(filter *types.ChannelFilter) (types.ChannelSet, error) + Find(types.ChannelFilter) (types.ChannelSet, types.ChannelFilter, error) Create(channel *types.Channel) (*types.Channel, error) Update(channel *types.Channel) (*types.Channel, error) @@ -125,22 +125,15 @@ func (svc *channel) findByID(ID uint64) (ch *types.Channel, err error) { return } -func (svc *channel) Find(filter *types.ChannelFilter) (cc types.ChannelSet, err error) { +func (svc *channel) Find(filter types.ChannelFilter) (set types.ChannelSet, f types.ChannelFilter, err error) { filter.CurrentUserID = auth.GetIdentityFromContext(svc.ctx).Identity() - return cc, svc.db.Transaction(func() (err error) { - if cc, err = svc.channel.Find(filter); err != nil { - return - } else if err = svc.preloadExtras(cc); err != nil { - return - } + set, f, err = svc.channel.Find(filter) + if err == nil { + err = svc.preloadExtras(set) + } - cc, err = cc.Filter(func(c *types.Channel) (b bool, e error) { - return svc.ac.CanReadChannel(svc.ctx, c), nil - }) - - return - }) + return } // preloadExtras pre-loads channel's members, views diff --git a/messaging/types/channel.go b/messaging/types/channel.go index c70bfb032..c4491e928 100644 --- a/messaging/types/channel.go +++ b/messaging/types/channel.go @@ -6,6 +6,7 @@ import ( "github.com/jmoiron/sqlx/types" "github.com/cortezaproject/corteza-server/pkg/permissions" + "github.com/cortezaproject/corteza-server/pkg/rh" ) type ( @@ -58,6 +59,11 @@ type ( // Do not filter out deleted channels IncludeDeleted bool + + Sort string `json:"sort"` + + // Standard paging fields & helpers + rh.PageFilter } ChannelMembershipPolicy string diff --git a/messaging/websocket/session.go b/messaging/websocket/session.go index 1313e4bdd..49253c0c1 100644 --- a/messaging/websocket/session.go +++ b/messaging/websocket/session.go @@ -84,7 +84,7 @@ func (sess *Session) connected() (err error) { // Push user info about all channels he has access to... // @todo filter out all muted/non-joined channels - if cc, err = sess.svc.ch.With(sess.ctx).Find(&types.ChannelFilter{}); err != nil { + if cc, _, err = sess.svc.ch.With(sess.ctx).Find(types.ChannelFilter{}); err != nil { sess.log(zap.Error(err)).Error("Could not load subscribed channels") } else { sess.log().Debug( diff --git a/messaging/websocket/session_incoming_channel.go b/messaging/websocket/session_incoming_channel.go index 187ae521b..ee9a7c16d 100644 --- a/messaging/websocket/session_incoming_channel.go +++ b/messaging/websocket/session_incoming_channel.go @@ -45,7 +45,7 @@ func (s *Session) channelPart(ctx context.Context, p *incoming.ChannelPart) erro } func (s *Session) channelList(ctx context.Context, p *incoming.Channels) error { - channels, err := s.svc.ch.With(ctx).Find(&types.ChannelFilter{}) + channels, _, err := s.svc.ch.With(ctx).Find(types.ChannelFilter{}) if err != nil { return err }