Skip to content

Commit ba466a7

Browse files
AchoArnoldCopilot
andcommitted
fix(api): return updated thread atomically
Use PostgreSQL RETURNING so concurrent activity cannot make status update responses stale. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: bf00a0ac-e11f-4015-b295-3cdd9b491229
1 parent b5a7db2 commit ba466a7

8 files changed

Lines changed: 52 additions & 51 deletions

api/pkg/handlers/message_thread_handler_test.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -33,8 +33,8 @@ func (stub *messageThreadHandlerRepositoryStub) UpdateActivity(context.Context,
3333
return nil
3434
}
3535

36-
func (stub *messageThreadHandlerRepositoryStub) UpdateStatus(context.Context, entities.UserID, uuid.UUID, repositories.MessageThreadStatusUpdate) error {
37-
return stacktrace.PropagateWithCode(gorm.ErrRecordNotFound, repositories.ErrCodeNotFound, "not found")
36+
func (stub *messageThreadHandlerRepositoryStub) UpdateStatus(context.Context, entities.UserID, uuid.UUID, repositories.MessageThreadStatusUpdate) (*entities.MessageThread, error) {
37+
return nil, stacktrace.PropagateWithCode(gorm.ErrRecordNotFound, repositories.ErrCodeNotFound, "not found")
3838
}
3939

4040
func (stub *messageThreadHandlerRepositoryStub) LoadByOwnerContact(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) {

api/pkg/listeners/read_receipts_test_helpers_test.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -36,8 +36,8 @@ func (repository *listenerMessageThreadRepository) UpdateActivity(_ context.Cont
3636
return nil
3737
}
3838

39-
func (repository *listenerMessageThreadRepository) UpdateStatus(context.Context, entities.UserID, uuid.UUID, repositories.MessageThreadStatusUpdate) error {
40-
return nil
39+
func (repository *listenerMessageThreadRepository) UpdateStatus(_ context.Context, _ entities.UserID, threadID uuid.UUID, _ repositories.MessageThreadStatusUpdate) (*entities.MessageThread, error) {
40+
return &entities.MessageThread{ID: threadID}, nil
4141
}
4242

4343
func (repository *listenerMessageThreadRepository) UpdateAfterDeletedMessage(context.Context, repositories.MessageThreadDeletedUpdate) error {

api/pkg/repositories/gorm_message_thread_repository.go

Lines changed: 17 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -40,14 +40,22 @@ func messageThreadActivityUpdates(params MessageThreadActivityUpdate) map[string
4040
"order_timestamp": params.Timestamp,
4141
"last_message_id": params.MessageID,
4242
"last_message_content": params.Content,
43-
"status": string(params.Status),
43+
"status": params.Status,
4444
}
4545
if params.Unarchive {
4646
updates["is_archived"] = false
4747
}
4848
return updates
4949
}
5050

51+
func messageThreadDeletedUpdates(params MessageThreadDeletedUpdate) map[string]any {
52+
return map[string]any{
53+
"last_message_id": params.LastMessageID,
54+
"last_message_content": params.LastMessageContent,
55+
"status": params.LastMessageStatus,
56+
}
57+
}
58+
5159
func messageThreadStatusUpdates(params MessageThreadStatusUpdate) map[string]any {
5260
updates := make(map[string]any)
5361
if params.IsArchived != nil {
@@ -95,11 +103,7 @@ func (repository *gormMessageThreadRepository) UpdateAfterDeletedMessage(ctx con
95103
Model(&entities.MessageThread{}).
96104
Where("user_id = ?", params.UserID).
97105
Where("id = ?", params.MessageThreadID).
98-
Updates(map[string]any{
99-
"last_message_id": params.LastMessageID,
100-
"last_message_content": params.LastMessageContent,
101-
"status": string(params.LastMessageStatus),
102-
})
106+
Updates(messageThreadDeletedUpdates(params))
103107
if result.Error != nil {
104108
return repository.tracer.WrapErrorSpan(span, stacktrace.Propagate(result.Error, "cannot update deleted-message metadata for thread [%s]", params.MessageThreadID))
105109
}
@@ -183,23 +187,25 @@ func (repository *gormMessageThreadRepository) UpdateStatus(
183187
userID entities.UserID,
184188
messageThreadID uuid.UUID,
185189
params MessageThreadStatusUpdate,
186-
) error {
190+
) (*entities.MessageThread, error) {
187191
ctx, span := repository.tracer.Start(ctx)
188192
defer span.End()
189193

194+
thread := new(entities.MessageThread)
190195
result := repository.db.WithContext(ctx).
191-
Model(&entities.MessageThread{}).
196+
Model(thread).
197+
Clauses(clause.Returning{}).
192198
Where("user_id = ?", userID).
193199
Where("id = ?", messageThreadID).
194200
Updates(messageThreadStatusUpdates(params))
195201
if result.Error != nil {
196-
return repository.tracer.WrapErrorSpan(span, stacktrace.Propagate(result.Error, "cannot update status for thread [%s]", messageThreadID))
202+
return nil, repository.tracer.WrapErrorSpan(span, stacktrace.Propagate(result.Error, "cannot update status for thread [%s]", messageThreadID))
197203
}
198204
if result.RowsAffected == 0 {
199-
return repository.tracer.WrapErrorSpan(span, stacktrace.PropagateWithCode(gorm.ErrRecordNotFound, ErrCodeNotFound, "thread with id [%s] not found", messageThreadID))
205+
return nil, repository.tracer.WrapErrorSpan(span, stacktrace.PropagateWithCode(gorm.ErrRecordNotFound, ErrCodeNotFound, "thread with id [%s] not found", messageThreadID))
200206
}
201207

202-
return nil
208+
return thread, nil
203209
}
204210

205211
// LoadByOwnerContact a thread between 2 users

api/pkg/repositories/gorm_message_thread_repository_test.go

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -116,13 +116,29 @@ func TestMessageThreadActivityUpdatesOwnOnlyMessageColumns(t *testing.T) {
116116
"order_timestamp": time.Date(2026, 7, 18, 7, 0, 0, 0, time.UTC),
117117
"last_message_id": messageID,
118118
"last_message_content": "hello",
119-
"status": entities.MessageStatusReceived,
119+
"status": entities.MessageStatus(entities.MessageStatusReceived),
120120
}, updates)
121121
assert.NotContains(t, updates, "is_read")
122122
assert.NotContains(t, updates, "is_archived")
123123
assert.NotContains(t, updates, "last_read_at")
124124
}
125125

126+
func TestMessageThreadDeletedUpdatesPreserveStatusType(t *testing.T) {
127+
messageID := uuid.New()
128+
content := "previous message"
129+
updates := messageThreadDeletedUpdates(MessageThreadDeletedUpdate{
130+
LastMessageID: &messageID,
131+
LastMessageContent: &content,
132+
LastMessageStatus: entities.MessageStatusDelivered,
133+
})
134+
135+
assert.Equal(t, map[string]any{
136+
"last_message_id": &messageID,
137+
"last_message_content": &content,
138+
"status": entities.MessageStatus(entities.MessageStatusDelivered),
139+
}, updates)
140+
}
141+
126142
func TestMessageThreadStatusUpdatesReadOnly(t *testing.T) {
127143
isRead := true
128144
readAt := time.Date(2026, 7, 18, 7, 1, 0, 0, time.UTC)

api/pkg/repositories/message_thread_repository.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -44,7 +44,7 @@ type MessageThreadRepository interface {
4444
UpdateActivity(ctx context.Context, params MessageThreadActivityUpdate) error
4545

4646
// UpdateStatus persists archive/read status fields for a thread
47-
UpdateStatus(ctx context.Context, userID entities.UserID, messageThreadID uuid.UUID, params MessageThreadStatusUpdate) error
47+
UpdateStatus(ctx context.Context, userID entities.UserID, messageThreadID uuid.UUID, params MessageThreadStatusUpdate) (*entities.MessageThread, error)
4848

4949
// LoadByOwnerContact fetches a thread between owner and contact
5050
LoadByOwnerContact(ctx context.Context, userID entities.UserID, owner string, contact string) (*entities.MessageThread, error)

api/pkg/services/message_thread_service.go

Lines changed: 3 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -128,7 +128,7 @@ func (service *MessageThreadService) UpdateThread(ctx context.Context, params Me
128128
return service.tracer.WrapErrorSpan(span, stacktrace.PropagateWithCode(err, stacktrace.GetCode(err), "cannot update message thread with id [%s] after adding message [%s]", thread.ID, params.MessageID))
129129
}
130130

131-
ctxLogger.Info(fmt.Sprintf("thread with id [%s] updated with last message [%s] and status [%s]", thread.ID, thread.LastMessageID, thread.Status))
131+
ctxLogger.Info(fmt.Sprintf("thread with id [%s] updated with last message [%s] and status [%s]", thread.ID, params.MessageID, params.Status))
132132
return nil
133133
}
134134

@@ -145,30 +145,16 @@ func (service *MessageThreadService) UpdateStatus(ctx context.Context, params Me
145145
ctx, span := service.tracer.Start(ctx)
146146
defer span.End()
147147

148-
thread, err := service.repository.Load(ctx, params.UserID, params.MessageThreadID)
149-
if err != nil {
150-
return nil, service.tracer.WrapErrorSpan(span, stacktrace.Propagate(err, "cannot find thread with id [%s]", params.MessageThreadID))
151-
}
152-
153148
update := repositories.MessageThreadStatusUpdate{
154149
IsArchived: params.IsArchived,
155150
IsRead: params.IsRead,
156151
ReadAt: time.Now().UTC(),
157152
}
158-
if err = service.repository.UpdateStatus(ctx, params.UserID, params.MessageThreadID, update); err != nil {
153+
thread, err := service.repository.UpdateStatus(ctx, params.UserID, params.MessageThreadID, update)
154+
if err != nil {
159155
return nil, service.tracer.WrapErrorSpan(span, stacktrace.PropagateWithCode(err, stacktrace.GetCode(err), "cannot update message thread with id [%s]", params.MessageThreadID))
160156
}
161157

162-
if params.IsArchived != nil {
163-
thread.IsArchived = *params.IsArchived
164-
}
165-
if params.IsRead != nil {
166-
thread.IsRead = *params.IsRead
167-
if *params.IsRead {
168-
thread.LastReadAt = update.ReadAt
169-
}
170-
}
171-
172158
return thread, nil
173159
}
174160

api/pkg/services/message_thread_service_test.go

Lines changed: 7 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,7 @@ type messageThreadRepositoryStub struct {
2020
load func(context.Context, entities.UserID, uuid.UUID) (*entities.MessageThread, error)
2121
store func(context.Context, *entities.MessageThread) error
2222
updateActivity func(context.Context, repositories.MessageThreadActivityUpdate) error
23-
updateStatus func(context.Context, entities.UserID, uuid.UUID, repositories.MessageThreadStatusUpdate) error
23+
updateStatus func(context.Context, entities.UserID, uuid.UUID, repositories.MessageThreadStatusUpdate) (*entities.MessageThread, error)
2424
}
2525

2626
func (stub *messageThreadRepositoryStub) Store(ctx context.Context, thread *entities.MessageThread) error {
@@ -37,11 +37,11 @@ func (stub *messageThreadRepositoryStub) UpdateActivity(ctx context.Context, par
3737
return nil
3838
}
3939

40-
func (stub *messageThreadRepositoryStub) UpdateStatus(ctx context.Context, userID entities.UserID, threadID uuid.UUID, params repositories.MessageThreadStatusUpdate) error {
40+
func (stub *messageThreadRepositoryStub) UpdateStatus(ctx context.Context, userID entities.UserID, threadID uuid.UUID, params repositories.MessageThreadStatusUpdate) (*entities.MessageThread, error) {
4141
if stub.updateStatus != nil {
4242
return stub.updateStatus(ctx, userID, threadID, params)
4343
}
44-
return nil
44+
return &entities.MessageThread{ID: threadID}, nil
4545
}
4646

4747
func (stub *messageThreadRepositoryStub) UpdateAfterDeletedMessage(context.Context, repositories.MessageThreadDeletedUpdate) error {
@@ -181,16 +181,10 @@ func TestUpdateStatusChangesOnlyRequestedState(t *testing.T) {
181181
threadID := uuid.New()
182182
isRead := false
183183
var captured repositories.MessageThreadStatusUpdate
184-
var calls []string
185184
repository := &messageThreadRepositoryStub{
186-
updateStatus: func(_ context.Context, _ entities.UserID, _ uuid.UUID, params repositories.MessageThreadStatusUpdate) error {
187-
calls = append(calls, "update")
185+
updateStatus: func(_ context.Context, _ entities.UserID, _ uuid.UUID, params repositories.MessageThreadStatusUpdate) (*entities.MessageThread, error) {
188186
captured = params
189-
return nil
190-
},
191-
load: func(context.Context, entities.UserID, uuid.UUID) (*entities.MessageThread, error) {
192-
calls = append(calls, "load")
193-
return &entities.MessageThread{ID: threadID, IsArchived: true, IsRead: true}, nil
187+
return &entities.MessageThread{ID: threadID, IsArchived: true, IsRead: false}, nil
194188
},
195189
}
196190

@@ -207,16 +201,12 @@ func TestUpdateStatusChangesOnlyRequestedState(t *testing.T) {
207201
assert.False(t, captured.ReadAt.IsZero())
208202
assert.True(t, thread.IsArchived)
209203
assert.False(t, thread.IsRead)
210-
assert.Equal(t, []string{"load", "update"}, calls)
211204
}
212205

213206
func TestUpdateStatusPreservesNotFoundCode(t *testing.T) {
214207
repository := &messageThreadRepositoryStub{
215-
load: func(context.Context, entities.UserID, uuid.UUID) (*entities.MessageThread, error) {
216-
return &entities.MessageThread{ID: uuid.New()}, nil
217-
},
218-
updateStatus: func(context.Context, entities.UserID, uuid.UUID, repositories.MessageThreadStatusUpdate) error {
219-
return stacktrace.PropagateWithCode(gorm.ErrRecordNotFound, repositories.ErrCodeNotFound, "not found")
208+
updateStatus: func(context.Context, entities.UserID, uuid.UUID, repositories.MessageThreadStatusUpdate) (*entities.MessageThread, error) {
209+
return nil, stacktrace.PropagateWithCode(gorm.ErrRecordNotFound, repositories.ErrCodeNotFound, "not found")
220210
},
221211
}
222212

tests/read_receipts_test.go

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -158,6 +158,9 @@ func TestMessageThreadReadReceipts(t *testing.T) {
158158

159159
updated := markMessageThreadRead(ctx, t, thread.ID)
160160
assert.True(t, updated.IsRead)
161+
assert.Equal(t, contact, updated.Contact)
162+
require.NotNil(t, updated.LastMessageContent)
163+
assert.Equal(t, "Unread inbound message", *updated.LastMessageContent)
161164
waitForMessageThread(ctx, t, phone.PhoneNumber, contact, 10*time.Second, func(thread integrationMessageThread) bool {
162165
return thread.IsRead
163166
})

0 commit comments

Comments
 (0)