From 85391de22a68ceb19f6893cb90f977260ba0cc3a Mon Sep 17 00:00:00 2001 From: catalintomai <56169943+catalintomai@users.noreply.github.com> Date: Mon, 16 Jun 2025 02:30:21 +0200 Subject: [PATCH] MM-57326: [Shared Channels] Message priority, acknowledgement and persistent notifications need to be synced (#30736) --- .../api4/shared_channel_metadata_test.go | 659 ++++++++++++++++++ .../api4/shared_channel_test_utils.go | 245 +++++++ .../app/platform/shared_channel_notifier.go | 2 + server/channels/app/post.go | 11 +- server/channels/app/post_acknowledgements.go | 232 +++++- .../app/post_acknowledgements_test.go | 18 +- server/channels/db/migrations/migrations.list | 4 + ...hannelid_to_post_acknowledgements.down.sql | 29 + ..._channelid_to_post_acknowledgements.up.sql | 29 + ...hannelid_to_post_acknowledgements.down.sql | 2 + ..._channelid_to_post_acknowledgements.up.sql | 2 + .../channels/store/retrylayer/retrylayer.go | 130 +++- .../sqlstore/post_acknowledgements_store.go | 228 +++++- .../store/sqlstore/post_priority_store.go | 134 +++- server/channels/store/store.go | 8 +- .../mocks/PostAcknowledgementStore.go | 128 +++- .../storetest/mocks/PostPriorityStore.go | 48 ++ .../storetest/post_acknowledgements_store.go | 250 ++++++- .../channels/store/timerlayer/timerlayer.go | 102 ++- server/i18n/en.json | 16 + .../sharedchannel/mock_AppIface_test.go | 136 ++++ .../services/sharedchannel/service.go | 5 + .../services/sharedchannel/sync_recv.go | 228 +++++- .../services/sharedchannel/sync_send.go | 3 + .../sharedchannel/sync_send_remote.go | 90 ++- server/public/model/post_acknowledgement.go | 25 +- server/public/model/shared_channel.go | 4 + 27 files changed, 2667 insertions(+), 101 deletions(-) create mode 100644 server/channels/api4/shared_channel_metadata_test.go create mode 100644 server/channels/api4/shared_channel_test_utils.go create mode 100644 server/channels/db/migrations/mysql/000141_add_remoteid_channelid_to_post_acknowledgements.down.sql create mode 100644 server/channels/db/migrations/mysql/000141_add_remoteid_channelid_to_post_acknowledgements.up.sql create mode 100644 server/channels/db/migrations/postgres/000141_add_remoteid_channelid_to_post_acknowledgements.down.sql create mode 100644 server/channels/db/migrations/postgres/000141_add_remoteid_channelid_to_post_acknowledgements.up.sql diff --git a/server/channels/api4/shared_channel_metadata_test.go b/server/channels/api4/shared_channel_metadata_test.go new file mode 100644 index 0000000000..adfbdfeb28 --- /dev/null +++ b/server/channels/api4/shared_channel_metadata_test.go @@ -0,0 +1,659 @@ +// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved. +// See LICENSE.txt for license information. + +package api4 + +import ( + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/v8/platform/services/remotecluster" + "github.com/mattermost/mattermost/server/v8/platform/services/sharedchannel" +) + +// setupTestEnvironment sets up a common test environment for shared channel metadata tests +func setupTestEnvironment(t *testing.T) (*TestHelper, *sharedchannel.Service) { + th := setupForSharedChannels(t).InitBasic() + ss := th.App.Srv().Store() + EnsureCleanState(t, th, ss) + + // Set license with all enterprise features + license := model.NewTestLicense() + license.SkuShortName = model.LicenseShortSkuEnterprise + th.App.Srv().SetLicense(license) + + // Enable post priorities and persistent notifications + th.App.UpdateConfig(func(cfg *model.Config) { + *cfg.ServiceSettings.PostPriority = true + *cfg.ServiceSettings.AllowPersistentNotifications = true + *cfg.ServiceSettings.AllowPersistentNotificationsForGuests = true + *cfg.ServiceSettings.PersistentNotificationMaxRecipients = 100 + }) + + // Verify license and settings + require.NotNil(t, th.App.Srv().License(), "License should be active") + postPriorityEnabled := *th.App.Config().ServiceSettings.PostPriority + require.True(t, postPriorityEnabled, "Post priorities should be enabled") + + // Get the shared channel service and cast to concrete type + scsInterface := th.App.Srv().GetSharedChannelSyncService() + service, ok := scsInterface.(*sharedchannel.Service) + require.True(t, ok, "Expected sharedchannel.Service concrete type") + require.True(t, service.Active(), "SharedChannel service should be active") + + // Ensure services are running + err := service.Start() + require.NoError(t, err) + + rcService := th.App.Srv().GetRemoteClusterService() + if rcService != nil { + _ = rcService.Start() + if rc, ok := rcService.(*remotecluster.Service); ok { + rc.SetActive(true) + } + require.True(t, rcService.Active(), "RemoteClusterService should be active") + } + + return th, service +} + +// createSharedChannelSetup creates a shared channel with remote cluster for testing +func createSharedChannelSetup(t *testing.T, th *TestHelper, service *sharedchannel.Service, testServer *httptest.Server) (*model.Channel, *model.RemoteCluster) { + // Create remote cluster + selfCluster := &model.RemoteCluster{ + RemoteId: model.NewId(), + Name: "test-cluster-" + model.NewId()[:8], + SiteURL: testServer.URL, + CreateAt: model.GetMillis(), + LastPingAt: model.GetMillis(), + Token: model.NewId(), + CreatorId: th.BasicUser.Id, + RemoteToken: model.NewId(), + } + var err error + selfCluster, err = th.App.Srv().Store().RemoteCluster().Save(selfCluster) + require.NoError(t, err) + + // Create channel with users + testChannel := th.CreatePublicChannel() + _, appErr := th.App.AddUserToChannel(th.Context, th.BasicUser, testChannel, false) + require.Nil(t, appErr) + _, appErr = th.App.AddUserToChannel(th.Context, th.BasicUser2, testChannel, false) + require.Nil(t, appErr) + + // Create shared channel + sc := &model.SharedChannel{ + ChannelId: testChannel.Id, + TeamId: testChannel.TeamId, + Home: true, + ShareName: "test_sync_" + model.NewId()[:8], + CreatorId: th.BasicUser.Id, + RemoteId: selfCluster.RemoteId, + } + sc, err = th.App.ShareChannel(th.Context, sc) + require.NoError(t, err) + + // Create shared channel remote + scr := &model.SharedChannelRemote{ + Id: model.NewId(), + ChannelId: sc.ChannelId, + CreatorId: sc.CreatorId, + IsInviteAccepted: true, + IsInviteConfirmed: true, + RemoteId: sc.RemoteId, + } + _, err = th.App.SaveSharedChannelRemote(scr) + require.NoError(t, err) + + return testChannel, selfCluster +} + +func TestSharedChannelPostMetadataSync(t *testing.T) { + th, service := setupTestEnvironment(t) + defer th.TearDown() + + t.Run("Post Priority Metadata Self-Referential Sync", func(t *testing.T) { + var syncedPosts []*model.Post + + // Create test HTTP server using self-referential approach + testServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + writeOKResponse(w) + })) + defer testServer.Close() + + testChannel, selfCluster := createSharedChannelSetup(t, th, service, testServer) + + // Initialize sync handler + syncHandler := NewSelfReferentialSyncHandler(t, service, selfCluster) + syncHandler.OnPostSync = func(post *model.Post) { + t.Logf("Received synced post: ID=%s, Message=%s, HasMetadata=%v", post.Id, post.Message, post.Metadata != nil) + if post.Metadata != nil && post.Metadata.Priority != nil { + t.Logf("Post has priority metadata: Priority=%v, RequestedAck=%v, PersistentNotifications=%v", + post.Metadata.Priority.Priority, + post.Metadata.Priority.RequestedAck, + post.Metadata.Priority.PersistentNotifications) + } + syncedPosts = append(syncedPosts, post) + } + + // Update test server to use the sync handler + testServer.Config.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + syncHandler.HandleRequest(w, r) + }) + + // Create a local post with priority metadata + originalPost, appErr := th.App.CreatePost(th.Context, &model.Post{ + UserId: th.BasicUser.Id, + ChannelId: testChannel.Id, + Message: "Test post with priority metadata @" + th.BasicUser2.Username, + Metadata: &model.PostMetadata{ + Priority: &model.PostPriority{ + Priority: model.NewPointer(model.PostPriorityUrgent), + RequestedAck: model.NewPointer(true), + PersistentNotifications: model.NewPointer(true), + }, + }, + }, testChannel, model.CreatePostFlags{}) + require.Nil(t, appErr) + require.NotNil(t, originalPost) + + // Trigger sync + t.Logf("Triggering sync for channel: %s", testChannel.Id) + service.NotifyChannelChanged(testChannel.Id) + + // Wait for sync completion using Eventually pattern + require.Eventually(t, func() bool { + return len(syncedPosts) >= 2 + }, 5*time.Second, 100*time.Millisecond, "Should receive synced posts via self-referential handler") + + // Verify priority metadata is preserved through the complete sync flow + t.Logf("Found %d synced posts", len(syncedPosts)) + syncedPost := syncedPosts[len(syncedPosts)-1] + require.NotNil(t, syncedPost.Metadata, "Post metadata should be preserved") + require.NotNil(t, syncedPost.Metadata.Priority, "Priority metadata should be preserved") + assert.Equal(t, model.PostPriorityUrgent, *syncedPost.Metadata.Priority.Priority, "Priority should be preserved") + assert.True(t, *syncedPost.Metadata.Priority.RequestedAck, "RequestedAck should be preserved") + assert.True(t, *syncedPost.Metadata.Priority.PersistentNotifications, "PersistentNotifications should be preserved") + }) + + t.Run("Post Acknowledgement Metadata Self-Referential Sync", func(t *testing.T) { + EnsureCleanState(t, th, th.App.Srv().Store()) + var syncedPosts []*model.Post + + testServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + writeOKResponse(w) + })) + defer testServer.Close() + + testChannel, selfCluster := createSharedChannelSetup(t, th, service, testServer) + + syncHandler := NewSelfReferentialSyncHandler(t, service, selfCluster) + syncHandler.OnPostSync = func(post *model.Post) { + syncedPosts = append(syncedPosts, post) + } + + testServer.Config.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + syncHandler.HandleRequest(w, r) + }) + + // Create post with acknowledgement request + originalPost, appErr := th.App.CreatePost(th.Context, &model.Post{ + UserId: th.BasicUser.Id, + ChannelId: testChannel.Id, + Message: "Test post requesting acknowledgements @" + th.BasicUser2.Username, + Metadata: &model.PostMetadata{ + Priority: &model.PostPriority{ + Priority: model.NewPointer(model.PostPriorityUrgent), + RequestedAck: model.NewPointer(true), + PersistentNotifications: model.NewPointer(false), + }, + }, + }, testChannel, model.CreatePostFlags{}) + require.Nil(t, appErr) + + // Add acknowledgement to the post + ack := &model.PostAcknowledgement{ + PostId: originalPost.Id, + UserId: th.BasicUser2.Id, + ChannelId: originalPost.ChannelId, + } + _, appErr = th.App.SaveAcknowledgementForPostWithModel(th.Context, ack) + require.Nil(t, appErr) + + // Trigger sync + service.NotifyChannelChanged(testChannel.Id) + + // Wait for sync completion + require.Eventually(t, func() bool { + return len(syncedPosts) >= 2 + }, 5*time.Second, 100*time.Millisecond, "Should receive synced posts via self-referential handler") + + // Verify acknowledgement metadata is preserved + syncedPost := syncedPosts[len(syncedPosts)-1] + require.NotNil(t, syncedPost.Metadata, "Post metadata should be preserved") + require.NotNil(t, syncedPost.Metadata.Priority, "Priority metadata should be preserved") + assert.Equal(t, model.PostPriorityUrgent, *syncedPost.Metadata.Priority.Priority, "Priority should be preserved") + assert.True(t, *syncedPost.Metadata.Priority.RequestedAck, "RequestedAck should be preserved") + assert.False(t, *syncedPost.Metadata.Priority.PersistentNotifications, "PersistentNotifications should be preserved") + }) + + t.Run("Acknowledgement Count Sync Back to Sender", func(t *testing.T) { + EnsureCleanState(t, th, th.App.Srv().Store()) + var syncedPosts []*model.Post + var postIdToSync string + + testServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + writeOKResponse(w) + })) + defer testServer.Close() + + testChannel, selfCluster := createSharedChannelSetup(t, th, service, testServer) + + syncHandler := NewSelfReferentialSyncHandler(t, service, selfCluster) + syncHandler.OnPostSync = func(post *model.Post) { + if post.Id == postIdToSync { + t.Logf("Received sync for target post: ID=%s, HasAcks=%v", post.Id, + post.Metadata != nil && post.Metadata.Acknowledgements != nil) + if post.Metadata != nil && post.Metadata.Acknowledgements != nil { + t.Logf("Acknowledgement count in sync: %d", len(post.Metadata.Acknowledgements)) + } + } + syncedPosts = append(syncedPosts, post) + } + + testServer.Config.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + syncHandler.HandleRequest(w, r) + }) + + // Create post with acknowledgement request + originalPost, appErr := th.App.CreatePost(th.Context, &model.Post{ + UserId: th.BasicUser.Id, + ChannelId: testChannel.Id, + Message: "Test post for ack count sync @" + th.BasicUser2.Username, + Metadata: &model.PostMetadata{ + Priority: &model.PostPriority{ + Priority: model.NewPointer(model.PostPriorityUrgent), + RequestedAck: model.NewPointer(true), + PersistentNotifications: model.NewPointer(false), + }, + }, + }, testChannel, model.CreatePostFlags{}) + require.Nil(t, appErr) + postIdToSync = originalPost.Id + + // Verify initial state - no acknowledgements + acks, appErr := th.App.GetAcknowledgementsForPost(originalPost.Id) + require.Nil(t, appErr) + require.Empty(t, acks, "Should have no acknowledgements initially") + + // Trigger initial sync + service.NotifyChannelChanged(testChannel.Id) + + // Wait for initial sync + require.Eventually(t, func() bool { + return len(syncedPosts) >= 2 + }, 5*time.Second, 100*time.Millisecond, "Should complete initial sync") + + // Add acknowledgement + ackForSync := &model.PostAcknowledgement{ + PostId: originalPost.Id, + UserId: th.BasicUser2.Id, + ChannelId: originalPost.ChannelId, + AcknowledgedAt: model.GetMillis(), + } + _, appErr = th.App.SaveAcknowledgementForPostWithModel(th.Context, ackForSync) + require.Nil(t, appErr) + + // Clear previous synced posts and trigger sync + syncedPosts = syncedPosts[:0] + service.NotifyChannelChanged(testChannel.Id) + + // Wait for acknowledgement sync + require.Eventually(t, func() bool { + for _, post := range syncedPosts { + if post.Id == postIdToSync && + post.Metadata != nil && + post.Metadata.Acknowledgements != nil && + len(post.Metadata.Acknowledgements) > 0 { + return true + } + } + return false + }, 5*time.Second, 100*time.Millisecond, "Should sync acknowledgements back") + + // Verify acknowledgement was synced + var syncedPostWithAcks *model.Post + for _, post := range syncedPosts { + if post.Id == postIdToSync && post.Metadata != nil && post.Metadata.Acknowledgements != nil { + syncedPostWithAcks = post + break + } + } + + require.NotNil(t, syncedPostWithAcks, "Should find synced post with acknowledgements") + require.NotNil(t, syncedPostWithAcks.Metadata.Acknowledgements, "Acknowledgements should exist") + require.Len(t, syncedPostWithAcks.Metadata.Acknowledgements, 1, "Should have exactly 1 acknowledgement") + + ack := syncedPostWithAcks.Metadata.Acknowledgements[0] + assert.Equal(t, th.BasicUser2.Id, ack.UserId, "Acknowledgement should be from BasicUser2") + assert.Equal(t, originalPost.Id, ack.PostId, "Acknowledgement should be for the original post") + assert.Greater(t, ack.AcknowledgedAt, int64(0), "Acknowledgement should have a timestamp") + }) + + t.Run("Persistent Notifications Self-Referential Sync", func(t *testing.T) { + EnsureCleanState(t, th, th.App.Srv().Store()) + var syncedPosts []*model.Post + + testServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + writeOKResponse(w) + })) + defer testServer.Close() + + testChannel, selfCluster := createSharedChannelSetup(t, th, service, testServer) + + syncHandler := NewSelfReferentialSyncHandler(t, service, selfCluster) + syncHandler.OnPostSync = func(post *model.Post) { + syncedPosts = append(syncedPosts, post) + } + + testServer.Config.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + syncHandler.HandleRequest(w, r) + }) + + // Create post with persistent notifications enabled + _, appErr := th.App.CreatePost(th.Context, &model.Post{ + UserId: th.BasicUser.Id, + ChannelId: testChannel.Id, + Message: "Test post with persistent notifications @" + th.BasicUser2.Username, + Metadata: &model.PostMetadata{ + Priority: &model.PostPriority{ + Priority: model.NewPointer(model.PostPriorityUrgent), + RequestedAck: model.NewPointer(true), + PersistentNotifications: model.NewPointer(true), + }, + }, + }, testChannel, model.CreatePostFlags{}) + require.Nil(t, appErr) + + // Trigger sync + service.NotifyChannelChanged(testChannel.Id) + + // Wait for sync completion + require.Eventually(t, func() bool { + return len(syncedPosts) >= 2 + }, 5*time.Second, 100*time.Millisecond, "Should receive synced posts via self-referential handler") + + // Verify persistent notifications setting is preserved + syncedPost := syncedPosts[len(syncedPosts)-1] + require.NotNil(t, syncedPost.Metadata, "Post metadata should be preserved") + require.NotNil(t, syncedPost.Metadata.Priority, "Priority metadata should be preserved") + assert.Equal(t, model.PostPriorityUrgent, *syncedPost.Metadata.Priority.Priority, "Priority should be preserved") + assert.True(t, *syncedPost.Metadata.Priority.RequestedAck, "RequestedAck should be preserved") + assert.True(t, *syncedPost.Metadata.Priority.PersistentNotifications, "PersistentNotifications should be preserved") + }) + + t.Run("Cross-Cluster Acknowledgement End-to-End Flow", func(t *testing.T) { + EnsureCleanState(t, th, th.App.Srv().Store()) + var syncedPostsServerA []*model.Post + var syncedPostsServerB []*model.Post + var postIdToTrack string + + // Create test HTTP servers for both "clusters" + testServerA := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + writeOKResponse(w) + })) + defer testServerA.Close() + + testServerB := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + writeOKResponse(w) + })) + defer testServerB.Close() + + // Create remote clusters for both "servers" + clusterA := &model.RemoteCluster{ + RemoteId: model.NewId(), + Name: "cluster-a-ack-flow", + SiteURL: testServerA.URL, + CreateAt: model.GetMillis(), + LastPingAt: model.GetMillis(), + Token: model.NewId(), + CreatorId: th.BasicUser.Id, + RemoteToken: model.NewId(), + } + var err error + clusterA, err = th.App.Srv().Store().RemoteCluster().Save(clusterA) + require.NoError(t, err) + + clusterB := &model.RemoteCluster{ + RemoteId: model.NewId(), + Name: "cluster-b-ack-flow", + SiteURL: testServerB.URL, + CreateAt: model.GetMillis(), + LastPingAt: model.GetMillis(), + Token: model.NewId(), + CreatorId: th.BasicUser.Id, + RemoteToken: model.NewId(), + } + clusterB, err = th.App.Srv().Store().RemoteCluster().Save(clusterB) + require.NoError(t, err) + + // Create test channel and add local user + testChannel := th.CreatePublicChannel() + _, appErr := th.App.AddUserToChannel(th.Context, th.BasicUser, testChannel, false) + require.Nil(t, appErr) + + // Create remote user from Cluster B + remoteUserFromClusterB := &model.User{ + Email: "remote-user-b@example.com", + Username: "remoteuserb" + model.NewId()[:4], + Password: "password123", + EmailVerified: true, + RemoteId: &clusterB.RemoteId, + } + remoteUserFromClusterB, appErr = th.App.CreateUser(th.Context, remoteUserFromClusterB) + require.Nil(t, appErr) + + // Add remote user to team and channel + _, _, appErr = th.App.AddUserToTeam(th.Context, testChannel.TeamId, remoteUserFromClusterB.Id, "") + require.Nil(t, appErr) + _, appErr = th.App.AddUserToChannel(th.Context, remoteUserFromClusterB, testChannel, false) + require.Nil(t, appErr) + + // Create shared channel + sc := &model.SharedChannel{ + ChannelId: testChannel.Id, + TeamId: testChannel.TeamId, + Home: true, + ShareName: "test_cross_cluster_ack", + CreatorId: th.BasicUser.Id, + RemoteId: "", + } + sc, err = th.App.ShareChannel(th.Context, sc) + require.NoError(t, err) + + // Create shared channel remotes for both clusters + scrA := &model.SharedChannelRemote{ + Id: model.NewId(), + ChannelId: sc.ChannelId, + CreatorId: sc.CreatorId, + IsInviteAccepted: true, + IsInviteConfirmed: true, + RemoteId: clusterA.RemoteId, + } + _, err = th.App.SaveSharedChannelRemote(scrA) + require.NoError(t, err) + + scrB := &model.SharedChannelRemote{ + Id: model.NewId(), + ChannelId: sc.ChannelId, + CreatorId: sc.CreatorId, + IsInviteAccepted: true, + IsInviteConfirmed: true, + RemoteId: clusterB.RemoteId, + } + _, err = th.App.SaveSharedChannelRemote(scrB) + require.NoError(t, err) + + // Initialize sync handlers for both clusters + syncHandlerA := NewSelfReferentialSyncHandler(t, service, clusterA) + syncHandlerA.OnPostSync = func(post *model.Post) { + t.Logf("Cluster A received sync: ID=%s, Message=%s, HasAcks=%v", + post.Id, post.Message, + post.Metadata != nil && post.Metadata.Acknowledgements != nil) + if post.Metadata != nil && post.Metadata.Acknowledgements != nil { + t.Logf(" Cluster A sees %d acknowledgements", len(post.Metadata.Acknowledgements)) + } + syncedPostsServerA = append(syncedPostsServerA, post) + } + + syncHandlerB := NewSelfReferentialSyncHandler(t, service, clusterB) + syncHandlerB.OnPostSync = func(post *model.Post) { + t.Logf("Cluster B received sync: ID=%s, Message=%s, RequestedAck=%v", + post.Id, post.Message, + post.Metadata != nil && post.Metadata.Priority != nil && post.Metadata.Priority.RequestedAck != nil && *post.Metadata.Priority.RequestedAck) + syncedPostsServerB = append(syncedPostsServerB, post) + } + + // Update test servers to use sync handlers + testServerA.Config.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + syncHandlerA.HandleRequest(w, r) + }) + + testServerB.Config.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + syncHandlerB.HandleRequest(w, r) + }) + + // STEP 1: Server A creates a post with acknowledgement request + t.Log("=== STEP 1: Server A creates post with ack request ===") + originalPost, appErr := th.App.CreatePost(th.Context, &model.Post{ + UserId: th.BasicUser.Id, + ChannelId: testChannel.Id, + Message: "Cross-cluster ack test - please acknowledge", + Metadata: &model.PostMetadata{ + Priority: &model.PostPriority{ + Priority: model.NewPointer(model.PostPriorityUrgent), + RequestedAck: model.NewPointer(true), + PersistentNotifications: model.NewPointer(false), + }, + }, + }, testChannel, model.CreatePostFlags{}) + require.Nil(t, appErr) + postIdToTrack = originalPost.Id + + // Verify initial state - no acknowledgements + acks, appErr := th.App.GetAcknowledgementsForPost(originalPost.Id) + require.Nil(t, appErr) + require.Empty(t, acks, "Should have no acknowledgements initially") + + // STEP 2: Post syncs from Server A to Server B + t.Log("=== STEP 2: Post syncs from Server A to Server B ===") + service.NotifyChannelChanged(testChannel.Id) + + // Wait for Server B to receive the post + var syncedPostIdOnServerB string + require.Eventually(t, func() bool { + for _, post := range syncedPostsServerB { + if post.Message == originalPost.Message && post.Metadata != nil && post.Metadata.Priority != nil && + post.Metadata.Priority.RequestedAck != nil && *post.Metadata.Priority.RequestedAck { + syncedPostIdOnServerB = post.Id + t.Logf("Server B received post %s with ack request (original was %s)", post.Id, postIdToTrack) + return true + } + } + return false + }, 5*time.Second, 100*time.Millisecond, "Server B should receive post with ack request") + + // STEP 3: User on Server B acknowledges the post + t.Log("=== STEP 3: User on Server B acknowledges the post ===") + ackFromServerB := &model.PostAcknowledgement{ + PostId: syncedPostIdOnServerB, + UserId: remoteUserFromClusterB.Id, + ChannelId: testChannel.Id, + AcknowledgedAt: model.GetMillis(), + } + _, appErr = th.App.SaveAcknowledgementForPostWithModel(th.Context, ackFromServerB) + require.Nil(t, appErr) + + // Verify acknowledgement was saved locally + acksAfterSave, appErr := th.App.GetAcknowledgementsForPost(syncedPostIdOnServerB) + require.Nil(t, appErr) + require.Len(t, acksAfterSave, 1, "Should have exactly 1 acknowledgement after user B acks") + require.Equal(t, remoteUserFromClusterB.Id, acksAfterSave[0].UserId) + + // STEP 4: Acknowledgement syncs back from Server B to Server A + t.Log("=== STEP 4: Acknowledgement syncs back from Server B to Server A ===") + // Clear previous sync data to focus on acknowledgement sync + syncedPostsServerA = syncedPostsServerA[:0] + syncedPostsServerB = syncedPostsServerB[:0] + + // Trigger sync to send acknowledgement back to Server A + service.NotifyChannelChanged(testChannel.Id) + + // Wait for Server A to receive the acknowledgement sync + require.Eventually(t, func() bool { + for _, post := range syncedPostsServerA { + if post.Id == postIdToTrack && post.Metadata != nil && post.Metadata.Acknowledgements != nil { + t.Logf("Server A received post %s with %d acknowledgements", post.Id, len(post.Metadata.Acknowledgements)) + return len(post.Metadata.Acknowledgements) > 0 + } + } + return false + }, 5*time.Second, 100*time.Millisecond, "Server A should receive acknowledgement sync") + + // STEP 5: Verify the complete acknowledgement flow + t.Log("=== STEP 5: Verify complete acknowledgement flow ===") + var serverAPostWithAcks *model.Post + for _, post := range syncedPostsServerA { + if post.Id == postIdToTrack && post.Metadata != nil && post.Metadata.Acknowledgements != nil { + serverAPostWithAcks = post + break + } + } + + require.NotNil(t, serverAPostWithAcks, "Server A should receive post with acknowledgements") + require.NotNil(t, serverAPostWithAcks.Metadata.Acknowledgements, "Acknowledgements should exist") + require.Len(t, serverAPostWithAcks.Metadata.Acknowledgements, 1, "Should have exactly 1 acknowledgement") + + // Verify acknowledgement details + ack := serverAPostWithAcks.Metadata.Acknowledgements[0] + assert.Equal(t, remoteUserFromClusterB.Id, ack.UserId, "Acknowledgement should be from remote user") + assert.Equal(t, postIdToTrack, ack.PostId, "Acknowledgement should be for the correct post") + assert.Greater(t, ack.AcknowledgedAt, int64(0), "Acknowledgement should have a timestamp") + + // Verify priority metadata is preserved + require.NotNil(t, serverAPostWithAcks.Metadata.Priority, "Priority metadata should be preserved") + assert.Equal(t, model.PostPriorityUrgent, *serverAPostWithAcks.Metadata.Priority.Priority, "Priority should be preserved") + assert.True(t, *serverAPostWithAcks.Metadata.Priority.RequestedAck, "RequestedAck should be preserved") + + // STEP 6: Test echo prevention - verify no duplicate acknowledgements + t.Log("=== STEP 6: Test echo prevention ===") + syncedPostsServerA = syncedPostsServerA[:0] + + // Trigger another sync to ensure no duplicates are created + service.NotifyChannelChanged(testChannel.Id) + + // Verify acknowledgement count remains 1 (no duplicates) + require.Eventually(t, func() bool { + for _, post := range syncedPostsServerA { + if post.Id == postIdToTrack && post.Metadata != nil && post.Metadata.Acknowledgements != nil { + return len(post.Metadata.Acknowledgements) == 1 + } + } + return len(syncedPostsServerA) > 0 + }, 3*time.Second, 100*time.Millisecond, "Should maintain single acknowledgement after resync") + + t.Logf("✅ Cross-cluster acknowledgement flow completed successfully:") + t.Logf(" 1. Server A created post with ack request: %s", postIdToTrack) + t.Logf(" 2. Post synced to Server B with priority metadata intact") + t.Logf(" 3. User on Server B acknowledged the post: %s", ack.UserId) + t.Logf(" 4. Acknowledgement synced back to Server A") + t.Logf(" 5. Server A shows acknowledgement in post metadata") + t.Logf(" 6. Echo prevention verified - no duplicates created") + }) +} diff --git a/server/channels/api4/shared_channel_test_utils.go b/server/channels/api4/shared_channel_test_utils.go new file mode 100644 index 0000000000..ddd0e41a56 --- /dev/null +++ b/server/channels/api4/shared_channel_test_utils.go @@ -0,0 +1,245 @@ +// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved. +// See LICENSE.txt for license information. + +package api4 + +import ( + "encoding/json" + "io" + "net/http" + "sync/atomic" + "testing" + "time" + + "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/v8/channels/store" + "github.com/mattermost/mattermost/server/v8/channels/store/sqlstore" + "github.com/mattermost/mattermost/server/v8/platform/services/remotecluster" + "github.com/mattermost/mattermost/server/v8/platform/services/sharedchannel" + "github.com/stretchr/testify/require" +) + +// writeOKResponse writes a standard OK JSON response in the format expected by remotecluster +func writeOKResponse(w http.ResponseWriter) { + response := &remotecluster.Response{ + Status: "OK", + Err: "", + } + + // Set empty sync response as payload + syncResp := &model.SyncResponse{} + _ = response.SetPayload(syncResp) + + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + respBytes, _ := json.Marshal(response) + _, _ = w.Write(respBytes) +} + +// SelfReferentialSyncHandler handles incoming sync messages for self-referential tests. +type SelfReferentialSyncHandler struct { + t *testing.T + service *sharedchannel.Service + selfCluster *model.RemoteCluster + syncMessageCount *int32 + + // Callbacks for capturing sync data + OnPostSync func(post *model.Post) + OnAcknowledgementSync func(ack *model.PostAcknowledgement) +} + +// NewSelfReferentialSyncHandler creates a new handler for processing sync messages in tests +func NewSelfReferentialSyncHandler(t *testing.T, service *sharedchannel.Service, selfCluster *model.RemoteCluster) *SelfReferentialSyncHandler { + count := int32(0) + return &SelfReferentialSyncHandler{ + t: t, + service: service, + selfCluster: selfCluster, + syncMessageCount: &count, + } +} + +// HandleRequest processes incoming HTTP requests for the test server. +func (h *SelfReferentialSyncHandler) HandleRequest(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/v4/remotecluster/msg": + atomic.AddInt32(h.syncMessageCount, 1) + + // Read and process the sync message + body, _ := io.ReadAll(r.Body) + + // The message is wrapped in a RemoteClusterFrame + var frame model.RemoteClusterFrame + err := json.Unmarshal(body, &frame) + if err == nil { + // Process the message to update cursor + response := &remotecluster.Response{} + processErr := h.service.OnReceiveSyncMessageForTesting(frame.Msg, h.selfCluster, response) + if processErr != nil { + response.Status = "ERROR" + response.Err = processErr.Error() + h.t.Logf("Sync processing error: %v", processErr) + } else { + response.Status = "OK" + response.Err = "" + + var syncMsg model.SyncMsg + if unmarshalErr := json.Unmarshal(frame.Msg.Payload, &syncMsg); unmarshalErr == nil { + // Handle posts - call callback for verification + if len(syncMsg.Posts) > 0 && h.OnPostSync != nil { + for _, post := range syncMsg.Posts { + h.OnPostSync(post) + } + } + + // Handle acknowledgements - call callback for verification + if len(syncMsg.Acknowledgements) > 0 && h.OnAcknowledgementSync != nil { + for _, ack := range syncMsg.Acknowledgements { + h.OnAcknowledgementSync(ack) + } + } + + // Create success response + syncResp := &model.SyncResponse{ + UsersSyncd: make([]string, 0), + } + _ = response.SetPayload(syncResp) + } + } + + // Send the response + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + respBytes, _ := json.Marshal(response) + _, _ = w.Write(respBytes) + return + } + + writeOKResponse(w) + + case "/api/v4/remotecluster/ping": + writeOKResponse(w) + + case "/api/v4/remotecluster/confirm_invite": + writeOKResponse(w) + + default: + writeOKResponse(w) + } +} + +// GetSyncMessageCount returns the current count of sync messages received +func (h *SelfReferentialSyncHandler) GetSyncMessageCount() int32 { + return atomic.LoadInt32(h.syncMessageCount) +} + +// ensureCleanState ensures a clean test state by removing all shared channels and remote clusters. +// This helps prevent state pollution between tests. +func EnsureCleanState(t *testing.T, th *TestHelper, ss store.Store) { + t.Helper() + + // First, wait for any pending async tasks to complete, then shutdown services + scsInterface := th.App.Srv().GetSharedChannelSyncService() + if scsInterface != nil && scsInterface.Active() { + // Cast to concrete type to access testing methods + if service, ok := scsInterface.(*sharedchannel.Service); ok { + // Wait for any pending tasks from previous tests to complete + require.Eventually(t, func() bool { + return !service.HasPendingTasksForTesting() + }, 10*time.Second, 100*time.Millisecond, "All pending sync tasks should complete before cleanup") + } + + // Shutdown the shared channel service to stop any async operations + _ = scsInterface.Shutdown() + + // Wait for shutdown to complete with more time + require.Eventually(t, func() bool { + return !scsInterface.Active() + }, 5*time.Second, 100*time.Millisecond, "Shared channel service should be inactive after shutdown") + } + + // Clear all shared channels and remotes from previous tests + allSharedChannels, _ := ss.SharedChannel().GetAll(0, 1000, model.SharedChannelFilterOpts{}) + for _, sc := range allSharedChannels { + // Delete all remotes for this channel + remotes, _ := ss.SharedChannel().GetRemotes(0, 100, model.SharedChannelRemoteFilterOpts{ChannelId: sc.ChannelId}) + for _, remote := range remotes { + _, _ = ss.SharedChannel().DeleteRemote(remote.Id) + } + // Delete the shared channel + _, _ = ss.SharedChannel().Delete(sc.ChannelId) + } + + // Delete all remote clusters + allRemoteClusters, _ := ss.RemoteCluster().GetAll(0, 1000, model.RemoteClusterQueryFilter{}) + for _, rc := range allRemoteClusters { + _, _ = ss.RemoteCluster().Delete(rc.RemoteId) + } + + // Clear all acknowledgements from previous tests by setting AcknowledgedAt to 0 + // This uses raw SQL to ensure all acknowledgements are cleared regardless of their state + if sqlStore, ok := ss.(*sqlstore.SqlStore); ok { + _, _ = sqlStore.GetMaster().Exec("UPDATE PostAcknowledgements SET AcknowledgedAt = 0") + } + + // Remove all channel members from test channels (except the basic team/channel setup) + channels, _ := ss.Channel().GetAll(th.BasicTeam.Id) + for _, channel := range channels { + // Skip direct message and group channels, and skip the default channels + if channel.Type != model.ChannelTypeDirect && channel.Type != model.ChannelTypeGroup && + channel.Id != th.BasicChannel.Id { + members, _ := ss.Channel().GetMembers(model.ChannelMembersGetOptions{ + ChannelID: channel.Id, + }) + for _, member := range members { + _ = ss.Channel().RemoveMember(th.Context, channel.Id, member.UserId) + } + } + } + + // Get all active users and deactivate non-basic ones + options := &model.UserGetOptions{ + Page: 0, + PerPage: 200, + Active: true, + } + users, _ := ss.User().GetAllProfiles(options) + for _, user := range users { + // Keep only the basic test users active + if user.Id != th.BasicUser.Id && user.Id != th.BasicUser2.Id && + user.Id != th.SystemAdminUser.Id { + // Deactivate the user (soft delete) + user.DeleteAt = model.GetMillis() + _, _ = ss.User().Update(th.Context, user, true) + } + } + + // Verify cleanup is complete + require.Eventually(t, func() bool { + sharedChannels, _ := ss.SharedChannel().GetAll(0, 1000, model.SharedChannelFilterOpts{}) + remoteClusters, _ := ss.RemoteCluster().GetAll(0, 1000, model.RemoteClusterQueryFilter{}) + return len(sharedChannels) == 0 && len(remoteClusters) == 0 + }, 2*time.Second, 100*time.Millisecond, "Failed to clean up shared channels and remote clusters") + + // Restart services and ensure they are running and ready + if scsInterface != nil { + // Restart the shared channel service + _ = scsInterface.Start() + + if scs, ok := scsInterface.(*sharedchannel.Service); ok { + require.Eventually(t, func() bool { + return scs.Active() + }, 5*time.Second, 100*time.Millisecond, "Shared channel service should be active after restart") + } + } + + rcService := th.App.Srv().GetRemoteClusterService() + if rcService != nil { + if rc, ok := rcService.(*remotecluster.Service); ok { + rc.SetActive(true) + } + require.Eventually(t, func() bool { + return rcService.Active() + }, 5*time.Second, 100*time.Millisecond, "Remote cluster service should be active") + } +} diff --git a/server/channels/app/platform/shared_channel_notifier.go b/server/channels/app/platform/shared_channel_notifier.go index 84764748bc..3739b81b9c 100644 --- a/server/channels/app/platform/shared_channel_notifier.go +++ b/server/channels/app/platform/shared_channel_notifier.go @@ -21,6 +21,8 @@ var sharedChannelEventsForSync = []model.WebsocketEventType{ model.WebsocketEventPostDeleted, model.WebsocketEventReactionAdded, model.WebsocketEventReactionRemoved, + model.WebsocketEventAcknowledgementAdded, + model.WebsocketEventAcknowledgementRemoved, } var sharedChannelEventsForInvitation = []model.WebsocketEventType{ diff --git a/server/channels/app/post.go b/server/channels/app/post.go index b817c18281..eeeb0580aa 100644 --- a/server/channels/app/post.go +++ b/server/channels/app/post.go @@ -773,9 +773,14 @@ func (a *App) UpdatePost(c request.CTX, receivedUpdatedPost *model.Post, updateP if newPost == nil { return nil, model.NewAppError("UpdatePost", "Post rejected by plugin. "+rejectionReason, nil, "", http.StatusBadRequest) } - // Restore the post metadata that was stripped by the plugin. Set it to - // the last known good. - newPost.Metadata = oldPost.Metadata + // Always use incoming metadata when provided, otherwise retain existing + if receivedUpdatedPost.Metadata != nil { + newPost.Metadata = receivedUpdatedPost.Metadata.Copy() + } else { + // Restore the post metadata that was stripped by the plugin. Set it to + // the last known good. + newPost.Metadata = oldPost.Metadata + } rpost, nErr := a.Srv().Store().Post().Update(c, newPost, oldPost) if nErr != nil { diff --git a/server/channels/app/post_acknowledgements.go b/server/channels/app/post_acknowledgements.go index b39b456eb0..b10ecfb049 100644 --- a/server/channels/app/post_acknowledgements.go +++ b/server/channels/app/post_acknowledgements.go @@ -15,9 +15,19 @@ import ( ) func (a *App) SaveAcknowledgementForPost(c request.CTX, postID, userID string) (*model.PostAcknowledgement, *model.AppError) { - post, err := a.GetSinglePost(c, postID, false) - if err != nil { - return nil, err + return a.saveAcknowledgementForPostWithPost(c, nil, userID, postID) +} + +func (a *App) saveAcknowledgementForPostWithPost(c request.CTX, post *model.Post, userID string, postID ...string) (*model.PostAcknowledgement, *model.AppError) { + if post == nil { + if len(postID) == 0 { + return nil, model.NewAppError("SaveAcknowledgementForPost", "app.acknowledgement.save.missing_post.app_error", nil, "", http.StatusBadRequest) + } + var err *model.AppError + post, err = a.GetSinglePost(c, postID[0], false) + if err != nil { + return nil, err + } } channel, err := a.GetChannel(c, post.ChannelId) @@ -29,8 +39,14 @@ func (a *App) SaveAcknowledgementForPost(c request.CTX, postID, userID string) ( return nil, model.NewAppError("SaveAcknowledgementForPost", "api.acknowledgement.save.archived_channel.app_error", nil, "", http.StatusForbidden) } - acknowledgedAt := model.GetMillis() - acknowledgement, nErr := a.Srv().Store().PostAcknowledgement().Save(postID, userID, acknowledgedAt) + // Pre-populate the ChannelId to save a DB call in store + acknowledgement := &model.PostAcknowledgement{ + PostId: post.Id, + UserId: userID, + ChannelId: post.ChannelId, + } + + savedAck, nErr := a.Srv().Store().PostAcknowledgement().SaveWithModel(acknowledgement) if nErr != nil { var appErr *model.AppError switch { @@ -56,15 +72,28 @@ func (a *App) SaveAcknowledgementForPost(c request.CTX, postID, userID string) ( // The post is always modified since the UpdateAt always changes a.Srv().Store().Post().InvalidateLastPostTimeCache(channel.Id) - a.sendAcknowledgementEvent(c, model.WebsocketEventAcknowledgementAdded, acknowledgement, post) + a.sendAcknowledgementEvent(c, model.WebsocketEventAcknowledgementAdded, savedAck, post) - return acknowledgement, nil + // Trigger post updated event to ensure shared channel sync + a.sendPostUpdateEvent(c, post) + + return savedAck, nil } func (a *App) DeleteAcknowledgementForPost(c request.CTX, postID, userID string) *model.AppError { - post, err := a.GetSinglePost(c, postID, false) - if err != nil { - return err + return a.deleteAcknowledgementForPostWithPost(c, nil, userID, postID) +} + +func (a *App) deleteAcknowledgementForPostWithPost(c request.CTX, post *model.Post, userID string, postID ...string) *model.AppError { + if post == nil { + if len(postID) == 0 { + return model.NewAppError("DeleteAcknowledgementForPost", "app.acknowledgement.delete.missing_post.app_error", nil, "", http.StatusBadRequest) + } + var err *model.AppError + post, err = a.GetSinglePost(c, postID[0], false) + if err != nil { + return err + } } channel, err := a.GetChannel(c, post.ChannelId) @@ -76,7 +105,7 @@ func (a *App) DeleteAcknowledgementForPost(c request.CTX, postID, userID string) return model.NewAppError("DeleteAcknowledgementForPost", "api.acknowledgement.delete.archived_channel.app_error", nil, "", http.StatusForbidden) } - oldAck, nErr := a.Srv().Store().PostAcknowledgement().Get(postID, userID) + oldAck, nErr := a.Srv().Store().PostAcknowledgement().Get(post.Id, userID) if nErr != nil { var nfErr *store.ErrNotFound @@ -102,6 +131,9 @@ func (a *App) DeleteAcknowledgementForPost(c request.CTX, postID, userID string) a.sendAcknowledgementEvent(c, model.WebsocketEventAcknowledgementRemoved, oldAck, post) + // Trigger post updated event to ensure shared channel sync + a.sendPostUpdateEvent(c, post) + return nil } @@ -130,6 +162,80 @@ func (a *App) GetAcknowledgementsForPostList(postList *model.PostList) (map[stri return acknowledgementsMap, nil } +// SaveAcknowledgementsForPost saves multiple acknowledgements for a post in a single operation. +func (a *App) SaveAcknowledgementsForPost(c request.CTX, postID string, userIDs []string) ([]*model.PostAcknowledgement, *model.AppError) { + if len(userIDs) == 0 { + return []*model.PostAcknowledgement{}, nil + } + + post, err := a.GetSinglePost(c, postID, false) + if err != nil { + return nil, err + } + + channel, err := a.GetChannel(c, post.ChannelId) + if err != nil { + return nil, err + } + + if channel.DeleteAt > 0 { + return nil, model.NewAppError("SaveAcknowledgementsForPost", "api.acknowledgement.save.archived_channel.app_error", nil, "", http.StatusForbidden) + } + + // Create acknowledgements with current timestamp + acknowledgedAt := model.GetMillis() + var acknowledgements []*model.PostAcknowledgement + + for _, userID := range userIDs { + acknowledgements = append(acknowledgements, &model.PostAcknowledgement{ + PostId: post.Id, + UserId: userID, + ChannelId: post.ChannelId, + AcknowledgedAt: acknowledgedAt, + }) + } + + // Save all acknowledgements + savedAcks, nErr := a.Srv().Store().PostAcknowledgement().BatchSave(acknowledgements) + if nErr != nil { + var appErr *model.AppError + switch { + case errors.As(nErr, &appErr): + return nil, appErr + default: + return nil, model.NewAppError("SaveAcknowledgementsForPost", "app.acknowledgement.batch_save.app_error", nil, "", http.StatusInternalServerError).Wrap(nErr) + } + } + + // Resolve persistent notifications for each user + for _, userID := range userIDs { + if appErr := a.ResolvePersistentNotification(c, post, userID); appErr != nil { + a.CountNotificationReason(model.NotificationStatusError, model.NotificationTypeWebsocket, model.NotificationReasonResolvePersistentNotificationError, model.NotificationNoPlatform) + a.NotificationsLog().Error("Error resolving persistent notification", + mlog.String("sender_id", userID), + mlog.String("post_id", post.RootId), + mlog.String("status", model.NotificationStatusError), + mlog.String("reason", model.NotificationReasonResolvePersistentNotificationError), + mlog.Err(appErr), + ) + // We continue processing other acknowledgements even if one fails + } + } + + // The post is always modified since the UpdateAt always changes + a.Srv().Store().Post().InvalidateLastPostTimeCache(channel.Id) + + // Send WebSocket events for each acknowledgement + for _, ack := range savedAcks { + a.sendAcknowledgementEvent(c, model.WebsocketEventAcknowledgementAdded, ack, post) + } + + // Trigger post updated event to ensure shared channel sync + a.sendPostUpdateEvent(c, post) + + return savedAcks, nil +} + func (a *App) sendAcknowledgementEvent(rctx request.CTX, event model.WebsocketEventType, acknowledgement *model.PostAcknowledgement, post *model.Post) { // send out that a acknowledgement has been added/removed message := model.NewWebSocketEvent(event, "", post.ChannelId, "", nil, "") @@ -141,3 +247,107 @@ func (a *App) sendAcknowledgementEvent(rctx request.CTX, event model.WebsocketEv message.Add("acknowledgement", string(acknowledgementJSON)) a.Publish(message) } + +func (a *App) SaveAcknowledgementForPostWithModel(c request.CTX, acknowledgement *model.PostAcknowledgement) (*model.PostAcknowledgement, *model.AppError) { + // Get the post to verify it exists and get the channel + post, err := a.GetSinglePost(c, acknowledgement.PostId, false) + if err != nil { + return nil, err + } + + channel, err := a.GetChannel(c, post.ChannelId) + if err != nil { + return nil, err + } + + if channel.DeleteAt > 0 { + return nil, model.NewAppError("SaveAcknowledgementForPostWithModel", "api.acknowledgement.save.archived_channel.app_error", nil, "", http.StatusForbidden) + } + + // Make sure ChannelId is set + if acknowledgement.ChannelId == "" { + acknowledgement.ChannelId = post.ChannelId + } + + savedAck, nErr := a.Srv().Store().PostAcknowledgement().SaveWithModel(acknowledgement) + if nErr != nil { + var appErr *model.AppError + switch { + case errors.As(nErr, &appErr): + return nil, appErr + default: + return nil, model.NewAppError("SaveAcknowledgementForPostWithModel", "app.acknowledgement.save.save.app_error", nil, "", http.StatusInternalServerError).Wrap(nErr) + } + } + + if appErr := a.ResolvePersistentNotification(c, post, acknowledgement.UserId); appErr != nil { + a.CountNotificationReason(model.NotificationStatusError, model.NotificationTypeWebsocket, model.NotificationReasonResolvePersistentNotificationError, model.NotificationNoPlatform) + a.NotificationsLog().Error("Error resolving persistent notification", + mlog.String("sender_id", acknowledgement.UserId), + mlog.String("post_id", post.RootId), + mlog.String("status", model.NotificationStatusError), + mlog.String("reason", model.NotificationReasonResolvePersistentNotificationError), + mlog.Err(appErr), + ) + return nil, appErr + } + + // The post is always modified since the UpdateAt always changes + a.Srv().Store().Post().InvalidateLastPostTimeCache(channel.Id) + + a.sendAcknowledgementEvent(c, model.WebsocketEventAcknowledgementAdded, savedAck, post) + + // Trigger post updated event to ensure shared channel sync + a.sendPostUpdateEvent(c, post) + + return savedAck, nil +} + +func (a *App) DeleteAcknowledgementForPostWithModel(c request.CTX, acknowledgement *model.PostAcknowledgement) *model.AppError { + // Get the post to verify it exists and get the channel + post, err := a.GetSinglePost(c, acknowledgement.PostId, false) + if err != nil { + return err + } + + channel, err := a.GetChannel(c, post.ChannelId) + if err != nil { + return err + } + + if channel.DeleteAt > 0 { + return model.NewAppError("DeleteAcknowledgementForPostWithModel", "api.acknowledgement.delete.archived_channel.app_error", nil, "", http.StatusForbidden) + } + + nErr := a.Srv().Store().PostAcknowledgement().Delete(acknowledgement) + if nErr != nil { + return model.NewAppError("DeleteAcknowledgementForPostWithModel", "app.acknowledgement.delete.app_error", nil, "", http.StatusInternalServerError).Wrap(nErr) + } + + // The post is always modified since the UpdateAt always changes + a.Srv().Store().Post().InvalidateLastPostTimeCache(channel.Id) + + a.sendAcknowledgementEvent(c, model.WebsocketEventAcknowledgementRemoved, acknowledgement, post) + + // Trigger post updated event to ensure shared channel sync + a.sendPostUpdateEvent(c, post) + + return nil +} + +func (a *App) sendPostUpdateEvent(c request.CTX, post *model.Post) { + if post == nil { + c.Logger().Warn("sendPostUpdateEvent called with nil post") + return + } + + // Send a post edited event to trigger shared channel sync + message := model.NewWebSocketEvent(model.WebsocketEventPostEdited, "", post.ChannelId, "", nil, "") + + // Prepare the post with metadata for the event + preparedPost := a.PreparePostForClient(c, post, false, true, true) + + if appErr := a.publishWebsocketEventForPost(c, preparedPost, message); appErr != nil { + c.Logger().Warn("Failed to send post update event for acknowledgement sync", mlog.String("post_id", post.Id), mlog.Err(appErr)) + } +} diff --git a/server/channels/app/post_acknowledgements_test.go b/server/channels/app/post_acknowledgements_test.go index 8918abdd2c..660b8cec4b 100644 --- a/server/channels/app/post_acknowledgements_test.go +++ b/server/channels/app/post_acknowledgements_test.go @@ -109,7 +109,13 @@ func testDeleteAcknowledgementForPost(t *testing.T) { }) t.Run("delete acknowledgment for post after 5 min after acknowledged should not delete", func(t *testing.T) { - _, nErr := th.App.Srv().Store().PostAcknowledgement().Save(post.Id, th.BasicUser.Id, model.GetMillis()-int64(6*60*1000)) + acknowledgement := &model.PostAcknowledgement{ + PostId: post.Id, + UserId: th.BasicUser.Id, + AcknowledgedAt: model.GetMillis() - int64(6*60*1000), + ChannelId: post.ChannelId, + } + _, nErr := th.App.Srv().Store().PostAcknowledgement().SaveWithModel(acknowledgement) require.NoError(t, nErr) acknowledgments, err := th.App.GetAcknowledgementsForPost(post.Id) @@ -180,13 +186,13 @@ func testGetAcknowledgementsForPostList(t *testing.T) { acknowledgementsMap, err := th.App.GetAcknowledgementsForPostList(postList) require.Nil(t, err) - expected := map[string][]*model.PostAcknowledgement{ - p1.Id: acks1, - p2.Id: acks2, - } - require.Equal(t, expected, acknowledgementsMap) + // Verify p1 acknowledgements (order-agnostic) require.Len(t, acknowledgementsMap[p1.Id], 2) + require.ElementsMatch(t, acks1, acknowledgementsMap[p1.Id]) + + // Verify p2 acknowledgements (order-agnostic) require.Len(t, acknowledgementsMap[p2.Id], 1) + require.ElementsMatch(t, acks2, acknowledgementsMap[p2.Id]) require.Nil(t, acknowledgementsMap[p3.Id]) }) } diff --git a/server/channels/db/migrations/migrations.list b/server/channels/db/migrations/migrations.list index 0aad4d3c7c..026d392d38 100644 --- a/server/channels/db/migrations/migrations.list +++ b/server/channels/db/migrations/migrations.list @@ -277,6 +277,8 @@ channels/db/migrations/mysql/000139_remoteclusters_add_last_global_user_sync_at. channels/db/migrations/mysql/000139_remoteclusters_add_last_global_user_sync_at.up.sql channels/db/migrations/mysql/000140_add_lastmemberssyncat_to_sharedchannelremotes.down.sql channels/db/migrations/mysql/000140_add_lastmemberssyncat_to_sharedchannelremotes.up.sql +channels/db/migrations/mysql/000141_add_remoteid_channelid_to_post_acknowledgements.down.sql +channels/db/migrations/mysql/000141_add_remoteid_channelid_to_post_acknowledgements.up.sql channels/db/migrations/postgres/000001_create_teams.down.sql channels/db/migrations/postgres/000001_create_teams.up.sql channels/db/migrations/postgres/000002_create_team_members.down.sql @@ -555,3 +557,5 @@ channels/db/migrations/postgres/000139_remoteclusters_add_last_global_user_sync_ channels/db/migrations/postgres/000139_remoteclusters_add_last_global_user_sync_at.up.sql channels/db/migrations/postgres/000140_add_lastmemberssyncat_to_sharedchannelremotes.down.sql channels/db/migrations/postgres/000140_add_lastmemberssyncat_to_sharedchannelremotes.up.sql +channels/db/migrations/postgres/000141_add_remoteid_channelid_to_post_acknowledgements.down.sql +channels/db/migrations/postgres/000141_add_remoteid_channelid_to_post_acknowledgements.up.sql diff --git a/server/channels/db/migrations/mysql/000141_add_remoteid_channelid_to_post_acknowledgements.down.sql b/server/channels/db/migrations/mysql/000141_add_remoteid_channelid_to_post_acknowledgements.down.sql new file mode 100644 index 0000000000..2e5577f6e5 --- /dev/null +++ b/server/channels/db/migrations/mysql/000141_add_remoteid_channelid_to_post_acknowledgements.down.sql @@ -0,0 +1,29 @@ +SET @preparedStatement1 = (SELECT IF( + ( + SELECT COUNT(*) FROM INFORMATION_SCHEMA.COLUMNS + WHERE table_name = 'PostAcknowledgements' + AND table_schema = DATABASE() + AND column_name = 'RemoteId' + ) > 0, + 'ALTER TABLE PostAcknowledgements DROP COLUMN RemoteId;', + 'SELECT 1' +)); + +PREPARE alterIfExists FROM @preparedStatement1; +EXECUTE alterIfExists; +DEALLOCATE PREPARE alterIfExists; + +SET @preparedStatement2 = (SELECT IF( + ( + SELECT COUNT(*) FROM INFORMATION_SCHEMA.COLUMNS + WHERE table_name = 'PostAcknowledgements' + AND table_schema = DATABASE() + AND column_name = 'ChannelId' + ) > 0, + 'ALTER TABLE PostAcknowledgements DROP COLUMN ChannelId;', + 'SELECT 1' +)); + +PREPARE alterIfExists FROM @preparedStatement2; +EXECUTE alterIfExists; +DEALLOCATE PREPARE alterIfExists; \ No newline at end of file diff --git a/server/channels/db/migrations/mysql/000141_add_remoteid_channelid_to_post_acknowledgements.up.sql b/server/channels/db/migrations/mysql/000141_add_remoteid_channelid_to_post_acknowledgements.up.sql new file mode 100644 index 0000000000..8350895646 --- /dev/null +++ b/server/channels/db/migrations/mysql/000141_add_remoteid_channelid_to_post_acknowledgements.up.sql @@ -0,0 +1,29 @@ +SET @preparedStatement1 = (SELECT IF( + ( + SELECT COUNT(*) FROM INFORMATION_SCHEMA.COLUMNS + WHERE table_name = 'PostAcknowledgements' + AND table_schema = DATABASE() + AND column_name = 'RemoteId' + ) > 0, + 'SELECT 1;', + 'ALTER TABLE PostAcknowledgements ADD COLUMN RemoteId varchar(26) DEFAULT \'\';' +)); + +PREPARE addColumnIfNotExists FROM @preparedStatement1; +EXECUTE addColumnIfNotExists; +DEALLOCATE PREPARE addColumnIfNotExists; + +SET @preparedStatement2 = (SELECT IF( + ( + SELECT COUNT(*) FROM INFORMATION_SCHEMA.COLUMNS + WHERE table_name = 'PostAcknowledgements' + AND table_schema = DATABASE() + AND column_name = 'ChannelId' + ) > 0, + 'SELECT 1;', + 'ALTER TABLE PostAcknowledgements ADD COLUMN ChannelId varchar(26) DEFAULT \'\';' +)); + +PREPARE addColumnIfNotExists FROM @preparedStatement2; +EXECUTE addColumnIfNotExists; +DEALLOCATE PREPARE addColumnIfNotExists; \ No newline at end of file diff --git a/server/channels/db/migrations/postgres/000141_add_remoteid_channelid_to_post_acknowledgements.down.sql b/server/channels/db/migrations/postgres/000141_add_remoteid_channelid_to_post_acknowledgements.down.sql new file mode 100644 index 0000000000..5cd6650b2f --- /dev/null +++ b/server/channels/db/migrations/postgres/000141_add_remoteid_channelid_to_post_acknowledgements.down.sql @@ -0,0 +1,2 @@ +ALTER TABLE postacknowledgements DROP COLUMN IF EXISTS remoteid; +ALTER TABLE postacknowledgements DROP COLUMN IF EXISTS channelid; \ No newline at end of file diff --git a/server/channels/db/migrations/postgres/000141_add_remoteid_channelid_to_post_acknowledgements.up.sql b/server/channels/db/migrations/postgres/000141_add_remoteid_channelid_to_post_acknowledgements.up.sql new file mode 100644 index 0000000000..8f5c5d94b4 --- /dev/null +++ b/server/channels/db/migrations/postgres/000141_add_remoteid_channelid_to_post_acknowledgements.up.sql @@ -0,0 +1,2 @@ +ALTER TABLE postacknowledgements ADD COLUMN IF NOT EXISTS remoteid VARCHAR(26) DEFAULT ''; +ALTER TABLE postacknowledgements ADD COLUMN IF NOT EXISTS channelid VARCHAR(26) DEFAULT ''; \ No newline at end of file diff --git a/server/channels/store/retrylayer/retrylayer.go b/server/channels/store/retrylayer/retrylayer.go index a3ad8e2ce8..0a347f6b86 100644 --- a/server/channels/store/retrylayer/retrylayer.go +++ b/server/channels/store/retrylayer/retrylayer.go @@ -8570,6 +8570,48 @@ func (s *RetryLayerPostStore) Update(rctx request.CTX, newPost *model.Post, oldP } +func (s *RetryLayerPostAcknowledgementStore) BatchDelete(acknowledgements []*model.PostAcknowledgement) error { + + tries := 0 + for { + err := s.PostAcknowledgementStore.BatchDelete(acknowledgements) + if err == nil { + return nil + } + if !isRepeatableError(err) { + return err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return err + } + timepkg.Sleep(100 * timepkg.Millisecond) + } + +} + +func (s *RetryLayerPostAcknowledgementStore) BatchSave(acknowledgements []*model.PostAcknowledgement) ([]*model.PostAcknowledgement, error) { + + tries := 0 + for { + result, err := s.PostAcknowledgementStore.BatchSave(acknowledgements) + if err == nil { + return result, nil + } + if !isRepeatableError(err) { + return result, err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return result, err + } + timepkg.Sleep(100 * timepkg.Millisecond) + } + +} + func (s *RetryLayerPostAcknowledgementStore) Delete(acknowledgement *model.PostAcknowledgement) error { tries := 0 @@ -8633,6 +8675,27 @@ func (s *RetryLayerPostAcknowledgementStore) GetForPost(postID string) ([]*model } +func (s *RetryLayerPostAcknowledgementStore) GetForPostSince(postID string, since int64, excludeRemoteID string, inclDeleted bool) ([]*model.PostAcknowledgement, error) { + + tries := 0 + for { + result, err := s.PostAcknowledgementStore.GetForPostSince(postID, since, excludeRemoteID, inclDeleted) + if err == nil { + return result, nil + } + if !isRepeatableError(err) { + return result, err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return result, err + } + timepkg.Sleep(100 * timepkg.Millisecond) + } + +} + func (s *RetryLayerPostAcknowledgementStore) GetForPosts(postIds []string) ([]*model.PostAcknowledgement, error) { tries := 0 @@ -8654,11 +8717,32 @@ func (s *RetryLayerPostAcknowledgementStore) GetForPosts(postIds []string) ([]*m } -func (s *RetryLayerPostAcknowledgementStore) Save(postID string, userID string, acknowledgedAt int64) (*model.PostAcknowledgement, error) { +func (s *RetryLayerPostAcknowledgementStore) GetSingle(userID string, postID string, remoteID string) (*model.PostAcknowledgement, error) { tries := 0 for { - result, err := s.PostAcknowledgementStore.Save(postID, userID, acknowledgedAt) + result, err := s.PostAcknowledgementStore.GetSingle(userID, postID, remoteID) + if err == nil { + return result, nil + } + if !isRepeatableError(err) { + return result, err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return result, err + } + timepkg.Sleep(100 * timepkg.Millisecond) + } + +} + +func (s *RetryLayerPostAcknowledgementStore) SaveWithModel(acknowledgement *model.PostAcknowledgement) (*model.PostAcknowledgement, error) { + + tries := 0 + for { + result, err := s.PostAcknowledgementStore.SaveWithModel(acknowledgement) if err == nil { return result, nil } @@ -8822,6 +8906,27 @@ func (s *RetryLayerPostPersistentNotificationStore) UpdateLastActivity(postIds [ } +func (s *RetryLayerPostPriorityStore) Delete(postID string) error { + + tries := 0 + for { + err := s.PostPriorityStore.Delete(postID) + if err == nil { + return nil + } + if !isRepeatableError(err) { + return err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return err + } + timepkg.Sleep(100 * timepkg.Millisecond) + } + +} + func (s *RetryLayerPostPriorityStore) GetForPost(postID string) (*model.PostPriority, error) { tries := 0 @@ -8864,6 +8969,27 @@ func (s *RetryLayerPostPriorityStore) GetForPosts(ids []string) ([]*model.PostPr } +func (s *RetryLayerPostPriorityStore) Save(priority *model.PostPriority) (*model.PostPriority, error) { + + tries := 0 + for { + result, err := s.PostPriorityStore.Save(priority) + if err == nil { + return result, nil + } + if !isRepeatableError(err) { + return result, err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return result, err + } + timepkg.Sleep(100 * timepkg.Millisecond) + } + +} + func (s *RetryLayerPreferenceStore) CleanupFlagsBatch(limit int64) (int64, error) { tries := 0 diff --git a/server/channels/store/sqlstore/post_acknowledgements_store.go b/server/channels/store/sqlstore/post_acknowledgements_store.go index 0f70b5d9b7..cf67e4e631 100644 --- a/server/channels/store/sqlstore/post_acknowledgements_store.go +++ b/server/channels/store/sqlstore/post_acknowledgements_store.go @@ -23,7 +23,7 @@ func newSqlPostAcknowledgementStore(sqlStore *SqlStore) store.PostAcknowledgemen func (s *SqlPostAcknowledgementStore) Get(postID, userID string) (*model.PostAcknowledgement, error) { query := s.getQueryBuilder(). - Select("PostId", "UserId", "AcknowledgedAt"). + Select("PostId", "UserId", "ChannelId", "AcknowledgedAt", "RemoteId"). From("PostAcknowledgements"). Where(sq.And{ sq.Eq{"PostId": postID}, @@ -44,38 +44,20 @@ func (s *SqlPostAcknowledgementStore) Get(postID, userID string) (*model.PostAck return &acknowledgement, nil } -func (s *SqlPostAcknowledgementStore) Save(postID, userID string, acknowledgedAt int64) (*model.PostAcknowledgement, error) { - if acknowledgedAt == 0 { - acknowledgedAt = model.GetMillis() - } - - acknowledgement := &model.PostAcknowledgement{ - UserId: userID, - PostId: postID, - AcknowledgedAt: acknowledgedAt, - } - +func (s *SqlPostAcknowledgementStore) SaveWithModel(acknowledgement *model.PostAcknowledgement) (*model.PostAcknowledgement, error) { if err := acknowledgement.IsValid(); err != nil { return nil, err } + acknowledgement.PreSave() + transaction, err := s.GetMaster().Beginx() if err != nil { return nil, errors.Wrap(err, "begin_transaction") } defer finalizeTransactionX(transaction, &err) - query := s.getQueryBuilder(). - Insert("PostAcknowledgements"). - Columns("PostId", "UserId", "AcknowledgedAt"). - Values(acknowledgement.PostId, acknowledgement.UserId, acknowledgement.AcknowledgedAt) - - if s.DriverName() == model.DatabaseDriverMysql { - query = query.SuffixExpr(sq.Expr("ON DUPLICATE KEY UPDATE AcknowledgedAt = ?", acknowledgement.AcknowledgedAt)) - } else { - query = query.SuffixExpr(sq.Expr("ON CONFLICT (postid, userid) DO UPDATE SET AcknowledgedAt = ?", acknowledgement.AcknowledgedAt)) - } - + query := s.buildUpsertQuery(acknowledgement) _, err = transaction.ExecBuilder(query) if err != nil { return nil, err @@ -131,7 +113,7 @@ func (s *SqlPostAcknowledgementStore) GetForPost(postID string) ([]*model.PostAc var acknowledgements []*model.PostAcknowledgement query := s.getQueryBuilder(). - Select("PostId", "UserId", "AcknowledgedAt"). + Select("PostId", "UserId", "ChannelId", "AcknowledgedAt", "RemoteId"). From("PostAcknowledgements"). Where(sq.And{ sq.NotEq{"AcknowledgedAt": 0}, @@ -157,7 +139,7 @@ func (s *SqlPostAcknowledgementStore) GetForPosts(postIds []string) ([]*model.Po } query := s.getQueryBuilder(). - Select("PostId", "UserId", "AcknowledgedAt"). + Select("PostId", "UserId", "ChannelId", "AcknowledgedAt", "RemoteId"). From("PostAcknowledgements"). Where(sq.And{ sq.Eq{"PostId": postIds[i:j]}, @@ -176,6 +158,89 @@ func (s *SqlPostAcknowledgementStore) GetForPosts(postIds []string) ([]*model.Po return acknowledgements, nil } +func (s *SqlPostAcknowledgementStore) GetForPostSince(postID string, since int64, excludeRemoteID string, inclDeleted bool) ([]*model.PostAcknowledgement, error) { + var acknowledgements []*model.PostAcknowledgement + + query := s.getQueryBuilder(). + Select("PostId", "UserId", "ChannelId", "AcknowledgedAt", "RemoteId"). + From("PostAcknowledgements"). + Where(sq.Eq{"PostId": postID}) + + if !inclDeleted { + query = query.Where(sq.NotEq{"AcknowledgedAt": 0}) + } + + if since > 0 { + query = query.Where(sq.Gt{"AcknowledgedAt": since}) + } + + if excludeRemoteID != "" { + query = query.Where(sq.NotEq{"COALESCE(RemoteId, '')": excludeRemoteID}) + } + + err := s.GetReplica().SelectBuilder(&acknowledgements, query) + if err != nil { + return nil, errors.Wrapf(err, "failed to get PostAcknowledgements for postID=%s since=%d", postID, since) + } + + return acknowledgements, nil +} + +func (s *SqlPostAcknowledgementStore) GetSingle(userID, postID, remoteID string) (*model.PostAcknowledgement, error) { + query := s.getQueryBuilder(). + Select("PostId", "UserId", "ChannelId", "AcknowledgedAt", "RemoteId"). + From("PostAcknowledgements"). + Where(sq.And{ + sq.Eq{"PostId": postID}, + sq.Eq{"UserId": userID}, + }) + + if remoteID != "" { + query = query.Where(sq.Eq{"RemoteId": remoteID}) + } else { + query = query.Where(sq.Or{ + sq.Eq{"RemoteId": ""}, + sq.Eq{"RemoteId": nil}, + }) + } + + var acknowledgement model.PostAcknowledgement + err := s.GetReplica().GetBuilder(&acknowledgement, query) + if err != nil { + if err == sql.ErrNoRows { + return nil, store.NewErrNotFound("PostAcknowledgement", postID) + } + return nil, err + } + + return &acknowledgement, nil +} + +// buildUpsertQuery creates an upsert query for a PostAcknowledgement +func (s *SqlPostAcknowledgementStore) buildUpsertQuery(acknowledgement *model.PostAcknowledgement) sq.InsertBuilder { + columnsToInsert := []string{"PostId", "UserId", "ChannelId", "AcknowledgedAt", "RemoteId"} + var remoteIdValue any + if acknowledgement.RemoteId != nil { + remoteIdValue = *acknowledgement.RemoteId + } else { + remoteIdValue = nil + } + valuesToInsert := []any{acknowledgement.PostId, acknowledgement.UserId, acknowledgement.ChannelId, acknowledgement.AcknowledgedAt, remoteIdValue} + + query := s.getQueryBuilder(). + Insert("PostAcknowledgements"). + Columns(columnsToInsert...). + Values(valuesToInsert...) + + if s.DriverName() == model.DatabaseDriverMysql { + query = query.SuffixExpr(sq.Expr("ON DUPLICATE KEY UPDATE AcknowledgedAt = ?", acknowledgement.AcknowledgedAt)) + } else { + query = query.SuffixExpr(sq.Expr("ON CONFLICT (postid, userid) DO UPDATE SET AcknowledgedAt = ?", acknowledgement.AcknowledgedAt)) + } + + return query +} + func updatePost(transaction *sqlxTxWrapper, postId string) error { _, err := transaction.Exec( `UPDATE @@ -190,3 +255,116 @@ func updatePost(transaction *sqlxTxWrapper, postId string) error { return err } + +func (s *SqlPostAcknowledgementStore) BatchSave(acknowledgements []*model.PostAcknowledgement) ([]*model.PostAcknowledgement, error) { + if len(acknowledgements) == 0 { + return []*model.PostAcknowledgement{}, nil + } + + // Populate missing ChannelId fields and validate all acknowledgements + for _, ack := range acknowledgements { + // If ChannelId is not set, look it up from the post + if ack.ChannelId == "" { + postQuery := s.getQueryBuilder(). + Select("ChannelId"). + From("Posts"). + Where(sq.Eq{"Id": ack.PostId}) + + var channelId string + err := s.GetReplica().GetBuilder(&channelId, postQuery) + if err != nil { + return nil, errors.Wrapf(err, "failed to get channel id for post %s", ack.PostId) + } + ack.ChannelId = channelId + } + + if err := ack.IsValid(); err != nil { + return nil, err + } + } + + transaction, err := s.GetMaster().Beginx() + if err != nil { + return nil, errors.Wrap(err, "begin_transaction") + } + defer finalizeTransactionX(transaction, &err) + + // Keep track of which posts need to be updated + postsToUpdate := make(map[string]bool) + + // Insert all acknowledgements + for _, ack := range acknowledgements { + ack.PreSave() + + query := s.buildUpsertQuery(ack) + _, err = transaction.ExecBuilder(query) + if err != nil { + return nil, err + } + + postsToUpdate[ack.PostId] = true + } + + // Update the UpdateAt timestamp for all affected posts + for postID := range postsToUpdate { + err = updatePost(transaction, postID) + if err != nil { + return nil, err + } + } + + err = transaction.Commit() + if err != nil { + return nil, errors.Wrap(err, "commit_transaction") + } + + return acknowledgements, nil +} + +func (s *SqlPostAcknowledgementStore) BatchDelete(acknowledgements []*model.PostAcknowledgement) error { + if len(acknowledgements) == 0 { + return nil + } + + transaction, err := s.GetMaster().Beginx() + if err != nil { + return errors.Wrap(err, "begin_transaction") + } + defer finalizeTransactionX(transaction, &err) + + // Keep track of which posts need to be updated + postsToUpdate := make(map[string]bool) + + // Set AcknowledgedAt to 0 for all acknowledgements + for _, ack := range acknowledgements { + query := s.getQueryBuilder(). + Update("PostAcknowledgements"). + Set("AcknowledgedAt", 0). + Where(sq.And{ + sq.Eq{"PostId": ack.PostId}, + sq.Eq{"UserId": ack.UserId}, + }) + + _, err = transaction.ExecBuilder(query) + if err != nil { + return err + } + + postsToUpdate[ack.PostId] = true + } + + // Update the UpdateAt timestamp for all affected posts + for postID := range postsToUpdate { + err = updatePost(transaction, postID) + if err != nil { + return err + } + } + + err = transaction.Commit() + if err != nil { + return errors.Wrap(err, "commit_transaction") + } + + return nil +} diff --git a/server/channels/store/sqlstore/post_priority_store.go b/server/channels/store/sqlstore/post_priority_store.go index c1d137aded..7df2479199 100644 --- a/server/channels/store/sqlstore/post_priority_store.go +++ b/server/channels/store/sqlstore/post_priority_store.go @@ -5,6 +5,7 @@ package sqlstore import ( sq "github.com/mattermost/squirrel" + "github.com/pkg/errors" "github.com/mattermost/mattermost/server/public/model" "github.com/mattermost/mattermost/server/v8/channels/store" @@ -22,7 +23,7 @@ func newSqlPostPriorityStore(sqlStore *SqlStore) store.PostPriorityStore { func (s *SqlPostPriorityStore) GetForPost(postId string) (*model.PostPriority, error) { query := s.getQueryBuilder(). - Select("Priority", "RequestedAck", "PersistentNotifications"). + Select("PostId", "ChannelId", "Priority", "RequestedAck", "PersistentNotifications"). From("PostsPriority"). Where(sq.Eq{"PostId": postId}) @@ -46,12 +47,12 @@ func (s *SqlPostPriorityStore) GetForPosts(postIds []string) ([]*model.PostPrior } query := s.getQueryBuilder(). - Select("PostId", "Priority", "RequestedAck", "PersistentNotifications"). + Select("PostId", "ChannelId", "Priority", "RequestedAck", "PersistentNotifications"). From("PostsPriority"). Where(sq.Eq{"PostId": postIds[i:j]}) var priorityBatch []*model.PostPriority - err := s.GetReplica().SelectBuilder(&priority, query) + err := s.GetReplica().SelectBuilder(&priorityBatch, query) if err != nil { return nil, err @@ -62,3 +63,130 @@ func (s *SqlPostPriorityStore) GetForPosts(postIds []string) ([]*model.PostPrior return priority, nil } + +func (s *SqlPostPriorityStore) Save(priority *model.PostPriority) (*model.PostPriority, error) { + tx, err := s.GetMaster().Beginx() + if err != nil { + return nil, errors.Wrap(err, "begin_transaction") + } + defer finalizeTransactionX(tx, &err) + + // Delete existing priority + deleteQuery := s.getQueryBuilder(). + Delete("PostsPriority"). + Where(sq.Eq{"PostId": priority.PostId}) + + if _, err := tx.ExecBuilder(deleteQuery); err != nil { + return nil, errors.Wrap(err, "delete_existing_priority") + } + + // Insert new priority + insertQuery := s.getQueryBuilder(). + Insert("PostsPriority"). + Columns("PostId", "ChannelId", "Priority", "RequestedAck", "PersistentNotifications"). + Values(priority.PostId, priority.ChannelId, priority.Priority, priority.RequestedAck, priority.PersistentNotifications) + + if _, err := tx.ExecBuilder(insertQuery); err != nil { + return nil, errors.Wrap(err, "insert_priority") + } + + // Handle persistent notifications - always delete first, then insert if enabled + deletePersistentQuery := s.getQueryBuilder(). + Delete("PersistentNotifications"). + Where(sq.Eq{"PostId": priority.PostId}) + + if _, err := tx.ExecBuilder(deletePersistentQuery); err != nil { + return nil, errors.Wrap(err, "delete_persistent_notification") + } + + if priority.PersistentNotifications != nil && *priority.PersistentNotifications { + insertPersistentQuery := s.getQueryBuilder(). + Insert("PersistentNotifications"). + Columns("PostId", "CreateAt", "LastSentAt", "DeleteAt", "SentCount"). + Values(priority.PostId, model.GetMillis(), 0, 0, 0) + + if _, err := tx.ExecBuilder(insertPersistentQuery); err != nil { + return nil, errors.Wrap(err, "insert_persistent_notification") + } + } + + // Clear acknowledgements if not requested + if priority.RequestedAck == nil || !*priority.RequestedAck { + clearAckQuery := s.getQueryBuilder(). + Update("PostAcknowledgements"). + Set("AcknowledgedAt", 0). + Where(sq.Eq{"PostId": priority.PostId}) + + if _, err := tx.ExecBuilder(clearAckQuery); err != nil { + return nil, errors.Wrap(err, "clear_acknowledgements") + } + } + + // Update the post's UpdateAt to trigger clients to refresh + updatePostQuery := s.getQueryBuilder(). + Update("Posts"). + Set("UpdateAt", model.GetMillis()). + Where(sq.Eq{"Id": priority.PostId}) + + if _, err := tx.ExecBuilder(updatePostQuery); err != nil { + return nil, errors.Wrap(err, "update_post") + } + + if err := tx.Commit(); err != nil { + return nil, errors.Wrap(err, "commit_transaction") + } + + return priority, nil +} + +func (s *SqlPostPriorityStore) Delete(postId string) error { + tx, err := s.GetMaster().Beginx() + if err != nil { + return errors.Wrap(err, "begin_transaction") + } + defer finalizeTransactionX(tx, &err) + + // Delete from PostsPriority + deletePriorityQuery := s.getQueryBuilder(). + Delete("PostsPriority"). + Where(sq.Eq{"PostId": postId}) + + if _, err := tx.ExecBuilder(deletePriorityQuery); err != nil { + return errors.Wrap(err, "delete_priority") + } + + // Delete from PersistentNotifications + deletePersistentQuery := s.getQueryBuilder(). + Delete("PersistentNotifications"). + Where(sq.Eq{"PostId": postId}) + + if _, err := tx.ExecBuilder(deletePersistentQuery); err != nil { + return errors.Wrap(err, "delete_persistent_notification") + } + + // Clear acknowledgements + clearAckQuery := s.getQueryBuilder(). + Update("PostAcknowledgements"). + Set("AcknowledgedAt", 0). + Where(sq.Eq{"PostId": postId}) + + if _, err := tx.ExecBuilder(clearAckQuery); err != nil { + return errors.Wrap(err, "clear_acknowledgements") + } + + // Update post + updatePostQuery := s.getQueryBuilder(). + Update("Posts"). + Set("UpdateAt", model.GetMillis()). + Where(sq.Eq{"Id": postId}) + + if _, err := tx.ExecBuilder(updatePostQuery); err != nil { + return errors.Wrap(err, "update_post") + } + + if err := tx.Commit(); err != nil { + return errors.Wrap(err, "commit_transaction") + } + + return nil +} diff --git a/server/channels/store/store.go b/server/channels/store/store.go index 1e02b9793c..dd1c45c800 100644 --- a/server/channels/store/store.go +++ b/server/channels/store/store.go @@ -1035,6 +1035,8 @@ type SharedChannelStore interface { type PostPriorityStore interface { GetForPost(postID string) (*model.PostPriority, error) GetForPosts(ids []string) ([]*model.PostPriority, error) + Save(priority *model.PostPriority) (*model.PostPriority, error) + Delete(postID string) error } type DraftStore interface { @@ -1053,8 +1055,12 @@ type PostAcknowledgementStore interface { Get(postID, userID string) (*model.PostAcknowledgement, error) GetForPost(postID string) ([]*model.PostAcknowledgement, error) GetForPosts(postIds []string) ([]*model.PostAcknowledgement, error) - Save(postID, userID string, acknowledgedAt int64) (*model.PostAcknowledgement, error) + GetForPostSince(postID string, since int64, excludeRemoteID string, inclDeleted bool) ([]*model.PostAcknowledgement, error) + GetSingle(userID, postID, remoteID string) (*model.PostAcknowledgement, error) + SaveWithModel(acknowledgement *model.PostAcknowledgement) (*model.PostAcknowledgement, error) + BatchSave(acknowledgements []*model.PostAcknowledgement) ([]*model.PostAcknowledgement, error) Delete(acknowledgement *model.PostAcknowledgement) error + BatchDelete(acknowledgements []*model.PostAcknowledgement) error } type PostPersistentNotificationStore interface { diff --git a/server/channels/store/storetest/mocks/PostAcknowledgementStore.go b/server/channels/store/storetest/mocks/PostAcknowledgementStore.go index c3d48fe9e0..6b56eac0a1 100644 --- a/server/channels/store/storetest/mocks/PostAcknowledgementStore.go +++ b/server/channels/store/storetest/mocks/PostAcknowledgementStore.go @@ -14,6 +14,54 @@ type PostAcknowledgementStore struct { mock.Mock } +// BatchDelete provides a mock function with given fields: acknowledgements +func (_m *PostAcknowledgementStore) BatchDelete(acknowledgements []*model.PostAcknowledgement) error { + ret := _m.Called(acknowledgements) + + if len(ret) == 0 { + panic("no return value specified for BatchDelete") + } + + var r0 error + if rf, ok := ret.Get(0).(func([]*model.PostAcknowledgement) error); ok { + r0 = rf(acknowledgements) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// BatchSave provides a mock function with given fields: acknowledgements +func (_m *PostAcknowledgementStore) BatchSave(acknowledgements []*model.PostAcknowledgement) ([]*model.PostAcknowledgement, error) { + ret := _m.Called(acknowledgements) + + if len(ret) == 0 { + panic("no return value specified for BatchSave") + } + + var r0 []*model.PostAcknowledgement + var r1 error + if rf, ok := ret.Get(0).(func([]*model.PostAcknowledgement) ([]*model.PostAcknowledgement, error)); ok { + return rf(acknowledgements) + } + if rf, ok := ret.Get(0).(func([]*model.PostAcknowledgement) []*model.PostAcknowledgement); ok { + r0 = rf(acknowledgements) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.PostAcknowledgement) + } + } + + if rf, ok := ret.Get(1).(func([]*model.PostAcknowledgement) error); ok { + r1 = rf(acknowledgements) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // Delete provides a mock function with given fields: acknowledgement func (_m *PostAcknowledgementStore) Delete(acknowledgement *model.PostAcknowledgement) error { ret := _m.Called(acknowledgement) @@ -92,6 +140,36 @@ func (_m *PostAcknowledgementStore) GetForPost(postID string) ([]*model.PostAckn return r0, r1 } +// GetForPostSince provides a mock function with given fields: postID, since, excludeRemoteID, inclDeleted +func (_m *PostAcknowledgementStore) GetForPostSince(postID string, since int64, excludeRemoteID string, inclDeleted bool) ([]*model.PostAcknowledgement, error) { + ret := _m.Called(postID, since, excludeRemoteID, inclDeleted) + + if len(ret) == 0 { + panic("no return value specified for GetForPostSince") + } + + var r0 []*model.PostAcknowledgement + var r1 error + if rf, ok := ret.Get(0).(func(string, int64, string, bool) ([]*model.PostAcknowledgement, error)); ok { + return rf(postID, since, excludeRemoteID, inclDeleted) + } + if rf, ok := ret.Get(0).(func(string, int64, string, bool) []*model.PostAcknowledgement); ok { + r0 = rf(postID, since, excludeRemoteID, inclDeleted) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.PostAcknowledgement) + } + } + + if rf, ok := ret.Get(1).(func(string, int64, string, bool) error); ok { + r1 = rf(postID, since, excludeRemoteID, inclDeleted) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // GetForPosts provides a mock function with given fields: postIds func (_m *PostAcknowledgementStore) GetForPosts(postIds []string) ([]*model.PostAcknowledgement, error) { ret := _m.Called(postIds) @@ -122,29 +200,59 @@ func (_m *PostAcknowledgementStore) GetForPosts(postIds []string) ([]*model.Post return r0, r1 } -// Save provides a mock function with given fields: postID, userID, acknowledgedAt -func (_m *PostAcknowledgementStore) Save(postID string, userID string, acknowledgedAt int64) (*model.PostAcknowledgement, error) { - ret := _m.Called(postID, userID, acknowledgedAt) +// GetSingle provides a mock function with given fields: userID, postID, remoteID +func (_m *PostAcknowledgementStore) GetSingle(userID string, postID string, remoteID string) (*model.PostAcknowledgement, error) { + ret := _m.Called(userID, postID, remoteID) if len(ret) == 0 { - panic("no return value specified for Save") + panic("no return value specified for GetSingle") } var r0 *model.PostAcknowledgement var r1 error - if rf, ok := ret.Get(0).(func(string, string, int64) (*model.PostAcknowledgement, error)); ok { - return rf(postID, userID, acknowledgedAt) + if rf, ok := ret.Get(0).(func(string, string, string) (*model.PostAcknowledgement, error)); ok { + return rf(userID, postID, remoteID) } - if rf, ok := ret.Get(0).(func(string, string, int64) *model.PostAcknowledgement); ok { - r0 = rf(postID, userID, acknowledgedAt) + if rf, ok := ret.Get(0).(func(string, string, string) *model.PostAcknowledgement); ok { + r0 = rf(userID, postID, remoteID) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.PostAcknowledgement) } } - if rf, ok := ret.Get(1).(func(string, string, int64) error); ok { - r1 = rf(postID, userID, acknowledgedAt) + if rf, ok := ret.Get(1).(func(string, string, string) error); ok { + r1 = rf(userID, postID, remoteID) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// SaveWithModel provides a mock function with given fields: acknowledgement +func (_m *PostAcknowledgementStore) SaveWithModel(acknowledgement *model.PostAcknowledgement) (*model.PostAcknowledgement, error) { + ret := _m.Called(acknowledgement) + + if len(ret) == 0 { + panic("no return value specified for SaveWithModel") + } + + var r0 *model.PostAcknowledgement + var r1 error + if rf, ok := ret.Get(0).(func(*model.PostAcknowledgement) (*model.PostAcknowledgement, error)); ok { + return rf(acknowledgement) + } + if rf, ok := ret.Get(0).(func(*model.PostAcknowledgement) *model.PostAcknowledgement); ok { + r0 = rf(acknowledgement) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*model.PostAcknowledgement) + } + } + + if rf, ok := ret.Get(1).(func(*model.PostAcknowledgement) error); ok { + r1 = rf(acknowledgement) } else { r1 = ret.Error(1) } diff --git a/server/channels/store/storetest/mocks/PostPriorityStore.go b/server/channels/store/storetest/mocks/PostPriorityStore.go index 0f351a27db..852c2570ef 100644 --- a/server/channels/store/storetest/mocks/PostPriorityStore.go +++ b/server/channels/store/storetest/mocks/PostPriorityStore.go @@ -14,6 +14,24 @@ type PostPriorityStore struct { mock.Mock } +// Delete provides a mock function with given fields: postID +func (_m *PostPriorityStore) Delete(postID string) error { + ret := _m.Called(postID) + + if len(ret) == 0 { + panic("no return value specified for Delete") + } + + var r0 error + if rf, ok := ret.Get(0).(func(string) error); ok { + r0 = rf(postID) + } else { + r0 = ret.Error(0) + } + + return r0 +} + // GetForPost provides a mock function with given fields: postID func (_m *PostPriorityStore) GetForPost(postID string) (*model.PostPriority, error) { ret := _m.Called(postID) @@ -74,6 +92,36 @@ func (_m *PostPriorityStore) GetForPosts(ids []string) ([]*model.PostPriority, e return r0, r1 } +// Save provides a mock function with given fields: priority +func (_m *PostPriorityStore) Save(priority *model.PostPriority) (*model.PostPriority, error) { + ret := _m.Called(priority) + + if len(ret) == 0 { + panic("no return value specified for Save") + } + + var r0 *model.PostPriority + var r1 error + if rf, ok := ret.Get(0).(func(*model.PostPriority) (*model.PostPriority, error)); ok { + return rf(priority) + } + if rf, ok := ret.Get(0).(func(*model.PostPriority) *model.PostPriority); ok { + r0 = rf(priority) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*model.PostPriority) + } + } + + if rf, ok := ret.Get(1).(func(*model.PostPriority) error); ok { + r1 = rf(priority) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // NewPostPriorityStore creates a new instance of PostPriorityStore. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. // The first argument is typically a *testing.T value. func NewPostPriorityStore(t interface { diff --git a/server/channels/store/storetest/post_acknowledgements_store.go b/server/channels/store/storetest/post_acknowledgements_store.go index c94fffef1e..478ae2b26e 100644 --- a/server/channels/store/storetest/post_acknowledgements_store.go +++ b/server/channels/store/storetest/post_acknowledgements_store.go @@ -17,6 +17,8 @@ func TestPostAcknowledgementsStore(t *testing.T, rctx request.CTX, ss store.Stor t.Run("Save", func(t *testing.T) { testPostAcknowledgementsStoreSave(t, rctx, ss) }) t.Run("GetForPost", func(t *testing.T) { testPostAcknowledgementsStoreGetForPost(t, rctx, ss) }) t.Run("GetForPosts", func(t *testing.T) { testPostAcknowledgementsStoreGetForPosts(t, rctx, ss) }) + t.Run("BatchSave", func(t *testing.T) { testPostAcknowledgementsStoreBatchSave(t, rctx, ss) }) + t.Run("BatchDelete", func(t *testing.T) { testPostAcknowledgementsStoreBatchDelete(t, rctx, ss) }) } func testPostAcknowledgementsStoreSave(t *testing.T, rctx request.CTX, ss store.Store) { @@ -37,13 +39,16 @@ func testPostAcknowledgementsStoreSave(t *testing.T, rctx request.CTX, ss store. require.NoError(t, err) t.Run("consecutive saves should just update the acknowledged at", func(t *testing.T) { - _, err := ss.PostAcknowledgement().Save(post.Id, userID1, 0) + ack := &model.PostAcknowledgement{PostId: post.Id, UserId: userID1, AcknowledgedAt: 0, ChannelId: post.ChannelId} + _, err := ss.PostAcknowledgement().SaveWithModel(ack) require.NoError(t, err) - _, err = ss.PostAcknowledgement().Save(post.Id, userID1, 0) + ack = &model.PostAcknowledgement{PostId: post.Id, UserId: userID1, AcknowledgedAt: 0, ChannelId: post.ChannelId} + _, err = ss.PostAcknowledgement().SaveWithModel(ack) require.NoError(t, err) - ack1, err := ss.PostAcknowledgement().Save(post.Id, userID1, 0) + ack1 := &model.PostAcknowledgement{PostId: post.Id, UserId: userID1, AcknowledgedAt: 0, ChannelId: post.ChannelId} + ack1, err = ss.PostAcknowledgement().SaveWithModel(ack1) require.NoError(t, err) acknowledgements, err := ss.PostAcknowledgement().GetForPost(post.Id) @@ -53,7 +58,8 @@ func testPostAcknowledgementsStoreSave(t *testing.T, rctx request.CTX, ss store. t.Run("saving should update the update at of the post", func(t *testing.T) { oldUpdateAt := post.UpdateAt - _, err := ss.PostAcknowledgement().Save(post.Id, userID1, 0) + ack := &model.PostAcknowledgement{PostId: post.Id, UserId: userID1, AcknowledgedAt: 0, ChannelId: post.ChannelId} + _, err := ss.PostAcknowledgement().SaveWithModel(ack) require.NoError(t, err) post, err = ss.Post().GetSingle(rctx, post.Id, false) @@ -82,11 +88,14 @@ func testPostAcknowledgementsStoreGetForPost(t *testing.T, rctx request.CTX, ss require.NoError(t, err) t.Run("get acknowledgements for post", func(t *testing.T) { - ack1, err := ss.PostAcknowledgement().Save(p1.Id, userID1, 0) + ack1 := &model.PostAcknowledgement{PostId: p1.Id, UserId: userID1, AcknowledgedAt: 0, ChannelId: p1.ChannelId} + ack1, err := ss.PostAcknowledgement().SaveWithModel(ack1) require.NoError(t, err) - ack2, err := ss.PostAcknowledgement().Save(p1.Id, userID2, 0) + ack2 := &model.PostAcknowledgement{PostId: p1.Id, UserId: userID2, AcknowledgedAt: 0, ChannelId: p1.ChannelId} + ack2, err = ss.PostAcknowledgement().SaveWithModel(ack2) require.NoError(t, err) - ack3, err := ss.PostAcknowledgement().Save(p1.Id, userID3, 0) + ack3 := &model.PostAcknowledgement{PostId: p1.Id, UserId: userID3, AcknowledgedAt: 0, ChannelId: p1.ChannelId} + ack3, err = ss.PostAcknowledgement().SaveWithModel(ack3) require.NoError(t, err) acknowledgements, err := ss.PostAcknowledgement().GetForPost(p1.Id) @@ -145,13 +154,17 @@ func testPostAcknowledgementsStoreGetForPosts(t *testing.T, rctx request.CTX, ss require.Equal(t, -1, errIdx) t.Run("get acknowledgements for post", func(t *testing.T) { - ack1, err := ss.PostAcknowledgement().Save(p1.Id, userID1, 0) + ack1 := &model.PostAcknowledgement{PostId: p1.Id, UserId: userID1, AcknowledgedAt: 0, ChannelId: p1.ChannelId} + ack1, err := ss.PostAcknowledgement().SaveWithModel(ack1) require.NoError(t, err) - ack2, err := ss.PostAcknowledgement().Save(p1.Id, userID2, 0) + ack2 := &model.PostAcknowledgement{PostId: p1.Id, UserId: userID2, AcknowledgedAt: 0, ChannelId: p1.ChannelId} + ack2, err = ss.PostAcknowledgement().SaveWithModel(ack2) require.NoError(t, err) - ack3, err := ss.PostAcknowledgement().Save(p2.Id, userID2, 0) + ack3 := &model.PostAcknowledgement{PostId: p2.Id, UserId: userID2, AcknowledgedAt: 0, ChannelId: p2.ChannelId} + ack3, err = ss.PostAcknowledgement().SaveWithModel(ack3) require.NoError(t, err) - ack4, err := ss.PostAcknowledgement().Save(p2.Id, userID3, 0) + ack4 := &model.PostAcknowledgement{PostId: p2.Id, UserId: userID3, AcknowledgedAt: 0, ChannelId: p2.ChannelId} + ack4, err = ss.PostAcknowledgement().SaveWithModel(ack4) require.NoError(t, err) acknowledgements, err := ss.PostAcknowledgement().GetForPosts([]string{p1.Id}) @@ -191,3 +204,218 @@ func testPostAcknowledgementsStoreGetForPosts(t *testing.T, rctx request.CTX, ss require.Empty(t, acknowledgements) }) } + +func testPostAcknowledgementsStoreBatchSave(t *testing.T, rctx request.CTX, ss store.Store) { + userID1 := model.NewId() + userID2 := model.NewId() + userID3 := model.NewId() + + p1 := model.Post{} + p1.ChannelId = model.NewId() + p1.UserId = model.NewId() + p1.Message = NewTestID() + post, err := ss.Post().Save(rctx, &p1) + require.NoError(t, err) + + t.Run("batch save acknowledgements for a post", func(t *testing.T) { + // Create a batch of acknowledgements + acks := []*model.PostAcknowledgement{ + { + PostId: post.Id, + UserId: userID1, + AcknowledgedAt: model.GetMillis(), + }, + { + PostId: post.Id, + UserId: userID2, + AcknowledgedAt: model.GetMillis(), + }, + { + PostId: post.Id, + UserId: userID3, + AcknowledgedAt: model.GetMillis(), + }, + } + + // Save the batch + savedAcks, err := ss.PostAcknowledgement().BatchSave(acks) + require.NoError(t, err) + require.Len(t, savedAcks, 3) + + // Verify all were saved correctly + retrievedAcks, err := ss.PostAcknowledgement().GetForPost(post.Id) + require.NoError(t, err) + require.Len(t, retrievedAcks, 3) + + // Verify all users are in the saved acknowledgements + userIDMap := make(map[string]bool) + for _, ack := range retrievedAcks { + userIDMap[ack.UserId] = true + require.Equal(t, post.Id, ack.PostId) + require.Greater(t, ack.AcknowledgedAt, int64(0)) + } + + require.True(t, userIDMap[userID1]) + require.True(t, userIDMap[userID2]) + require.True(t, userIDMap[userID3]) + }) + + t.Run("batch save empty list of acknowledgements", func(t *testing.T) { + // Create an empty batch of acknowledgements + acks := []*model.PostAcknowledgement{} + + // Save the empty batch + savedAcks, err := ss.PostAcknowledgement().BatchSave(acks) + require.NoError(t, err) + require.Empty(t, savedAcks) + }) + + t.Run("batch save should update existing acknowledgements", func(t *testing.T) { + // First, delete all existing acknowledgements + acks, err := ss.PostAcknowledgement().GetForPost(post.Id) + require.NoError(t, err) + + for _, ack := range acks { + err = ss.PostAcknowledgement().Delete(ack) + require.NoError(t, err) + } + + // Create initial acknowledgement + ack := &model.PostAcknowledgement{PostId: post.Id, UserId: userID1, AcknowledgedAt: model.GetMillis(), ChannelId: post.ChannelId} + ack, err = ss.PostAcknowledgement().SaveWithModel(ack) + require.NoError(t, err) + + initialAckTime := ack.AcknowledgedAt + + // Create a batch with updated timestamp + newTimestamp := model.GetMillis() + 1000 + updatedAcks := []*model.PostAcknowledgement{ + { + PostId: post.Id, + UserId: userID1, + AcknowledgedAt: newTimestamp, + }, + } + + // Batch save should update the existing acknowledgement + savedAcks, err := ss.PostAcknowledgement().BatchSave(updatedAcks) + require.NoError(t, err) + require.Len(t, savedAcks, 1) + require.Equal(t, newTimestamp, savedAcks[0].AcknowledgedAt) + require.Greater(t, savedAcks[0].AcknowledgedAt, initialAckTime) + + // Verify the acknowledgement was updated + retrievedAcks, err := ss.PostAcknowledgement().GetForPost(post.Id) + require.NoError(t, err) + require.Len(t, retrievedAcks, 1) + require.Equal(t, newTimestamp, retrievedAcks[0].AcknowledgedAt) + }) + + t.Run("batch save should update post's update_at", func(t *testing.T) { + // First, check the current post update timestamp + currentPost, err := ss.Post().GetSingle(rctx, post.Id, false) + require.NoError(t, err) + oldUpdateAt := currentPost.UpdateAt + + // Create a batch of new acknowledgements + acks := []*model.PostAcknowledgement{ + { + PostId: post.Id, + UserId: model.NewId(), + AcknowledgedAt: model.GetMillis(), + }, + } + + // Save the batch + _, err = ss.PostAcknowledgement().BatchSave(acks) + require.NoError(t, err) + + // Verify post's update_at was updated + updatedPost, err := ss.Post().GetSingle(rctx, post.Id, false) + require.NoError(t, err) + require.Greater(t, updatedPost.UpdateAt, oldUpdateAt) + }) +} + +func testPostAcknowledgementsStoreBatchDelete(t *testing.T, rctx request.CTX, ss store.Store) { + userID1 := model.NewId() + userID2 := model.NewId() + userID3 := model.NewId() + + p1 := model.Post{} + p1.ChannelId = model.NewId() + p1.UserId = model.NewId() + p1.Message = NewTestID() + post, err := ss.Post().Save(rctx, &p1) + require.NoError(t, err) + + t.Run("batch delete all acknowledgements for a post", func(t *testing.T) { + // Create multiple acknowledgements + ack1 := &model.PostAcknowledgement{PostId: post.Id, UserId: userID1, AcknowledgedAt: 0, ChannelId: post.ChannelId} + ack1, err = ss.PostAcknowledgement().SaveWithModel(ack1) + require.NoError(t, err) + ack2 := &model.PostAcknowledgement{PostId: post.Id, UserId: userID2, AcknowledgedAt: 0, ChannelId: post.ChannelId} + ack2, err = ss.PostAcknowledgement().SaveWithModel(ack2) + require.NoError(t, err) + ack3 := &model.PostAcknowledgement{PostId: post.Id, UserId: userID3, AcknowledgedAt: 0, ChannelId: post.ChannelId} + ack3, err = ss.PostAcknowledgement().SaveWithModel(ack3) + require.NoError(t, err) + + // Verify acknowledgements were created + acks, pErr := ss.PostAcknowledgement().GetForPost(post.Id) + require.NoError(t, pErr) + require.Len(t, acks, 3) + + // Delete all acknowledgements in batch + err = ss.PostAcknowledgement().BatchDelete([]*model.PostAcknowledgement{ack1, ack2, ack3}) + require.NoError(t, err) + + // Verify all acknowledgements were deleted + acks, err = ss.PostAcknowledgement().GetForPost(post.Id) + require.NoError(t, err) + require.Empty(t, acks) + }) + + t.Run("batch delete should update post's update_at", func(t *testing.T) { + // Create acknowledgements + ack1 := &model.PostAcknowledgement{PostId: post.Id, UserId: userID1, AcknowledgedAt: 0, ChannelId: post.ChannelId} + ack1, err = ss.PostAcknowledgement().SaveWithModel(ack1) + require.NoError(t, err) + ack2 := &model.PostAcknowledgement{PostId: post.Id, UserId: userID2, AcknowledgedAt: 0, ChannelId: post.ChannelId} + ack2, err = ss.PostAcknowledgement().SaveWithModel(ack2) + require.NoError(t, err) + + // Get current post update timestamp + currentPost, err := ss.Post().GetSingle(rctx, post.Id, false) + require.NoError(t, err) + oldUpdateAt := currentPost.UpdateAt + + // Delete acknowledgements in batch + err = ss.PostAcknowledgement().BatchDelete([]*model.PostAcknowledgement{ack1, ack2}) + require.NoError(t, err) + + // Verify post's update_at was updated + updatedPost, err := ss.Post().GetSingle(rctx, post.Id, false) + require.NoError(t, err) + require.Greater(t, updatedPost.UpdateAt, oldUpdateAt) + }) + + t.Run("batch delete with empty list should not error", func(t *testing.T) { + // Delete with empty list should not error + err := ss.PostAcknowledgement().BatchDelete([]*model.PostAcknowledgement{}) + require.NoError(t, err) + }) + + t.Run("batch delete with non-existent acknowledgements should not error", func(t *testing.T) { + // Create non-existent acknowledgements + nonExistentAck := &model.PostAcknowledgement{ + PostId: model.NewId(), + UserId: model.NewId(), + AcknowledgedAt: model.GetMillis(), + } + + // Delete non-existent acknowledgement should not error + err := ss.PostAcknowledgement().BatchDelete([]*model.PostAcknowledgement{nonExistentAck}) + require.NoError(t, err) + }) +} diff --git a/server/channels/store/timerlayer/timerlayer.go b/server/channels/store/timerlayer/timerlayer.go index 97c94ae355..727f75b806 100644 --- a/server/channels/store/timerlayer/timerlayer.go +++ b/server/channels/store/timerlayer/timerlayer.go @@ -6821,6 +6821,38 @@ func (s *TimerLayerPostStore) Update(rctx request.CTX, newPost *model.Post, oldP return result, err } +func (s *TimerLayerPostAcknowledgementStore) BatchDelete(acknowledgements []*model.PostAcknowledgement) error { + start := time.Now() + + err := s.PostAcknowledgementStore.BatchDelete(acknowledgements) + + elapsed := float64(time.Since(start)) / float64(time.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("PostAcknowledgementStore.BatchDelete", success, elapsed) + } + return err +} + +func (s *TimerLayerPostAcknowledgementStore) BatchSave(acknowledgements []*model.PostAcknowledgement) ([]*model.PostAcknowledgement, error) { + start := time.Now() + + result, err := s.PostAcknowledgementStore.BatchSave(acknowledgements) + + elapsed := float64(time.Since(start)) / float64(time.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("PostAcknowledgementStore.BatchSave", success, elapsed) + } + return result, err +} + func (s *TimerLayerPostAcknowledgementStore) Delete(acknowledgement *model.PostAcknowledgement) error { start := time.Now() @@ -6869,6 +6901,22 @@ func (s *TimerLayerPostAcknowledgementStore) GetForPost(postID string) ([]*model return result, err } +func (s *TimerLayerPostAcknowledgementStore) GetForPostSince(postID string, since int64, excludeRemoteID string, inclDeleted bool) ([]*model.PostAcknowledgement, error) { + start := time.Now() + + result, err := s.PostAcknowledgementStore.GetForPostSince(postID, since, excludeRemoteID, inclDeleted) + + elapsed := float64(time.Since(start)) / float64(time.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("PostAcknowledgementStore.GetForPostSince", success, elapsed) + } + return result, err +} + func (s *TimerLayerPostAcknowledgementStore) GetForPosts(postIds []string) ([]*model.PostAcknowledgement, error) { start := time.Now() @@ -6885,10 +6933,10 @@ func (s *TimerLayerPostAcknowledgementStore) GetForPosts(postIds []string) ([]*m return result, err } -func (s *TimerLayerPostAcknowledgementStore) Save(postID string, userID string, acknowledgedAt int64) (*model.PostAcknowledgement, error) { +func (s *TimerLayerPostAcknowledgementStore) GetSingle(userID string, postID string, remoteID string) (*model.PostAcknowledgement, error) { start := time.Now() - result, err := s.PostAcknowledgementStore.Save(postID, userID, acknowledgedAt) + result, err := s.PostAcknowledgementStore.GetSingle(userID, postID, remoteID) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -6896,7 +6944,23 @@ func (s *TimerLayerPostAcknowledgementStore) Save(postID string, userID string, if err == nil { success = "true" } - s.Root.Metrics.ObserveStoreMethodDuration("PostAcknowledgementStore.Save", success, elapsed) + s.Root.Metrics.ObserveStoreMethodDuration("PostAcknowledgementStore.GetSingle", success, elapsed) + } + return result, err +} + +func (s *TimerLayerPostAcknowledgementStore) SaveWithModel(acknowledgement *model.PostAcknowledgement) (*model.PostAcknowledgement, error) { + start := time.Now() + + result, err := s.PostAcknowledgementStore.SaveWithModel(acknowledgement) + + elapsed := float64(time.Since(start)) / float64(time.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("PostAcknowledgementStore.SaveWithModel", success, elapsed) } return result, err } @@ -7013,6 +7077,22 @@ func (s *TimerLayerPostPersistentNotificationStore) UpdateLastActivity(postIds [ return err } +func (s *TimerLayerPostPriorityStore) Delete(postID string) error { + start := time.Now() + + err := s.PostPriorityStore.Delete(postID) + + elapsed := float64(time.Since(start)) / float64(time.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("PostPriorityStore.Delete", success, elapsed) + } + return err +} + func (s *TimerLayerPostPriorityStore) GetForPost(postID string) (*model.PostPriority, error) { start := time.Now() @@ -7045,6 +7125,22 @@ func (s *TimerLayerPostPriorityStore) GetForPosts(ids []string) ([]*model.PostPr return result, err } +func (s *TimerLayerPostPriorityStore) Save(priority *model.PostPriority) (*model.PostPriority, error) { + start := time.Now() + + result, err := s.PostPriorityStore.Save(priority) + + elapsed := float64(time.Since(start)) / float64(time.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("PostPriorityStore.Save", success, elapsed) + } + return result, err +} + func (s *TimerLayerPreferenceStore) CleanupFlagsBatch(limit int64) (int64, error) { start := time.Now() diff --git a/server/i18n/en.json b/server/i18n/en.json index f2343e4929..a1aabf8284 100644 --- a/server/i18n/en.json +++ b/server/i18n/en.json @@ -4506,10 +4506,18 @@ "id": "api4.plugin.reattachPlugin.invalid_request", "translation": "Failed to parse request" }, + { + "id": "app.acknowledgement.batch_save.app_error", + "translation": "Failed to save the batch of acknowledgement objects" + }, { "id": "app.acknowledgement.delete.app_error", "translation": "Unable to delete acknowledgement." }, + { + "id": "app.acknowledgement.delete.missing_post.app_error", + "translation": "Cannot delete acknowledgement for missing post" + }, { "id": "app.acknowledgement.get.app_error", "translation": "Unable to get acknowledgement." @@ -4518,6 +4526,10 @@ "id": "app.acknowledgement.getforpost.get.app_error", "translation": "Unable to get acknowledgement for post." }, + { + "id": "app.acknowledgement.save.missing_post.app_error", + "translation": "Cannot save acknowledgement for missing post" + }, { "id": "app.acknowledgement.save.save.app_error", "translation": "Unable to save acknowledgement for post." @@ -8752,6 +8764,10 @@ "id": "model.access_policy.is_valid.version.app_error", "translation": "Version is not valid for this access control policy." }, + { + "id": "model.acknowledgement.is_valid.channel_id.app_error", + "translation": "Invalid channel id." + }, { "id": "model.acknowledgement.is_valid.post_id.app_error", "translation": "Invalid post id." diff --git a/server/platform/services/sharedchannel/mock_AppIface_test.go b/server/platform/services/sharedchannel/mock_AppIface_test.go index b6a5afc8a4..8bc2ff593a 100644 --- a/server/platform/services/sharedchannel/mock_AppIface_test.go +++ b/server/platform/services/sharedchannel/mock_AppIface_test.go @@ -205,6 +205,26 @@ func (_m *MockAppIface) CreateUploadSession(c request.CTX, us *model.UploadSessi return r0, r1 } +// DeleteAcknowledgementForPostWithModel provides a mock function with given fields: c, acknowledgement +func (_m *MockAppIface) DeleteAcknowledgementForPostWithModel(c request.CTX, acknowledgement *model.PostAcknowledgement) *model.AppError { + ret := _m.Called(c, acknowledgement) + + if len(ret) == 0 { + panic("no return value specified for DeleteAcknowledgementForPostWithModel") + } + + var r0 *model.AppError + if rf, ok := ret.Get(0).(func(request.CTX, *model.PostAcknowledgement) *model.AppError); ok { + r0 = rf(c, acknowledgement) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*model.AppError) + } + } + + return r0 +} + // DeletePost provides a mock function with given fields: c, postID, deleteByID func (_m *MockAppIface) DeletePost(c request.CTX, postID string, deleteByID string) (*model.Post, *model.AppError) { ret := _m.Called(c, postID, deleteByID) @@ -289,6 +309,38 @@ func (_m *MockAppIface) FileReader(path string) (filestore.ReadCloseSeeker, *mod return r0, r1 } +// GetAcknowledgementsForPost provides a mock function with given fields: postID +func (_m *MockAppIface) GetAcknowledgementsForPost(postID string) ([]*model.PostAcknowledgement, *model.AppError) { + ret := _m.Called(postID) + + if len(ret) == 0 { + panic("no return value specified for GetAcknowledgementsForPost") + } + + var r0 []*model.PostAcknowledgement + var r1 *model.AppError + if rf, ok := ret.Get(0).(func(string) ([]*model.PostAcknowledgement, *model.AppError)); ok { + return rf(postID) + } + if rf, ok := ret.Get(0).(func(string) []*model.PostAcknowledgement); ok { + r0 = rf(postID) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.PostAcknowledgement) + } + } + + if rf, ok := ret.Get(1).(func(string) *model.AppError); ok { + r1 = rf(postID) + } else { + if ret.Get(1) != nil { + r1 = ret.Get(1).(*model.AppError) + } + } + + return r0, r1 +} + // GetOrCreateDirectChannel provides a mock function with given fields: c, userId, otherUserId, channelOptions func (_m *MockAppIface) GetOrCreateDirectChannel(c request.CTX, userId string, otherUserId string, channelOptions ...model.ChannelOption) (*model.Channel, *model.AppError) { _va := make([]interface{}, len(channelOptions)) @@ -508,6 +560,26 @@ func (_m *MockAppIface) PermanentDeleteChannel(c request.CTX, channel *model.Cha return r0 } +// PreparePostForClient provides a mock function with given fields: c, post, isNewPost, includeDeleted, includePriority +func (_m *MockAppIface) PreparePostForClient(c request.CTX, post *model.Post, isNewPost bool, includeDeleted bool, includePriority bool) *model.Post { + ret := _m.Called(c, post, isNewPost, includeDeleted, includePriority) + + if len(ret) == 0 { + panic("no return value specified for PreparePostForClient") + } + + var r0 *model.Post + if rf, ok := ret.Get(0).(func(request.CTX, *model.Post, bool, bool, bool) *model.Post); ok { + r0 = rf(c, post, isNewPost, includeDeleted, includePriority) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*model.Post) + } + } + + return r0 +} + // Publish provides a mock function with given fields: message func (_m *MockAppIface) Publish(message *model.WebSocketEvent) { _m.Called(message) @@ -533,6 +605,70 @@ func (_m *MockAppIface) RemoveUserFromChannel(c request.CTX, userID string, remo return r0 } +// SaveAcknowledgementForPostWithModel provides a mock function with given fields: c, acknowledgement +func (_m *MockAppIface) SaveAcknowledgementForPostWithModel(c request.CTX, acknowledgement *model.PostAcknowledgement) (*model.PostAcknowledgement, *model.AppError) { + ret := _m.Called(c, acknowledgement) + + if len(ret) == 0 { + panic("no return value specified for SaveAcknowledgementForPostWithModel") + } + + var r0 *model.PostAcknowledgement + var r1 *model.AppError + if rf, ok := ret.Get(0).(func(request.CTX, *model.PostAcknowledgement) (*model.PostAcknowledgement, *model.AppError)); ok { + return rf(c, acknowledgement) + } + if rf, ok := ret.Get(0).(func(request.CTX, *model.PostAcknowledgement) *model.PostAcknowledgement); ok { + r0 = rf(c, acknowledgement) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*model.PostAcknowledgement) + } + } + + if rf, ok := ret.Get(1).(func(request.CTX, *model.PostAcknowledgement) *model.AppError); ok { + r1 = rf(c, acknowledgement) + } else { + if ret.Get(1) != nil { + r1 = ret.Get(1).(*model.AppError) + } + } + + return r0, r1 +} + +// SaveAcknowledgementsForPost provides a mock function with given fields: c, postID, userIDs +func (_m *MockAppIface) SaveAcknowledgementsForPost(c request.CTX, postID string, userIDs []string) ([]*model.PostAcknowledgement, *model.AppError) { + ret := _m.Called(c, postID, userIDs) + + if len(ret) == 0 { + panic("no return value specified for SaveAcknowledgementsForPost") + } + + var r0 []*model.PostAcknowledgement + var r1 *model.AppError + if rf, ok := ret.Get(0).(func(request.CTX, string, []string) ([]*model.PostAcknowledgement, *model.AppError)); ok { + return rf(c, postID, userIDs) + } + if rf, ok := ret.Get(0).(func(request.CTX, string, []string) []*model.PostAcknowledgement); ok { + r0 = rf(c, postID, userIDs) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.PostAcknowledgement) + } + } + + if rf, ok := ret.Get(1).(func(request.CTX, string, []string) *model.AppError); ok { + r1 = rf(c, postID, userIDs) + } else { + if ret.Get(1) != nil { + r1 = ret.Get(1).(*model.AppError) + } + } + + return r0, r1 +} + // SaveAndBroadcastStatus provides a mock function with given fields: status func (_m *MockAppIface) SaveAndBroadcastStatus(status *model.Status) { _m.Called(status) diff --git a/server/platform/services/sharedchannel/service.go b/server/platform/services/sharedchannel/service.go index 503d4cbb66..65b976e449 100644 --- a/server/platform/services/sharedchannel/service.go +++ b/server/platform/services/sharedchannel/service.go @@ -78,6 +78,11 @@ type AppIface interface { OnSharedChannelsAttachmentSyncMsg(fi *model.FileInfo, post *model.Post, rc *model.RemoteCluster) error OnSharedChannelsProfileImageSyncMsg(user *model.User, rc *model.RemoteCluster) error Publish(message *model.WebSocketEvent) + SaveAcknowledgementForPostWithModel(c request.CTX, acknowledgement *model.PostAcknowledgement) (*model.PostAcknowledgement, *model.AppError) + DeleteAcknowledgementForPostWithModel(c request.CTX, acknowledgement *model.PostAcknowledgement) *model.AppError + SaveAcknowledgementsForPost(c request.CTX, postID string, userIDs []string) ([]*model.PostAcknowledgement, *model.AppError) + GetAcknowledgementsForPost(postID string) ([]*model.PostAcknowledgement, *model.AppError) + PreparePostForClient(c request.CTX, post *model.Post, isNewPost, includeDeleted, includePriority bool) *model.Post } // errNotFound allows checking against Store.ErrNotFound errors without making Store a dependency. diff --git a/server/platform/services/sharedchannel/sync_recv.go b/server/platform/services/sharedchannel/sync_recv.go index 36d943cc45..c864ecc9c6 100644 --- a/server/platform/services/sharedchannel/sync_recv.go +++ b/server/platform/services/sharedchannel/sync_recv.go @@ -82,10 +82,11 @@ func (scs *Service) processSyncMessage(c request.CTX, syncMsg *model.SyncMsg, rc var err error syncResp := model.SyncResponse{ - UserErrors: make([]string, 0), - UsersSyncd: make([]string, 0), - PostErrors: make([]string, 0), - ReactionErrors: make([]string, 0), + UserErrors: make([]string, 0), + UsersSyncd: make([]string, 0), + PostErrors: make([]string, 0), + ReactionErrors: make([]string, 0), + AcknowledgementErrors: make([]string, 0), } // Check if feature flag is enabled for membership changes @@ -103,6 +104,7 @@ func (scs *Service) processSyncMessage(c request.CTX, syncMsg *model.SyncMsg, rc mlog.Int("user_count", len(syncMsg.Users)), mlog.Int("post_count", len(syncMsg.Posts)), mlog.Int("reaction_count", len(syncMsg.Reactions)), + mlog.Int("acknowledgement_count", len(syncMsg.Acknowledgements)), mlog.Int("status_count", len(syncMsg.Statuses)), mlog.Int("membership_change_count", len(syncMsg.MembershipChanges)), ) @@ -234,6 +236,31 @@ func (scs *Service) processSyncMessage(c request.CTX, syncMsg *model.SyncMsg, rc } } + // add/remove acknowledgements + for _, acknowledgement := range syncMsg.Acknowledgements { + if _, err := scs.upsertSyncAcknowledgement(acknowledgement, targetChannel, rc); err != nil { + scs.server.Log().Log(mlog.LvlSharedChannelServiceError, "Error upserting sync acknowledgement", + mlog.String("remote", rc.Name), + mlog.String("user_id", acknowledgement.UserId), + mlog.String("post_id", acknowledgement.PostId), + mlog.Int("acknowledged_at", acknowledgement.AcknowledgedAt), + mlog.Err(err), + ) + syncResp.AcknowledgementErrors = append(syncResp.AcknowledgementErrors, acknowledgement.PostId) + } else { + scs.server.Log().Log(mlog.LvlSharedChannelServiceDebug, "Acknowledgement upserted via sync", + mlog.String("remote", rc.Name), + mlog.String("user_id", acknowledgement.UserId), + mlog.String("post_id", acknowledgement.PostId), + mlog.Int("acknowledged_at", acknowledgement.AcknowledgedAt), + ) + + if syncResp.AcknowledgementsLastUpdateAt < acknowledgement.AcknowledgedAt { + syncResp.AcknowledgementsLastUpdateAt = acknowledgement.AcknowledgedAt + } + } + } + for _, status := range syncMsg.Statuses { scs.app.SaveAndBroadcastStatus(status) } @@ -431,7 +458,6 @@ func (scs *Service) upsertSyncPost(post *model.Post, targetChannel *model.Channe post.RemoteId = model.NewPointer(rc.RemoteId) rctx := request.EmptyContext(scs.server.Log()) - rpost, err := scs.server.GetStore().Post().GetSingle(rctx, post.Id, true) if err != nil { if _, ok := err.(errNotFound); !ok { @@ -460,8 +486,7 @@ func (scs *Service) upsertSyncPost(post *model.Post, targetChannel *model.Channe if appErr == nil { scs.server.Log().Log(mlog.LvlSharedChannelServiceDebug, "Created sync post", mlog.String("post_id", post.Id), - mlog.String("channel_id", post.ChannelId), - ) + mlog.String("channel_id", post.ChannelId)) } } else if post.DeleteAt > 0 { // delete post @@ -472,14 +497,37 @@ func (scs *Service) upsertSyncPost(post *model.Post, targetChannel *model.Channe mlog.String("channel_id", post.ChannelId), ) } - } else if post.EditAt > rpost.EditAt || post.Message != rpost.Message { - // update post - rpost, appErr = scs.app.UpdatePost(request.EmptyContext(scs.server.Log()), post, nil) - if appErr == nil { - scs.server.Log().Log(mlog.LvlSharedChannelServiceDebug, "Updated sync post", - mlog.String("post_id", post.Id), - mlog.String("channel_id", post.ChannelId), - ) + } else if post.EditAt > rpost.EditAt || post.Message != rpost.Message || post.UpdateAt > rpost.UpdateAt || post.Metadata != nil { + var priority *model.PostPriority + var acknowledgements []*model.PostAcknowledgement + + if post.Metadata != nil { + // Save the received priority + if post.Metadata.Priority != nil { + priority = post.Metadata.Priority + } + + // Save the received acknowledgements + if post.Metadata.Acknowledgements != nil { + acknowledgements = post.Metadata.Acknowledgements + } + } + + // First update the basic post + rpost, appErr = scs.app.UpdatePost(rctx, post, nil) + if appErr != nil { + rerr := errors.New(appErr.Error()) + return nil, rerr + } + + // Handle priority metadata separately if needed + if priority != nil { + rpost = scs.syncRemotePriorityMetadata(rctx, post, priority, rpost) + } + + // Handle acknowledgements metadata separately if needed + if acknowledgements != nil { + rpost = scs.syncRemoteAcknowledgementsMetadata(rctx, post, acknowledgements, rpost) } } else { // nothing to update @@ -496,6 +544,105 @@ func (scs *Service) upsertSyncPost(post *model.Post, targetChannel *model.Channe return rpost, rerr } +// syncRemotePriorityMetadata handles syncing priority metadata from a remote post. +// It completely replaces existing priority settings with the ones from the remote post, +// regardless of update type. +func (scs *Service) syncRemotePriorityMetadata(rctx request.CTX, post *model.Post, priority *model.PostPriority, rpost *model.Post) *model.Post { + // First, create a new priority object with proper post and channel IDs + newPriority := &model.PostPriority{ + PostId: post.Id, + ChannelId: post.ChannelId, + } + + // Copy the priority values from the remote post + if priority.Priority != nil { + newPriority.Priority = priority.Priority + } + + if priority.RequestedAck != nil { + newPriority.RequestedAck = priority.RequestedAck + } + + if priority.PersistentNotifications != nil { + newPriority.PersistentNotifications = priority.PersistentNotifications + } + + // Save the new priority - this will replace any existing priority for the post + savedPriority, priorityErr := scs.server.GetStore().PostPriority().Save(newPriority) + if priorityErr != nil { + scs.server.Log().Log(mlog.LvlSharedChannelServiceError, "Error saving post priority from remote", + mlog.String("post_id", post.Id), + mlog.String("channel_id", post.ChannelId), + mlog.Err(priorityErr), + ) + } else { + // If the priority is successfully saved, ensure it's in the returned post + if rpost.Metadata == nil { + rpost.Metadata = &model.PostMetadata{} + } + // Use the saved priority from the database operation + rpost.Metadata.Priority = savedPriority + } + + return rpost +} + +// syncRemoteAcknowledgementsMetadata handles syncing acknowledgements metadata from a remote post. +// It replaces all existing acknowledgements with the ones from the remote post. +func (scs *Service) syncRemoteAcknowledgementsMetadata(rctx request.CTX, post *model.Post, acknowledgements []*model.PostAcknowledgement, rpost *model.Post) *model.Post { + // When syncing from remote, we completely replace the existing acknowledgements + // with the ones received from the remote, regardless of update type + + // Get existing acknowledgements and delete them using batch operation + existingAcks, appErrGet := scs.app.GetAcknowledgementsForPost(post.Id) + if appErrGet != nil { + scs.server.Log().Log(mlog.LvlSharedChannelServiceError, "Error getting existing acknowledgements for remote sync", + mlog.String("post_id", post.Id), + mlog.Err(appErrGet), + ) + } else if len(existingAcks) > 0 { + // Use batch delete for better performance + if nErr := scs.server.GetStore().PostAcknowledgement().BatchDelete(existingAcks); nErr != nil { + scs.server.Log().Log(mlog.LvlSharedChannelServiceError, "Error batch deleting acknowledgements for remote sync", + mlog.String("post_id", post.Id), + mlog.Int("count", len(existingAcks)), + mlog.Err(nErr), + ) + } + } + + // Extract all user IDs from acknowledgements for batch processing + userIDs := make([]string, 0, len(acknowledgements)) + for _, ack := range acknowledgements { + userIDs = append(userIDs, ack.UserId) + } + + // Use batch operation to save all acknowledgements at once + var savedAcks []*model.PostAcknowledgement + + if len(userIDs) > 0 { + var appErrAck *model.AppError + savedAcks, appErrAck = scs.app.SaveAcknowledgementsForPost(rctx, post.Id, userIDs) + if appErrAck != nil { + scs.server.Log().Log(mlog.LvlSharedChannelServiceError, "Error syncing remote post acknowledgements", + mlog.String("post_id", post.Id), + mlog.Int("count", len(userIDs)), + mlog.Err(appErrAck), + ) + // Fall back to original acknowledgements if batch save fails + savedAcks = acknowledgements + } + } + + // Update acknowledgements in the returned post + if rpost.Metadata == nil { + rpost.Metadata = &model.PostMetadata{} + } + rpost.Metadata.Acknowledgements = savedAcks + + return rpost +} + func (scs *Service) upsertSyncReaction(reaction *model.Reaction, targetChannel *model.Channel, rc *model.RemoteCluster) (*model.Reaction, error) { savedReaction := reaction var appErr *model.AppError @@ -542,3 +689,54 @@ func (scs *Service) upsertSyncReaction(reaction *model.Reaction, targetChannel * } return savedReaction, retErr } + +func (scs *Service) upsertSyncAcknowledgement(acknowledgement *model.PostAcknowledgement, targetChannel *model.Channel, rc *model.RemoteCluster) (*model.PostAcknowledgement, error) { + savedAcknowledgement := acknowledgement + var appErr *model.AppError + + // check that the acknowledgement's post is in the target channel. This ensures the acknowledgement can only be associated with a post + // that is in a channel shared with the remote. + rctx := request.EmptyContext(scs.server.Log()) + post, err := scs.server.GetStore().Post().GetSingle(rctx, acknowledgement.PostId, true) + if err != nil { + return nil, fmt.Errorf("error fetching post for acknowledgement sync: %w", err) + } + if post.ChannelId != targetChannel.Id { + return nil, fmt.Errorf("acknowledgement sync failed: %w", ErrChannelIDMismatch) + } + + existingAcknowledgement, err := scs.server.GetStore().PostAcknowledgement().GetSingle(acknowledgement.UserId, acknowledgement.PostId, rc.RemoteId) + if err != nil && !isNotFoundError(err) { + return nil, fmt.Errorf("error fetching acknowledgement for sync: %w", err) + } + + if existingAcknowledgement == nil { + // acknowledgement does not exist; check that user belongs to remote and create acknowledgement + // this is not done for delete since deletion can be done by admins on the remote + user, err := scs.server.GetStore().User().Get(context.TODO(), acknowledgement.UserId) + if err != nil { + return nil, fmt.Errorf("error fetching user for acknowledgement sync: %w", err) + } + if user.GetRemoteID() != rc.RemoteId { + return nil, fmt.Errorf("acknowledgement sync failed: %w", ErrRemoteIDMismatch) + } + acknowledgement.RemoteId = model.NewPointer(rc.RemoteId) + acknowledgement.ChannelId = targetChannel.Id + savedAcknowledgement, appErr = scs.app.SaveAcknowledgementForPostWithModel(request.EmptyContext(scs.server.Log()), acknowledgement) + } else { + // make sure the acknowledgement being deleted is owned by the remote + if existingAcknowledgement.GetRemoteID() != rc.RemoteId { + return nil, fmt.Errorf("acknowledgement sync failed: %w", ErrRemoteIDMismatch) + } + if acknowledgement.AcknowledgedAt == 0 { + // Delete the acknowledgement + appErr = scs.app.DeleteAcknowledgementForPostWithModel(request.EmptyContext(scs.server.Log()), acknowledgement) + } + } + + var retErr error + if appErr != nil { + retErr = errors.New(appErr.Error()) + } + return savedAcknowledgement, retErr +} diff --git a/server/platform/services/sharedchannel/sync_send.go b/server/platform/services/sharedchannel/sync_send.go index 6e856224ea..f47b009c5d 100644 --- a/server/platform/services/sharedchannel/sync_send.go +++ b/server/platform/services/sharedchannel/sync_send.go @@ -469,6 +469,9 @@ func (scs *Service) handlePostError(postId string, task syncTask, rc *model.Remo return } + // Populate metadata for the retry post + post = scs.app.PreparePostForClient(request.EmptyContext(scs.server.Log()), post, false, false, true) + syncMsg := model.NewSyncMsg(task.channelID) syncMsg.Posts = []*model.Post{post} diff --git a/server/platform/services/sharedchannel/sync_send_remote.go b/server/platform/services/sharedchannel/sync_send_remote.go index f8c9ebd865..adf2b91be3 100644 --- a/server/platform/services/sharedchannel/sync_send_remote.go +++ b/server/platform/services/sharedchannel/sync_send_remote.go @@ -32,12 +32,13 @@ type syncData struct { rc *model.RemoteCluster scr *model.SharedChannelRemote - users map[string]*model.User - profileImages map[string]*model.User - posts []*model.Post - reactions []*model.Reaction - statuses []*model.Status - attachments []attachment + users map[string]*model.User + profileImages map[string]*model.User + posts []*model.Post + reactions []*model.Reaction + acknowledgements []*model.PostAcknowledgement + statuses []*model.Status + attachments []attachment resultRepeat bool resultNextCursor model.GetPostsSinceForSyncCursor @@ -59,7 +60,7 @@ func newSyncData(task syncTask, rc *model.RemoteCluster, scr *model.SharedChanne } func (sd *syncData) isEmpty() bool { - return len(sd.users) == 0 && len(sd.profileImages) == 0 && len(sd.posts) == 0 && len(sd.reactions) == 0 && len(sd.attachments) == 0 + return len(sd.users) == 0 && len(sd.profileImages) == 0 && len(sd.posts) == 0 && len(sd.reactions) == 0 && len(sd.acknowledgements) == 0 && len(sd.attachments) == 0 } func (sd *syncData) isCursorChanged() bool { @@ -75,6 +76,7 @@ func (sd *syncData) setDataFromMsg(msg *model.SyncMsg) { sd.users = msg.Users sd.posts = msg.Posts sd.reactions = msg.Reactions + sd.acknowledgements = msg.Acknowledgements sd.statuses = msg.Statuses } @@ -185,6 +187,11 @@ func (scs *Service) syncForRemote(task syncTask, rc *model.RemoteCluster) error return fmt.Errorf("cannot fetch reactions for sync %v: %w", sd, err) } + // fetch acknowledgements for posts + if err := scs.fetchAcknowledgementsForSync(sd); err != nil { + return fmt.Errorf("cannot fetch acknowledgements for sync %v: %w", sd, err) + } + // fetch users associated with posts & reactions if err := scs.fetchPostUsersForSync(sd); err != nil { return fmt.Errorf("cannot fetch post users for sync %v: %w", sd, err) @@ -218,6 +225,7 @@ func (scs *Service) syncForRemote(task syncTask, rc *model.RemoteCluster) error mlog.Int("images", len(sd.profileImages)), mlog.Int("posts", len(sd.posts)), mlog.Int("reactions", len(sd.reactions)), + mlog.Int("acknowledgements", len(sd.acknowledgements)), mlog.Int("attachments", len(sd.attachments)), ) @@ -319,6 +327,13 @@ func (scs *Service) fetchPostsForSync(sd *syncData) error { sd.posts = appendPosts(sd.posts, posts, scs.server.GetStore().Post(), cursor.LastPostUpdateAt, scs.server.Log()) } + // Populate metadata for all posts before syncing + for i, post := range sd.posts { + if post != nil { + sd.posts[i] = scs.app.PreparePostForClient(request.EmptyContext(scs.server.Log()), post, false, false, true) + } + } + sd.resultNextCursor = nextCursor sd.resultRepeat = count >= maxPostsPerSync @@ -377,6 +392,28 @@ func (scs *Service) fetchReactionsForSync(sd *syncData) error { return merr.ErrorOrNil() } +// fetchAcknowledgementsForSync populates the sync data with any new acknowledgements since the last sync. +func (scs *Service) fetchAcknowledgementsForSync(sd *syncData) error { + start := time.Now() + defer func() { + if metrics := scs.server.GetMetrics(); metrics != nil { + metrics.ObserveSharedChannelsSyncCollectionStepDuration(sd.rc.RemoteId, "Acknowledgements", time.Since(start).Seconds()) + } + }() + + merr := merror.New() + for _, post := range sd.posts { + // any acknowledgements originating from the remote cluster are filtered out + acknowledgements, err := scs.server.GetStore().PostAcknowledgement().GetForPostSince(post.Id, sd.scr.LastPostUpdateAt, sd.rc.RemoteId, true) + if err != nil { + merr.Append(fmt.Errorf("could not get acknowledgements for post %s: %w", post.Id, err)) + continue + } + sd.acknowledgements = append(sd.acknowledgements, acknowledgements...) + } + return merr.ErrorOrNil() +} + // fetchPostUsersForSync populates the sync data with all users associated with posts. func (scs *Service) fetchPostUsersForSync(sd *syncData) error { start := time.Now() @@ -402,6 +439,10 @@ func (scs *Service) fetchPostUsersForSync(sd *syncData) error { userIDs[reaction.UserId] = p2mm{} } + for _, acknowledgement := range sd.acknowledgements { + userIDs[acknowledgement.UserId] = p2mm{} + } + for _, post := range sd.posts { // add author userIDs[post.UserId] = p2mm{} @@ -495,7 +536,9 @@ func (scs *Service) filterPostsForSync(sd *syncData) { // - new posts (EditAt == 0) // - edited posts (EditAt >= LastPostUpdateAt) // - deleted posts (DeleteAt > 0) - if p.EditAt > 0 && p.EditAt < sd.scr.LastPostUpdateAt && p.DeleteAt == 0 { + // - posts with metadata changes (acknowledgements/priority) + hasMetadataChanges := p.Metadata != nil && (p.Metadata.Acknowledgements != nil || p.Metadata.Priority != nil) + if p.EditAt > 0 && p.EditAt < sd.scr.LastPostUpdateAt && p.DeleteAt == 0 && !hasMetadataChanges { continue } @@ -514,6 +557,7 @@ func (scs *Service) filterPostsForSync(sd *syncData) { filtered = append(filtered, p) } + sd.posts = filtered } @@ -552,6 +596,13 @@ func (scs *Service) sendSyncData(sd *syncData) error { scs.updateCursorForRemote(sd.scr.Id, sd.rc, sd.resultNextCursor) } + // send acknowledgements + if len(sd.acknowledgements) != 0 { + if err := scs.sendAcknowledgementSyncData(sd); err != nil { + merr.Append(fmt.Errorf("cannot send acknowledgement sync data: %w", err)) + } + } + // send reactions if len(sd.reactions) != 0 { if err := scs.sendReactionSyncData(sd); err != nil { @@ -679,6 +730,29 @@ func (scs *Service) sendReactionSyncData(sd *syncData) error { }) } +// sendAcknowledgementSyncData sends the collected acknowledgement updates to the remote cluster. +func (scs *Service) sendAcknowledgementSyncData(sd *syncData) error { + start := time.Now() + defer func() { + if metrics := scs.server.GetMetrics(); metrics != nil { + metrics.ObserveSharedChannelsSyncSendStepDuration(sd.rc.RemoteId, "Acknowledgements", time.Since(start).Seconds()) + } + }() + + msg := model.NewSyncMsg(sd.task.channelID) + msg.Acknowledgements = sd.acknowledgements + + return scs.sendSyncMsgToRemote(msg, sd.rc, func(syncResp model.SyncResponse, errResp error) { + if len(syncResp.AcknowledgementErrors) != 0 { + scs.server.Log().Log(mlog.LvlSharedChannelServiceError, "Response indicates error for acknowledgement(s) sync", + mlog.String("channel_id", sd.task.channelID), + mlog.String("remote_id", sd.rc.RemoteId), + mlog.Array("acknowledgement_posts", syncResp.AcknowledgementErrors), + ) + } + }) +} + // sendStatusSyncData sends the collected status updates to the remote cluster. func (scs *Service) sendStatusSyncData(sd *syncData) error { msg := model.NewSyncMsg(sd.task.channelID) diff --git a/server/public/model/post_acknowledgement.go b/server/public/model/post_acknowledgement.go index 227a678e6b..3d343cbc34 100644 --- a/server/public/model/post_acknowledgement.go +++ b/server/public/model/post_acknowledgement.go @@ -6,9 +6,11 @@ package model import "net/http" type PostAcknowledgement struct { - UserId string `json:"user_id"` - PostId string `json:"post_id"` - AcknowledgedAt int64 `json:"acknowledged_at"` + UserId string `json:"user_id"` + PostId string `json:"post_id"` + AcknowledgedAt int64 `json:"acknowledged_at"` + ChannelId string `json:"channel_id"` + RemoteId *string `json:"remote_id,omitempty"` } func (o *PostAcknowledgement) IsValid() *AppError { @@ -20,5 +22,22 @@ func (o *PostAcknowledgement) IsValid() *AppError { return NewAppError("PostAcknowledgement.IsValid", "model.acknowledgement.is_valid.post_id.app_error", nil, "post_id="+o.PostId, http.StatusBadRequest) } + if !IsValidId(o.ChannelId) { + return NewAppError("PostAcknowledgement.IsValid", "model.acknowledgement.is_valid.channel_id.app_error", nil, "channel_id="+o.ChannelId, http.StatusBadRequest) + } + return nil } + +func (o *PostAcknowledgement) GetRemoteID() string { + if o.RemoteId != nil { + return *o.RemoteId + } + return "" +} + +func (o *PostAcknowledgement) PreSave() { + if o.AcknowledgedAt == 0 { + o.AcknowledgedAt = GetMillis() + } +} diff --git a/server/public/model/shared_channel.go b/server/public/model/shared_channel.go index e09b3746d1..3b7c8dc6ae 100644 --- a/server/public/model/shared_channel.go +++ b/server/public/model/shared_channel.go @@ -291,6 +291,7 @@ type SyncMsg struct { Reactions []*Reaction `json:"reactions,omitempty"` Statuses []*Status `json:"statuses,omitempty"` MembershipChanges []*MembershipChangeMsg `json:"membership_changes,omitempty"` + Acknowledgements []*PostAcknowledgement `json:"acknowledgements,omitempty"` } func NewSyncMsg(channelID string) *SyncMsg { @@ -328,6 +329,9 @@ type SyncResponse struct { ReactionsLastUpdateAt int64 `json:"reactions_last_update_at"` ReactionErrors []string `json:"reaction_errors"` + AcknowledgementsLastUpdateAt int64 `json:"acknowledgements_last_update_at"` + AcknowledgementErrors []string `json:"acknowledgement_errors"` + StatusErrors []string `json:"status_errors"` // user IDs for which the status sync failed }