Add WebSocket basic comm

This commit is contained in:
Denis Arh
2018-07-25 13:20:02 +02:00
parent a7a5d0011f
commit d0ea82b683
11 changed files with 158 additions and 90 deletions
+13 -2
View File
@@ -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())
+1 -1
View File
@@ -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,
})
}
+3 -3
View File
@@ -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
+5 -3
View File
@@ -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
}
)
+57 -9
View File
@@ -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
}
+6 -18
View File
@@ -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(),
}
}
+15 -29
View File
@@ -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"`
}
+23 -13
View File
@@ -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),
}}
}
+14 -3
View File
@@ -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"`
}
)
+11 -9
View File
@@ -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")
}
+10
View File
@@ -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
}
}
}