From 7e8cb4b04a10bfe32b9d1df8c632d932864d06e5 Mon Sep 17 00:00:00 2001 From: Tit Petric Date: Wed, 3 Apr 2019 14:55:59 +0200 Subject: [PATCH] fix(messaging): repository stutters --- messaging/internal/repository/channel.go | 46 ++++++++++----------- messaging/internal/repository/message.go | 31 +++++++++----- messaging/internal/repository/pubsub.go | 2 +- messaging/internal/repository/repository.go | 3 ++ messaging/internal/repository/util.go | 12 ------ 5 files changed, 46 insertions(+), 48 deletions(-) diff --git a/messaging/internal/repository/channel.go b/messaging/internal/repository/channel.go index f1130fcf2..08b9990fd 100644 --- a/messaging/internal/repository/channel.go +++ b/messaging/internal/repository/channel.go @@ -15,16 +15,17 @@ type ( ChannelRepository interface { With(ctx context.Context, db *factory.DB) ChannelRepository - FindChannelByID(id uint64) (*types.Channel, error) - FindChannelByMemberSet(memberID ...uint64) (*types.Channel, error) - FindChannels(filter *types.ChannelFilter) ([]*types.Channel, error) - CreateChannel(mod *types.Channel) (*types.Channel, error) - UpdateChannel(mod *types.Channel) (*types.Channel, error) + FindByID(id uint64) (*types.Channel, error) + FindByMemberSet(memberID ...uint64) (*types.Channel, error) + Find(filter *types.ChannelFilter) ([]*types.Channel, error) - ArchiveChannelByID(id uint64) error - UnarchiveChannelByID(id uint64) error - DeleteChannelByID(id uint64) error - UndeleteChannelByID(id uint64) error + Create(mod *types.Channel) (*types.Channel, error) + Update(mod *types.Channel) (*types.Channel, error) + + ArchiveByID(id uint64) error + UnarchiveByID(id uint64) error + DeleteByID(id uint64) error + UndeleteByID(id uint64) error CountCreated(userID uint64) (c int, err error) ChangeCreator(userID, target uint64) error @@ -87,7 +88,7 @@ func (r *channel) With(ctx context.Context, db *factory.DB) ChannelRepository { } } -func (r *channel) FindChannelByID(id uint64) (*types.Channel, error) { +func (r *channel) FindByID(id uint64) (*types.Channel, error) { mod := &types.Channel{} sql := sqlChannelSelect + " AND id = ?" @@ -95,7 +96,7 @@ func (r *channel) FindChannelByID(id uint64) (*types.Channel, error) { } // FindChannelByMemberSet searches for channel (group!) with exactly the same membership structure -func (r *channel) FindChannelByMemberSet(memberIDs ...uint64) (*types.Channel, error) { +func (r *channel) FindByMemberSet(memberIDs ...uint64) (*types.Channel, error) { mod := &types.Channel{} sort.Slice(memberIDs, func(i, j int) bool { @@ -111,7 +112,7 @@ func (r *channel) FindChannelByMemberSet(memberIDs ...uint64) (*types.Channel, e return mod, isFound(r.db().Get(mod, sqlChannelGroupByMemberSet, types.ChannelTypeGroup, len(memberIDs), membersConcat), mod.ID > 0, ErrChannelNotFound) } -func (r *channel) FindChannels(filter *types.ChannelFilter) ([]*types.Channel, error) { +func (r *channel) Find(filter *types.ChannelFilter) ([]*types.Channel, error) { // @todo: actual searching (filter.Query) not just a full select params := make([]interface{}, 0) @@ -136,11 +137,10 @@ func (r *channel) FindChannels(filter *types.ChannelFilter) ([]*types.Channel, e return rval, r.db().Select(&rval, sql, params...) } -func (r *channel) CreateChannel(mod *types.Channel) (*types.Channel, error) { +func (r *channel) Create(mod *types.Channel) (*types.Channel, error) { mod.ID = factory.Sonyflake.NextID() mod.CreatedAt = time.Now() mod.UpdatedAt = nil - mod.Meta = coalesceJson(mod.Meta, []byte("{}")) if mod.Type == "" { mod.Type = types.ChannelTypePublic @@ -149,39 +149,35 @@ func (r *channel) CreateChannel(mod *types.Channel) (*types.Channel, error) { return mod, r.db().Insert("messaging_channel", mod) } -func (r *channel) UpdateChannel(mod *types.Channel) (*types.Channel, error) { +func (r *channel) Update(mod *types.Channel) (*types.Channel, error) { mod.UpdatedAt = timeNowPtr() - mod.Meta = coalesceJson(mod.Meta, []byte("{}")) if mod.Type == "" { mod.Type = types.ChannelTypePublic } whitelist := []string{"id", "name", "type", "topic", "meta", "updated_at"} - return mod, r.db(). - UpdatePartial("messaging_channel", mod, whitelist, "id") + return mod, r.db().UpdatePartial("messaging_channel", mod, whitelist, "id") } -func (r *channel) ArchiveChannelByID(id uint64) error { +func (r *channel) ArchiveByID(id uint64) error { return r.updateColumnByID("messaging_channel", "archived_at", time.Now(), id) } -func (r *channel) UnarchiveChannelByID(id uint64) error { +func (r *channel) UnarchiveByID(id uint64) error { return r.updateColumnByID("messaging_channel", "archived_at", nil, id) } -func (r *channel) DeleteChannelByID(id uint64) error { +func (r *channel) DeleteByID(id uint64) error { return r.updateColumnByID("messaging_channel", "deleted_at", time.Now(), id) } -func (r *channel) UndeleteChannelByID(id uint64) error { +func (r *channel) UndeleteByID(id uint64) error { return r.updateColumnByID("messaging_channel", "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) + return c, r.db().Get(&c, "SELECT COUNT(*) FROM messaging_channel WHERE rel_creator = ?", userID) } func (r *channel) ChangeCreator(userID, target uint64) error { diff --git a/messaging/internal/repository/message.go b/messaging/internal/repository/message.go index 6774dc135..072ee4045 100644 --- a/messaging/internal/repository/message.go +++ b/messaging/internal/repository/message.go @@ -3,6 +3,7 @@ package repository import ( "context" "fmt" + "io" "time" "github.com/jmoiron/sqlx" @@ -15,14 +16,18 @@ type ( MessageRepository interface { With(ctx context.Context, db *factory.DB) MessageRepository - FindMessageByID(id uint64) (*types.Message, error) - FindMessages(filter *types.MessageFilter) (types.MessageSet, error) + FindByID(id uint64) (*types.Message, error) + Find(filter *types.MessageFilter) (types.MessageSet, error) FindThreads(filter *types.MessageFilter) (types.MessageSet, error) CountFromMessageID(channelID, threadID, messageID uint64) (uint32, error) PrefillThreadParticipants(mm types.MessageSet) error - CreateMessage(mod *types.Message) (*types.Message, error) - UpdateMessage(mod *types.Message) (*types.Message, error) - DeleteMessageByID(ID uint64) error + + Create(mod *types.Message) (*types.Message, error) + Update(mod *types.Message) (*types.Message, error) + DeleteByID(ID uint64) error + + BindAvatar(message *types.Message, avatar io.Reader) (*types.Message, error) + IncReplyCount(ID uint64) error DecReplyCount(ID uint64) error @@ -98,14 +103,14 @@ func (r *message) With(ctx context.Context, db *factory.DB) MessageRepository { } } -func (r *message) FindMessageByID(id uint64) (*types.Message, error) { +func (r *message) FindByID(id uint64) (*types.Message, error) { mod := &types.Message{} sql := sqlMessagesSelect + " AND id = ?" return mod, isFound(r.db().Get(mod, sql, id), mod.ID > 0, ErrMessageNotFound) } -func (r *message) FindMessages(filter *types.MessageFilter) (types.MessageSet, error) { +func (r *message) Find(filter *types.MessageFilter) (types.MessageSet, error) { r.sanitizeFilter(filter) params := make([]interface{}, 0) @@ -252,20 +257,26 @@ func (r *message) sanitizeFilter(filter *types.MessageFilter) { } } -func (r *message) CreateMessage(mod *types.Message) (*types.Message, error) { +func (r *message) Create(mod *types.Message) (*types.Message, error) { mod.ID = factory.Sonyflake.NextID() mod.CreatedAt = time.Now() return mod, r.db().Insert("messaging_message", mod) } -func (r *message) UpdateMessage(mod *types.Message) (*types.Message, error) { +func (r *message) Update(mod *types.Message) (*types.Message, error) { mod.UpdatedAt = timeNowPtr() return mod, r.db().Replace("messaging_message", mod) } -func (r *message) DeleteMessageByID(ID uint64) error { +func (svc *message) BindAvatar(in *types.Message, avatar io.Reader) (*types.Message, error) { + // @todo: implement setting avatar on a message + in.Meta.Avatar = "" + return in, nil +} + +func (r *message) DeleteByID(ID uint64) error { return r.updateColumnByID("messaging_message", "deleted_at", time.Now(), ID) } diff --git a/messaging/internal/repository/pubsub.go b/messaging/internal/repository/pubsub.go index bcf88b094..9a166cefe 100644 --- a/messaging/internal/repository/pubsub.go +++ b/messaging/internal/repository/pubsub.go @@ -40,7 +40,7 @@ func (PubSub) New() *PubSub { } // create isntances based on mode - if flags.PubSub.Mode == "redis" { + if flags != nil && flags.PubSub.Mode == "redis" { return save(PubSubRedis{}.New(flags.PubSub)) } return save(PubSubMemory{}.New(flags.PubSub)) diff --git a/messaging/internal/repository/repository.go b/messaging/internal/repository/repository.go index 1bf300546..f94974138 100644 --- a/messaging/internal/repository/repository.go +++ b/messaging/internal/repository/repository.go @@ -22,16 +22,19 @@ func DB(ctx context.Context) *factory.DB { return factory.Database.MustGet("messaging").With(ctx) } +// Identity returns the User ID from context func Identity(ctx context.Context) uint64 { return auth.GetIdentityFromContext(ctx).Identity() } +// Organisation returns the Organisation from context func Organization(ctx context.Context) *types.Organisation { return &types.Organisation{ organization.GetFromContext(ctx), } } +// With updates repository and database contexts func (r *repository) With(ctx context.Context, db *factory.DB) *repository { return &repository{ ctx: ctx, diff --git a/messaging/internal/repository/util.go b/messaging/internal/repository/util.go index deb2a3ecf..320f0f153 100644 --- a/messaging/internal/repository/util.go +++ b/messaging/internal/repository/util.go @@ -3,8 +3,6 @@ package repository import ( "fmt" "time" - - "encoding/json" ) func (r repository) updateColumnByID(tableName, columnName string, value interface{}, id uint64) (err error) { @@ -33,13 +31,3 @@ func timeNowPtr() *time.Time { n := time.Now() return &n } - -func coalesceJson(vals ...json.RawMessage) json.RawMessage { - for _, val := range vals { - if val != nil { - return val - } - } - - return nil -}