From 49c4b0c4340ee3cecb9b9feccbe6a80f5e591300 Mon Sep 17 00:00:00 2001 From: Tit Petric Date: Thu, 20 Sep 2018 16:39:32 +0200 Subject: [PATCH] upd(sam): refactor repository to split responsibility, L1 --- crm/repository/db.go | 19 ------ crm/repository/repository.go | 16 ++++- sam/repository/attachment.go | 33 ++++++++--- sam/repository/channel.go | 44 +++++++++----- sam/repository/event_queue.go | 55 ------------------ sam/repository/events.go | 97 +++++++++++++++++++++++++++++++ sam/repository/events_test.go | 8 +++ sam/repository/message.go | 33 ++++++++--- sam/repository/organisation.go | 37 +++++++++--- sam/repository/reaction.go | 49 ++++++++++------ sam/repository/repository.go | 90 ++++++++-------------------- sam/repository/repository_test.go | 11 ++++ sam/repository/team.go | 41 +++++++++---- sam/repository/util.go | 3 +- sam/service/attachment.go | 63 ++++++++++---------- sam/service/channel.go | 76 ++++++++++++------------ sam/service/message.go | 39 ++++++------- sam/service/organisation.go | 9 +-- sam/service/team.go | 29 +++++---- sam/websocket/event_queue.go | 18 ++---- sam/websocket/router.go | 6 +- 21 files changed, 434 insertions(+), 342 deletions(-) delete mode 100644 crm/repository/db.go delete mode 100644 sam/repository/event_queue.go create mode 100644 sam/repository/events.go create mode 100644 sam/repository/events_test.go create mode 100644 sam/repository/repository_test.go diff --git a/crm/repository/db.go b/crm/repository/db.go deleted file mode 100644 index 828d80514..000000000 --- a/crm/repository/db.go +++ /dev/null @@ -1,19 +0,0 @@ -package repository - -import ( - "context" - "github.com/titpetric/factory" -) - -var _db *factory.DB - -func DB(ctxs ...context.Context) *factory.DB { - if _db == nil { - _db = factory.Database.MustGet() - } - for _, ctx := range ctxs { - _db = _db.With(ctx) - break - } - return _db -} diff --git a/crm/repository/repository.go b/crm/repository/repository.go index 682fd6a10..2ca484a49 100644 --- a/crm/repository/repository.go +++ b/crm/repository/repository.go @@ -14,6 +14,20 @@ type ( } ) +var _db *factory.DB + +// DB returns a repository-wide singleton DB handle +func DB(ctxs ...context.Context) *factory.DB { + if _db == nil { + _db = factory.Database.MustGet() + } + for _, ctx := range ctxs { + _db = _db.With(ctx) + break + } + return _db +} + // With updates repository and database contexts func (r *repository) With(ctx context.Context) *repository { res := &repository{ @@ -26,7 +40,7 @@ func (r *repository) With(ctx context.Context) *repository { return res } -// Return context-aware db handle +// db returns context-aware db handle func (r *repository) db() *factory.DB { return r.dbh(r.ctx) } diff --git a/sam/repository/attachment.go b/sam/repository/attachment.go index 56d2b7378..e9b080baf 100644 --- a/sam/repository/attachment.go +++ b/sam/repository/attachment.go @@ -1,20 +1,29 @@ package repository import ( - "github.com/crusttech/crust/sam/types" + "context" + "time" + "github.com/jmoiron/sqlx" "github.com/titpetric/factory" - "time" + + "github.com/crusttech/crust/sam/types" ) type ( Attachment interface { + With(ctx context.Context) Attachment + FindAttachmentByID(id uint64) (*types.Attachment, error) FindAttachmentByMessageID(IDs ...uint64) (types.MessageAttachmentSet, error) CreateAttachment(mod *types.Attachment) (*types.Attachment, error) DeleteAttachmentByID(id uint64) error BindAttachment(attachmentId, messageId uint64) error } + + attachment struct { + *repository + } ) const ( @@ -23,16 +32,24 @@ const ( ErrAttachmentNotFound = repositoryError("AttachmentNotFound") ) -var _ Attachment = &repository{} +func NewAttachment(ctx context.Context) Attachment { + return (&attachment{}).With(ctx) +} -func (r *repository) FindAttachmentByID(id uint64) (*types.Attachment, error) { +func (r *attachment) With(ctx context.Context) Attachment { + return &attachment{ + repository: r.repository.With(ctx), + } +} + +func (r *attachment) FindAttachmentByID(id uint64) (*types.Attachment, error) { sql := "SELECT * FROM attachments WHERE id = ? AND " + sqlAttachmentScope mod := &types.Attachment{} return mod, isFound(r.db().Get(mod, sql, id), mod.ID > 0, ErrAttachmentNotFound) } -func (r *repository) FindAttachmentByMessageID(IDs ...uint64) (rval types.MessageAttachmentSet, err error) { +func (r *attachment) FindAttachmentByMessageID(IDs ...uint64) (rval types.MessageAttachmentSet, err error) { rval = make([]*types.MessageAttachment, 0) if len(IDs) == 0 { @@ -51,7 +68,7 @@ func (r *repository) FindAttachmentByMessageID(IDs ...uint64) (rval types.Messag } } -func (r *repository) CreateAttachment(mod *types.Attachment) (*types.Attachment, error) { +func (r *attachment) CreateAttachment(mod *types.Attachment) (*types.Attachment, error) { if mod.ID == 0 { mod.ID = factory.Sonyflake.NextID() } @@ -61,11 +78,11 @@ func (r *repository) CreateAttachment(mod *types.Attachment) (*types.Attachment, return mod, r.db().Insert("attachments", mod) } -func (r *repository) DeleteAttachmentByID(id uint64) error { +func (r *attachment) DeleteAttachmentByID(id uint64) error { return r.updateColumnByID("attachments", "deleted_at", nil, id) } -func (r *repository) BindAttachment(attachmentId, messageId uint64) error { +func (r *attachment) BindAttachment(attachmentId, messageId uint64) error { bond := struct { RelAttachment uint64 `db:"rel_attachment"` RelMessage uint64 `db:"rel_message"` diff --git a/sam/repository/channel.go b/sam/repository/channel.go index d98415f76..7dc58ab88 100644 --- a/sam/repository/channel.go +++ b/sam/repository/channel.go @@ -1,14 +1,18 @@ package repository import ( + "context" "time" - "github.com/crusttech/crust/sam/types" "github.com/titpetric/factory" + + "github.com/crusttech/crust/sam/types" ) type ( Channel interface { + With(ctx context.Context) Channel + FindChannelByID(id uint64) (*types.Channel, error) FindDirectChannelByUserID(fromUserID, toUserID uint64) (*types.Channel, error) FindChannels(filter *types.ChannelFilter) ([]*types.Channel, error) @@ -22,6 +26,10 @@ type ( UnarchiveChannelByID(id uint64) error DeleteChannelByID(id uint64) error } + + channel struct { + *repository + } ) const ( @@ -50,14 +58,24 @@ const ( ErrChannelNotFound = repositoryError("ChannelNotFound") ) -func (r *repository) FindChannelByID(id uint64) (*types.Channel, error) { +func NewChannel(ctx context.Context) Channel { + return (&channel{}).With(ctx) +} + +func (r *channel) With(ctx context.Context) Channel { + return &channel{ + repository: r.repository.With(ctx), + } +} + +func (r *channel) FindChannelByID(id uint64) (*types.Channel, error) { mod := &types.Channel{} sql := sqlChannelSelect + " AND id = ?" return mod, isFound(r.db().Get(mod, sql, id), mod.ID > 0, ErrChannelNotFound) } -func (r *repository) FindDirectChannelByUserID(fromUserID, toUserID uint64) (*types.Channel, error) { +func (r *channel) FindDirectChannelByUserID(fromUserID, toUserID uint64) (*types.Channel, error) { mod := &types.Channel{} if fromUserID == toUserID { @@ -75,7 +93,7 @@ func (r *repository) FindDirectChannelByUserID(fromUserID, toUserID uint64) (*ty return mod, isFound(r.db().Get(mod, sqlChannelDirect, types.ChannelTypeDirect, fromUserID, toUserID), mod.ID > 0, ErrChannelNotFound) } -func (r *repository) FindChannels(filter *types.ChannelFilter) ([]*types.Channel, error) { +func (r *channel) FindChannels(filter *types.ChannelFilter) ([]*types.Channel, error) { // @todo: actual searching (filter.Query) not just a full select params := make([]interface{}, 0) @@ -95,7 +113,7 @@ func (r *repository) FindChannels(filter *types.ChannelFilter) ([]*types.Channel return rval, r.db().Select(&rval, sql, params...) } -func (r *repository) CreateChannel(mod *types.Channel) (*types.Channel, error) { +func (r *channel) CreateChannel(mod *types.Channel) (*types.Channel, error) { mod.ID = factory.Sonyflake.NextID() mod.CreatedAt = time.Now() mod.Meta = coalesceJson(mod.Meta, []byte("{}")) @@ -107,7 +125,7 @@ func (r *repository) CreateChannel(mod *types.Channel) (*types.Channel, error) { return mod, r.db().Insert("channels", mod) } -func (r *repository) UpdateChannel(mod *types.Channel) (*types.Channel, error) { +func (r *channel) UpdateChannel(mod *types.Channel) (*types.Channel, error) { mod.UpdatedAt = timeNowPtr() mod.Meta = coalesceJson(mod.Meta, []byte("{}")) if mod.Type == "" { @@ -120,36 +138,36 @@ func (r *repository) UpdateChannel(mod *types.Channel) (*types.Channel, error) { UpdatePartial("channels", mod, whitelist, "id") } -func (r *repository) FindChannelsMembershipsByMemberId(memberId uint64) ([]*types.ChannelMember, error) { +func (r *channel) FindChannelsMembershipsByMemberId(memberId uint64) ([]*types.ChannelMember, error) { var rval = make([]*types.ChannelMember, 0) return rval, r.db().Select(&rval, sqlChannelMemberships+" AND cm.rel_user = ? ", memberId) } -func (r *repository) AddChannelMember(mod *types.ChannelMember) (*types.ChannelMember, error) { +func (r *channel) AddChannelMember(mod *types.ChannelMember) (*types.ChannelMember, error) { sql := `INSERT INTO channel_members (rel_channel, rel_user) VALUES (?, ?)` mod.CreatedAt = time.Now() return mod, exec(r.db().Exec(sql, mod.ChannelID, mod.UserID)) } -func (r *repository) RemoveChannelMember(channelID, userID uint64) error { +func (r *channel) RemoveChannelMember(channelID, userID uint64) error { sql := `DELETE FROM channel_members WHERE rel_channel = ? AND rel_user = ?` return exec(r.db().Exec(sql, channelID, userID)) } -func (r *repository) ArchiveChannelByID(id uint64) error { +func (r *channel) ArchiveChannelByID(id uint64) error { return r.updateColumnByID("channels", "archived_at", time.Now(), id) } -func (r *repository) UnarchiveChannelByID(id uint64) error { +func (r *channel) UnarchiveChannelByID(id uint64) error { return r.updateColumnByID("channels", "archived_at", nil, id) } -func (r *repository) DeleteChannelByID(id uint64) error { +func (r *channel) DeleteChannelByID(id uint64) error { return r.updateColumnByID("channels", "deleted_at", time.Now(), id) } -func (r *repository) RecoverChannelByID(id uint64) error { +func (r *channel) RecoverChannelByID(id uint64) error { return r.updateColumnByID("channels", "deleted_at", nil, id) } diff --git a/sam/repository/event_queue.go b/sam/repository/event_queue.go deleted file mode 100644 index d7d4e7ad1..000000000 --- a/sam/repository/event_queue.go +++ /dev/null @@ -1,55 +0,0 @@ -package repository - -import ( - "github.com/crusttech/crust/sam/types" - "github.com/titpetric/factory" -) - -type ( - EventQueue interface { - EventQueuePull(origin uint64) ([]*types.EventQueueItem, error) - EventQueuePush(eqi *types.EventQueueItem) error - EventQueueSync(origin, id uint64) error - } -) - -func (r *repository) EventQueuePull(origin uint64) ([]*types.EventQueueItem, error) { - var ee = make([]*types.EventQueueItem, 0) - - return ee, r.db().Quiet().Select(&ee, ` - SELECT * - FROM event_queue - WHERE origin <> ? - AND id > GREATEST(COALESCE((SELECT rel_last FROM event_queue_synced WHERE origin = ?), 0), ?) - LIMIT 50`, origin, origin, origin) -} - -func (r *repository) EventQueuePush(eqi *types.EventQueueItem) error { - eqi.ID = factory.Sonyflake.NextID() - return r.db().Quiet().Insert("event_queue", eqi) -} - -func (r *repository) EventQueueSync(origin, id uint64) error { - type evqs struct { - Origin uint64 `db:"origin"` - LastEvent uint64 `db:"rel_last"` - } - - // @todo do we even need this? - return r.db().Quiet().Replace("event_queue_synced", evqs{ - Origin: origin, - LastEvent: id, - }) -} - -func (r *repository) EventQueueCleanup() error { - return exec(r.db().Exec("DELETE FROM event_queue WHERE id < (SELECT MIN(rel_last) FROM event_queue_synced)")) -} - -/* - -do we need event_queue_synced?? -do we need stable server id or can it be regenerad on each run? - - -*/ diff --git a/sam/repository/events.go b/sam/repository/events.go new file mode 100644 index 000000000..ecac31d07 --- /dev/null +++ b/sam/repository/events.go @@ -0,0 +1,97 @@ +package repository + +import ( + "context" + + "github.com/titpetric/factory" + + "github.com/crusttech/crust/sam/types" +) + +/* + +The event queue table is used for a multi-server broadcast scenario. + +If we have two servers, one channel, which have users [1,2,3] and [4,5,6], +the event queue table holds the broadcast message which should be sent +to all these users in the channel; + +The reading of the event queue table is triggered by pubsub. + +- mostly, as the servers send out all the data, the contents of the + event queue table can be discarded, +- the events queue table might eventually be not needed if we can + solve everything on the level of pubsub, (@todo) +- when a client reloads the browser, the events queue table isn't + read, everything should be in messages +- the event queue has a server id for messages originating from the + websocket; the rest api should broadcast to all websocket connected + clients, while the websocket API (currently), performs a local + broadcast, triggering the event poll only on other servers + +*/ + +type ( + Events interface { + With(ctx context.Context) Events + + Pull(origin uint64) ([]*types.EventQueueItem, error) + Push(eqi *types.EventQueueItem) error + Sync(origin, id uint64) error + } + + events struct { + *repository + } +) + +func NewEvents(ctx context.Context) Events { + return (&events{}).With(ctx) +} + +func (r *events) With(ctx context.Context) Events { + return &events{ + repository: r.repository.With(ctx), + } +} + +func (r *events) Pull(origin uint64) ([]*types.EventQueueItem, error) { + var ee = make([]*types.EventQueueItem, 0) + + return ee, r.db().Quiet().Select(&ee, ` + SELECT * + FROM event_queue + WHERE origin <> ? + AND id > GREATEST(COALESCE((SELECT rel_last FROM event_queue_synced WHERE origin = ?), 0), ?) + LIMIT 50`, origin, origin, origin) +} + +func (r *events) Push(eqi *types.EventQueueItem) error { + eqi.ID = factory.Sonyflake.NextID() + return r.db().Quiet().Insert("event_queue", eqi) +} + +func (r *events) Sync(origin, id uint64) error { + type evqs struct { + Origin uint64 `db:"origin"` + LastEvent uint64 `db:"rel_last"` + } + + // @todo do we even need this? + return r.db().Quiet().Replace("event_queue_synced", evqs{ + Origin: origin, + LastEvent: id, + }) +} + +func (r *events) Cleanup() error { + return exec(r.db().Exec("DELETE FROM event_queue WHERE id < (SELECT MIN(rel_last) FROM event_queue_synced)")) +} + +/* + +do we need event_queue_synced?? +do we need stable server id or can it be regenerad on each run? + + +*/ diff --git a/sam/repository/events_test.go b/sam/repository/events_test.go new file mode 100644 index 000000000..a5bb6e7dd --- /dev/null +++ b/sam/repository/events_test.go @@ -0,0 +1,8 @@ +package repository + +import ( + "testing" +) + +func TestEvents(t *testing.T) { +} diff --git a/sam/repository/message.go b/sam/repository/message.go index f01777054..35256af76 100644 --- a/sam/repository/message.go +++ b/sam/repository/message.go @@ -1,19 +1,28 @@ package repository import ( - "github.com/crusttech/crust/sam/types" - "github.com/titpetric/factory" + "context" "time" + + "github.com/titpetric/factory" + + "github.com/crusttech/crust/sam/types" ) type ( Message interface { + With(ctx context.Context) Message + FindMessageByID(id uint64) (*types.Message, error) FindMessages(filter *types.MessageFilter) (types.MessageSet, error) CreateMessage(mod *types.Message) (*types.Message, error) UpdateMessage(mod *types.Message) (*types.Message, error) DeleteMessageByID(id uint64) error } + + message struct { + *repository + } ) const ( @@ -34,14 +43,24 @@ const ( ErrMessageNotFound = repositoryError("MessageNotFound") ) -func (r *repository) FindMessageByID(id uint64) (*types.Message, error) { +func NewMessage(ctx context.Context) Message { + return (&message{}).With(ctx) +} + +func (r *message) With(ctx context.Context) Message { + return &message{ + repository: r.repository.With(ctx), + } +} + +func (r *message) FindMessageByID(id uint64) (*types.Message, error) { mod := &types.Message{} sql := sqlMessagesSelect + " AND id = ?" return mod, isFound(r.db().Get(mod, sql, id), mod.ID > 0, ErrMessageNotFound) } -func (r *repository) FindMessages(filter *types.MessageFilter) (types.MessageSet, error) { +func (r *message) FindMessages(filter *types.MessageFilter) (types.MessageSet, error) { params := make([]interface{}, 0) rval := make(types.MessageSet, 0) @@ -79,19 +98,19 @@ func (r *repository) FindMessages(filter *types.MessageFilter) (types.MessageSet return rval, r.db().Select(&rval, sql, params...) } -func (r *repository) CreateMessage(mod *types.Message) (*types.Message, error) { +func (r *message) CreateMessage(mod *types.Message) (*types.Message, error) { mod.ID = factory.Sonyflake.NextID() mod.CreatedAt = time.Now() return mod, r.db().Insert("messages", mod) } -func (r *repository) UpdateMessage(mod *types.Message) (*types.Message, error) { +func (r *message) UpdateMessage(mod *types.Message) (*types.Message, error) { mod.UpdatedAt = timeNowPtr() return mod, r.db().Replace("messages", mod) } -func (r *repository) DeleteMessageByID(id uint64) error { +func (r *message) DeleteMessageByID(id uint64) error { return r.updateColumnByID("messages", "deleted_at", nil, id) } diff --git a/sam/repository/organisation.go b/sam/repository/organisation.go index bf7afd8cc..0c1313d1b 100644 --- a/sam/repository/organisation.go +++ b/sam/repository/organisation.go @@ -1,13 +1,18 @@ package repository import ( - "github.com/crusttech/crust/sam/types" - "github.com/titpetric/factory" + "context" "time" + + "github.com/titpetric/factory" + + "github.com/crusttech/crust/sam/types" ) type ( Organisation interface { + With(ctx context.Context) Organisation + FindOrganisationByID(id uint64) (*types.Organisation, error) FindOrganisations(filter *types.OrganisationFilter) ([]*types.Organisation, error) CreateOrganisation(mod *types.Organisation) (*types.Organisation, error) @@ -16,6 +21,10 @@ type ( UnarchiveOrganisationByID(id uint64) error DeleteOrganisationByID(id uint64) error } + + organisation struct { + *repository + } ) const ( @@ -24,14 +33,24 @@ const ( ErrOrganisationNotFound = repositoryError("OrganisationNotFound") ) -func (r *repository) FindOrganisationByID(id uint64) (*types.Organisation, error) { +func NewOrganisation(ctx context.Context) Organisation { + return (&organisation{}).With(ctx) +} + +func (r *organisation) With(ctx context.Context) Organisation { + return &organisation{ + repository: r.repository.With(ctx), + } +} + +func (r *organisation) FindOrganisationByID(id uint64) (*types.Organisation, error) { sql := "SELECT * FROM organisations WHERE id = ? AND " + sqlOrganisationScope mod := &types.Organisation{} return mod, isFound(r.db().Get(mod, sql, id), mod.ID > 0, ErrOrganisationNotFound) } -func (r *repository) FindOrganisations(filter *types.OrganisationFilter) ([]*types.Organisation, error) { +func (r *organisation) FindOrganisations(filter *types.OrganisationFilter) ([]*types.Organisation, error) { rval := make([]*types.Organisation, 0) params := make([]interface{}, 0) sql := "SELECT * FROM organisations WHERE " + sqlOrganisationScope @@ -48,27 +67,27 @@ func (r *repository) FindOrganisations(filter *types.OrganisationFilter) ([]*typ return rval, r.db().Select(&rval, sql, params...) } -func (r *repository) CreateOrganisation(mod *types.Organisation) (*types.Organisation, error) { +func (r *organisation) CreateOrganisation(mod *types.Organisation) (*types.Organisation, error) { mod.ID = factory.Sonyflake.NextID() mod.CreatedAt = time.Now() return mod, r.db().Insert("organisations", mod) } -func (r *repository) UpdateOrganisation(mod *types.Organisation) (*types.Organisation, error) { +func (r *organisation) UpdateOrganisation(mod *types.Organisation) (*types.Organisation, error) { mod.UpdatedAt = timeNowPtr() return mod, r.db().Replace("organisations", mod) } -func (r *repository) ArchiveOrganisationByID(id uint64) error { +func (r *organisation) ArchiveOrganisationByID(id uint64) error { return r.updateColumnByID("organisations", "archived_at", time.Now(), id) } -func (r *repository) UnarchiveOrganisationByID(id uint64) error { +func (r *organisation) UnarchiveOrganisationByID(id uint64) error { return r.updateColumnByID("organisations", "archived_at", nil, id) } -func (r *repository) DeleteOrganisationByID(id uint64) error { +func (r *organisation) DeleteOrganisationByID(id uint64) error { return r.updateColumnByID("organisations", "deleted_at", nil, id) } diff --git a/sam/repository/reaction.go b/sam/repository/reaction.go index e93781972..89b67c3e4 100644 --- a/sam/repository/reaction.go +++ b/sam/repository/reaction.go @@ -1,50 +1,61 @@ package repository import ( - "github.com/crusttech/crust/sam/types" - "github.com/titpetric/factory" + "context" "time" + + "github.com/titpetric/factory" + + "github.com/crusttech/crust/sam/types" ) type ( Reaction interface { + With(ctx context.Context) Reaction + FindReactionByID(id uint64) (*types.Reaction, error) FindReactionsByRange(channelID, fromReactionID, toReactionID uint64) ([]*types.Reaction, error) CreateReaction(mod *types.Reaction) (*types.Reaction, error) DeleteReactionByID(id uint64) error } + + reaction struct { + *repository + } ) const ( ErrReactionNotFound = repositoryError("ReactionNotFound") ) -func (r *repository) FindReactionByID(id uint64) (*types.Reaction, error) { - sql := "SELECT * FROM reactions WHERE id = ?" - mod := &types.Reaction{} - - return mod, isFound(r.db().Get(mod, sql, id), mod.ID > 0, ErrReactionNotFound) - +func NewReaction(ctx context.Context) Reaction { + return (&reaction{}).With(ctx) } -func (r *repository) FindReactionsByRange(channelID, fromReactionID, toReactionID uint64) ([]*types.Reaction, error) { - rval := make([]*types.Reaction, 0) - sql := ` - SELECT * - FROM reactions - WHERE rel_reaction BETWEEN ? AND ? - AND rel_channel = ?` +func (r *reaction) With(ctx context.Context) Reaction { + return &reaction{ + repository: r.repository.With(ctx), + } +} +func (r *reaction) FindReactionByID(id uint64) (*types.Reaction, error) { + sql := "SELECT * FROM reactions WHERE id=?" + mod := &types.Reaction{} + return mod, isFound(r.db().Get(mod, sql, id), mod.ID > 0, ErrReactionNotFound) +} + +func (r *reaction) FindReactionsByRange(channelID, fromReactionID, toReactionID uint64) ([]*types.Reaction, error) { + rval := make([]*types.Reaction, 0) + sql := `SELECT * FROM reactions WHERE rel_reaction BETWEEN ? AND ? AND rel_channel=?` return rval, r.db().Select(&rval, sql, fromReactionID, toReactionID, channelID) } -func (r *repository) CreateReaction(mod *types.Reaction) (*types.Reaction, error) { +func (r *reaction) CreateReaction(mod *types.Reaction) (*types.Reaction, error) { mod.ID = factory.Sonyflake.NextID() mod.CreatedAt = time.Now() - return mod, r.db().Insert("reactions", mod) } -func (r *repository) DeleteReactionByID(id uint64) error { - return exec(r.db().Exec("DELETE FROM reactions WHERE id = ?", id)) +func (r *reaction) DeleteReactionByID(id uint64) error { + return exec(r.db().Exec("DELETE FROM reactions WHERE id=?", id)) } diff --git a/sam/repository/repository.go b/sam/repository/repository.go index 49e871c8b..778106e9f 100644 --- a/sam/repository/repository.go +++ b/sam/repository/repository.go @@ -2,7 +2,7 @@ package repository import ( "context" - "github.com/pkg/errors" + "github.com/titpetric/factory" ) @@ -10,80 +10,38 @@ type ( repository struct { ctx context.Context - // Current transaction - tx *factory.DB + // Get database handle + dbh func(ctxs ...context.Context) *factory.DB } - - Transactionable interface { - BeginWith(ctx context.Context, callback BeginCallback) error - Begin() error - Rollback() error - Commit() error - } - - Contextable interface { - WithCtx(ctx context.Context) Interfaces - } - - Interfaces interface { - Transactionable - Contextable - - Attachment - Channel - Message - Organisation - Reaction - Team - EventQueue - } - - BeginCallback func(r Interfaces) error ) -func New() *repository { - return &repository{ctx: context.Background()} -} +var _db *factory.DB -func (r *repository) WithCtx(ctx context.Context) Interfaces { - return &repository{ctx: ctx, tx: r.tx} -} - -func (r *repository) BeginWith(ctx context.Context, callback BeginCallback) error { - - txr := &repository{ctx: ctx} - - if err := txr.Begin(); err != nil { - return err +// DB returns a repository-wide singleton DB handle +func DB(ctxs ...context.Context) *factory.DB { + if _db == nil { + _db = factory.Database.MustGet() } - - if err := callback(txr); err != nil { - if err := txr.Rollback(); err != nil { - return err - } - - return err + for _, ctx := range ctxs { + _db = _db.With(ctx) + break } - - return txr.Commit() + return _db } -func (r *repository) Begin() error { - return r.db().Begin() -} - -func (r *repository) Commit() error { - return errors.Wrap(r.db().Commit(), "Can not commit changes") -} - -func (r *repository) Rollback() error { - return errors.Wrap(r.db().Rollback(), "Can not rollback changes") +// With updates repository and database contexts +func (r *repository) With(ctx context.Context) *repository { + res := &repository{ + ctx: ctx, + dbh: DB, + } + if r != nil { + res.dbh = r.dbh + } + return res } +// db returns context-aware db handle func (r *repository) db() *factory.DB { - if r.tx == nil { - r.tx = factory.Database.MustGet().With(r.ctx) - } - - return r.tx + return r.dbh(r.ctx) } diff --git a/sam/repository/repository_test.go b/sam/repository/repository_test.go new file mode 100644 index 000000000..e72b56b2b --- /dev/null +++ b/sam/repository/repository_test.go @@ -0,0 +1,11 @@ +package repository + +import ( + "context" + "testing" +) + +func TestEvents(t *testing.T) { + repo = &repository{} + repo.With(context.Background()) +} diff --git a/sam/repository/team.go b/sam/repository/team.go index f7f52603b..e933128c8 100644 --- a/sam/repository/team.go +++ b/sam/repository/team.go @@ -1,13 +1,18 @@ package repository import ( - "github.com/crusttech/crust/sam/types" - "github.com/titpetric/factory" + "context" "time" + + "github.com/titpetric/factory" + + "github.com/crusttech/crust/sam/types" ) type ( Team interface { + With(ctx context.Context) Team + FindTeamByID(id uint64) (*types.Team, error) FindTeams(filter *types.TeamFilter) ([]*types.Team, error) CreateTeam(mod *types.Team) (*types.Team, error) @@ -18,6 +23,10 @@ type ( MergeTeamByID(id, targetTeamID uint64) error MoveTeamByID(id, targetOrganisationID uint64) error } + + team struct { + *repository + } ) const ( @@ -26,14 +35,24 @@ const ( ErrTeamNotFound = repositoryError("TeamNotFound") ) -func (r *repository) FindTeamByID(id uint64) (*types.Team, error) { +func NewTeam(ctx context.Context) Team { + return (&team{}).With(ctx) +} + +func (r *team) With(ctx context.Context) Team { + return &team{ + repository: r.repository.With(ctx), + } +} + +func (r *team) FindTeamByID(id uint64) (*types.Team, error) { sql := "SELECT * FROM teams WHERE id = ? AND " + sqlTeamScope mod := &types.Team{} return mod, isFound(r.db().Get(mod, sql, id), mod.ID > 0, ErrTeamNotFound) } -func (r *repository) FindTeams(filter *types.TeamFilter) ([]*types.Team, error) { +func (r *team) FindTeams(filter *types.TeamFilter) ([]*types.Team, error) { rval := make([]*types.Team, 0) params := make([]interface{}, 0) @@ -51,35 +70,35 @@ func (r *repository) FindTeams(filter *types.TeamFilter) ([]*types.Team, error) return rval, r.db().Select(&rval, sql, params...) } -func (r *repository) CreateTeam(mod *types.Team) (*types.Team, error) { +func (r *team) CreateTeam(mod *types.Team) (*types.Team, error) { mod.ID = factory.Sonyflake.NextID() mod.CreatedAt = time.Now() return mod, r.db().Insert("teams", mod) } -func (r *repository) UpdateTeam(mod *types.Team) (*types.Team, error) { +func (r *team) UpdateTeam(mod *types.Team) (*types.Team, error) { mod.UpdatedAt = timeNowPtr() return mod, r.db().Replace("teams", mod) } -func (r *repository) ArchiveTeamByID(id uint64) error { +func (r *team) ArchiveTeamByID(id uint64) error { return r.updateColumnByID("teams", "archived_at", time.Now(), id) } -func (r *repository) UnarchiveTeamByID(id uint64) error { +func (r *team) UnarchiveTeamByID(id uint64) error { return r.updateColumnByID("teams", "archived_at", nil, id) } -func (r *repository) DeleteTeamByID(id uint64) error { +func (r *team) DeleteTeamByID(id uint64) error { return r.updateColumnByID("teams", "deleted_at", nil, id) } -func (r *repository) MergeTeamByID(id, targetTeamID uint64) error { +func (r *team) MergeTeamByID(id, targetTeamID uint64) error { return ErrNotImplemented } -func (r *repository) MoveTeamByID(id, targetOrganisationID uint64) error { +func (r *team) MoveTeamByID(id, targetOrganisationID uint64) error { return ErrNotImplemented } diff --git a/sam/repository/util.go b/sam/repository/util.go index 0e80f2d4d..deb2a3ecf 100644 --- a/sam/repository/util.go +++ b/sam/repository/util.go @@ -1,9 +1,10 @@ package repository import ( - "encoding/json" "fmt" "time" + + "encoding/json" ) func (r repository) updateColumnByID(tableName, columnName string, value interface{}, id uint64) (err error) { diff --git a/sam/service/attachment.go b/sam/service/attachment.go index 75465b840..ca648baff 100644 --- a/sam/service/attachment.go +++ b/sam/service/attachment.go @@ -3,23 +3,27 @@ package service import ( "context" "fmt" + "io" + "log" + "path" + "strings" + + "net/http" + "net/url" + + "github.com/titpetric/factory" + "github.com/crusttech/crust/internal/auth" "github.com/crusttech/crust/internal/store" "github.com/crusttech/crust/sam/repository" "github.com/crusttech/crust/sam/types" - "github.com/titpetric/factory" - "io" - "log" - "net/http" - "net/url" - "path" - "strings" ) type ( attachment struct { - rpo attachmentRepository - sto store.Store + attachment repository.Attachment + message repository.Message + store store.Store config struct { url string @@ -34,34 +38,29 @@ type ( OpenOriginal(att *types.Attachment) (io.ReadSeeker, error) OpenPreview(att *types.Attachment) (io.ReadSeeker, error) } - - attachmentRepository interface { - repository.Transactionable - repository.Attachment - } ) func Attachment(store store.Store) *attachment { - svc := &attachment{} - + svc := &attachment{ + attachment: repository.NewAttachment(context.Background()), + message: repository.NewMessage(context.Background()), + store: store, + } svc.config.url = "/attachment/%d/%s" svc.config.previewUrl = "/attachment/%d/%s/preview" - svc.rpo = repository.New() - svc.sto = store - return svc } func (svc attachment) FindByID(id uint64) (*types.Attachment, error) { - return svc.rpo.FindAttachmentByID(id) + return svc.attachment.FindAttachmentByID(id) } func (svc attachment) OpenOriginal(att *types.Attachment) (io.ReadSeeker, error) { - return svc.sto.Open(att.Url) + return svc.store.Open(att.Url) } func (svc attachment) OpenPreview(att *types.Attachment) (io.ReadSeeker, error) { - return svc.sto.Open(att.PreviewUrl) + return svc.store.Open(att.PreviewUrl) } @@ -74,7 +73,7 @@ func (svc attachment) LoadFromMessages(ctx context.Context, mm types.MessageSet) return nil }) - if set, err := svc.rpo.FindAttachmentByMessageID(ids...); err != nil { + if set, err := svc.attachment.FindAttachmentByMessageID(ids...); err != nil { return err } else { return set.Walk(func(a *types.MessageAttachment) error { @@ -115,9 +114,9 @@ func (svc attachment) Create(ctx context.Context, channelId uint64, name string, log.Printf("Processing uploaded file (name: %s, size: %d, mime: %s)", att.Name, att.Size, att.Mimetype) - if svc.sto != nil { - att.Url = svc.sto.Original(att.ID, ext) - if err = svc.sto.Save(att.Url, fh); err != nil { + if svc.store != nil { + att.Url = svc.store.Original(att.ID, ext) + if err = svc.store.Save(att.Url, fh); err != nil { log.Print(err.Error()) return } @@ -128,9 +127,9 @@ func (svc attachment) Create(ctx context.Context, channelId uint64, name string, log.Printf("File %s stored as %s", att.Name, att.Url) - return att, svc.rpo.BeginWith(ctx, func(r repository.Interfaces) (err error) { + return att, repository.DB().Transaction(func() (err error) { - if att, err = r.CreateAttachment(att); err != nil { + if att, err = svc.attachment.CreateAttachment(att); err != nil { return } @@ -147,11 +146,11 @@ func (svc attachment) Create(ctx context.Context, channelId uint64, name string, // Create the first message, doing this directly with repository to circumvent // message service constraints - if msg, err = r.CreateMessage(msg); err != nil { + if msg, err = svc.message.CreateMessage(msg); err != nil { return } - if err = r.BindAttachment(att.ID, msg.ID); err != nil { + if err = svc.attachment.BindAttachment(att.ID, msg.ID); err != nil { return } @@ -216,9 +215,9 @@ func (svc attachment) makePreview(att *types.Attachment, original io.ReadSeeker) // Can and how we make a preview of this attachment? var ext = "jpg" - att.PreviewUrl = svc.sto.Preview(att.ID, ext) + att.PreviewUrl = svc.store.Preview(att.ID, ext) - return svc.sto.Save(att.PreviewUrl, original) + return svc.store.Save(att.PreviewUrl, original) } var _ AttachmentService = &attachment{} diff --git a/sam/service/channel.go b/sam/service/channel.go index fcb17ffe4..f1dc65781 100644 --- a/sam/service/channel.go +++ b/sam/service/channel.go @@ -11,7 +11,8 @@ import ( type ( channel struct { - rpo channelRepository + channel repository.Channel + message repository.Message } ChannelService interface { @@ -26,27 +27,22 @@ type ( archiver } - channelRepository interface { - repository.Transactionable - repository.Channel - } - //channelSecurity interface { // CanRead(ctx context.Context, ch *types.Channel) bool //} ) func Channel() *channel { - var svc = &channel{} - - svc.rpo = repository.New() - //svc.sec.ch = ChannelSecurity(svc.rpo) - + var svc = &channel{ + channel: repository.NewChannel(context.Background()), + message: repository.NewMessage(context.Background()), + } + //svc.sec.ch = ChannelSecurity(svc.channel) return svc } func (svc channel) FindByID(ctx context.Context, id uint64) (ch *types.Channel, err error) { - ch, err = svc.rpo.FindChannelByID(id) + ch, err = svc.channel.FindChannelByID(id) if err != nil { return } @@ -60,7 +56,7 @@ func (svc channel) FindByID(ctx context.Context, id uint64) (ch *types.Channel, func (svc channel) Find(ctx context.Context, filter *types.ChannelFilter) ([]*types.Channel, error) { // @todo: permission check to return only channels that channel has access to - if cc, err := svc.rpo.FindChannels(filter); err != nil { + if cc, err := svc.channel.FindChannels(filter); err != nil { return nil, err } else { return cc, svc.preloadMembers(ctx, cc) @@ -74,15 +70,15 @@ func (svc channel) preloadMembers(ctx context.Context, set types.ChannelSet) err // Returns all channels with membership info func (svc channel) FindByMembership(ctx context.Context) (rval []*types.Channel, err error) { - return rval, svc.rpo.BeginWith(ctx, func(r repository.Interfaces) error { + return rval, repository.DB().Transaction(func() error { var chMemberId = auth.GetIdentityFromContext(ctx).Identity() var mm []*types.ChannelMember - if mm, err = r.FindChannelsMembershipsByMemberId(chMemberId); err != nil { + if mm, err = svc.channel.FindChannelsMembershipsByMemberId(chMemberId); err != nil { return err } - if rval, err = r.FindChannels(nil); err != nil { + if rval, err = svc.channel.FindChannels(nil); err != nil { return err } @@ -101,7 +97,7 @@ func (svc channel) FindByMembership(ctx context.Context) (rval []*types.Channel, func (svc channel) Create(ctx context.Context, in *types.Channel) (out *types.Channel, err error) { // @todo: [SECURITY] permission check if user can add channel - return out, svc.rpo.BeginWith(ctx, func(r repository.Interfaces) (err error) { + return out, repository.DB().Transaction(func() (err error) { var msg *types.Message // @todo get organisation from somewhere @@ -139,12 +135,12 @@ func (svc channel) Create(ctx context.Context, in *types.Channel) (out *types.Ch } // Save the channel - if out, err = r.CreateChannel(out); err != nil { + if out, err = svc.channel.CreateChannel(out); err != nil { return } // Join current user as an member & owner - _, err = r.AddChannelMember(&types.ChannelMember{ + _, err = svc.channel.AddChannelMember(&types.ChannelMember{ ChannelID: out.ID, UserID: chCreatorID, Type: types.ChannelMembershipTypeOwner, @@ -157,7 +153,7 @@ func (svc channel) Create(ctx context.Context, in *types.Channel) (out *types.Ch // Create the first message, doing this directly with repository to circumvent // message service constraints - msg, err = r.CreateMessage(svc.makeSystemMessage( + msg, err = svc.message.CreateMessage(svc.makeSystemMessage( out, "@%d created new %s channel, topic is: %s", chCreatorID, @@ -178,11 +174,11 @@ func (svc channel) Create(ctx context.Context, in *types.Channel) (out *types.Ch } func (svc channel) Update(ctx context.Context, in *types.Channel) (out *types.Channel, err error) { - return out, svc.rpo.BeginWith(ctx, func(r repository.Interfaces) (err error) { + return out, repository.DB().Transaction(func() (err error) { var msgs types.MessageSet // @todo [SECURITY] can user access this channel? - if out, err = r.FindChannelByID(in.ID); err != nil { + if out, err = svc.channel.FindChannelByID(in.ID); err != nil { return } @@ -236,7 +232,7 @@ func (svc channel) Update(ctx context.Context, in *types.Channel) (out *types.Ch } // Save the updated channel - if out, err = r.UpdateChannel(in); err != nil { + if out, err = svc.channel.UpdateChannel(in); err != nil { return } @@ -245,7 +241,7 @@ func (svc channel) Update(ctx context.Context, in *types.Channel) (out *types.Ch // Create the first message, doing this directly with repository to circumvent // message service constraints for _, msg := range msgs { - if msg, err = r.CreateMessage(msg); err != nil { + if msg, err = svc.message.CreateMessage(msg); err != nil { // @todo send new msg to the event-loop return err } @@ -261,12 +257,12 @@ func (svc channel) Update(ctx context.Context, in *types.Channel) (out *types.Ch } func (svc channel) Delete(ctx context.Context, id uint64) error { - return svc.rpo.BeginWith(ctx, func(r repository.Interfaces) (err error) { + return repository.DB().Transaction(func() (err error) { var userID = auth.GetIdentityFromContext(ctx).Identity() var ch *types.Channel // @todo [SECURITY] can user access this channel? - if ch, err = r.FindChannelByID(id); err != nil { + if ch, err = svc.channel.FindChannelByID(id); err != nil { return } @@ -276,19 +272,19 @@ func (svc channel) Delete(ctx context.Context, id uint64) error { return errors.New("Channel already deleted") } - _, err = r.CreateMessage(svc.makeSystemMessage(ch, "@%d deleted this channel", userID)) + _, err = svc.message.CreateMessage(svc.makeSystemMessage(ch, "@%d deleted this channel", userID)) - return r.DeleteChannelByID(id) + return svc.channel.DeleteChannelByID(id) }) } func (svc channel) Recover(ctx context.Context, id uint64) error { - return svc.rpo.BeginWith(ctx, func(r repository.Interfaces) (err error) { + return repository.DB().Transaction(func() (err error) { var userID = auth.GetIdentityFromContext(ctx).Identity() var ch *types.Channel // @todo [SECURITY] can user access this channel? - if ch, err = r.FindChannelByID(id); err != nil { + if ch, err = svc.channel.FindChannelByID(id); err != nil { return } @@ -298,19 +294,19 @@ func (svc channel) Recover(ctx context.Context, id uint64) error { return errors.New("Channel not deleted") } - _, err = r.CreateMessage(svc.makeSystemMessage(ch, "@%d recovered this channel", userID)) + _, err = svc.message.CreateMessage(svc.makeSystemMessage(ch, "@%d recovered this channel", userID)) - return r.DeleteChannelByID(id) + return svc.channel.DeleteChannelByID(id) }) } func (svc channel) Archive(ctx context.Context, id uint64) error { - return svc.rpo.BeginWith(ctx, func(r repository.Interfaces) (err error) { + return repository.DB().Transaction(func() (err error) { var userID = auth.GetIdentityFromContext(ctx).Identity() var ch *types.Channel // @todo [SECURITY] can user access this channel? - if ch, err = r.FindChannelByID(id); err != nil { + if ch, err = svc.channel.FindChannelByID(id); err != nil { return } @@ -320,19 +316,19 @@ func (svc channel) Archive(ctx context.Context, id uint64) error { return errors.New("Channel already archived") } - _, err = r.CreateMessage(svc.makeSystemMessage(ch, "@%d archived this channel", userID)) + _, err = svc.message.CreateMessage(svc.makeSystemMessage(ch, "@%d archived this channel", userID)) - return r.ArchiveChannelByID(id) + return svc.channel.ArchiveChannelByID(id) }) } func (svc channel) Unarchive(ctx context.Context, id uint64) error { - return svc.rpo.BeginWith(ctx, func(r repository.Interfaces) (err error) { + return repository.DB().Transaction(func() (err error) { var userID = auth.GetIdentityFromContext(ctx).Identity() var ch *types.Channel // @todo [SECURITY] can user access this channel? - if ch, err = r.FindChannelByID(id); err != nil { + if ch, err = svc.channel.FindChannelByID(id); err != nil { return } @@ -342,9 +338,9 @@ func (svc channel) Unarchive(ctx context.Context, id uint64) error { return errors.New("Channel not archived") } - _, err = r.CreateMessage(svc.makeSystemMessage(ch, "@%d unarchived this channel", userID)) + _, err = svc.message.CreateMessage(svc.makeSystemMessage(ch, "@%d unarchived this channel", userID)) - return r.ArchiveChannelByID(id) + return svc.channel.ArchiveChannelByID(id) }) } diff --git a/sam/service/message.go b/sam/service/message.go index 8816ed819..1de09f6d7 100644 --- a/sam/service/message.go +++ b/sam/service/message.go @@ -11,7 +11,10 @@ import ( type ( message struct { - rpo messageRepository + channel repository.Channel + message repository.Message + reaction repository.Reaction + att AttachmentService } @@ -34,20 +37,12 @@ type ( deleter } - - messageRepository interface { - repository.Transactionable - repository.Message - repository.Reaction - repository.Attachment - repository.Channel - } ) func Message(attSvc AttachmentService) *message { m := &message{ att: attSvc, - rpo: repository.New(), + message: repository.NewMessage(context.Background()), } return m } @@ -60,7 +55,7 @@ func (svc message) Find(ctx context.Context, filter *types.MessageFilter) (mm ty _ = currentUserID _ = filter.ChannelID - mm, err = svc.rpo.FindMessages(filter) + mm, err = svc.message.FindMessages(filter) if err != nil { return nil, err } @@ -69,7 +64,7 @@ func (svc message) Find(ctx context.Context, filter *types.MessageFilter) (mm ty } func (svc message) Direct(ctx context.Context, recipientID uint64, in *types.Message) (out *types.Message, err error) { - return out, svc.rpo.BeginWith(ctx, func(r repository.Interfaces) (err error) { + return out, repository.DB().Transaction(func() (err error) { var currentUserID = auth.GetIdentityFromContext(ctx).Identity() // @todo [SECURITY] verify if current user can send direct messages to anyone? @@ -82,9 +77,9 @@ func (svc message) Direct(ctx context.Context, recipientID uint64, in *types.Mes return errors.New("Not allowed to send direct messages to this user") } - dch, err := r.FindDirectChannelByUserID(currentUserID, recipientID) + dch, err := svc.channel.FindDirectChannelByUserID(currentUserID, recipientID) if err == repository.ErrChannelNotFound { - dch, err = r.CreateChannel(&types.Channel{ + dch, err = svc.channel.CreateChannel(&types.Channel{ Type: types.ChannelTypeDirect, }) @@ -96,13 +91,13 @@ func (svc message) Direct(ctx context.Context, recipientID uint64, in *types.Mes membership.UserID = currentUserID spew.Dump(membership) - if _, err = r.AddChannelMember(membership); err != nil { + if _, err = svc.channel.AddChannelMember(membership); err != nil { return } spew.Dump(membership) membership.UserID = recipientID - if _, err = r.AddChannelMember(membership); err != nil { + if _, err = svc.channel.AddChannelMember(membership); err != nil { return } @@ -118,7 +113,7 @@ func (svc message) Direct(ctx context.Context, recipientID uint64, in *types.Mes spew.Dump(in) // @todo send new msg to the event-loop - out, err = r.CreateMessage(in) + out, err = svc.message.CreateMessage(in) return }) } @@ -131,7 +126,7 @@ func (svc message) Create(ctx context.Context, mod *types.Message) (*types.Messa mod.UserID = currentUserID - message, err := svc.rpo.CreateMessage(mod) + message, err := svc.message.CreateMessage(mod) if err == nil { PubSub().Event(ctx, "new message added") } @@ -149,7 +144,7 @@ func (svc message) Update(ctx context.Context, mod *types.Message) (*types.Messa // @todo verify ownership - return svc.rpo.UpdateMessage(mod) + return svc.message.UpdateMessage(mod) } func (svc message) Delete(ctx context.Context, id uint64) error { @@ -163,7 +158,7 @@ func (svc message) Delete(ctx context.Context, id uint64) error { // @todo verify ownership - return svc.rpo.DeleteMessageByID(id) + return svc.message.DeleteMessageByID(id) } func (svc message) React(ctx context.Context, messageID uint64, reaction string) error { @@ -182,7 +177,7 @@ func (svc message) React(ctx context.Context, messageID uint64, reaction string) Reaction: reaction, } - if _, err := svc.rpo.CreateReaction(r); err != nil { + if _, err := svc.reaction.CreateReaction(r); err != nil { return err } @@ -199,7 +194,7 @@ func (svc message) Unreact(ctx context.Context, messageID uint64, reaction strin // @todo load reaction and verify ownership var r *types.Reaction - return svc.rpo.DeleteReactionByID(r.ID) + return svc.reaction.DeleteReactionByID(r.ID) } func (svc message) Pin(ctx context.Context, messageID uint64) error { diff --git a/sam/service/organisation.go b/sam/service/organisation.go index 59bcef36e..4f80be418 100644 --- a/sam/service/organisation.go +++ b/sam/service/organisation.go @@ -8,7 +8,7 @@ import ( type ( organisation struct { - rpo organisationRepository + rpo repository.Organisation } OrganisationService interface { @@ -21,15 +21,10 @@ type ( deleter archiver } - - organisationRepository interface { - repository.Transactionable - repository.Organisation - } ) func Organisation() *organisation { - return &organisation{rpo: repository.New()} + return &organisation{rpo: repository.NewOrganisation(context.Background())} } func (svc organisation) FindByID(ctx context.Context, id uint64) (*types.Organisation, error) { diff --git a/sam/service/team.go b/sam/service/team.go index 90f7def73..d783597d8 100644 --- a/sam/service/team.go +++ b/sam/service/team.go @@ -8,7 +8,7 @@ import ( type ( team struct { - rpo teamRepository + team repository.Team } TeamService interface { @@ -23,69 +23,66 @@ type ( deleter archiver } - - teamRepository interface { - repository.Transactionable - repository.Team - } ) func Team() *team { - return &team{rpo: repository.New()} + return &team{ + team: repository.NewTeam(context.Background()), + } } func (svc team) FindByID(ctx context.Context, id uint64) (*types.Team, error) { // @todo: permission check if current user has access to this team - return svc.rpo.FindTeamByID(id) + return svc.team.FindTeamByID(id) } func (svc team) Find(ctx context.Context, filter *types.TeamFilter) ([]*types.Team, error) { // @todo: permission check to return only teams that current user has access to - return svc.rpo.FindTeams(filter) + return svc.team.FindTeams(filter) } func (svc team) Create(ctx context.Context, mod *types.Team) (*types.Team, error) { // @todo: permission check if current user can add/edit team - return svc.rpo.CreateTeam(mod) + return svc.team.CreateTeam(mod) } func (svc team) Update(ctx context.Context, mod *types.Team) (*types.Team, error) { // @todo: permission check if current user can add/edit team // @todo: make sure archived & deleted entries can not be edited - return svc.rpo.UpdateTeam(mod) + return svc.team.UpdateTeam(mod) } func (svc team) Delete(ctx context.Context, id uint64) error { // @todo: make history unavailable // @todo: notify users that team has been removed (remove from web UI) // @todo: permissions check if current user can remove team - return svc.rpo.DeleteTeamByID(id) + return svc.team.DeleteTeamByID(id) } func (svc team) Archive(ctx context.Context, id uint64) error { // @todo: make history unavailable // @todo: notify users that team has been removed (remove from web UI) // @todo: permissions check if current user can remove team - return svc.rpo.ArchiveTeamByID(id) + return svc.team.ArchiveTeamByID(id) } func (svc team) Unarchive(ctx context.Context, id uint64) error { // @todo: permissions check if current user can unarchive team // @todo: make history accessible // @todo: notify users that team has been unarchived - return svc.rpo.UnarchiveTeamByID(id) + return svc.team.UnarchiveTeamByID(id) } func (svc team) Merge(ctx context.Context, id, targetTeamID uint64) error { // @todo: permission check if current user can merge team - return svc.rpo.MergeTeamByID(id, targetTeamID) + return svc.team.MergeTeamByID(id, targetTeamID) } func (svc team) Move(ctx context.Context, id, targetOrganisationID uint64) error { // @todo: permission check if current user can move team to another organisation - return svc.rpo.MoveTeamByID(id, targetOrganisationID) + return svc.team.MoveTeamByID(id, targetOrganisationID) } var _ TeamService = &team{} diff --git a/sam/websocket/event_queue.go b/sam/websocket/event_queue.go index eb059df47..ca412a8a9 100644 --- a/sam/websocket/event_queue.go +++ b/sam/websocket/event_queue.go @@ -11,14 +11,6 @@ import ( ) type ( - eventQueuePuller interface { - EventQueuePull(origin uint64) ([]*types.EventQueueItem, error) - EventQueueSync(origin uint64, ID uint64) error - } - eventQueuePusher interface { - EventQueuePush(*types.EventQueueItem) error - } - eventQueueWalker interface { Walk(func(session *Session)) } @@ -46,19 +38,19 @@ func EventQueue(origin uint64) *eventQueue { } } -func (eq *eventQueue) store(ctx context.Context, qp eventQueuePusher) { +func (eq *eventQueue) store(ctx context.Context, qp repository.Events) { go func() { for { select { case <-ctx.Done(): case eqi := <-eq.queue: - qp.EventQueuePush(eqi) + qp.Push(eqi) } } }() } -func (eq *eventQueue) feedSessions(ctx context.Context, config *repository.Flags, qp eventQueuePuller, store eventQueueWalker) error { +func (eq *eventQueue) feedSessions(ctx context.Context, config *repository.Flags, qp repository.Events, store eventQueueWalker) error { newMessageEvent := make(chan struct{}, eventQueueBacklog) done := make(chan error, 1) @@ -76,7 +68,7 @@ func (eq *eventQueue) feedSessions(ctx context.Context, config *repository.Flags poll := func() error { for { - items, err := qp.EventQueuePull(eq.origin) + items, err := qp.Pull(eq.origin) if err != nil { return err } @@ -106,7 +98,7 @@ func (eq *eventQueue) feedSessions(ctx context.Context, config *repository.Flags } if lastSyncedId > 0 { - qp.EventQueueSync(eq.origin, lastSyncedId) + qp.Sync(eq.origin, lastSyncedId) } } } diff --git a/sam/websocket/router.go b/sam/websocket/router.go index b0276bbce..27da6d3e6 100644 --- a/sam/websocket/router.go +++ b/sam/websocket/router.go @@ -11,14 +11,14 @@ import ( func MountRoutes(ctx context.Context, config *repository.Flags) func(chi.Router) { return func(r chi.Router) { - repo := repository.New() + events := repository.NewEvents(ctx) go func() { - if err := eq.feedSessions(ctx, config, repo, store); err != nil { + if err := eq.feedSessions(ctx, config, events, store); err != nil { panic(fmt.Sprintf("Error when starting sessions event feed: %+v", err)) } }() - eq.store(ctx, repo) + eq.store(ctx, events) websocket := Websocket{}.New(config) r.Group(func(r chi.Router) {