From 4e215e64d3d411b0c34121740ee99d7b116d168a Mon Sep 17 00:00:00 2001 From: Denis Arh Date: Sat, 28 Jul 2018 14:32:15 +0200 Subject: [PATCH] Improve websocket request handling (value of a concret request instead of ptr to payload) Each request handler (eg channelJoin()) now takes a specific params as a value (eg incoming.ChannelJoin) --- sam/websocket/incoming.go | 14 ++++++++------ sam/websocket/incoming_channel.go | 29 ++++++++++++----------------- sam/websocket/incoming_message.go | 27 ++++++++++++--------------- 3 files changed, 32 insertions(+), 38 deletions(-) diff --git a/sam/websocket/incoming.go b/sam/websocket/incoming.go index 49d32d676..f52ebc848 100644 --- a/sam/websocket/incoming.go +++ b/sam/websocket/incoming.go @@ -17,19 +17,21 @@ func (s *Session) dispatch(raw []byte) (err error) { // message actions case p.MessageCreate != nil: - return s.messageCreate(ctx, p) + return s.messageCreate(ctx, *p.MessageCreate) case p.MessageUpdate != nil: - return s.messageUpdate(ctx, p) + return s.messageUpdate(ctx, *p.MessageUpdate) case p.MessageDelete != nil: - return s.messageDelete(ctx, p) + return s.messageDelete(ctx, *p.MessageDelete) // channel actions case p.ChannelJoin != nil: - return s.channelJoin(ctx, p) + return s.channelJoin(ctx, *p.ChannelJoin) case p.ChannelPart != nil: - return s.channelPart(ctx, p) + return s.channelPart(ctx, *p.ChannelPart) + case p.ChannelPart != nil: + return s.channelPartAll(ctx, *p.ChannelPartAll) case p.ChannelOpen != nil: - return s.channelOpen(ctx, p) + return s.channelOpen(ctx, *p.ChannelOpen) } diff --git a/sam/websocket/incoming_channel.go b/sam/websocket/incoming_channel.go index 736966278..734fa9083 100644 --- a/sam/websocket/incoming_channel.go +++ b/sam/websocket/incoming_channel.go @@ -7,37 +7,32 @@ import ( "github.com/crusttech/crust/sam/websocket/incoming" ) -func (s *Session) channelJoin(ctx context.Context, payload *incoming.Payload) error { - var ( - request = payload.ChannelJoin - ) +func (s *Session) channelJoin(ctx context.Context, p incoming.ChannelJoin) error { + var () // @todo: check access to channel - s.subs.Add(request.ChannelID, &Subscription{}) + s.subs.Add(p.ChannelID, &Subscription{}) return nil } -func (s *Session) channelPart(ctx context.Context, payload *incoming.Payload) error { - var ( - request = payload.ChannelJoin - ) +func (s *Session) channelPart(ctx context.Context, p incoming.ChannelPart) error { + var () // @todo: check access to channel - s.subs.Delete(request.ChannelID) + s.subs.Delete(p.ChannelID) return nil } -func (s *Session) channelPartAll(ctx context.Context, payload *incoming.Payload) error { - if payload.ChannelPartAll.Leave { +func (s *Session) channelPartAll(ctx context.Context, p incoming.ChannelPartAll) error { + if p.Leave { s.subs.DeleteAll() } return nil } -func (s *Session) channelOpen(ctx context.Context, payload *incoming.Payload) error { +func (s *Session) channelOpen(ctx context.Context, p incoming.ChannelOpen) error { var ( - request = payload.ChannelOpen - filter = &types.MessageFilter{ - ChannelID: parseUInt64(request.ChannelID), - FromMessageID: parseUInt64(request.Since), + filter = &types.MessageFilter{ + ChannelID: parseUInt64(p.ChannelID), + FromMessageID: parseUInt64(p.Since), } ) diff --git a/sam/websocket/incoming_message.go b/sam/websocket/incoming_message.go index c81127af3..0ac5488eb 100644 --- a/sam/websocket/incoming_message.go +++ b/sam/websocket/incoming_message.go @@ -8,12 +8,11 @@ import ( "github.com/crusttech/crust/sam/websocket/outgoing" ) -func (s *Session) messageCreate(ctx context.Context, payload *incoming.Payload) error { +func (s *Session) messageCreate(ctx context.Context, p incoming.MessageCreate) error { var ( - request = payload.MessageCreate - msg = &types.Message{ - ChannelID: parseUInt64(request.ChannelID), - Message: request.Message, + msg = &types.Message{ + ChannelID: parseUInt64(p.ChannelID), + Message: p.Message, } ) @@ -24,30 +23,28 @@ func (s *Session) messageCreate(ctx context.Context, payload *incoming.Payload) return s.sendMessageChannel(uint64toa(msg.ChannelID), payloadFromMessage(msg)) } -func (s *Session) messageUpdate(ctx context.Context, payload *incoming.Payload) error { +func (s *Session) messageUpdate(ctx context.Context, p incoming.MessageUpdate) error { var ( - request = payload.MessageUpdate - msg = &types.Message{ - ID: parseUInt64(request.ID), - Message: request.Message, + msg = &types.Message{ + ID: parseUInt64(p.ID), + Message: p.Message, } ) msg, err := service.Message().Update(ctx, msg) if err != nil { return err } - return s.sendMessageChannel(uint64toa(msg.ChannelID), &outgoing.MessageUpdate{ID: request.ID, Message: msg.Message}) + return s.sendMessageChannel(uint64toa(msg.ChannelID), &outgoing.MessageUpdate{ID: p.ID, Message: msg.Message}) } -func (s *Session) messageDelete(ctx context.Context, payload *incoming.Payload) error { +func (s *Session) messageDelete(ctx context.Context, p incoming.MessageDelete) error { var ( - request = payload.MessageDelete - id = parseUInt64(request.ID) + id = parseUInt64(p.ID) ) if err := service.Message().Delete(ctx, id); err != nil { return err } - return s.sendMessageChannel(request.ChannelID, &outgoing.MessageDelete{ID: request.ID}) + return s.sendMessageChannel(p.ChannelID, &outgoing.MessageDelete{ID: p.ID}) }