From d0ea82b683ebb657edff9434a91e1c2f9dba97f0 Mon Sep 17 00:00:00 2001 From: Denis Arh Date: Wed, 25 Jul 2018 13:20:02 +0200 Subject: [PATCH] Add WebSocket basic comm --- sam/repository/message.go | 15 ++++++- sam/rest/message.go | 2 +- sam/types/channel.go | 6 +-- sam/types/message_filter.go | 8 ++-- sam/websocket/incoming.go | 66 ++++++++++++++++++++++++++----- sam/websocket/incoming/message.go | 24 +++-------- sam/websocket/incoming/types.go | 44 +++++++-------------- sam/websocket/outgoing/message.go | 36 +++++++++++------ sam/websocket/outgoing/types.go | 17 ++++++-- sam/websocket/session.go | 20 +++++----- sam/websocket/store.go | 10 +++++ 11 files changed, 158 insertions(+), 90 deletions(-) diff --git a/sam/repository/message.go b/sam/repository/message.go index 50597de48..a5f149722 100644 --- a/sam/repository/message.go +++ b/sam/repository/message.go @@ -54,13 +54,24 @@ func (r message) Find(ctx context.Context, filter *types.MessageFilter) ([]*type params = append(params, filter.ChannelId) } - if filter.LastMessageId > 0 { + if filter.FromMessageId > 0 { sql += " AND id > ? " - params = append(params, filter.LastMessageId) + params = append(params, filter.FromMessageId) + } + + if filter.UntilMessageId > 0 { + sql += " AND id < ? " + params = append(params, filter.UntilMessageId) } sql += " ORDER BY id ASC" + if filter.Limit > 0 { + // @todo implement some kind of protection + sql += " LIMIT ? " + params = append(params, filter.Limit) + } + rval := make([]*types.Message, 0) if err := db.SelectContext(ctx, &rval, sql, params...); err != nil { return nil, errors.Wrap(err, ErrDatabaseError.String()) diff --git a/sam/rest/message.go b/sam/rest/message.go index bf6960e12..a04774a12 100644 --- a/sam/rest/message.go +++ b/sam/rest/message.go @@ -54,7 +54,7 @@ func (ctrl *Message) Create(ctx context.Context, r *server.MessageCreateRequest) func (ctrl *Message) History(ctx context.Context, r *server.MessageHistoryRequest) (interface{}, error) { return ctrl.service.Find(ctx, &types.MessageFilter{ ChannelId: r.ChannelId, - LastMessageId: r.LastMessageId, + FromMessageId: r.LastMessageId, }) } diff --git a/sam/types/channel.go b/sam/types/channel.go index 8351b9945..f2a9d3fc6 100644 --- a/sam/types/channel.go +++ b/sam/types/channel.go @@ -94,15 +94,15 @@ func (c *Channel) SetMeta(value json.RawMessage) *Channel { return c } -// Get the value of LastMessageId +// Get the value of FromMessageId func (c *Channel) GetLastMessageId() uint64 { return c.LastMessageId } -// Set the value of LastMessageId +// Set the value of FromMessageId func (c *Channel) SetLastMessageId(value uint64) *Channel { if c.LastMessageId != value { - c.changed = append(c.changed, "LastMessageId") + c.changed = append(c.changed, "FromMessageId") c.LastMessageId = value } return c diff --git a/sam/types/message_filter.go b/sam/types/message_filter.go index e6fe225d2..2688cae95 100644 --- a/sam/types/message_filter.go +++ b/sam/types/message_filter.go @@ -2,8 +2,10 @@ package types type ( MessageFilter struct { - Query string - ChannelId uint64 - LastMessageId uint64 + Query string + ChannelId uint64 + FromMessageId uint64 + UntilMessageId uint64 + Limit uint } ) diff --git a/sam/websocket/incoming.go b/sam/websocket/incoming.go index 975227e1b..780eec123 100644 --- a/sam/websocket/incoming.go +++ b/sam/websocket/incoming.go @@ -1,22 +1,70 @@ package websocket import ( + "context" "encoding/json" - "log" - + "github.com/crusttech/crust/sam/service" + "github.com/crusttech/crust/sam/types" "github.com/crusttech/crust/sam/websocket/incoming" + "github.com/crusttech/crust/sam/websocket/outgoing" "github.com/pkg/errors" + "strconv" ) -func (s *Session) dispatch(raw []byte) error { - log.Printf("%s> %s", s.remoteAddr, string(raw)) - - msg := incoming.Message{}.New() - if err := json.Unmarshal(raw, msg); err != nil { - return errors.Wrap(err, "Session.incoming: malformed json payload") +func (s *Session) dispatch(raw []byte) (err error) { + var payload = &incoming.Message{} + if err = json.Unmarshal(raw, payload); err != nil { + return errors.Wrap(err, "Session.incoming: payload malformed") } - // @todo: do stuff with msg + if p := payload.MessageCreate; p != nil { + return s.dispatchMessageCreate(p) + } + + if p := payload.ChannelOpen; p != nil { + return s.dispatchChannelOpen(p) + } return nil } + +func (s *Session) dispatchMessageCreate(payload *incoming.MessageCreate) (err error) { + var ( + msg = &types.Message{Message: payload.Message} + ) + + if msg.ChannelId, err = strconv.ParseUint(payload.ChannelId, 10, 64); err != nil { + return + } + + if msg, err = service.Message().Create(context.TODO(), msg); err != nil { + return + } else { + // @todo move this to outgoing.FromMessage(*types.Message) *outgoing.WsMessage + store.MessageFanout(outgoing.FromMessage(msg)) + } + + return +} + +func (s *Session) dispatchChannelOpen(payload *incoming.ChannelOpen) (err error) { + var filter = &types.MessageFilter{} + + if filter.ChannelId, err = strconv.ParseUint(payload.ChannelId, 10, 64); err != nil { + return + } + + if filter.FromMessageId, err = strconv.ParseUint(payload.Since, 10, 64); err != nil { + return + } + + if messages, err := service.Message().Find(context.TODO(), filter); err != nil { + return err + } else { + for _, msg := range messages { + s.send <- outgoing.FromMessage(msg) + } + } + + return +} diff --git a/sam/websocket/incoming/message.go b/sam/websocket/incoming/message.go index 2cbcd64e0..84e0a6068 100644 --- a/sam/websocket/incoming/message.go +++ b/sam/websocket/incoming/message.go @@ -5,29 +5,17 @@ import ( ) type Message struct { - // User login - Login *Login `json:"login"` - // Channel actions - Join *Join `json:"join"` - Leave *Leave `json:"leave"` + *ChannelJoin `json:"chjoin"` + *ChannelPart `json:"chpart"` // Get channel message history - History *History `json:"history"` + *ChannelOpen `json:"chopen"` // Message actions - Create *Create `json:"create"` - Edit *Edit `json:"edit"` - Delete *Delete `json:"delete"` - - // Client notifications (message received, message read, typing indicator) - Note *Note `json:"note"` + *MessageCreate `json:"msgcre"` + *MessageUpdate `json:"msgupd"` + *MessageDelete `json:"msgdel"` timestamp time.Time } - -func (Message) New() *Message { - return &Message{ - timestamp: time.Now().UTC(), - } -} \ No newline at end of file diff --git a/sam/websocket/incoming/types.go b/sam/websocket/incoming/types.go index 783ce4f2e..055d4df52 100644 --- a/sam/websocket/incoming/types.go +++ b/sam/websocket/incoming/types.go @@ -1,44 +1,30 @@ package incoming -type Login struct { - Username string `json:"username,omitempty"` - Password []byte `json:"password"` +type ChannelJoin struct { + ChannelId string `json:"cid"` } -type Join struct { - Topic string `json:"topic"` +type ChannelPart struct { + ChannelId string `json:"cid"` } -type Leave struct { - Topic string `json:"topic"` +type ChannelOpen struct { + ChannelId string `json:"cid"` + Since string `json:"since,omitempty"` + Until string `json:"until,omitempty"` } -type History struct { - Topic string `json:"topic"` - - // if 0 = last 50 messages, else where message.id < Since - Since uint64 `json:"since,omitempty"` - - // @todo: extend API (search,...) +type MessageCreate struct { + ChannelId string `json:"cid"` + Message string `json:"msg"` } -type Create struct { - Topic string `json:"topic"` - Content interface{} `json:"content"` +type MessageUpdate struct { + ID string `json:"id"` + Message string `json:"msg"` } -type Edit struct { - ID string `json:"id"` - Topic string `json:"topic"` - Content interface{} `json:"content"` -} - -type Delete struct { +type MessageDelete struct { ID string `json:"id"` Topic string `json:"topic"` } - -type Note struct { - Topic string `json:"topic"` - Event string `json:"what"` -} diff --git a/sam/websocket/outgoing/message.go b/sam/websocket/outgoing/message.go index 923924ed6..04cf77071 100644 --- a/sam/websocket/outgoing/message.go +++ b/sam/websocket/outgoing/message.go @@ -1,28 +1,38 @@ package outgoing import ( + "github.com/crusttech/crust/sam/types" + "strconv" "time" - - "github.com/titpetric/factory" ) -type Message struct { +type WsMessage struct { Error *Error `json:"error,omitempty"` - // @todo: implement outgoing message types + *Message `json:"m"` - id uint64 + // @todo: implement outgoing message types timestamp time.Time } -func (Message) New() *Message { - return &Message{ - id: factory.Sonyflake.NextID(), - timestamp: time.Now().UTC(), - } +//func (WsMessage) New() *WsMessage { +// return &WsMessage{ +// //id: factory.Sonyflake.NextID(), +// timestamp: time.Now().UTC(), +// } +//} + +func NewError(err error) *WsMessage { + return &WsMessage{Error: &Error{Message: err.Error()}} } -func (m *Message) FromError(err error) *Message { - m.Error = &Error{err.Error()} - return m +func FromMessage(msg *types.Message) *WsMessage { + return &WsMessage{Message: &Message{ + Message: msg.Message, + Id: strconv.FormatUint(msg.ID, 10), + ChannelId: strconv.FormatUint(msg.ChannelId, 10), + Type: msg.Type, + UserId: strconv.FormatUint(msg.UserId, 10), + ReplyTo: strconv.FormatUint(msg.ReplyTo, 10), + }} } diff --git a/sam/websocket/outgoing/types.go b/sam/websocket/outgoing/types.go index 254aec7fe..6d2708b09 100644 --- a/sam/websocket/outgoing/types.go +++ b/sam/websocket/outgoing/types.go @@ -1,5 +1,16 @@ package outgoing -type Error struct { - Message string `json:"message"` -} +type ( + Error struct { + Message string `json:"m"` + } + + Message struct { + Id string `json:"id"` + ChannelId string `json:"cid""` + Message string `json:"m"` + Type string `json:"t"` + ReplyTo string `json:"rid"` + UserId string `json:"uid"` + } +) diff --git a/sam/websocket/session.go b/sam/websocket/session.go index e5c3250ce..abd21dc9b 100644 --- a/sam/websocket/session.go +++ b/sam/websocket/session.go @@ -66,9 +66,10 @@ func (sess *Session) readLoop() error { if err != nil { return errors.Wrap(err, "sess.readLoop") } + if err := sess.dispatch(raw); err != nil { // @todo: log error? - sess.send <- outgoing.Message{}.New().FromError(err) + sess.send <- outgoing.NewError(err) } } } @@ -82,15 +83,15 @@ func (sess *Session) writeLoop() error { }() write := func(mt int, msg interface{}) error { - var bits []byte - if msg != nil { - bits = msg.([]byte) - } else { - // PintMessage = empty frame - bits = []byte{} - } sess.conn.SetWriteDeadline(time.Now().Add(sess.config.writeTimeout)) - return sess.conn.WriteMessage(mt, bits) + + switch msg := msg.(type) { + case *outgoing.WsMessage: + return sess.conn.WriteJSON(msg) + default: + return sess.conn.WriteMessage(mt, nil) + } + } for { @@ -100,6 +101,7 @@ func (sess *Session) writeLoop() error { // channel closed return nil } + if err := write(websocket.TextMessage, msg); err != nil { return errors.Wrap(err, "writeLoop send") } diff --git a/sam/websocket/store.go b/sam/websocket/store.go index 2108acaa9..8978966ee 100644 --- a/sam/websocket/store.go +++ b/sam/websocket/store.go @@ -1,6 +1,7 @@ package websocket import ( + "github.com/crusttech/crust/sam/websocket/outgoing" "github.com/titpetric/factory" "sync" ) @@ -42,3 +43,12 @@ func (s *Store) Delete(id uint64) { defer s.Unlock() delete(s.Sessions, id) } + +func (s *Store) MessageFanout(messages ...*outgoing.WsMessage) { + // @todo this should probably implement some logic behind... + for _, message := range messages { + for _, sess := range s.Sessions { + sess.send <- message + } + } +}