diff --git a/api/docs/docs.go b/api/docs/docs.go index 8eb6dafe..b1a92b6b 100644 --- a/api/docs/docs.go +++ b/api/docs/docs.go @@ -4129,12 +4129,12 @@ const docTemplate = `{ "created_at", "id", "is_archived", - "is_read", "last_message_content", "last_message_id", "order_timestamp", "owner", "status", + "unread_count", "updated_at", "user_id" ], @@ -4167,10 +4167,6 @@ const docTemplate = `{ "type": "boolean", "example": false }, - "is_read": { - "type": "boolean", - "example": true - }, "last_message_content": { "type": "string", "example": "This is a sample message content" @@ -4191,6 +4187,10 @@ const docTemplate = `{ "type": "string", "example": "PENDING" }, + "unread_count": { + "type": "integer", + "example": 2 + }, "updated_at": { "type": "string", "example": "2022-06-05T14:26:09.527976+03:00" @@ -4920,9 +4920,11 @@ const docTemplate = `{ "type": "boolean", "example": true }, - "is_read": { - "type": "boolean", - "example": true + "unread_count": { + "type": "integer", + "maximum": 0, + "minimum": 0, + "example": 0 } } }, diff --git a/api/docs/swagger.json b/api/docs/swagger.json index 21c6bbf8..2f8706d9 100644 --- a/api/docs/swagger.json +++ b/api/docs/swagger.json @@ -4126,12 +4126,12 @@ "created_at", "id", "is_archived", - "is_read", "last_message_content", "last_message_id", "order_timestamp", "owner", "status", + "unread_count", "updated_at", "user_id" ], @@ -4164,10 +4164,6 @@ "type": "boolean", "example": false }, - "is_read": { - "type": "boolean", - "example": true - }, "last_message_content": { "type": "string", "example": "This is a sample message content" @@ -4188,6 +4184,10 @@ "type": "string", "example": "PENDING" }, + "unread_count": { + "type": "integer", + "example": 2 + }, "updated_at": { "type": "string", "example": "2022-06-05T14:26:09.527976+03:00" @@ -4917,9 +4917,11 @@ "type": "boolean", "example": true }, - "is_read": { - "type": "boolean", - "example": true + "unread_count": { + "type": "integer", + "maximum": 0, + "minimum": 0, + "example": 0 } } }, diff --git a/api/docs/swagger.yaml b/api/docs/swagger.yaml index 4fe2f2b7..fb88f704 100644 --- a/api/docs/swagger.yaml +++ b/api/docs/swagger.yaml @@ -367,9 +367,6 @@ definitions: is_archived: example: false type: boolean - is_read: - example: true - type: boolean last_message_content: example: This is a sample message content type: string @@ -385,6 +382,9 @@ definitions: status: example: PENDING type: string + unread_count: + example: 2 + type: integer updated_at: example: "2022-06-05T14:26:09.527976+03:00" type: string @@ -397,12 +397,12 @@ definitions: - created_at - id - is_archived - - is_read - last_message_content - last_message_id - order_timestamp - owner - status + - unread_count - updated_at - user_id type: object @@ -956,9 +956,11 @@ definitions: is_archived: example: true type: boolean - is_read: - example: true - type: boolean + unread_count: + example: 0 + maximum: 0 + minimum: 0 + type: integer type: object requests.PhoneAPIKeyStoreRequest: properties: diff --git a/api/pkg/di/container.go b/api/pkg/di/container.go index ebe57662..910fa9a8 100644 --- a/api/pkg/di/container.go +++ b/api/pkg/di/container.go @@ -380,7 +380,39 @@ ALTER TABLE discords ADD CONSTRAINT IF NOT EXISTS uni_discords_server_id CHECK ( } if err = db.AutoMigrate(&entities.MessageThread{}); err != nil { - container.logger.Fatal(stacktrace.Propagatef(err, "cannot migrate %T", &entities.MessageThread{})) + container.logger.Fatal(stacktrace.Propagate(err, "cannot migrate message thread schema")) + } + + if db.Migrator().HasColumn(&entities.MessageThread{}, "is_read") { + if err = db. + Model(&entities.MessageThread{}). + Where("is_read = ?", false). + Where("unread_count = ?", 0). + Update("unread_count", 1). + Error; err != nil { + container.logger.Fatal(stacktrace.Propagate(err, "cannot backfill message thread unread counts")) + } + if err = db.Migrator().DropColumn(&entities.MessageThread{}, "is_read"); err != nil { + container.logger.Fatal(stacktrace.Propagate(err, "cannot drop legacy message thread read state")) + } + } + + if db.Migrator().HasColumn(&entities.MessageThread{}, "last_read_at") { + if err = db.Migrator().DropColumn(&entities.MessageThread{}, "last_read_at"); err != nil { + container.logger.Fatal(stacktrace.Propagate(err, "cannot drop message thread read watermark")) + } + } + + if db.Migrator().HasTable("message_thread_unread_items") { + if err = db.Migrator().DropTable("message_thread_unread_items"); err != nil { + container.logger.Fatal(stacktrace.Propagate(err, "cannot drop message thread unread item ledger")) + } + } + + if db.Migrator().HasTable("message_thread_deleted_items") { + if err = db.Migrator().DropTable("message_thread_deleted_items"); err != nil { + container.logger.Fatal(stacktrace.Propagate(err, "cannot drop deleted message item ledger")) + } } if err = db.AutoMigrate(&entities.User{}); err != nil { diff --git a/api/pkg/entities/message_thread.go b/api/pkg/entities/message_thread.go index d4d893a4..c1f4abb5 100644 --- a/api/pkg/entities/message_thread.go +++ b/api/pkg/entities/message_thread.go @@ -9,12 +9,11 @@ import ( // MessageThread represents a message thread between 2 phone numbers type MessageThread struct { ID uuid.UUID `json:"id" gorm:"primaryKey;type:uuid;" example:"32343a19-da5e-4b1b-a767-3298a73703ca"` - Owner string `json:"owner" example:"+18005550199"` - Contact string `json:"contact" example:"+18005550100"` + Owner string `json:"owner" gorm:"uniqueIndex:idx_message_threads_conversation" example:"+18005550199"` + Contact string `json:"contact" gorm:"uniqueIndex:idx_message_threads_conversation" example:"+18005550100"` IsArchived bool `json:"is_archived" example:"false"` - IsRead bool `json:"is_read" gorm:"not null;default:true" example:"true"` - LastReadAt time.Time `json:"-" gorm:"not null;default:CURRENT_TIMESTAMP"` - UserID UserID `json:"user_id" example:"WB7DRDWrJZRGbYrv2CKGkqbzvqdC"` + UnreadCount uint `json:"unread_count" gorm:"not null;default:0" example:"2"` + UserID UserID `json:"user_id" gorm:"uniqueIndex:idx_message_threads_conversation" example:"WB7DRDWrJZRGbYrv2CKGkqbzvqdC"` Color string `json:"color" example:"indigo"` Status MessageStatus `json:"status" example:"PENDING"` LastMessageContent *string `json:"last_message_content" example:"This is a sample message content"` diff --git a/api/pkg/entities/message_thread_test.go b/api/pkg/entities/message_thread_test.go index b587bacb..767bee21 100644 --- a/api/pkg/entities/message_thread_test.go +++ b/api/pkg/entities/message_thread_test.go @@ -9,20 +9,20 @@ import ( "github.com/stretchr/testify/require" ) -func TestMessageThreadReadFieldsHaveBackwardCompatibleDefaults(t *testing.T) { +func TestMessageThreadUnreadFields(t *testing.T) { threadType := reflect.TypeOf(MessageThread{}) - isRead, ok := threadType.FieldByName("IsRead") - require.True(t, ok) - assert.Contains(t, isRead.Tag.Get("gorm"), "not null") - assert.Contains(t, isRead.Tag.Get("gorm"), "default:true") - assert.Equal(t, "is_read", isRead.Tag.Get("json")) + _, hasIsRead := threadType.FieldByName("IsRead") + assert.False(t, hasIsRead) - lastReadAt, ok := threadType.FieldByName("LastReadAt") + unreadCount, ok := threadType.FieldByName("UnreadCount") require.True(t, ok) - assert.Contains(t, lastReadAt.Tag.Get("gorm"), "not null") - assert.Contains(t, lastReadAt.Tag.Get("gorm"), "default:CURRENT_TIMESTAMP") - assert.Equal(t, "-", lastReadAt.Tag.Get("json")) + assert.Equal(t, "unread_count", unreadCount.Tag.Get("json")) + assert.Contains(t, unreadCount.Tag.Get("gorm"), "not null") + assert.Contains(t, unreadCount.Tag.Get("gorm"), "default:0") + + _, hasLastReadAt := threadType.FieldByName("LastReadAt") + assert.False(t, hasLastReadAt) } func TestMessageThreadContactDetailsAreTransientAndOmittedWhenNil(t *testing.T) { @@ -41,3 +41,13 @@ func TestMessageThreadContactDetailsAreTransientAndOmittedWhenNil(t *testing.T) require.NoError(t, json.Unmarshal(data, &payload)) assert.NotContains(t, payload, "contact_details") } + +func TestMessageThreadConversationHasCompositeUniqueIndex(t *testing.T) { + threadType := reflect.TypeOf(MessageThread{}) + + for _, fieldName := range []string{"UserID", "Owner", "Contact"} { + field, ok := threadType.FieldByName(fieldName) + require.True(t, ok) + assert.Contains(t, field.Tag.Get("gorm"), "uniqueIndex:idx_message_threads_conversation") + } +} diff --git a/api/pkg/handlers/message_thread_handler_test.go b/api/pkg/handlers/message_thread_handler_test.go index 56dfa0c7..7f4b53ad 100644 --- a/api/pkg/handlers/message_thread_handler_test.go +++ b/api/pkg/handlers/message_thread_handler_test.go @@ -25,7 +25,7 @@ import ( type messageThreadHandlerRepositoryStub struct{} -func (stub *messageThreadHandlerRepositoryStub) Store(context.Context, *entities.MessageThread) error { +func (stub *messageThreadHandlerRepositoryStub) Store(context.Context, repositories.MessageThreadStoreParams) error { return nil } @@ -75,7 +75,7 @@ func TestMessageThreadHandlerUpdate_ReturnsNotFoundWhenThreadIsMissing(t *testin handler.RegisterRoutes(app) messageThreadID := uuid.New() - req := httptest.NewRequest(http.MethodPut, "/v1/message-threads/"+messageThreadID.String(), bytes.NewBufferString(`{"is_read":true}`)) + req := httptest.NewRequest(http.MethodPut, "/v1/message-threads/"+messageThreadID.String(), bytes.NewBufferString(`{"unread_count":0}`)) req.Header.Set("Content-Type", "application/json") resp, err := app.Test(req, fiber.TestConfig{Timeout: time.Second}) @@ -90,6 +90,36 @@ func TestMessageThreadHandlerUpdate_ReturnsNotFoundWhenThreadIsMissing(t *testin require.Equal(t, "cannot find message thread with ID ["+messageThreadID.String()+"]", payload.Message) } +func TestMessageThreadHandlerUpdate_RejectsLegacyIsReadPayload(t *testing.T) { + logger := &messageThreadHandlerNoopLogger{} + tracer := telemetry.NewOtelLogger("test", logger) + service := services.NewMessageThreadService(logger, tracer, &messageThreadHandlerRepositoryStub{}, nil, nil, nil) + handler := NewMessageThreadHandler(logger, tracer, validators.NewMessageThreadHandlerValidator(logger, tracer), service) + + app := fiber.New() + app.Use(func(c fiber.Ctx) error { + c.Locals(middlewares.ContextKeyAuthUserID, entities.AuthContext{ID: entities.UserID("user-id"), Email: "user@example.com"}) + return c.Next() + }) + handler.RegisterRoutes(app) + + req := httptest.NewRequest(http.MethodPut, "/v1/message-threads/"+uuid.NewString(), bytes.NewBufferString(`{"is_read":true}`)) + req.Header.Set("Content-Type", "application/json") + + resp, err := app.Test(req, fiber.TestConfig{Timeout: time.Second}) + + require.NoError(t, err) + require.Equal(t, http.StatusUnprocessableEntity, resp.StatusCode) + + var payload struct { + Message string `json:"message"` + Data map[string][]string `json:"data"` + } + require.NoError(t, json.NewDecoder(resp.Body).Decode(&payload)) + require.Equal(t, "validation errors while updating message thread", payload.Message) + require.Equal(t, []string{"at least one of is_archived or unread_count is required"}, payload.Data["payload"]) +} + type messageThreadHandlerNoopLogger struct{} var _ telemetry.Logger = (*messageThreadHandlerNoopLogger)(nil) diff --git a/api/pkg/listeners/message_thread_listener.go b/api/pkg/listeners/message_thread_listener.go index f6d2c52c..2b1c9734 100644 --- a/api/pkg/listeners/message_thread_listener.go +++ b/api/pkg/listeners/message_thread_listener.go @@ -210,15 +210,14 @@ func (listener *MessageThreadListener) OnMessagePhoneReceived(ctx context.Contex } updateParams := services.MessageThreadUpdateParams{ - Owner: payload.Owner, - Contact: payload.Contact, - Timestamp: payload.Timestamp, - UserID: payload.UserID, - Status: entities.MessageStatusReceived, - Content: payload.Content, - MessageID: payload.MessageID, - MarkAsUnread: true, - EventTimestamp: event.Time(), + Owner: payload.Owner, + Contact: payload.Contact, + Timestamp: payload.Timestamp, + UserID: payload.UserID, + Status: entities.MessageStatusReceived, + Content: payload.Content, + MessageID: payload.MessageID, + CountAsUnread: true, } if err := listener.service.UpdateThread(ctx, updateParams); err != nil { @@ -239,15 +238,14 @@ func (listener *MessageThreadListener) OnMessageCallMissed(ctx context.Context, } params := services.MessageThreadUpdateParams{ - Owner: payload.Owner, - Contact: payload.Contact, - UserID: payload.UserID, - Status: entities.MessageStatusReceived, - Timestamp: payload.Timestamp, - Content: "Missed phone call", - MessageID: payload.MessageID, - MarkAsUnread: true, - EventTimestamp: event.Time(), + Owner: payload.Owner, + Contact: payload.Contact, + UserID: payload.UserID, + Status: entities.MessageStatusReceived, + Timestamp: payload.Timestamp, + Content: "Missed phone call", + MessageID: payload.MessageID, + CountAsUnread: true, } if err := listener.service.UpdateThread(ctx, params); err != nil { return listener.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot update thread for missed call [%s] on event [%s]", payload.MessageID, event.ID())) diff --git a/api/pkg/listeners/message_thread_listener_test.go b/api/pkg/listeners/message_thread_listener_test.go index f0410dfa..0fadd105 100644 --- a/api/pkg/listeners/message_thread_listener_test.go +++ b/api/pkg/listeners/message_thread_listener_test.go @@ -34,8 +34,7 @@ func TestMessageThreadListenerMarksInboundMessageUnread(t *testing.T) { err := routes[events.EventTypeMessagePhoneReceived](context.Background(), event) require.NoError(t, err) - assert.True(t, repository.activity.MarkAsUnread) - assert.Equal(t, event.Time(), repository.activity.EventTimestamp) + assert.True(t, repository.activity.CountAsUnread) } func TestMessageThreadListenerMarksMissedCallUnread(t *testing.T) { @@ -57,9 +56,38 @@ func TestMessageThreadListenerMarksMissedCallUnread(t *testing.T) { err := routes[events.MessageCallMissed](context.Background(), event) require.NoError(t, err) - assert.True(t, repository.activity.MarkAsUnread) + assert.True(t, repository.activity.CountAsUnread) assert.Equal(t, "Missed phone call", repository.activity.Content) - assert.Equal(t, event.Time(), repository.activity.EventTimestamp) +} + +func TestMessageThreadListenerDeletesNonLastUnreadMessage(t *testing.T) { + repository, routes := newMessageThreadListenerForTest() + deletedMessageID := uuid.New() + previousMessageID := uuid.New() + previousStatus := entities.MessageStatus(entities.MessageStatusDelivered) + previousContent := "previous" + event := cloudevents.NewEvent() + event.SetID(uuid.NewString()) + event.SetSource("/v1/messages/deleted") + event.SetType(events.MessageAPIDeleted) + require.NoError(t, event.SetData(cloudevents.ApplicationJSON, events.MessageAPIDeletedPayload{ + MessageID: deletedMessageID, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + PreviousMessageID: &previousMessageID, + PreviousMessageStatus: &previousStatus, + PreviousMessageContent: &previousContent, + })) + + err := routes[events.MessageAPIDeleted](context.Background(), event) + + require.NoError(t, err) + assert.Equal(t, deletedMessageID, repository.deletedUpdate.DeletedMessageID) + assert.Equal(t, "+18005550199", repository.deletedUpdate.Owner) + assert.Equal(t, "+18005550100", repository.deletedUpdate.Contact) + require.NotNil(t, repository.deletedUpdate.LastMessageID) + assert.Equal(t, previousMessageID, *repository.deletedUpdate.LastMessageID) } func newMessageThreadListenerForTest() (*listenerMessageThreadRepository, map[string]events.EventListener) { diff --git a/api/pkg/listeners/read_receipts_test_helpers_test.go b/api/pkg/listeners/read_receipts_test_helpers_test.go index 60cd27e6..e812a297 100644 --- a/api/pkg/listeners/read_receipts_test_helpers_test.go +++ b/api/pkg/listeners/read_receipts_test_helpers_test.go @@ -24,10 +24,12 @@ func (logger *noopListenerLogger) Fatal(error) { func (logger *noopListenerLogger) Printf(string, ...interface{}) {} type listenerMessageThreadRepository struct { - activity repositories.MessageThreadActivityUpdate + activity repositories.MessageThreadActivityUpdate + deletedUpdate repositories.MessageThreadDeletedUpdate + thread *entities.MessageThread } -func (repository *listenerMessageThreadRepository) Store(context.Context, *entities.MessageThread) error { +func (repository *listenerMessageThreadRepository) Store(context.Context, repositories.MessageThreadStoreParams) error { return nil } @@ -40,12 +42,21 @@ func (repository *listenerMessageThreadRepository) UpdateStatus(_ context.Contex return &entities.MessageThread{ID: threadID}, nil } -func (repository *listenerMessageThreadRepository) UpdateAfterDeletedMessage(context.Context, repositories.MessageThreadDeletedUpdate) error { +func (repository *listenerMessageThreadRepository) UpdateAfterDeletedMessage(_ context.Context, params repositories.MessageThreadDeletedUpdate) error { + repository.deletedUpdate = params return nil } func (repository *listenerMessageThreadRepository) LoadByOwnerContact(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { - return &entities.MessageThread{ID: uuid.New()}, nil + if repository.thread != nil { + return repository.thread, nil + } + return &entities.MessageThread{ + ID: uuid.New(), + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + }, nil } func (repository *listenerMessageThreadRepository) Load(context.Context, entities.UserID, uuid.UUID) (*entities.MessageThread, error) { diff --git a/api/pkg/listeners/websocket_listener.go b/api/pkg/listeners/websocket_listener.go index bdc43f93..0215418d 100644 --- a/api/pkg/listeners/websocket_listener.go +++ b/api/pkg/listeners/websocket_listener.go @@ -19,6 +19,10 @@ type WebsocketListener struct { client *pusher.Client } +type websocketMessagePayload struct { + MessageID string `json:"message_id"` +} + // NewWebsocketListener creates a new instance of WebsocketListener func NewWebsocketListener( logger telemetry.Logger, @@ -49,7 +53,9 @@ func (listener *WebsocketListener) onMessageCallMissed(ctx context.Context, even return listener.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot decode [%s] into [%T]", event.Data(), payload)) } - if err := listener.client.Trigger(payload.UserID.String(), event.Type(), event.ID()); err != nil { + if err := listener.client.Trigger(payload.UserID.String(), event.Type(), websocketMessagePayload{ + MessageID: payload.MessageID.String(), + }); err != nil { return listener.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot trigger websocket [%s] event with ID [%s] for user with ID [%s]", event.Type(), event.ID(), payload.UserID)) } return nil @@ -82,7 +88,9 @@ func (listener *WebsocketListener) onMessagePhoneReceived(ctx context.Context, e return listener.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot decode [%s] into [%T]", event.Data(), payload)) } - if err := listener.client.Trigger(payload.UserID.String(), event.Type(), event.ID()); err != nil { + if err := listener.client.Trigger(payload.UserID.String(), event.Type(), websocketMessagePayload{ + MessageID: payload.MessageID.String(), + }); err != nil { return listener.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot trigger websocket [%s] event with ID [%s] for user with ID [%s]", event.Type(), event.ID(), payload.UserID)) } diff --git a/api/pkg/listeners/websocket_listener_test.go b/api/pkg/listeners/websocket_listener_test.go index fdaf1c6d..903e5799 100644 --- a/api/pkg/listeners/websocket_listener_test.go +++ b/api/pkg/listeners/websocket_listener_test.go @@ -1,12 +1,22 @@ package listeners import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" "testing" + "github.com/NdoleStudio/httpsms/pkg/entities" "github.com/NdoleStudio/httpsms/pkg/events" "github.com/NdoleStudio/httpsms/pkg/telemetry" + cloudevents "github.com/cloudevents/sdk-go/v2" + "github.com/google/uuid" "github.com/pusher/pusher-http-go/v5" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestWebsocketListenerRegistersMissedCalls(t *testing.T) { @@ -16,3 +26,74 @@ func TestWebsocketListenerRegistersMissedCalls(t *testing.T) { assert.Contains(t, routes, events.MessageCallMissed) } + +func TestWebsocketListenerPublishesReceivedMessageID(t *testing.T) { + userID := "user-id" + messageID := uuid.New() + event := cloudevents.NewEvent() + event.SetID(uuid.NewString()) + event.SetType(events.EventTypeMessagePhoneReceived) + require.NoError(t, event.SetData(cloudevents.ApplicationJSON, events.MessagePhoneReceivedPayload{ + MessageID: messageID, + UserID: entities.UserID(userID), + })) + + payload := captureWebsocketPayload(t, event, userID) + + assert.Equal(t, messageID.String(), payload.MessageID) +} + +func TestWebsocketListenerPublishesMissedCallMessageID(t *testing.T) { + userID := "user-id" + messageID := uuid.New() + event := cloudevents.NewEvent() + event.SetID(uuid.NewString()) + event.SetType(events.MessageCallMissed) + require.NoError(t, event.SetData(cloudevents.ApplicationJSON, events.MessageCallMissedPayload{ + MessageID: messageID, + UserID: entities.UserID(userID), + })) + + payload := captureWebsocketPayload(t, event, userID) + + assert.Equal(t, messageID.String(), payload.MessageID) +} + +func captureWebsocketPayload(t *testing.T, event cloudevents.Event, userID string) websocketMessagePayload { + t.Helper() + + requestBody := make(chan []byte, 1) + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + body, err := io.ReadAll(request.Body) + require.NoError(t, err) + requestBody <- body + writer.Header().Set("Content-Type", "application/json") + _, err = writer.Write([]byte(`{}`)) + require.NoError(t, err) + })) + t.Cleanup(server.Close) + + logger := &noopListenerLogger{} + tracer := telemetry.NewOtelLogger("test", logger) + client := &pusher.Client{ + AppID: "app-id", + Key: "key", + Secret: "secret", + Host: strings.TrimPrefix(server.URL, "http://"), + HTTPClient: server.Client(), + } + _, routes := NewWebsocketListener(logger, tracer, client) + + require.NoError(t, routes[event.Type()](context.Background(), event)) + + var trigger struct { + Channels []string `json:"channels"` + Data string `json:"data"` + } + require.NoError(t, json.Unmarshal(<-requestBody, &trigger)) + assert.Equal(t, []string{userID}, trigger.Channels) + + var payload websocketMessagePayload + require.NoError(t, json.Unmarshal([]byte(trigger.Data), &payload)) + return payload +} diff --git a/api/pkg/repositories/gorm_message_thread_repository.go b/api/pkg/repositories/gorm_message_thread_repository.go index 09241804..ffed79f1 100644 --- a/api/pkg/repositories/gorm_message_thread_repository.go +++ b/api/pkg/repositories/gorm_message_thread_repository.go @@ -35,32 +35,24 @@ func NewGormMessageThreadRepository( } } -func messageThreadActivityUpdates(params MessageThreadActivityUpdate) map[string]any { - updates := map[string]any{ - "order_timestamp": params.Timestamp, - "last_message_id": params.MessageID, - "last_message_content": params.Content, - "status": params.Status, - } - if params.Unarchive { - updates["is_archived"] = false +func messageThreadDeletedUpdates(params MessageThreadDeletedUpdate) (map[string]any, error) { + if params.LastMessageContent == nil { + return nil, stacktrace.NewErrorf( + "last message content is required when replacing deleted message [%s]", + params.DeletedMessageID, + ) } - if params.MarkAsUnread { - updates["is_read"] = gorm.Expr( - "CASE WHEN last_read_at < ? THEN ? ELSE is_read END", - params.EventTimestamp, - false, + if params.LastMessageStatus == nil { + return nil, stacktrace.NewErrorf( + "last message status is required when replacing deleted message [%s]", + params.DeletedMessageID, ) } - return updates -} - -func messageThreadDeletedUpdates(params MessageThreadDeletedUpdate) map[string]any { return map[string]any{ "last_message_id": params.LastMessageID, "last_message_content": params.LastMessageContent, - "status": params.LastMessageStatus, - } + "status": *params.LastMessageStatus, + }, nil } func messageThreadStatusUpdates(params MessageThreadStatusUpdate) map[string]any { @@ -68,15 +60,71 @@ func messageThreadStatusUpdates(params MessageThreadStatusUpdate) map[string]any if params.IsArchived != nil { updates["is_archived"] = *params.IsArchived } - if params.IsRead != nil { - updates["is_read"] = *params.IsRead - if *params.IsRead { - updates["last_read_at"] = params.ReadAt - } + if params.UnreadCount != nil { + updates["unread_count"] = *params.UnreadCount } return updates } +func applyMessageThreadActivity(db *gorm.DB, params MessageThreadActivityUpdate) (*entities.MessageThread, error) { + const activityCondition = "order_timestamp <= ? AND (status <> ? OR last_message_id IS DISTINCT FROM ?)" + + conditionArgs := []any{ + params.Timestamp, + entities.MessageStatusDelivered, + params.MessageID, + } + updates := map[string]any{ + "order_timestamp": gorm.Expr( + "CASE WHEN "+activityCondition+" THEN ? ELSE order_timestamp END", + append(conditionArgs, params.Timestamp)..., + ), + "last_message_id": gorm.Expr( + "CASE WHEN "+activityCondition+" THEN ? ELSE last_message_id END", + append(conditionArgs, params.MessageID)..., + ), + "last_message_content": gorm.Expr( + "CASE WHEN "+activityCondition+" THEN ? ELSE last_message_content END", + append(conditionArgs, params.Content)..., + ), + "status": gorm.Expr( + "CASE WHEN "+activityCondition+" THEN ? ELSE status END", + append(conditionArgs, params.Status)..., + ), + } + if params.Unarchive { + updates["is_archived"] = false + } + if params.CountAsUnread { + updates["unread_count"] = gorm.Expr("unread_count + ?", 1) + } + + thread := new(entities.MessageThread) + result := db. + Model(thread). + Clauses(clause.Returning{}). + Where("user_id = ?", params.UserID). + Where("id = ?", params.MessageThreadID). + Updates(updates) + if result.Error != nil { + return nil, stacktrace.Propagatef( + result.Error, + "cannot update message activity for thread [%s] and user [%s]", + params.MessageThreadID, + params.UserID, + ) + } + if result.RowsAffected == 0 { + return nil, stacktrace.PropagateWithCodef( + gorm.ErrRecordNotFound, + ErrCodeNotFound, + "thread with id [%s] not found", + params.MessageThreadID, + ) + } + return thread, nil +} + func (repository *gormMessageThreadRepository) DeleteAllForUser(ctx context.Context, userID entities.UserID) error { ctx, span := repository.tracer.Start(ctx) defer span.End() @@ -106,44 +154,78 @@ func (repository *gormMessageThreadRepository) UpdateAfterDeletedMessage(ctx con ctx, span := repository.tracer.Start(ctx) defer span.End() - result := repository.db.WithContext(ctx). - Model(&entities.MessageThread{}). - Where("user_id = ?", params.UserID). - Where("id = ?", params.MessageThreadID). - Updates(messageThreadDeletedUpdates(params)) - if result.Error != nil { - return repository.tracer.WrapErrorSpan(span, stacktrace.Propagatef(result.Error, "cannot update deleted-message metadata for thread [%s]", params.MessageThreadID)) + thread, err := repository.LoadByOwnerContact(ctx, params.UserID, params.Owner, params.Contact) + if stacktrace.GetCode(err) == ErrCodeNotFound { + return nil + } + if err != nil { + return repository.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot load thread after deleting message [%s]", params.DeletedMessageID)) + } + if !thread.HasLastMessage(params.DeletedMessageID) { + return nil + } + if params.LastMessageID == nil { + if err := repository.db.WithContext(ctx).Session(&gorm.Session{SkipDefaultTransaction: true}). + Where("user_id = ?", params.UserID). + Where("id = ?", thread.ID). + Delete(&entities.MessageThread{}). + Error; err != nil { + return repository.tracer.WrapErrorSpan(span, stacktrace.Propagatef( + err, + "cannot delete message thread [%s] for user [%s] after deleting final message [%s]", + thread.ID, + params.UserID, + params.DeletedMessageID, + )) + } + return nil } + updates, err := messageThreadDeletedUpdates(params) + if err != nil { + return repository.tracer.WrapErrorSpan(span, err) + } + if err := repository.db.WithContext(ctx).Session(&gorm.Session{SkipDefaultTransaction: true}). + Model(thread). + Where("user_id = ?", params.UserID). + Where("id = ?", thread.ID). + Updates(updates). + Error; err != nil { + return repository.tracer.WrapErrorSpan(span, stacktrace.Propagatef( + err, + "cannot update deleted-message metadata for thread [%s] and user [%s]", + thread.ID, + params.UserID, + )) + } return nil } // Store a new entities.MessageThread -func (repository *gormMessageThreadRepository) Store(ctx context.Context, thread *entities.MessageThread) error { +func (repository *gormMessageThreadRepository) Store(ctx context.Context, params MessageThreadStoreParams) error { ctx, span := repository.tracer.Start(ctx) defer span.End() - isRead := thread.IsRead - err := repository.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { - result := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(thread) - thread.IsRead = isRead - if result.Error != nil { - return result.Error - } - if result.RowsAffected == 0 || isRead { - return nil - } - - return tx.Model(&entities.MessageThread{}). - Where("user_id = ?", thread.UserID). - Where("id = ?", thread.ID). - UpdateColumn("is_read", false). - Error - }) - if err != nil { - return repository.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot save message thread with ID [%s]", thread.ID)) + candidate := *params.Thread + db := repository.db.WithContext(ctx).Session(&gorm.Session{SkipDefaultTransaction: true}) + onConflict := clause.OnConflict{ + Columns: []clause.Column{ + {Name: "user_id"}, + {Name: "owner"}, + {Name: "contact"}, + }, + DoNothing: !params.CountAsUnread, + } + if params.CountAsUnread { + onConflict.DoUpdates = clause.Assignments(map[string]any{ + "unread_count": gorm.Expr("unread_count + ?", 1), + }) } + result := db.Clauses(onConflict).Create(&candidate) + if result.Error != nil { + return repository.tracer.WrapErrorSpan(span, stacktrace.Propagatef(result.Error, "cannot insert message thread with ID [%s]", params.Thread.ID)) + } return nil } @@ -152,22 +234,26 @@ func (repository *gormMessageThreadRepository) UpdateActivity(ctx context.Contex ctx, span := repository.tracer.Start(ctx) defer span.End() - result := repository.db.WithContext(ctx). - Model(&entities.MessageThread{}). - Where("user_id = ?", params.UserID). - Where("id = ?", params.MessageThreadID). - Updates(messageThreadActivityUpdates(params)) - if result.Error != nil { - return repository.tracer.WrapErrorSpan(span, stacktrace.Propagatef(result.Error, "cannot update message activity for thread [%s]", params.MessageThreadID)) - } - if result.RowsAffected == 0 { - return repository.tracer.WrapErrorSpan(span, stacktrace.PropagateWithCodef(gorm.ErrRecordNotFound, ErrCodeNotFound, "thread with id [%s] not found", params.MessageThreadID)) + _, err := applyMessageThreadActivity( + repository.db.WithContext(ctx).Session(&gorm.Session{SkipDefaultTransaction: true}), + params, + ) + if err != nil { + return repository.tracer.WrapErrorSpan( + span, + stacktrace.Propagatef( + err, + "cannot update message activity for thread [%s] and user [%s]", + params.MessageThreadID, + params.UserID, + ), + ) } return nil } -// UpdateStatus persists archive/read status fields for a thread +// UpdateStatus persists archive/unread status fields for a thread func (repository *gormMessageThreadRepository) UpdateStatus( ctx context.Context, userID entities.UserID, @@ -177,18 +263,41 @@ func (repository *gormMessageThreadRepository) UpdateStatus( ctx, span := repository.tracer.Start(ctx) defer span.End() - thread := new(entities.MessageThread) - result := repository.db.WithContext(ctx). - Model(thread). - Clauses(clause.Returning{}). - Where("user_id = ?", userID). - Where("id = ?", messageThreadID). - Updates(messageThreadStatusUpdates(params)) - if result.Error != nil { - return nil, repository.tracer.WrapErrorSpan(span, stacktrace.Propagatef(result.Error, "cannot update status for thread [%s] and user [%s]", messageThreadID, userID)) + if params.UnreadCount != nil && *params.UnreadCount != 0 { + return nil, repository.tracer.WrapErrorSpan(span, stacktrace.NewErrorf( + "cannot set unread count to [%d] for thread [%s] and user [%s]: only zero is supported", + *params.UnreadCount, + messageThreadID, + userID, + )) } - if result.RowsAffected == 0 { - return nil, repository.tracer.WrapErrorSpan(span, stacktrace.PropagateWithCodef(gorm.ErrRecordNotFound, ErrCodeNotFound, "thread with id [%s] not found for user with ID [%s]", messageThreadID, userID)) + + thread, err := repository.Load(ctx, userID, messageThreadID) + if err != nil { + return nil, repository.tracer.WrapErrorSpan(span, err) + } + + updates := messageThreadStatusUpdates(params) + if len(updates) != 0 { + if err := repository.db.WithContext(ctx).Session(&gorm.Session{SkipDefaultTransaction: true}). + Model(thread). + Where("user_id = ?", userID). + Where("id = ?", messageThreadID). + Updates(updates). + Error; err != nil { + return nil, repository.tracer.WrapErrorSpan(span, stacktrace.Propagatef( + err, + "cannot update status for thread [%s] and user [%s]", + messageThreadID, + userID, + )) + } + } + if params.IsArchived != nil { + thread.IsArchived = *params.IsArchived + } + if params.UnreadCount != nil { + thread.UnreadCount = *params.UnreadCount } return thread, nil diff --git a/api/pkg/repositories/gorm_message_thread_repository_test.go b/api/pkg/repositories/gorm_message_thread_repository_test.go index 31dc43ed..91c18a92 100644 --- a/api/pkg/repositories/gorm_message_thread_repository_test.go +++ b/api/pkg/repositories/gorm_message_thread_repository_test.go @@ -5,12 +5,14 @@ import ( "database/sql" "database/sql/driver" "errors" + "io" "strings" "testing" "time" "github.com/NdoleStudio/httpsms/pkg/entities" "github.com/NdoleStudio/httpsms/pkg/telemetry" + "github.com/NdoleStudio/stacktrace" "github.com/google/uuid" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -25,7 +27,14 @@ type messageThreadTestStatement struct { } type messageThreadTestConnPool struct { - statements []messageThreadTestStatement + statements []messageThreadTestStatement + thread *entities.MessageThread + rowsAffected func(query string) int64 + execError func(query string) error + queryDB *sql.DB + begins int + commits int + rollbacks int } func (messageThreadTestConnPool) PrepareContext(context.Context, string) (*sql.Stmt, error) { @@ -37,11 +46,23 @@ func (pool *messageThreadTestConnPool) ExecContext(_ context.Context, query stri query: query, args: append([]any(nil), args...), }) + if pool.execError != nil { + if err := pool.execError(query); err != nil { + return nil, err + } + } + if pool.rowsAffected != nil { + return driver.RowsAffected(pool.rowsAffected(query)), nil + } return driver.RowsAffected(1), nil } -func (messageThreadTestConnPool) QueryContext(context.Context, string, ...any) (*sql.Rows, error) { - return nil, errors.New("unexpected QueryContext") +func (pool *messageThreadTestConnPool) QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + pool.statements = append(pool.statements, messageThreadTestStatement{ + query: query, + args: append([]any(nil), args...), + }) + return pool.queryDB.QueryContext(ctx, query, args...) } func (messageThreadTestConnPool) QueryRowContext(context.Context, string, ...any) *sql.Row { @@ -49,14 +70,134 @@ func (messageThreadTestConnPool) QueryRowContext(context.Context, string, ...any } func (pool *messageThreadTestConnPool) BeginTx(context.Context, *sql.TxOptions) (gorm.ConnPool, error) { - return pool, nil + pool.begins++ + return &messageThreadTestTxPool{pool: pool}, nil +} + +type messageThreadTestTxPool struct { + pool *messageThreadTestConnPool +} + +func (tx *messageThreadTestTxPool) PrepareContext(ctx context.Context, query string) (*sql.Stmt, error) { + return tx.pool.PrepareContext(ctx, query) +} + +func (tx *messageThreadTestTxPool) ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) { + return tx.pool.ExecContext(ctx, query, args...) +} + +func (tx *messageThreadTestTxPool) QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + return tx.pool.QueryContext(ctx, query, args...) +} + +func (tx *messageThreadTestTxPool) QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row { + return tx.pool.QueryRowContext(ctx, query, args...) +} + +func (tx *messageThreadTestTxPool) Commit() error { + tx.pool.commits++ + return nil +} + +func (tx *messageThreadTestTxPool) Rollback() error { + tx.pool.rollbacks++ + return nil +} + +type messageThreadRowsConnector struct { + pool *messageThreadTestConnPool +} + +func (connector *messageThreadRowsConnector) Connect(context.Context) (driver.Conn, error) { + return &messageThreadRowsConn{pool: connector.pool}, nil +} + +func (*messageThreadRowsConnector) Driver() driver.Driver { + return messageThreadRowsDriver{} +} + +type messageThreadRowsDriver struct{} + +func (messageThreadRowsDriver) Open(string) (driver.Conn, error) { + return nil, errors.New("message thread test driver requires a connector") +} + +type messageThreadRowsConn struct { + pool *messageThreadTestConnPool +} + +func (*messageThreadRowsConn) Prepare(string) (driver.Stmt, error) { + return nil, errors.New("unexpected Prepare") } -func (*messageThreadTestConnPool) Commit() error { +func (*messageThreadRowsConn) Close() error { return nil } -func (*messageThreadTestConnPool) Rollback() error { +func (*messageThreadRowsConn) Begin() (driver.Tx, error) { + return nil, errors.New("unexpected Begin") +} + +func (conn *messageThreadRowsConn) QueryContext(_ context.Context, query string, _ []driver.NamedValue) (driver.Rows, error) { + rows := &messageThreadDriverRows{ + columns: []string{ + "id", + "user_id", + "owner", + "contact", + "is_archived", + "unread_count", + "last_message_id", + "last_message_content", + "status", + "order_timestamp", + }, + } + if conn.pool.thread != nil { + var lastMessageID driver.Value + if conn.pool.thread.LastMessageID != nil { + lastMessageID = conn.pool.thread.LastMessageID.String() + } + var lastMessageContent driver.Value + if conn.pool.thread.LastMessageContent != nil { + lastMessageContent = *conn.pool.thread.LastMessageContent + } + rows.values = []driver.Value{ + conn.pool.thread.ID.String(), + string(conn.pool.thread.UserID), + conn.pool.thread.Owner, + conn.pool.thread.Contact, + conn.pool.thread.IsArchived, + int64(conn.pool.thread.UnreadCount), + lastMessageID, + lastMessageContent, + string(conn.pool.thread.Status), + conn.pool.thread.OrderTimestamp, + } + } + return rows, nil +} + +type messageThreadDriverRows struct { + columns []string + values []driver.Value + read bool +} + +func (rows *messageThreadDriverRows) Columns() []string { + return rows.columns +} + +func (*messageThreadDriverRows) Close() error { + return nil +} + +func (rows *messageThreadDriverRows) Next(dest []driver.Value) error { + if rows.read || rows.values == nil { + return io.EOF + } + copy(dest, rows.values) + rows.read = true return nil } @@ -75,8 +216,14 @@ func (logger *messageThreadTestLogger) Debug(string) func (logger *messageThreadTestLogger) Fatal(error) {} func (logger *messageThreadTestLogger) Printf(string, ...interface{}) {} -func TestMessageThreadStorePreservesExplicitUnreadState(t *testing.T) { - pool := &messageThreadTestConnPool{} +func newMessageThreadTestRepository(t *testing.T, pool *messageThreadTestConnPool) *gormMessageThreadRepository { + t.Helper() + + pool.queryDB = sql.OpenDB(&messageThreadRowsConnector{pool: pool}) + t.Cleanup(func() { + require.NoError(t, pool.queryDB.Close()) + }) + db, err := gorm.Open( postgres.New(postgres.Config{ Conn: pool, @@ -87,87 +234,410 @@ func TestMessageThreadStorePreservesExplicitUnreadState(t *testing.T) { require.NoError(t, err) logger := &messageThreadTestLogger{} - repository := NewGormMessageThreadRepository(logger, telemetry.NewOtelLogger("test", logger), db) + repository, ok := NewGormMessageThreadRepository(logger, telemetry.NewOtelLogger("test", logger), db).(*gormMessageThreadRepository) + require.True(t, ok) + return repository +} + +func messageThreadStatementIndex(pool *messageThreadTestConnPool, fragment string) int { + for index, statement := range pool.statements { + if strings.Contains(statement.query, fragment) { + return index + } + } + return -1 +} + +func messageThreadStatementCount(pool *messageThreadTestConnPool, fragment string) int { + count := 0 + for _, statement := range pool.statements { + if strings.Contains(statement.query, fragment) { + count++ + } + } + return count +} + +func messageThreadStatementIndexAfter(pool *messageThreadTestConnPool, fragment string, after int) int { + for index := after + 1; index < len(pool.statements); index++ { + if strings.Contains(pool.statements[index].query, fragment) { + return index + } + } + return -1 +} + +func TestMessageThreadUnreadStoreUsesThreadCounterOnly(t *testing.T) { + threadID := uuid.New() + messageID := uuid.New() + pool := &messageThreadTestConnPool{} + repository := newMessageThreadTestRepository(t, pool) thread := &entities.MessageThread{ - ID: uuid.New(), - IsRead: false, + ID: threadID, + UserID: entities.UserID("user-id"), + UnreadCount: 1, + LastMessageID: &messageID, } - require.NoError(t, repository.Store(context.Background(), thread)) - assert.False(t, thread.IsRead) + require.NoError(t, repository.Store(context.Background(), MessageThreadStoreParams{ + Thread: thread, + CountAsUnread: true, + })) + require.Zero(t, pool.begins) + require.Zero(t, pool.commits) + require.Zero(t, pool.rollbacks) - require.NotEmpty(t, pool.statements) - update := pool.statements[len(pool.statements)-1] - assert.True(t, strings.HasPrefix(update.query, `UPDATE "message_threads"`)) - assert.Contains(t, update.query, `"is_read"=$1`) - assert.Contains(t, update.args, false) + threadInsert := messageThreadStatementIndex(pool, `INSERT INTO "message_threads"`) + require.NotEqual(t, -1, threadInsert) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_unread_items`)) } -func TestMessageThreadActivityUpdatesOwnOnlyMessageColumns(t *testing.T) { +func TestMessageThreadStoreConflictOnlyIncrementsUnreadCount(t *testing.T) { + winnerThreadID := uuid.New() + losingThreadID := uuid.New() messageID := uuid.New() - updates := messageThreadActivityUpdates(MessageThreadActivityUpdate{ - Timestamp: time.Date(2026, 7, 18, 7, 0, 0, 0, time.UTC), - MessageID: messageID, - Content: "hello", - Status: entities.MessageStatusReceived, + userID := entities.UserID("user-id") + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: winnerThreadID, + UserID: userID, + UnreadCount: 1, + }, + rowsAffected: func(query string) int64 { + if strings.Contains(query, `INSERT INTO "message_threads"`) { + return 0 + } + return 1 + }, + } + repository := newMessageThreadTestRepository(t, pool) + content := "losing message" + eventTimestamp := time.Date(2026, 8, 21, 10, 0, 1, 0, time.UTC) + thread := &entities.MessageThread{ + ID: losingThreadID, + UserID: userID, + Owner: "+18005550199", + Contact: "+18005550100", + UnreadCount: 1, + LastMessageID: &messageID, + LastMessageContent: &content, + Status: entities.MessageStatusReceived, + OrderTimestamp: eventTimestamp, + } + + err := repository.Store(context.Background(), MessageThreadStoreParams{ + Thread: thread, + CountAsUnread: true, }) - assert.Equal(t, map[string]any{ - "order_timestamp": time.Date(2026, 7, 18, 7, 0, 0, 0, time.UTC), - "last_message_id": messageID, - "last_message_content": "hello", - "status": entities.MessageStatus(entities.MessageStatusReceived), - }, updates) - assert.NotContains(t, updates, "is_read") - assert.NotContains(t, updates, "is_archived") - assert.NotContains(t, updates, "last_read_at") + require.NoError(t, err) + insert := messageThreadStatementIndex(pool, `INSERT INTO "message_threads"`) + require.NotEqual(t, -1, insert) + assert.Contains(t, pool.statements[insert].query, `ON CONFLICT ("user_id","owner","contact") DO UPDATE`) + assert.Contains(t, pool.statements[insert].query, `"unread_count"=unread_count +`) + assert.Equal(t, 1, messageThreadStatementCount(pool, `INSERT INTO "message_threads"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `SELECT * FROM "message_threads"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `UPDATE "message_threads"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "FOR UPDATE")) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_unread_items`)) } -func TestUpdateActivityMarksUnreadWithOneQuery(t *testing.T) { - pool := &messageThreadTestConnPool{} - db, err := gorm.Open( - postgres.New(postgres.Config{ - Conn: pool, - WithoutReturning: true, - }), - &gorm.Config{DisableAutomaticPing: true}, - ) +func TestMessageThreadStoreConflictStaleActivityPreservesWinnerMetadataAndCountsUnread(t *testing.T) { + winnerThreadID := uuid.New() + losingThreadID := uuid.New() + winnerMessageID := uuid.New() + losingMessageID := uuid.New() + userID := entities.UserID("user-id") + winnerContent := "winner" + winnerTimestamp := time.Date(2026, 8, 21, 10, 0, 2, 0, time.UTC) + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: winnerThreadID, + UserID: userID, + UnreadCount: 1, + LastMessageID: &winnerMessageID, + LastMessageContent: &winnerContent, + Status: entities.MessageStatusDelivered, + OrderTimestamp: winnerTimestamp, + }, + rowsAffected: func(query string) int64 { + if strings.Contains(query, `INSERT INTO "message_threads"`) { + return 0 + } + return 1 + }, + } + repository := newMessageThreadTestRepository(t, pool) + losingContent := "loser" + losingTimestamp := winnerTimestamp.Add(-time.Second) + + err := repository.Store(context.Background(), MessageThreadStoreParams{ + Thread: &entities.MessageThread{ + ID: losingThreadID, + UserID: userID, + Owner: "+18005550199", + Contact: "+18005550100", + UnreadCount: 1, + LastMessageID: &losingMessageID, + LastMessageContent: &losingContent, + Status: entities.MessageStatusReceived, + OrderTimestamp: losingTimestamp, + }, + CountAsUnread: true, + }) + require.NoError(t, err) + insert := messageThreadStatementIndex(pool, `INSERT INTO "message_threads"`) + require.NotEqual(t, -1, insert) + assert.Contains(t, pool.statements[insert].query, `DO UPDATE SET "unread_count"=unread_count +`) + assert.NotContains(t, pool.statements[insert].query, `DO UPDATE SET "order_timestamp"`) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `SELECT * FROM "message_threads"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `UPDATE "message_threads"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "FOR UPDATE")) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_unread_items`)) +} - logger := &messageThreadTestLogger{} - repository := NewGormMessageThreadRepository(logger, telemetry.NewOtelLogger("test", logger), db) +func TestMessageThreadActivityIncrementsWithoutLocking(t *testing.T) { + threadID := uuid.New() + messageID := uuid.New() + userID := entities.UserID("user-id") + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: userID, + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ + MessageThreadID: threadID, + UserID: userID, + Timestamp: time.Date(2026, 7, 19, 10, 0, 0, 0, time.UTC), + MessageID: messageID, + Content: "hello", + Status: entities.MessageStatusReceived, + CountAsUnread: true, + }) + + require.NoError(t, err) + require.Zero(t, pool.begins) + require.Zero(t, pool.commits) + require.Zero(t, pool.rollbacks) + + update := messageThreadStatementIndex(pool, `UPDATE "message_threads"`) + require.NotEqual(t, -1, update) + assert.Contains(t, pool.statements[update].query, `"order_timestamp"`) + assert.Contains(t, pool.statements[update].query, `"unread_count"=unread_count +`) + assert.Equal(t, 1, messageThreadStatementCount(pool, `UPDATE "message_threads"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `SELECT * FROM "message_threads"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "FOR UPDATE")) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_unread_items`)) +} + +func TestMessageThreadStaleActivityPreservesPreviewAndCountsUnread(t *testing.T) { + threadID := uuid.New() + currentMessageID := uuid.New() + incomingMessageID := uuid.New() + userID := entities.UserID("user-id") + content := "current" + currentTimestamp := time.Date(2026, 8, 21, 10, 0, 2, 0, time.UTC) + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: userID, + LastMessageID: ¤tMessageID, + LastMessageContent: &content, + Status: entities.MessageStatusDelivered, + OrderTimestamp: currentTimestamp, + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ + MessageThreadID: threadID, + UserID: userID, + Timestamp: currentTimestamp.Add(-time.Second), + MessageID: incomingMessageID, + Content: "stale", + Status: entities.MessageStatusReceived, + CountAsUnread: true, + }) - err = repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ - MessageThreadID: uuid.New(), + require.NoError(t, err) + update := messageThreadStatementIndex(pool, `UPDATE "message_threads"`) + require.NotEqual(t, -1, update) + assert.Contains(t, pool.statements[update].query, `"order_timestamp"=CASE WHEN order_timestamp <=`) + assert.Contains(t, pool.statements[update].query, `"unread_count"=unread_count +`) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "FOR UPDATE")) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_unread_items`)) +} + +func TestMessageThreadStaleActivityStillUnarchives(t *testing.T) { + threadID := uuid.New() + currentMessageID := uuid.New() + content := "current" + currentTimestamp := time.Date(2026, 8, 21, 10, 0, 2, 0, time.UTC) + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + IsArchived: true, + LastMessageID: ¤tMessageID, + LastMessageContent: &content, + Status: entities.MessageStatusDelivered, + OrderTimestamp: currentTimestamp, + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ + MessageThreadID: threadID, + UserID: entities.UserID("user-id"), + Timestamp: currentTimestamp.Add(-time.Second), + MessageID: uuid.New(), + Content: "stale", + Status: entities.MessageStatusReceived, + Unarchive: true, + }) + + require.NoError(t, err) + unarchive := messageThreadStatementIndex(pool, `"is_archived"`) + require.NotEqual(t, -1, unarchive) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "FOR UPDATE")) + assert.Contains(t, pool.statements[unarchive].query, `"order_timestamp"=CASE WHEN order_timestamp <=`) + assert.Contains(t, pool.statements[unarchive].query, `"last_message_id"=CASE WHEN order_timestamp <=`) + assert.Contains(t, pool.statements[unarchive].query, `"last_message_content"=CASE WHEN order_timestamp <=`) + assert.Contains(t, pool.statements[unarchive].query, `"status"=CASE WHEN order_timestamp <=`) +} + +func TestMessageThreadDeliveredActivityDoesNotRegressSameMessage(t *testing.T) { + threadID := uuid.New() + messageID := uuid.New() + content := "delivered" + timestamp := time.Date(2026, 8, 21, 10, 0, 2, 0, time.UTC) + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + LastMessageID: &messageID, + LastMessageContent: &content, + Status: entities.MessageStatusDelivered, + OrderTimestamp: timestamp, + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ + MessageThreadID: threadID, + UserID: entities.UserID("user-id"), + Timestamp: timestamp.Add(time.Second), + MessageID: messageID, + Content: "regressed", + Status: entities.MessageStatusSent, + }) + + require.NoError(t, err) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "FOR UPDATE")) + update := messageThreadStatementIndex(pool, `UPDATE "message_threads"`) + require.NotEqual(t, -1, update) + assert.Contains(t, pool.statements[update].query, `status <>`) + assert.Contains(t, pool.statements[update].query, `last_message_id IS DISTINCT FROM`) +} + +func TestMessageThreadActivityAlwaysIncrementsIncomingMessages(t *testing.T) { + threadID := uuid.New() + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ + MessageThreadID: threadID, UserID: entities.UserID("user-id"), Timestamp: time.Date(2026, 7, 19, 10, 0, 0, 0, time.UTC), MessageID: uuid.New(), Content: "hello", Status: entities.MessageStatusReceived, - MarkAsUnread: true, - EventTimestamp: time.Date(2026, 7, 19, 10, 0, 1, 0, time.UTC), + CountAsUnread: true, }) require.NoError(t, err) - var updates []messageThreadTestStatement - for _, statement := range pool.statements { - if strings.HasPrefix(statement.query, `UPDATE "message_threads"`) { - updates = append(updates, statement) - } + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_unread_items`)) + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, "unread_count +")) +} + +func TestMessageThreadActivityDoesNotUseReadWatermark(t *testing.T) { + timestamp := time.Date(2026, 7, 19, 10, 0, 1, 0, time.UTC) + threadID := uuid.New() + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + }, } - require.Len(t, updates, 1) - assert.Contains(t, updates[0].query, `"is_read"=CASE WHEN last_read_at <`) + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ + MessageThreadID: threadID, + UserID: entities.UserID("user-id"), + Timestamp: timestamp, + MessageID: uuid.New(), + Content: "hello", + Status: entities.MessageStatusReceived, + CountAsUnread: true, + }) + + require.NoError(t, err) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_unread_items`)) + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, "unread_count +")) +} + +func TestMessageThreadActivityMissingThreadReturnsScopedNotFound(t *testing.T) { + threadID := uuid.New() + userID := entities.UserID("user-id") + pool := &messageThreadTestConnPool{ + rowsAffected: func(query string) int64 { + if strings.Contains(query, `UPDATE "message_threads"`) { + return 0 + } + return 1 + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ + MessageThreadID: threadID, + UserID: userID, + }) + + require.Error(t, err) + assert.Equal(t, ErrCodeNotFound, stacktrace.GetCode(err)) + assert.Contains(t, err.Error(), threadID.String()) + assert.Contains(t, err.Error(), string(userID)) + require.Zero(t, pool.begins) + require.Zero(t, pool.commits) + require.Zero(t, pool.rollbacks) + query := messageThreadStatementIndex(pool, `UPDATE "message_threads"`) + require.NotEqual(t, -1, query) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `SELECT * FROM "message_threads"`)) + assert.Contains(t, pool.statements[query].query, "user_id =") + assert.Contains(t, pool.statements[query].query, "id =") } func TestMessageThreadDeletedUpdatesPreserveStatusType(t *testing.T) { messageID := uuid.New() content := "previous message" - updates := messageThreadDeletedUpdates(MessageThreadDeletedUpdate{ + status := entities.MessageStatus(entities.MessageStatusDelivered) + updates, err := messageThreadDeletedUpdates(MessageThreadDeletedUpdate{ LastMessageID: &messageID, LastMessageContent: &content, - LastMessageStatus: entities.MessageStatusDelivered, + LastMessageStatus: &status, }) + require.NoError(t, err) assert.Equal(t, map[string]any{ "last_message_id": &messageID, "last_message_content": &content, @@ -175,19 +645,358 @@ func TestMessageThreadDeletedUpdatesPreserveStatusType(t *testing.T) { }, updates) } -func TestMessageThreadStatusUpdatesReadOnly(t *testing.T) { - isRead := true - readAt := time.Date(2026, 7, 18, 7, 1, 0, 0, time.UTC) +func TestMessageThreadDeletedUpdatesRequirePreviousContent(t *testing.T) { + status := entities.MessageStatus(entities.MessageStatusDelivered) + updates, err := messageThreadDeletedUpdates(MessageThreadDeletedUpdate{ + DeletedMessageID: uuid.New(), + LastMessageStatus: &status, + }) + + require.Error(t, err) + assert.Nil(t, updates) + assert.Contains(t, err.Error(), "content") +} + +func TestMessageThreadDeletedMessagePreservesUnreadCount(t *testing.T) { + threadID := uuid.New() + messageID := uuid.New() + previousMessageID := uuid.New() + previousContent := "previous" + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + UnreadCount: 1, + LastMessageID: &messageID, + }, + } + repository := newMessageThreadTestRepository(t, pool) + previousStatus := entities.MessageStatus(entities.MessageStatusDelivered) + + err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + DeletedMessageID: messageID, + LastMessageID: &previousMessageID, + LastMessageContent: &previousContent, + LastMessageStatus: &previousStatus, + }) + + require.NoError(t, err) + require.Zero(t, pool.begins) + require.Zero(t, pool.commits) + require.Zero(t, pool.rollbacks) + metadata := messageThreadStatementIndex(pool, `"last_message_id"`) + require.NotEqual(t, -1, metadata) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "FOR UPDATE")) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_deleted_items`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `UPDATE "message_thread_unread_items"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "unread_count")) +} + +func TestMessageThreadDeletedStaleReplacementPreservesNewerLastActivity(t *testing.T) { + threadID := uuid.New() + deletedMessageID := uuid.New() + newerMessageID := uuid.New() + previousMessageID := uuid.New() + previousContent := "previous" + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + UnreadCount: 1, + LastMessageID: &newerMessageID, + }, + } + repository := newMessageThreadTestRepository(t, pool) + previousStatus := entities.MessageStatus(entities.MessageStatusDelivered) + + err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + DeletedMessageID: deletedMessageID, + LastMessageID: &previousMessageID, + LastMessageContent: &previousContent, + LastMessageStatus: &previousStatus, + }) + + require.NoError(t, err) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_deleted_items`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `UPDATE "message_thread_unread_items"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "unread_count")) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `"last_message_id"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `DELETE FROM "message_threads"`)) +} + +func TestMessageThreadDeletedStaleFinalMessagePreservesNewerLastActivity(t *testing.T) { + threadID := uuid.New() + deletedMessageID := uuid.New() + newerMessageID := uuid.New() + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + UnreadCount: 1, + LastMessageID: &newerMessageID, + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + DeletedMessageID: deletedMessageID, + }) + + require.NoError(t, err) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_deleted_items`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `UPDATE "message_thread_unread_items"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "unread_count")) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `"last_message_id"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `DELETE FROM "message_threads"`)) +} + +func TestMessageThreadDeletedCurrentFinalMessageDeletesThread(t *testing.T) { + threadID := uuid.New() + deletedMessageID := uuid.New() + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + UnreadCount: 1, + LastMessageID: &deletedMessageID, + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + DeletedMessageID: deletedMessageID, + }) + + require.NoError(t, err) + require.Zero(t, pool.begins) + require.Zero(t, pool.commits) + require.Zero(t, pool.rollbacks) + threadDelete := messageThreadStatementIndex(pool, `DELETE FROM "message_threads"`) + require.NotEqual(t, -1, threadDelete) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "FOR UPDATE")) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_deleted_items`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `UPDATE "message_thread_unread_items"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "unread_count")) +} + +func TestMessageThreadDeletedCurrentMessageRequiresPreviousStatus(t *testing.T) { + threadID := uuid.New() + deletedMessageID := uuid.New() + previousMessageID := uuid.New() + previousContent := "previous" + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + LastMessageID: &deletedMessageID, + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + DeletedMessageID: deletedMessageID, + LastMessageID: &previousMessageID, + LastMessageContent: &previousContent, + }) + + require.Error(t, err) + assert.Contains(t, err.Error(), "status") + require.Zero(t, pool.rollbacks) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `"last_message_id"`)) +} + +func TestMessageThreadDeletionDoesNotDeduplicateReplay(t *testing.T) { + threadID := uuid.New() + messageID := uuid.New() + userID := entities.UserID("user-id") + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: userID, + UnreadCount: 1, + }, + } + repository := newMessageThreadTestRepository(t, pool) + + require.NoError(t, repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + UserID: userID, + Owner: "+18005550199", + Contact: "+18005550100", + DeletedMessageID: messageID, + })) + require.NoError(t, repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ + MessageThreadID: threadID, + UserID: userID, + Timestamp: time.Date(2026, 8, 21, 10, 0, 1, 0, time.UTC), + MessageID: messageID, + Content: "replayed", + Status: entities.MessageStatusReceived, + CountAsUnread: true, + })) + + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_deleted_items`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_unread_items`)) + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, "unread_count +")) +} + +func TestMessageThreadDeletionBeforeThreadDoesNotBlockLaterStore(t *testing.T) { + messageID := uuid.New() + threadID := uuid.New() + userID := entities.UserID("user-id") + pool := &messageThreadTestConnPool{} + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + UserID: userID, + Owner: "+18005550199", + Contact: "+18005550100", + DeletedMessageID: messageID, + }) + + require.NoError(t, err) + query := messageThreadStatementIndex(pool, `SELECT * FROM "message_threads"`) + require.NotEqual(t, -1, query) + assert.NotContains(t, pool.statements[query].query, "FOR UPDATE") + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_deleted_items`)) + + pool.statements = nil + content := "deleted inbound" + err = repository.Store(context.Background(), MessageThreadStoreParams{ + Thread: &entities.MessageThread{ + ID: threadID, + UserID: userID, + Owner: "+18005550199", + Contact: "+18005550100", + UnreadCount: 1, + LastMessageID: &messageID, + LastMessageContent: &content, + Status: entities.MessageStatusReceived, + }, + CountAsUnread: true, + }) + + require.NoError(t, err) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_deleted_items`)) + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, `INSERT INTO "message_threads"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_unread_items`)) +} + +func TestMessageThreadDeletedInboundReplayUpdatesExistingThread(t *testing.T) { + messageID := uuid.New() + newerMessageID := uuid.New() + threadID := uuid.New() + userID := entities.UserID("user-id") + currentContent := "newer preview" + currentTimestamp := time.Date(2026, 8, 21, 10, 0, 0, 0, time.UTC) + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: userID, + Owner: "+18005550199", + Contact: "+18005550100", + IsArchived: true, + LastMessageID: &newerMessageID, + LastMessageContent: ¤tContent, + OrderTimestamp: currentTimestamp, + }, + } + repository := newMessageThreadTestRepository(t, pool) + + require.NoError(t, repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + UserID: userID, + Owner: "+18005550199", + Contact: "+18005550100", + DeletedMessageID: messageID, + })) + + pool.statements = nil + require.NoError(t, repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ + MessageThreadID: threadID, + UserID: userID, + Timestamp: currentTimestamp.Add(time.Second), + MessageID: messageID, + Content: "replayed deleted inbound", + Status: entities.MessageStatusReceived, + CountAsUnread: true, + Unarchive: true, + })) + + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_deleted_items`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "FOR UPDATE")) + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, `"order_timestamp"`)) + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, `"is_archived"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_unread_items`)) + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, "unread_count +")) +} + +func TestMessageThreadFinalDeletionAllowsStoreReplay(t *testing.T) { + messageID := uuid.New() + threadID := uuid.New() + userID := entities.UserID("user-id") + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: userID, + Owner: "+18005550199", + Contact: "+18005550100", + LastMessageID: &messageID, + }, + } + repository := newMessageThreadTestRepository(t, pool) + + require.NoError(t, repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + UserID: userID, + Owner: "+18005550199", + Contact: "+18005550100", + DeletedMessageID: messageID, + })) + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, `DELETE FROM "message_threads"`)) + + pool.thread = nil + pool.statements = nil + content := "replayed final message" + require.NoError(t, repository.Store(context.Background(), MessageThreadStoreParams{ + Thread: &entities.MessageThread{ + ID: uuid.New(), + UserID: userID, + Owner: "+18005550199", + Contact: "+18005550100", + UnreadCount: 1, + LastMessageID: &messageID, + LastMessageContent: &content, + Status: entities.MessageStatusReceived, + }, + CountAsUnread: true, + })) + + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_deleted_items`)) + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, `INSERT INTO "message_threads"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_unread_items`)) +} + +func TestMessageThreadStatusUpdatesResetUnreadCount(t *testing.T) { + zero := uint(0) updates := messageThreadStatusUpdates(MessageThreadStatusUpdate{ - IsRead: &isRead, - ReadAt: readAt, + UnreadCount: &zero, }) - assert.Equal(t, map[string]any{ - "is_read": true, - "last_read_at": readAt, - }, updates) + assert.Equal(t, map[string]any{"unread_count": uint(0)}, updates) assert.NotContains(t, updates, "is_archived") } @@ -199,6 +1008,61 @@ func TestMessageThreadStatusUpdatesArchiveOnly(t *testing.T) { }) assert.Equal(t, map[string]any{"is_archived": true}, updates) - assert.NotContains(t, updates, "is_read") - assert.NotContains(t, updates, "last_read_at") + assert.NotContains(t, updates, "unread_count") +} + +func TestMessageThreadStatusResetUpdatesCounterWithoutLocking(t *testing.T) { + threadID := uuid.New() + zero := uint(0) + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + UnreadCount: 2, + }, + } + repository := newMessageThreadTestRepository(t, pool) + + thread, err := repository.UpdateStatus( + context.Background(), + entities.UserID("user-id"), + threadID, + MessageThreadStatusUpdate{ + UnreadCount: &zero, + }, + ) + + require.NoError(t, err) + require.NotNil(t, thread) + assert.Zero(t, thread.UnreadCount) + require.Zero(t, pool.begins) + require.Zero(t, pool.commits) + require.Zero(t, pool.rollbacks) + statusUpdate := messageThreadStatementIndex(pool, `"unread_count"`) + require.NotEqual(t, -1, statusUpdate) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "FOR UPDATE")) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `message_thread_unread_items`)) +} + +func TestMessageThreadStatusRejectsNonzeroUnreadCount(t *testing.T) { + threadID := uuid.New() + one := uint(1) + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + }, + } + repository := newMessageThreadTestRepository(t, pool) + + thread, err := repository.UpdateStatus( + context.Background(), + entities.UserID("user-id"), + threadID, + MessageThreadStatusUpdate{UnreadCount: &one}, + ) + + require.Error(t, err) + assert.Nil(t, thread) + assert.Contains(t, err.Error(), "unread count") } diff --git a/api/pkg/repositories/message_thread_repository.go b/api/pkg/repositories/message_thread_repository.go index e6093141..afe137e9 100644 --- a/api/pkg/repositories/message_thread_repository.go +++ b/api/pkg/repositories/message_thread_repository.go @@ -12,39 +12,43 @@ import ( type MessageThreadActivityUpdate struct { MessageThreadID uuid.UUID UserID entities.UserID - // Timestamp controls thread activity ordering; EventTimestamp is the server-side unread watermark. - Timestamp time.Time - MessageID uuid.UUID - Content string - Status entities.MessageStatus - MarkAsUnread bool - EventTimestamp time.Time - Unarchive bool + Timestamp time.Time + MessageID uuid.UUID + Content string + Status entities.MessageStatus + CountAsUnread bool + Unarchive bool +} + +type MessageThreadStoreParams struct { + Thread *entities.MessageThread + CountAsUnread bool } type MessageThreadStatusUpdate struct { - IsArchived *bool - IsRead *bool - ReadAt time.Time + IsArchived *bool + UnreadCount *uint } type MessageThreadDeletedUpdate struct { - MessageThreadID uuid.UUID UserID entities.UserID + Owner string + Contact string + DeletedMessageID uuid.UUID LastMessageID *uuid.UUID LastMessageContent *string - LastMessageStatus entities.MessageStatus + LastMessageStatus *entities.MessageStatus } // MessageThreadRepository loads and persists an entities.MessageThread type MessageThreadRepository interface { // Store a new entities.MessageThread - Store(ctx context.Context, thread *entities.MessageThread) error + Store(ctx context.Context, params MessageThreadStoreParams) error // UpdateActivity persists the last-message activity fields for a thread UpdateActivity(ctx context.Context, params MessageThreadActivityUpdate) error - // UpdateStatus persists archive/read status fields for a thread + // UpdateStatus persists archive/unread status fields for a thread UpdateStatus(ctx context.Context, userID entities.UserID, messageThreadID uuid.UUID, params MessageThreadStatusUpdate) (*entities.MessageThread, error) // LoadByOwnerContact fetches a thread between owner and contact diff --git a/api/pkg/requests/message_thread_update_request.go b/api/pkg/requests/message_thread_update_request.go index 3309fc91..95f0095f 100644 --- a/api/pkg/requests/message_thread_update_request.go +++ b/api/pkg/requests/message_thread_update_request.go @@ -10,8 +10,8 @@ import ( // MessageThreadUpdate is the payload for updating a message thread type MessageThreadUpdate struct { request - IsArchived *bool `json:"is_archived,omitempty" example:"true"` - IsRead *bool `json:"is_read,omitempty" example:"true"` + IsArchived *bool `json:"is_archived,omitempty" example:"true"` + UnreadCount *uint `json:"unread_count,omitempty" example:"0" minimum:"0" maximum:"0"` MessageThreadID string `json:"messageThreadID" swaggerignore:"true"` // used internally for validation } @@ -22,6 +22,6 @@ func (input *MessageThreadUpdate) ToUpdateParams(userID entities.UserID) service UserID: userID, MessageThreadID: uuid.MustParse(input.MessageThreadID), IsArchived: input.IsArchived, - IsRead: input.IsRead, + UnreadCount: input.UnreadCount, } } diff --git a/api/pkg/requests/message_thread_update_request_test.go b/api/pkg/requests/message_thread_update_request_test.go index 9f9579fd..d2f3e62a 100644 --- a/api/pkg/requests/message_thread_update_request_test.go +++ b/api/pkg/requests/message_thread_update_request_test.go @@ -1,25 +1,58 @@ package requests import ( + "encoding/json" + "reflect" "testing" "github.com/NdoleStudio/httpsms/pkg/entities" "github.com/google/uuid" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) +func TestMessageThreadUpdateJSONDistinguishesUnreadCountZeroFromOmitted(t *testing.T) { + t.Run("zero value is preserved as a pointer", func(t *testing.T) { + var input MessageThreadUpdate + + err := json.Unmarshal([]byte(`{"unread_count":0}`), &input) + + require.NoError(t, err) + require.NotNil(t, input.UnreadCount) + assert.Equal(t, uint(0), *input.UnreadCount) + }) + + t.Run("omitted unread count stays nil", func(t *testing.T) { + var input MessageThreadUpdate + + err := json.Unmarshal([]byte(`{"is_archived":true}`), &input) + + require.NoError(t, err) + assert.Nil(t, input.UnreadCount) + }) +} + func TestMessageThreadUpdateToUpdateParamsPreservesOptionalFields(t *testing.T) { threadID := uuid.New() - isRead := true + isArchived := true + unreadCount := uint(0) input := MessageThreadUpdate{ MessageThreadID: threadID.String(), - IsRead: &isRead, + IsArchived: &isArchived, + UnreadCount: &unreadCount, } params := input.ToUpdateParams(entities.UserID("user-id")) assert.Equal(t, threadID, params.MessageThreadID) assert.Equal(t, entities.UserID("user-id"), params.UserID) - assert.Nil(t, params.IsArchived) - assert.Same(t, &isRead, params.IsRead) + assert.Same(t, &isArchived, params.IsArchived) + assert.Same(t, &unreadCount, params.UnreadCount) +} + +func TestMessageThreadUpdateUnreadCountSwaggerAllowsExactlyZero(t *testing.T) { + field, ok := reflect.TypeOf(MessageThreadUpdate{}).FieldByName("UnreadCount") + require.True(t, ok) + assert.Equal(t, "0", field.Tag.Get("minimum")) + assert.Equal(t, "0", field.Tag.Get("maximum")) } diff --git a/api/pkg/services/message_thread_service.go b/api/pkg/services/message_thread_service.go index 8431672b..f55da3bf 100644 --- a/api/pkg/services/message_thread_service.go +++ b/api/pkg/services/message_thread_service.go @@ -51,16 +51,14 @@ func NewMessageThreadService( // MessageThreadUpdateParams are parameters for updating a thread type MessageThreadUpdateParams struct { - Owner string - Status entities.MessageStatus - Contact string - Content string - UserID entities.UserID - MessageID uuid.UUID - // Timestamp controls thread activity ordering; EventTimestamp is the server-side unread watermark. - Timestamp time.Time - MarkAsUnread bool - EventTimestamp time.Time + Owner string + Status entities.MessageStatus + Contact string + Content string + UserID entities.UserID + MessageID uuid.UUID + Timestamp time.Time + CountAsUnread bool } // shouldCheckUnarchive reports whether a thread update is a new inbound message @@ -101,16 +99,6 @@ func (service *MessageThreadService) UpdateThread(ctx context.Context, params Me return service.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot find thread with owner [%s], and contact [%s]. creating new thread", params.Owner, params.Contact)) } - if thread.OrderTimestamp.Unix() > params.Timestamp.Unix() && thread.Status != entities.MessageStatusSending && thread.HasLastMessage(params.MessageID) { - ctxLogger.Warn(stacktrace.NewErrorf("thread [%s] has timestamp [%s] and status [%s] which is greater than timestamp [%s] for message [%s] and status [%s]", thread.ID, thread.OrderTimestamp, thread.Status, params.Timestamp, params.MessageID, params.Status)) - return nil - } - - if thread.Status == entities.MessageStatusDelivered && thread.LastMessageID != nil && thread.HasLastMessage(params.MessageID) { - ctxLogger.Warn(stacktrace.NewErrorf("thread [%s] already has status [%s] not updating with status [%s] for message [%s]", thread.ID, thread.Status, params.Status, params.MessageID)) - return nil - } - activity := repositories.MessageThreadActivityUpdate{ MessageThreadID: thread.ID, UserID: params.UserID, @@ -118,8 +106,7 @@ func (service *MessageThreadService) UpdateThread(ctx context.Context, params Me MessageID: params.MessageID, Content: params.Content, Status: params.Status, - MarkAsUnread: params.MarkAsUnread, - EventTimestamp: params.EventTimestamp, + CountAsUnread: params.CountAsUnread, } if service.shouldCheckUnarchive(thread, params) { @@ -143,7 +130,7 @@ func (service *MessageThreadService) UpdateThread(ctx context.Context, params Me // MessageThreadStatusParams are parameters for updating a thread status type MessageThreadStatusParams struct { IsArchived *bool - IsRead *bool + UnreadCount *uint UserID entities.UserID MessageThreadID uuid.UUID } @@ -154,9 +141,8 @@ func (service *MessageThreadService) UpdateStatus(ctx context.Context, params Me defer span.End() update := repositories.MessageThreadStatusUpdate{ - IsArchived: params.IsArchived, - IsRead: params.IsRead, - ReadAt: time.Now().UTC(), + IsArchived: params.IsArchived, + UnreadCount: params.UnreadCount, } thread, err := service.repository.UpdateStatus(ctx, params.UserID, params.MessageThreadID, update) if err != nil { @@ -171,38 +157,32 @@ func (service *MessageThreadService) UpdateAfterDeletedMessage(ctx context.Conte ctx, span, ctxLogger := service.tracer.StartWithLogger(ctx, service.logger) defer span.End() - thread, err := service.repository.LoadByOwnerContact(ctx, payload.UserID, payload.Owner, payload.Contact) - if err != nil { - return service.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot find thread for user [%s] with owner [%s], and contact [%s]", payload.UserID, payload.Owner, payload.Contact)) - } - - if payload.PreviousMessageID == nil { - if err = service.repository.Delete(ctx, thread.UserID, thread.ID); err != nil { - ctxLogger.Error(stacktrace.Propagatef(err, "cannot delete thread with ID [%s] for user [%s] and owner [%s]", thread.ID, thread.UserID, thread.Owner)) - return nil - } - msg := fmt.Sprintf("previous message ID is nil for thread with ID [%s] and user [%s]", thread.ID, thread.UserID) - ctxLogger.Info(msg) - return nil - } - - if thread.LastMessageID != nil && *thread.LastMessageID != payload.MessageID { - msg := fmt.Sprintf("last message ID [%s] does not match message ID [%s] for thread with ID [%s]", *thread.LastMessageID, payload.MessageID, thread.ID) - ctxLogger.Info(msg) - return nil - } - - if err = service.repository.UpdateAfterDeletedMessage(ctx, repositories.MessageThreadDeletedUpdate{ - MessageThreadID: thread.ID, - UserID: thread.UserID, + if err := service.repository.UpdateAfterDeletedMessage(ctx, repositories.MessageThreadDeletedUpdate{ + UserID: payload.UserID, + Owner: payload.Owner, + Contact: payload.Contact, + DeletedMessageID: payload.MessageID, LastMessageID: payload.PreviousMessageID, LastMessageContent: payload.PreviousMessageContent, - LastMessageStatus: *payload.PreviousMessageStatus, + LastMessageStatus: payload.PreviousMessageStatus, }); err != nil { - return service.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot update thread with ID [%s] for user with ID [%s]", thread.ID, thread.UserID)) + return service.tracer.WrapErrorSpan(span, stacktrace.Propagatef( + err, + "cannot apply deleted message [%s] to conversation [%s/%s] for user [%s]", + payload.MessageID, + payload.Owner, + payload.Contact, + payload.UserID, + )) } - ctxLogger.Info(fmt.Sprintf("last message has been removed from thread with ID [%s] and userID [%s]", thread.ID, thread.UserID)) + ctxLogger.Info(fmt.Sprintf( + "message [%s] has been removed from conversation [%s/%s] for user [%s]", + payload.MessageID, + payload.Owner, + payload.Contact, + payload.UserID, + )) return nil } @@ -219,8 +199,7 @@ func (service *MessageThreadService) createThread(ctx context.Context, params Me Contact: params.Contact, UserID: params.UserID, IsArchived: false, - IsRead: !params.MarkAsUnread, - LastReadAt: now, + UnreadCount: 0, Color: service.getColor(), LastMessageContent: ¶ms.Content, Status: params.Status, @@ -230,7 +209,14 @@ func (service *MessageThreadService) createThread(ctx context.Context, params Me OrderTimestamp: params.Timestamp, } - if err := service.repository.Store(ctx, thread); err != nil { + if params.CountAsUnread { + thread.UnreadCount = 1 + } + + if err := service.repository.Store(ctx, repositories.MessageThreadStoreParams{ + Thread: thread, + CountAsUnread: params.CountAsUnread, + }); err != nil { return service.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot store thread with id [%s] for message with ID [%s]", thread.ID, params.MessageID)) } diff --git a/api/pkg/services/message_thread_service_test.go b/api/pkg/services/message_thread_service_test.go index 44a3bcfd..95af1e25 100644 --- a/api/pkg/services/message_thread_service_test.go +++ b/api/pkg/services/message_thread_service_test.go @@ -2,10 +2,13 @@ package services import ( "context" + "errors" + "reflect" "testing" "time" "github.com/NdoleStudio/httpsms/pkg/entities" + "github.com/NdoleStudio/httpsms/pkg/events" "github.com/NdoleStudio/httpsms/pkg/repositories" "github.com/NdoleStudio/httpsms/pkg/telemetry" "github.com/NdoleStudio/stacktrace" @@ -18,14 +21,16 @@ import ( type messageThreadRepositoryStub struct { loadByOwnerContact func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) load func(context.Context, entities.UserID, uuid.UUID) (*entities.MessageThread, error) - store func(context.Context, *entities.MessageThread) error + store func(context.Context, repositories.MessageThreadStoreParams) error updateActivity func(context.Context, repositories.MessageThreadActivityUpdate) error updateStatus func(context.Context, entities.UserID, uuid.UUID, repositories.MessageThreadStatusUpdate) (*entities.MessageThread, error) + updateAfterDelete func(context.Context, repositories.MessageThreadDeletedUpdate) error + delete func(context.Context, entities.UserID, uuid.UUID) error } -func (stub *messageThreadRepositoryStub) Store(ctx context.Context, thread *entities.MessageThread) error { +func (stub *messageThreadRepositoryStub) Store(ctx context.Context, params repositories.MessageThreadStoreParams) error { if stub.store != nil { - return stub.store(ctx, thread) + return stub.store(ctx, params) } return nil } @@ -44,7 +49,10 @@ func (stub *messageThreadRepositoryStub) UpdateStatus(ctx context.Context, userI return &entities.MessageThread{ID: threadID}, nil } -func (stub *messageThreadRepositoryStub) UpdateAfterDeletedMessage(context.Context, repositories.MessageThreadDeletedUpdate) error { +func (stub *messageThreadRepositoryStub) UpdateAfterDeletedMessage(ctx context.Context, params repositories.MessageThreadDeletedUpdate) error { + if stub.updateAfterDelete != nil { + return stub.updateAfterDelete(ctx, params) + } return nil } @@ -61,7 +69,10 @@ func (stub *messageThreadRepositoryStub) Index(context.Context, entities.UserID, return &threads, nil } -func (stub *messageThreadRepositoryStub) Delete(context.Context, entities.UserID, uuid.UUID) error { +func (stub *messageThreadRepositoryStub) Delete(ctx context.Context, userID entities.UserID, threadID uuid.UUID) error { + if stub.delete != nil { + return stub.delete(ctx, userID, threadID) + } return nil } @@ -69,12 +80,54 @@ func (stub *messageThreadRepositoryStub) DeleteAllForUser(context.Context, entit return nil } +type messageThreadPhoneRepositoryStub struct { + load func(context.Context, entities.UserID, string) (*entities.Phone, error) +} + +func (stub *messageThreadPhoneRepositoryStub) Save(context.Context, *entities.Phone) error { + return nil +} + +func (stub *messageThreadPhoneRepositoryStub) Index(context.Context, entities.UserID, repositories.IndexParams) (*[]entities.Phone, error) { + phones := []entities.Phone{} + return &phones, nil +} + +func (stub *messageThreadPhoneRepositoryStub) Load(ctx context.Context, userID entities.UserID, phoneNumber string) (*entities.Phone, error) { + if stub.load != nil { + return stub.load(ctx, userID, phoneNumber) + } + return &entities.Phone{}, nil +} + +func (stub *messageThreadPhoneRepositoryStub) LoadByID(context.Context, entities.UserID, uuid.UUID) (*entities.Phone, error) { + return &entities.Phone{}, nil +} + +func (stub *messageThreadPhoneRepositoryStub) Delete(context.Context, entities.UserID, uuid.UUID) error { + return nil +} + +func (stub *messageThreadPhoneRepositoryStub) NullifyScheduleID(context.Context, entities.UserID, uuid.UUID) error { + return nil +} + +func (stub *messageThreadPhoneRepositoryStub) DeleteAllForUser(context.Context, entities.UserID) error { + return nil +} + func newMessageThreadServiceForTest(repository repositories.MessageThreadRepository) *MessageThreadService { logger := &noopLogger{} tracer := telemetry.NewOtelLogger("test", logger) return NewMessageThreadService(logger, tracer, repository, nil, nil, nil) } +func newMessageThreadServiceWithPhoneForTest(repository repositories.MessageThreadRepository, phoneRepository repositories.PhoneRepository) *MessageThreadService { + logger := &noopLogger{} + tracer := telemetry.NewOtelLogger("test", logger) + return NewMessageThreadService(logger, tracer, repository, phoneRepository, nil, nil) +} + func TestUpdateThreadPassesUnreadWatermarkForInboundActivity(t *testing.T) { threadID := uuid.New() eventTimestamp := time.Date(2026, 7, 18, 7, 0, 0, 0, time.UTC) @@ -91,27 +144,25 @@ func TestUpdateThreadPassesUnreadWatermarkForInboundActivity(t *testing.T) { service := newMessageThreadServiceForTest(repository) err := service.UpdateThread(context.Background(), MessageThreadUpdateParams{ - UserID: entities.UserID("user-id"), - Owner: "+18005550199", - Contact: "+18005550100", - MessageID: uuid.New(), - Content: "hello", - Status: entities.MessageStatusReceived, - Timestamp: eventTimestamp, - MarkAsUnread: true, - EventTimestamp: eventTimestamp, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + MessageID: uuid.New(), + Content: "hello", + Status: entities.MessageStatusReceived, + Timestamp: eventTimestamp, + CountAsUnread: true, }) require.NoError(t, err) - assert.True(t, captured.MarkAsUnread) - assert.Equal(t, eventTimestamp, captured.EventTimestamp) + assert.True(t, captured.CountAsUnread) } func TestUpdateThreadPreservesReadStateForOutboundActivity(t *testing.T) { var captured repositories.MessageThreadActivityUpdate repository := &messageThreadRepositoryStub{ loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { - return &entities.MessageThread{ID: uuid.New(), IsRead: false}, nil + return &entities.MessageThread{ID: uuid.New()}, nil }, updateActivity: func(_ context.Context, params repositories.MessageThreadActivityUpdate) error { captured = params @@ -131,60 +182,182 @@ func TestUpdateThreadPreservesReadStateForOutboundActivity(t *testing.T) { }) require.NoError(t, err) - assert.False(t, captured.MarkAsUnread) + assert.False(t, captured.CountAsUnread) } -func TestCreateThreadSetsReadStateFromActivityDirection(t *testing.T) { +func TestUpdateThreadUnarchivesArchivedInboundMessageWhenPhoneSettingEnabled(t *testing.T) { + var captured repositories.MessageThreadActivityUpdate + repository := &messageThreadRepositoryStub{ + loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { + return &entities.MessageThread{ID: uuid.New(), IsArchived: true}, nil + }, + updateActivity: func(_ context.Context, params repositories.MessageThreadActivityUpdate) error { + captured = params + return nil + }, + } + phoneRepository := &messageThreadPhoneRepositoryStub{ + load: func(context.Context, entities.UserID, string) (*entities.Phone, error) { + return &entities.Phone{UnarchiveThread: true}, nil + }, + } + + service := newMessageThreadServiceWithPhoneForTest(repository, phoneRepository) + err := service.UpdateThread(context.Background(), MessageThreadUpdateParams{ + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + MessageID: uuid.New(), + Content: "hello", + Status: entities.MessageStatusReceived, + Timestamp: time.Now().UTC(), + CountAsUnread: true, + }) + + require.NoError(t, err) + assert.True(t, captured.Unarchive) +} + +func TestUpdateThreadIgnoresPhoneLookupErrorsWhenCheckingUnarchive(t *testing.T) { + var captured repositories.MessageThreadActivityUpdate + repository := &messageThreadRepositoryStub{ + loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { + return &entities.MessageThread{ID: uuid.New(), IsArchived: true}, nil + }, + updateActivity: func(_ context.Context, params repositories.MessageThreadActivityUpdate) error { + captured = params + return nil + }, + } + phoneRepository := &messageThreadPhoneRepositoryStub{ + load: func(context.Context, entities.UserID, string) (*entities.Phone, error) { + return nil, stacktrace.NewError("load failed") + }, + } + + service := newMessageThreadServiceWithPhoneForTest(repository, phoneRepository) + err := service.UpdateThread(context.Background(), MessageThreadUpdateParams{ + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + MessageID: uuid.New(), + Content: "hello", + Status: entities.MessageStatusReceived, + Timestamp: time.Now().UTC(), + CountAsUnread: true, + }) + + require.NoError(t, err) + assert.False(t, captured.Unarchive) +} + +func TestUpdateThreadDelegatesStaleDeliveredActivityAfterResolvingUnarchive(t *testing.T) { + threadID := uuid.New() + messageID := uuid.New() + phoneLoaded := false + var captured repositories.MessageThreadActivityUpdate + repository := &messageThreadRepositoryStub{ + loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { + return &entities.MessageThread{ + ID: threadID, + IsArchived: true, + LastMessageID: &messageID, + Status: entities.MessageStatusDelivered, + OrderTimestamp: time.Date(2026, 8, 21, 10, 0, 2, 0, time.UTC), + }, nil + }, + updateActivity: func(_ context.Context, params repositories.MessageThreadActivityUpdate) error { + captured = params + return nil + }, + } + phoneRepository := &messageThreadPhoneRepositoryStub{ + load: func(context.Context, entities.UserID, string) (*entities.Phone, error) { + phoneLoaded = true + return &entities.Phone{UnarchiveThread: true}, nil + }, + } + + service := newMessageThreadServiceWithPhoneForTest(repository, phoneRepository) + err := service.UpdateThread(context.Background(), MessageThreadUpdateParams{ + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + MessageID: messageID, + Content: "stale", + Status: entities.MessageStatusReceived, + Timestamp: time.Date(2026, 8, 21, 10, 0, 1, 0, time.UTC), + CountAsUnread: true, + }) + + require.NoError(t, err) + assert.True(t, phoneLoaded) + assert.Equal(t, threadID, captured.MessageThreadID) + assert.Equal(t, messageID, captured.MessageID) + assert.True(t, captured.CountAsUnread) + assert.True(t, captured.Unarchive) +} + +func TestCreateThreadSetsUnreadCountFromActivityDirection(t *testing.T) { tests := []struct { - name string - marksUnread bool - wantRead bool + name string + status entities.MessageStatus + countAsUnread bool + wantUnreadCount uint + wantUnreadMessage bool }{ - {name: "inbound", marksUnread: true, wantRead: false}, - {name: "outbound", marksUnread: false, wantRead: true}, + {name: "inbound", status: entities.MessageStatusReceived, countAsUnread: true, wantUnreadCount: 1, wantUnreadMessage: true}, + {name: "outbound", status: entities.MessageStatusSent, countAsUnread: false, wantUnreadCount: 0, wantUnreadMessage: false}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { - var stored *entities.MessageThread + var stored repositories.MessageThreadStoreParams repository := &messageThreadRepositoryStub{ loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { return nil, stacktrace.PropagateWithCodef(gorm.ErrRecordNotFound, repositories.ErrCodeNotFound, "not found") }, - store: func(_ context.Context, thread *entities.MessageThread) error { - stored = thread + store: func(_ context.Context, params repositories.MessageThreadStoreParams) error { + stored = params return nil }, } + messageID := uuid.New() service := newMessageThreadServiceForTest(repository) err := service.UpdateThread(context.Background(), MessageThreadUpdateParams{ - UserID: entities.UserID("user-id"), - Owner: "+18005550199", - Contact: "+18005550100", - MessageID: uuid.New(), - Content: "hello", - Status: entities.MessageStatusReceived, - Timestamp: time.Now().UTC(), - MarkAsUnread: test.marksUnread, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + MessageID: messageID, + Content: "hello", + Status: test.status, + Timestamp: time.Now().UTC(), + CountAsUnread: test.countAsUnread, }) require.NoError(t, err) - require.NotNil(t, stored) - assert.Equal(t, test.wantRead, stored.IsRead) - assert.False(t, stored.LastReadAt.IsZero()) + require.NotNil(t, stored.Thread) + assert.Equal(t, test.wantUnreadCount, stored.Thread.UnreadCount) + assert.Equal(t, test.countAsUnread, stored.CountAsUnread) + if test.wantUnreadMessage { + require.NotNil(t, stored.Thread.LastMessageID) + assert.Equal(t, messageID, *stored.Thread.LastMessageID) + } else { + assert.False(t, stored.CountAsUnread) + } }) } } func TestUpdateStatusChangesOnlyRequestedState(t *testing.T) { threadID := uuid.New() - isRead := false + unreadCount := uint(0) var captured repositories.MessageThreadStatusUpdate repository := &messageThreadRepositoryStub{ updateStatus: func(_ context.Context, _ entities.UserID, _ uuid.UUID, params repositories.MessageThreadStatusUpdate) (*entities.MessageThread, error) { captured = params - return &entities.MessageThread{ID: threadID, IsArchived: true, IsRead: false}, nil + return &entities.MessageThread{ID: threadID, IsArchived: true, UnreadCount: 0}, nil }, } @@ -192,15 +365,16 @@ func TestUpdateStatusChangesOnlyRequestedState(t *testing.T) { thread, err := service.UpdateStatus(context.Background(), MessageThreadStatusParams{ UserID: entities.UserID("user-id"), MessageThreadID: threadID, - IsRead: &isRead, + UnreadCount: &unreadCount, }) require.NoError(t, err) assert.Nil(t, captured.IsArchived) - assert.Same(t, &isRead, captured.IsRead) - assert.False(t, captured.ReadAt.IsZero()) + assert.Same(t, &unreadCount, captured.UnreadCount) + _, hasReadAt := reflect.TypeOf(captured).FieldByName("ReadAt") + assert.False(t, hasReadAt) assert.True(t, thread.IsArchived) - assert.False(t, thread.IsRead) + assert.Zero(t, thread.UnreadCount) } func TestUpdateStatusPreservesNotFoundCode(t *testing.T) { @@ -211,16 +385,205 @@ func TestUpdateStatusPreservesNotFoundCode(t *testing.T) { } service := newMessageThreadServiceForTest(repository) - isRead := true + unreadCount := uint(0) _, err := service.UpdateStatus(context.Background(), MessageThreadStatusParams{ UserID: entities.UserID("user-id"), MessageThreadID: uuid.New(), - IsRead: &isRead, + UnreadCount: &unreadCount, }) assert.Equal(t, repositories.ErrCodeNotFound, stacktrace.GetCode(err)) } +func TestUpdateAfterDeletedMessageDelegatesAllDecisionsToRepository(t *testing.T) { + deletedMessageID := uuid.New() + previousMessageID := uuid.New() + previousStatus := entities.MessageStatus(entities.MessageStatusDelivered) + previousContent := "previous" + var captured repositories.MessageThreadDeletedUpdate + + repository := &messageThreadRepositoryStub{ + updateAfterDelete: func(_ context.Context, params repositories.MessageThreadDeletedUpdate) error { + captured = params + return nil + }, + } + + service := newMessageThreadServiceForTest(repository) + err := service.UpdateAfterDeletedMessage(context.Background(), &events.MessageAPIDeletedPayload{ + MessageID: deletedMessageID, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + PreviousMessageID: &previousMessageID, + PreviousMessageStatus: &previousStatus, + PreviousMessageContent: &previousContent, + }) + + require.NoError(t, err) + assert.Equal(t, entities.UserID("user-id"), captured.UserID) + assert.Equal(t, "+18005550199", captured.Owner) + assert.Equal(t, "+18005550100", captured.Contact) + assert.Equal(t, deletedMessageID, captured.DeletedMessageID) + require.NotNil(t, captured.LastMessageID) + assert.Equal(t, previousMessageID, *captured.LastMessageID) + require.NotNil(t, captured.LastMessageContent) + assert.Equal(t, previousContent, *captured.LastMessageContent) + require.NotNil(t, captured.LastMessageStatus) + assert.Equal(t, previousStatus, *captured.LastMessageStatus) +} + +func TestUpdateAfterDeletedMessageDelegatesWithoutLoadingThread(t *testing.T) { + deletedMessageID := uuid.New() + loadCalled := false + var captured repositories.MessageThreadDeletedUpdate + repository := &messageThreadRepositoryStub{ + loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { + loadCalled = true + return &entities.MessageThread{ + ID: uuid.New(), + UserID: entities.UserID("user-id"), + }, nil + }, + updateAfterDelete: func(_ context.Context, params repositories.MessageThreadDeletedUpdate) error { + captured = params + return nil + }, + } + + service := newMessageThreadServiceForTest(repository) + err := service.UpdateAfterDeletedMessage(context.Background(), &events.MessageAPIDeletedPayload{ + MessageID: deletedMessageID, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + }) + + require.NoError(t, err) + assert.False(t, loadCalled) + assert.Equal(t, deletedMessageID, captured.DeletedMessageID) + assert.Equal(t, "+18005550199", captured.Owner) + assert.Equal(t, "+18005550100", captured.Contact) +} + +func TestUpdateAfterDeletedMessagePassesPreviousMessageMetadata(t *testing.T) { + deletedMessageID := uuid.New() + previousMessageID := uuid.New() + previousStatus := entities.MessageStatus(entities.MessageStatusDelivered) + previousContent := "previous" + var captured repositories.MessageThreadDeletedUpdate + + repository := &messageThreadRepositoryStub{ + updateAfterDelete: func(_ context.Context, params repositories.MessageThreadDeletedUpdate) error { + captured = params + return nil + }, + } + + service := newMessageThreadServiceForTest(repository) + err := service.UpdateAfterDeletedMessage(context.Background(), &events.MessageAPIDeletedPayload{ + MessageID: deletedMessageID, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + PreviousMessageID: &previousMessageID, + PreviousMessageStatus: &previousStatus, + PreviousMessageContent: &previousContent, + }) + require.NoError(t, err) + assert.Equal(t, deletedMessageID, captured.DeletedMessageID) + assert.Equal(t, deletedMessageID, captured.DeletedMessageID) + require.NotNil(t, captured.LastMessageID) + assert.Equal(t, previousMessageID, *captured.LastMessageID) + require.NotNil(t, captured.LastMessageStatus) + assert.Equal(t, previousStatus, *captured.LastMessageStatus) +} + +func TestUpdateAfterDeletedMessageDelegatesFinalMessageDeletion(t *testing.T) { + deletedMessageID := uuid.New() + var captured repositories.MessageThreadDeletedUpdate + + repository := &messageThreadRepositoryStub{ + delete: func(context.Context, entities.UserID, uuid.UUID) error { + t.Fatal("service must not delete a thread outside the repository transaction") + return nil + }, + updateAfterDelete: func(_ context.Context, params repositories.MessageThreadDeletedUpdate) error { + captured = params + return nil + }, + } + + service := newMessageThreadServiceForTest(repository) + err := service.UpdateAfterDeletedMessage(context.Background(), &events.MessageAPIDeletedPayload{ + MessageID: deletedMessageID, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + }) + + require.NoError(t, err) + assert.Equal(t, "+18005550199", captured.Owner) + assert.Equal(t, "+18005550100", captured.Contact) + assert.Equal(t, deletedMessageID, captured.DeletedMessageID) + assert.Nil(t, captured.LastMessageID) + assert.Nil(t, captured.LastMessageContent) +} + +func TestUpdateAfterDeletedMessagePropagatesRepositoryError(t *testing.T) { + deletedMessageID := uuid.New() + updateErr := errors.New("update failed") + repository := &messageThreadRepositoryStub{ + delete: func(context.Context, entities.UserID, uuid.UUID) error { + t.Fatal("service must not delete a thread outside the repository transaction") + return nil + }, + updateAfterDelete: func(context.Context, repositories.MessageThreadDeletedUpdate) error { + return updateErr + }, + } + + service := newMessageThreadServiceForTest(repository) + err := service.UpdateAfterDeletedMessage(context.Background(), &events.MessageAPIDeletedPayload{ + MessageID: deletedMessageID, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + }) + + require.Error(t, err) + assert.ErrorIs(t, err, updateErr) + assert.Contains(t, err.Error(), deletedMessageID.String()) + assert.Contains(t, err.Error(), "+18005550199") + assert.Contains(t, err.Error(), "+18005550100") +} + +func TestUpdateAfterDeletedMessagePassesNilPreviousStatusWithoutPanicking(t *testing.T) { + previousMessageID := uuid.New() + previousContent := "previous" + updateErr := errors.New("missing previous status") + repository := &messageThreadRepositoryStub{ + updateAfterDelete: func(context.Context, repositories.MessageThreadDeletedUpdate) error { + return updateErr + }, + } + + service := newMessageThreadServiceForTest(repository) + var err error + require.NotPanics(t, func() { + err = service.UpdateAfterDeletedMessage(context.Background(), &events.MessageAPIDeletedPayload{ + MessageID: uuid.New(), + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + PreviousMessageID: &previousMessageID, + PreviousMessageContent: &previousContent, + }) + }) + require.Error(t, err) + assert.ErrorIs(t, err, updateErr) +} + func TestShouldCheckUnarchive(t *testing.T) { service := &MessageThreadService{} diff --git a/api/pkg/validators/message_thread_handler_validator.go b/api/pkg/validators/message_thread_handler_validator.go index 72a64194..4c165e98 100644 --- a/api/pkg/validators/message_thread_handler_validator.go +++ b/api/pkg/validators/message_thread_handler_validator.go @@ -73,11 +73,18 @@ func (validator *MessageThreadHandlerValidator) ValidateUpdate(_ context.Context }) errors := v.ValidateStruct() - if request.IsArchived == nil && request.IsRead == nil { + if request.IsArchived == nil && request.UnreadCount == nil { if errors == nil { errors = url.Values{} } - errors.Add("payload", "at least one of is_archived or is_read is required") + errors.Add("payload", "at least one of is_archived or unread_count is required") + } + + if request.UnreadCount != nil && *request.UnreadCount != 0 { + if errors == nil { + errors = url.Values{} + } + errors.Add("unread_count", "must be 0") } return errors diff --git a/api/pkg/validators/message_thread_handler_validator_test.go b/api/pkg/validators/message_thread_handler_validator_test.go index e5ec7c1b..543b2bb8 100644 --- a/api/pkg/validators/message_thread_handler_validator_test.go +++ b/api/pkg/validators/message_thread_handler_validator_test.go @@ -7,6 +7,7 @@ import ( "github.com/NdoleStudio/httpsms/pkg/requests" "github.com/google/uuid" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestValidateUpdateRequiresAtLeastOneStatusField(t *testing.T) { @@ -20,15 +21,56 @@ func TestValidateUpdateRequiresAtLeastOneStatusField(t *testing.T) { assert.NotEmpty(t, errors.Get("payload")) } -func TestValidateUpdateAcceptsReadOnlyUpdate(t *testing.T) { +func TestValidateUpdateAcceptsUnreadCountReset(t *testing.T) { validator := &MessageThreadHandlerValidator{} - isRead := true + zero := uint(0) request := requests.MessageThreadUpdate{ MessageThreadID: uuid.NewString(), - IsRead: &isRead, + UnreadCount: &zero, } errors := validator.ValidateUpdate(context.Background(), request) assert.Empty(t, errors) } + +func TestValidateUpdateRejectsUnreadCountValuesOtherThanZero(t *testing.T) { + validator := &MessageThreadHandlerValidator{} + one := uint(1) + request := requests.MessageThreadUpdate{ + MessageThreadID: uuid.NewString(), + UnreadCount: &one, + } + + errors := validator.ValidateUpdate(context.Background(), request) + + require.NotNil(t, errors) + assert.Contains(t, errors, "unread_count") + assert.Equal(t, "must be 0", errors.Get("unread_count")) +} + +func TestValidateUpdateAcceptsArchiveOnlyAndCombinedPayloads(t *testing.T) { + validator := &MessageThreadHandlerValidator{} + isArchived := true + zero := uint(0) + + testCases := map[string]requests.MessageThreadUpdate{ + "archive only": { + MessageThreadID: uuid.NewString(), + IsArchived: &isArchived, + }, + "combined payload": { + MessageThreadID: uuid.NewString(), + IsArchived: &isArchived, + UnreadCount: &zero, + }, + } + + for name, request := range testCases { + t.Run(name, func(t *testing.T) { + errors := validator.ValidateUpdate(context.Background(), request) + + assert.Empty(t, errors) + }) + } +} diff --git a/docs/superpowers/plans/2026-08-21-unread-count-concurrency-fixes.md b/docs/superpowers/plans/2026-08-21-unread-count-concurrency-fixes.md new file mode 100644 index 00000000..77a7044c --- /dev/null +++ b/docs/superpowers/plans/2026-08-21-unread-count-concurrency-fixes.md @@ -0,0 +1,241 @@ +# Unread Count Concurrency Fixes Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Make unread-count persistence idempotent under deletion replay, read/reset races, Cockroach transaction retries, and concurrent first-message creation. + +**Architecture:** Keep `message_threads.unread_count` as the cached public count and retain unread ledger rows as tombstones after individual deletion. All unread mutations lock the thread first and run through `crdbgorm.ExecuteTx`; startup migration adds a non-destructive unique conversation index only after proving no duplicate identities exist. + +**Tech Stack:** Go 1.25.8, GORM 1.31.2, CockroachDB `crdbgorm` 2.4.3, PostgreSQL test fakes, Testify, Swag 1.16.6, Nuxt 2/pnpm. + +## Global Constraints + +- Work only in `C:\Users\achoa\Work\NdoleStudio\httpsms-unread-message-count`. +- Follow strict RED-GREEN-REFACTOR and record exact RED output. +- Do not destructively deduplicate existing message threads. +- Use GORM query builders with context propagation and wrap repository/service errors with stacktrace. +- Commit once with subject `fix(api): harden unread count concurrency` and the required Copilot trailers. + +--- + +### Task 1: Retain deleted-item tombstones + +**Files:** +- Modify: `api/pkg/entities/message_thread_unread_item.go` +- Modify: `api/pkg/entities/message_thread_test.go` +- Modify: `api/pkg/repositories/gorm_message_thread_repository.go` +- Test: `api/pkg/repositories/gorm_message_thread_repository_test.go` + +**Interfaces:** +- Consumes: `MessageThreadDeletedUpdate`, `MessageThreadActivityUpdate` +- Produces: `MessageThreadUnreadItem.Counted bool` + +- [ ] **Step 1: Write the failing tests** + +Add an entity-tag test for `Counted` and a repository behavior test that calls deletion followed by replay. Assert deletion emits a conditional `counted=true -> false` update, decrements once, replay uses conflict-ignore, and replay does not increment. + +- [ ] **Step 2: Verify RED** + +Run: `cd api && go test ./pkg/entities ./pkg/repositories -run 'TestMessageThreadUnreadItemCountedState|TestMessageThreadDeletedItemReplayDoesNotIncrement'` + +Expected: FAIL because `Counted` and tombstone transition do not exist. + +- [ ] **Step 3: Implement minimal tombstone state** + +```go +type MessageThreadUnreadItem struct { + MessageID uuid.UUID `gorm:"primaryKey;type:uuid"` + MessageThreadID uuid.UUID `gorm:"not null;type:uuid;index"` + Counted bool `gorm:"not null;default:true"` + MessageThread MessageThread `gorm:"constraint:OnDelete:CASCADE;"` +} +``` + +Replace ledger deletion with a conditional `Update("counted", false)` and decrement only when that update affects one row. Keep reset deleting all rows and insert using `ON CONFLICT DO NOTHING`. + +- [ ] **Step 4: Verify GREEN** + +Run the Task 1 command and all repository tests. + +### Task 2: Own reset watermarks after locking + +**Files:** +- Modify: `api/pkg/repositories/message_thread_repository.go` +- Modify: `api/pkg/repositories/gorm_message_thread_repository.go` +- Modify: `api/pkg/services/message_thread_service.go` +- Test: `api/pkg/repositories/gorm_message_thread_repository_test.go` +- Test: `api/pkg/services/message_thread_service_test.go` + +**Interfaces:** +- Produces: `MessageThreadStatusUpdate{IsArchived *bool, UnreadCount *uint}` +- Produces: repository-private `now func() time.Time` + +- [ ] **Step 1: Write the failing tests** + +Add a repository test whose injected clock asserts the `FOR UPDATE` statement has already executed and returns a non-UTC fixed-zone time. Assert the persisted and returned watermark is the UTC conversion. Update the service test to require no public `ReadAt` value. + +- [ ] **Step 2: Verify RED** + +Run: `cd api && go test ./pkg/repositories ./pkg/services -run 'TestMessageThreadStatusResetCreatesUTCWatermarkAfterLock|TestUpdateStatusForwardsOnlyPublicState'` + +Expected: FAIL because the service currently creates `ReadAt`. + +- [ ] **Step 3: Implement minimal ownership change** + +Remove `ReadAt` from `MessageThreadStatusUpdate`, inject `time.Now` in `gormMessageThreadRepository`, and call `repository.now().UTC()` only after `lockMessageThread` succeeds. Use the same value in SQL updates and the returned entity. + +- [ ] **Step 4: Verify GREEN** + +Run the Task 2 command and all repository/service tests. + +### Task 3: Retry all unread transactions + +**Files:** +- Modify: `api/pkg/repositories/gorm_message_thread_repository.go` +- Test: `api/pkg/repositories/gorm_message_thread_repository_test.go` + +**Interfaces:** +- Consumes: `crdbgorm.ExecuteTx(ctx, repository.db, nil, func(tx *gorm.DB) error)` + +- [ ] **Step 1: Write/adjust behavior tests** + +Teach the repository fake to accept Cockroach savepoint statements. Keep assertions on lock/mutation ordering and committed outcomes, not on implementation function names. + +- [ ] **Step 2: Verify RED** + +Run focused repository tests after replacing one transaction at a time; an unadapted fake must expose any retry/savepoint incompatibility. + +- [ ] **Step 3: Implement retry-safe closures** + +Replace the four `db.Transaction` calls in `Store`, `UpdateActivity`, `UpdateStatus`, and `UpdateAfterDeletedMessage`. Reinitialize closure-local loaded/output thread state at each attempt; assign returned state only from the successful attempt. + +- [ ] **Step 4: Verify GREEN** + +Run: `cd api && go test ./pkg/repositories` + +### Task 4: Resolve concurrent first-message creation + +**Files:** +- Modify: `api/pkg/repositories/message_thread_repository.go` +- Modify: `api/pkg/repositories/gorm_message_thread_repository.go` +- Modify: `api/pkg/services/message_thread_service.go` +- Test: `api/pkg/repositories/gorm_message_thread_repository_test.go` +- Test: `api/pkg/services/message_thread_service_test.go` + +**Interfaces:** +- Produces: `MessageThreadStoreParams{Thread *entities.MessageThread, CountAsUnread bool, EventTimestamp time.Time}` + +- [ ] **Step 1: Write the failing tests** + +Add a Store conflict test where the conversation insert affects zero rows. Assert the winning `(user_id, owner, contact)` row is locked, losing activity is applied, and its ledger/count is applied idempotently. Add a service test asserting `EventTimestamp` and `CountAsUnread` reach Store. + +- [ ] **Step 2: Verify RED** + +Run: `cd api && go test ./pkg/repositories ./pkg/services -run 'TestMessageThreadStoreConflictAppliesLosingActivity|TestCreateThreadForwardsStoreUnreadIntent'` + +Expected: FAIL because Store currently returns success without applying the losing event. + +- [ ] **Step 3: Implement minimal conflict fallback** + +Store the thread with `ON CONFLICT DO NOTHING`. On conflict, lock the winner by conversation identity, apply the losing activity, compare the event watermark, insert the ledger row with conflict-ignore, and increment only on insertion. Initialize brand-new threads with a stable pre-event watermark so concurrent first events are countable while a post-create read reset still wins by lock order. + +- [ ] **Step 4: Verify GREEN** + +Run the Task 4 command and all repository/service tests. + +### Task 5: Propagate final-message delete failures + +**Files:** +- Modify: `api/pkg/services/message_thread_service.go` +- Test: `api/pkg/services/message_thread_service_test.go` + +- [ ] **Step 1: Write the failing test** + +Make repository `Delete` return a sentinel error when `PreviousMessageID == nil`; assert `UpdateAfterDeletedMessage` returns an error containing both the sentinel and thread context. + +- [ ] **Step 2: Verify RED** + +Run: `cd api && go test ./pkg/services -run TestUpdateAfterDeletedMessagePropagatesFinalThreadDeleteError` + +Expected: FAIL because the service logs the error and returns nil. + +- [ ] **Step 3: Implement minimal propagation** + +Return the wrapped delete error instead of logging a success-shaped result. + +- [ ] **Step 4: Verify GREEN** + +Run all message-thread service tests. + +### Task 6: Make schema migration non-destructive and idempotent + +**Files:** +- Modify: `api/pkg/migrations/message_thread_unread_count.go` +- Test: `api/pkg/migrations/message_thread_unread_count_test.go` + +**Interfaces:** +- Produces: unique index `idx_message_threads_conversation` over `(user_id, owner, contact)` + +- [ ] **Step 1: Write the failing migration tests** + +Extend the fake schema state to cover legacy `is_read`, existing indexes, and duplicate identities. Assert backfill precedes drop, a second run skips both, counted-column migration occurs, duplicate identities return a precise error before index creation, and a clean schema creates the unique index. + +- [ ] **Step 2: Verify RED** + +Run: `cd api && go test ./pkg/migrations` + +Expected: FAIL because current coverage cannot model state transitions and no unique migration exists. + +- [ ] **Step 3: Implement safest migration** + +Auto-migrate the table/ledger columns, backfill then drop `is_read`, preflight duplicate conversation identities using a GORM grouped query, and create the composite unique index through `Migrator.CreateIndex` only when absent and safe. Never delete or merge duplicates. + +- [ ] **Step 4: Verify GREEN** + +Run all migration tests. Record that no real database migration was run. + +### Task 7: Constrain and regenerate the API contract + +**Files:** +- Modify: `api/pkg/requests/message_thread_update_request.go` +- Regenerate: `api/docs/docs.go` +- Regenerate: `api/docs/swagger.json` +- Regenerate: `api/docs/swagger.yaml` +- Regenerate: `web/shared/types/api.ts` + +- [ ] **Step 1: Write the contract assertion** + +Add or update a request reflection/generated-contract test requiring exactly-zero Swagger metadata. + +- [ ] **Step 2: Verify RED** + +Run the focused request test and confirm the generated Swagger lacks the constraint. + +- [ ] **Step 3: Implement and regenerate** + +Use `minimum:"0" maximum:"0"` on `UnreadCount`. Run pinned `go run github.com/swaggo/swag/cmd/swag@v1.16.6 init --requiredByDefault --parseDependency --parseInternal`, then `pnpm api:models`. + +- [ ] **Step 4: Verify GREEN** + +Assert generated Swagger has minimum and maximum zero and the web type remains `unread_count?: number`. + +### Task 8: Validate, review, commit, and report + +**Files:** +- Create: `C:\Users\achoa\Work\NdoleStudio\httpsms\.git\sdd\unread-concurrency-fix-report.md` + +- [ ] **Step 1: Format and run required validation** + +Run gofumpt on changed Go files, `cd api && go test ./...`, `cd web && pnpm lint && pnpm run generate`, and `cd tests && go test -run '^$' ./...`. + +- [ ] **Step 2: Review invariants** + +Review retry closure state, lock order, tombstone lifecycle, conflict fallback, migration safety, generated contracts, and unrelated diffs. + +- [ ] **Step 3: Commit** + +Commit all intended changes with the exact requested subject and trailers. + +- [ ] **Step 4: Write report** + +Write exact RED/GREEN output, changed files, validations, migration limitations, concerns, and commit hash to the required report path. diff --git a/docs/superpowers/plans/2026-08-21-unread-message-count.md b/docs/superpowers/plans/2026-08-21-unread-message-count.md new file mode 100644 index 00000000..7626c157 --- /dev/null +++ b/docs/superpowers/plans/2026-08-21-unread-message-count.md @@ -0,0 +1,728 @@ +# Message Thread Unread Count Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Replace binary thread read state with an idempotent unread-item count for received SMS messages and missed calls. + +**Architecture:** Store `unread_count` on each message thread for cheap list reads and maintain an internal message-ID ledger so retried inbound events cannot double-count. Repository transactions lock the thread and update activity, ledger, count, and read watermark atomically; the existing update endpoint permits clients only to reset the count to zero. + +**Tech Stack:** Go, Fiber, GORM, PostgreSQL/CockroachDB, CloudEvents, Testify, Nuxt 4, Vue 3, Pinia, Vuetify, TypeScript. + +## Global Constraints + +- `unread_count` is the sole public unread-state field; remove `is_read` from API and UI contracts. +- Received SMS messages and missed calls each increment once; outbound/status events do not change the count. +- `PUT /v1/message-threads/{id}` accepts only client value `unread_count: 0`. +- Existing `is_read=false` rows migrate to `unread_count=1`; existing read rows migrate to zero. +- Preserve `last_read_at` as an internal race-resolution watermark. +- Deleting a counted unread item decrements once without a preliminary lookup. +- Badge values are exact through 99 and display `99+` above 99. +- Use GORM with context propagation and `stacktrace.Propagate`; do not introduce raw SQL. +- Format Go with gofumpt and web code with the existing lint configuration. + +--- + +## File Structure + +### New files + +- `api/pkg/entities/message_thread_unread_item.go`: internal ledger entity keyed by message ID. +- `api/pkg/migrations/message_thread_unread_count.go`: idempotent schema/backfill transition from `is_read`. +- `api/pkg/migrations/message_thread_unread_count_test.go`: migration decision/helper coverage. + +### Modified API files + +- `api/pkg/entities/message_thread.go`: replace `IsRead` with `UnreadCount`. +- `api/pkg/entities/message_thread_test.go`: assert count schema/public contract. +- `api/pkg/di/container.go`: run the unread-count migration and ledger auto-migration. +- `api/pkg/repositories/message_thread_repository.go`: rename count intent and define ledger-aware update inputs. +- `api/pkg/repositories/gorm_message_thread_repository.go`: atomic ledger/counter/reset/deletion transactions. +- `api/pkg/repositories/gorm_message_thread_repository_test.go`: helper SQL ownership, idempotency intent, reset, and deletion tests. +- `api/pkg/services/message_thread_service.go`: pass count intent, initialize new counts, and process non-last deletions. +- `api/pkg/services/message_thread_service_test.go`: service contract and deletion routing tests. +- `api/pkg/listeners/message_thread_listener.go`: mark received SMS and missed calls as countable. +- `api/pkg/listeners/message_thread_listener_test.go`: listener count-intent tests. +- `api/pkg/listeners/read_receipts_test_helpers_test.go`: update repository test stub signatures. +- `api/pkg/requests/message_thread_update_request.go`: replace `is_read` with optional `unread_count`. +- `api/pkg/requests/message_thread_update_request_test.go`: conversion tests. +- `api/pkg/validators/message_thread_handler_validator.go`: require archive/reset and reject nonzero counts. +- `api/pkg/validators/message_thread_handler_validator_test.go`: request validation tests. +- `api/pkg/handlers/message_thread_handler_test.go`: response and removed-contract tests. +- `api/docs/docs.go`, `api/docs/swagger.json`, `api/docs/swagger.yaml`: regenerated API contract. + +### Modified web and integration files + +- `web/shared/types/api.ts`: regenerated `unread_count` types. +- `web/app/stores/threads.ts`: reset count and use zero/nonzero state. +- `web/app/pages/threads/[id]/index.vue`: rename mark-read functions to count reset. +- `web/app/components/MessageThread.vue`: numeric badge and `unread_count > 0` styling. +- `tests/read_receipts_test.go`: count-based end-to-end coverage. +- `tests/README.md`: describe unread-count integration coverage. + +--- + +### Task 1: Add unread-count schema and migration + +**Files:** +- Create: `api/pkg/entities/message_thread_unread_item.go` +- Create: `api/pkg/migrations/message_thread_unread_count.go` +- Create: `api/pkg/migrations/message_thread_unread_count_test.go` +- Modify: `api/pkg/entities/message_thread.go` +- Modify: `api/pkg/entities/message_thread_test.go` +- Modify: `api/pkg/di/container.go` + +**Interfaces:** +- Produces: `entities.MessageThread.UnreadCount uint` +- Produces: `entities.MessageThreadUnreadItem{MessageID, MessageThreadID}` +- Produces: `migrations.MigrateMessageThreadUnreadCount(db *gorm.DB) error` + +- [ ] **Step 1: Replace the entity test with count and ledger schema assertions** + +```go +func TestMessageThreadUnreadFields(t *testing.T) { + threadType := reflect.TypeOf(MessageThread{}) + _, hasIsRead := threadType.FieldByName("IsRead") + assert.False(t, hasIsRead) + + unreadCount, ok := threadType.FieldByName("UnreadCount") + require.True(t, ok) + assert.Equal(t, "unread_count", unreadCount.Tag.Get("json")) + assert.Contains(t, unreadCount.Tag.Get("gorm"), "not null") + assert.Contains(t, unreadCount.Tag.Get("gorm"), "default:0") + + lastReadAt, ok := threadType.FieldByName("LastReadAt") + require.True(t, ok) + assert.Equal(t, "-", lastReadAt.Tag.Get("json")) +} + +func TestMessageThreadUnreadItemUsesMessageIDAsPrimaryKey(t *testing.T) { + itemType := reflect.TypeOf(MessageThreadUnreadItem{}) + messageID, ok := itemType.FieldByName("MessageID") + require.True(t, ok) + assert.Contains(t, messageID.Tag.Get("gorm"), "primaryKey") +} +``` + +- [ ] **Step 2: Run the entity tests and verify failure** + +Run: + +```bash +cd api +go test ./pkg/entities -run 'TestMessageThreadUnread' -count=1 +``` + +Expected: FAIL because `UnreadCount` and `MessageThreadUnreadItem` do not exist. + +- [ ] **Step 3: Add the count and ledger entities** + +```go +// MessageThread fields +IsArchived bool `json:"is_archived" example:"false"` +UnreadCount uint `json:"unread_count" gorm:"not null;default:0" example:"2"` +LastReadAt time.Time `json:"-" gorm:"not null;default:CURRENT_TIMESTAMP"` +``` + +```go +package entities + +import "github.com/google/uuid" + +// MessageThreadUnreadItem records an inbound item currently counted as unread. +type MessageThreadUnreadItem struct { + MessageID uuid.UUID `gorm:"primaryKey;type:uuid"` + MessageThreadID uuid.UUID `gorm:"not null;type:uuid;index"` + MessageThread MessageThread `gorm:"constraint:OnDelete:CASCADE;"` +} +``` + +- [ ] **Step 4: Add an idempotent GORM migration** + +Implement `MigrateMessageThreadUnreadCount` so it: + +```go +func MigrateMessageThreadUnreadCount(db *gorm.DB) error { + if err := db.AutoMigrate(&entities.MessageThread{}, &entities.MessageThreadUnreadItem{}); err != nil { + return stacktrace.Propagate(err, "cannot migrate message thread unread count schema") + } + if !db.Migrator().HasColumn("message_threads", "is_read") { + return nil + } + if err := db.Table("message_threads"). + Where("is_read = ?", false). + Where("unread_count = ?", 0). + Update("unread_count", 1).Error; err != nil { + return stacktrace.Propagate(err, "cannot backfill message thread unread counts") + } + if err := db.Migrator().DropColumn("message_threads", "is_read"); err != nil { + return stacktrace.Propagate(err, "cannot drop legacy message thread is_read column") + } + return nil +} +``` + +Add helper-level tests proving the migration skips the backfill when the legacy +column is absent and propagates migration errors; use the repository's existing +GORM fake-connection pattern rather than a new test dependency. + +- [ ] **Step 5: Wire the migration into the DI container** + +Replace the direct `AutoMigrate(&entities.MessageThread{})` call with: + +```go +if err = migrations.MigrateMessageThreadUnreadCount(db); err != nil { + container.logger.Fatal(stacktrace.Propagate(err, "cannot migrate message thread unread counts")) +} +``` + +- [ ] **Step 6: Run focused tests and format** + +Run: + +```bash +cd api +gofumpt -w pkg/entities/message_thread.go pkg/entities/message_thread_unread_item.go pkg/entities/message_thread_test.go pkg/migrations/message_thread_unread_count.go pkg/migrations/message_thread_unread_count_test.go pkg/di/container.go +go test ./pkg/entities ./pkg/migrations -count=1 +``` + +Expected: PASS. + +- [ ] **Step 7: Commit** + +```bash +git add api/pkg/entities api/pkg/migrations api/pkg/di/container.go +git commit -m "feat(api): add unread count schema" +``` + +--- + +### Task 2: Implement ledger-backed repository updates + +**Files:** +- Modify: `api/pkg/repositories/message_thread_repository.go` +- Modify: `api/pkg/repositories/gorm_message_thread_repository.go` +- Modify: `api/pkg/repositories/gorm_message_thread_repository_test.go` + +**Interfaces:** +- Produces: `MessageThreadActivityUpdate.CountAsUnread bool` +- Produces: `MessageThreadStatusUpdate.UnreadCount *uint` +- Produces: `MessageThreadDeletedUpdate.DeletedMessageID uuid.UUID` +- Consumes: `entities.MessageThreadUnreadItem` + +- [ ] **Step 1: Write failing repository contract/helper tests** + +Add tests that assert: + +```go +func TestMessageThreadStatusUpdatesResetUnreadCount(t *testing.T) { + zero := uint(0) + readAt := time.Date(2026, 8, 21, 10, 0, 0, 0, time.UTC) + updates := messageThreadStatusUpdates(MessageThreadStatusUpdate{ + UnreadCount: &zero, + ReadAt: readAt, + }) + assert.Equal(t, map[string]any{ + "unread_count": 0, + "last_read_at": readAt, + }, updates) +} + +func TestMessageThreadActivityUpdatesDoNotOwnUnreadColumns(t *testing.T) { + updates := messageThreadActivityUpdates(MessageThreadActivityUpdate{ /* activity fields */ }) + assert.NotContains(t, updates, "unread_count") + assert.NotContains(t, updates, "last_read_at") +} +``` + +Also add transaction tests using the existing fake connection to verify that: + +- countable activity emits a ledger insert and count increment; +- duplicate ledger insert (`RowsAffected == 0`) does not increment; +- reset deletes ledger rows and updates `last_read_at`; +- deletion decrements only when ledger deletion affects one row; +- decrement uses `GREATEST(unread_count - 1, 0)`. + +- [ ] **Step 2: Run repository tests and verify failure** + +Run: + +```bash +cd api +go test ./pkg/repositories -run 'TestMessageThread(Activity|Status|Unread|Deleted)' -count=1 +``` + +Expected: FAIL on missing count fields and ledger behavior. + +- [ ] **Step 3: Update repository input types** + +```go +type MessageThreadActivityUpdate struct { + MessageThreadID uuid.UUID + UserID entities.UserID + Timestamp time.Time + MessageID uuid.UUID + Content string + Status entities.MessageStatus + CountAsUnread bool + EventTimestamp time.Time + Unarchive bool +} + +type MessageThreadStatusUpdate struct { + IsArchived *bool + UnreadCount *uint + ReadAt time.Time +} + +type MessageThreadDeletedUpdate struct { + MessageThreadID uuid.UUID + UserID entities.UserID + DeletedMessageID uuid.UUID + UpdateLastMessage bool + LastMessageID *uuid.UUID + LastMessageContent *string + LastMessageStatus entities.MessageStatus +} +``` + +- [ ] **Step 4: Add shared lock and ledger helpers** + +Use `clause.Locking{Strength: "UPDATE"}` with `WithContext(ctx)` and always +scope the thread by `user_id` and ID. Add private helpers that accept the +transaction: + +```go +func lockMessageThread(tx *gorm.DB, userID entities.UserID, threadID uuid.UUID) (*entities.MessageThread, error) +func insertUnreadItem(tx *gorm.DB, item entities.MessageThreadUnreadItem) (bool, error) +func deleteUnreadItem(tx *gorm.DB, messageID uuid.UUID, threadID uuid.UUID) (bool, error) +``` + +`insertUnreadItem` uses `clause.OnConflict{DoNothing: true}` and returns +`RowsAffected == 1`. + +- [ ] **Step 5: Implement atomic activity counting** + +Inside `UpdateActivity`: + +```go +return repository.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + thread, err := lockMessageThread(tx, params.UserID, params.MessageThreadID) + if err != nil { return err } + + if err := tx.Model(thread).Updates(messageThreadActivityUpdates(params)).Error; err != nil { + return err + } + if !params.CountAsUnread || !params.EventTimestamp.After(thread.LastReadAt) { + return nil + } + inserted, err := insertUnreadItem(tx, entities.MessageThreadUnreadItem{ + MessageID: params.MessageID, MessageThreadID: params.MessageThreadID, + }) + if err != nil || !inserted { return err } + return tx.Model(thread).UpdateColumn( + "unread_count", gorm.Expr("unread_count + ?", 1), + ).Error +}) +``` + +Wrap all returned errors with the existing tracer/stacktrace pattern and map a +missing locked thread to `ErrCodeNotFound`. + +- [ ] **Step 6: Implement reset and deletion transactions** + +`UpdateStatus` locks the thread. When `UnreadCount != nil`, update count and +watermark, delete all ledger rows for the thread, and return the updated entity. +The validator guarantees the value is zero, but the repository must reject or +return an error for a nonzero value rather than silently writing it. + +`UpdateAfterDeletedMessage` locks the thread, deletes the matching ledger row, +conditionally decrements, and applies last-message fields only when +`UpdateLastMessage` is true. + +- [ ] **Step 7: Update new-thread storage** + +Change `Store` to accept an optional unread item ID: + +```go +Store(ctx context.Context, thread *entities.MessageThread, unreadMessageID *uuid.UUID) error +``` + +Create the thread and initial ledger row in one transaction. Outbound threads +pass nil. Inbound threads pass their message ID and store `UnreadCount=1`. + +- [ ] **Step 8: Run and commit** + +Run: + +```bash +cd api +gofumpt -w pkg/repositories/message_thread_repository.go pkg/repositories/gorm_message_thread_repository.go pkg/repositories/gorm_message_thread_repository_test.go +go test ./pkg/repositories -count=1 +``` + +Expected: PASS. + +```bash +git add api/pkg/repositories +git commit -m "feat(api): count unread thread items" +``` + +--- + +### Task 3: Update service and listener flows + +**Files:** +- Modify: `api/pkg/services/message_thread_service.go` +- Modify: `api/pkg/services/message_thread_service_test.go` +- Modify: `api/pkg/listeners/message_thread_listener.go` +- Modify: `api/pkg/listeners/message_thread_listener_test.go` +- Modify: `api/pkg/listeners/read_receipts_test_helpers_test.go` + +**Interfaces:** +- Consumes: repository contracts from Task 2. +- Produces: `MessageThreadUpdateParams.CountAsUnread bool` +- Produces: `MessageThreadStatusParams.UnreadCount *uint` + +- [ ] **Step 1: Write failing service/listener tests** + +Cover: + +```go +assert.True(t, captured.CountAsUnread) // received SMS +assert.True(t, captured.CountAsUnread) // missed call +assert.Equal(t, uint(1), stored.UnreadCount) // new inbound +assert.Equal(t, uint(0), stored.UnreadCount) // new outbound +``` + +Add deletion tests proving a non-last deleted message still calls +`UpdateAfterDeletedMessage` with `UpdateLastMessage=false`, while deleting the +last message sets it true. + +- [ ] **Step 2: Run focused tests and verify failure** + +```bash +cd api +go test ./pkg/services ./pkg/listeners -run 'Test(MessageThread|UpdateThread|CreateThread|UpdateAfterDeleted)' -count=1 +``` + +- [ ] **Step 3: Rename count intent and initialize new threads** + +Replace `MarkAsUnread` with `CountAsUnread` throughout service/listener inputs. +For new threads: + +```go +thread.UnreadCount = 0 +var unreadMessageID *uuid.UUID +if params.CountAsUnread { + thread.UnreadCount = 1 + unreadMessageID = ¶ms.MessageID +} +err := service.repository.Store(ctx, thread, unreadMessageID) +``` + +- [ ] **Step 4: Make deletion cleanup unconditional** + +Keep the existing whole-thread deletion when no previous message remains. +Otherwise always call the repository: + +```go +updateLastMessage := thread.LastMessageID != nil && *thread.LastMessageID == payload.MessageID +err = service.repository.UpdateAfterDeletedMessage(ctx, repositories.MessageThreadDeletedUpdate{ + MessageThreadID: thread.ID, + UserID: thread.UserID, + DeletedMessageID: payload.MessageID, + UpdateLastMessage: updateLastMessage, + LastMessageID: payload.PreviousMessageID, + LastMessageContent: payload.PreviousMessageContent, + LastMessageStatus: *payload.PreviousMessageStatus, +}) +``` + +- [ ] **Step 5: Run, format, and commit** + +```bash +cd api +gofumpt -w pkg/services/message_thread_service.go pkg/services/message_thread_service_test.go pkg/listeners/message_thread_listener.go pkg/listeners/message_thread_listener_test.go pkg/listeners/read_receipts_test_helpers_test.go +go test ./pkg/services ./pkg/listeners -count=1 +git add pkg/services pkg/listeners +git commit -m "feat(api): route unread count activity" +``` + +--- + +### Task 4: Replace the update API contract + +**Files:** +- Modify: `api/pkg/requests/message_thread_update_request.go` +- Modify: `api/pkg/requests/message_thread_update_request_test.go` +- Modify: `api/pkg/validators/message_thread_handler_validator.go` +- Modify: `api/pkg/validators/message_thread_handler_validator_test.go` +- Modify: `api/pkg/handlers/message_thread_handler_test.go` + +**Interfaces:** +- Produces: request field `UnreadCount *uint` +- Consumes: `services.MessageThreadStatusParams.UnreadCount *uint` + +- [ ] **Step 1: Write failing request and validator tests** + +Add cases for: + +```go +zero := uint(0) +request := requests.MessageThreadUpdate{MessageThreadID: uuid.NewString(), UnreadCount: &zero} +assert.Empty(t, validator.ValidateUpdate(context.Background(), request)) +``` + +```go +one := uint(1) +errors := validator.ValidateUpdate(context.Background(), requests.MessageThreadUpdate{ + MessageThreadID: uuid.NewString(), UnreadCount: &one, +}) +assert.Contains(t, errors, "unread_count") +``` + +Also verify an `is_read`-only JSON body returns 422 because it contains no +supported update field. + +- [ ] **Step 2: Run focused tests and verify failure** + +```bash +cd api +go test ./pkg/requests ./pkg/validators ./pkg/handlers -run 'Test(MessageThreadUpdate|ValidateUpdate|MessageThreadHandler)' -count=1 +``` + +- [ ] **Step 3: Replace the request and validation fields** + +```go +type MessageThreadUpdate struct { + request + IsArchived *bool `json:"is_archived,omitempty" example:"true"` + UnreadCount *uint `json:"unread_count,omitempty" example:"0"` + MessageThreadID string `json:"messageThreadID" swaggerignore:"true"` +} +``` + +Validation requires at least one supported pointer. If `UnreadCount != nil && +*UnreadCount != 0`, add `"unread_count": "must be 0"`. + +- [ ] **Step 4: Run, format, and commit** + +```bash +cd api +gofumpt -w pkg/requests/message_thread_update_request.go pkg/requests/message_thread_update_request_test.go pkg/validators/message_thread_handler_validator.go pkg/validators/message_thread_handler_validator_test.go pkg/handlers/message_thread_handler_test.go +go test ./pkg/requests ./pkg/validators ./pkg/handlers -count=1 +git add pkg/requests pkg/validators pkg/handlers +git commit -m "feat(api): expose unread count reset" +``` + +--- + +### Task 5: Regenerate Swagger and web API types + +**Files:** +- Modify: `api/docs/docs.go` +- Modify: `api/docs/swagger.json` +- Modify: `api/docs/swagger.yaml` +- Modify: `web/shared/types/api.ts` + +**Interfaces:** +- Produces: `EntitiesMessageThread.unread_count: number` +- Produces: `RequestsMessageThreadUpdate.unread_count?: number` +- Removes: both generated `is_read` properties. + +- [ ] **Step 1: Regenerate Swagger** + +```bash +cd api +swag init --requiredByDefault --parseDependency --parseInternal +``` + +Expected: generated docs contain `unread_count` and no message-thread +`is_read`. + +- [ ] **Step 2: Regenerate web types** + +```bash +cd web +pnpm api:models +``` + +- [ ] **Step 3: Verify generated contracts** + +```bash +rg -n '"?unread_count"?|"?is_read"?' api/docs web/shared/types/api.ts +``` + +Expected: message-thread schemas contain `unread_count`; `is_read` has no +message-thread contract matches. + +- [ ] **Step 4: Commit** + +```bash +git add api/docs web/shared/types/api.ts +git commit -m "docs(api): publish unread counts" +``` + +--- + +### Task 6: Update the web store, detail page, and badge + +**Files:** +- Modify: `web/app/stores/threads.ts` +- Modify: `web/app/pages/threads/[id]/index.vue` +- Modify: `web/app/components/MessageThread.vue` + +**Interfaces:** +- Consumes: generated `EntitiesMessageThread.unread_count`. +- Produces: `resetThreadUnreadCount(threadId: string, force?: boolean): Promise`. + +- [ ] **Step 1: Replace the store action** + +```ts +async function resetThreadUnreadCount(threadId: string, force = false) { + const thread = threads.value.find((item) => item.id === threadId) + if (!thread) throw new Error(`Cannot find thread with id ${threadId}`) + if (!force && thread.unread_count === 0) return + + const response = await apiFetch<{ data: EntitiesMessageThread }>( + `/v1/message-threads/${threadId}`, + { method: 'PUT', body: { unread_count: 0 } }, + ) + replaceThread(response.data) +} +``` + +Preserve the existing `try/catch`, notification, reload, and `AggregateError` +behavior around the request. Export the renamed action. + +- [ ] **Step 2: Rename detail-page read calls** + +Rename `markCurrentThreadRead` to `resetCurrentThreadUnreadCount` and call +`threadsStore.resetThreadUnreadCount`. Preserve forced realtime resets for +received SMS and missed-call events. + +- [ ] **Step 3: Render the numeric badge** + +Add: + +```ts +function unreadBadge(count: number): false | { color: string; content: string } { + if (count === 0) return false + return { color: 'primary', content: count > 99 ? '99+' : String(count) } +} +``` + +Use `thread.unread_count > 0` for bold classes and bind the avatar badge to +`unreadBadge(thread.unread_count)`. + +- [ ] **Step 4: Run web validation** + +```bash +cd web +pnpm lint +pnpm run generate +``` + +Expected: both commands pass. + +- [ ] **Step 5: Commit** + +```bash +git add web/app/stores/threads.ts web/app/pages/threads/[id]/index.vue web/app/components/MessageThread.vue +git commit -m "feat(web): show unread message counts" +``` + +--- + +### Task 7: Update integration coverage and run full validation + +**Files:** +- Modify: `tests/read_receipts_test.go` +- Modify: `tests/README.md` + +**Interfaces:** +- Consumes: public `unread_count` response and reset request. + +- [ ] **Step 1: Convert the integration model and reset helper** + +```go +type integrationMessageThread struct { + ID string `json:"id"` + Contact string `json:"contact"` + UnreadCount uint `json:"unread_count"` + LastMessageContent *string `json:"last_message_content"` +} +``` + +Reset with: + +```go +map[string]any{"unread_count": 0} +``` + +- [ ] **Step 2: Extend count assertions** + +Exercise: + +- first received SMS reaches count 1; +- second received SMS reaches count 2; +- reset returns and persists zero; +- missed call reaches count 1; +- outbound activity preserves count 1; +- deleting the unread missed-call item returns count to zero when the existing + integration API provides the created message ID. + +Do not attempt to replay internal CloudEvents through a public endpoint. Keep +duplicate-event idempotency in repository/listener tests. + +- [ ] **Step 3: Update integration coverage documentation** + +Change the read-receipts entry in `tests/README.md` to state that the test +covers unread SMS/missed-call counts, reset, and outbound preservation. + +- [ ] **Step 4: Run API tests** + +```bash +cd api +go test ./... +``` + +Expected: PASS. + +- [ ] **Step 5: Run web checks** + +```bash +cd web +pnpm lint +pnpm run generate +``` + +Expected: PASS. + +- [ ] **Step 6: Run targeted integration tests when the Docker stack is available** + +```bash +cd tests +go test -v -timeout 120s -run TestMessageThreadReadReceipts ./... +``` + +Expected: PASS. If the stack is unavailable, record the connection failure +without treating it as product behavior. + +- [ ] **Step 7: Scan the active contract and inspect the final diff** + +```bash +rg -n 'IsRead|is_read|MarkAsUnread|markThreadRead' api/pkg web/app web/shared/types/api.ts tests +git diff --check +git status --short +``` + +Expected: no stale active-contract matches; historical specs/plans may still +mention `is_read`. + +- [ ] **Step 8: Commit** + +```bash +git add tests/read_receipts_test.go tests/README.md +git commit -m "test: cover unread message counts" +``` diff --git a/docs/superpowers/specs/2026-08-21-unread-message-count-design.md b/docs/superpowers/specs/2026-08-21-unread-message-count-design.md new file mode 100644 index 00000000..380ecb72 --- /dev/null +++ b/docs/superpowers/specs/2026-08-21-unread-message-count-design.md @@ -0,0 +1,358 @@ +# Message Thread Unread Count + +- Date: 2026-08-21 +- Status: Approved (design) +- Scope: `api/` Go backend and `web/` Nuxt frontend. Android is unchanged. +- Branch: `feat/unread-message-count`, based on `origin/main` + +## Problem + +Message threads currently expose only a binary `is_read` state. Users can tell +that a thread contains unread activity, but not how many inbound items they have +not opened. + +Replace the binary state with an exact, server-owned unread count. Received SMS +messages and missed calls each contribute one unread item. Opening a thread +resets its count to zero. + +## Decisions + +- `unread_count` is the sole public unread-state field. +- Remove `is_read` from the Go entity, API requests, API responses, generated + web types, and UI logic. +- A received SMS increments the count once. +- A missed call increments the count once. +- Duplicate or retried events do not increment the count twice. +- Outbound messages and later delivery/status updates preserve the count. +- Opening a thread resets the count to zero. +- The existing thread update endpoint accepts `unread_count: 0`; clients cannot + assign a nonzero count. +- Deleting a still-unread inbound item decrements the count without an extra + preliminary database read. +- Existing `is_read=false` threads migrate to `unread_count=1`; exact counting + starts for new activity after deployment. +- The thread list displays a numeric badge through `99`, then displays `99+`. +- The existing internal `last_read_at` watermark remains to resolve races. + +## Architecture + +Use a ledger-backed cached counter: + +1. `message_threads.unread_count` is the value returned by the API and rendered + by the web UI. +2. An internal unread-item ledger stores the message ID for every post-deploy + inbound item that currently contributes to the count. +3. Thread updates, ledger changes, and count changes occur in one transaction. + +The ledger provides idempotency. Its message ID is unique, so replaying the same +received-SMS or missed-call event cannot increment the count again. The cached +counter keeps thread-list reads as cheap as they are today and avoids correlated +message-count queries on every page load. + +## Persistence + +### Message thread + +Replace `IsRead` with: + +```go +UnreadCount uint `json:"unread_count" gorm:"not null;default:0" example:"2"` +``` + +Keep: + +```go +LastReadAt time.Time `json:"-" gorm:"not null;default:CURRENT_TIMESTAMP"` +``` + +`LastReadAt` is not exposed to clients. It remains the ordering watermark used +to stop a delayed inbound listener from restoring unread state after the user +has opened the thread. + +### Unread-item ledger + +Add an internal entity with: + +- `MessageID` as its UUID primary key; +- `MessageThreadID` as an indexed UUID foreign key; +- cascading deletion when the thread is deleted. + +The ledger does not need a public API. A globally unique message ID identifies +both received SMS messages and stored missed-call messages. + +### Schema transition + +The startup migration performs these steps before normal service traffic: + +1. Add `message_threads.unread_count` with a non-null default of zero. +2. Create the unread-item ledger table and indexes. +3. Where the legacy column exists, set `unread_count=1` for rows whose + `is_read=false`; leave previously read rows at zero. +4. Drop the legacy `is_read` column after the backfill succeeds. + +The transition is idempotent: it checks schema state before each one-time step. +Migration errors remain fatal. The application must not serve a mixed contract +or silently skip a failed backfill. + +Existing unread rows intentionally have no synthetic ledger record. Their +preserved count of one remains until the thread is opened. All inbound activity +processed after deployment is tracked exactly in the ledger. + +## API Components + +### Listener inputs + +Rename the service/repository intent from `MarkAsUnread` to `CountAsUnread`. +`MessageThreadUpdateParams` continues to carry: + +- the message ID; +- the activity timestamp used for thread ordering; +- the CloudEvent timestamp used as the unread watermark; +- whether this event represents a countable inbound item. + +Received-SMS and missed-call listeners set `CountAsUnread=true`. Outbound, +sending, delivery, failure, scheduling, and expiry listeners leave it false. + +### Service + +`MessageThreadService` continues to coordinate: + +- loading or creating the thread; +- last-message metadata; +- optional unarchiving for inbound activity; +- repository calls. + +The service does not implement ledger or counter arithmetic. Those details stay +inside the repository transaction. + +New threads start with: + +- `unread_count=1` and one ledger row when created from countable inbound + activity; +- `unread_count=0` and no ledger row for outbound activity. + +### Repository + +For a countable existing-thread update, the repository transaction: + +1. locks the thread row; +2. updates the normal last-message activity fields; +3. compares the CloudEvent timestamp with `last_read_at`; +4. inserts the message ID into the ledger with conflict-ignore when the event + is newer than the read watermark; +5. increments `unread_count` only when the insert affected one row. + +For a non-countable event, the repository updates only the existing activity +and optional unarchive fields. It does not touch the ledger, count, or read +watermark. + +For a read reset, the repository transaction: + +1. locks and updates the authenticated user's thread; +2. sets `unread_count=0` and `last_read_at` to the same UTC timestamp; +3. deletes all ledger rows for the thread; +4. returns the updated thread. + +For deletion of an individual message, the service must process unread-ledger +cleanup even when the deleted item was not the thread's last message. If the +deleted item was the last message, the repository also applies the existing +last-message replacement fields. In the same transaction it: + +1. removes the matching ledger row; +2. decrements `unread_count` with a floor of zero only if a ledger row was + removed; +3. updates last-message metadata only when the deleted item was the thread's + current last message. + +This deletion path needs no extra lookup and no extension to the deletion event: +the payload already includes the deleted message ID. + +Deleting a thread or user cascades or explicitly deletes its ledger rows as part +of the existing deletion operation. + +## Update Endpoint + +Keep: + +```text +PUT /v1/message-threads/{messageThreadID} +``` + +Replace the optional `is_read` request field with: + +```go +UnreadCount *uint `json:"unread_count,omitempty" example:"0"` +``` + +Validation rules: + +- at least one of `is_archived` or `unread_count` is present; +- if `unread_count` is present, its only valid value is zero; +- archive-only updates preserve unread count and `last_read_at`; +- unread-reset-only updates preserve archive state; +- combined archive/reset updates apply atomically; +- invalid IDs and unsupported or empty payloads return the existing bad-request + response; +- a thread outside the authenticated user scope returns the existing not-found + response. + +API responses contain `unread_count` and no `is_read`. + +## Concurrency and Idempotency + +All operations that mutate unread state lock the thread first and use the same +lock order. This prevents receive, read, and delete transactions from producing +a counter/ledger mismatch. + +The required race behavior is: + +- duplicate inbound event: the ledger conflict prevents a second increment; +- inbound event committed before read: the later read clears its ledger row and + resets the count; +- old inbound event processed after read: its CloudEvent timestamp is not newer + than `last_read_at`, so it does not insert or increment; +- genuinely new inbound event after read: it inserts and increments; +- deletion after read: the cleared ledger has no matching row, so the count + remains zero; +- repeated deletion: only the first successful ledger deletion can decrement; +- every decrement uses a floor of zero. + +Transaction failures roll back activity metadata, ledger mutations, and counter +changes together. + +## Web Components + +### Store + +Replace binary read checks with `thread.unread_count > 0`. + +The current mark-read action becomes an unread-count reset: + +```json +{ "unread_count": 0 } +``` + +The action: + +- skips the request when the local count is already zero unless a realtime + refresh explicitly forces reconciliation; +- replaces the matching local thread from the successful API response; +- does not optimistically clear the count; +- preserves the existing notification and reload behavior when the request + fails. + +Opening a thread invokes the reset as part of the existing message-loading flow. +Inbound realtime activity for the currently open thread forces the idempotent +reset so the thread does not remain unread while visible. + +### Thread list + +`MessageThread.vue` uses `unread_count > 0` for: + +- bold contact text; +- bold message preview text; +- displaying the primary-color avatar badge. + +The badge displays: + +- no badge for zero; +- the exact count from 1 through 99; +- `99+` for counts above 99. + +The same shared component continues to cover mobile, desktop, inbox, and +archived thread lists. + +### Generated contracts + +After changing API annotations: + +1. Run `swag init --requiredByDefault --parseDependency --parseInternal` in + `api/`. +2. Run `pnpm api:models` in `web/`. + +Commit generated Swagger files and `web/shared/types/api.ts`. + +## Error Handling + +- Repository and service errors continue to use `stacktrace.Propagate`. +- Repository transactions return errors instead of falling back to + success-shaped state. +- Missing user-owned threads preserve the repository not-found code. +- Migration and backfill failures stop startup. +- The web UI clears a badge only from a successful response or subsequent + reload. +- Failed automatic resets remain visible through the existing notification + path and do not block message display. + +## Testing + +### API unit and repository tests + +Cover: + +- entity schema defaults and removal of the public `is_read` field; +- migration of legacy read rows to zero and unread rows to one; +- new inbound thread initialization with count one and a ledger row; +- new outbound thread initialization with count zero; +- received-SMS increments; +- missed-call increments; +- duplicate inbound-event idempotency; +- outbound and delivery/status preservation; +- read reset clearing both counter and ledger; +- archive-only and reset-only field isolation; +- combined archive/reset atomicity; +- old delayed inbound event losing to a newer read watermark; +- genuinely new inbound event incrementing after a read; +- unread-item deletion decrementing once for both last and non-last messages; +- deletion after read and repeated deletion preserving zero; +- counter underflow protection; +- request validation accepting only `unread_count=0`; +- handler responses exposing `unread_count` and preserving not-found errors. + +Run: + +```bash +cd api +go test ./... +``` + +### Web validation + +Cover the store and component behavior with existing frontend test facilities +where available, including zero/nonzero logic, exact badge values, and `99+`. +Then run: + +```bash +cd web +pnpm lint +pnpm run generate +``` + +### Integration + +Extend the read-receipts integration coverage to exercise: + +1. a received SMS increments a thread to one; +2. replaying that event does not increment again; +3. another received SMS increments to two; +4. opening/resetting the thread returns the count to zero; +5. a missed call increments to one; +6. deleting that unread missed-call message returns the count to zero; +7. outbound/status activity does not change the count. + +Run: + +```bash +cd tests +go test -v -timeout 120s -run TestMessageThreadReadReceipts ./... +``` + +## Out of Scope + +- Android unread-count UI. +- Per-device or per-user-within-a-shared-account read positions. +- Manual "mark unread" behavior. +- Client-assigned nonzero counts. +- Unread-count badges outside the thread list. +- Reconstructing an exact historical count for threads that were already + unread before deployment. diff --git a/tests/README.md b/tests/README.md index 20e4f5e1..79f1fbb1 100644 --- a/tests/README.md +++ b/tests/README.md @@ -53,7 +53,7 @@ The API's Firebase SDK is configured (via `FCM_ENDPOINT` env var) to redirect al - [x] **Send SMS E2E** — Full send lifecycle: API → FCM push → emulator responds with SENT/DELIVERED events → message reaches `delivered` status - [x] **Receive SMS E2E** — Phone submits received message to API → message is stored and retrievable via GET endpoint -- [x] **Message thread read receipts E2E** — Incoming SMS and missed calls mark a thread unread, the existing thread update endpoint marks it read, and outbound activity preserves unread state +- [x] **Message thread unread count E2E** — Incoming SMS and missed calls increment the unread count, the existing thread update endpoint resets it, outbound activity preserves it, and deleting an unread item decrements it - [x] **Unarchive Thread on Receive E2E** — Archived thread returns to the inbox on inbound message when the phone's `unarchive_thread` setting is enabled, and stays archived when disabled - [x] **Contacts E2E** — JSON CRUD, search and pagination totals, CSV import normalization, and contact details attached to message threads diff --git a/tests/read_receipts_test.go b/tests/read_receipts_test.go index a65e7eb7..8e709967 100644 --- a/tests/read_receipts_test.go +++ b/tests/read_receipts_test.go @@ -19,10 +19,14 @@ import ( type integrationMessageThread struct { ID string `json:"id"` Contact string `json:"contact"` - IsRead bool `json:"is_read"` + UnreadCount uint `json:"unread_count"` LastMessageContent *string `json:"last_message_content"` } +type integrationMessage struct { + ID string `json:"id"` +} + func requestJSON( ctx context.Context, t *testing.T, @@ -109,7 +113,7 @@ func waitForMessageThread( return integrationMessageThread{} } -func markMessageThreadRead(ctx context.Context, t *testing.T, threadID string) integrationMessageThread { +func resetMessageThreadUnreadCount(ctx context.Context, t *testing.T, threadID string) integrationMessageThread { t.Helper() var response struct { @@ -121,7 +125,7 @@ func markMessageThreadRead(ctx context.Context, t *testing.T, threadID string) i http.MethodPut, "/v1/message-threads/"+threadID, userAPIKey, - map[string]any{"is_read": true}, + map[string]any{"unread_count": 0}, http.StatusOK, &response, ) @@ -152,19 +156,45 @@ func TestMessageThreadReadReceipts(t *testing.T) { ) thread := waitForMessageThread(ctx, t, phone.PhoneNumber, contact, 20*time.Second, func(thread integrationMessageThread) bool { - return !thread.IsRead + return thread.UnreadCount == 1 }) - assert.False(t, thread.IsRead) + assert.Equal(t, uint(1), thread.UnreadCount) + + requestJSON( + ctx, + t, + http.MethodPost, + "/v1/messages/receive", + phone.PhoneAPIKey, + map[string]any{ + "from": contact, + "to": phone.PhoneNumber, + "content": "Second unread inbound message", + "encrypted": false, + "sim": "SIM1", + "timestamp": time.Now().UTC().Format(time.RFC3339Nano), + }, + http.StatusOK, + nil, + ) - updated := markMessageThreadRead(ctx, t, thread.ID) - assert.True(t, updated.IsRead) + thread = waitForMessageThread(ctx, t, phone.PhoneNumber, contact, 20*time.Second, func(thread integrationMessageThread) bool { + return thread.UnreadCount == 2 + }) + assert.Equal(t, uint(2), thread.UnreadCount) + + updated := resetMessageThreadUnreadCount(ctx, t, thread.ID) + assert.Zero(t, updated.UnreadCount) assert.Equal(t, contact, updated.Contact) require.NotNil(t, updated.LastMessageContent) - assert.Equal(t, "Unread inbound message", *updated.LastMessageContent) + assert.Equal(t, "Second unread inbound message", *updated.LastMessageContent) waitForMessageThread(ctx, t, phone.PhoneNumber, contact, 10*time.Second, func(thread integrationMessageThread) bool { - return thread.IsRead + return thread.UnreadCount == 0 }) + var missedCallResponse struct { + Data integrationMessage `json:"data"` + } requestJSON( ctx, t, @@ -178,15 +208,16 @@ func TestMessageThreadReadReceipts(t *testing.T) { "timestamp": time.Now().UTC().Format(time.RFC3339Nano), }, http.StatusOK, - nil, + &missedCallResponse, ) + require.NotEmpty(t, missedCallResponse.Data.ID) thread = waitForMessageThread(ctx, t, phone.PhoneNumber, contact, 20*time.Second, func(thread integrationMessageThread) bool { - return !thread.IsRead && + return thread.UnreadCount == 1 && thread.LastMessageContent != nil && *thread.LastMessageContent == "Missed phone call" }) - assert.False(t, thread.IsRead) + assert.Equal(t, uint(1), thread.UnreadCount) outboundContent := "Outbound activity preserves unread" client := newAPIClient() @@ -199,8 +230,26 @@ func TestMessageThreadReadReceipts(t *testing.T) { require.Equal(t, http.StatusOK, response.HTTPResponse.StatusCode) thread = waitForMessageThread(ctx, t, phone.PhoneNumber, contact, 20*time.Second, func(thread integrationMessageThread) bool { - return thread.LastMessageContent != nil && + return thread.UnreadCount == 1 && + thread.LastMessageContent != nil && + *thread.LastMessageContent == outboundContent + }) + assert.Equal(t, uint(1), thread.UnreadCount, "outbound activity must preserve unread count") + + requestJSON( + ctx, + t, + http.MethodDelete, + "/v1/messages/"+missedCallResponse.Data.ID, + userAPIKey, + nil, + http.StatusNoContent, + nil, + ) + thread = waitForMessageThread(ctx, t, phone.PhoneNumber, contact, 20*time.Second, func(thread integrationMessageThread) bool { + return thread.UnreadCount == 0 && + thread.LastMessageContent != nil && *thread.LastMessageContent == outboundContent }) - assert.False(t, thread.IsRead, "outbound activity must not clear unread state") + assert.Zero(t, thread.UnreadCount) } diff --git a/web/app/components/MessageThread.vue b/web/app/components/MessageThread.vue index 51de93fc..93c77324 100644 --- a/web/app/components/MessageThread.vue +++ b/web/app/components/MessageThread.vue @@ -22,6 +22,22 @@ function threadDate(date: string): string { }) } +function hasUnreadMessages(unreadCount: number): boolean { + return unreadCount > 0 +} + +function unreadBadge( + unreadCount: number, +): false | { color: string; content: string; dot: false } { + if (unreadCount === 0) return false + + return { + color: 'primary', + content: unreadCount > 99 ? '99+' : String(unreadCount), + dot: false, + } +} + function onInstallApp() { notificationsStore.addNotification({ type: 'info', @@ -118,7 +134,7 @@ function threadAvatarInitial(thread: EntitiesMessageThread): string { {{ mdiAccount @@ -128,12 +144,20 @@ function threadAvatarInitial(thread: EntitiesMessageThread): string { }} - {{ - threadContactTitle(thread) - }} + + {{ threadContactTitle(thread) }} + {{ thread.last_message_content }} diff --git a/web/app/pages/threads/[id]/index.vue b/web/app/pages/threads/[id]/index.vue index f933973d..d80f91c7 100644 --- a/web/app/pages/threads/[id]/index.vue +++ b/web/app/pages/threads/[id]/index.vue @@ -62,6 +62,10 @@ const formMessageRules = [ let webhookChannel: Channel | null = null +interface WebsocketMessageEvent { + message_id: string +} + const contactIsPhoneNumber = computed(() => { const thread = currentThread.value if (!thread) return false @@ -128,19 +132,43 @@ function scrollToElement() { hideMessages.value = false } -async function markCurrentThreadRead(force = false) { +async function resetCurrentThreadUnreadCount(force = false) { const threadId = route.params.id as string try { - await threadsStore.markThreadRead(threadId, force) + await threadsStore.resetThreadUnreadCount(threadId, force) } catch (error) { console.error(error) } } -function loadMessages(hide = true, markRead = true) { +async function handleInboundMessage(event: WebsocketMessageEvent) { + if (loadingMessages.value) return + + try { + const message = await messagesStore.getMessage(event.message_id) + await threadsStore.loadThreads() + + const thread = currentThread.value + if ( + !thread || + message.owner !== thread.owner || + message.contact !== thread.contact || + loadingMessages.value + ) { + return + } + + await resetCurrentThreadUnreadCount(true) + loadMessages(false, false) + } catch (error) { + console.error(error) + } +} + +function loadMessages(hide = true, resetUnreadCount = true) { loadingMessages.value = true const threadId = route.params.id as string - if (markRead) void markCurrentThreadRead() + if (resetUnreadCount) void resetCurrentThreadUnreadCount() threadsStore .loadThreadMessages(threadId) .then((response: EntitiesMessage[]) => { @@ -257,17 +285,14 @@ onMounted(async () => { webhookChannel.bind('message.send.failed', () => { if (!loadingMessages.value) loadMessages(false) }) - webhookChannel.bind('message.phone.received', () => { - if (!loadingMessages.value) { - void markCurrentThreadRead(true) - loadMessages(false, false) - } - }) - webhookChannel.bind('message.call.missed', () => { - if (!loadingMessages.value) { - void markCurrentThreadRead(true) - loadMessages(false, false) - } + webhookChannel.bind( + 'message.phone.received', + (event: WebsocketMessageEvent) => { + void handleInboundMessage(event) + }, + ) + webhookChannel.bind('message.call.missed', (event: WebsocketMessageEvent) => { + void handleInboundMessage(event) }) }) diff --git a/web/app/stores/messages.ts b/web/app/stores/messages.ts index abaf47b6..ae0df175 100644 --- a/web/app/stores/messages.ts +++ b/web/app/stores/messages.ts @@ -48,6 +48,13 @@ export const useMessagesStore = defineStore('messages', () => { }) } + async function getMessage(messageId: string): Promise { + const response = await apiFetch<{ data: EntitiesMessage }>( + `/v1/messages/${messageId}`, + ) + return response.data + } + async function searchMessages( payload: SearchMessagesRequest, ): Promise { @@ -88,6 +95,7 @@ export const useMessagesStore = defineStore('messages', () => { return { sendMessage, deleteMessage, + getMessage, searchMessages, sendBulkMessages, fetchBulkMessageOrders, diff --git a/web/app/stores/threads.ts b/web/app/stores/threads.ts index fdc7094e..81d04f0b 100644 --- a/web/app/stores/threads.ts +++ b/web/app/stores/threads.ts @@ -113,17 +113,17 @@ export const useThreadsStore = defineStore('threads', () => { }) } - async function markThreadRead(threadId: string, force = false) { + async function resetThreadUnreadCount(threadId: string, force = false) { const thread = threads.value.find((item) => item.id === threadId) if (!thread) throw new Error(`Cannot find thread with id ${threadId}`) - if (!force && thread.is_read) return + if (!force && thread.unread_count === 0) return try { const response = await apiFetch<{ data: EntitiesMessageThread }>( `/v1/message-threads/${threadId}`, { method: 'PUT', - body: { is_read: true }, + body: { unread_count: 0 }, }, ) replaceThread(response.data) @@ -174,7 +174,7 @@ export const useThreadsStore = defineStore('threads', () => { setThreadId, toggleArchive, updateThread, - markThreadRead, + resetThreadUnreadCount, deleteThread, resetState, } diff --git a/web/shared/types/api.ts b/web/shared/types/api.ts index 3f85c4fb..702b7a47 100644 --- a/web/shared/types/api.ts +++ b/web/shared/types/api.ts @@ -225,8 +225,6 @@ export interface EntitiesMessageThread { id: string; /** @example false */ is_archived: boolean; - /** @example true */ - is_read: boolean; /** @example "This is a sample message content" */ last_message_content: string; /** @example "32343a19-da5e-4b1b-a767-3298a73703ca" */ @@ -237,6 +235,8 @@ export interface EntitiesMessageThread { owner: string; /** @example "PENDING" */ status: string; + /** @example 2 */ + unread_count: number; /** @example "2022-06-05T14:26:09.527976+03:00" */ updated_at: string; /** @example "WB7DRDWrJZRGbYrv2CKGkqbzvqdC" */ @@ -527,8 +527,12 @@ export interface RequestsMessageSendScheduleWindow { export interface RequestsMessageThreadUpdate { /** @example true */ is_archived?: boolean; - /** @example true */ - is_read?: boolean; + /** + * @min 0 + * @max 0 + * @example 0 + */ + unread_count?: number; } export interface RequestsPhoneAPIKeyStoreRequest {