fix: sync support private chat content protection

This commit is contained in:
iamxvbaba 2026-07-28 16:22:33 +08:00
parent 74c9249091
commit 037ce017d4
25 changed files with 1395 additions and 36 deletions

View file

@ -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 {

View file

@ -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...)

View 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()
}

View 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)
}
}

View file

@ -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]