3
0

Add basic support for unread messages

This commit is contained in:
Denis Arh
2018-10-13 18:02:21 +02:00
parent b2fcd83f1f
commit 9f419c2d47
14 changed files with 300 additions and 28 deletions

View File

@@ -11,6 +11,11 @@ type (
ChannelID string `json:"id"`
}
ChannelViewRecord struct {
ChannelID string `json:"channelID"`
LastMessageID string `json:"lastMessageID"`
}
ChannelCreate struct {
Name *string `json:"name"`
Topic *string `json:"topic"`

View File

@@ -10,6 +10,8 @@ type Payload struct {
*ChannelUpdate `json:"updateChannel"`
*ChannelDelete `json:"deleteChannel"`
*ChannelViewRecord `json:"recordChannelView"`
// Get channel message history
*Messages `json:"messages"`

View File

@@ -48,6 +48,7 @@ func Channel(ch *sam.Channel) *outgoing.Channel {
Topic: ch.Topic,
Type: string(ch.Type),
Members: Uint64stoa(ch.Members),
View: ChannelView(ch.View),
CreatedAt: ch.CreatedAt,
UpdatedAt: ch.UpdatedAt,
@@ -83,6 +84,17 @@ func ChannelMembers(members sam.ChannelMemberSet) *outgoing.ChannelMemberSet {
return &retval
}
func ChannelView(v *sam.ChannelView) *outgoing.ChannelView {
if v == nil {
return nil
}
return &outgoing.ChannelView{
LastMessageID: Uint64toa(v.LastMessageID),
NewMessagesCount: v.NewMessagesCount,
}
}
func User(user *auth.User) *outgoing.User {
if user == nil {
return nil

View File

@@ -32,12 +32,13 @@ type (
Channel struct {
// Channel to part (nil) for ALL channels
ID string `json:"ID"`
Name string `json:"name"`
Topic string `json:"topic"`
Type string `json:"type"`
LastMessageID string `json:"lastMessageID"`
Members []string `json:"members,omitempty"`
ID string `json:"ID"`
Name string `json:"name"`
Topic string `json:"topic"`
Type string `json:"type"`
LastMessageID string `json:"lastMessageID"`
Members []string `json:"members,omitempty"`
View *ChannelView `json:"view,omitempty"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt *time.Time `json:"updatedAt,omitempty"`

View File

@@ -0,0 +1,9 @@
package outgoing
type (
ChannelView struct {
// Channel to part (nil) for ALL channels
LastMessageID string `json:"lastMessageID"`
NewMessagesCount uint32 `json:"newMessagesCount"`
}
)

File diff suppressed because one or more lines are too long

View File

@@ -0,0 +1,20 @@
ALTER TABLE channel_views DROP viewed_at;
ALTER TABLE channel_views ADD rel_last_message_id BIGINT UNSIGNED;
ALTER TABLE channel_views RENAME COLUMN new_since TO new_messages_count;
-- Table structure after these changes:
-- +---------------------+---------------------+------+-----+---------+-------+
-- | Field | Type | Null | Key | Default | Extra |
-- +---------------------+---------------------+------+-----+---------+-------+
-- | rel_channel | bigint(20) unsigned | NO | PRI | NULL | |
-- | rel_user | bigint(20) unsigned | NO | PRI | NULL | |
-- | rel_last_message_id | bigint(20) unsigned | YES | | NULL | |
-- | new_messages_count | int(10) unsigned | NO | | 0 | |
-- +---------------------+---------------------+------+-----+---------+-------+
-- Prefill with data
INSERT INTO channel_views (rel_channel, rel_user, rel_last_message_id)
SELECT cm.rel_channel, cm.rel_user, max(m.ID)
FROM channel_members AS cm INNER JOIN messages AS m ON (m.rel_channel = cm.rel_channel)
GROUP BY cm.rel_channel, cm.rel_user;

View File

@@ -0,0 +1,96 @@
package repository
import (
"context"
"github.com/titpetric/factory"
"github.com/crusttech/crust/sam/types"
)
type (
// ChannelViewRepository interface to channel member repository
ChannelViewRepository interface {
With(ctx context.Context, db *factory.DB) ChannelViewRepository
Find(filter *types.ChannelViewFilter) (types.ChannelViewSet, error)
Record(channelID, userID, lastMessageID uint64, count uint32) error
Inc(channelID, userID uint64) error
Dec(channelID, userID uint64) error
}
channelViews struct {
*repository
}
)
const (
// Fetching channel members of all channels a specific user has access to
sqlChannelViewsSelect = `SELECT rel_channel, rel_user, new_messages_count, rel_last_message_id
FROM channel_views
WHERE true `
sqlChannelViewsIncCount = `UPDATE channel_views
SET new_messages_count = new_messages_count + 1
WHERE rel_channel = ? AND rel_user <> ?`
sqlChannelViewsDecCount = `UPDATE channel_views
SET new_messages_count = new_messages_count - 1
WHERE rel_channel = ? AND rel_user <> ?`
)
// ChannelView creates new instance of channel member repository
func ChannelView(ctx context.Context, db *factory.DB) ChannelViewRepository {
return (&channelViews{}).With(ctx, db)
}
// With context...
func (r *channelViews) With(ctx context.Context, db *factory.DB) ChannelViewRepository {
return &channelViews{
repository: r.repository.With(ctx, db),
}
}
// FindMembers fetches membership info
//
// If channelID > 0 it returns members of a specific channel
// If userID > 0 it returns members of all channels this user is member of
func (r *channelViews) Find(filter *types.ChannelViewFilter) (types.ChannelViewSet, error) {
params := make([]interface{}, 0)
vv := types.ChannelViewSet{}
sql := sqlChannelViewsSelect
if filter != nil {
if filter.UserID > 0 {
// scope: only channel we have access to
sql += ` AND rel_user = ?`
params = append(params, filter.UserID)
}
}
return vv, r.db().Select(&vv, sql, params...)
}
// Records channel view
func (r *channelViews) Record(channelID, userID, lastMessageID uint64, count uint32) error {
mod := &types.ChannelView{
ChannelID: channelID,
UserID: userID,
LastMessageID: lastMessageID,
NewMessagesCount: count,
}
return r.db().Replace("channel_views", mod)
}
// Increments unread (new) message count on a channel for all but one user
func (r *channelViews) Inc(channelID, userID uint64) error {
_, err := r.db().Exec(sqlChannelViewsIncCount, channelID, userID)
return err
}
// Increments unread (new) message count on a channel for all but one user
func (r *channelViews) Dec(channelID, userID uint64) error {
_, err := r.db().Exec(sqlChannelViewsDecCount, channelID, userID)
return err
}

View File

@@ -4,6 +4,7 @@ import (
"context"
"fmt"
"github.com/davecgh/go-spew/spew"
"github.com/pkg/errors"
"github.com/titpetric/factory"
@@ -24,6 +25,7 @@ type (
channel repository.ChannelRepository
cmember repository.ChannelMemberRepository
cview repository.ChannelViewRepository
message repository.MessageRepository
sysmsgs types.MessageSet
@@ -48,6 +50,7 @@ type (
Archive(ID uint64) error
Unarchive(ID uint64) error
Delete(ID uint64) error
RecordView(channelID, userID, lastMessageID uint64) error
}
// channelSecurity interface {
@@ -73,6 +76,7 @@ func (svc *channel) With(ctx context.Context) ChannelService {
channel: repository.Channel(ctx, db),
cmember: repository.ChannelMember(ctx, db),
cview: repository.ChannelView(ctx, db),
message: repository.Message(ctx, db),
// System messages should be flushed at the end of each session
@@ -93,14 +97,33 @@ func (svc *channel) FindByID(id uint64) (ch *types.Channel, err error) {
return
}
func (svc *channel) Find(filter *types.ChannelFilter) (types.ChannelSet, error) {
func (svc *channel) Find(filter *types.ChannelFilter) (cc types.ChannelSet, err error) {
filter.CurrentUserID = auth.GetIdentityFromContext(svc.ctx).Identity()
if cc, err := svc.channel.FindChannels(filter); err != nil {
return nil, err
} else {
return cc, svc.preloadMembers(cc)
return cc, svc.db.Transaction(func() (err error) {
if cc, err = svc.channel.FindChannels(filter); err != nil {
return
} else if err = svc.preloadExtras(cc); err != nil {
return
}
return
})
}
// preloadExtras pre-loads channel's members, views
func (svc *channel) preloadExtras(cc types.ChannelSet) (err error) {
if err = svc.preloadMembers(cc); err != nil {
return
}
if err = svc.preloadViews(cc); err != nil {
return
}
spew.Dump(cc.FindById(55955117148471560))
return
}
func (svc *channel) preloadMembers(cc types.ChannelSet) error {
@@ -118,6 +141,21 @@ func (svc *channel) preloadMembers(cc types.ChannelSet) error {
return nil
}
func (svc *channel) preloadViews(cc types.ChannelSet) error {
var userID = auth.GetIdentityFromContext(svc.ctx).Identity()
if vv, err := svc.cview.Find(&types.ChannelViewFilter{UserID: userID}); err != nil {
return err
} else {
cc.Walk(func(ch *types.Channel) error {
ch.View = vv.FindByChannelId(ch.ID)
return nil
})
}
return nil
}
// FindMembers loads all members (and full users) for a specific channel
func (svc *channel) FindMembers(channelID uint64) (out types.ChannelMemberSet, err error) {
var userID = auth.GetIdentityFromContext(svc.ctx).Identity()
@@ -144,8 +182,8 @@ func (svc *channel) FindMembers(channelID uint64) (out types.ChannelMemberSet, e
}
// Returns all channels with membership info
func (svc *channel) FindByMembership() (rval []*types.Channel, err error) {
return rval, svc.db.Transaction(func() error {
func (svc *channel) FindByMembership() (cc []*types.Channel, err error) {
return cc, svc.db.Transaction(func() error {
var chMemberId = repository.Identity(svc.ctx)
var mm []*types.ChannelMember
@@ -154,12 +192,14 @@ func (svc *channel) FindByMembership() (rval []*types.Channel, err error) {
return err
}
if rval, err = svc.channel.FindChannels(nil); err != nil {
if cc, err = svc.channel.FindChannels(nil); err != nil {
return err
} else if err = svc.preloadExtras(cc); err != nil {
return err
}
for _, m := range mm {
for _, c := range rval {
for _, c := range cc {
if c.ID == m.ChannelID {
c.Member = m
}
@@ -592,6 +632,12 @@ func (svc *channel) DeleteMember(channelID uint64, memberIDs ...uint64) (err err
})
}
func (svc *channel) RecordView(channelID, userID, lastMessageID uint64) error {
return svc.db.Transaction(func() (err error) {
return svc.cview.Record(channelID, userID, lastMessageID, 0)
})
}
func (svc *channel) scheduleSystemMessage(ch *types.Channel, format string, a ...interface{}) {
svc.sysmsgs = append(svc.sysmsgs, &types.Message{
ChannelID: ch.ID,

View File

@@ -19,6 +19,7 @@ type (
attachment repository.AttachmentRepository
channel repository.ChannelRepository
cmember repository.ChannelMemberRepository
cview repository.ChannelViewRepository
message repository.MessageRepository
reaction repository.ReactionRepository
@@ -68,6 +69,7 @@ func (svc *message) With(ctx context.Context) MessageService {
attachment: repository.Attachment(ctx, db),
channel: repository.Channel(ctx, db),
cmember: repository.ChannelMember(ctx, db),
cview: repository.ChannelView(ctx, db),
message: repository.Message(ctx, db),
reaction: repository.Reaction(ctx, db),
}
@@ -169,11 +171,15 @@ func (svc *message) Direct(recipientID uint64, in *types.Message) (out *types.Me
return
}
if err = svc.cview.Inc(in.ChannelID, in.UserID); err != nil {
return err
}
return svc.sendEvent(out)
})
}
func (svc *message) Create(mod *types.Message) (*types.Message, error) {
func (svc *message) Create(mod *types.Message) (message *types.Message, err error) {
// @todo get user from context
var currentUserID uint64 = repository.Identity(svc.ctx)
@@ -181,13 +187,17 @@ func (svc *message) Create(mod *types.Message) (*types.Message, error) {
mod.UserID = currentUserID
message, err := svc.message.CreateMessage(mod)
return message, svc.db.Transaction(func() (err error) {
if message, err = svc.message.CreateMessage(mod); err != nil {
return err
}
if err != nil {
return nil, err
}
if err = svc.cview.Inc(message.ChannelID, message.UserID); err != nil {
return err
}
return message, svc.sendEvent(message)
return svc.sendEvent(message)
})
}
func (svc *message) Update(mod *types.Message) (*types.Message, error) {
@@ -220,13 +230,21 @@ func (svc *message) Delete(id uint64) error {
// @todo load current message
// @todo verify ownership
err := svc.message.DeleteMessageByID(id)
return svc.db.Transaction(func() (err error) {
if err = svc.message.DeleteMessageByID(id); err != nil {
return
}
//if err == nil {
// err = svc.evl.Message(message)
//}
// if err == nil {
// err = svc.evl.Message(message)
// }
return err
// if err = svc.cview.Dec(message.ChannelID, message.UserID); err != nil {
// return err
// }
return nil
})
}
func (svc *message) React(messageID uint64, reaction string) error {

View File

@@ -26,6 +26,7 @@ type (
Member *ChannelMember `json:"-" db:"-"`
Members []uint64 `json:"-" db:"-"`
View *ChannelView `json:"-" db:"-"`
}
ChannelFilter struct {
@@ -67,6 +68,16 @@ func (cc ChannelSet) Walk(w func(*Channel) error) (err error) {
return
}
func (cc ChannelSet) FindById(ID uint64) *Channel {
for i := range cc {
if cc[i].ID == ID {
return cc[i]
}
}
return nil
}
const (
ChannelTypePublic ChannelType = "public"
ChannelTypePrivate = "private"

37
sam/types/channel_view.go Normal file
View File

@@ -0,0 +1,37 @@
package types
type (
ChannelView struct {
ChannelID uint64 `db:"rel_channel"`
UserID uint64 `db:"rel_user"`
LastMessageID uint64 `db:"rel_last_message_id"`
NewMessagesCount uint32 `db:"new_messages_count"`
}
ChannelViewFilter struct {
UserID uint64
}
ChannelViewSet []*ChannelView
)
func (mm ChannelViewSet) Walk(w func(*ChannelView) error) (err error) {
for i := range mm {
if err = w(mm[i]); err != nil {
return
}
}
return
}
func (uu ChannelViewSet) FindByChannelId(channelID uint64) *ChannelView {
for i := range uu {
if uu[i].ChannelID == channelID {
return uu[i]
}
}
return nil
}

View File

@@ -14,7 +14,6 @@ func (s *Session) dispatch(raw []byte) error {
ctx := s.Context()
switch {
// message actions
case p.MessageCreate != nil:
return s.messageCreate(ctx, p.MessageCreate)
@@ -38,6 +37,8 @@ func (s *Session) dispatch(raw []byte) error {
return s.channelDelete(ctx, p.ChannelDelete)
case p.ChannelUpdate != nil:
return s.channelUpdate(ctx, p.ChannelUpdate)
case p.ChannelViewRecord != nil:
return s.channelViewRecord(ctx, p.ChannelViewRecord)
case p.Users != nil:
return s.userList(ctx, p.Users)

View File

@@ -105,3 +105,17 @@ func (s *Session) channelUpdate(ctx context.Context, p *incoming.ChannelUpdate)
_, err = s.svc.ch.With(ctx).Update(ch)
return err
}
func (s *Session) channelViewRecord(ctx context.Context, p *incoming.ChannelViewRecord) error {
var (
channelID = payload.ParseUInt64(p.ChannelID)
lastMessageID = payload.ParseUInt64(p.LastMessageID)
userID = auth.GetIdentityFromContext(ctx).Identity()
)
if channelID == 0 || lastMessageID == 0 {
return nil
}
return s.svc.ch.With(ctx).RecordView(channelID, userID, lastMessageID)
}