Refactor channel repo
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
+122
-86
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user