From fa0cd8c876d4eb802d75dedfb0082ad0d16455b1 Mon Sep 17 00:00:00 2001 From: Acho Arnold Ewin Date: Fri, 21 Aug 2026 12:41:05 +0300 Subject: [PATCH 01/12] docs: design unread message count Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../2026-08-21-unread-message-count-design.md | 358 ++++++++++++++++++ 1 file changed, 358 insertions(+) create mode 100644 docs/superpowers/specs/2026-08-21-unread-message-count-design.md 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. From 2a741b5e8d5879311ff484eadb5c4d6c021d65d5 Mon Sep 17 00:00:00 2001 From: Acho Arnold Ewin Date: Fri, 21 Aug 2026 12:43:28 +0300 Subject: [PATCH 02/12] docs: plan unread message count Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../plans/2026-08-21-unread-message-count.md | 728 ++++++++++++++++++ 1 file changed, 728 insertions(+) create mode 100644 docs/superpowers/plans/2026-08-21-unread-message-count.md 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" +``` From c264f8382d5d8350a617839ac52b424a260700aa Mon Sep 17 00:00:00 2001 From: Acho Arnold Ewin Date: Fri, 21 Aug 2026 12:52:51 +0300 Subject: [PATCH 03/12] feat(api): add unread count schema Introduce unread_count and unread-item ledger schema plus an idempotent startup migration that backfills legacy unread threads before dropping message_threads.is_read. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 332cfd7b-8a60-4e28-9ea8-61854aa712c3 --- api/pkg/di/container.go | 5 +- api/pkg/entities/message_thread.go | 2 +- api/pkg/entities/message_thread_test.go | 23 +- .../entities/message_thread_unread_item.go | 10 + .../migrations/message_thread_unread_count.go | 31 +++ .../message_thread_unread_count_test.go | 201 ++++++++++++++++++ 6 files changed, 262 insertions(+), 10 deletions(-) create mode 100644 api/pkg/entities/message_thread_unread_item.go create mode 100644 api/pkg/migrations/message_thread_unread_count.go create mode 100644 api/pkg/migrations/message_thread_unread_count_test.go diff --git a/api/pkg/di/container.go b/api/pkg/di/container.go index ad9a1885..2c832762 100644 --- a/api/pkg/di/container.go +++ b/api/pkg/di/container.go @@ -65,6 +65,7 @@ import ( "github.com/NdoleStudio/httpsms/pkg/entities" "github.com/NdoleStudio/httpsms/pkg/listeners" + "github.com/NdoleStudio/httpsms/pkg/migrations" "github.com/NdoleStudio/httpsms/pkg/repositories" "github.com/NdoleStudio/httpsms/pkg/services" "github.com/NdoleStudio/stacktrace" @@ -374,8 +375,8 @@ ALTER TABLE discords ADD CONSTRAINT IF NOT EXISTS uni_discords_server_id CHECK ( container.logger.Fatal(stacktrace.Propagatef(err, "cannot migrate %T", &entities.Message{})) } - if err = db.AutoMigrate(&entities.MessageThread{}); err != nil { - container.logger.Fatal(stacktrace.Propagatef(err, "cannot migrate %T", &entities.MessageThread{})) + if err = migrations.MigrateMessageThreadUnreadCount(db); err != nil { + container.logger.Fatal(stacktrace.Propagate(err, "cannot migrate message thread unread counts")) } 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 3766acc7..d7606769 100644 --- a/api/pkg/entities/message_thread.go +++ b/api/pkg/entities/message_thread.go @@ -12,7 +12,7 @@ type MessageThread struct { Owner string `json:"owner" example:"+18005550199"` Contact string `json:"contact" example:"+18005550100"` IsArchived bool `json:"is_archived" example:"false"` - IsRead bool `json:"is_read" gorm:"not null;default:true" example:"true"` + UnreadCount uint `json:"unread_count" gorm:"not null;default:0" example:"2"` LastReadAt time.Time `json:"-" gorm:"not null;default:CURRENT_TIMESTAMP"` UserID UserID `json:"user_id" example:"WB7DRDWrJZRGbYrv2CKGkqbzvqdC"` Color string `json:"color" example:"indigo"` diff --git a/api/pkg/entities/message_thread_test.go b/api/pkg/entities/message_thread_test.go index 1409dff2..67630b72 100644 --- a/api/pkg/entities/message_thread_test.go +++ b/api/pkg/entities/message_thread_test.go @@ -8,18 +8,27 @@ 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") + _, hasIsRead := threadType.FieldByName("IsRead") + assert.False(t, hasIsRead) + + unreadCount, ok := threadType.FieldByName("UnreadCount") 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")) + 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.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")) } + +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") +} diff --git a/api/pkg/entities/message_thread_unread_item.go b/api/pkg/entities/message_thread_unread_item.go new file mode 100644 index 00000000..949d86f8 --- /dev/null +++ b/api/pkg/entities/message_thread_unread_item.go @@ -0,0 +1,10 @@ +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;"` +} diff --git a/api/pkg/migrations/message_thread_unread_count.go b/api/pkg/migrations/message_thread_unread_count.go new file mode 100644 index 00000000..0264c646 --- /dev/null +++ b/api/pkg/migrations/message_thread_unread_count.go @@ -0,0 +1,31 @@ +package migrations + +import ( + "github.com/NdoleStudio/httpsms/pkg/entities" + "github.com/NdoleStudio/stacktrace" + "gorm.io/gorm" +) + +// MigrateMessageThreadUnreadCount migrates message thread unread count schema. +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 +} diff --git a/api/pkg/migrations/message_thread_unread_count_test.go b/api/pkg/migrations/message_thread_unread_count_test.go new file mode 100644 index 00000000..17cbb598 --- /dev/null +++ b/api/pkg/migrations/message_thread_unread_count_test.go @@ -0,0 +1,201 @@ +package migrations + +import ( + "context" + "database/sql" + "database/sql/driver" + "errors" + "io" + "strings" + "testing" + + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/driver/postgres" + "gorm.io/gorm" +) + +func TestMigrateMessageThreadUnreadCountSkipsLegacyBackfillWhenIsReadColumnMissing(t *testing.T) { + db, recorder := newMigrationTestDB(t, migrationTestDBOptions{}) + + err := MigrateMessageThreadUnreadCount(db) + + require.NoError(t, err) + assert.NotEmpty(t, recorder.execs) + assert.NotContains(t, strings.Join(recorder.execs, "\n"), `UPDATE "message_threads" SET "unread_count"=$1 WHERE is_read = $2 AND unread_count = $3`) + assert.NotContains(t, strings.Join(recorder.execs, "\n"), `ALTER TABLE "message_threads" DROP COLUMN "is_read"`) +} + +func TestMigrateMessageThreadUnreadCountPropagatesSchemaErrors(t *testing.T) { + db, _ := newMigrationTestDB(t, migrationTestDBOptions{ + failExecContains: `CREATE TABLE "message_threads"`, + }) + + err := MigrateMessageThreadUnreadCount(db) + + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot migrate message thread unread count schema") + assert.Contains(t, err.Error(), `create table failed: CREATE TABLE "message_threads"`) +} + +type migrationTestDBOptions struct { + failExecContains string +} + +type migrationTestRecorder struct { + execs []string + failExecContains string +} + +type migrationTestDriver struct { + recorder *migrationTestRecorder +} + +func (driver *migrationTestDriver) Open(string) (driver.Conn, error) { + return &migrationTestConn{recorder: driver.recorder}, nil +} + +type migrationTestConn struct { + recorder *migrationTestRecorder +} + +func (*migrationTestConn) Close() error { + return nil +} + +func (*migrationTestConn) Begin() (driver.Tx, error) { + return migrationTestTx{}, nil +} + +func (*migrationTestConn) BeginTx(context.Context, driver.TxOptions) (driver.Tx, error) { + return migrationTestTx{}, nil +} + +func (conn *migrationTestConn) Prepare(query string) (driver.Stmt, error) { + return &migrationTestStmt{conn: conn, query: query}, nil +} + +func (conn *migrationTestConn) ExecContext(_ context.Context, query string, _ []driver.NamedValue) (driver.Result, error) { + conn.recorder.execs = append(conn.recorder.execs, query) + if conn.recorder.failExecContains != "" && strings.Contains(query, conn.recorder.failExecContains) { + return nil, errors.New("create table failed: " + query) + } + return driver.RowsAffected(1), nil +} + +func (conn *migrationTestConn) QueryContext(_ context.Context, query string, _ []driver.NamedValue) (driver.Rows, error) { + upperQuery := strings.ToUpper(query) + switch { + case strings.Contains(upperQuery, "SELECT CURRENT_DATABASE()"): + return &migrationTestRows{ + columns: []string{"current_database"}, + values: [][]driver.Value{{"httpsms_test"}}, + }, nil + case strings.Contains(upperQuery, "FROM INFORMATION_SCHEMA.TABLES"): + return &migrationTestRows{ + columns: []string{"count"}, + values: [][]driver.Value{{int64(0)}}, + }, nil + case strings.Contains(upperQuery, "FROM INFORMATION_SCHEMA.COLUMNS"): + return &migrationTestRows{ + columns: []string{"count"}, + values: [][]driver.Value{{int64(0)}}, + }, nil + default: + return nil, errors.New("unexpected query: " + query) + } +} + +type migrationTestStmt struct { + conn *migrationTestConn + query string +} + +func (*migrationTestStmt) Close() error { + return nil +} + +func (*migrationTestStmt) NumInput() int { + return -1 +} + +func (stmt *migrationTestStmt) Exec(args []driver.Value) (driver.Result, error) { + return stmt.conn.ExecContext(context.Background(), stmt.query, migrationNamedValues(args)) +} + +func (stmt *migrationTestStmt) Query(args []driver.Value) (driver.Rows, error) { + return stmt.conn.QueryContext(context.Background(), stmt.query, migrationNamedValues(args)) +} + +type migrationTestTx struct{} + +func (migrationTestTx) Commit() error { + return nil +} + +func (migrationTestTx) Rollback() error { + return nil +} + +type migrationTestRows struct { + columns []string + values [][]driver.Value + index int +} + +func (rows *migrationTestRows) Columns() []string { + return rows.columns +} + +func (*migrationTestRows) Close() error { + return nil +} + +func (rows *migrationTestRows) Next(dest []driver.Value) error { + if rows.index >= len(rows.values) { + return io.EOF + } + + copy(dest, rows.values[rows.index]) + rows.index++ + return nil +} + +func migrationNamedValues(values []driver.Value) []driver.NamedValue { + namedValues := make([]driver.NamedValue, len(values)) + for index, value := range values { + namedValues[index] = driver.NamedValue{ + Ordinal: index + 1, + Value: value, + } + } + return namedValues +} + +func newMigrationTestDB(t *testing.T, options migrationTestDBOptions) (*gorm.DB, *migrationTestRecorder) { + t.Helper() + + recorder := &migrationTestRecorder{ + failExecContains: options.failExecContains, + } + driverName := "migration-test-" + strings.ReplaceAll(uuid.NewString(), "-", "") + sql.Register(driverName, &migrationTestDriver{recorder: recorder}) + + sqlDB, err := sql.Open(driverName, "") + require.NoError(t, err) + t.Cleanup(func() { + assert.NoError(t, sqlDB.Close()) + }) + + db, err := gorm.Open( + postgres.New(postgres.Config{ + Conn: sqlDB, + WithoutReturning: true, + }), + &gorm.Config{DisableAutomaticPing: true}, + ) + require.NoError(t, err) + + return db, recorder +} From e8c382658628daff2b9ca58c330401c783dc8dac Mon Sep 17 00:00:00 2001 From: Acho Arnold Ewin Date: Fri, 21 Aug 2026 13:02:39 +0300 Subject: [PATCH 04/12] feat(api): count unread thread items Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 332cfd7b-8a60-4e28-9ea8-61854aa712c3 --- .../gorm_message_thread_repository.go | 304 +++++++++-- .../gorm_message_thread_repository_test.go | 472 ++++++++++++++++-- .../repositories/message_thread_repository.go | 14 +- 3 files changed, 685 insertions(+), 105 deletions(-) diff --git a/api/pkg/repositories/gorm_message_thread_repository.go b/api/pkg/repositories/gorm_message_thread_repository.go index 09241804..9c5f2ab2 100644 --- a/api/pkg/repositories/gorm_message_thread_repository.go +++ b/api/pkg/repositories/gorm_message_thread_repository.go @@ -45,17 +45,13 @@ func messageThreadActivityUpdates(params MessageThreadActivityUpdate) map[string if params.Unarchive { updates["is_archived"] = false } - if params.MarkAsUnread { - updates["is_read"] = gorm.Expr( - "CASE WHEN last_read_at < ? THEN ? ELSE is_read END", - params.EventTimestamp, - false, - ) - } return updates } func messageThreadDeletedUpdates(params MessageThreadDeletedUpdate) map[string]any { + if !params.UpdateLastMessage { + return map[string]any{} + } return map[string]any{ "last_message_id": params.LastMessageID, "last_message_content": params.LastMessageContent, @@ -68,15 +64,65 @@ 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"] = 0 + updates["last_read_at"] = params.ReadAt } return updates } +func lockMessageThread(tx *gorm.DB, userID entities.UserID, threadID uuid.UUID) (*entities.MessageThread, error) { + thread := new(entities.MessageThread) + err := tx. + Clauses(clause.Locking{Strength: "UPDATE"}). + Where("user_id = ?", userID). + Where("id = ?", threadID). + First(thread). + Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, stacktrace.PropagateWithCodef( + err, + ErrCodeNotFound, + "message thread with ID [%s] for user [%s] does not exist", + threadID, + userID, + ) + } + if err != nil { + return nil, stacktrace.Propagatef(err, "cannot lock message thread with ID [%s] for user [%s]", threadID, userID) + } + return thread, nil +} + +func insertUnreadItem(tx *gorm.DB, item entities.MessageThreadUnreadItem) (bool, error) { + result := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&item) + if result.Error != nil { + return false, stacktrace.Propagatef( + result.Error, + "cannot insert unread ledger item for message [%s] in thread [%s]", + item.MessageID, + item.MessageThreadID, + ) + } + return result.RowsAffected == 1, nil +} + +func deleteUnreadItem(tx *gorm.DB, messageID uuid.UUID, threadID uuid.UUID) (bool, error) { + result := tx. + Where("message_id = ?", messageID). + Where("message_thread_id = ?", threadID). + Delete(&entities.MessageThreadUnreadItem{}) + if result.Error != nil { + return false, stacktrace.Propagatef( + result.Error, + "cannot delete unread ledger item for message [%s] in thread [%s]", + messageID, + threadID, + ) + } + return result.RowsAffected == 1, nil +} + func (repository *gormMessageThreadRepository) DeleteAllForUser(ctx context.Context, userID entities.UserID) error { ctx, span := repository.tracer.Start(ctx) defer span.End() @@ -106,39 +152,98 @@ 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)) + err := repository.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + thread, err := lockMessageThread(tx, params.UserID, params.MessageThreadID) + if err != nil { + return err + } + + deleted, err := deleteUnreadItem(tx, params.DeletedMessageID, params.MessageThreadID) + if err != nil { + return err + } + if deleted { + if err := tx. + Model(thread). + Where("user_id = ?", params.UserID). + Where("id = ?", params.MessageThreadID). + UpdateColumn("unread_count", gorm.Expr("GREATEST(unread_count - 1, 0)")). + Error; err != nil { + return stacktrace.Propagatef( + err, + "cannot decrement unread count for thread [%s] and user [%s]", + params.MessageThreadID, + params.UserID, + ) + } + if thread.UnreadCount > 0 { + thread.UnreadCount-- + } + } + + if !params.UpdateLastMessage { + return nil + } + if err := tx. + Model(thread). + Where("user_id = ?", params.UserID). + Where("id = ?", params.MessageThreadID). + Updates(messageThreadDeletedUpdates(params)). + Error; err != nil { + return stacktrace.Propagatef( + err, + "cannot update deleted-message metadata for thread [%s] and user [%s]", + params.MessageThreadID, + params.UserID, + ) + } + return nil + }) + if err != nil { + return repository.tracer.WrapErrorSpan( + span, + stacktrace.Propagatef( + err, + "cannot apply deleted message [%s] to thread [%s] for user [%s]", + params.DeletedMessageID, + params.MessageThreadID, + 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, thread *entities.MessageThread, unreadMessageID *uuid.UUID) 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 + return stacktrace.Propagatef(result.Error, "cannot insert message thread with ID [%s]", thread.ID) } - if result.RowsAffected == 0 || isRead { + if result.RowsAffected == 0 || unreadMessageID == nil { return nil } - return tx.Model(&entities.MessageThread{}). - Where("user_id = ?", thread.UserID). - Where("id = ?", thread.ID). - UpdateColumn("is_read", false). - Error + inserted, err := insertUnreadItem(tx, entities.MessageThreadUnreadItem{ + MessageID: *unreadMessageID, + MessageThreadID: thread.ID, + }) + if err != nil { + return err + } + if !inserted { + return stacktrace.NewErrorf( + "unread ledger item for message [%s] was not inserted for new thread [%s]", + *unreadMessageID, + thread.ID, + ) + } + return nil }) if err != nil { return repository.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot save message thread with ID [%s]", thread.ID)) @@ -152,16 +257,63 @@ 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 := 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). + Where("user_id = ?", params.UserID). + Where("id = ?", params.MessageThreadID). + Updates(messageThreadActivityUpdates(params)). + Error; err != nil { + return stacktrace.Propagatef( + err, + "cannot update message activity for thread [%s] and user [%s]", + params.MessageThreadID, + params.UserID, + ) + } + 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 + } + + if err := tx. + Model(thread). + Where("user_id = ?", params.UserID). + Where("id = ?", params.MessageThreadID). + UpdateColumn("unread_count", gorm.Expr("unread_count + ?", 1)). + Error; err != nil { + return stacktrace.Propagatef( + err, + "cannot increment unread count for thread [%s] and user [%s]", + params.MessageThreadID, + params.UserID, + ) + } + thread.UnreadCount++ + return nil + }) + 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 @@ -177,18 +329,70 @@ 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)) + + var thread *entities.MessageThread + err := repository.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var err error + thread, err = lockMessageThread(tx, userID, messageThreadID) + if err != nil { + return err + } + + updates := messageThreadStatusUpdates(params) + if len(updates) > 0 { + if err := tx. + Model(thread). + Clauses(clause.Returning{}). + Where("user_id = ?", userID). + Where("id = ?", messageThreadID). + Updates(updates). + Error; err != nil { + return 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 { + return nil + } + + thread.UnreadCount = *params.UnreadCount + thread.LastReadAt = params.ReadAt + if err := tx. + Where("message_thread_id = ?", messageThreadID). + Delete(&entities.MessageThreadUnreadItem{}). + Error; err != nil { + return stacktrace.Propagatef( + err, + "cannot clear unread ledger items for thread [%s] and user [%s]", + messageThreadID, + userID, + ) + } + return nil + }) + if err != nil { + return nil, repository.tracer.WrapErrorSpan( + span, + stacktrace.Propagatef(err, "cannot update status for thread [%s] and user [%s]", messageThreadID, userID), + ) } 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..403982c3 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,13 @@ type messageThreadTestStatement struct { } type messageThreadTestConnPool struct { - statements []messageThreadTestStatement + statements []messageThreadTestStatement + thread *entities.MessageThread + rowsAffected func(query string) int64 + queryDB *sql.DB + begins int + commits int + rollbacks int } func (messageThreadTestConnPool) PrepareContext(context.Context, string) (*sql.Stmt, error) { @@ -37,11 +45,18 @@ func (pool *messageThreadTestConnPool) ExecContext(_ context.Context, query stri query: query, args: append([]any(nil), args...), }) + 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 +64,109 @@ 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 (*messageThreadTestConnPool) Commit() error { +func (*messageThreadRowsConn) Prepare(string) (driver.Stmt, error) { + return nil, errors.New("unexpected Prepare") +} + +func (*messageThreadRowsConn) Close() error { + return nil +} + +func (*messageThreadRowsConn) Begin() (driver.Tx, error) { + return nil, errors.New("unexpected Begin") +} + +func (conn *messageThreadRowsConn) QueryContext(context.Context, string, []driver.NamedValue) (driver.Rows, error) { + rows := &messageThreadDriverRows{ + columns: []string{"id", "user_id", "last_read_at", "unread_count"}, + } + if conn.pool.thread != nil { + rows.values = []driver.Value{ + conn.pool.thread.ID.String(), + string(conn.pool.thread.UserID), + conn.pool.thread.LastReadAt, + int64(conn.pool.thread.UnreadCount), + } + } + 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 (*messageThreadTestConnPool) Rollback() error { +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 +185,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) MessageThreadRepository { + 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,23 +203,45 @@ func TestMessageThreadStorePreservesExplicitUnreadState(t *testing.T) { require.NoError(t, err) logger := &messageThreadTestLogger{} - repository := NewGormMessageThreadRepository(logger, telemetry.NewOtelLogger("test", logger), db) - thread := &entities.MessageThread{ - ID: uuid.New(), - IsRead: false, + return NewGormMessageThreadRepository(logger, telemetry.NewOtelLogger("test", logger), db) +} + +func messageThreadStatementIndex(pool *messageThreadTestConnPool, fragment string) int { + for index, statement := range pool.statements { + if strings.Contains(statement.query, fragment) { + return index + } } + return -1 +} - require.NoError(t, repository.Store(context.Background(), thread)) - assert.False(t, thread.IsRead) +func TestMessageThreadUnreadStoreCreatesInitialLedgerItem(t *testing.T) { + threadID := uuid.New() + messageID := uuid.New() + pool := &messageThreadTestConnPool{} + repository := newMessageThreadTestRepository(t, pool) + thread := &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + UnreadCount: 1, + } - 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) + require.NoError(t, repository.Store(context.Background(), thread, &messageID)) + require.Equal(t, 1, pool.begins) + require.Equal(t, 1, pool.commits) + require.Zero(t, pool.rollbacks) + + threadInsert := messageThreadStatementIndex(pool, `INSERT INTO "message_threads"`) + ledgerInsert := messageThreadStatementIndex(pool, `INSERT INTO "message_thread_unread_items"`) + require.NotEqual(t, -1, threadInsert) + require.NotEqual(t, -1, ledgerInsert) + assert.Less(t, threadInsert, ledgerInsert) + assert.Contains(t, pool.statements[ledgerInsert].query, "ON CONFLICT DO NOTHING") + assert.Contains(t, pool.statements[ledgerInsert].query, `"message_id"`) + assert.Contains(t, pool.statements[ledgerInsert].query, `"message_thread_id"`) } -func TestMessageThreadActivityUpdatesOwnOnlyMessageColumns(t *testing.T) { +func TestMessageThreadActivityUpdatesDoNotOwnUnreadColumns(t *testing.T) { messageID := uuid.New() updates := messageThreadActivityUpdates(MessageThreadActivityUpdate{ Timestamp: time.Date(2026, 7, 18, 7, 0, 0, 0, time.UTC), @@ -118,45 +256,139 @@ func TestMessageThreadActivityUpdatesOwnOnlyMessageColumns(t *testing.T) { "last_message_content": "hello", "status": entities.MessageStatus(entities.MessageStatusReceived), }, updates) - assert.NotContains(t, updates, "is_read") + assert.NotContains(t, updates, "unread_count") assert.NotContains(t, updates, "is_archived") assert.NotContains(t, updates, "last_read_at") } -func TestUpdateActivityMarksUnreadWithOneQuery(t *testing.T) { - pool := &messageThreadTestConnPool{} - db, err := gorm.Open( - postgres.New(postgres.Config{ - Conn: pool, - WithoutReturning: true, - }), - &gorm.Config{DisableAutomaticPing: true}, - ) +func TestMessageThreadActivityCountableItemInsertsLedgerAndIncrements(t *testing.T) { + threadID := uuid.New() + messageID := uuid.New() + userID := entities.UserID("user-id") + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: userID, + LastReadAt: time.Date(2026, 7, 19, 10, 0, 0, 0, time.UTC), + }, + } + 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, + EventTimestamp: time.Date(2026, 7, 19, 10, 0, 1, 0, time.UTC), + }) + require.NoError(t, err) + require.Equal(t, 1, pool.begins) + require.Equal(t, 1, pool.commits) + require.Zero(t, pool.rollbacks) + + lock := messageThreadStatementIndex(pool, "FOR UPDATE") + activity := messageThreadStatementIndex(pool, `"order_timestamp"`) + ledger := messageThreadStatementIndex(pool, `INSERT INTO "message_thread_unread_items"`) + increment := messageThreadStatementIndex(pool, "unread_count +") + require.NotEqual(t, -1, lock) + require.NotEqual(t, -1, activity) + require.NotEqual(t, -1, ledger) + require.NotEqual(t, -1, increment) + assert.Less(t, lock, activity) + assert.Less(t, activity, ledger) + assert.Less(t, ledger, increment) + assert.Contains(t, pool.statements[lock].query, "user_id =") + assert.Contains(t, pool.statements[lock].query, "id =") + assert.Contains(t, pool.statements[ledger].query, "ON CONFLICT DO NOTHING") +} - logger := &messageThreadTestLogger{} - repository := NewGormMessageThreadRepository(logger, telemetry.NewOtelLogger("test", logger), db) +func TestMessageThreadActivityDuplicateLedgerItemDoesNotIncrement(t *testing.T) { + threadID := uuid.New() + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + LastReadAt: time.Date(2026, 7, 19, 10, 0, 0, 0, time.UTC), + }, + rowsAffected: func(query string) int64 { + if strings.Contains(query, `INSERT INTO "message_thread_unread_items"`) { + return 0 + } + return 1 + }, + } + repository := newMessageThreadTestRepository(t, pool) - err = repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ - MessageThreadID: uuid.New(), + 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, + CountAsUnread: true, EventTimestamp: time.Date(2026, 7, 19, 10, 0, 1, 0, time.UTC), }) 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.NotEqual(t, -1, messageThreadStatementIndex(pool, `INSERT INTO "message_thread_unread_items"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "unread_count +")) +} + +func TestMessageThreadActivityAtWatermarkDoesNotCount(t *testing.T) { + watermark := 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"), + LastReadAt: watermark, + }, } - 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: watermark, + MessageID: uuid.New(), + Content: "hello", + Status: entities.MessageStatusReceived, + CountAsUnread: true, + EventTimestamp: watermark, + }) + + require.NoError(t, err) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `INSERT INTO "message_thread_unread_items"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "unread_count +")) +} + +func TestMessageThreadActivityMissingThreadReturnsScopedNotFound(t *testing.T) { + threadID := uuid.New() + userID := entities.UserID("user-id") + pool := &messageThreadTestConnPool{} + 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.Equal(t, 1, pool.begins) + require.Zero(t, pool.commits) + require.Equal(t, 1, pool.rollbacks) + lock := messageThreadStatementIndex(pool, "FOR UPDATE") + require.NotEqual(t, -1, lock) + assert.Contains(t, pool.statements[lock].query, "user_id =") + assert.Contains(t, pool.statements[lock].query, "id =") } func TestMessageThreadDeletedUpdatesPreserveStatusType(t *testing.T) { @@ -166,6 +398,7 @@ func TestMessageThreadDeletedUpdatesPreserveStatusType(t *testing.T) { LastMessageID: &messageID, LastMessageContent: &content, LastMessageStatus: entities.MessageStatusDelivered, + UpdateLastMessage: true, }) assert.Equal(t, map[string]any{ @@ -175,17 +408,93 @@ 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 TestMessageThreadDeletedUpdatesSkipLastMessageWhenNotRequested(t *testing.T) { + assert.Empty(t, messageThreadDeletedUpdates(MessageThreadDeletedUpdate{})) +} + +func TestMessageThreadDeletedMessageDecrementsWhenLedgerDeleted(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, + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + MessageThreadID: threadID, + UserID: entities.UserID("user-id"), + DeletedMessageID: messageID, + UpdateLastMessage: true, + LastMessageID: &previousMessageID, + LastMessageContent: &previousContent, + LastMessageStatus: entities.MessageStatusDelivered, + }) + + require.NoError(t, err) + require.Equal(t, 1, pool.begins) + require.Equal(t, 1, pool.commits) + require.Zero(t, pool.rollbacks) + lock := messageThreadStatementIndex(pool, "FOR UPDATE") + ledgerDelete := messageThreadStatementIndex(pool, `DELETE FROM "message_thread_unread_items"`) + decrement := messageThreadStatementIndex(pool, "GREATEST(unread_count - 1, 0)") + metadata := messageThreadStatementIndex(pool, `"last_message_id"`) + require.NotEqual(t, -1, lock) + require.NotEqual(t, -1, ledgerDelete) + require.NotEqual(t, -1, decrement) + require.NotEqual(t, -1, metadata) + assert.Less(t, lock, ledgerDelete) + assert.Less(t, ledgerDelete, decrement) + assert.Less(t, decrement, metadata) + assert.Contains(t, pool.statements[ledgerDelete].query, "message_id =") + assert.Contains(t, pool.statements[ledgerDelete].query, "message_thread_id =") +} + +func TestMessageThreadDeletedMessageWithoutLedgerDoesNotDecrement(t *testing.T) { + threadID := uuid.New() + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + UnreadCount: 0, + }, + rowsAffected: func(query string) int64 { + if strings.Contains(query, `DELETE FROM "message_thread_unread_items"`) { + return 0 + } + return 1 + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + MessageThreadID: threadID, + UserID: entities.UserID("user-id"), + DeletedMessageID: uuid.New(), + }) + + require.NoError(t, err) + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, `DELETE FROM "message_thread_unread_items"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "GREATEST")) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `"last_message_id"`)) +} + +func TestMessageThreadStatusUpdatesResetUnreadCount(t *testing.T) { + zero := uint(0) + readAt := time.Date(2026, 8, 21, 10, 0, 0, 0, time.UTC) updates := messageThreadStatusUpdates(MessageThreadStatusUpdate{ - IsRead: &isRead, - ReadAt: readAt, + UnreadCount: &zero, + ReadAt: readAt, }) assert.Equal(t, map[string]any{ - "is_read": true, + "unread_count": 0, "last_read_at": readAt, }, updates) assert.NotContains(t, updates, "is_archived") @@ -199,6 +508,71 @@ 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, "unread_count") assert.NotContains(t, updates, "last_read_at") } + +func TestMessageThreadStatusResetDeletesLedgerRows(t *testing.T) { + threadID := uuid.New() + readAt := time.Date(2026, 8, 21, 10, 0, 0, 0, time.UTC) + 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, + ReadAt: readAt, + }, + ) + + require.NoError(t, err) + require.NotNil(t, thread) + assert.Zero(t, thread.UnreadCount) + assert.Equal(t, readAt, thread.LastReadAt) + require.Equal(t, 1, pool.begins) + require.Equal(t, 1, pool.commits) + require.Zero(t, pool.rollbacks) + lock := messageThreadStatementIndex(pool, "FOR UPDATE") + statusUpdate := messageThreadStatementIndex(pool, `"unread_count"`) + ledgerDelete := messageThreadStatementIndex(pool, `DELETE FROM "message_thread_unread_items"`) + require.NotEqual(t, -1, lock) + require.NotEqual(t, -1, statusUpdate) + require.NotEqual(t, -1, ledgerDelete) + assert.Less(t, lock, statusUpdate) + assert.Less(t, statusUpdate, ledgerDelete) + assert.Contains(t, pool.statements[statusUpdate].query, `"last_read_at"`) + assert.Contains(t, pool.statements[ledgerDelete].query, "message_thread_id =") +} + +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..876e60a2 100644 --- a/api/pkg/repositories/message_thread_repository.go +++ b/api/pkg/repositories/message_thread_repository.go @@ -17,20 +17,22 @@ type MessageThreadActivityUpdate struct { MessageID uuid.UUID Content string Status entities.MessageStatus - MarkAsUnread bool + CountAsUnread bool EventTimestamp time.Time Unarchive bool } type MessageThreadStatusUpdate struct { - IsArchived *bool - IsRead *bool - ReadAt time.Time + 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 @@ -39,12 +41,12 @@ type MessageThreadDeletedUpdate struct { // 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, thread *entities.MessageThread, unreadMessageID *uuid.UUID) 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 From 0835caef05d3a7a5cc031558f66e4161eba17af7 Mon Sep 17 00:00:00 2001 From: Acho Arnold Ewin Date: Fri, 21 Aug 2026 17:31:35 +0300 Subject: [PATCH 05/12] feat(api): route unread count activity Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 332cfd7b-8a60-4e28-9ea8-61854aa712c3 --- api/pkg/listeners/message_thread_listener.go | 4 +- .../listeners/message_thread_listener_test.go | 42 ++- .../read_receipts_test_helpers_test.go | 19 +- .../gorm_message_thread_repository.go | 2 +- api/pkg/services/message_thread_service.go | 31 +- .../services/message_thread_service_test.go | 323 ++++++++++++++++-- 6 files changed, 366 insertions(+), 55 deletions(-) diff --git a/api/pkg/listeners/message_thread_listener.go b/api/pkg/listeners/message_thread_listener.go index f6d2c52c..61e8bcb1 100644 --- a/api/pkg/listeners/message_thread_listener.go +++ b/api/pkg/listeners/message_thread_listener.go @@ -217,7 +217,7 @@ func (listener *MessageThreadListener) OnMessagePhoneReceived(ctx context.Contex Status: entities.MessageStatusReceived, Content: payload.Content, MessageID: payload.MessageID, - MarkAsUnread: true, + CountAsUnread: true, EventTimestamp: event.Time(), } @@ -246,7 +246,7 @@ func (listener *MessageThreadListener) OnMessageCallMissed(ctx context.Context, Timestamp: payload.Timestamp, Content: "Missed phone call", MessageID: payload.MessageID, - MarkAsUnread: true, + CountAsUnread: true, EventTimestamp: event.Time(), } if err := listener.service.UpdateThread(ctx, params); err != nil { diff --git a/api/pkg/listeners/message_thread_listener_test.go b/api/pkg/listeners/message_thread_listener_test.go index 39a444cc..3159058d 100644 --- a/api/pkg/listeners/message_thread_listener_test.go +++ b/api/pkg/listeners/message_thread_listener_test.go @@ -34,7 +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.True(t, repository.activity.CountAsUnread) assert.Equal(t, event.Time(), repository.activity.EventTimestamp) } @@ -57,11 +57,49 @@ 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() + currentLastMessageID := uuid.New() + repository.thread = &entities.MessageThread{ + ID: uuid.New(), + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + LastMessageID: ¤tLastMessageID, + } + + 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.False(t, repository.deletedUpdate.UpdateLastMessage) + require.NotNil(t, repository.deletedUpdate.LastMessageID) + assert.Equal(t, previousMessageID, *repository.deletedUpdate.LastMessageID) +} + func newMessageThreadListenerForTest() (*listenerMessageThreadRepository, map[string]events.EventListener) { repository := &listenerMessageThreadRepository{} logger := &noopListenerLogger{} diff --git a/api/pkg/listeners/read_receipts_test_helpers_test.go b/api/pkg/listeners/read_receipts_test_helpers_test.go index 60cd27e6..7128b52b 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, *entities.MessageThread, *uuid.UUID) 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/repositories/gorm_message_thread_repository.go b/api/pkg/repositories/gorm_message_thread_repository.go index 9c5f2ab2..98c68f9a 100644 --- a/api/pkg/repositories/gorm_message_thread_repository.go +++ b/api/pkg/repositories/gorm_message_thread_repository.go @@ -319,7 +319,7 @@ func (repository *gormMessageThreadRepository) UpdateActivity(ctx context.Contex 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, diff --git a/api/pkg/services/message_thread_service.go b/api/pkg/services/message_thread_service.go index 532a223e..f3c773b6 100644 --- a/api/pkg/services/message_thread_service.go +++ b/api/pkg/services/message_thread_service.go @@ -52,7 +52,7 @@ type MessageThreadUpdateParams struct { MessageID uuid.UUID // Timestamp controls thread activity ordering; EventTimestamp is the server-side unread watermark. Timestamp time.Time - MarkAsUnread bool + CountAsUnread bool EventTimestamp time.Time } @@ -111,7 +111,7 @@ func (service *MessageThreadService) UpdateThread(ctx context.Context, params Me MessageID: params.MessageID, Content: params.Content, Status: params.Status, - MarkAsUnread: params.MarkAsUnread, + CountAsUnread: params.CountAsUnread, EventTimestamp: params.EventTimestamp, } @@ -136,7 +136,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 } @@ -147,9 +147,9 @@ 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, + ReadAt: time.Now().UTC(), } thread, err := service.repository.UpdateStatus(ctx, params.UserID, params.MessageThreadID, update) if err != nil { @@ -179,15 +179,12 @@ func (service *MessageThreadService) UpdateAfterDeletedMessage(ctx context.Conte 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 - } - + updateLastMessage := thread.LastMessageID != nil && *thread.LastMessageID == payload.MessageID if 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, @@ -212,7 +209,7 @@ func (service *MessageThreadService) createThread(ctx context.Context, params Me Contact: params.Contact, UserID: params.UserID, IsArchived: false, - IsRead: !params.MarkAsUnread, + UnreadCount: 0, LastReadAt: now, Color: service.getColor(), LastMessageContent: ¶ms.Content, @@ -223,7 +220,13 @@ func (service *MessageThreadService) createThread(ctx context.Context, params Me OrderTimestamp: params.Timestamp, } - if err := service.repository.Store(ctx, thread); err != nil { + var unreadMessageID *uuid.UUID + if params.CountAsUnread { + thread.UnreadCount = 1 + unreadMessageID = ¶ms.MessageID + } + + if err := service.repository.Store(ctx, thread, unreadMessageID); 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 8cc8bdcb..ebecdac8 100644 --- a/api/pkg/services/message_thread_service_test.go +++ b/api/pkg/services/message_thread_service_test.go @@ -6,6 +6,7 @@ import ( "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 +19,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, *entities.MessageThread, *uuid.UUID) 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, thread *entities.MessageThread, unreadMessageID *uuid.UUID) error { if stub.store != nil { - return stub.store(ctx, thread) + return stub.store(ctx, thread, unreadMessageID) } return nil } @@ -44,7 +47,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 +67,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 +78,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) } +func newMessageThreadServiceWithPhoneForTest(repository repositories.MessageThreadRepository, phoneRepository repositories.PhoneRepository) *MessageThreadService { + logger := &noopLogger{} + tracer := telemetry.NewOtelLogger("test", logger) + return NewMessageThreadService(logger, tracer, repository, phoneRepository, nil) +} + func TestUpdateThreadPassesUnreadWatermarkForInboundActivity(t *testing.T) { threadID := uuid.New() eventTimestamp := time.Date(2026, 7, 18, 7, 0, 0, 0, time.UTC) @@ -98,12 +149,12 @@ func TestUpdateThreadPassesUnreadWatermarkForInboundActivity(t *testing.T) { Content: "hello", Status: entities.MessageStatusReceived, Timestamp: eventTimestamp, - MarkAsUnread: true, + CountAsUnread: true, EventTimestamp: eventTimestamp, }) require.NoError(t, err) - assert.True(t, captured.MarkAsUnread) + assert.True(t, captured.CountAsUnread) assert.Equal(t, eventTimestamp, captured.EventTimestamp) } @@ -111,7 +162,7 @@ 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,139 @@ func TestUpdateThreadPreservesReadStateForOutboundActivity(t *testing.T) { }) require.NoError(t, err) - assert.False(t, captured.MarkAsUnread) + assert.False(t, captured.CountAsUnread) +} + +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, + EventTimestamp: time.Now().UTC(), + }) + + 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, + EventTimestamp: time.Now().UTC(), + }) + + require.NoError(t, err) + assert.False(t, captured.Unarchive) } -func TestCreateThreadSetsReadStateFromActivityDirection(t *testing.T) { +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 unreadMessageID *uuid.UUID 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 { + store: func(_ context.Context, thread *entities.MessageThread, messageID *uuid.UUID) error { stored = thread + unreadMessageID = messageID 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.Equal(t, test.wantUnreadCount, stored.UnreadCount) assert.False(t, stored.LastReadAt.IsZero()) + if test.wantUnreadMessage { + require.NotNil(t, unreadMessageID) + assert.Equal(t, messageID, *unreadMessageID) + } else { + assert.Nil(t, unreadMessageID) + } }) } } 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 +322,15 @@ 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.Same(t, &unreadCount, captured.UnreadCount) assert.False(t, captured.ReadAt.IsZero()) assert.True(t, thread.IsArchived) - assert.False(t, thread.IsRead) + assert.Zero(t, thread.UnreadCount) } func TestUpdateStatusPreservesNotFoundCode(t *testing.T) { @@ -211,16 +341,145 @@ 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 TestUpdateAfterDeletedMessageCleansUnreadLedgerForNonLastMessage(t *testing.T) { + threadID := uuid.New() + deletedMessageID := uuid.New() + currentLastMessageID := uuid.New() + previousMessageID := uuid.New() + previousStatus := entities.MessageStatus(entities.MessageStatusDelivered) + previousContent := "previous" + var captured repositories.MessageThreadDeletedUpdate + + repository := &messageThreadRepositoryStub{ + loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { + return &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + LastMessageID: ¤tLastMessageID, + }, 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", + PreviousMessageID: &previousMessageID, + PreviousMessageStatus: &previousStatus, + PreviousMessageContent: &previousContent, + }) + + require.NoError(t, err) + assert.Equal(t, threadID, captured.MessageThreadID) + assert.Equal(t, entities.UserID("user-id"), captured.UserID) + assert.Equal(t, deletedMessageID, captured.DeletedMessageID) + assert.False(t, captured.UpdateLastMessage) + require.NotNil(t, captured.LastMessageID) + assert.Equal(t, previousMessageID, *captured.LastMessageID) + require.NotNil(t, captured.LastMessageContent) + assert.Equal(t, previousContent, *captured.LastMessageContent) + assert.Equal(t, previousStatus, captured.LastMessageStatus) +} + +func TestUpdateAfterDeletedMessageUpdatesLastMessageWhenDeletedMessageIsCurrent(t *testing.T) { + threadID := uuid.New() + deletedMessageID := uuid.New() + previousMessageID := uuid.New() + previousStatus := entities.MessageStatus(entities.MessageStatusDelivered) + previousContent := "previous" + var captured repositories.MessageThreadDeletedUpdate + + repository := &messageThreadRepositoryStub{ + loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { + return &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + LastMessageID: &deletedMessageID, + }, 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", + PreviousMessageID: &previousMessageID, + PreviousMessageStatus: &previousStatus, + PreviousMessageContent: &previousContent, + }) + + require.NoError(t, err) + assert.Equal(t, deletedMessageID, captured.DeletedMessageID) + assert.True(t, captured.UpdateLastMessage) + require.NotNil(t, captured.LastMessageID) + assert.Equal(t, previousMessageID, *captured.LastMessageID) + assert.Equal(t, previousStatus, captured.LastMessageStatus) +} + +func TestUpdateAfterDeletedMessageDeletesThreadWhenNoPreviousMessageExists(t *testing.T) { + threadID := uuid.New() + deletedMessageID := uuid.New() + deleted := false + + repository := &messageThreadRepositoryStub{ + loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { + return &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + }, nil + }, + delete: func(_ context.Context, userID entities.UserID, id uuid.UUID) error { + deleted = true + assert.Equal(t, entities.UserID("user-id"), userID) + assert.Equal(t, threadID, id) + return nil + }, + updateAfterDelete: func(context.Context, repositories.MessageThreadDeletedUpdate) error { + t.Fatal("expected whole-thread delete instead of metadata update") + 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.True(t, deleted) +} + func TestShouldCheckUnarchive(t *testing.T) { service := &MessageThreadService{} From 44065d4fb1f8d4acc7e19d7359af781f600e60b9 Mon Sep 17 00:00:00 2001 From: Acho Arnold Ewin Date: Fri, 21 Aug 2026 17:39:52 +0300 Subject: [PATCH 06/12] feat(api): expose unread count reset Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 332cfd7b-8a60-4e28-9ea8-61854aa712c3 --- .../handlers/message_thread_handler_test.go | 34 ++++++++++++- .../requests/message_thread_update_request.go | 6 +-- .../message_thread_update_request_test.go | 33 +++++++++++-- .../message_thread_handler_validator.go | 11 ++++- .../message_thread_handler_validator_test.go | 48 +++++++++++++++++-- 5 files changed, 118 insertions(+), 14 deletions(-) diff --git a/api/pkg/handlers/message_thread_handler_test.go b/api/pkg/handlers/message_thread_handler_test.go index fefb8189..97037a44 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, *entities.MessageThread, *uuid.UUID) 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) + 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/requests/message_thread_update_request.go b/api/pkg/requests/message_thread_update_request.go index 3309fc91..0832ad81 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"` 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..40e9860f 100644 --- a/api/pkg/requests/message_thread_update_request_test.go +++ b/api/pkg/requests/message_thread_update_request_test.go @@ -1,25 +1,50 @@ package requests import ( + "encoding/json" "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) } 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) + }) + } +} From ced78d4bd68071fdcbedb36396ea82d863acc0ef Mon Sep 17 00:00:00 2001 From: Acho Arnold Ewin Date: Fri, 21 Aug 2026 17:44:47 +0300 Subject: [PATCH 07/12] docs(api): publish unread counts Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 332cfd7b-8a60-4e28-9ea8-61854aa712c3 --- api/docs/docs.go | 16 ++++++++-------- api/docs/swagger.json | 16 ++++++++-------- api/docs/swagger.yaml | 14 +++++++------- web/shared/types/api.ts | 8 ++++---- 4 files changed, 27 insertions(+), 27 deletions(-) diff --git a/api/docs/docs.go b/api/docs/docs.go index 8cd408cb..c40a094a 100644 --- a/api/docs/docs.go +++ b/api/docs/docs.go @@ -3697,12 +3697,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" ], @@ -3727,10 +3727,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" @@ -3751,6 +3747,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" @@ -4404,9 +4404,9 @@ const docTemplate = `{ "type": "boolean", "example": true }, - "is_read": { - "type": "boolean", - "example": true + "unread_count": { + "type": "integer", + "example": 0 } } }, diff --git a/api/docs/swagger.json b/api/docs/swagger.json index ac9c12c9..8db64d9e 100644 --- a/api/docs/swagger.json +++ b/api/docs/swagger.json @@ -3694,12 +3694,12 @@ "created_at", "id", "is_archived", - "is_read", "last_message_content", "last_message_id", "order_timestamp", "owner", "status", + "unread_count", "updated_at", "user_id" ], @@ -3724,10 +3724,6 @@ "type": "boolean", "example": false }, - "is_read": { - "type": "boolean", - "example": true - }, "last_message_content": { "type": "string", "example": "This is a sample message content" @@ -3748,6 +3744,10 @@ "type": "string", "example": "PENDING" }, + "unread_count": { + "type": "integer", + "example": 2 + }, "updated_at": { "type": "string", "example": "2022-06-05T14:26:09.527976+03:00" @@ -4401,9 +4401,9 @@ "type": "boolean", "example": true }, - "is_read": { - "type": "boolean", - "example": true + "unread_count": { + "type": "integer", + "example": 0 } } }, diff --git a/api/docs/swagger.yaml b/api/docs/swagger.yaml index ddc6ae70..838b3737 100644 --- a/api/docs/swagger.yaml +++ b/api/docs/swagger.yaml @@ -319,9 +319,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 @@ -337,6 +334,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 @@ -349,12 +349,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 @@ -857,9 +857,9 @@ definitions: is_archived: example: true type: boolean - is_read: - example: true - type: boolean + unread_count: + example: 0 + type: integer type: object requests.PhoneAPIKeyStoreRequest: properties: diff --git a/web/shared/types/api.ts b/web/shared/types/api.ts index 27e2db11..46da6f72 100644 --- a/web/shared/types/api.ts +++ b/web/shared/types/api.ts @@ -205,8 +205,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" */ @@ -217,6 +215,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" */ @@ -487,8 +487,8 @@ export interface RequestsMessageSendScheduleWindow { export interface RequestsMessageThreadUpdate { /** @example true */ is_archived?: boolean; - /** @example true */ - is_read?: boolean; + /** @example 0 */ + unread_count?: number; } export interface RequestsPhoneAPIKeyStoreRequest { From bbce8a96b058ebbb44a3caf3393b483aa808ad76 Mon Sep 17 00:00:00 2001 From: Acho Arnold Ewin Date: Fri, 21 Aug 2026 21:28:39 +0300 Subject: [PATCH 08/12] feat(web): show unread message counts Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 332cfd7b-8a60-4e28-9ea8-61854aa712c3 --- web/app/components/MessageThread.vue | 34 ++++++++++++++++++++++++---- web/app/pages/threads/[id]/index.vue | 12 +++++----- web/app/stores/threads.ts | 8 +++---- 3 files changed, 39 insertions(+), 15 deletions(-) diff --git a/web/app/components/MessageThread.vue b/web/app/components/MessageThread.vue index 2a983431..63c3e558 100644 --- a/web/app/components/MessageThread.vue +++ b/web/app/components/MessageThread.vue @@ -21,6 +21,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', @@ -102,7 +118,7 @@ function onInstallApp() { {{ mdiAccount @@ -112,12 +128,20 @@ function onInstallApp() { }} - {{ - formatPhoneNumber(thread.contact) - }} + + {{ formatPhoneNumber(thread.contact) }} + {{ thread.last_message_content }} diff --git a/web/app/pages/threads/[id]/index.vue b/web/app/pages/threads/[id]/index.vue index 73632a64..ca8701ce 100644 --- a/web/app/pages/threads/[id]/index.vue +++ b/web/app/pages/threads/[id]/index.vue @@ -112,19 +112,19 @@ 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) { +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[]) => { @@ -243,13 +243,13 @@ onMounted(async () => { }) webhookChannel.bind('message.phone.received', () => { if (!loadingMessages.value) { - void markCurrentThreadRead(true) + void resetCurrentThreadUnreadCount(true) loadMessages(false, false) } }) webhookChannel.bind('message.call.missed', () => { if (!loadingMessages.value) { - void markCurrentThreadRead(true) + void resetCurrentThreadUnreadCount(true) loadMessages(false, false) } }) diff --git a/web/app/stores/threads.ts b/web/app/stores/threads.ts index 7e300262..9e9fab92 100644 --- a/web/app/stores/threads.ts +++ b/web/app/stores/threads.ts @@ -100,17 +100,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) @@ -161,7 +161,7 @@ export const useThreadsStore = defineStore('threads', () => { setThreadId, toggleArchive, updateThread, - markThreadRead, + resetThreadUnreadCount, deleteThread, resetState, } From c543b4d2201f73da902372cd4e44adb520492122 Mon Sep 17 00:00:00 2001 From: Acho Arnold Ewin Date: Fri, 21 Aug 2026 21:29:45 +0300 Subject: [PATCH 09/12] test: cover unread message counts Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 332cfd7b-8a60-4e28-9ea8-61854aa712c3 --- tests/README.md | 2 +- tests/read_receipts_test.go | 77 ++++++++++++++++++++++++++++++------- 2 files changed, 64 insertions(+), 15 deletions(-) diff --git a/tests/README.md b/tests/README.md index 45bbb278..75263ad5 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 ## Prerequisites 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) } From 71246f82bf1a28114ca981d88079c3e8b956c713 Mon Sep 17 00:00:00 2001 From: Acho Arnold Ewin Date: Fri, 21 Aug 2026 21:52:03 +0300 Subject: [PATCH 10/12] fix(api): harden unread count concurrency Retain unread tombstones, generate read watermarks under lock, and retry all counter transactions on CockroachDB serialization failures. Refuse duplicate conversations before destructive migration, then create the unique identity needed for concurrent first-message fallback. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 332cfd7b-8a60-4e28-9ea8-61854aa712c3 --- api/docs/docs.go | 2 + api/docs/swagger.json | 2 + api/docs/swagger.yaml | 2 + api/pkg/entities/message_thread_test.go | 9 + .../entities/message_thread_unread_item.go | 1 + .../handlers/message_thread_handler_test.go | 2 +- .../read_receipts_test_helpers_test.go | 2 +- .../migrations/message_thread_unread_count.go | 56 ++- .../message_thread_unread_count_test.go | 131 ++++++- .../gorm_message_thread_repository.go | 221 +++++++---- .../gorm_message_thread_repository_test.go | 343 +++++++++++++++++- .../repositories/message_thread_repository.go | 9 +- .../requests/message_thread_update_request.go | 2 +- .../message_thread_update_request_test.go | 8 + api/pkg/services/message_thread_service.go | 23 +- .../services/message_thread_service_test.go | 81 +++-- ...26-08-21-unread-count-concurrency-fixes.md | 241 ++++++++++++ web/shared/types/api.ts | 6 +- 18 files changed, 1009 insertions(+), 132 deletions(-) create mode 100644 docs/superpowers/plans/2026-08-21-unread-count-concurrency-fixes.md diff --git a/api/docs/docs.go b/api/docs/docs.go index c40a094a..01b382fd 100644 --- a/api/docs/docs.go +++ b/api/docs/docs.go @@ -4406,6 +4406,8 @@ const docTemplate = `{ }, "unread_count": { "type": "integer", + "maximum": 0, + "minimum": 0, "example": 0 } } diff --git a/api/docs/swagger.json b/api/docs/swagger.json index 8db64d9e..75a6587f 100644 --- a/api/docs/swagger.json +++ b/api/docs/swagger.json @@ -4403,6 +4403,8 @@ }, "unread_count": { "type": "integer", + "maximum": 0, + "minimum": 0, "example": 0 } } diff --git a/api/docs/swagger.yaml b/api/docs/swagger.yaml index 838b3737..9a433c7a 100644 --- a/api/docs/swagger.yaml +++ b/api/docs/swagger.yaml @@ -859,6 +859,8 @@ definitions: type: boolean unread_count: example: 0 + maximum: 0 + minimum: 0 type: integer type: object requests.PhoneAPIKeyStoreRequest: diff --git a/api/pkg/entities/message_thread_test.go b/api/pkg/entities/message_thread_test.go index 67630b72..6010e11e 100644 --- a/api/pkg/entities/message_thread_test.go +++ b/api/pkg/entities/message_thread_test.go @@ -32,3 +32,12 @@ func TestMessageThreadUnreadItemUsesMessageIDAsPrimaryKey(t *testing.T) { require.True(t, ok) assert.Contains(t, messageID.Tag.Get("gorm"), "primaryKey") } + +func TestMessageThreadUnreadItemRetainsCountedState(t *testing.T) { + itemType := reflect.TypeOf(MessageThreadUnreadItem{}) + + counted, ok := itemType.FieldByName("Counted") + require.True(t, ok) + assert.Contains(t, counted.Tag.Get("gorm"), "not null") + assert.Contains(t, counted.Tag.Get("gorm"), "default:true") +} diff --git a/api/pkg/entities/message_thread_unread_item.go b/api/pkg/entities/message_thread_unread_item.go index 949d86f8..fa32e20b 100644 --- a/api/pkg/entities/message_thread_unread_item.go +++ b/api/pkg/entities/message_thread_unread_item.go @@ -6,5 +6,6 @@ import "github.com/google/uuid" 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;"` } diff --git a/api/pkg/handlers/message_thread_handler_test.go b/api/pkg/handlers/message_thread_handler_test.go index 97037a44..de935490 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, *uuid.UUID) error { +func (stub *messageThreadHandlerRepositoryStub) Store(context.Context, repositories.MessageThreadStoreParams) error { return nil } diff --git a/api/pkg/listeners/read_receipts_test_helpers_test.go b/api/pkg/listeners/read_receipts_test_helpers_test.go index 7128b52b..e812a297 100644 --- a/api/pkg/listeners/read_receipts_test_helpers_test.go +++ b/api/pkg/listeners/read_receipts_test_helpers_test.go @@ -29,7 +29,7 @@ type listenerMessageThreadRepository struct { thread *entities.MessageThread } -func (repository *listenerMessageThreadRepository) Store(context.Context, *entities.MessageThread, *uuid.UUID) error { +func (repository *listenerMessageThreadRepository) Store(context.Context, repositories.MessageThreadStoreParams) error { return nil } diff --git a/api/pkg/migrations/message_thread_unread_count.go b/api/pkg/migrations/message_thread_unread_count.go index 0264c646..e7d6f3db 100644 --- a/api/pkg/migrations/message_thread_unread_count.go +++ b/api/pkg/migrations/message_thread_unread_count.go @@ -6,25 +6,63 @@ import ( "gorm.io/gorm" ) +const messageThreadConversationIndexName = "idx_message_threads_conversation" + +type messageThreadConversationIndex struct { + UserID entities.UserID `gorm:"column:user_id;uniqueIndex:idx_message_threads_conversation"` + Owner string `gorm:"column:owner;uniqueIndex:idx_message_threads_conversation"` + Contact string `gorm:"column:contact;uniqueIndex:idx_message_threads_conversation"` +} + +func (messageThreadConversationIndex) TableName() string { + return "message_threads" +} + // MigrateMessageThreadUnreadCount migrates message thread unread count schema. 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 + needsConversationIndex := !db.Migrator().HasIndex("message_threads", messageThreadConversationIndexName) + if needsConversationIndex { + var duplicates []messageThreadConversationIndex + if err := db. + Model(&entities.MessageThread{}). + Select("user_id", "owner", "contact"). + Group("user_id, owner, contact"). + Having("COUNT(*) > ?", 1). + Limit(1). + Find(&duplicates). + Error; err != nil { + return stacktrace.Propagate(err, "cannot check duplicate message thread conversations") + } + if len(duplicates) != 0 { + return stacktrace.NewError( + "cannot create unique message thread conversation index: duplicate message thread conversations exist", + ) + } } - 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 db.Migrator().HasColumn("message_threads", "is_read") { + 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") + } + } + + if !needsConversationIndex { + return nil } - if err := db.Migrator().DropColumn("message_threads", "is_read"); err != nil { - return stacktrace.Propagate(err, "cannot drop legacy message thread is_read column") + if err := db.Migrator().CreateIndex(&messageThreadConversationIndex{}, messageThreadConversationIndexName); err != nil { + return stacktrace.Propagate(err, "cannot create unique message thread conversation index") } return nil diff --git a/api/pkg/migrations/message_thread_unread_count_test.go b/api/pkg/migrations/message_thread_unread_count_test.go index 17cbb598..63b787d4 100644 --- a/api/pkg/migrations/message_thread_unread_count_test.go +++ b/api/pkg/migrations/message_thread_unread_count_test.go @@ -25,6 +25,61 @@ func TestMigrateMessageThreadUnreadCountSkipsLegacyBackfillWhenIsReadColumnMissi assert.NotEmpty(t, recorder.execs) assert.NotContains(t, strings.Join(recorder.execs, "\n"), `UPDATE "message_threads" SET "unread_count"=$1 WHERE is_read = $2 AND unread_count = $3`) assert.NotContains(t, strings.Join(recorder.execs, "\n"), `ALTER TABLE "message_threads" DROP COLUMN "is_read"`) + assert.Contains(t, strings.Join(recorder.execs, "\n"), `"counted" boolean NOT NULL DEFAULT true`) +} + +func TestMigrateMessageThreadUnreadCountBackfillsBeforeDropAndSkipsOnSecondRun(t *testing.T) { + db, recorder := newMigrationTestDB(t, migrationTestDBOptions{hasLegacyIsRead: true}) + + require.NoError(t, MigrateMessageThreadUnreadCount(db)) + + backfill := migrationExecIndex(recorder, `UPDATE "message_threads" SET "unread_count"=$1 WHERE is_read = $2 AND unread_count = $3`) + drop := migrationExecIndex(recorder, `ALTER TABLE "message_threads" DROP COLUMN "is_read"`) + require.NotEqual(t, -1, backfill) + require.NotEqual(t, -1, drop) + assert.Less(t, backfill, drop) + assert.Equal(t, 1, migrationExecCount(recorder, `UPDATE "message_threads" SET "unread_count"=$1 WHERE is_read = $2 AND unread_count = $3`)) + assert.Equal(t, 1, migrationExecCount(recorder, `ALTER TABLE "message_threads" DROP COLUMN "is_read"`)) + + require.NoError(t, MigrateMessageThreadUnreadCount(db)) + assert.Equal(t, 1, migrationExecCount(recorder, `UPDATE "message_threads" SET "unread_count"=$1 WHERE is_read = $2 AND unread_count = $3`)) + assert.Equal(t, 1, migrationExecCount(recorder, `ALTER TABLE "message_threads" DROP COLUMN "is_read"`)) +} + +func TestMigrateMessageThreadUnreadCountRejectsDuplicateConversationIdentity(t *testing.T) { + db, recorder := newMigrationTestDB(t, migrationTestDBOptions{hasDuplicateConversations: true}) + + err := MigrateMessageThreadUnreadCount(db) + + require.Error(t, err) + assert.Contains(t, err.Error(), "duplicate message thread conversations") + assert.Equal(t, 0, migrationExecCount(recorder, `CREATE UNIQUE INDEX IF NOT EXISTS "idx_message_threads_conversation"`)) +} + +func TestMigrateMessageThreadUnreadCountRejectsDuplicatesBeforeLegacyMutation(t *testing.T) { + db, recorder := newMigrationTestDB(t, migrationTestDBOptions{ + hasLegacyIsRead: true, + hasDuplicateConversations: true, + }) + + err := MigrateMessageThreadUnreadCount(db) + + require.Error(t, err) + assert.Contains(t, err.Error(), "duplicate message thread conversations") + assert.Equal(t, 0, migrationExecCount(recorder, `UPDATE "message_threads" SET "unread_count"`)) + assert.Equal(t, 0, migrationExecCount(recorder, `ALTER TABLE "message_threads" DROP COLUMN "is_read"`)) +} + +func TestMigrateMessageThreadUnreadCountCreatesConversationIndexOnce(t *testing.T) { + db, recorder := newMigrationTestDB(t, migrationTestDBOptions{}) + + require.NoError(t, MigrateMessageThreadUnreadCount(db)) + require.NoError(t, MigrateMessageThreadUnreadCount(db)) + + assert.Equal(t, 1, migrationExecCount(recorder, `CREATE UNIQUE INDEX IF NOT EXISTS "idx_message_threads_conversation"`)) + index := migrationExecIndex(recorder, `CREATE UNIQUE INDEX IF NOT EXISTS "idx_message_threads_conversation"`) + require.NotEqual(t, -1, index) + assert.Contains(t, recorder.execs[index], `("user_id","owner","contact")`) } func TestMigrateMessageThreadUnreadCountPropagatesSchemaErrors(t *testing.T) { @@ -40,12 +95,17 @@ func TestMigrateMessageThreadUnreadCountPropagatesSchemaErrors(t *testing.T) { } type migrationTestDBOptions struct { - failExecContains string + failExecContains string + hasLegacyIsRead bool + hasDuplicateConversations bool } type migrationTestRecorder struct { - execs []string - failExecContains string + execs []string + failExecContains string + hasLegacyIsRead bool + hasDuplicateConversations bool + hasConversationIndex bool } type migrationTestDriver struct { @@ -81,10 +141,16 @@ func (conn *migrationTestConn) ExecContext(_ context.Context, query string, _ [] if conn.recorder.failExecContains != "" && strings.Contains(query, conn.recorder.failExecContains) { return nil, errors.New("create table failed: " + query) } + if strings.Contains(query, `ALTER TABLE "message_threads" DROP COLUMN "is_read"`) { + conn.recorder.hasLegacyIsRead = false + } + if strings.Contains(query, `CREATE UNIQUE INDEX IF NOT EXISTS "idx_message_threads_conversation"`) { + conn.recorder.hasConversationIndex = true + } return driver.RowsAffected(1), nil } -func (conn *migrationTestConn) QueryContext(_ context.Context, query string, _ []driver.NamedValue) (driver.Rows, error) { +func (conn *migrationTestConn) QueryContext(_ context.Context, query string, args []driver.NamedValue) (driver.Rows, error) { upperQuery := strings.ToUpper(query) switch { case strings.Contains(upperQuery, "SELECT CURRENT_DATABASE()"): @@ -98,10 +164,33 @@ func (conn *migrationTestConn) QueryContext(_ context.Context, query string, _ [ values: [][]driver.Value{{int64(0)}}, }, nil case strings.Contains(upperQuery, "FROM INFORMATION_SCHEMA.COLUMNS"): + count := int64(0) + if migrationArgsContain(args, "is_read") && conn.recorder.hasLegacyIsRead { + count = 1 + } return &migrationTestRows{ columns: []string{"count"}, - values: [][]driver.Value{{int64(0)}}, + values: [][]driver.Value{{count}}, + }, nil + case strings.Contains(upperQuery, "FROM PG_INDEXES"): + count := int64(0) + if conn.recorder.hasConversationIndex { + count = 1 + } + return &migrationTestRows{ + columns: []string{"count"}, + values: [][]driver.Value{{count}}, }, nil + case strings.Contains(upperQuery, `FROM "MESSAGE_THREADS"`) && + strings.Contains(upperQuery, "GROUP BY") && + strings.Contains(upperQuery, "HAVING"): + rows := &migrationTestRows{ + columns: []string{"user_id", "owner", "contact"}, + } + if conn.recorder.hasDuplicateConversations { + rows.values = [][]driver.Value{{"user-id", "+18005550199", "+18005550100"}} + } + return rows, nil default: return nil, errors.New("unexpected query: " + query) } @@ -173,11 +262,41 @@ func migrationNamedValues(values []driver.Value) []driver.NamedValue { return namedValues } +func migrationArgsContain(args []driver.NamedValue, value string) bool { + for _, arg := range args { + if arg.Value == value { + return true + } + } + return false +} + +func migrationExecIndex(recorder *migrationTestRecorder, fragment string) int { + for index, query := range recorder.execs { + if strings.Contains(query, fragment) { + return index + } + } + return -1 +} + +func migrationExecCount(recorder *migrationTestRecorder, fragment string) int { + count := 0 + for _, query := range recorder.execs { + if strings.Contains(query, fragment) { + count++ + } + } + return count +} + func newMigrationTestDB(t *testing.T, options migrationTestDBOptions) (*gorm.DB, *migrationTestRecorder) { t.Helper() recorder := &migrationTestRecorder{ - failExecContains: options.failExecContains, + failExecContains: options.failExecContains, + hasLegacyIsRead: options.hasLegacyIsRead, + hasDuplicateConversations: options.hasDuplicateConversations, } driverName := "migration-test-" + strings.ReplaceAll(uuid.NewString(), "-", "") sql.Register(driverName, &migrationTestDriver{recorder: recorder}) diff --git a/api/pkg/repositories/gorm_message_thread_repository.go b/api/pkg/repositories/gorm_message_thread_repository.go index 98c68f9a..12417df8 100644 --- a/api/pkg/repositories/gorm_message_thread_repository.go +++ b/api/pkg/repositories/gorm_message_thread_repository.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "time" "github.com/google/uuid" @@ -12,6 +13,7 @@ import ( "github.com/NdoleStudio/httpsms/pkg/entities" "github.com/NdoleStudio/httpsms/pkg/telemetry" "github.com/NdoleStudio/stacktrace" + "github.com/cockroachdb/cockroach-go/v2/crdb/crdbgorm" "gorm.io/gorm" ) @@ -20,6 +22,7 @@ type gormMessageThreadRepository struct { logger telemetry.Logger tracer telemetry.Tracer db *gorm.DB + now func() time.Time } // NewGormMessageThreadRepository creates the GORM version of the MessageRepository @@ -32,6 +35,7 @@ func NewGormMessageThreadRepository( logger: logger.WithService(fmt.Sprintf("%T", &gormMessageThreadRepository{})), tracer: tracer, db: db, + now: time.Now, } } @@ -59,14 +63,14 @@ func messageThreadDeletedUpdates(params MessageThreadDeletedUpdate) map[string]a } } -func messageThreadStatusUpdates(params MessageThreadStatusUpdate) map[string]any { +func messageThreadStatusUpdates(params MessageThreadStatusUpdate, readAt time.Time) map[string]any { updates := make(map[string]any) if params.IsArchived != nil { updates["is_archived"] = *params.IsArchived } if params.UnreadCount != nil { updates["unread_count"] = 0 - updates["last_read_at"] = params.ReadAt + updates["last_read_at"] = readAt } return updates } @@ -94,6 +98,37 @@ func lockMessageThread(tx *gorm.DB, userID entities.UserID, threadID uuid.UUID) return thread, nil } +func lockMessageThreadByConversation(tx *gorm.DB, userID entities.UserID, owner string, contact string) (*entities.MessageThread, error) { + thread := new(entities.MessageThread) + err := tx. + Clauses(clause.Locking{Strength: "UPDATE"}). + Where("user_id = ?", userID). + Where("owner = ?", owner). + Where("contact = ?", contact). + First(thread). + Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, stacktrace.PropagateWithCodef( + err, + ErrCodeNotFound, + "message thread for user [%s], owner [%s], and contact [%s] does not exist", + userID, + owner, + contact, + ) + } + if err != nil { + return nil, stacktrace.Propagatef( + err, + "cannot lock message thread for user [%s], owner [%s], and contact [%s]", + userID, + owner, + contact, + ) + } + return thread, nil +} + func insertUnreadItem(tx *gorm.DB, item entities.MessageThreadUnreadItem) (bool, error) { result := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&item) if result.Error != nil { @@ -107,15 +142,17 @@ func insertUnreadItem(tx *gorm.DB, item entities.MessageThreadUnreadItem) (bool, return result.RowsAffected == 1, nil } -func deleteUnreadItem(tx *gorm.DB, messageID uuid.UUID, threadID uuid.UUID) (bool, error) { +func markUnreadItemDeleted(tx *gorm.DB, messageID uuid.UUID, threadID uuid.UUID) (bool, error) { result := tx. + Model(&entities.MessageThreadUnreadItem{}). Where("message_id = ?", messageID). Where("message_thread_id = ?", threadID). - Delete(&entities.MessageThreadUnreadItem{}) + Where("counted = ?", true). + Update("counted", false) if result.Error != nil { return false, stacktrace.Propagatef( result.Error, - "cannot delete unread ledger item for message [%s] in thread [%s]", + "cannot mark unread ledger item deleted for message [%s] in thread [%s]", messageID, threadID, ) @@ -123,6 +160,49 @@ func deleteUnreadItem(tx *gorm.DB, messageID uuid.UUID, threadID uuid.UUID) (boo return result.RowsAffected == 1, nil } +func applyMessageThreadActivity(tx *gorm.DB, thread *entities.MessageThread, params MessageThreadActivityUpdate) error { + if err := tx. + Model(thread). + Where("user_id = ?", params.UserID). + Where("id = ?", thread.ID). + Updates(messageThreadActivityUpdates(params)). + Error; err != nil { + return stacktrace.Propagatef( + err, + "cannot update message activity for thread [%s] and user [%s]", + thread.ID, + params.UserID, + ) + } + if !params.CountAsUnread || !params.EventTimestamp.After(thread.LastReadAt) { + return nil + } + + inserted, err := insertUnreadItem(tx, entities.MessageThreadUnreadItem{ + MessageID: params.MessageID, + MessageThreadID: thread.ID, + }) + if err != nil || !inserted { + return err + } + + if err := tx. + Model(thread). + Where("user_id = ?", params.UserID). + Where("id = ?", thread.ID). + UpdateColumn("unread_count", gorm.Expr("unread_count + ?", 1)). + Error; err != nil { + return stacktrace.Propagatef( + err, + "cannot increment unread count for thread [%s] and user [%s]", + thread.ID, + params.UserID, + ) + } + thread.UnreadCount++ + return nil +} + func (repository *gormMessageThreadRepository) DeleteAllForUser(ctx context.Context, userID entities.UserID) error { ctx, span := repository.tracer.Start(ctx) defer span.End() @@ -152,13 +232,14 @@ func (repository *gormMessageThreadRepository) UpdateAfterDeletedMessage(ctx con ctx, span := repository.tracer.Start(ctx) defer span.End() - err := repository.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + err := crdbgorm.ExecuteTx(ctx, repository.db, nil, func(tx *gorm.DB) error { + tx = tx.WithContext(ctx) thread, err := lockMessageThread(tx, params.UserID, params.MessageThreadID) if err != nil { return err } - deleted, err := deleteUnreadItem(tx, params.DeletedMessageID, params.MessageThreadID) + deleted, err := markUnreadItemDeleted(tx, params.DeletedMessageID, params.MessageThreadID) if err != nil { return err } @@ -216,22 +297,61 @@ func (repository *gormMessageThreadRepository) UpdateAfterDeletedMessage(ctx con } // Store a new entities.MessageThread -func (repository *gormMessageThreadRepository) Store(ctx context.Context, thread *entities.MessageThread, unreadMessageID *uuid.UUID) error { +func (repository *gormMessageThreadRepository) Store(ctx context.Context, params MessageThreadStoreParams) error { ctx, span := repository.tracer.Start(ctx) defer span.End() - err := repository.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { - result := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(thread) + err := crdbgorm.ExecuteTx(ctx, repository.db, nil, func(tx *gorm.DB) error { + tx = tx.WithContext(ctx) + candidate := *params.Thread + result := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&candidate) if result.Error != nil { - return stacktrace.Propagatef(result.Error, "cannot insert message thread with ID [%s]", thread.ID) + return stacktrace.Propagatef(result.Error, "cannot insert message thread with ID [%s]", params.Thread.ID) } - if result.RowsAffected == 0 || unreadMessageID == nil { + if result.RowsAffected == 0 { + thread, err := lockMessageThreadByConversation( + tx, + params.Thread.UserID, + params.Thread.Owner, + params.Thread.Contact, + ) + if err != nil { + return err + } + if params.Thread.LastMessageID == nil { + return stacktrace.NewErrorf( + "cannot apply conflicting thread [%s] without a last message ID", + params.Thread.ID, + ) + } + content := "" + if params.Thread.LastMessageContent != nil { + content = *params.Thread.LastMessageContent + } + return applyMessageThreadActivity(tx, thread, MessageThreadActivityUpdate{ + MessageThreadID: thread.ID, + UserID: params.Thread.UserID, + Timestamp: params.Thread.OrderTimestamp, + MessageID: *params.Thread.LastMessageID, + Content: content, + Status: params.Thread.Status, + CountAsUnread: params.CountAsUnread, + EventTimestamp: params.EventTimestamp, + }) + } + if !params.CountAsUnread { return nil } + if params.Thread.LastMessageID == nil { + return stacktrace.NewErrorf( + "cannot store unread ledger item for new thread [%s] without a last message ID", + params.Thread.ID, + ) + } inserted, err := insertUnreadItem(tx, entities.MessageThreadUnreadItem{ - MessageID: *unreadMessageID, - MessageThreadID: thread.ID, + MessageID: *params.Thread.LastMessageID, + MessageThreadID: params.Thread.ID, }) if err != nil { return err @@ -239,14 +359,14 @@ func (repository *gormMessageThreadRepository) Store(ctx context.Context, thread if !inserted { return stacktrace.NewErrorf( "unread ledger item for message [%s] was not inserted for new thread [%s]", - *unreadMessageID, - thread.ID, + *params.Thread.LastMessageID, + params.Thread.ID, ) } return nil }) if err != nil { - return repository.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot save message thread with ID [%s]", thread.ID)) + return repository.tracer.WrapErrorSpan(span, stacktrace.Propagatef(err, "cannot save message thread with ID [%s]", params.Thread.ID)) } return nil @@ -257,52 +377,14 @@ func (repository *gormMessageThreadRepository) UpdateActivity(ctx context.Contex ctx, span := repository.tracer.Start(ctx) defer span.End() - err := repository.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + err := crdbgorm.ExecuteTx(ctx, repository.db, nil, func(tx *gorm.DB) error { + tx = tx.WithContext(ctx) thread, err := lockMessageThread(tx, params.UserID, params.MessageThreadID) if err != nil { return err } - if err := tx. - Model(thread). - Where("user_id = ?", params.UserID). - Where("id = ?", params.MessageThreadID). - Updates(messageThreadActivityUpdates(params)). - Error; err != nil { - return stacktrace.Propagatef( - err, - "cannot update message activity for thread [%s] and user [%s]", - params.MessageThreadID, - params.UserID, - ) - } - 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 - } - - if err := tx. - Model(thread). - Where("user_id = ?", params.UserID). - Where("id = ?", params.MessageThreadID). - UpdateColumn("unread_count", gorm.Expr("unread_count + ?", 1)). - Error; err != nil { - return stacktrace.Propagatef( - err, - "cannot increment unread count for thread [%s] and user [%s]", - params.MessageThreadID, - params.UserID, - ) - } - thread.UnreadCount++ - return nil + return applyMessageThreadActivity(tx, thread, params) }) if err != nil { return repository.tracer.WrapErrorSpan( @@ -342,17 +424,22 @@ func (repository *gormMessageThreadRepository) UpdateStatus( } var thread *entities.MessageThread - err := repository.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { - var err error - thread, err = lockMessageThread(tx, userID, messageThreadID) + err := crdbgorm.ExecuteTx(ctx, repository.db, nil, func(tx *gorm.DB) error { + tx = tx.WithContext(ctx) + thread = nil + lockedThread, err := lockMessageThread(tx, userID, messageThreadID) if err != nil { return err } - updates := messageThreadStatusUpdates(params) + var readAt time.Time + if params.UnreadCount != nil { + readAt = repository.now().UTC() + } + updates := messageThreadStatusUpdates(params, readAt) if len(updates) > 0 { if err := tx. - Model(thread). + Model(lockedThread). Clauses(clause.Returning{}). Where("user_id = ?", userID). Where("id = ?", messageThreadID). @@ -367,14 +454,15 @@ func (repository *gormMessageThreadRepository) UpdateStatus( } } if params.IsArchived != nil { - thread.IsArchived = *params.IsArchived + lockedThread.IsArchived = *params.IsArchived } if params.UnreadCount == nil { + thread = lockedThread return nil } - thread.UnreadCount = *params.UnreadCount - thread.LastReadAt = params.ReadAt + lockedThread.UnreadCount = *params.UnreadCount + lockedThread.LastReadAt = readAt if err := tx. Where("message_thread_id = ?", messageThreadID). Delete(&entities.MessageThreadUnreadItem{}). @@ -386,6 +474,7 @@ func (repository *gormMessageThreadRepository) UpdateStatus( userID, ) } + thread = lockedThread return nil }) if err != nil { diff --git a/api/pkg/repositories/gorm_message_thread_repository_test.go b/api/pkg/repositories/gorm_message_thread_repository_test.go index 403982c3..6b053b27 100644 --- a/api/pkg/repositories/gorm_message_thread_repository_test.go +++ b/api/pkg/repositories/gorm_message_thread_repository_test.go @@ -30,6 +30,7 @@ type messageThreadTestConnPool struct { statements []messageThreadTestStatement thread *entities.MessageThread rowsAffected func(query string) int64 + execError func(query string) error queryDB *sql.DB begins int commits int @@ -45,6 +46,11 @@ 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 } @@ -185,7 +191,17 @@ func (logger *messageThreadTestLogger) Debug(string) func (logger *messageThreadTestLogger) Fatal(error) {} func (logger *messageThreadTestLogger) Printf(string, ...interface{}) {} -func newMessageThreadTestRepository(t *testing.T, pool *messageThreadTestConnPool) MessageThreadRepository { +type messageThreadRetryableError struct{} + +func (messageThreadRetryableError) Error() string { + return "restart transaction" +} + +func (messageThreadRetryableError) SQLState() string { + return "40001" +} + +func newMessageThreadTestRepository(t *testing.T, pool *messageThreadTestConnPool) *gormMessageThreadRepository { t.Helper() pool.queryDB = sql.OpenDB(&messageThreadRowsConnector{pool: pool}) @@ -203,7 +219,9 @@ func newMessageThreadTestRepository(t *testing.T, pool *messageThreadTestConnPoo require.NoError(t, err) logger := &messageThreadTestLogger{} - return 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 { @@ -215,18 +233,165 @@ func messageThreadStatementIndex(pool *messageThreadTestConnPool, fragment strin 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 TestMessageThreadMutationsRetrySerializationFailures(t *testing.T) { + t.Run("store", func(t *testing.T) { + attempts := 0 + pool := &messageThreadTestConnPool{ + execError: func(query string) error { + if strings.Contains(query, `INSERT INTO "message_threads"`) { + attempts++ + if attempts == 1 { + return messageThreadRetryableError{} + } + } + return nil + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.Store(context.Background(), MessageThreadStoreParams{ + Thread: &entities.MessageThread{ + ID: uuid.New(), + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + }, + }) + + require.NoError(t, err) + assert.Equal(t, 2, attempts) + }) + + t.Run("activity", func(t *testing.T) { + attempts := 0 + threadID := uuid.New() + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ID: threadID, UserID: entities.UserID("user-id")}, + execError: func(query string) error { + if strings.Contains(query, `UPDATE "message_threads"`) && strings.Contains(query, `"order_timestamp"`) { + attempts++ + if attempts == 1 { + return messageThreadRetryableError{} + } + } + return nil + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ + MessageThreadID: threadID, + UserID: entities.UserID("user-id"), + Timestamp: time.Date(2026, 8, 21, 10, 0, 0, 0, time.UTC), + MessageID: uuid.New(), + Content: "hello", + Status: entities.MessageStatusReceived, + }) + + require.NoError(t, err) + assert.Equal(t, 2, attempts) + }) + + t.Run("status", func(t *testing.T) { + attempts := 0 + clockCalls := 0 + threadID := uuid.New() + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ID: threadID, UserID: entities.UserID("user-id"), UnreadCount: 1}, + execError: func(query string) error { + if strings.Contains(query, `UPDATE "message_threads"`) && strings.Contains(query, `"unread_count"`) { + attempts++ + if attempts == 1 { + return messageThreadRetryableError{} + } + } + return nil + }, + } + repository := newMessageThreadTestRepository(t, pool) + repository.now = func() time.Time { + clockCalls++ + return time.Date(2026, 8, 21, 10, 0, clockCalls, 0, time.UTC) + } + zero := uint(0) + + thread, err := repository.UpdateStatus( + context.Background(), + entities.UserID("user-id"), + threadID, + MessageThreadStatusUpdate{UnreadCount: &zero}, + ) + + require.NoError(t, err) + require.NotNil(t, thread) + assert.Equal(t, 2, attempts) + assert.Equal(t, 2, clockCalls) + assert.Equal(t, time.Date(2026, 8, 21, 10, 0, 2, 0, time.UTC), thread.LastReadAt) + }) + + t.Run("deleted message", func(t *testing.T) { + attempts := 0 + threadID := uuid.New() + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ID: threadID, UserID: entities.UserID("user-id"), UnreadCount: 1}, + execError: func(query string) error { + if strings.Contains(query, `UPDATE "message_thread_unread_items"`) { + attempts++ + if attempts == 1 { + return messageThreadRetryableError{} + } + } + return nil + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + MessageThreadID: threadID, + UserID: entities.UserID("user-id"), + DeletedMessageID: uuid.New(), + }) + + require.NoError(t, err) + assert.Equal(t, 2, attempts) + }) +} + func TestMessageThreadUnreadStoreCreatesInitialLedgerItem(t *testing.T) { threadID := uuid.New() messageID := uuid.New() pool := &messageThreadTestConnPool{} repository := newMessageThreadTestRepository(t, pool) thread := &entities.MessageThread{ - ID: threadID, - UserID: entities.UserID("user-id"), - UnreadCount: 1, + ID: threadID, + UserID: entities.UserID("user-id"), + UnreadCount: 1, + LastMessageID: &messageID, } - require.NoError(t, repository.Store(context.Background(), thread, &messageID)) + require.NoError(t, repository.Store(context.Background(), MessageThreadStoreParams{ + Thread: thread, + CountAsUnread: true, + })) require.Equal(t, 1, pool.begins) require.Equal(t, 1, pool.commits) require.Zero(t, pool.rollbacks) @@ -241,6 +406,65 @@ func TestMessageThreadUnreadStoreCreatesInitialLedgerItem(t *testing.T) { assert.Contains(t, pool.statements[ledgerInsert].query, `"message_thread_id"`) } +func TestMessageThreadStoreConflictAppliesLosingActivity(t *testing.T) { + winnerThreadID := uuid.New() + losingThreadID := uuid.New() + messageID := uuid.New() + userID := entities.UserID("user-id") + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: winnerThreadID, + UserID: userID, + UnreadCount: 1, + LastReadAt: time.Unix(0, 0).UTC(), + }, + 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, + LastReadAt: time.Unix(0, 0).UTC(), + LastMessageID: &messageID, + LastMessageContent: &content, + Status: entities.MessageStatusReceived, + OrderTimestamp: eventTimestamp, + } + + err := repository.Store(context.Background(), MessageThreadStoreParams{ + Thread: thread, + CountAsUnread: true, + EventTimestamp: eventTimestamp, + }) + + require.NoError(t, err) + lock := messageThreadStatementIndex(pool, "FOR UPDATE") + activity := messageThreadStatementIndexAfter(pool, `"order_timestamp"`, lock) + ledger := messageThreadStatementIndex(pool, `INSERT INTO "message_thread_unread_items"`) + increment := messageThreadStatementIndex(pool, "unread_count +") + require.NotEqual(t, -1, lock) + require.NotEqual(t, -1, activity) + require.NotEqual(t, -1, ledger) + require.NotEqual(t, -1, increment) + assert.Less(t, lock, activity) + assert.Less(t, activity, ledger) + assert.Less(t, ledger, increment) + assert.Contains(t, pool.statements[lock].query, "user_id =") + assert.Contains(t, pool.statements[lock].query, "owner =") + assert.Contains(t, pool.statements[lock].query, "contact =") + assert.Contains(t, pool.statements[ledger].args, winnerThreadID) +} + func TestMessageThreadActivityUpdatesDoNotOwnUnreadColumns(t *testing.T) { messageID := uuid.New() updates := messageThreadActivityUpdates(MessageThreadActivityUpdate{ @@ -441,7 +665,7 @@ func TestMessageThreadDeletedMessageDecrementsWhenLedgerDeleted(t *testing.T) { require.Equal(t, 1, pool.commits) require.Zero(t, pool.rollbacks) lock := messageThreadStatementIndex(pool, "FOR UPDATE") - ledgerDelete := messageThreadStatementIndex(pool, `DELETE FROM "message_thread_unread_items"`) + ledgerDelete := messageThreadStatementIndex(pool, `UPDATE "message_thread_unread_items"`) decrement := messageThreadStatementIndex(pool, "GREATEST(unread_count - 1, 0)") metadata := messageThreadStatementIndex(pool, `"last_message_id"`) require.NotEqual(t, -1, lock) @@ -453,6 +677,7 @@ func TestMessageThreadDeletedMessageDecrementsWhenLedgerDeleted(t *testing.T) { assert.Less(t, decrement, metadata) assert.Contains(t, pool.statements[ledgerDelete].query, "message_id =") assert.Contains(t, pool.statements[ledgerDelete].query, "message_thread_id =") + assert.Contains(t, pool.statements[ledgerDelete].query, "counted =") } func TestMessageThreadDeletedMessageWithoutLedgerDoesNotDecrement(t *testing.T) { @@ -464,7 +689,7 @@ func TestMessageThreadDeletedMessageWithoutLedgerDoesNotDecrement(t *testing.T) UnreadCount: 0, }, rowsAffected: func(query string) int64 { - if strings.Contains(query, `DELETE FROM "message_thread_unread_items"`) { + if strings.Contains(query, `UPDATE "message_thread_unread_items"`) { return 0 } return 1 @@ -479,19 +704,77 @@ func TestMessageThreadDeletedMessageWithoutLedgerDoesNotDecrement(t *testing.T) }) require.NoError(t, err) - assert.NotEqual(t, -1, messageThreadStatementIndex(pool, `DELETE FROM "message_thread_unread_items"`)) + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, `UPDATE "message_thread_unread_items"`)) assert.Equal(t, -1, messageThreadStatementIndex(pool, "GREATEST")) assert.Equal(t, -1, messageThreadStatementIndex(pool, `"last_message_id"`)) } +func TestMessageThreadDeletedItemReplayDoesNotIncrement(t *testing.T) { + threadID := uuid.New() + messageID := uuid.New() + userID := entities.UserID("user-id") + tombstoned := false + pool := &messageThreadTestConnPool{ + thread: &entities.MessageThread{ + ID: threadID, + UserID: userID, + UnreadCount: 1, + LastReadAt: time.Date(2026, 8, 21, 10, 0, 0, 0, time.UTC), + }, + rowsAffected: func(query string) int64 { + switch { + case strings.Contains(query, `UPDATE "message_thread_unread_items"`) && + strings.Contains(query, `"counted"`): + if tombstoned { + return 0 + } + tombstoned = true + return 1 + case strings.Contains(query, `INSERT INTO "message_thread_unread_items"`): + if tombstoned { + return 0 + } + return 1 + default: + return 1 + } + }, + } + repository := newMessageThreadTestRepository(t, pool) + + require.NoError(t, repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + MessageThreadID: threadID, + UserID: userID, + 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, + EventTimestamp: time.Date(2026, 8, 21, 10, 0, 1, 0, time.UTC), + })) + + tombstone := messageThreadStatementIndex(pool, `UPDATE "message_thread_unread_items"`) + replay := messageThreadStatementIndex(pool, `INSERT INTO "message_thread_unread_items"`) + require.NotEqual(t, -1, tombstone) + require.NotEqual(t, -1, replay) + assert.Less(t, tombstone, replay) + assert.Contains(t, pool.statements[tombstone].query, `"counted"`) + assert.Zero(t, messageThreadStatementCount(pool, "unread_count +")) + assert.Equal(t, 1, messageThreadStatementCount(pool, "GREATEST(unread_count - 1, 0)")) +} + 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, - }) + }, readAt) assert.Equal(t, map[string]any{ "unread_count": 0, @@ -505,7 +788,7 @@ func TestMessageThreadStatusUpdatesArchiveOnly(t *testing.T) { updates := messageThreadStatusUpdates(MessageThreadStatusUpdate{ IsArchived: &isArchived, - }) + }, time.Time{}) assert.Equal(t, map[string]any{"is_archived": true}, updates) assert.NotContains(t, updates, "unread_count") @@ -514,7 +797,6 @@ func TestMessageThreadStatusUpdatesArchiveOnly(t *testing.T) { func TestMessageThreadStatusResetDeletesLedgerRows(t *testing.T) { threadID := uuid.New() - readAt := time.Date(2026, 8, 21, 10, 0, 0, 0, time.UTC) zero := uint(0) pool := &messageThreadTestConnPool{ thread: &entities.MessageThread{ @@ -531,14 +813,13 @@ func TestMessageThreadStatusResetDeletesLedgerRows(t *testing.T) { threadID, MessageThreadStatusUpdate{ UnreadCount: &zero, - ReadAt: readAt, }, ) require.NoError(t, err) require.NotNil(t, thread) assert.Zero(t, thread.UnreadCount) - assert.Equal(t, readAt, thread.LastReadAt) + assert.False(t, thread.LastReadAt.IsZero()) require.Equal(t, 1, pool.begins) require.Equal(t, 1, pool.commits) require.Zero(t, pool.rollbacks) @@ -554,6 +835,38 @@ func TestMessageThreadStatusResetDeletesLedgerRows(t *testing.T) { assert.Contains(t, pool.statements[ledgerDelete].query, "message_thread_id =") } +func TestMessageThreadStatusResetCreatesUTCWatermarkAfterLock(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) + sourceWatermark := time.Date(2026, 8, 21, 12, 30, 0, 123, time.FixedZone("UTC+2", 2*60*60)) + repository.now = func() time.Time { + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, "FOR UPDATE")) + return sourceWatermark + } + + thread, err := repository.UpdateStatus( + context.Background(), + entities.UserID("user-id"), + threadID, + MessageThreadStatusUpdate{UnreadCount: &zero}, + ) + + require.NoError(t, err) + require.NotNil(t, thread) + assert.Equal(t, sourceWatermark.UTC(), thread.LastReadAt) + statusUpdate := messageThreadStatementIndex(pool, `"last_read_at"`) + require.NotEqual(t, -1, statusUpdate) + assert.Contains(t, pool.statements[statusUpdate].args, sourceWatermark.UTC()) +} + func TestMessageThreadStatusRejectsNonzeroUnreadCount(t *testing.T) { threadID := uuid.New() one := uint(1) diff --git a/api/pkg/repositories/message_thread_repository.go b/api/pkg/repositories/message_thread_repository.go index 876e60a2..92ae1b46 100644 --- a/api/pkg/repositories/message_thread_repository.go +++ b/api/pkg/repositories/message_thread_repository.go @@ -22,10 +22,15 @@ type MessageThreadActivityUpdate struct { Unarchive bool } +type MessageThreadStoreParams struct { + Thread *entities.MessageThread + CountAsUnread bool + EventTimestamp time.Time +} + type MessageThreadStatusUpdate struct { IsArchived *bool UnreadCount *uint - ReadAt time.Time } type MessageThreadDeletedUpdate struct { @@ -41,7 +46,7 @@ type MessageThreadDeletedUpdate struct { // MessageThreadRepository loads and persists an entities.MessageThread type MessageThreadRepository interface { // Store a new entities.MessageThread - Store(ctx context.Context, thread *entities.MessageThread, unreadMessageID *uuid.UUID) 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 diff --git a/api/pkg/requests/message_thread_update_request.go b/api/pkg/requests/message_thread_update_request.go index 0832ad81..95f0095f 100644 --- a/api/pkg/requests/message_thread_update_request.go +++ b/api/pkg/requests/message_thread_update_request.go @@ -11,7 +11,7 @@ import ( type MessageThreadUpdate struct { request IsArchived *bool `json:"is_archived,omitempty" example:"true"` - UnreadCount *uint `json:"unread_count,omitempty" example:"0"` + UnreadCount *uint `json:"unread_count,omitempty" example:"0" minimum:"0" maximum:"0"` MessageThreadID string `json:"messageThreadID" swaggerignore:"true"` // used internally for validation } diff --git a/api/pkg/requests/message_thread_update_request_test.go b/api/pkg/requests/message_thread_update_request_test.go index 40e9860f..d2f3e62a 100644 --- a/api/pkg/requests/message_thread_update_request_test.go +++ b/api/pkg/requests/message_thread_update_request_test.go @@ -2,6 +2,7 @@ package requests import ( "encoding/json" + "reflect" "testing" "github.com/NdoleStudio/httpsms/pkg/entities" @@ -48,3 +49,10 @@ func TestMessageThreadUpdateToUpdateParamsPreservesOptionalFields(t *testing.T) 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 f3c773b6..e74163cc 100644 --- a/api/pkg/services/message_thread_service.go +++ b/api/pkg/services/message_thread_service.go @@ -149,7 +149,6 @@ func (service *MessageThreadService) UpdateStatus(ctx context.Context, params Me update := repositories.MessageThreadStatusUpdate{ IsArchived: params.IsArchived, UnreadCount: params.UnreadCount, - ReadAt: time.Now().UTC(), } thread, err := service.repository.UpdateStatus(ctx, params.UserID, params.MessageThreadID, update) if err != nil { @@ -171,8 +170,16 @@ func (service *MessageThreadService) UpdateAfterDeletedMessage(ctx context.Conte 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 + return service.tracer.WrapErrorSpan( + span, + stacktrace.Propagatef( + err, + "cannot delete thread with ID [%s] for user [%s] and owner [%s]", + thread.ID, + thread.UserID, + thread.Owner, + ), + ) } msg := fmt.Sprintf("previous message ID is nil for thread with ID [%s] and user [%s]", thread.ID, thread.UserID) ctxLogger.Info(msg) @@ -210,7 +217,7 @@ func (service *MessageThreadService) createThread(ctx context.Context, params Me UserID: params.UserID, IsArchived: false, UnreadCount: 0, - LastReadAt: now, + LastReadAt: time.Unix(0, 0).UTC(), Color: service.getColor(), LastMessageContent: ¶ms.Content, Status: params.Status, @@ -220,13 +227,15 @@ func (service *MessageThreadService) createThread(ctx context.Context, params Me OrderTimestamp: params.Timestamp, } - var unreadMessageID *uuid.UUID if params.CountAsUnread { thread.UnreadCount = 1 - unreadMessageID = ¶ms.MessageID } - if err := service.repository.Store(ctx, thread, unreadMessageID); err != nil { + if err := service.repository.Store(ctx, repositories.MessageThreadStoreParams{ + Thread: thread, + CountAsUnread: params.CountAsUnread, + EventTimestamp: params.EventTimestamp, + }); 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 ebecdac8..27a15243 100644 --- a/api/pkg/services/message_thread_service_test.go +++ b/api/pkg/services/message_thread_service_test.go @@ -2,6 +2,8 @@ package services import ( "context" + "errors" + "reflect" "testing" "time" @@ -19,16 +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, *uuid.UUID) 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, unreadMessageID *uuid.UUID) error { +func (stub *messageThreadRepositoryStub) Store(ctx context.Context, params repositories.MessageThreadStoreParams) error { if stub.store != nil { - return stub.store(ctx, thread, unreadMessageID) + return stub.store(ctx, params) } return nil } @@ -267,41 +269,43 @@ func TestCreateThreadSetsUnreadCountFromActivityDirection(t *testing.T) { for _, test := range tests { t.Run(test.name, func(t *testing.T) { - var stored *entities.MessageThread - var unreadMessageID *uuid.UUID + 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, messageID *uuid.UUID) error { - stored = thread - unreadMessageID = messageID + store: func(_ context.Context, params repositories.MessageThreadStoreParams) error { + stored = params return nil }, } messageID := uuid.New() + eventTimestamp := time.Date(2026, 8, 21, 10, 0, 0, 0, time.UTC) service := newMessageThreadServiceForTest(repository) err := service.UpdateThread(context.Background(), MessageThreadUpdateParams{ - UserID: entities.UserID("user-id"), - Owner: "+18005550199", - Contact: "+18005550100", - MessageID: messageID, - Content: "hello", - Status: test.status, - Timestamp: time.Now().UTC(), - CountAsUnread: test.countAsUnread, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + MessageID: messageID, + Content: "hello", + Status: test.status, + Timestamp: time.Now().UTC(), + CountAsUnread: test.countAsUnread, + EventTimestamp: eventTimestamp, }) require.NoError(t, err) - require.NotNil(t, stored) - assert.Equal(t, test.wantUnreadCount, stored.UnreadCount) - assert.False(t, stored.LastReadAt.IsZero()) + require.NotNil(t, stored.Thread) + assert.Equal(t, test.wantUnreadCount, stored.Thread.UnreadCount) + assert.Equal(t, time.Unix(0, 0).UTC(), stored.Thread.LastReadAt) + assert.Equal(t, test.countAsUnread, stored.CountAsUnread) + assert.Equal(t, eventTimestamp, stored.EventTimestamp) if test.wantUnreadMessage { - require.NotNil(t, unreadMessageID) - assert.Equal(t, messageID, *unreadMessageID) + require.NotNil(t, stored.Thread.LastMessageID) + assert.Equal(t, messageID, *stored.Thread.LastMessageID) } else { - assert.Nil(t, unreadMessageID) + assert.False(t, stored.CountAsUnread) } }) } @@ -328,7 +332,8 @@ func TestUpdateStatusChangesOnlyRequestedState(t *testing.T) { require.NoError(t, err) assert.Nil(t, captured.IsArchived) assert.Same(t, &unreadCount, captured.UnreadCount) - assert.False(t, captured.ReadAt.IsZero()) + _, hasReadAt := reflect.TypeOf(captured).FieldByName("ReadAt") + assert.False(t, hasReadAt) assert.True(t, thread.IsArchived) assert.Zero(t, thread.UnreadCount) } @@ -480,6 +485,36 @@ func TestUpdateAfterDeletedMessageDeletesThreadWhenNoPreviousMessageExists(t *te assert.True(t, deleted) } +func TestUpdateAfterDeletedMessagePropagatesFinalThreadDeleteError(t *testing.T) { + threadID := uuid.New() + deleteErr := errors.New("delete failed") + repository := &messageThreadRepositoryStub{ + loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { + return &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + }, nil + }, + delete: func(context.Context, entities.UserID, uuid.UUID) error { + return deleteErr + }, + } + + service := newMessageThreadServiceForTest(repository) + err := service.UpdateAfterDeletedMessage(context.Background(), &events.MessageAPIDeletedPayload{ + MessageID: uuid.New(), + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + }) + + require.Error(t, err) + assert.ErrorIs(t, err, deleteErr) + assert.Contains(t, err.Error(), threadID.String()) +} + func TestShouldCheckUnarchive(t *testing.T) { service := &MessageThreadService{} 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/web/shared/types/api.ts b/web/shared/types/api.ts index 46da6f72..bb3b8b3b 100644 --- a/web/shared/types/api.ts +++ b/web/shared/types/api.ts @@ -487,7 +487,11 @@ export interface RequestsMessageSendScheduleWindow { export interface RequestsMessageThreadUpdate { /** @example true */ is_archived?: boolean; - /** @example 0 */ + /** + * @min 0 + * @max 0 + * @example 0 + */ unread_count?: number; } From b34167c38f4a62b9e4c0383f8c792ac95ee18b6f Mon Sep 17 00:00:00 2001 From: Acho Arnold Ewin Date: Fri, 21 Aug 2026 22:08:32 +0300 Subject: [PATCH 11/12] fix(api): serialize thread metadata updates Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 332cfd7b-8a60-4e28-9ea8-61854aa712c3 --- .../listeners/message_thread_listener_test.go | 1 - .../gorm_message_thread_repository.go | 82 ++-- .../gorm_message_thread_repository_test.go | 372 +++++++++++++++++- .../repositories/message_thread_repository.go | 3 +- api/pkg/services/message_thread_service.go | 32 +- .../services/message_thread_service_test.go | 128 +++++- 6 files changed, 531 insertions(+), 87 deletions(-) diff --git a/api/pkg/listeners/message_thread_listener_test.go b/api/pkg/listeners/message_thread_listener_test.go index 3159058d..2a6eccad 100644 --- a/api/pkg/listeners/message_thread_listener_test.go +++ b/api/pkg/listeners/message_thread_listener_test.go @@ -95,7 +95,6 @@ func TestMessageThreadListenerDeletesNonLastUnreadMessage(t *testing.T) { require.NoError(t, err) assert.Equal(t, deletedMessageID, repository.deletedUpdate.DeletedMessageID) - assert.False(t, repository.deletedUpdate.UpdateLastMessage) require.NotNil(t, repository.deletedUpdate.LastMessageID) assert.Equal(t, previousMessageID, *repository.deletedUpdate.LastMessageID) } diff --git a/api/pkg/repositories/gorm_message_thread_repository.go b/api/pkg/repositories/gorm_message_thread_repository.go index 12417df8..cef1113c 100644 --- a/api/pkg/repositories/gorm_message_thread_repository.go +++ b/api/pkg/repositories/gorm_message_thread_repository.go @@ -46,21 +46,27 @@ func messageThreadActivityUpdates(params MessageThreadActivityUpdate) map[string "last_message_content": params.Content, "status": params.Status, } - if params.Unarchive { - updates["is_archived"] = false - } return updates } -func messageThreadDeletedUpdates(params MessageThreadDeletedUpdate) map[string]any { - if !params.UpdateLastMessage { - return map[string]any{} +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.LastMessageStatus == nil { + return nil, stacktrace.NewErrorf( + "last message status is required when replacing deleted message [%s]", + params.DeletedMessageID, + ) } 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, readAt time.Time) map[string]any { @@ -161,18 +167,29 @@ func markUnreadItemDeleted(tx *gorm.DB, messageID uuid.UUID, threadID uuid.UUID) } func applyMessageThreadActivity(tx *gorm.DB, thread *entities.MessageThread, params MessageThreadActivityUpdate) error { - if err := tx. - Model(thread). - Where("user_id = ?", params.UserID). - Where("id = ?", thread.ID). - Updates(messageThreadActivityUpdates(params)). - Error; err != nil { - return stacktrace.Propagatef( - err, - "cannot update message activity for thread [%s] and user [%s]", - thread.ID, - params.UserID, - ) + updates := make(map[string]any) + isStale := thread.OrderTimestamp.After(params.Timestamp) + isDeliveredMessage := thread.Status == entities.MessageStatusDelivered && thread.HasLastMessage(params.MessageID) + if !isStale && !isDeliveredMessage { + updates = messageThreadActivityUpdates(params) + } + if params.Unarchive { + updates["is_archived"] = false + } + if len(updates) != 0 { + if err := tx. + Model(thread). + Where("user_id = ?", params.UserID). + Where("id = ?", thread.ID). + Updates(updates). + Error; err != nil { + return stacktrace.Propagatef( + err, + "cannot update message activity for thread [%s] and user [%s]", + thread.ID, + params.UserID, + ) + } } if !params.CountAsUnread || !params.EventTimestamp.After(thread.LastReadAt) { return nil @@ -262,14 +279,35 @@ func (repository *gormMessageThreadRepository) UpdateAfterDeletedMessage(ctx con } } - if !params.UpdateLastMessage { + if !thread.HasLastMessage(params.DeletedMessageID) { + return nil + } + if params.LastMessageID == nil { + if err := tx. + Where("user_id = ?", params.UserID). + Where("id = ?", params.MessageThreadID). + Delete(&entities.MessageThread{}). + Error; err != nil { + return stacktrace.Propagatef( + err, + "cannot delete message thread [%s] for user [%s] after deleting final message [%s]", + params.MessageThreadID, + params.UserID, + params.DeletedMessageID, + ) + } return nil } + + updates, err := messageThreadDeletedUpdates(params) + if err != nil { + return err + } if err := tx. Model(thread). Where("user_id = ?", params.UserID). Where("id = ?", params.MessageThreadID). - Updates(messageThreadDeletedUpdates(params)). + Updates(updates). Error; err != nil { return stacktrace.Propagatef( err, diff --git a/api/pkg/repositories/gorm_message_thread_repository_test.go b/api/pkg/repositories/gorm_message_thread_repository_test.go index 6b053b27..0e1f8485 100644 --- a/api/pkg/repositories/gorm_message_thread_repository_test.go +++ b/api/pkg/repositories/gorm_message_thread_repository_test.go @@ -140,14 +140,41 @@ func (*messageThreadRowsConn) Begin() (driver.Tx, error) { func (conn *messageThreadRowsConn) QueryContext(context.Context, string, []driver.NamedValue) (driver.Rows, error) { rows := &messageThreadDriverRows{ - columns: []string{"id", "user_id", "last_read_at", "unread_count"}, + columns: []string{ + "id", + "user_id", + "owner", + "contact", + "is_archived", + "last_read_at", + "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, conn.pool.thread.LastReadAt, int64(conn.pool.thread.UnreadCount), + lastMessageID, + lastMessageContent, + string(conn.pool.thread.Status), + conn.pool.thread.OrderTimestamp, } } return rows, nil @@ -465,6 +492,68 @@ func TestMessageThreadStoreConflictAppliesLosingActivity(t *testing.T) { assert.Contains(t, pool.statements[ledger].args, winnerThreadID) } +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, + LastReadAt: time.Unix(0, 0).UTC(), + 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, + LastReadAt: time.Unix(0, 0).UTC(), + LastMessageID: &losingMessageID, + LastMessageContent: &losingContent, + Status: entities.MessageStatusReceived, + OrderTimestamp: losingTimestamp, + }, + CountAsUnread: true, + EventTimestamp: losingTimestamp, + }) + + require.NoError(t, err) + lock := messageThreadStatementIndex(pool, "FOR UPDATE") + metadata := messageThreadStatementIndexAfter(pool, `"order_timestamp"`, lock) + ledger := messageThreadStatementIndex(pool, `INSERT INTO "message_thread_unread_items"`) + increment := messageThreadStatementIndex(pool, "unread_count +") + require.NotEqual(t, -1, lock) + assert.Equal(t, -1, metadata) + require.NotEqual(t, -1, ledger) + require.NotEqual(t, -1, increment) + assert.Less(t, lock, ledger) + assert.Less(t, ledger, increment) + assert.Contains(t, pool.statements[ledger].args, losingMessageID) + assert.Contains(t, pool.statements[ledger].args, winnerThreadID) +} + func TestMessageThreadActivityUpdatesDoNotOwnUnreadColumns(t *testing.T) { messageID := uuid.New() updates := messageThreadActivityUpdates(MessageThreadActivityUpdate{ @@ -530,6 +619,121 @@ func TestMessageThreadActivityCountableItemInsertsLedgerAndIncrements(t *testing assert.Contains(t, pool.statements[ledger].query, "ON CONFLICT DO NOTHING") } +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, + LastReadAt: time.Unix(0, 0).UTC(), + 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, + EventTimestamp: currentTimestamp, + }) + + require.NoError(t, err) + lock := messageThreadStatementIndex(pool, "FOR UPDATE") + metadata := messageThreadStatementIndexAfter(pool, `"order_timestamp"`, lock) + ledger := messageThreadStatementIndex(pool, `INSERT INTO "message_thread_unread_items"`) + increment := messageThreadStatementIndex(pool, "unread_count +") + require.NotEqual(t, -1, lock) + assert.Equal(t, -1, metadata) + require.NotEqual(t, -1, ledger) + require.NotEqual(t, -1, increment) + assert.Less(t, lock, ledger) + assert.Less(t, ledger, increment) +} + +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) + lock := messageThreadStatementIndex(pool, "FOR UPDATE") + unarchive := messageThreadStatementIndexAfter(pool, `"is_archived"`, lock) + require.NotEqual(t, -1, lock) + require.NotEqual(t, -1, unarchive) + assert.NotContains(t, pool.statements[unarchive].query, `"order_timestamp"`) + assert.NotContains(t, pool.statements[unarchive].query, `"last_message_id"`) + assert.NotContains(t, pool.statements[unarchive].query, `"last_message_content"`) + assert.NotContains(t, pool.statements[unarchive].query, `"status"`) +} + +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) + lock := messageThreadStatementIndex(pool, "FOR UPDATE") + require.NotEqual(t, -1, lock) + assert.Equal(t, -1, messageThreadStatementIndexAfter(pool, `"order_timestamp"`, lock)) +} + func TestMessageThreadActivityDuplicateLedgerItemDoesNotIncrement(t *testing.T) { threadID := uuid.New() pool := &messageThreadTestConnPool{ @@ -618,13 +822,14 @@ func TestMessageThreadActivityMissingThreadReturnsScopedNotFound(t *testing.T) { 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, - UpdateLastMessage: true, + LastMessageStatus: &status, }) + require.NoError(t, err) assert.Equal(t, map[string]any{ "last_message_id": &messageID, "last_message_content": &content, @@ -632,8 +837,16 @@ func TestMessageThreadDeletedUpdatesPreserveStatusType(t *testing.T) { }, updates) } -func TestMessageThreadDeletedUpdatesSkipLastMessageWhenNotRequested(t *testing.T) { - assert.Empty(t, messageThreadDeletedUpdates(MessageThreadDeletedUpdate{})) +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 TestMessageThreadDeletedMessageDecrementsWhenLedgerDeleted(t *testing.T) { @@ -643,21 +856,22 @@ func TestMessageThreadDeletedMessageDecrementsWhenLedgerDeleted(t *testing.T) { previousContent := "previous" pool := &messageThreadTestConnPool{ thread: &entities.MessageThread{ - ID: threadID, - UserID: entities.UserID("user-id"), - UnreadCount: 1, + 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{ MessageThreadID: threadID, UserID: entities.UserID("user-id"), DeletedMessageID: messageID, - UpdateLastMessage: true, LastMessageID: &previousMessageID, LastMessageContent: &previousContent, - LastMessageStatus: entities.MessageStatusDelivered, + LastMessageStatus: &previousStatus, }) require.NoError(t, err) @@ -680,6 +894,142 @@ func TestMessageThreadDeletedMessageDecrementsWhenLedgerDeleted(t *testing.T) { assert.Contains(t, pool.statements[ledgerDelete].query, "counted =") } +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{ + MessageThreadID: threadID, + UserID: entities.UserID("user-id"), + DeletedMessageID: deletedMessageID, + LastMessageID: &previousMessageID, + LastMessageContent: &previousContent, + LastMessageStatus: &previousStatus, + }) + + require.NoError(t, err) + tombstone := messageThreadStatementIndex(pool, `UPDATE "message_thread_unread_items"`) + decrement := messageThreadStatementIndex(pool, "GREATEST(unread_count - 1, 0)") + require.NotEqual(t, -1, tombstone) + require.NotEqual(t, -1, decrement) + assert.Less(t, tombstone, decrement) + 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{ + MessageThreadID: threadID, + UserID: entities.UserID("user-id"), + DeletedMessageID: deletedMessageID, + }) + + require.NoError(t, err) + tombstone := messageThreadStatementIndex(pool, `UPDATE "message_thread_unread_items"`) + decrement := messageThreadStatementIndex(pool, "GREATEST(unread_count - 1, 0)") + require.NotEqual(t, -1, tombstone) + require.NotEqual(t, -1, decrement) + assert.Less(t, tombstone, decrement) + 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{ + MessageThreadID: threadID, + UserID: entities.UserID("user-id"), + DeletedMessageID: deletedMessageID, + }) + + require.NoError(t, err) + require.Equal(t, 1, pool.begins) + require.Equal(t, 1, pool.commits) + require.Zero(t, pool.rollbacks) + lock := messageThreadStatementIndex(pool, "FOR UPDATE") + tombstone := messageThreadStatementIndex(pool, `UPDATE "message_thread_unread_items"`) + decrement := messageThreadStatementIndex(pool, "GREATEST(unread_count - 1, 0)") + threadDelete := messageThreadStatementIndex(pool, `DELETE FROM "message_threads"`) + require.NotEqual(t, -1, lock) + require.NotEqual(t, -1, tombstone) + require.NotEqual(t, -1, decrement) + require.NotEqual(t, -1, threadDelete) + assert.Less(t, lock, tombstone) + assert.Less(t, tombstone, decrement) + assert.Less(t, decrement, threadDelete) +} + +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, + }, + rowsAffected: func(query string) int64 { + if strings.Contains(query, `UPDATE "message_thread_unread_items"`) { + return 0 + } + return 1 + }, + } + repository := newMessageThreadTestRepository(t, pool) + + err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + MessageThreadID: threadID, + UserID: entities.UserID("user-id"), + DeletedMessageID: deletedMessageID, + LastMessageID: &previousMessageID, + LastMessageContent: &previousContent, + }) + + require.Error(t, err) + assert.Contains(t, err.Error(), "status") + require.Equal(t, 1, pool.rollbacks) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `"last_message_id"`)) +} + func TestMessageThreadDeletedMessageWithoutLedgerDoesNotDecrement(t *testing.T) { threadID := uuid.New() pool := &messageThreadTestConnPool{ diff --git a/api/pkg/repositories/message_thread_repository.go b/api/pkg/repositories/message_thread_repository.go index 92ae1b46..c4e7230d 100644 --- a/api/pkg/repositories/message_thread_repository.go +++ b/api/pkg/repositories/message_thread_repository.go @@ -37,10 +37,9 @@ type MessageThreadDeletedUpdate struct { MessageThreadID uuid.UUID UserID entities.UserID DeletedMessageID uuid.UUID - UpdateLastMessage bool LastMessageID *uuid.UUID LastMessageContent *string - LastMessageStatus entities.MessageStatus + LastMessageStatus *entities.MessageStatus } // MessageThreadRepository loads and persists an entities.MessageThread diff --git a/api/pkg/services/message_thread_service.go b/api/pkg/services/message_thread_service.go index e74163cc..c886f795 100644 --- a/api/pkg/services/message_thread_service.go +++ b/api/pkg/services/message_thread_service.go @@ -94,16 +94,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, @@ -168,33 +158,13 @@ func (service *MessageThreadService) UpdateAfterDeletedMessage(ctx context.Conte 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 { - return service.tracer.WrapErrorSpan( - span, - stacktrace.Propagatef( - err, - "cannot delete thread with ID [%s] for user [%s] and owner [%s]", - thread.ID, - thread.UserID, - thread.Owner, - ), - ) - } - 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 - } - - updateLastMessage := thread.LastMessageID != nil && *thread.LastMessageID == payload.MessageID if 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, + 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)) } diff --git a/api/pkg/services/message_thread_service_test.go b/api/pkg/services/message_thread_service_test.go index 27a15243..3eeb36d0 100644 --- a/api/pkg/services/message_thread_service_test.go +++ b/api/pkg/services/message_thread_service_test.go @@ -255,6 +255,54 @@ func TestUpdateThreadIgnoresPhoneLookupErrorsWhenCheckingUnarchive(t *testing.T) 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, + EventTimestamp: time.Date(2026, 8, 21, 10, 0, 1, 0, time.UTC), + }) + + 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 @@ -356,7 +404,7 @@ func TestUpdateStatusPreservesNotFoundCode(t *testing.T) { assert.Equal(t, repositories.ErrCodeNotFound, stacktrace.GetCode(err)) } -func TestUpdateAfterDeletedMessageCleansUnreadLedgerForNonLastMessage(t *testing.T) { +func TestUpdateAfterDeletedMessageDelegatesAllDecisionsToRepository(t *testing.T) { threadID := uuid.New() deletedMessageID := uuid.New() currentLastMessageID := uuid.New() @@ -396,15 +444,15 @@ func TestUpdateAfterDeletedMessageCleansUnreadLedgerForNonLastMessage(t *testing assert.Equal(t, threadID, captured.MessageThreadID) assert.Equal(t, entities.UserID("user-id"), captured.UserID) assert.Equal(t, deletedMessageID, captured.DeletedMessageID) - assert.False(t, captured.UpdateLastMessage) require.NotNil(t, captured.LastMessageID) assert.Equal(t, previousMessageID, *captured.LastMessageID) require.NotNil(t, captured.LastMessageContent) assert.Equal(t, previousContent, *captured.LastMessageContent) - assert.Equal(t, previousStatus, captured.LastMessageStatus) + require.NotNil(t, captured.LastMessageStatus) + assert.Equal(t, previousStatus, *captured.LastMessageStatus) } -func TestUpdateAfterDeletedMessageUpdatesLastMessageWhenDeletedMessageIsCurrent(t *testing.T) { +func TestUpdateAfterDeletedMessageDoesNotUseLoadedLastMessageSnapshot(t *testing.T) { threadID := uuid.New() deletedMessageID := uuid.New() previousMessageID := uuid.New() @@ -438,19 +486,19 @@ func TestUpdateAfterDeletedMessageUpdatesLastMessageWhenDeletedMessageIsCurrent( PreviousMessageStatus: &previousStatus, PreviousMessageContent: &previousContent, }) - require.NoError(t, err) assert.Equal(t, deletedMessageID, captured.DeletedMessageID) - assert.True(t, captured.UpdateLastMessage) + assert.Equal(t, deletedMessageID, captured.DeletedMessageID) require.NotNil(t, captured.LastMessageID) assert.Equal(t, previousMessageID, *captured.LastMessageID) - assert.Equal(t, previousStatus, captured.LastMessageStatus) + require.NotNil(t, captured.LastMessageStatus) + assert.Equal(t, previousStatus, *captured.LastMessageStatus) } -func TestUpdateAfterDeletedMessageDeletesThreadWhenNoPreviousMessageExists(t *testing.T) { +func TestUpdateAfterDeletedMessageDelegatesFinalMessageDeletion(t *testing.T) { threadID := uuid.New() deletedMessageID := uuid.New() - deleted := false + var captured repositories.MessageThreadDeletedUpdate repository := &messageThreadRepositoryStub{ loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { @@ -461,14 +509,12 @@ func TestUpdateAfterDeletedMessageDeletesThreadWhenNoPreviousMessageExists(t *te Contact: "+18005550100", }, nil }, - delete: func(_ context.Context, userID entities.UserID, id uuid.UUID) error { - deleted = true - assert.Equal(t, entities.UserID("user-id"), userID) - assert.Equal(t, threadID, id) + 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 { - t.Fatal("expected whole-thread delete instead of metadata update") + updateAfterDelete: func(_ context.Context, params repositories.MessageThreadDeletedUpdate) error { + captured = params return nil }, } @@ -482,12 +528,15 @@ func TestUpdateAfterDeletedMessageDeletesThreadWhenNoPreviousMessageExists(t *te }) require.NoError(t, err) - assert.True(t, deleted) + assert.Equal(t, threadID, captured.MessageThreadID) + assert.Equal(t, deletedMessageID, captured.DeletedMessageID) + assert.Nil(t, captured.LastMessageID) + assert.Nil(t, captured.LastMessageContent) } -func TestUpdateAfterDeletedMessagePropagatesFinalThreadDeleteError(t *testing.T) { +func TestUpdateAfterDeletedMessagePropagatesRepositoryError(t *testing.T) { threadID := uuid.New() - deleteErr := errors.New("delete failed") + updateErr := errors.New("update failed") repository := &messageThreadRepositoryStub{ loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { return &entities.MessageThread{ @@ -498,7 +547,11 @@ func TestUpdateAfterDeletedMessagePropagatesFinalThreadDeleteError(t *testing.T) }, nil }, delete: func(context.Context, entities.UserID, uuid.UUID) error { - return deleteErr + t.Fatal("service must not delete a thread outside the repository transaction") + return nil + }, + updateAfterDelete: func(context.Context, repositories.MessageThreadDeletedUpdate) error { + return updateErr }, } @@ -511,10 +564,45 @@ func TestUpdateAfterDeletedMessagePropagatesFinalThreadDeleteError(t *testing.T) }) require.Error(t, err) - assert.ErrorIs(t, err, deleteErr) + assert.ErrorIs(t, err, updateErr) assert.Contains(t, err.Error(), threadID.String()) } +func TestUpdateAfterDeletedMessagePassesNilPreviousStatusWithoutPanicking(t *testing.T) { + threadID := uuid.New() + previousMessageID := uuid.New() + previousContent := "previous" + updateErr := errors.New("missing previous status") + repository := &messageThreadRepositoryStub{ + loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { + return &entities.MessageThread{ + ID: threadID, + UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", + }, nil + }, + 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{} From a09ef6d9d63fd62f93d712c3535e221562beb771 Mon Sep 17 00:00:00 2001 From: Acho Arnold Ewin Date: Fri, 21 Aug 2026 22:22:36 +0300 Subject: [PATCH 12/12] fix(api): preserve deleted message markers Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 332cfd7b-8a60-4e28-9ea8-61854aa712c3 --- .../entities/message_thread_deleted_item.go | 8 + api/pkg/entities/message_thread_test.go | 12 + .../listeners/message_thread_listener_test.go | 11 +- .../migrations/message_thread_unread_count.go | 6 +- .../message_thread_unread_count_test.go | 13 ++ .../gorm_message_thread_repository.go | 62 +++++- .../gorm_message_thread_repository_test.go | 206 ++++++++++++++++-- .../repositories/message_thread_repository.go | 3 +- api/pkg/services/message_thread_service.go | 29 ++- .../services/message_thread_service_test.go | 96 ++++---- 10 files changed, 345 insertions(+), 101 deletions(-) create mode 100644 api/pkg/entities/message_thread_deleted_item.go diff --git a/api/pkg/entities/message_thread_deleted_item.go b/api/pkg/entities/message_thread_deleted_item.go new file mode 100644 index 00000000..dbae9bf6 --- /dev/null +++ b/api/pkg/entities/message_thread_deleted_item.go @@ -0,0 +1,8 @@ +package entities + +import "github.com/google/uuid" + +// MessageThreadDeletedItem records a permanently deleted message activity. +type MessageThreadDeletedItem struct { + MessageID uuid.UUID `gorm:"primaryKey;type:uuid"` +} diff --git a/api/pkg/entities/message_thread_test.go b/api/pkg/entities/message_thread_test.go index 6010e11e..652c9904 100644 --- a/api/pkg/entities/message_thread_test.go +++ b/api/pkg/entities/message_thread_test.go @@ -41,3 +41,15 @@ func TestMessageThreadUnreadItemRetainsCountedState(t *testing.T) { assert.Contains(t, counted.Tag.Get("gorm"), "not null") assert.Contains(t, counted.Tag.Get("gorm"), "default:true") } + +func TestMessageThreadDeletedItemIsIndependentFromThreadLifecycle(t *testing.T) { + itemType := reflect.TypeOf(MessageThreadDeletedItem{}) + + require.Equal(t, 1, itemType.NumField()) + messageID, ok := itemType.FieldByName("MessageID") + require.True(t, ok) + assert.Contains(t, messageID.Tag.Get("gorm"), "primaryKey") + assert.Contains(t, messageID.Tag.Get("gorm"), "type:uuid") + _, hasMessageThreadID := itemType.FieldByName("MessageThreadID") + assert.False(t, hasMessageThreadID) +} diff --git a/api/pkg/listeners/message_thread_listener_test.go b/api/pkg/listeners/message_thread_listener_test.go index 2a6eccad..b38a2b7a 100644 --- a/api/pkg/listeners/message_thread_listener_test.go +++ b/api/pkg/listeners/message_thread_listener_test.go @@ -64,15 +64,6 @@ func TestMessageThreadListenerMarksMissedCallUnread(t *testing.T) { func TestMessageThreadListenerDeletesNonLastUnreadMessage(t *testing.T) { repository, routes := newMessageThreadListenerForTest() - currentLastMessageID := uuid.New() - repository.thread = &entities.MessageThread{ - ID: uuid.New(), - UserID: entities.UserID("user-id"), - Owner: "+18005550199", - Contact: "+18005550100", - LastMessageID: ¤tLastMessageID, - } - deletedMessageID := uuid.New() previousMessageID := uuid.New() previousStatus := entities.MessageStatus(entities.MessageStatusDelivered) @@ -95,6 +86,8 @@ func TestMessageThreadListenerDeletesNonLastUnreadMessage(t *testing.T) { 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) } diff --git a/api/pkg/migrations/message_thread_unread_count.go b/api/pkg/migrations/message_thread_unread_count.go index e7d6f3db..1851df89 100644 --- a/api/pkg/migrations/message_thread_unread_count.go +++ b/api/pkg/migrations/message_thread_unread_count.go @@ -20,7 +20,11 @@ func (messageThreadConversationIndex) TableName() string { // MigrateMessageThreadUnreadCount migrates message thread unread count schema. func MigrateMessageThreadUnreadCount(db *gorm.DB) error { - if err := db.AutoMigrate(&entities.MessageThread{}, &entities.MessageThreadUnreadItem{}); err != nil { + if err := db.AutoMigrate( + &entities.MessageThread{}, + &entities.MessageThreadUnreadItem{}, + &entities.MessageThreadDeletedItem{}, + ); err != nil { return stacktrace.Propagate(err, "cannot migrate message thread unread count schema") } diff --git a/api/pkg/migrations/message_thread_unread_count_test.go b/api/pkg/migrations/message_thread_unread_count_test.go index 63b787d4..317d42f9 100644 --- a/api/pkg/migrations/message_thread_unread_count_test.go +++ b/api/pkg/migrations/message_thread_unread_count_test.go @@ -28,6 +28,19 @@ func TestMigrateMessageThreadUnreadCountSkipsLegacyBackfillWhenIsReadColumnMissi assert.Contains(t, strings.Join(recorder.execs, "\n"), `"counted" boolean NOT NULL DEFAULT true`) } +func TestMigrateMessageThreadUnreadCountCreatesIndependentDeletedItemSchema(t *testing.T) { + db, recorder := newMigrationTestDB(t, migrationTestDBOptions{}) + + require.NoError(t, MigrateMessageThreadUnreadCount(db)) + + deletedItems := migrationExecIndex(recorder, `CREATE TABLE "message_thread_deleted_items"`) + require.NotEqual(t, -1, deletedItems) + assert.Contains(t, recorder.execs[deletedItems], `"message_id" uuid`) + assert.Contains(t, recorder.execs[deletedItems], `PRIMARY KEY`) + assert.NotContains(t, recorder.execs[deletedItems], `"message_thread_id"`) + assert.NotContains(t, recorder.execs[deletedItems], `FOREIGN KEY`) +} + func TestMigrateMessageThreadUnreadCountBackfillsBeforeDropAndSkipsOnSecondRun(t *testing.T) { db, recorder := newMigrationTestDB(t, migrationTestDBOptions{hasLegacyIsRead: true}) diff --git a/api/pkg/repositories/gorm_message_thread_repository.go b/api/pkg/repositories/gorm_message_thread_repository.go index cef1113c..c84e3384 100644 --- a/api/pkg/repositories/gorm_message_thread_repository.go +++ b/api/pkg/repositories/gorm_message_thread_repository.go @@ -148,6 +148,28 @@ func insertUnreadItem(tx *gorm.DB, item entities.MessageThreadUnreadItem) (bool, return result.RowsAffected == 1, nil } +func insertDeletedItem(tx *gorm.DB, messageID uuid.UUID) error { + if err := tx. + Clauses(clause.OnConflict{DoNothing: true}). + Create(&entities.MessageThreadDeletedItem{MessageID: messageID}). + Error; err != nil { + return stacktrace.Propagatef(err, "cannot insert deleted message marker for message [%s]", messageID) + } + return nil +} + +func isDeletedItem(tx *gorm.DB, messageID uuid.UUID) (bool, error) { + item := new(entities.MessageThreadDeletedItem) + result := tx. + Where("message_id = ?", messageID). + Limit(1). + Find(item) + if result.Error != nil { + return false, stacktrace.Propagatef(result.Error, "cannot load deleted message marker for message [%s]", messageID) + } + return result.RowsAffected != 0, nil +} + func markUnreadItemDeleted(tx *gorm.DB, messageID uuid.UUID, threadID uuid.UUID) (bool, error) { result := tx. Model(&entities.MessageThreadUnreadItem{}). @@ -251,12 +273,19 @@ func (repository *gormMessageThreadRepository) UpdateAfterDeletedMessage(ctx con err := crdbgorm.ExecuteTx(ctx, repository.db, nil, func(tx *gorm.DB) error { tx = tx.WithContext(ctx) - thread, err := lockMessageThread(tx, params.UserID, params.MessageThreadID) + if err := insertDeletedItem(tx, params.DeletedMessageID); err != nil { + return err + } + + thread, err := lockMessageThreadByConversation(tx, params.UserID, params.Owner, params.Contact) + if stacktrace.GetCode(err) == ErrCodeNotFound { + return nil + } if err != nil { return err } - deleted, err := markUnreadItemDeleted(tx, params.DeletedMessageID, params.MessageThreadID) + deleted, err := markUnreadItemDeleted(tx, params.DeletedMessageID, thread.ID) if err != nil { return err } @@ -264,13 +293,13 @@ func (repository *gormMessageThreadRepository) UpdateAfterDeletedMessage(ctx con if err := tx. Model(thread). Where("user_id = ?", params.UserID). - Where("id = ?", params.MessageThreadID). + Where("id = ?", thread.ID). UpdateColumn("unread_count", gorm.Expr("GREATEST(unread_count - 1, 0)")). Error; err != nil { return stacktrace.Propagatef( err, "cannot decrement unread count for thread [%s] and user [%s]", - params.MessageThreadID, + thread.ID, params.UserID, ) } @@ -285,13 +314,13 @@ func (repository *gormMessageThreadRepository) UpdateAfterDeletedMessage(ctx con if params.LastMessageID == nil { if err := tx. Where("user_id = ?", params.UserID). - Where("id = ?", params.MessageThreadID). + Where("id = ?", thread.ID). Delete(&entities.MessageThread{}). Error; err != nil { return stacktrace.Propagatef( err, "cannot delete message thread [%s] for user [%s] after deleting final message [%s]", - params.MessageThreadID, + thread.ID, params.UserID, params.DeletedMessageID, ) @@ -306,13 +335,13 @@ func (repository *gormMessageThreadRepository) UpdateAfterDeletedMessage(ctx con if err := tx. Model(thread). Where("user_id = ?", params.UserID). - Where("id = ?", params.MessageThreadID). + Where("id = ?", thread.ID). Updates(updates). Error; err != nil { return stacktrace.Propagatef( err, "cannot update deleted-message metadata for thread [%s] and user [%s]", - params.MessageThreadID, + thread.ID, params.UserID, ) } @@ -323,9 +352,10 @@ func (repository *gormMessageThreadRepository) UpdateAfterDeletedMessage(ctx con span, stacktrace.Propagatef( err, - "cannot apply deleted message [%s] to thread [%s] for user [%s]", + "cannot apply deleted message [%s] to conversation [%s/%s] for user [%s]", params.DeletedMessageID, - params.MessageThreadID, + params.Owner, + params.Contact, params.UserID, ), ) @@ -341,6 +371,13 @@ func (repository *gormMessageThreadRepository) Store(ctx context.Context, params err := crdbgorm.ExecuteTx(ctx, repository.db, nil, func(tx *gorm.DB) error { tx = tx.WithContext(ctx) + if params.Thread.LastMessageID != nil { + deleted, err := isDeletedItem(tx, *params.Thread.LastMessageID) + if err != nil || deleted { + return err + } + } + candidate := *params.Thread result := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&candidate) if result.Error != nil { @@ -417,6 +454,11 @@ func (repository *gormMessageThreadRepository) UpdateActivity(ctx context.Contex err := crdbgorm.ExecuteTx(ctx, repository.db, nil, func(tx *gorm.DB) error { tx = tx.WithContext(ctx) + deleted, err := isDeletedItem(tx, params.MessageID) + if err != nil || deleted { + return err + } + thread, err := lockMessageThread(tx, params.UserID, params.MessageThreadID) if err != nil { return err diff --git a/api/pkg/repositories/gorm_message_thread_repository_test.go b/api/pkg/repositories/gorm_message_thread_repository_test.go index 0e1f8485..c3fe24b1 100644 --- a/api/pkg/repositories/gorm_message_thread_repository_test.go +++ b/api/pkg/repositories/gorm_message_thread_repository_test.go @@ -27,14 +27,15 @@ type messageThreadTestStatement struct { } type messageThreadTestConnPool struct { - statements []messageThreadTestStatement - thread *entities.MessageThread - rowsAffected func(query string) int64 - execError func(query string) error - queryDB *sql.DB - begins int - commits int - rollbacks int + statements []messageThreadTestStatement + thread *entities.MessageThread + deletedMessageID *uuid.UUID + 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) { @@ -138,7 +139,17 @@ func (*messageThreadRowsConn) Begin() (driver.Tx, error) { return nil, errors.New("unexpected Begin") } -func (conn *messageThreadRowsConn) QueryContext(context.Context, string, []driver.NamedValue) (driver.Rows, error) { +func (conn *messageThreadRowsConn) QueryContext(_ context.Context, query string, _ []driver.NamedValue) (driver.Rows, error) { + if strings.Contains(query, `"message_thread_deleted_items"`) { + rows := &messageThreadDriverRows{ + columns: []string{"message_id"}, + } + if conn.pool.deletedMessageID != nil { + rows.values = []driver.Value{conn.pool.deletedMessageID.String()} + } + return rows, nil + } + rows := &messageThreadDriverRows{ columns: []string{ "id", @@ -393,8 +404,9 @@ func TestMessageThreadMutationsRetrySerializationFailures(t *testing.T) { repository := newMessageThreadTestRepository(t, pool) err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ - MessageThreadID: threadID, UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", DeletedMessageID: uuid.New(), }) @@ -866,8 +878,9 @@ func TestMessageThreadDeletedMessageDecrementsWhenLedgerDeleted(t *testing.T) { previousStatus := entities.MessageStatus(entities.MessageStatusDelivered) err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ - MessageThreadID: threadID, UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", DeletedMessageID: messageID, LastMessageID: &previousMessageID, LastMessageContent: &previousContent, @@ -912,8 +925,9 @@ func TestMessageThreadDeletedStaleReplacementPreservesNewerLastActivity(t *testi previousStatus := entities.MessageStatus(entities.MessageStatusDelivered) err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ - MessageThreadID: threadID, UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", DeletedMessageID: deletedMessageID, LastMessageID: &previousMessageID, LastMessageContent: &previousContent, @@ -945,8 +959,9 @@ func TestMessageThreadDeletedStaleFinalMessagePreservesNewerLastActivity(t *test repository := newMessageThreadTestRepository(t, pool) err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ - MessageThreadID: threadID, UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", DeletedMessageID: deletedMessageID, }) @@ -974,8 +989,9 @@ func TestMessageThreadDeletedCurrentFinalMessageDeletesThread(t *testing.T) { repository := newMessageThreadTestRepository(t, pool) err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ - MessageThreadID: threadID, UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", DeletedMessageID: deletedMessageID, }) @@ -1017,8 +1033,9 @@ func TestMessageThreadDeletedCurrentMessageRequiresPreviousStatus(t *testing.T) repository := newMessageThreadTestRepository(t, pool) err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ - MessageThreadID: threadID, UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", DeletedMessageID: deletedMessageID, LastMessageID: &previousMessageID, LastMessageContent: &previousContent, @@ -1048,8 +1065,9 @@ func TestMessageThreadDeletedMessageWithoutLedgerDoesNotDecrement(t *testing.T) repository := newMessageThreadTestRepository(t, pool) err := repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ - MessageThreadID: threadID, UserID: entities.UserID("user-id"), + Owner: "+18005550199", + Contact: "+18005550100", DeletedMessageID: uuid.New(), }) @@ -1093,8 +1111,9 @@ func TestMessageThreadDeletedItemReplayDoesNotIncrement(t *testing.T) { repository := newMessageThreadTestRepository(t, pool) require.NoError(t, repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ - MessageThreadID: threadID, UserID: userID, + Owner: "+18005550199", + Contact: "+18005550100", DeletedMessageID: messageID, })) require.NoError(t, repository.UpdateActivity(context.Background(), MessageThreadActivityUpdate{ @@ -1118,6 +1137,159 @@ func TestMessageThreadDeletedItemReplayDoesNotIncrement(t *testing.T) { assert.Equal(t, 1, messageThreadStatementCount(pool, "GREATEST(unread_count - 1, 0)")) } +func TestMessageThreadDeletionBeforeThreadBlocksLaterStore(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) + marker := messageThreadStatementIndex(pool, `INSERT INTO "message_thread_deleted_items"`) + lock := messageThreadStatementIndex(pool, "FOR UPDATE") + require.NotEqual(t, -1, marker) + require.NotEqual(t, -1, lock) + assert.Less(t, marker, lock) + + pool.deletedMessageID = &messageID + 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.NotEqual(t, -1, messageThreadStatementIndex(pool, `FROM "message_thread_deleted_items"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `INSERT INTO "message_threads"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `INSERT INTO "message_thread_unread_items"`)) +} + +func TestMessageThreadDeletedInboundReplayPreservesExistingThread(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, + LastReadAt: time.Unix(0, 0).UTC(), + }, + rowsAffected: func(query string) int64 { + if strings.Contains(query, `UPDATE "message_thread_unread_items"`) { + return 0 + } + return 1 + }, + } + repository := newMessageThreadTestRepository(t, pool) + + require.NoError(t, repository.UpdateAfterDeletedMessage(context.Background(), MessageThreadDeletedUpdate{ + UserID: userID, + Owner: "+18005550199", + Contact: "+18005550100", + DeletedMessageID: messageID, + })) + + pool.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, + EventTimestamp: currentTimestamp.Add(time.Second), + Unarchive: true, + })) + + assert.NotEqual(t, -1, messageThreadStatementIndex(pool, `FROM "message_thread_deleted_items"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "FOR UPDATE")) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `"order_timestamp"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `"is_archived"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `INSERT INTO "message_thread_unread_items"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, "unread_count +")) +} + +func TestMessageThreadFinalDeletionBlocksStoreReplay(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, + }, + rowsAffected: func(query string) int64 { + if strings.Contains(query, `UPDATE "message_thread_unread_items"`) { + return 0 + } + return 1 + }, + } + 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.deletedMessageID = &messageID + 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.NotEqual(t, -1, messageThreadStatementIndex(pool, `FROM "message_thread_deleted_items"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `INSERT INTO "message_threads"`)) + assert.Equal(t, -1, messageThreadStatementIndex(pool, `INSERT INTO "message_thread_unread_items"`)) +} + func TestMessageThreadStatusUpdatesResetUnreadCount(t *testing.T) { zero := uint(0) readAt := time.Date(2026, 8, 21, 10, 0, 0, 0, time.UTC) diff --git a/api/pkg/repositories/message_thread_repository.go b/api/pkg/repositories/message_thread_repository.go index c4e7230d..3874f57f 100644 --- a/api/pkg/repositories/message_thread_repository.go +++ b/api/pkg/repositories/message_thread_repository.go @@ -34,8 +34,9 @@ type MessageThreadStatusUpdate struct { } type MessageThreadDeletedUpdate struct { - MessageThreadID uuid.UUID UserID entities.UserID + Owner string + Contact string DeletedMessageID uuid.UUID LastMessageID *uuid.UUID LastMessageContent *string diff --git a/api/pkg/services/message_thread_service.go b/api/pkg/services/message_thread_service.go index c886f795..317d3c6d 100644 --- a/api/pkg/services/message_thread_service.go +++ b/api/pkg/services/message_thread_service.go @@ -153,23 +153,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 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, }); 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 } diff --git a/api/pkg/services/message_thread_service_test.go b/api/pkg/services/message_thread_service_test.go index 3eeb36d0..49eb1fb0 100644 --- a/api/pkg/services/message_thread_service_test.go +++ b/api/pkg/services/message_thread_service_test.go @@ -405,24 +405,13 @@ func TestUpdateStatusPreservesNotFoundCode(t *testing.T) { } func TestUpdateAfterDeletedMessageDelegatesAllDecisionsToRepository(t *testing.T) { - threadID := uuid.New() deletedMessageID := uuid.New() - currentLastMessageID := uuid.New() previousMessageID := uuid.New() previousStatus := entities.MessageStatus(entities.MessageStatusDelivered) previousContent := "previous" var captured repositories.MessageThreadDeletedUpdate repository := &messageThreadRepositoryStub{ - loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { - return &entities.MessageThread{ - ID: threadID, - UserID: entities.UserID("user-id"), - Owner: "+18005550199", - Contact: "+18005550100", - LastMessageID: ¤tLastMessageID, - }, nil - }, updateAfterDelete: func(_ context.Context, params repositories.MessageThreadDeletedUpdate) error { captured = params return nil @@ -441,8 +430,9 @@ func TestUpdateAfterDeletedMessageDelegatesAllDecisionsToRepository(t *testing.T }) require.NoError(t, err) - assert.Equal(t, threadID, captured.MessageThreadID) 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) @@ -452,22 +442,16 @@ func TestUpdateAfterDeletedMessageDelegatesAllDecisionsToRepository(t *testing.T assert.Equal(t, previousStatus, *captured.LastMessageStatus) } -func TestUpdateAfterDeletedMessageDoesNotUseLoadedLastMessageSnapshot(t *testing.T) { - threadID := uuid.New() +func TestUpdateAfterDeletedMessageDelegatesWithoutLoadingThread(t *testing.T) { deletedMessageID := uuid.New() - previousMessageID := uuid.New() - previousStatus := entities.MessageStatus(entities.MessageStatusDelivered) - previousContent := "previous" + 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: threadID, - UserID: entities.UserID("user-id"), - Owner: "+18005550199", - Contact: "+18005550100", - LastMessageID: &deletedMessageID, + ID: uuid.New(), + UserID: entities.UserID("user-id"), }, nil }, updateAfterDelete: func(_ context.Context, params repositories.MessageThreadDeletedUpdate) error { @@ -476,6 +460,35 @@ func TestUpdateAfterDeletedMessageDoesNotUseLoadedLastMessageSnapshot(t *testing }, } + 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, @@ -496,19 +509,10 @@ func TestUpdateAfterDeletedMessageDoesNotUseLoadedLastMessageSnapshot(t *testing } func TestUpdateAfterDeletedMessageDelegatesFinalMessageDeletion(t *testing.T) { - threadID := uuid.New() deletedMessageID := uuid.New() var captured repositories.MessageThreadDeletedUpdate repository := &messageThreadRepositoryStub{ - loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { - return &entities.MessageThread{ - ID: threadID, - UserID: entities.UserID("user-id"), - Owner: "+18005550199", - Contact: "+18005550100", - }, nil - }, delete: func(context.Context, entities.UserID, uuid.UUID) error { t.Fatal("service must not delete a thread outside the repository transaction") return nil @@ -528,24 +532,17 @@ func TestUpdateAfterDeletedMessageDelegatesFinalMessageDeletion(t *testing.T) { }) require.NoError(t, err) - assert.Equal(t, threadID, captured.MessageThreadID) + 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) { - threadID := uuid.New() + deletedMessageID := uuid.New() updateErr := errors.New("update failed") repository := &messageThreadRepositoryStub{ - loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { - return &entities.MessageThread{ - ID: threadID, - UserID: entities.UserID("user-id"), - Owner: "+18005550199", - Contact: "+18005550100", - }, nil - }, delete: func(context.Context, entities.UserID, uuid.UUID) error { t.Fatal("service must not delete a thread outside the repository transaction") return nil @@ -557,7 +554,7 @@ func TestUpdateAfterDeletedMessagePropagatesRepositoryError(t *testing.T) { service := newMessageThreadServiceForTest(repository) err := service.UpdateAfterDeletedMessage(context.Background(), &events.MessageAPIDeletedPayload{ - MessageID: uuid.New(), + MessageID: deletedMessageID, UserID: entities.UserID("user-id"), Owner: "+18005550199", Contact: "+18005550100", @@ -565,23 +562,16 @@ func TestUpdateAfterDeletedMessagePropagatesRepositoryError(t *testing.T) { require.Error(t, err) assert.ErrorIs(t, err, updateErr) - assert.Contains(t, err.Error(), threadID.String()) + assert.Contains(t, err.Error(), deletedMessageID.String()) + assert.Contains(t, err.Error(), "+18005550199") + assert.Contains(t, err.Error(), "+18005550100") } func TestUpdateAfterDeletedMessagePassesNilPreviousStatusWithoutPanicking(t *testing.T) { - threadID := uuid.New() previousMessageID := uuid.New() previousContent := "previous" updateErr := errors.New("missing previous status") repository := &messageThreadRepositoryStub{ - loadByOwnerContact: func(context.Context, entities.UserID, string, string) (*entities.MessageThread, error) { - return &entities.MessageThread{ - ID: threadID, - UserID: entities.UserID("user-id"), - Owner: "+18005550199", - Contact: "+18005550100", - }, nil - }, updateAfterDelete: func(context.Context, repositories.MessageThreadDeletedUpdate) error { return updateErr },