feat(notifications): tighten chatwoot scoping

This commit is contained in:
2026-06-06 20:57:14 +08:00
parent 4a7df91556
commit 3aa21996f6
5 changed files with 251 additions and 43 deletions
+48 -14
View File
@@ -2,7 +2,9 @@ package v1
import (
"encoding/json"
"fmt"
"net/http"
"strconv"
"strings"
"time"
@@ -81,7 +83,9 @@ func (h *NotificationHandler) Get(c *gin.Context) {
return
}
notification, svcErr := h.notificationService.GetNotification(c.Request.Context(), notificationID)
accountID := getAccountID(c)
userID := getUserID(c)
notification, svcErr := h.notificationService.GetNotificationByAccount(c.Request.Context(), notificationID, userID, accountID)
if svcErr != nil {
handleServiceError(c, svcErr)
return
@@ -99,15 +103,10 @@ func (h *NotificationHandler) Update(c *gin.Context) {
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid notification id")
return
}
accountID := getAccountID(c)
userID := getUserID(c)
// MarkRead only returns error; need to fetch updated notification for response
if svcErr := h.notificationService.MarkRead(c.Request.Context(), notificationID); svcErr != nil {
handleServiceError(c, svcErr)
return
}
// Return the updated notification (Chatwoot renders json: @notification)
notification, svcErr := h.notificationService.GetNotification(c.Request.Context(), notificationID)
notification, svcErr := h.notificationService.MarkReadByAccount(c.Request.Context(), notificationID, userID, accountID)
if svcErr != nil {
handleServiceError(c, svcErr)
return
@@ -147,9 +146,10 @@ func (h *NotificationHandler) MarkAllRead(c *gin.Context) {
// GET /api/v1/accounts/:account_id/notifications/unread_count
// Reference: Chatwoot unread_count — render json: @unread_count
func (h *NotificationHandler) UnreadCount(c *gin.Context) {
accountID := getAccountID(c)
userID := getUserID(c)
count, err := h.notificationService.GetUnreadCount(c.Request.Context(), userID)
count, err := h.notificationService.GetUnreadCountByAccount(c.Request.Context(), userID, accountID)
if err != nil {
applogger.L().Errorf("UnreadCount notifications: %v", err)
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "Failed to count unread notifications")
@@ -173,14 +173,23 @@ func (h *NotificationHandler) Snooze(c *gin.Context) {
userID := getUserID(c)
var req struct {
SnoozedUntil string `json:"snoozed_until"`
SnoozedUntil interface{} `json:"snoozed_until"`
}
if err := c.ShouldBindJSON(&req); err != nil {
if err := bindOptionalNotificationJSON(c, &req); err != nil {
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid request body")
return
}
if req.SnoozedUntil == nil || fmt.Sprint(req.SnoozedUntil) == "" {
notification, svcErr := h.notificationService.GetNotificationByAccount(c.Request.Context(), notificationID, userID, accountID)
if svcErr != nil {
handleServiceError(c, svcErr)
return
}
c.JSON(http.StatusOK, serializeNotification(notification))
return
}
snoozedUntil, parseErr := time.Parse(time.RFC3339, req.SnoozedUntil)
snoozedUntil, parseErr := parseNotificationUnixTime(req.SnoozedUntil)
if parseErr != nil {
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid snoozed_until format")
return
@@ -227,7 +236,9 @@ func (h *NotificationHandler) Destroy(c *gin.Context) {
return
}
if svcErr := h.notificationService.DeleteNotification(c.Request.Context(), notificationID); svcErr != nil {
accountID := getAccountID(c)
userID := getUserID(c)
if svcErr := h.notificationService.DeleteNotificationByAccount(c.Request.Context(), notificationID, userID, accountID); svcErr != nil {
handleServiceError(c, svcErr)
return
}
@@ -341,3 +352,26 @@ func bindOptionalNotificationJSON(c *gin.Context, target interface{}) error {
}
return c.ShouldBindJSON(target)
}
func parseNotificationUnixTime(value interface{}) (time.Time, error) {
switch v := value.(type) {
case float64:
return time.Unix(int64(v), 0).UTC(), nil
case int64:
return time.Unix(v, 0).UTC(), nil
case int:
return time.Unix(int64(v), 0).UTC(), nil
case string:
seconds, err := strconv.ParseInt(v, 10, 64)
if err != nil {
return time.Time{}, err
}
return time.Unix(seconds, 0).UTC(), nil
default:
seconds, err := strconv.ParseInt(fmt.Sprint(v), 10, 64)
if err != nil {
return time.Time{}, err
}
return time.Unix(seconds, 0).UTC(), nil
}
}
@@ -54,14 +54,16 @@ func setupNotificationRouter(handler *NotificationHandler) *gin.Engine {
c.Next()
})
router.GET("/api/v1/accounts/:account_id/notifications", handler.List)
router.GET("/api/v1/accounts/:account_id/notifications/:notification_id", handler.Get)
router.POST("/api/v1/accounts/:account_id/notifications/read_all", handler.MarkAllRead)
router.GET("/api/v1/accounts/:account_id/notifications/unread_count", handler.UnreadCount)
router.POST("/api/v1/accounts/:account_id/notifications/destroy_all", handler.DestroyAll)
router.DELETE("/api/v1/accounts/:account_id/notifications/destroy_all", handler.DestroyAll)
router.GET("/api/v1/accounts/:account_id/notifications/:notification_id", handler.Get)
router.PUT("/api/v1/accounts/:account_id/notifications/:notification_id", handler.Update)
router.DELETE("/api/v1/accounts/:account_id/notifications/:notification_id", handler.Destroy)
// G8 extension routes
router.POST("/api/v1/accounts/:account_id/notifications/:notification_id/snooze", handler.Snooze)
router.POST("/api/v1/accounts/:account_id/notifications/:notification_id/unread", handler.Unread)
router.POST("/api/v1/accounts/:account_id/notifications/destroy_all", handler.DestroyAll)
router.DELETE("/api/v1/accounts/:account_id/notifications/destroy_all", handler.DestroyAll)
return router
}
@@ -174,6 +176,87 @@ func TestNotificationGetDifferentID(t *testing.T) {
sqlDB.Close()
}
func TestNotificationMutationsAreScopedToCurrentUserAndAccount(t *testing.T) {
db := setupNotificationDB(t)
handler := setupNotificationHandler(t, db)
router := setupNotificationRouter(handler)
user := &model.User{Name: "Scoped User", Email: "scoped@example.com", Password: "pass", AccountID: 1}
otherUser := &model.User{Name: "Other User", Email: "scoped-other@example.com", Password: "pass", AccountID: 2}
require.NoError(t, db.Create(user).Error)
require.NoError(t, db.Create(otherUser).Error)
accountID := uint(1)
otherAccountID := uint(2)
own := &model.Notification{UserID: user.ID, AccountID: &accountID, NotificationType: "message_created", PrimaryActorType: "Conversation", PrimaryActorID: 1}
otherAccount := &model.Notification{UserID: user.ID, AccountID: &otherAccountID, NotificationType: "message_created", PrimaryActorType: "Conversation", PrimaryActorID: 2}
otherOwner := &model.Notification{UserID: otherUser.ID, AccountID: &accountID, NotificationType: "message_created", PrimaryActorType: "Conversation", PrimaryActorID: 3}
require.NoError(t, db.Create(own).Error)
require.NoError(t, db.Create(otherAccount).Error)
require.NoError(t, db.Create(otherOwner).Error)
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", fmt.Sprintf("/api/v1/accounts/1/notifications/%d", otherAccount.ID), nil)
req.Header.Set("X-User-ID", strconv.FormatUint(uint64(user.ID), 10))
router.ServeHTTP(w, req)
require.Equal(t, http.StatusNotFound, w.Code)
w = httptest.NewRecorder()
req, _ = http.NewRequest("PUT", fmt.Sprintf("/api/v1/accounts/1/notifications/%d", otherOwner.ID), nil)
req.Header.Set("X-User-ID", strconv.FormatUint(uint64(user.ID), 10))
router.ServeHTTP(w, req)
require.Equal(t, http.StatusNotFound, w.Code)
w = httptest.NewRecorder()
req, _ = http.NewRequest("PUT", fmt.Sprintf("/api/v1/accounts/1/notifications/%d", own.ID), nil)
req.Header.Set("X-User-ID", strconv.FormatUint(uint64(user.ID), 10))
router.ServeHTTP(w, req)
require.Equal(t, http.StatusOK, w.Code)
var reloaded model.Notification
require.NoError(t, db.First(&reloaded, own.ID).Error)
assert.NotNil(t, reloaded.ReadAt)
w = httptest.NewRecorder()
req, _ = http.NewRequest("DELETE", fmt.Sprintf("/api/v1/accounts/1/notifications/%d", otherAccount.ID), nil)
req.Header.Set("X-User-ID", strconv.FormatUint(uint64(user.ID), 10))
router.ServeHTTP(w, req)
require.Equal(t, http.StatusNotFound, w.Code)
var otherAccountReloaded model.Notification
require.NoError(t, db.First(&otherAccountReloaded, otherAccount.ID).Error)
w = httptest.NewRecorder()
req, _ = http.NewRequest("DELETE", fmt.Sprintf("/api/v1/accounts/1/notifications/%d", own.ID), nil)
req.Header.Set("X-User-ID", strconv.FormatUint(uint64(user.ID), 10))
router.ServeHTTP(w, req)
require.Equal(t, http.StatusOK, w.Code)
require.Error(t, db.First(&reloaded, own.ID).Error)
sqlDB, _ := db.DB()
sqlDB.Close()
}
func TestNotificationUnreadCountIsAccountScoped(t *testing.T) {
db := setupNotificationDB(t)
handler := setupNotificationHandler(t, db)
router := setupNotificationRouter(handler)
user := &model.User{Name: "Unread Scoped User", Email: "unread-scoped@example.com", Password: "pass", AccountID: 1}
require.NoError(t, db.Create(user).Error)
accountID := uint(1)
otherAccountID := uint(2)
require.NoError(t, db.Create(&model.Notification{UserID: user.ID, AccountID: &accountID, NotificationType: "message_created"}).Error)
require.NoError(t, db.Create(&model.Notification{UserID: user.ID, AccountID: &otherAccountID, NotificationType: "message_created"}).Error)
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/v1/accounts/1/notifications/unread_count", nil)
req.Header.Set("X-User-ID", strconv.FormatUint(uint64(user.ID), 10))
router.ServeHTTP(w, req)
require.Equal(t, http.StatusOK, w.Code)
assert.Equal(t, "1", strings.TrimSpace(w.Body.String()))
sqlDB, _ := db.DB()
sqlDB.Close()
}
func TestNotificationReadAll(t *testing.T) {
db := setupNotificationDB(t)
handler := setupNotificationHandler(t, db)
@@ -461,9 +544,9 @@ func TestNotificationHandler_SnoozeWithDB(t *testing.T) {
notif := &model.Notification{UserID: user.ID, AccountID: &accountID, NotificationType: "conversation_assignment", PrimaryActorType: "conversation", PrimaryActorID: 1}
db.Create(notif)
// Snooze the notification
snoozeTime := time.Now().Add(2 * time.Hour).Format(time.RFC3339)
body := fmt.Sprintf(`{"snoozed_until":"%s"}`, snoozeTime)
// Chatwoot DateRangeHelper parses snoozed_until as Unix seconds.
snoozeUnix := time.Now().Add(2 * time.Hour).Unix()
body := fmt.Sprintf(`{"snoozed_until":%d}`, snoozeUnix)
w := httptest.NewRecorder()
req, _ := http.NewRequest("POST", fmt.Sprintf("/api/v1/accounts/1/notifications/%d/snooze", notif.ID), strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
@@ -476,6 +559,10 @@ func TestNotificationHandler_SnoozeWithDB(t *testing.T) {
var updated model.Notification
require.NoError(t, db.First(&updated, notif.ID).Error)
assert.NotNil(t, updated.SnoozedUntil)
assert.Equal(t, snoozeUnix, updated.SnoozedUntil.Unix())
var meta map[string]any
require.NoError(t, json.Unmarshal(updated.AdditionalAttributes, &meta))
assert.Contains(t, meta, "last_snoozed_at")
sqlDB, _ := db.DB()
sqlDB.Close()
@@ -486,8 +573,8 @@ func TestNotificationHandler_Snooze_InvalidID(t *testing.T) {
handler := setupNotificationHandler(t, db)
router := setupNotificationRouter(handler)
snoozeTime := time.Now().Add(2 * time.Hour).Format(time.RFC3339)
body := fmt.Sprintf(`{"snoozed_until":"%s"}`, snoozeTime)
snoozeTime := time.Now().Add(2 * time.Hour).Unix()
body := fmt.Sprintf(`{"snoozed_until":%d}`, snoozeTime)
w := httptest.NewRecorder()
req, _ := http.NewRequest("POST", "/api/v1/accounts/1/notifications/abc/snooze", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
@@ -517,7 +604,7 @@ func TestNotificationHandler_Snooze_MissingBody(t *testing.T) {
req.Header.Set("X-User-ID", strconv.FormatUint(uint64(user.ID), 10))
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Equal(t, http.StatusOK, w.Code)
sqlDB, _ := db.DB()
sqlDB.Close()