fix: sync support private chat content protection
This commit is contained in:
parent
74c9249091
commit
037ce017d4
25 changed files with 1395 additions and 36 deletions
|
|
@ -20,6 +20,9 @@ func (s *MessageStore) ForwardPrivateMessages(ctx context.Context, req domain.Fo
|
|||
if req.Date == 0 {
|
||||
req.Date = int(time.Now().Unix())
|
||||
}
|
||||
if s.privateNoForwardsEnabled(req.OwnerUserID, req.FromPeer.ID) {
|
||||
return res, domain.ErrChatForwardsRestricted
|
||||
}
|
||||
s.mu.RLock()
|
||||
sources := make([]domain.Message, 0, len(req.MessageIDs))
|
||||
for _, id := range req.MessageIDs {
|
||||
|
|
|
|||
|
|
@ -114,10 +114,18 @@ func cloneRequestedPeerMedia(media *domain.MessageMedia) *domain.MessageMedia {
|
|||
video.Attributes = append([]domain.DocumentAttribute(nil), media.LivePhotoVideo.Attributes...)
|
||||
clone.LivePhotoVideo = &video
|
||||
}
|
||||
if media.ServiceAction == nil || media.ServiceAction.RequestedPeer == nil {
|
||||
if media.ServiceAction == nil {
|
||||
return &clone
|
||||
}
|
||||
action := *media.ServiceAction
|
||||
if media.ServiceAction.NoForwards != nil {
|
||||
noForwards := *media.ServiceAction.NoForwards
|
||||
action.NoForwards = &noForwards
|
||||
}
|
||||
if media.ServiceAction.RequestedPeer == nil {
|
||||
clone.ServiceAction = &action
|
||||
return &clone
|
||||
}
|
||||
requested := *media.ServiceAction.RequestedPeer
|
||||
requested.Peers = append([]domain.Peer(nil), requested.Peers...)
|
||||
requested.Details = append([]domain.MessageRequestedPeerDetails(nil), requested.Details...)
|
||||
|
|
|
|||
200
internal/store/memory/message_no_forwards.go
Normal file
200
internal/store/memory/message_no_forwards.go
Normal file
|
|
@ -0,0 +1,200 @@
|
|||
package memory
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
type privateNoForwardsPair struct {
|
||||
low int64
|
||||
high int64
|
||||
}
|
||||
|
||||
type memoryNoForwardsRequest struct {
|
||||
privateMessageID int64
|
||||
requesterUserID int64
|
||||
responderUserID int64
|
||||
expiresAt int
|
||||
handled bool
|
||||
}
|
||||
|
||||
func noForwardsPair(a, b int64) (privateNoForwardsPair, bool) {
|
||||
if a <= 0 || b <= 0 || a == b {
|
||||
return privateNoForwardsPair{}, false
|
||||
}
|
||||
if a > b {
|
||||
a, b = b, a
|
||||
}
|
||||
return privateNoForwardsPair{low: a, high: b}, true
|
||||
}
|
||||
|
||||
func (s *MessageStore) GetPrivateNoForwards(_ context.Context, viewerUserID, peerUserID int64) (domain.PrivateNoForwardsState, error) {
|
||||
pair, ok := noForwardsPair(viewerUserID, peerUserID)
|
||||
if !ok {
|
||||
return domain.PrivateNoForwardsState{}, domain.ErrMessageIDInvalid
|
||||
}
|
||||
s.noForwardsMu.Lock()
|
||||
defer s.noForwardsMu.Unlock()
|
||||
state := s.privateNoForwards[pair]
|
||||
state.UserLowID, state.UserHighID = pair.low, pair.high
|
||||
return state, nil
|
||||
}
|
||||
|
||||
func (s *MessageStore) TogglePrivateNoForwards(ctx context.Context, req domain.TogglePrivateNoForwardsRequest) (domain.TogglePrivateNoForwardsResult, error) {
|
||||
pair, ok := noForwardsPair(req.ActorUserID, req.PeerUserID)
|
||||
if !ok || req.RequestMsgID < 0 || req.RequestMsgID > domain.MaxMessageBoxID {
|
||||
return domain.TogglePrivateNoForwardsResult{}, domain.ErrMessageIDInvalid
|
||||
}
|
||||
if req.Date == 0 {
|
||||
req.Date = int(time.Now().Unix())
|
||||
}
|
||||
if req.RandomID == 0 {
|
||||
req.RandomID = time.Now().UnixNano()
|
||||
if req.RandomID == 0 {
|
||||
req.RandomID = 1
|
||||
}
|
||||
}
|
||||
|
||||
s.noForwardsMu.Lock()
|
||||
defer s.noForwardsMu.Unlock()
|
||||
|
||||
state := s.privateNoForwards[pair]
|
||||
state.UserLowID, state.UserHighID = pair.low, pair.high
|
||||
previousEnabled := state.Enabled()
|
||||
var (
|
||||
kind domain.MessageServiceActionKind
|
||||
action domain.MessageNoForwardsAction
|
||||
requestRecord *memoryNoForwardsRequest
|
||||
requestUID int64
|
||||
)
|
||||
|
||||
if req.RequestMsgID != 0 {
|
||||
s.mu.RLock()
|
||||
var source domain.Message
|
||||
for _, msg := range s.m[req.ActorUserID] {
|
||||
if msg.ID == req.RequestMsgID && msg.Peer == (domain.Peer{Type: domain.PeerTypeUser, ID: req.PeerUserID}) {
|
||||
source = msg
|
||||
break
|
||||
}
|
||||
}
|
||||
if source.ID != 0 {
|
||||
record := s.privateNoForwardsRequests[source.UID]
|
||||
requestRecord = &record
|
||||
requestUID = source.UID
|
||||
}
|
||||
s.mu.RUnlock()
|
||||
if source.ID == 0 || requestRecord == nil || requestRecord.privateMessageID != source.UID ||
|
||||
requestRecord.requesterUserID != req.PeerUserID || requestRecord.responderUserID != req.ActorUserID ||
|
||||
requestRecord.handled || requestRecord.expiresAt <= req.Date {
|
||||
return domain.TogglePrivateNoForwardsResult{}, domain.ErrNoForwardsRequestExpired
|
||||
}
|
||||
kind = domain.MessageServiceActionNoForwardsToggle
|
||||
action = domain.MessageNoForwardsAction{PrevValue: previousEnabled, NewValue: req.Enabled}
|
||||
if req.Enabled {
|
||||
state.EnabledByUserID = req.ActorUserID
|
||||
} else {
|
||||
state.EnabledByUserID = 0
|
||||
}
|
||||
} else if req.Enabled {
|
||||
if state.EnabledByUserID != 0 {
|
||||
return domain.TogglePrivateNoForwardsResult{State: state}, nil
|
||||
}
|
||||
kind = domain.MessageServiceActionNoForwardsToggle
|
||||
action = domain.MessageNoForwardsAction{PrevValue: false, NewValue: true}
|
||||
state.EnabledByUserID = req.ActorUserID
|
||||
} else {
|
||||
switch state.EnabledByUserID {
|
||||
case 0:
|
||||
return domain.TogglePrivateNoForwardsResult{State: state}, nil
|
||||
case req.ActorUserID:
|
||||
kind = domain.MessageServiceActionNoForwardsToggle
|
||||
action = domain.MessageNoForwardsAction{PrevValue: true, NewValue: false}
|
||||
state.EnabledByUserID = 0
|
||||
default:
|
||||
kind = domain.MessageServiceActionNoForwardsRequest
|
||||
action = domain.MessageNoForwardsAction{
|
||||
PrevValue: true,
|
||||
NewValue: false,
|
||||
ExpiresAt: req.Date + domain.PrivateNoForwardsRequestExpirePeriod,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
reply := (*domain.MessageReply)(nil)
|
||||
if req.RequestMsgID != 0 {
|
||||
reply = &domain.MessageReply{
|
||||
MessageID: req.RequestMsgID,
|
||||
Peer: domain.Peer{Type: domain.PeerTypeUser, ID: req.PeerUserID},
|
||||
}
|
||||
}
|
||||
send, err := s.SendPrivateText(ctx, domain.SendPrivateTextRequest{
|
||||
SenderUserID: req.ActorUserID,
|
||||
RecipientUserID: req.PeerUserID,
|
||||
RandomID: req.RandomID,
|
||||
Silent: true,
|
||||
Date: req.Date,
|
||||
OriginAuthKeyID: req.OriginAuthKeyID,
|
||||
OriginSessionID: req.OriginSessionID,
|
||||
ReplyTo: reply,
|
||||
Media: &domain.MessageMedia{
|
||||
Kind: domain.MessageMediaKindService,
|
||||
ServiceAction: &domain.MessageServiceAction{
|
||||
Kind: kind,
|
||||
NoForwards: &action,
|
||||
},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
if req.RequestMsgID != 0 && err == domain.ErrReplyMessageIDInvalid {
|
||||
return domain.TogglePrivateNoForwardsResult{}, domain.ErrNoForwardsRequestExpired
|
||||
}
|
||||
return domain.TogglePrivateNoForwardsResult{}, err
|
||||
}
|
||||
|
||||
s.privateNoForwards[pair] = state
|
||||
if kind == domain.MessageServiceActionNoForwardsRequest {
|
||||
s.privateNoForwardsRequests[send.SenderMessage.UID] = memoryNoForwardsRequest{
|
||||
privateMessageID: send.SenderMessage.UID,
|
||||
requesterUserID: req.ActorUserID,
|
||||
responderUserID: req.PeerUserID,
|
||||
expiresAt: action.ExpiresAt,
|
||||
}
|
||||
}
|
||||
if requestUID != 0 {
|
||||
record := s.privateNoForwardsRequests[requestUID]
|
||||
record.handled = true
|
||||
s.privateNoForwardsRequests[requestUID] = record
|
||||
s.markNoForwardsRequestExpired(requestUID)
|
||||
}
|
||||
return domain.TogglePrivateNoForwardsResult{State: state, Changed: true, Send: send}, nil
|
||||
}
|
||||
|
||||
func (s *MessageStore) markNoForwardsRequestExpired(privateMessageID int64) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
for ownerID, messages := range s.m {
|
||||
for i := range messages {
|
||||
action := messages[i].Media
|
||||
if messages[i].UID != privateMessageID || action == nil || action.ServiceAction == nil ||
|
||||
action.ServiceAction.Kind != domain.MessageServiceActionNoForwardsRequest ||
|
||||
action.ServiceAction.NoForwards == nil {
|
||||
continue
|
||||
}
|
||||
messages[i].Media = cloneRequestedPeerMedia(messages[i].Media)
|
||||
messages[i].Media.ServiceAction.NoForwards.Expired = true
|
||||
}
|
||||
s.m[ownerID] = messages
|
||||
}
|
||||
}
|
||||
|
||||
func (s *MessageStore) privateNoForwardsEnabled(a, b int64) bool {
|
||||
pair, ok := noForwardsPair(a, b)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
s.noForwardsMu.Lock()
|
||||
defer s.noForwardsMu.Unlock()
|
||||
return s.privateNoForwards[pair].Enabled()
|
||||
}
|
||||
162
internal/store/memory/message_no_forwards_test.go
Normal file
162
internal/store/memory/message_no_forwards_test.go
Normal file
|
|
@ -0,0 +1,162 @@
|
|||
package memory
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
func TestPrivateNoForwardsStateMachineAndForwardGate(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
messages := NewMessageStore()
|
||||
const alice, bob int64 = 1001, 1002
|
||||
|
||||
enable, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: alice, PeerUserID: bob, Enabled: true, RandomID: 11, Date: 100,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("enable: %v", err)
|
||||
}
|
||||
if !enable.Changed || enable.State.EnabledByUserID != alice ||
|
||||
enable.Send.SenderMessage.Pts != 1 || enable.Send.RecipientMessage.Pts != 1 ||
|
||||
enable.Send.SenderMessage.NoForwards {
|
||||
t.Fatalf("enable result = %+v", enable)
|
||||
}
|
||||
assertMemoryNoForwardsAction(t, enable.Send.SenderMessage, domain.MessageServiceActionNoForwardsToggle, false, true, false)
|
||||
|
||||
repeat, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: alice, PeerUserID: bob, Enabled: true, RandomID: 12, Date: 101,
|
||||
})
|
||||
if err != nil || repeat.Changed || repeat.State.EnabledByUserID != alice {
|
||||
t.Fatalf("repeat enable = %+v err=%v, want no-op", repeat, err)
|
||||
}
|
||||
|
||||
request, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: bob, PeerUserID: alice, Enabled: false, RandomID: 13, Date: 102,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("request disable: %v", err)
|
||||
}
|
||||
if request.State.EnabledByUserID != alice || request.Send.SenderMessage.Pts != 2 ||
|
||||
request.Send.RecipientMessage.Pts != 2 {
|
||||
t.Fatalf("request result = %+v", request)
|
||||
}
|
||||
assertMemoryNoForwardsAction(t, request.Send.SenderMessage, domain.MessageServiceActionNoForwardsRequest, true, false, false)
|
||||
|
||||
answer, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: alice,
|
||||
PeerUserID: bob,
|
||||
Enabled: false,
|
||||
RequestMsgID: request.Send.RecipientMessage.ID,
|
||||
RandomID: 14,
|
||||
Date: 103,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("accept request: %v", err)
|
||||
}
|
||||
if answer.State.Enabled() || answer.Send.SenderMessage.Pts != 3 || answer.Send.RecipientMessage.Pts != 3 {
|
||||
t.Fatalf("answer result = %+v", answer)
|
||||
}
|
||||
if answer.Send.SenderMessage.ReplyTo == nil ||
|
||||
answer.Send.SenderMessage.ReplyTo.MessageID != request.Send.RecipientMessage.ID ||
|
||||
answer.Send.RecipientMessage.ReplyTo == nil ||
|
||||
answer.Send.RecipientMessage.ReplyTo.MessageID != request.Send.SenderMessage.ID {
|
||||
t.Fatalf("answer reply mapping sender=%+v recipient=%+v", answer.Send.SenderMessage.ReplyTo, answer.Send.RecipientMessage.ReplyTo)
|
||||
}
|
||||
assertMemoryNoForwardsAction(t, answer.Send.SenderMessage, domain.MessageServiceActionNoForwardsToggle, true, false, false)
|
||||
|
||||
aliceHistory, err := messages.ListByUser(ctx, alice, domain.MessageFilter{
|
||||
HasPeer: true, Peer: domain.Peer{Type: domain.PeerTypeUser, ID: bob}, Limit: 20,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("alice history: %v", err)
|
||||
}
|
||||
var expired bool
|
||||
for _, msg := range aliceHistory.Messages {
|
||||
if msg.ID == request.Send.RecipientMessage.ID && msg.Media != nil && msg.Media.ServiceAction != nil &&
|
||||
msg.Media.ServiceAction.NoForwards != nil {
|
||||
expired = msg.Media.ServiceAction.NoForwards.Expired
|
||||
}
|
||||
}
|
||||
if !expired {
|
||||
t.Fatal("handled request was not projected expired")
|
||||
}
|
||||
if _, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: alice, PeerUserID: bob, RequestMsgID: request.Send.RecipientMessage.ID,
|
||||
RandomID: 15, Date: 104,
|
||||
}); !errors.Is(err, domain.ErrNoForwardsRequestExpired) {
|
||||
t.Fatalf("repeat answer err=%v, want ErrNoForwardsRequestExpired", err)
|
||||
}
|
||||
|
||||
source, err := messages.SendPrivateText(ctx, domain.SendPrivateTextRequest{
|
||||
SenderUserID: alice, RecipientUserID: bob, RandomID: 20, Message: "source", Date: 105,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("send source: %v", err)
|
||||
}
|
||||
if _, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: alice, PeerUserID: bob, Enabled: true, RandomID: 21, Date: 106,
|
||||
}); err != nil {
|
||||
t.Fatalf("re-enable: %v", err)
|
||||
}
|
||||
if _, err := messages.ForwardPrivateMessages(ctx, domain.ForwardPrivateMessagesRequest{
|
||||
OwnerUserID: alice,
|
||||
FromPeer: domain.Peer{Type: domain.PeerTypeUser, ID: bob},
|
||||
ToUserID: alice,
|
||||
MessageIDs: []int{source.SenderMessage.ID},
|
||||
RandomIDs: []int64{22},
|
||||
Date: 107,
|
||||
}); !errors.Is(err, domain.ErrChatForwardsRestricted) {
|
||||
t.Fatalf("forward protected chat err=%v, want ErrChatForwardsRestricted", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrivateNoForwardsRequestExpiresWithoutPTS(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
messages := NewMessageStore()
|
||||
const alice, bob int64 = 2001, 2002
|
||||
if _, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: alice, PeerUserID: bob, Enabled: true, RandomID: 31, Date: 200,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
request, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: bob, PeerUserID: alice, RandomID: 32, Date: 201,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: alice,
|
||||
PeerUserID: bob,
|
||||
RequestMsgID: request.Send.RecipientMessage.ID,
|
||||
RandomID: 33,
|
||||
Date: 201 + domain.PrivateNoForwardsRequestExpirePeriod,
|
||||
}); !errors.Is(err, domain.ErrNoForwardsRequestExpired) {
|
||||
t.Fatalf("expired answer err=%v", err)
|
||||
}
|
||||
state, _ := messages.GetPrivateNoForwards(ctx, alice, bob)
|
||||
if state.EnabledByUserID != alice {
|
||||
t.Fatalf("expired answer changed state = %+v", state)
|
||||
}
|
||||
history, _ := messages.ListByUser(ctx, alice, domain.MessageFilter{
|
||||
HasPeer: true, Peer: domain.Peer{Type: domain.PeerTypeUser, ID: bob}, Limit: 20,
|
||||
})
|
||||
if len(history.Messages) != 2 || history.Messages[0].Pts != 2 {
|
||||
t.Fatalf("expired answer allocated message/pts: %+v", history.Messages)
|
||||
}
|
||||
}
|
||||
|
||||
func assertMemoryNoForwardsAction(t *testing.T, msg domain.Message, kind domain.MessageServiceActionKind, prev, next, expired bool) {
|
||||
t.Helper()
|
||||
if msg.Media == nil || msg.Media.ServiceAction == nil || msg.Media.ServiceAction.Kind != kind ||
|
||||
msg.Media.ServiceAction.NoForwards == nil {
|
||||
t.Fatalf("message action = %+v, want %s", msg.Media, kind)
|
||||
}
|
||||
action := msg.Media.ServiceAction.NoForwards
|
||||
if action.PrevValue != prev || action.NewValue != next || action.Expired != expired {
|
||||
t.Fatalf("action = %+v, want prev=%v new=%v expired=%v", action, prev, next, expired)
|
||||
}
|
||||
}
|
||||
|
|
@ -8,6 +8,7 @@ import (
|
|||
// MessageStore 是 store.MessageStore 的内存实现。
|
||||
type MessageStore struct {
|
||||
mu sync.RWMutex
|
||||
noForwardsMu sync.Mutex
|
||||
m map[int64][]domain.Message
|
||||
nextUID int64
|
||||
nextBox map[int64]int
|
||||
|
|
@ -24,6 +25,11 @@ type MessageStore struct {
|
|||
polls *PollStore
|
||||
// savedPins 是收藏夹子会话置顶顺序(下标即 pinned_order,越小越前)。
|
||||
savedPins map[int64][]domain.Peer
|
||||
// privateNoForwards is keyed by the sorted user pair. Requests are keyed by
|
||||
// the shared logical private-message id so both local box ids resolve to one
|
||||
// one-shot response fact.
|
||||
privateNoForwards map[privateNoForwardsPair]domain.PrivateNoForwardsState
|
||||
privateNoForwardsRequests map[int64]memoryNoForwardsRequest
|
||||
}
|
||||
|
||||
// AttachPollStore 注入共享 poll 权威(与 ChannelStore 共用同一实例)。
|
||||
|
|
@ -40,18 +46,20 @@ type readOutboxDateKey struct {
|
|||
// NewMessageStore 创建内存 MessageStore。
|
||||
func NewMessageStore(dialogs ...*DialogStore) *MessageStore {
|
||||
s := &MessageStore{
|
||||
m: make(map[int64][]domain.Message),
|
||||
nextUID: 1,
|
||||
nextBox: make(map[int64]int),
|
||||
nextPts: make(map[int64]int),
|
||||
readOutboxDates: make(map[readOutboxDateKey]int),
|
||||
privateReactions: make(map[int64]map[int64][]domain.ChannelMessagePeerReaction),
|
||||
savedMessageTags: make(map[int64]map[int][]domain.MessageReaction),
|
||||
savedTagTitles: make(map[int64]map[string]string),
|
||||
privateSendDedup: make(map[privateSendDedupKey]privateSendDedupRecord),
|
||||
loginCodeDeliveries: make(map[[32]byte]loginCodeDeliveryRecord),
|
||||
albumGroups: make(map[albumGroupKey]albumGroupRecord),
|
||||
savedPins: make(map[int64][]domain.Peer),
|
||||
m: make(map[int64][]domain.Message),
|
||||
nextUID: 1,
|
||||
nextBox: make(map[int64]int),
|
||||
nextPts: make(map[int64]int),
|
||||
readOutboxDates: make(map[readOutboxDateKey]int),
|
||||
privateReactions: make(map[int64]map[int64][]domain.ChannelMessagePeerReaction),
|
||||
savedMessageTags: make(map[int64]map[int][]domain.MessageReaction),
|
||||
savedTagTitles: make(map[int64]map[string]string),
|
||||
privateSendDedup: make(map[privateSendDedupKey]privateSendDedupRecord),
|
||||
loginCodeDeliveries: make(map[[32]byte]loginCodeDeliveryRecord),
|
||||
albumGroups: make(map[albumGroupKey]albumGroupRecord),
|
||||
savedPins: make(map[int64][]domain.Peer),
|
||||
privateNoForwards: make(map[privateNoForwardsPair]domain.PrivateNoForwardsState),
|
||||
privateNoForwardsRequests: make(map[int64]memoryNoForwardsRequest),
|
||||
}
|
||||
if len(dialogs) > 0 {
|
||||
s.dialogs = dialogs[0]
|
||||
|
|
|
|||
|
|
@ -33,6 +33,13 @@ func (s *MessageStore) ForwardPrivateMessages(ctx context.Context, req domain.Fo
|
|||
if req.Date == 0 {
|
||||
req.Date = int(time.Now().Unix())
|
||||
}
|
||||
protected, err := s.privateNoForwardsEnabled(ctx, req.OwnerUserID, req.FromPeer.ID)
|
||||
if err != nil {
|
||||
return res, err
|
||||
}
|
||||
if protected {
|
||||
return res, domain.ErrChatForwardsRestricted
|
||||
}
|
||||
boxIDs := make([]int32, 0, len(req.MessageIDs))
|
||||
for i, id := range req.MessageIDs {
|
||||
if id <= 0 || id > domain.MaxMessageBoxID || req.RandomIDs[i] == 0 {
|
||||
|
|
|
|||
248
internal/store/postgres/message_no_forwards.go
Normal file
248
internal/store/postgres/message_no_forwards.go
Normal file
|
|
@ -0,0 +1,248 @@
|
|||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
var errPrivateNoForwardsNoop = errors.New("private no forwards no-op")
|
||||
|
||||
func pgNoForwardsPair(a, b int64) (low, high int64, ok bool) {
|
||||
if a <= 0 || b <= 0 || a == b {
|
||||
return 0, 0, false
|
||||
}
|
||||
if a > b {
|
||||
a, b = b, a
|
||||
}
|
||||
return a, b, true
|
||||
}
|
||||
|
||||
func (s *MessageStore) GetPrivateNoForwards(ctx context.Context, viewerUserID, peerUserID int64) (domain.PrivateNoForwardsState, error) {
|
||||
low, high, ok := pgNoForwardsPair(viewerUserID, peerUserID)
|
||||
if !ok {
|
||||
return domain.PrivateNoForwardsState{}, domain.ErrMessageIDInvalid
|
||||
}
|
||||
state := domain.PrivateNoForwardsState{UserLowID: low, UserHighID: high}
|
||||
err := s.db.QueryRow(ctx, `
|
||||
SELECT COALESCE(enabled_by_user_id, 0)
|
||||
FROM private_no_forwards_chats
|
||||
WHERE user_low_id = $1 AND user_high_id = $2`, low, high).Scan(&state.EnabledByUserID)
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return state, nil
|
||||
}
|
||||
if err != nil {
|
||||
return domain.PrivateNoForwardsState{}, fmt.Errorf("get private no forwards: %w", err)
|
||||
}
|
||||
return state, nil
|
||||
}
|
||||
|
||||
func (s *MessageStore) TogglePrivateNoForwards(ctx context.Context, req domain.TogglePrivateNoForwardsRequest) (domain.TogglePrivateNoForwardsResult, error) {
|
||||
low, high, ok := pgNoForwardsPair(req.ActorUserID, req.PeerUserID)
|
||||
if !ok || req.RequestMsgID < 0 || req.RequestMsgID > domain.MaxMessageBoxID {
|
||||
return domain.TogglePrivateNoForwardsResult{}, domain.ErrMessageIDInvalid
|
||||
}
|
||||
if req.Date == 0 {
|
||||
req.Date = int(time.Now().Unix())
|
||||
}
|
||||
if req.RandomID == 0 {
|
||||
req.RandomID = time.Now().UnixNano()
|
||||
if req.RandomID == 0 {
|
||||
req.RandomID = 1
|
||||
}
|
||||
}
|
||||
|
||||
state := domain.PrivateNoForwardsState{UserLowID: low, UserHighID: high}
|
||||
actionKind := domain.MessageServiceActionNoForwardsToggle
|
||||
action := domain.MessageNoForwardsAction{}
|
||||
var answeredRequestSenderID, answeredRequestMessageID int64
|
||||
|
||||
sendReq := domain.SendPrivateTextRequest{
|
||||
SenderUserID: req.ActorUserID,
|
||||
RecipientUserID: req.PeerUserID,
|
||||
RandomID: req.RandomID,
|
||||
Silent: true,
|
||||
Date: req.Date,
|
||||
OriginAuthKeyID: req.OriginAuthKeyID,
|
||||
OriginSessionID: req.OriginSessionID,
|
||||
// A non-empty placeholder is required before the send transaction starts.
|
||||
// The pair-locked before-hook replaces it with the authoritative action.
|
||||
Media: noForwardsServiceMedia(actionKind, action),
|
||||
}
|
||||
if req.RequestMsgID != 0 {
|
||||
sendReq.ReplyTo = &domain.MessageReply{
|
||||
MessageID: req.RequestMsgID,
|
||||
Peer: domain.Peer{Type: domain.PeerTypeUser, ID: req.PeerUserID},
|
||||
}
|
||||
}
|
||||
|
||||
hooks := privateSendTxHooks{
|
||||
before: func(ctx context.Context, tx pgx.Tx, send *domain.SendPrivateTextRequest) error {
|
||||
if _, err := tx.Exec(ctx, `
|
||||
INSERT INTO private_no_forwards_chats (user_low_id, user_high_id)
|
||||
VALUES ($1, $2)
|
||||
ON CONFLICT (user_low_id, user_high_id) DO NOTHING`, low, high); err != nil {
|
||||
return fmt.Errorf("ensure private no forwards state: %w", err)
|
||||
}
|
||||
if err := tx.QueryRow(ctx, `
|
||||
SELECT COALESCE(enabled_by_user_id, 0)
|
||||
FROM private_no_forwards_chats
|
||||
WHERE user_low_id = $1 AND user_high_id = $2
|
||||
FOR UPDATE`, low, high).Scan(&state.EnabledByUserID); err != nil {
|
||||
return fmt.Errorf("lock private no forwards state: %w", err)
|
||||
}
|
||||
|
||||
previousEnabled := state.Enabled()
|
||||
actionKind = domain.MessageServiceActionNoForwardsToggle
|
||||
action = domain.MessageNoForwardsAction{}
|
||||
if req.RequestMsgID != 0 {
|
||||
var expiresAt, handledAt int
|
||||
err := tx.QueryRow(ctx, `
|
||||
SELECT r.private_message_sender_user_id, r.private_message_id, r.expires_at, r.handled_at
|
||||
FROM message_boxes AS b
|
||||
JOIN private_no_forwards_requests AS r
|
||||
ON r.private_message_sender_user_id = b.message_sender_id
|
||||
AND r.private_message_id = b.private_message_id
|
||||
WHERE b.owner_user_id = $1
|
||||
AND b.box_id = $2
|
||||
AND b.peer_type = 'user'
|
||||
AND b.peer_id = $3
|
||||
AND r.requester_user_id = $3
|
||||
AND r.responder_user_id = $1
|
||||
FOR UPDATE OF r`,
|
||||
req.ActorUserID, req.RequestMsgID, req.PeerUserID,
|
||||
).Scan(&answeredRequestSenderID, &answeredRequestMessageID, &expiresAt, &handledAt)
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return domain.ErrNoForwardsRequestExpired
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("lock private no forwards request: %w", err)
|
||||
}
|
||||
if handledAt != 0 || expiresAt <= req.Date {
|
||||
return domain.ErrNoForwardsRequestExpired
|
||||
}
|
||||
action = domain.MessageNoForwardsAction{PrevValue: previousEnabled, NewValue: req.Enabled}
|
||||
if req.Enabled {
|
||||
state.EnabledByUserID = req.ActorUserID
|
||||
} else {
|
||||
state.EnabledByUserID = 0
|
||||
}
|
||||
if _, err := tx.Exec(ctx, `
|
||||
UPDATE private_no_forwards_requests
|
||||
SET handled_at = $3
|
||||
WHERE private_message_sender_user_id = $1
|
||||
AND private_message_id = $2
|
||||
AND handled_at = 0`,
|
||||
answeredRequestSenderID, answeredRequestMessageID, req.Date,
|
||||
); err != nil {
|
||||
return fmt.Errorf("handle private no forwards request: %w", err)
|
||||
}
|
||||
if err := expirePGNoForwardsRequest(ctx, tx, answeredRequestSenderID, answeredRequestMessageID); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if req.Enabled {
|
||||
if state.EnabledByUserID != 0 {
|
||||
return errPrivateNoForwardsNoop
|
||||
}
|
||||
action = domain.MessageNoForwardsAction{PrevValue: false, NewValue: true}
|
||||
state.EnabledByUserID = req.ActorUserID
|
||||
} else {
|
||||
switch state.EnabledByUserID {
|
||||
case 0:
|
||||
return errPrivateNoForwardsNoop
|
||||
case req.ActorUserID:
|
||||
action = domain.MessageNoForwardsAction{PrevValue: true, NewValue: false}
|
||||
state.EnabledByUserID = 0
|
||||
default:
|
||||
actionKind = domain.MessageServiceActionNoForwardsRequest
|
||||
action = domain.MessageNoForwardsAction{
|
||||
PrevValue: true,
|
||||
NewValue: false,
|
||||
ExpiresAt: req.Date + domain.PrivateNoForwardsRequestExpirePeriod,
|
||||
}
|
||||
}
|
||||
}
|
||||
var enabledBy any
|
||||
if state.EnabledByUserID != 0 {
|
||||
enabledBy = state.EnabledByUserID
|
||||
}
|
||||
if _, err := tx.Exec(ctx, `
|
||||
UPDATE private_no_forwards_chats
|
||||
SET enabled_by_user_id = $3, updated_at = now()
|
||||
WHERE user_low_id = $1 AND user_high_id = $2`, low, high, enabledBy); err != nil {
|
||||
return fmt.Errorf("update private no forwards state: %w", err)
|
||||
}
|
||||
send.Media = noForwardsServiceMedia(actionKind, action)
|
||||
return nil
|
||||
},
|
||||
after: func(ctx context.Context, tx pgx.Tx, result domain.SendPrivateTextResult) error {
|
||||
if actionKind != domain.MessageServiceActionNoForwardsRequest {
|
||||
return nil
|
||||
}
|
||||
if _, err := tx.Exec(ctx, `
|
||||
INSERT INTO private_no_forwards_requests (
|
||||
private_message_sender_user_id,
|
||||
private_message_id,
|
||||
requester_user_id,
|
||||
responder_user_id,
|
||||
expires_at
|
||||
) VALUES ($1, $2, $3, $4, $5)`,
|
||||
req.ActorUserID, result.SenderMessage.UID, req.ActorUserID, req.PeerUserID, action.ExpiresAt,
|
||||
); err != nil {
|
||||
return fmt.Errorf("create private no forwards request: %w", err)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
send, err := s.sendPrivateTextWithHooks(ctx, sendReq, hooks)
|
||||
if errors.Is(err, errPrivateNoForwardsNoop) {
|
||||
return domain.TogglePrivateNoForwardsResult{State: state}, nil
|
||||
}
|
||||
if errors.Is(err, domain.ErrReplyMessageIDInvalid) {
|
||||
return domain.TogglePrivateNoForwardsResult{}, domain.ErrNoForwardsRequestExpired
|
||||
}
|
||||
if err != nil {
|
||||
return domain.TogglePrivateNoForwardsResult{}, err
|
||||
}
|
||||
return domain.TogglePrivateNoForwardsResult{State: state, Changed: true, Send: send}, nil
|
||||
}
|
||||
|
||||
func noForwardsServiceMedia(kind domain.MessageServiceActionKind, action domain.MessageNoForwardsAction) *domain.MessageMedia {
|
||||
return &domain.MessageMedia{
|
||||
Kind: domain.MessageMediaKindService,
|
||||
ServiceAction: &domain.MessageServiceAction{
|
||||
Kind: kind,
|
||||
NoForwards: &action,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func expirePGNoForwardsRequest(ctx context.Context, tx pgx.Tx, senderUserID, privateMessageID int64) error {
|
||||
for _, statement := range []string{
|
||||
`UPDATE private_messages
|
||||
SET media = jsonb_set(media, '{service_action,no_forwards,expired}', 'true'::jsonb, true)
|
||||
WHERE sender_user_id = $1 AND id = $2`,
|
||||
`UPDATE message_boxes
|
||||
SET media = jsonb_set(media, '{service_action,no_forwards,expired}', 'true'::jsonb, true)
|
||||
WHERE message_sender_id = $1 AND private_message_id = $2`,
|
||||
} {
|
||||
tag, err := tx.Exec(ctx, statement, senderUserID, privateMessageID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("expire private no forwards request: %w", err)
|
||||
}
|
||||
if tag.RowsAffected() == 0 {
|
||||
return fmt.Errorf("expire private no forwards request: message disappeared")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *MessageStore) privateNoForwardsEnabled(ctx context.Context, a, b int64) (bool, error) {
|
||||
state, err := s.GetPrivateNoForwards(ctx, a, b)
|
||||
return state.Enabled(), err
|
||||
}
|
||||
229
internal/store/postgres/message_no_forwards_integration_test.go
Normal file
229
internal/store/postgres/message_no_forwards_integration_test.go
Normal file
|
|
@ -0,0 +1,229 @@
|
|||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
func TestPostgresPrivateNoForwardsAtomicStateAndDifference(t *testing.T) {
|
||||
pool := testPool(t)
|
||||
ctx := context.Background()
|
||||
suffix := randomSuffix(t)
|
||||
users := NewUserStore(pool)
|
||||
alice, err := users.Create(ctx, domain.User{AccessHash: 6101, Phone: "+1668" + suffix + "01", FirstName: "Alice"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
bob, err := users.Create(ctx, domain.User{AccessHash: 6102, Phone: "+1668" + suffix + "02", FirstName: "Bob"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_, _ = pool.Exec(ctx, "DELETE FROM users WHERE id = ANY($1::bigint[])", []int64{alice.ID, bob.ID})
|
||||
})
|
||||
|
||||
messages := NewMessageStore(pool)
|
||||
baseRandom := time.Now().UnixNano()
|
||||
enable, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: alice.ID, PeerUserID: bob.ID, Enabled: true, RandomID: baseRandom, Date: 1700100000,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("enable: %v", err)
|
||||
}
|
||||
request, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: bob.ID, PeerUserID: alice.ID, RandomID: baseRandom + 1, Date: 1700100001,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("request: %v", err)
|
||||
}
|
||||
answer, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: alice.ID, PeerUserID: bob.ID, RequestMsgID: request.Send.RecipientMessage.ID,
|
||||
RandomID: baseRandom + 2, Date: 1700100002,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("answer: %v", err)
|
||||
}
|
||||
if enable.Send.SenderMessage.Pts != 1 || request.Send.SenderMessage.Pts != 2 ||
|
||||
answer.Send.SenderMessage.Pts != 3 || answer.Send.SenderMessage.ReplyTo == nil ||
|
||||
answer.Send.SenderMessage.ReplyTo.MessageID != request.Send.RecipientMessage.ID ||
|
||||
answer.Send.RecipientMessage.ReplyTo == nil ||
|
||||
answer.Send.RecipientMessage.ReplyTo.MessageID != request.Send.SenderMessage.ID {
|
||||
t.Fatalf("pts/reply mapping enable=%+v request=%+v answer=%+v", enable.Send, request.Send, answer.Send)
|
||||
}
|
||||
state, err := messages.GetPrivateNoForwards(ctx, alice.ID, bob.ID)
|
||||
if err != nil || state.Enabled() {
|
||||
t.Fatalf("final state=%+v err=%v, want disabled", state, err)
|
||||
}
|
||||
if _, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: alice.ID, PeerUserID: bob.ID, RequestMsgID: request.Send.RecipientMessage.ID,
|
||||
RandomID: baseRandom + 3, Date: 1700100003,
|
||||
}); !errors.Is(err, domain.ErrNoForwardsRequestExpired) {
|
||||
t.Fatalf("repeat answer err=%v", err)
|
||||
}
|
||||
|
||||
for _, userID := range []int64{alice.ID, bob.ID} {
|
||||
events, err := NewUpdateEventStore(pool).ListAfter(ctx, userID, 0, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("events user %d: %v", userID, err)
|
||||
}
|
||||
if len(events) != 3 || events[0].Pts != 1 || events[1].Pts != 2 || events[2].Pts != 3 {
|
||||
t.Fatalf("events user %d = %+v, want continuous 1..3", userID, events)
|
||||
}
|
||||
}
|
||||
var eventCount, outboxCount int
|
||||
if err := pool.QueryRow(ctx, `
|
||||
SELECT
|
||||
(SELECT count(*) FROM user_update_events WHERE user_id = ANY($1::bigint[])),
|
||||
(SELECT count(*) FROM dispatch_outbox WHERE target_user_id = ANY($1::bigint[]))`,
|
||||
[]int64{alice.ID, bob.ID},
|
||||
).Scan(&eventCount, &outboxCount); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if eventCount != 6 || outboxCount != 6 {
|
||||
t.Fatalf("event/outbox count=%d/%d, want 6/6", eventCount, outboxCount)
|
||||
}
|
||||
var handledAt int
|
||||
var logicalExpired bool
|
||||
var expiredBoxes int
|
||||
if err := pool.QueryRow(ctx, `
|
||||
SELECT r.handled_at,
|
||||
COALESCE((pm.media #>> '{service_action,no_forwards,expired}')::boolean, false),
|
||||
(SELECT count(*)
|
||||
FROM message_boxes b
|
||||
WHERE b.message_sender_id = r.private_message_sender_user_id
|
||||
AND b.private_message_id = r.private_message_id
|
||||
AND COALESCE((b.media #>> '{service_action,no_forwards,expired}')::boolean, false))
|
||||
FROM private_no_forwards_requests r
|
||||
JOIN private_messages pm
|
||||
ON pm.sender_user_id = r.private_message_sender_user_id
|
||||
AND pm.id = r.private_message_id
|
||||
WHERE r.private_message_sender_user_id = $1
|
||||
AND r.private_message_id = $2`, bob.ID, request.Send.SenderMessage.UID,
|
||||
).Scan(&handledAt, &logicalExpired, &expiredBoxes); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if handledAt != 1700100002 || !logicalExpired || expiredBoxes != 2 {
|
||||
t.Fatalf("handled request handled_at=%d logical_expired=%v boxes=%d", handledAt, logicalExpired, expiredBoxes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostgresPrivateNoForwardsConcurrentOwnershipAndOneShotAnswer(t *testing.T) {
|
||||
pool := testPool(t)
|
||||
ctx := context.Background()
|
||||
suffix := randomSuffix(t)
|
||||
users := NewUserStore(pool)
|
||||
alice, err := users.Create(ctx, domain.User{AccessHash: 6201, Phone: "+1669" + suffix + "01", FirstName: "Alice"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
bob, err := users.Create(ctx, domain.User{AccessHash: 6202, Phone: "+1669" + suffix + "02", FirstName: "Bob"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_, _ = pool.Exec(ctx, "DELETE FROM users WHERE id = ANY($1::bigint[])", []int64{alice.ID, bob.ID})
|
||||
})
|
||||
|
||||
messages := NewMessageStore(pool)
|
||||
baseRandom := time.Now().UnixNano()
|
||||
enableResults := make([]domain.TogglePrivateNoForwardsResult, 2)
|
||||
enableErrors := make([]error, 2)
|
||||
actors := []int64{alice.ID, bob.ID}
|
||||
peers := []int64{bob.ID, alice.ID}
|
||||
var wg sync.WaitGroup
|
||||
for i := range actors {
|
||||
wg.Add(1)
|
||||
go func(i int) {
|
||||
defer wg.Done()
|
||||
enableResults[i], enableErrors[i] = messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: actors[i],
|
||||
PeerUserID: peers[i],
|
||||
Enabled: true,
|
||||
RandomID: baseRandom + int64(i),
|
||||
Date: 1700200000,
|
||||
})
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
changed := 0
|
||||
for i, err := range enableErrors {
|
||||
if err != nil {
|
||||
t.Fatalf("concurrent enable %d: %v", i, err)
|
||||
}
|
||||
if enableResults[i].Changed {
|
||||
changed++
|
||||
}
|
||||
}
|
||||
if changed != 1 {
|
||||
t.Fatalf("concurrent enable changed=%d, want exactly one service message", changed)
|
||||
}
|
||||
state, err := messages.GetPrivateNoForwards(ctx, alice.ID, bob.ID)
|
||||
if err != nil || (state.EnabledByUserID != alice.ID && state.EnabledByUserID != bob.ID) {
|
||||
t.Fatalf("concurrent enable state=%+v err=%v", state, err)
|
||||
}
|
||||
|
||||
ownerID := state.EnabledByUserID
|
||||
requesterID := alice.ID
|
||||
if ownerID == alice.ID {
|
||||
requesterID = bob.ID
|
||||
}
|
||||
request, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: requesterID,
|
||||
PeerUserID: ownerID,
|
||||
RandomID: baseRandom + 10,
|
||||
Date: 1700200001,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("create disable request: %v", err)
|
||||
}
|
||||
|
||||
answerResults := make([]domain.TogglePrivateNoForwardsResult, 2)
|
||||
answerErrors := make([]error, 2)
|
||||
for i := range answerResults {
|
||||
wg.Add(1)
|
||||
go func(i int) {
|
||||
defer wg.Done()
|
||||
answerResults[i], answerErrors[i] = messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||
ActorUserID: ownerID,
|
||||
PeerUserID: requesterID,
|
||||
Enabled: false,
|
||||
RequestMsgID: request.Send.RecipientMessage.ID,
|
||||
RandomID: baseRandom + 20 + int64(i),
|
||||
Date: 1700200002,
|
||||
})
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
successes, expired := 0, 0
|
||||
for i, err := range answerErrors {
|
||||
switch {
|
||||
case err == nil && answerResults[i].Changed:
|
||||
successes++
|
||||
case errors.Is(err, domain.ErrNoForwardsRequestExpired):
|
||||
expired++
|
||||
default:
|
||||
t.Fatalf("concurrent answer %d result=%+v err=%v", i, answerResults[i], err)
|
||||
}
|
||||
}
|
||||
if successes != 1 || expired != 1 {
|
||||
t.Fatalf("concurrent answers successes=%d expired=%d, want 1/1", successes, expired)
|
||||
}
|
||||
state, err = messages.GetPrivateNoForwards(ctx, alice.ID, bob.ID)
|
||||
if err != nil || state.Enabled() {
|
||||
t.Fatalf("state after concurrent answer=%+v err=%v, want disabled", state, err)
|
||||
}
|
||||
for _, userID := range []int64{alice.ID, bob.ID} {
|
||||
events, err := NewUpdateEventStore(pool).ListAfter(ctx, userID, 0, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("events user %d: %v", userID, err)
|
||||
}
|
||||
if len(events) != 3 || events[0].Pts != 1 || events[1].Pts != 2 || events[2].Pts != 3 {
|
||||
t.Fatalf("events user %d = %+v, want one enable/request/answer sequence", userID, events)
|
||||
}
|
||||
}
|
||||
}
|
||||
15
internal/store/private_no_forwards.go
Normal file
15
internal/store/private_no_forwards.go
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
// PrivateNoForwardsStore keeps the canonical pair state and its service-message
|
||||
// transition in one store transaction. It is optional so unrelated lightweight
|
||||
// MessageStore test doubles do not need to implement this capability.
|
||||
type PrivateNoForwardsStore interface {
|
||||
GetPrivateNoForwards(ctx context.Context, viewerUserID, peerUserID int64) (domain.PrivateNoForwardsState, error)
|
||||
TogglePrivateNoForwards(ctx context.Context, req domain.TogglePrivateNoForwardsRequest) (domain.TogglePrivateNoForwardsResult, error)
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue