fix: keep participant changes out of channel pts
(cherry picked from commit 07b2497664bd108dec84f6cfe43715540faf2688)
This commit is contained in:
parent
23a2b2aff7
commit
6fd690a06e
10 changed files with 184 additions and 115 deletions
|
|
@ -627,8 +627,7 @@ func (s *ChannelStore) EditChannelAdmin(_ context.Context, req domain.EditChanne
|
|||
})
|
||||
s.refreshChannelCountsLocked(req.ChannelID)
|
||||
channel = s.channels[req.ChannelID]
|
||||
event := s.appendParticipantEventLocked(channel, req.UserID, previous, member, req.Date)
|
||||
channel = s.channels[req.ChannelID]
|
||||
event := transientChannelParticipantEvent(channel.ID, req.UserID, previous, member, req.Date)
|
||||
if msg, ok := s.findMessageLocked(req.ChannelID, channel.TopMessageID); ok {
|
||||
s.upsertChannelDialogLocked(member.UserID, channel, msg, false)
|
||||
}
|
||||
|
|
@ -702,8 +701,7 @@ func (s *ChannelStore) EditChannelBanned(_ context.Context, req domain.EditChann
|
|||
})
|
||||
s.refreshChannelCountsLocked(req.ChannelID)
|
||||
channel = s.channels[req.ChannelID]
|
||||
event := s.appendParticipantEventLocked(channel, req.UserID, previous, member, req.Date)
|
||||
channel = s.channels[req.ChannelID]
|
||||
event := transientChannelParticipantEvent(channel.ID, req.UserID, previous, member, req.Date)
|
||||
if member.Status == domain.ChannelMemberActive {
|
||||
if msg, ok := s.findMessageLocked(req.ChannelID, channel.TopMessageID); ok {
|
||||
s.upsertChannelDialogLocked(member.UserID, channel, msg, false)
|
||||
|
|
@ -5142,25 +5140,6 @@ func (s *ChannelStore) nextChannelPtsNLocked(channelID int64, count int) int {
|
|||
return s.ptsSeq[channelID]
|
||||
}
|
||||
|
||||
func (s *ChannelStore) appendParticipantEventLocked(channel domain.Channel, actorUserID int64, previous, participant domain.ChannelMember, date int) domain.ChannelUpdateEvent {
|
||||
pts := s.nextChannelPtsLocked(channel.ID)
|
||||
channel.Pts = pts
|
||||
s.channels[channel.ID] = channel
|
||||
event := domain.ChannelUpdateEvent{
|
||||
ChannelID: channel.ID,
|
||||
Type: domain.ChannelUpdateParticipant,
|
||||
Pts: pts,
|
||||
PtsCount: 1,
|
||||
Date: date,
|
||||
SenderUserID: actorUserID,
|
||||
UserIDs: uniqueNonZeroInt64s(actorUserID, previous.UserID, previous.InviterUserID, participant.UserID, participant.InviterUserID),
|
||||
Previous: previous,
|
||||
Participant: participant,
|
||||
}
|
||||
s.events[channel.ID] = append(s.events[channel.ID], event)
|
||||
return cloneChannelEvent(event)
|
||||
}
|
||||
|
||||
func (s *ChannelStore) appendChannelServiceMessageLocked(channelID, senderUserID int64, date int, action domain.ChannelMessageAction) (domain.ChannelMessage, domain.ChannelUpdateEvent) {
|
||||
channel := s.channels[channelID]
|
||||
pts := s.nextChannelPtsLocked(channelID)
|
||||
|
|
@ -5189,6 +5168,18 @@ func (s *ChannelStore) appendChannelServiceMessageLocked(channelID, senderUserID
|
|||
return msg, event
|
||||
}
|
||||
|
||||
func transientChannelParticipantEvent(channelID, actorUserID int64, previous, participant domain.ChannelMember, date int) domain.ChannelUpdateEvent {
|
||||
return domain.ChannelUpdateEvent{
|
||||
ChannelID: channelID,
|
||||
Type: domain.ChannelUpdateParticipant,
|
||||
Date: date,
|
||||
SenderUserID: actorUserID,
|
||||
UserIDs: uniqueNonZeroInt64s(actorUserID, previous.UserID, previous.InviterUserID, participant.UserID, participant.InviterUserID),
|
||||
Previous: previous,
|
||||
Participant: participant,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *ChannelStore) channelForMemberLocked(userID, channelID int64) (domain.Channel, error) {
|
||||
channel, _, err := s.channelAndMemberLocked(userID, channelID)
|
||||
return channel, err
|
||||
|
|
|
|||
|
|
@ -39,6 +39,69 @@ func TestChannelRealtimeRecipientsAreCapped(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestChannelAdminAndBanDoNotAdvanceChannelPts(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := NewChannelStore()
|
||||
created, err := store.CreateChannel(ctx, domain.CreateChannelRequest{
|
||||
CreatorUserID: 1,
|
||||
Title: "participant state no pts",
|
||||
Megagroup: true,
|
||||
MemberUserIDs: []int64{2},
|
||||
Date: 1_700_000_120,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("create channel: %v", err)
|
||||
}
|
||||
channelID := created.Channel.ID
|
||||
ptsFloor := created.Channel.Pts
|
||||
|
||||
promoted, err := store.EditChannelAdmin(ctx, domain.EditChannelAdminRequest{
|
||||
UserID: 1,
|
||||
ChannelID: channelID,
|
||||
MemberID: 2,
|
||||
AdminRights: domain.ChannelAdminRights{
|
||||
InviteUsers: true,
|
||||
},
|
||||
Date: 1_700_000_121,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("edit admin: %v", err)
|
||||
}
|
||||
if promoted.Event.Pts != 0 || promoted.Event.PtsCount != 0 || promoted.Channel.Pts != ptsFloor {
|
||||
t.Fatalf("edit admin pts = event(%d,%d) channel %d, want unchanged %d", promoted.Event.Pts, promoted.Event.PtsCount, promoted.Channel.Pts, ptsFloor)
|
||||
}
|
||||
|
||||
banned, err := store.EditChannelBanned(ctx, domain.EditChannelBannedRequest{
|
||||
UserID: 1,
|
||||
ChannelID: channelID,
|
||||
Participant: domain.Peer{Type: domain.PeerTypeUser, ID: 2},
|
||||
BannedRights: domain.ChannelBannedRights{
|
||||
ViewMessages: true,
|
||||
UntilDate: 1_700_001_121,
|
||||
},
|
||||
Date: 1_700_000_122,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("edit banned: %v", err)
|
||||
}
|
||||
if banned.Event.Pts != 0 || banned.Event.PtsCount != 0 || banned.Channel.Pts != ptsFloor {
|
||||
t.Fatalf("edit banned pts = event(%d,%d) channel %d, want unchanged %d", banned.Event.Pts, banned.Event.PtsCount, banned.Channel.Pts, ptsFloor)
|
||||
}
|
||||
|
||||
diff, err := store.ListChannelDifference(ctx, domain.ChannelDifferenceRequest{
|
||||
UserID: 1,
|
||||
ChannelID: channelID,
|
||||
Pts: ptsFloor,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("list difference: %v", err)
|
||||
}
|
||||
if len(diff.Events) != 0 || diff.Pts != ptsFloor {
|
||||
t.Fatalf("difference after participant state change = %+v, want no durable events at pts %d", diff, ptsFloor)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPendingJoinRequestsSummaryAndInviteAdmins(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := NewChannelStore()
|
||||
|
|
|
|||
|
|
@ -798,11 +798,9 @@ func (s *ChannelStore) EditChannelAdmin(ctx context.Context, req domain.EditChan
|
|||
return domain.EditChannelAdminResult{}, fmt.Errorf("begin edit channel admin: %w", err)
|
||||
}
|
||||
committed := false
|
||||
var reserved []reservedChannelPts
|
||||
defer func() {
|
||||
if !committed {
|
||||
_ = tx.Rollback(ctx)
|
||||
s.recordChannelPtsGaps(ctx, reserved, req.Date)
|
||||
}
|
||||
}()
|
||||
channel, actor, err := s.getChannelForMember(ctx, tx, req.UserID, req.ChannelID)
|
||||
|
|
@ -873,10 +871,7 @@ func (s *ChannelStore) EditChannelAdmin(ctx context.Context, req domain.EditChan
|
|||
if err != nil {
|
||||
return domain.EditChannelAdminResult{}, err
|
||||
}
|
||||
event, channel, err := s.insertParticipantEventTx(ctx, tx, channel, req.UserID, previous, member, req.Date, &reserved)
|
||||
if err != nil {
|
||||
return domain.EditChannelAdminResult{}, err
|
||||
}
|
||||
event := transientChannelParticipantEvent(channel.ID, req.UserID, previous, member, req.Date)
|
||||
msg, _ := s.getChannelMessage(ctx, tx, req.ChannelID, channel.TopMessageID)
|
||||
if err := upsertChannelDialogTx(ctx, tx, member.UserID, channel, msg, member.ReadInboxMaxID, member.ReadOutboxMaxID); err != nil {
|
||||
return domain.EditChannelAdminResult{}, err
|
||||
|
|
@ -906,11 +901,9 @@ func (s *ChannelStore) EditChannelBanned(ctx context.Context, req domain.EditCha
|
|||
return domain.EditChannelBannedResult{}, fmt.Errorf("begin edit channel banned: %w", err)
|
||||
}
|
||||
committed := false
|
||||
var reserved []reservedChannelPts
|
||||
defer func() {
|
||||
if !committed {
|
||||
_ = tx.Rollback(ctx)
|
||||
s.recordChannelPtsGaps(ctx, reserved, req.Date)
|
||||
}
|
||||
}()
|
||||
channel, actor, err := s.getChannelForMember(ctx, tx, req.UserID, req.ChannelID)
|
||||
|
|
@ -979,10 +972,7 @@ func (s *ChannelStore) EditChannelBanned(ctx context.Context, req domain.EditCha
|
|||
if err != nil {
|
||||
return domain.EditChannelBannedResult{}, err
|
||||
}
|
||||
event, channel, err := s.insertParticipantEventTx(ctx, tx, channel, req.UserID, previous, member, req.Date, &reserved)
|
||||
if err != nil {
|
||||
return domain.EditChannelBannedResult{}, err
|
||||
}
|
||||
event := transientChannelParticipantEvent(channel.ID, req.UserID, previous, member, req.Date)
|
||||
if member.Status == domain.ChannelMemberActive {
|
||||
msg, _ := s.getChannelMessage(ctx, tx, req.ChannelID, channel.TopMessageID)
|
||||
if err := upsertChannelDialogTx(ctx, tx, member.UserID, channel, msg, member.ReadInboxMaxID, member.ReadOutboxMaxID); err != nil {
|
||||
|
|
@ -8025,31 +8015,16 @@ func (s *ChannelStore) insertServiceMessage(ctx context.Context, tx pgx.Tx, chan
|
|||
return msg, event, nil
|
||||
}
|
||||
|
||||
func (s *ChannelStore) insertParticipantEventTx(ctx context.Context, tx pgx.Tx, channel domain.Channel, actorUserID int64, previous, participant domain.ChannelMember, date int, reserved *[]reservedChannelPts) (domain.ChannelUpdateEvent, domain.Channel, error) {
|
||||
pts, err := s.pts.NextChannelPts(ctx, channel.ID)
|
||||
if err != nil {
|
||||
return domain.ChannelUpdateEvent{}, channel, fmt.Errorf("allocate channel participant pts: %w", err)
|
||||
}
|
||||
reserveChannelPts(reserved, channel.ID, pts, 1)
|
||||
event := domain.ChannelUpdateEvent{
|
||||
ChannelID: channel.ID,
|
||||
func transientChannelParticipantEvent(channelID, actorUserID int64, previous, participant domain.ChannelMember, date int) domain.ChannelUpdateEvent {
|
||||
return domain.ChannelUpdateEvent{
|
||||
ChannelID: channelID,
|
||||
Type: domain.ChannelUpdateParticipant,
|
||||
Pts: pts,
|
||||
PtsCount: 1,
|
||||
Date: date,
|
||||
SenderUserID: actorUserID,
|
||||
UserIDs: uniqueNonZeroInt64s(actorUserID, previous.UserID, previous.InviterUserID, participant.UserID, participant.InviterUserID),
|
||||
Previous: previous,
|
||||
Participant: participant,
|
||||
}
|
||||
if err := insertChannelEventTx(ctx, tx, event); err != nil {
|
||||
return domain.ChannelUpdateEvent{}, channel, err
|
||||
}
|
||||
if _, err := tx.Exec(ctx, `UPDATE channels SET pts = $2, updated_at = now() WHERE id = $1`, channel.ID, pts); err != nil {
|
||||
return domain.ChannelUpdateEvent{}, channel, fmt.Errorf("update channel participant pts: %w", err)
|
||||
}
|
||||
channel.Pts = pts
|
||||
return event, channel, nil
|
||||
}
|
||||
|
||||
func (s *ChannelStore) deleteChannelMessagesTx(ctx context.Context, tx pgx.Tx, channel domain.Channel, member domain.ChannelMember, ids []int, actorUserID int64, date int, reserved *[]reservedChannelPts) ([]int, domain.ChannelUpdateEvent, domain.Channel, error) {
|
||||
|
|
|
|||
|
|
@ -689,7 +689,8 @@ func TestChannelStoreJoinRejectsKickedMember(t *testing.T) {
|
|||
t.Fatalf("create channel: %v", err)
|
||||
}
|
||||
channelID = created.Channel.ID
|
||||
if _, err := channels.EditChannelBanned(ctx, domain.EditChannelBannedRequest{
|
||||
ptsFloor := created.Channel.Pts
|
||||
banned, err := channels.EditChannelBanned(ctx, domain.EditChannelBannedRequest{
|
||||
UserID: owner.ID,
|
||||
ChannelID: channelID,
|
||||
Participant: domain.Peer{Type: domain.PeerTypeUser, ID: member.ID},
|
||||
|
|
@ -698,9 +699,25 @@ func TestChannelStoreJoinRejectsKickedMember(t *testing.T) {
|
|||
UntilDate: 1700001300,
|
||||
},
|
||||
Date: 1700000306,
|
||||
}); err != nil {
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("kick member: %v", err)
|
||||
}
|
||||
if banned.Event.Pts != 0 || banned.Event.PtsCount != 0 || banned.Channel.Pts != ptsFloor {
|
||||
t.Fatalf("kick affected channel pts = event(%d,%d) channel %d, want no pts advance from %d", banned.Event.Pts, banned.Event.PtsCount, banned.Channel.Pts, ptsFloor)
|
||||
}
|
||||
banDiff, err := channels.ListChannelDifference(ctx, domain.ChannelDifferenceRequest{
|
||||
UserID: owner.ID,
|
||||
ChannelID: channelID,
|
||||
Pts: ptsFloor,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("difference after kick: %v", err)
|
||||
}
|
||||
if len(banDiff.Events) != 0 || banDiff.Pts != ptsFloor {
|
||||
t.Fatalf("difference after kick = %+v, want no durable participant event at pts %d", banDiff, ptsFloor)
|
||||
}
|
||||
if _, err := channels.JoinChannel(ctx, channelID, member.ID, 1700000307); !errors.Is(err, domain.ErrChannelUserBanned) {
|
||||
t.Fatalf("kicked JoinChannel err = %v, want ErrChannelUserBanned", err)
|
||||
}
|
||||
|
|
@ -1292,6 +1309,7 @@ func TestChannelStoreDifferenceStartsAtMemberAvailableMinPts(t *testing.T) {
|
|||
t.Fatalf("create channel: %v", err)
|
||||
}
|
||||
channelID = created.Channel.ID
|
||||
ptsFloor := created.Channel.Pts
|
||||
promoted, err := channels.EditChannelAdmin(ctx, domain.EditChannelAdminRequest{
|
||||
UserID: owner.ID,
|
||||
ChannelID: channelID,
|
||||
|
|
@ -1304,12 +1322,27 @@ func TestChannelStoreDifferenceStartsAtMemberAvailableMinPts(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("edit admin: %v", err)
|
||||
}
|
||||
if promoted.Event.Pts != 0 || promoted.Event.PtsCount != 0 || promoted.Channel.Pts != ptsFloor {
|
||||
t.Fatalf("promote affected channel pts = event(%d,%d) channel %d, want no pts advance from %d", promoted.Event.Pts, promoted.Event.PtsCount, promoted.Channel.Pts, ptsFloor)
|
||||
}
|
||||
adminDiff, err := channels.ListChannelDifference(ctx, domain.ChannelDifferenceRequest{
|
||||
UserID: member.ID,
|
||||
ChannelID: channelID,
|
||||
Pts: ptsFloor,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("difference after promote: %v", err)
|
||||
}
|
||||
if len(adminDiff.Events) != 0 || adminDiff.Pts != ptsFloor {
|
||||
t.Fatalf("difference after promote = %+v, want no durable participant event at pts %d", adminDiff, ptsFloor)
|
||||
}
|
||||
joined, err := channels.JoinChannel(ctx, channelID, joiner.ID, 1700000352)
|
||||
if err != nil {
|
||||
t.Fatalf("join channel: %v", err)
|
||||
}
|
||||
if len(joined.Members) != 1 || joined.Members[0].AvailableMinPts != promoted.Event.Pts {
|
||||
t.Fatalf("joined members = %+v, want available_min_pts %d", joined.Members, promoted.Event.Pts)
|
||||
if len(joined.Members) != 1 || joined.Members[0].AvailableMinPts != ptsFloor {
|
||||
t.Fatalf("joined members = %+v, want available_min_pts %d", joined.Members, ptsFloor)
|
||||
}
|
||||
diff, err := channels.ListChannelDifference(ctx, domain.ChannelDifferenceRequest{
|
||||
UserID: joiner.ID,
|
||||
|
|
@ -1324,13 +1357,13 @@ func TestChannelStoreDifferenceStartsAtMemberAvailableMinPts(t *testing.T) {
|
|||
t.Fatalf("diff pts = %d, want current channel pts %d", diff.Pts, joined.Channel.Pts)
|
||||
}
|
||||
for _, msg := range diff.NewMessages {
|
||||
if msg.Pts <= promoted.Event.Pts {
|
||||
t.Fatalf("diff leaks pre-join message %+v at or before available_min_pts %d", msg, promoted.Event.Pts)
|
||||
if msg.Pts <= ptsFloor {
|
||||
t.Fatalf("diff leaks pre-join message %+v at or before available_min_pts %d", msg, ptsFloor)
|
||||
}
|
||||
}
|
||||
for _, event := range diff.OtherUpdates {
|
||||
if event.Pts <= promoted.Event.Pts {
|
||||
t.Fatalf("diff leaks pre-join event %+v at or before available_min_pts %d", event, promoted.Event.Pts)
|
||||
if event.Pts <= ptsFloor {
|
||||
t.Fatalf("diff leaks pre-join event %+v at or before available_min_pts %d", event, ptsFloor)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue