Message length validation
This commit is contained in:
+49
-14
@@ -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
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
Reference in New Issue
Block a user