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

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

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

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