From deddcb1c5d03115dd22943219c1ee0f2abf5b1e0 Mon Sep 17 00:00:00 2001 From: Denis Arh Date: Tue, 30 Oct 2018 09:48:39 +0100 Subject: [PATCH] Message length validation --- sam/service/message.go | 63 ++++++++++++++++++++++++++++--------- sam/service/message_test.go | 45 ++++++++++++++++++++++++++ 2 files changed, 94 insertions(+), 14 deletions(-) create mode 100644 sam/service/message_test.go diff --git a/sam/service/message.go b/sam/service/message.go index a43b91e9c..fd227e702 100644 --- a/sam/service/message.go +++ b/sam/service/message.go @@ -2,10 +2,10 @@ package service import ( "context" + "strings" "time" "github.com/pkg/errors" - "github.com/titpetric/factory" authService "github.com/crusttech/crust/auth/service" authTypes "github.com/crusttech/crust/auth/types" @@ -16,7 +16,7 @@ import ( type ( message struct { - db *factory.DB + db db ctx context.Context attachment repository.AttachmentRepository @@ -52,6 +52,10 @@ type ( } ) +const ( + settingsMessageBodyLength = 4000 +) + func Message() MessageService { return &message{ usr: authService.DefaultUser, @@ -111,23 +115,36 @@ func (svc *message) FindThreads(filter *types.MessageFilter) (mm types.MessageSe return mm, svc.preloadAttachments(mm) } -func (svc *message) Create(mod *types.Message) (message *types.Message, err error) { +func (svc *message) Create(in *types.Message) (message *types.Message, err error) { + if in == nil { + in = &types.Message{} + } + + in.Message = strings.TrimSpace(in.Message) + var mlen = len(in.Message) + + if mlen == 0 { + return nil, errors.Errorf("Refusing to create message without contents") + } else if mlen > settingsMessageBodyLength { + return nil, errors.Errorf("Message length (%d characters) too long (max: %d)", mlen, settingsMessageBodyLength) + } + // @todo get user from context var currentUserID uint64 = repository.Identity(svc.ctx) - mod.UserID = currentUserID + in.UserID = currentUserID return message, svc.db.Transaction(func() (err error) { // Broadcast queue var bq = types.MessageSet{} - if mod.ReplyTo > 0 { + if in.ReplyTo > 0 { var original *types.Message - var replyTo = mod.ReplyTo + var replyTo = in.ReplyTo for replyTo > 0 { // Find original message - original, err = svc.message.FindMessageByID(mod.ReplyTo) + original, err = svc.message.FindMessageByID(in.ReplyTo) if err != nil { return } @@ -141,9 +158,9 @@ func (svc *message) Create(mod *types.Message) (message *types.Message, err erro // We do not want to have multi-level threads // Take original's reply-to and use it - mod.ReplyTo = original.ID + in.ReplyTo = original.ID - mod.ChannelID = original.ChannelID + in.ChannelID = original.ChannelID // Increment counter, on struct and in repostiry. original.Replies++ @@ -155,13 +172,13 @@ func (svc *message) Create(mod *types.Message) (message *types.Message, err erro bq = append(bq, original) } - if mod.ChannelID == 0 { + if in.ChannelID == 0 { return errors.New("ChannelID missing") } // @todo [SECURITY] verify if current user can access & write to this channel - if message, err = svc.message.CreateMessage(mod); err != nil { + if message, err = svc.message.CreateMessage(in); err != nil { return } @@ -173,7 +190,20 @@ func (svc *message) Create(mod *types.Message) (message *types.Message, err erro }) } -func (svc *message) Update(mod *types.Message) (message *types.Message, err error) { +func (svc *message) Update(in *types.Message) (message *types.Message, err error) { + if in == nil { + in = &types.Message{} + } + + in.Message = strings.TrimSpace(in.Message) + var mlen = len(in.Message) + + if mlen == 0 { + return nil, errors.Errorf("Refusing to update message without contents") + } else if mlen > settingsMessageBodyLength { + return nil, errors.Errorf("Message length (%d characters) too long (max: %d)", mlen, settingsMessageBodyLength) + } + // @todo get user from context var currentUserID uint64 = repository.Identity(svc.ctx) @@ -181,17 +211,22 @@ func (svc *message) Update(mod *types.Message) (message *types.Message, err erro _ = currentUserID return message, svc.db.Transaction(func() (err error) { - original, err := svc.message.FindMessageByID(mod.ID) + original, err := svc.message.FindMessageByID(in.ID) if err != nil { return err } + if original.Message == in.Message { + // Nothing changed + return nil + } + if original.UserID != currentUserID { return errors.New("Not an owner") } // Allow message content to be changed, ignore everything else - original.Message = mod.Message + original.Message = in.Message if message, err = svc.message.UpdateMessage(original); err != nil { return err diff --git a/sam/service/message_test.go b/sam/service/message_test.go new file mode 100644 index 000000000..56541a19d --- /dev/null +++ b/sam/service/message_test.go @@ -0,0 +1,45 @@ +package service + +import ( + "context" + "strings" + "testing" + + authTypes "github.com/crusttech/crust/auth/types" + "github.com/crusttech/crust/internal/auth" + "github.com/crusttech/crust/sam/types" +) + +// func TestChannelCreation(t *testing.T) { +// mockCtrl := gomock.NewController(t) +// defer mockCtrl.Finish() +// +// chRpoMock := NewMockRepository(mockCtrl) +// chRpoMock.EXPECT().WithCtx(gomock.Any()).AnyTimes().Return(chRpoMock) +// chRpoMock.EXPECT(). +// FindUserByID(usr.ID). +// Times(1). +// Return(usr, nil) +// +// svc := channel{ +// channel: +// } +// +// svc.Create() +// } + +func TesMessageLength(t *testing.T) { + // mockCtrl := gomock.NewController(t) + // defer mockCtrl.Finish() + + ctx := context.TODO() + auth.SetIdentityToContext(ctx, &authTypes.User{}) + + svc := message{db: &mockDB{}, ctx: ctx} + e := func(out *types.Message, err error) error { return err } + + longText := strings.Repeat("X", settingsMessageBodyLength+1) + + assert(t, e(svc.Create(&types.Message{})) != nil, "Should not allow to create unnamed channels") + assert(t, e(svc.Create(&types.Message{Message: longText})) != nil, "Should not allow to create channel with really long name") +}