diff --git a/internal/api/me_notifications.go b/internal/api/me_notifications.go index 885e816d..8c7ae485 100644 --- a/internal/api/me_notifications.go +++ b/internal/api/me_notifications.go @@ -3,6 +3,7 @@ package api import ( "encoding/json" "errors" + "io" "net/http" "strconv" "strings" @@ -170,7 +171,20 @@ func (h *handlers) handleMarkAllMyNotificationsRead(w http.ResponseWriter, r *ht if !ok { return } - if _, err := dbq.New(h.pool).MarkAllNotificationsRead(r.Context(), user.ID); err != nil { + // Optional {"up_to": RFC3339}: only what existed then. A client replaying + // an offline "mark all read" sends the moment the user asked. + var body struct { + UpTo *time.Time `json:"up_to"` + } + if err := json.NewDecoder(r.Body).Decode(&body); err != nil && !errors.Is(err, io.EOF) { + writeErr(w, apierror.BadRequest("bad_body", "invalid JSON body")) + return + } + params := dbq.MarkAllNotificationsReadParams{UserID: user.ID} + if body.UpTo != nil { + params.UpTo = pgtype.Timestamptz{Time: *body.UpTo, Valid: true} + } + if _, err := dbq.New(h.pool).MarkAllNotificationsRead(r.Context(), params); err != nil { writeErrWithLog(w, h.logger, "notifications: mark all read", apierror.Internal(err)) return } diff --git a/internal/api/me_notifications_test.go b/internal/api/me_notifications_test.go index 3bf34110..1839b3ab 100644 --- a/internal/api/me_notifications_test.go +++ b/internal/api/me_notifications_test.go @@ -7,6 +7,7 @@ import ( "net/http" "net/http/httptest" "testing" + "time" "github.com/go-chi/chi/v5" "github.com/jackc/pgx/v5/pgxpool" @@ -127,6 +128,33 @@ func TestMyNotifications_MarkReadIsOwnerScopedAndCountsDown(t *testing.T) { } } +// A "mark all read" replayed from an offline queue carries the moment the +// user asked; what arrived after it stays unread. +func TestMyNotifications_ReadAllUpToLeavesLaterOnesUnread(t *testing.T) { + h, pool := testHandlers(t) + alice := seedUser(t, pool, "notif-upto", "pw", false) + r := notificationsRouter(h) + + notifyN(t, pool, alice, 1) + var page notificationsPageResp + callAs(t, r, alice, http.MethodGet, "/api/me/notifications", "", &page) + seen := page.Items[0].CreatedAt + notifyN(t, pool, alice, 1) + + body := `{"up_to":"` + seen.Format(time.RFC3339Nano) + `"}` + if code := callAs(t, r, alice, http.MethodPost, "/api/me/notifications/read-all", body, nil); code != http.StatusNoContent { + t.Fatalf("read-all up_to = %d", code) + } + var count unreadCountResp + callAs(t, r, alice, http.MethodGet, "/api/me/notifications/unread-count", "", &count) + if count.UnreadCount != 1 { + t.Errorf("unread = %d, want 1 (the one that arrived later)", count.UnreadCount) + } + if code := callAs(t, r, alice, http.MethodPost, "/api/me/notifications/read-all", `{"up_to":"yesterday"}`, nil); code != http.StatusBadRequest { + t.Errorf("malformed up_to = %d, want 400", code) + } +} + func strPtr(s string) *string { return &s } func setSMTP(t *testing.T, pool *pgxpool.Pool, enabled bool) { diff --git a/internal/db/dbq/notifications.sql.go b/internal/db/dbq/notifications.sql.go index 82e43f42..2c12a44a 100644 --- a/internal/db/dbq/notifications.sql.go +++ b/internal/db/dbq/notifications.sql.go @@ -214,11 +214,22 @@ func (q *Queries) ListNotifications(ctx context.Context, arg ListNotificationsPa const markAllNotificationsRead = `-- name: MarkAllNotificationsRead :execrows UPDATE user_notifications SET read_at = now() - WHERE user_id = $1 AND read_at IS NULL + WHERE user_id = $1 + AND read_at IS NULL + AND ($2::timestamptz IS NULL OR created_at <= $2) ` -func (q *Queries) MarkAllNotificationsRead(ctx context.Context, userID pgtype.UUID) (int64, error) { - result, err := q.db.Exec(ctx, markAllNotificationsRead, userID) +type MarkAllNotificationsReadParams struct { + UserID pgtype.UUID + UpTo pgtype.Timestamptz +} + +// up_to, when set, limits it to what existed when the user asked: a "mark +// all read" queued offline and replayed later must not mark notices that +// arrived in between, which the user never saw. A coalesced row updated +// since then carries a newer created_at, so it stays unread too. +func (q *Queries) MarkAllNotificationsRead(ctx context.Context, arg MarkAllNotificationsReadParams) (int64, error) { + result, err := q.db.Exec(ctx, markAllNotificationsRead, arg.UserID, arg.UpTo) if err != nil { return 0, err } diff --git a/internal/db/queries/notifications.sql b/internal/db/queries/notifications.sql index 8abc10fe..2779f728 100644 --- a/internal/db/queries/notifications.sql +++ b/internal/db/queries/notifications.sql @@ -57,9 +57,15 @@ UPDATE user_notifications WHERE id = sqlc.arg(id) AND user_id = sqlc.arg(user_id); -- name: MarkAllNotificationsRead :execrows +-- up_to, when set, limits it to what existed when the user asked: a "mark +-- all read" queued offline and replayed later must not mark notices that +-- arrived in between, which the user never saw. A coalesced row updated +-- since then carries a newer created_at, so it stays unread too. UPDATE user_notifications SET read_at = now() - WHERE user_id = $1 AND read_at IS NULL; + WHERE user_id = sqlc.arg(user_id) + AND read_at IS NULL + AND (sqlc.narg(up_to)::timestamptz IS NULL OR created_at <= sqlc.narg(up_to)); -- name: TrimNotifications :execrows -- Retention: read rows go after read_cutoff, and anything at all after diff --git a/internal/library/notify_test.go b/internal/library/notify_test.go index eb7b4ba9..0ca635e2 100644 --- a/internal/library/notify_test.go +++ b/internal/library/notify_test.go @@ -151,7 +151,7 @@ func TestDuplicateSweep_NotifiesOnlyWhenSomethingNewIsProposed(t *testing.T) { } markAllRead := func() { t.Helper() - if _, err := q.MarkAllNotificationsRead(ctx, admin.ID); err != nil { + if _, err := q.MarkAllNotificationsRead(ctx, dbq.MarkAllNotificationsReadParams{UserID: admin.ID}); err != nil { t.Fatalf("mark read: %v", err) } }