From 76d93e3a763d648cf96c827b8b966604b674703e Mon Sep 17 00:00:00 2001 From: Tit Petric Date: Wed, 25 Jul 2018 19:28:56 +0200 Subject: [PATCH] upd(sam/websocket): general cleanup re: dispatch, prepare byte send to not re-encode msgs --- sam/websocket/incoming.go | 42 +++++++++++++----------------- sam/websocket/incoming/message.go | 8 +----- sam/websocket/session.go | 2 ++ sam/websocket/util.go | 43 +++++++++++++++++++++++++++++++ 4 files changed, 64 insertions(+), 31 deletions(-) create mode 100644 sam/websocket/util.go diff --git a/sam/websocket/incoming.go b/sam/websocket/incoming.go index e586d76dd..28f3ec583 100644 --- a/sam/websocket/incoming.go +++ b/sam/websocket/incoming.go @@ -7,36 +7,32 @@ import ( "github.com/crusttech/crust/sam/types" "github.com/crusttech/crust/sam/websocket/incoming" "github.com/pkg/errors" - "strconv" ) func (s *Session) dispatch(raw []byte) (err error) { - var payload = &incoming.Message{} - if err = json.Unmarshal(raw, payload); err != nil { + var p = &incoming.Payload{} + if err = json.Unmarshal(raw, p); err != nil { return errors.Wrap(err, "Session.incoming: payload malformed") } ctx := s.Context() - if p := payload.MessageCreate; p != nil { - return s.dispatchMessageCreate(ctx, p) - } - - if p := payload.ChannelOpen; p != nil { - return s.dispatchChannelOpen(ctx, p) + switch { + case p.MessageCreate != nil: + return s.messageCreate(ctx, p) + case p.ChannelOpen != nil: + return s.channelOpen(ctx, p) } return nil } -func (s *Session) dispatchMessageCreate(ctx context.Context, payload *incoming.MessageCreate) (err error) { +func (s *Session) messageCreate(ctx context.Context, payload *incoming.Payload) (err error) { var ( - msg = &types.Message{Message: payload.Message} + request = payload.MessageCreate + msg = &types.Message{Message: request.Message} ) - if msg.ChannelID, err = strconv.ParseUint(payload.ChannelID, 10, 64); err != nil { - return - } - + msg.ChannelID = parseUInt64(request.ChannelID) if msg, err = service.Message().Create(ctx, msg); err != nil { return } else { @@ -47,16 +43,14 @@ func (s *Session) dispatchMessageCreate(ctx context.Context, payload *incoming.M return } -func (s *Session) dispatchChannelOpen(ctx context.Context, payload *incoming.ChannelOpen) (err error) { - var filter = &types.MessageFilter{} +func (s *Session) channelOpen(ctx context.Context, payload *incoming.Payload) (err error) { + var ( + request = payload.ChannelOpen + 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 - } + filter.ChannelID = parseUInt64(request.ChannelID) + filter.FromMessageID = parseUInt64(request.Since) if messages, err := service.Message().Find(ctx, filter); err != nil { return err diff --git a/sam/websocket/incoming/message.go b/sam/websocket/incoming/message.go index 84e0a6068..e416feded 100644 --- a/sam/websocket/incoming/message.go +++ b/sam/websocket/incoming/message.go @@ -1,10 +1,6 @@ package incoming -import ( - "time" -) - -type Message struct { +type Payload struct { // Channel actions *ChannelJoin `json:"chjoin"` *ChannelPart `json:"chpart"` @@ -16,6 +12,4 @@ type Message struct { *MessageCreate `json:"msgcre"` *MessageUpdate `json:"msgupd"` *MessageDelete `json:"msgdel"` - - timestamp time.Time } diff --git a/sam/websocket/session.go b/sam/websocket/session.go index b3870419c..444f10b6c 100644 --- a/sam/websocket/session.go +++ b/sam/websocket/session.go @@ -95,6 +95,8 @@ func (sess *Session) writeLoop() error { switch msg := msg.(type) { case *outgoing.Payload: return sess.conn.WriteJSON(msg) + case []byte: + return sess.conn.WriteMessage(mt, msg) default: return sess.conn.WriteMessage(mt, nil) } diff --git a/sam/websocket/util.go b/sam/websocket/util.go new file mode 100644 index 000000000..435bfe172 --- /dev/null +++ b/sam/websocket/util.go @@ -0,0 +1,43 @@ +package websocket + +import ( + "regexp" + "strconv" + "strings" +) + +var truthy = regexp.MustCompile("^\\s*(t(rue)?|y(es)?|1)\\s*$") + +// parseInt64 parses an string to int64 +func parseInt64(s string) int64 { + if s == "" { + return 0 + } + i, _ := strconv.ParseInt(s, 10, 64) + + return i +} + +// parseUInt64 parses an string to uint64 +func parseUInt64(s string) uint64 { + if s == "" { + return 0 + } + i, _ := strconv.ParseUint(s, 10, 64) + return i +} + +// parseUInt64 parses an string to uint64 +func parseBool(s string) bool { + return truthy.MatchString(strings.ToLower(s)) +} + +// is checks if string s is contained in matches +func is(s string, matches ...string) bool { + for _, v := range matches { + if s == v { + return true + } + } + return false +}