diff --git a/deploy/migrations/0159_private_no_forwards.down.sql b/deploy/migrations/0159_private_no_forwards.down.sql new file mode 100644 index 00000000..62bbad4d --- /dev/null +++ b/deploy/migrations/0159_private_no_forwards.down.sql @@ -0,0 +1,2 @@ +DROP TABLE IF EXISTS private_no_forwards_requests; +DROP TABLE IF EXISTS private_no_forwards_chats; diff --git a/deploy/migrations/0159_private_no_forwards.up.sql b/deploy/migrations/0159_private_no_forwards.up.sql new file mode 100644 index 00000000..bfa7662d --- /dev/null +++ b/deploy/migrations/0159_private_no_forwards.up.sql @@ -0,0 +1,36 @@ +CREATE TABLE private_no_forwards_chats ( + user_low_id bigint NOT NULL REFERENCES users(id) ON DELETE CASCADE, + user_high_id bigint NOT NULL REFERENCES users(id) ON DELETE CASCADE, + enabled_by_user_id bigint REFERENCES users(id) ON DELETE CASCADE, + created_at timestamptz NOT NULL DEFAULT now(), + updated_at timestamptz NOT NULL DEFAULT now(), + PRIMARY KEY (user_low_id, user_high_id), + CONSTRAINT private_no_forwards_distinct_users CHECK (user_low_id < user_high_id), + CONSTRAINT private_no_forwards_enabled_participant CHECK ( + enabled_by_user_id IS NULL + OR enabled_by_user_id = user_low_id + OR enabled_by_user_id = user_high_id + ) +); + +CREATE TABLE private_no_forwards_requests ( + private_message_sender_user_id bigint NOT NULL, + private_message_id bigint NOT NULL, + requester_user_id bigint NOT NULL REFERENCES users(id) ON DELETE CASCADE, + responder_user_id bigint NOT NULL REFERENCES users(id) ON DELETE CASCADE, + expires_at integer NOT NULL, + handled_at integer NOT NULL DEFAULT 0, + created_at timestamptz NOT NULL DEFAULT now(), + PRIMARY KEY (private_message_sender_user_id, private_message_id), + CONSTRAINT private_no_forwards_request_message_fk + FOREIGN KEY (private_message_sender_user_id, private_message_id) + REFERENCES private_messages(sender_user_id, id) ON DELETE CASCADE, + CONSTRAINT private_no_forwards_request_distinct_users + CHECK (requester_user_id <> responder_user_id), + CONSTRAINT private_no_forwards_request_valid_expiry + CHECK (expires_at > 0 AND handled_at >= 0) +); + +CREATE INDEX private_no_forwards_requests_responder_expiry_idx + ON private_no_forwards_requests (responder_user_id, expires_at) + WHERE handled_at = 0; diff --git a/internal/app/help/service.go b/internal/app/help/service.go index 15721094..11a3a3cb 100644 --- a/internal/app/help/service.go +++ b/internal/app/help/service.go @@ -56,8 +56,9 @@ const tdesktopClient = "tdesktop" // WebK directly calls Array.some on fragment_prefixes while rendering user profiles, // so this compatibility key must always remain an array, even when it is empty. const tdesktopDefaultAppConfigBase = `{"chat_read_mark_expire_period":604800,"chat_read_mark_size_threshold":50,"pm_read_date_expire_period":604800,"quote_length_max":1024,"telegram_antispam_group_size_min":200,"telegram_antispam_user_id":"5434988373","fragment_prefixes":["888"],"forum_upgrade_participants_min":2,"reactions_default":{"_":"reactionEmoji","emoticon":"👍"},"reactions_uniq_max":11,"reactions_user_max_default":1,"reactions_user_max_premium":3,"reactions_in_chat_max":3,"boosts_channel_level_max":100,"rich_message_posting":"enabled","upload_markup_video":true,"emojies_send_dice":["🎲","🎯","🏀","⚽","⚽️","🎳","🎰"],"premium_purchase_blocked":false,"stars_purchase_blocked":false,"stargifts_blocked":false,"stargifts_pinned_to_top_limit":6,"stories_stealth_future_period":1500,"stories_stealth_past_period":300,"stories_stealth_cooldown_period":10800,"quick_replies_limit":100,"quick_reply_messages_limit":20,"business_chat_links_limit":100,"dialog_filters_enabled":true,"chatlist_update_period":3600,"chatlist_invites_limit_default":3,"chatlist_invites_limit_premium":20,"chatlists_joined_limit_default":2,"chatlists_joined_limit_premium":20,"about_length_limit_default":70,"about_length_limit_premium":140,"bot_verification_description_length_limit":70,"caption_length_limit_default":1024,"caption_length_limit_premium":4096,"channels_limit_default":500,"channels_limit_premium":1000,"channels_public_limit_default":10,"channels_public_limit_premium":20,"dialog_filters_limit_default":10,"dialog_filters_limit_premium":20,"dialog_filters_chats_limit_default":100,"dialog_filters_chats_limit_premium":200,"dialogs_pinned_limit_default":5,"dialogs_pinned_limit_premium":10,"dialogs_folder_pinned_limit_default":100,"dialogs_folder_pinned_limit_premium":200,"saved_dialogs_pinned_limit_default":5,"saved_dialogs_pinned_limit_premium":100,"saved_gifs_limit_default":200,"saved_gifs_limit_premium":400,"stickers_faved_limit_default":5,"stickers_faved_limit_premium":10,"recommended_channels_limit_default":10,"recommended_channels_limit_premium":100,"aicompose_tone_examples_num":3,"aicompose_tone_title_length_max":12,"aicompose_tone_prompt_length_max":1024,"aicompose_tone_saved_limit_default":5,"aicompose_tone_saved_limit_premium":20,"upload_max_fileparts_default":4000,"upload_max_fileparts_premium":8000` +const tdesktopNoForwardsAppConfig = `,"no_forwards_request_expire_period":86400` -const defaultAppConfigHash = 25 // 默认 app config 内容变更时必须递增,否则缓存端只会收到 notModified。 +const defaultAppConfigHash = 26 // 默认 app config 内容变更时必须递增,否则缓存端只会收到 notModified。 // Service 提供客户端启动配置与国家区号目录。 // @@ -117,14 +118,14 @@ func defaultAppConfig(mapboxToken string) domain.AppConfig { func defaultAppConfigJSON(mapboxToken string) []byte { if mapboxToken == "" { - return []byte(tdesktopDefaultAppConfigBase + `}`) + return []byte(tdesktopDefaultAppConfigBase + tdesktopNoForwardsAppConfig + `}`) } token, err := json.Marshal(mapboxToken) if err != nil { - return []byte(tdesktopDefaultAppConfigBase + `}`) + return []byte(tdesktopDefaultAppConfigBase + tdesktopNoForwardsAppConfig + `}`) } tokenJSON := string(token) - return []byte(tdesktopDefaultAppConfigBase + `,"tdesktop_config_map":{"maps":` + tokenJSON + `,"geo":` + tokenJSON + `,"bmaps":` + tokenJSON + `,"bgeo":` + tokenJSON + `}}`) + return []byte(tdesktopDefaultAppConfigBase + tdesktopNoForwardsAppConfig + `,"tdesktop_config_map":{"maps":` + tokenJSON + `,"geo":` + tokenJSON + `,"bmaps":` + tokenJSON + `,"bgeo":` + tokenJSON + `}}`) } func defaultAppConfigHashFor(mapboxToken string) int { diff --git a/internal/app/help/service_premium_test.go b/internal/app/help/service_premium_test.go index 27904f6b..6f58a6ea 100644 --- a/internal/app/help/service_premium_test.go +++ b/internal/app/help/service_premium_test.go @@ -4,6 +4,8 @@ import ( "context" "encoding/json" "testing" + + "telesrv/internal/domain" ) // TestAppConfigPremiumKeys 断言 premium / Stars 相关 key 完整下发且 hash 已递增: @@ -26,6 +28,9 @@ func TestAppConfigPremiumKeys(t *testing.T) { if err := json.Unmarshal(cfg.JSON, &decoded); err != nil { t.Fatalf("app config json invalid: %v", err) } + if period, ok := decoded["no_forwards_request_expire_period"].(float64); !ok || int(period) != domain.PrivateNoForwardsRequestExpirePeriod { + t.Fatalf("no_forwards_request_expire_period = %v, want %d", decoded["no_forwards_request_expire_period"], domain.PrivateNoForwardsRequestExpirePeriod) + } if blocked, ok := decoded["premium_purchase_blocked"].(bool); !ok || blocked { t.Fatalf("premium_purchase_blocked = %v, want false (star gift 送礼入口耦合此 flag)", decoded["premium_purchase_blocked"]) } diff --git a/internal/app/messages/service.go b/internal/app/messages/service.go index bf1e1ab2..c3db6bb0 100644 --- a/internal/app/messages/service.go +++ b/internal/app/messages/service.go @@ -222,6 +222,42 @@ func (s *Service) SetChatTheme(ctx context.Context, userID int64, req domain.Set return out, nil } +// GetPrivateNoForwards returns the canonical content-protection state for one +// ordinary private chat. +func (s *Service) GetPrivateNoForwards(ctx context.Context, userID, peerUserID int64) (domain.PrivateNoForwardsState, error) { + if s == nil || s.messages == nil || userID == 0 || peerUserID == 0 || userID == peerUserID { + return domain.PrivateNoForwardsState{}, domain.ErrMessageIDInvalid + } + backend, ok := s.messages.(store.PrivateNoForwardsStore) + if !ok { + return domain.PrivateNoForwardsState{}, nil + } + return backend.GetPrivateNoForwards(ctx, userID, peerUserID) +} + +// TogglePrivateNoForwards atomically mutates the pair state and appends the +// corresponding service message when the official state machine requires one. +func (s *Service) TogglePrivateNoForwards(ctx context.Context, userID int64, req domain.TogglePrivateNoForwardsRequest) (domain.TogglePrivateNoForwardsResult, error) { + if s == nil || s.messages == nil || userID == 0 { + return domain.TogglePrivateNoForwardsResult{}, domain.ErrMessageIDInvalid + } + if req.ActorUserID == 0 { + req.ActorUserID = userID + } + if req.ActorUserID != userID || req.PeerUserID == 0 || req.PeerUserID == userID || + req.RequestMsgID < 0 || req.RequestMsgID > domain.MaxMessageBoxID { + return domain.TogglePrivateNoForwardsResult{}, domain.ErrMessageIDInvalid + } + if err := s.ensureCanSend(ctx, userID); err != nil { + return domain.TogglePrivateNoForwardsResult{}, err + } + backend, ok := s.messages.(store.PrivateNoForwardsStore) + if !ok { + return domain.TogglePrivateNoForwardsResult{}, domain.ErrMessageIDInvalid + } + return backend.TogglePrivateNoForwards(ctx, req) +} + func chatThemeServiceMedia(emoticon string) *domain.MessageMedia { return &domain.MessageMedia{ Kind: domain.MessageMediaKindService, diff --git a/internal/domain/media.go b/internal/domain/media.go index 8875ca39..77120d05 100644 --- a/internal/domain/media.go +++ b/internal/domain/media.go @@ -558,6 +558,10 @@ const ( // MessageServiceActionSetChatTheme 映射 messageActionSetChatTheme: // 私聊双方共享的 chat theme token 变更。 MessageServiceActionSetChatTheme MessageServiceActionKind = "set_chat_theme" + // MessageServiceActionNoForwardsToggle / Request 映射私聊内容保护的 + // 状态切换与关闭请求。会话级保护不能写入普通消息的 NoForwards 字段。 + MessageServiceActionNoForwardsToggle MessageServiceActionKind = "no_forwards_toggle" + MessageServiceActionNoForwardsRequest MessageServiceActionKind = "no_forwards_request" // MessageServiceActionStarGift 映射 messageActionStarGift:收到一份 Star 礼物。 // 礼物快照(贴纸/星价)内嵌在 action 里,收礼人无需额外拉取即可渲染气泡。 MessageServiceActionStarGift MessageServiceActionKind = "star_gift" @@ -629,6 +633,15 @@ type MessageRequestedPeerDetails struct { Photo *Photo `json:"photo,omitempty"` } +// MessageNoForwardsAction 是私聊内容保护 service action 的协议中立载荷。 +// ExpiresAt 只用于 request 的读取时绝对过期投影;toggle 保持为 0。 +type MessageNoForwardsAction struct { + PrevValue bool `json:"prev_value"` + NewValue bool `json:"new_value"` + Expired bool `json:"expired,omitempty"` + ExpiresAt int `json:"expires_at,omitempty"` +} + // MessageServiceAction 是私聊服务消息动作的协议中立表示。 type MessageServiceAction struct { Kind MessageServiceActionKind `json:"kind"` @@ -639,6 +652,7 @@ type MessageServiceAction struct { WebViewData *MessageWebViewDataAction `json:"web_view_data,omitempty"` RequestedPeer *MessageRequestedPeerAction `json:"requested_peer,omitempty"` ChatThemeEmoticon string `json:"chat_theme_emoticon,omitempty"` + NoForwards *MessageNoForwardsAction `json:"no_forwards,omitempty"` StarGift *MessageStarGiftAction `json:"star_gift,omitempty"` StarGiftUnique *MessageStarGiftUniqueAction `json:"star_gift_unique,omitempty"` StarGiftOffer *MessageStarGiftOfferAction `json:"star_gift_offer,omitempty"` diff --git a/internal/domain/message.go b/internal/domain/message.go index 5334e2d0..f868c608 100644 --- a/internal/domain/message.go +++ b/internal/domain/message.go @@ -407,6 +407,47 @@ type ForwardPrivateMessagesResult struct { ReplayDeleteEvents []*UpdateEvent } +const PrivateNoForwardsRequestExpirePeriod = 24 * 60 * 60 + +// PrivateNoForwardsState 是一对普通用户唯一的内容保护权威。 +// EnabledByUserID 为 0 或参与者之一;非零时双方共享的会话均受保护。 +type PrivateNoForwardsState struct { + UserLowID int64 + UserHighID int64 + EnabledByUserID int64 +} + +func (s PrivateNoForwardsState) Enabled() bool { + return s.EnabledByUserID != 0 +} + +func (s PrivateNoForwardsState) ForViewer(viewerUserID int64) (myEnabled, peerEnabled bool) { + if s.EnabledByUserID == 0 { + return false, false + } + return s.EnabledByUserID == viewerUserID, s.EnabledByUserID != viewerUserID +} + +// TogglePrivateNoForwardsRequest 是私聊内容保护的原子状态+服务消息命令。 +type TogglePrivateNoForwardsRequest struct { + ActorUserID int64 + PeerUserID int64 + Enabled bool + RequestMsgID int + RandomID int64 + Date int + OriginAuthKeyID [8]byte + OriginSessionID int64 +} + +// TogglePrivateNoForwardsResult 同时返回提交后的权威状态与本次真实服务消息。 +// Changed=false 且 Send 为空表示官方定义的 no-op。 +type TogglePrivateNoForwardsResult struct { + State PrivateNoForwardsState + Changed bool + Send SendPrivateTextResult +} + // ReadHistoryRequest 是账号视角的 messages.readHistory 命令。 type ReadHistoryRequest struct { OwnerUserID int64 diff --git a/internal/domain/message_errors.go b/internal/domain/message_errors.go index 342c0ba2..f7b6e9d6 100644 --- a/internal/domain/message_errors.go +++ b/internal/domain/message_errors.go @@ -27,6 +27,7 @@ var ( ErrLoginCodeDeliveryCommitAmbiguous = errors.New("login code delivery commit ambiguous") ErrReplyMessageIDInvalid = errors.New("reply message id invalid") ErrChatForwardsRestricted = errors.New("chat forwards restricted") + ErrNoForwardsRequestExpired = errors.New("no forwards request expired") // ErrPinnedSavedDialogsTooMuch 映射 PINNED_TOO_MUCH:收藏夹子会话置顶 // 数量达到 MaxPinnedSavedDialogs 上限。 ErrPinnedSavedDialogsTooMuch = errors.New("pinned saved dialogs too much") diff --git a/internal/rpc/channels_legacy_chat.go b/internal/rpc/channels_legacy_chat.go index e3b16959..1ce197d6 100644 --- a/internal/rpc/channels_legacy_chat.go +++ b/internal/rpc/channels_legacy_chat.go @@ -625,25 +625,6 @@ func (r *Router) enqueueChannelWallpaperFanout(ctx context.Context, originUserID }) } -func (r *Router) onMessagesToggleNoForwards(ctx context.Context, req *tg.MessagesToggleNoForwardsRequest) (tg.UpdatesClass, error) { - if r.deps.Channels == nil { - return nil, notImplementedErr() - } - userID, _, err := r.currentUserID(ctx) - if err != nil { - return nil, internalErr() - } - channelID, err := r.channelIDFromLegacyInputPeerChecked(ctx, userID, req.Peer) - if err != nil { - return nil, err - } - channel, err := r.deps.Channels.SetNoForwards(ctx, userID, channelID, req.Enabled) - if err != nil { - return nil, channelAdminErr(err) - } - return r.channelStateMutationUpdates(ctx, userID, channel), nil -} - func (r *Router) onMessagesSetChatAvailableReactions(ctx context.Context, req *tg.MessagesSetChatAvailableReactionsRequest) (tg.UpdatesClass, error) { if r.deps.Channels == nil { return nil, notImplementedErr() diff --git a/internal/rpc/convert_messages.go b/internal/rpc/convert_messages.go index af793457..a33fdaff 100644 --- a/internal/rpc/convert_messages.go +++ b/internal/rpc/convert_messages.go @@ -1,7 +1,10 @@ package rpc import ( + "time" + "github.com/iamxvbaba/td/tg" + "telesrv/internal/domain" ) @@ -145,6 +148,25 @@ func tgMessageServiceAction(msg domain.Message) tg.MessageActionClass { return &tg.MessageActionSetChatTheme{ Theme: &tg.ChatTheme{Emoticon: m.ServiceAction.ChatThemeEmoticon}, } + case domain.MessageServiceActionNoForwardsToggle: + action := m.ServiceAction.NoForwards + if action == nil { + return &tg.MessageActionEmpty{} + } + return &tg.MessageActionNoForwardsToggle{ + PrevValue: action.PrevValue, + NewValue: action.NewValue, + } + case domain.MessageServiceActionNoForwardsRequest: + action := m.ServiceAction.NoForwards + if action == nil { + return &tg.MessageActionEmpty{} + } + return &tg.MessageActionNoForwardsRequest{ + Expired: action.Expired || (action.ExpiresAt > 0 && int(time.Now().Unix()) >= action.ExpiresAt), + PrevValue: action.PrevValue, + NewValue: action.NewValue, + } case domain.MessageServiceActionPhoneCall: if m.ServiceAction.Call == nil { return &tg.MessageActionEmpty{} diff --git a/internal/rpc/deps.go b/internal/rpc/deps.go index 1a885e68..3be483f3 100644 --- a/internal/rpc/deps.go +++ b/internal/rpc/deps.go @@ -604,6 +604,13 @@ type MessagesService interface { DeleteSavedHistory(ctx context.Context, userID int64, req domain.DeleteSavedHistoryRequest) (domain.DeleteSavedHistoryResult, error) } +// PrivateNoForwardsService is an optional messages capability used by the +// private-user branch of messages.toggleNoForwards and userFull projection. +type PrivateNoForwardsService interface { + GetPrivateNoForwards(ctx context.Context, userID, peerUserID int64) (domain.PrivateNoForwardsState, error) + TogglePrivateNoForwards(ctx context.Context, userID int64, req domain.TogglePrivateNoForwardsRequest) (domain.TogglePrivateNoForwardsResult, error) +} + // TranslationService owns read-only translation and the durable per-account // peer preference. It only exposes domain values to the RPC edge. type TranslationService interface { diff --git a/internal/rpc/errors.go b/internal/rpc/errors.go index 3811eef0..194b58ae 100644 --- a/internal/rpc/errors.go +++ b/internal/rpc/errors.go @@ -297,6 +297,8 @@ func replyMessageIDInvalidErr() error { return tgerr.New(400, "REPLY_MESSAGE_ID_ func chatForwardsRestrictedErr() error { return tgerr.New(400, "CHAT_FORWARDS_RESTRICTED") } +func requestMsgExpiredErr() error { return tgerr.New(400, "REQUEST_MSG_EXPIRED") } + func inputRequestInvalidErr() error { return tgerr.New(400, "INPUT_REQUEST_INVALID") } func inputRequestTooLongErr() error { return tgerr.New(400, "INPUT_REQUEST_TOO_LONG") } diff --git a/internal/rpc/messages_forward.go b/internal/rpc/messages_forward.go index 8254630a..a22a7caf 100644 --- a/internal/rpc/messages_forward.go +++ b/internal/rpc/messages_forward.go @@ -515,6 +515,15 @@ func (r *Router) forwardSourcesFromPrivateMessages(ctx context.Context, userID i if fromPeer.Type != domain.PeerTypeUser || fromPeer.ID == 0 { return nil, domain.ErrMessageIDInvalid } + if svc, ok := r.deps.Messages.(PrivateNoForwardsService); ok { + state, err := svc.GetPrivateNoForwards(ctx, userID, fromPeer.ID) + if err != nil { + return nil, err + } + if state.Enabled() { + return nil, domain.ErrChatForwardsRestricted + } + } byID := make(map[int]domain.Message, len(messages)) for _, msg := range messages { byID[msg.ID] = msg diff --git a/internal/rpc/messages_no_forwards.go b/internal/rpc/messages_no_forwards.go new file mode 100644 index 00000000..feb0ba8a --- /dev/null +++ b/internal/rpc/messages_no_forwards.go @@ -0,0 +1,144 @@ +package rpc + +import ( + "context" + cryptorand "crypto/rand" + "encoding/binary" + "errors" + "sync/atomic" + + "github.com/iamxvbaba/td/tg" + + "telesrv/internal/domain" +) + +var privateNoForwardsRandomFallback atomic.Uint64 + +func (r *Router) onMessagesToggleNoForwards(ctx context.Context, req *tg.MessagesToggleNoForwardsRequest) (tg.UpdatesClass, error) { + if req == nil { + return nil, inputRequestInvalidErr() + } + userID, _, err := r.currentUserID(ctx) + if err != nil { + return nil, internalErr() + } + peer, err := r.checkedDomainPeerFromInputPeer(ctx, userID, req.Peer) + if err != nil { + return nil, err + } + if peer.Type == domain.PeerTypeChannel { + if r.deps.Channels == nil { + return nil, notImplementedErr() + } + channel, err := r.deps.Channels.SetNoForwards(ctx, userID, peer.ID, req.Enabled) + if err != nil { + return nil, channelAdminErr(err) + } + return r.channelStateMutationUpdates(ctx, userID, channel), nil + } + input, ok := req.Peer.(*tg.InputPeerUser) + if !ok || input == nil || peer.Type != domain.PeerTypeUser || peer.ID == 0 || peer.ID == userID || + r.deps.Users == nil { + return nil, peerIDInvalidErr() + } + if err := r.validateInputUser(ctx, &tg.InputUser{UserID: input.UserID, AccessHash: input.AccessHash}); err != nil { + return nil, peerIDInvalidErr() + } + target, found, err := r.deps.Users.ByID(ctx, userID, peer.ID) + if err != nil { + return nil, internalErr() + } + if !found || target.Bot || target.Support || target.Deleted { + return nil, peerIDInvalidErr() + } + self, err := r.deps.Users.Self(ctx, userID) + if err != nil { + return nil, internalErr() + } + if self.Bot || self.Support || self.Deleted { + return nil, peerIDInvalidErr() + } + svc, ok := r.deps.Messages.(PrivateNoForwardsService) + if !ok { + return nil, notImplementedErr() + } + requestMsgID, hasRequestMsgID := req.GetRequestMsgID() + if hasRequestMsgID && (requestMsgID <= 0 || requestMsgID > domain.MaxMessageBoxID) { + return nil, requestMsgExpiredErr() + } + if !hasRequestMsgID { + requestMsgID = 0 + } + current, err := svc.GetPrivateNoForwards(ctx, userID, peer.ID) + if err != nil { + return nil, privateNoForwardsErr(err) + } + // Premium is required only to create a new protected state. Disabling, + // answering a request and no-op retries remain possible after expiry. + if req.Enabled && requestMsgID == 0 && !current.Enabled() && !self.PremiumActiveAt(r.clock.Now().Unix()) { + return nil, premiumAccountRequiredErr() + } + if requestMsgID != 0 || current.Enabled() != req.Enabled { + if err := r.checkSendRateLimit(ctx, userID, 1); err != nil { + return nil, err + } + } + sessionID, _ := SessionIDFrom(ctx) + result, err := svc.TogglePrivateNoForwards(ctx, userID, domain.TogglePrivateNoForwardsRequest{ + ActorUserID: userID, + PeerUserID: peer.ID, + Enabled: req.Enabled, + RequestMsgID: requestMsgID, + RandomID: newPrivateNoForwardsRandomID(), + Date: int(r.clock.Now().Unix()), + OriginAuthKeyID: rawAuthKeyIDForOrigin(ctx), + OriginSessionID: sessionID, + }) + if err != nil { + return nil, privateNoForwardsErr(err) + } + if result.Changed { + r.invalidateRPCProjectionForPeer(userID, peer) + r.invalidateRPCProjectionForPeer(peer.ID, domain.Peer{Type: domain.PeerTypeUser, ID: userID}) + } + if !result.Changed || result.Send.SenderMessage.ID == 0 { + return tgEmptyUpdates(int(r.clock.Now().Unix())), nil + } + return tgPrivateMessageUpdates( + result.Send.SenderEvent, + result.Send.SenderMessage, + 0, + false, + r.usersForMessageUpdate(ctx, userID, result.Send.SenderMessage), + []tg.ChatClass{}, + ), nil +} + +func privateNoForwardsErr(err error) error { + switch { + case errors.Is(err, domain.ErrNoForwardsRequestExpired), + errors.Is(err, domain.ErrReplyMessageIDInvalid): + return requestMsgExpiredErr() + case errors.Is(err, domain.ErrMessageIDInvalid): + return peerIDInvalidErr() + case errors.Is(err, domain.ErrChatForwardsRestricted): + return chatForwardsRestrictedErr() + case errors.Is(err, domain.ErrUserFrozen): + return frozenMethodInvalidErr() + case errors.Is(err, domain.ErrMessageRandomIDDuplicate): + return randomIDDuplicateErr() + default: + return internalErr() + } +} + +func newPrivateNoForwardsRandomID() int64 { + var raw [8]byte + if _, err := cryptorand.Read(raw[:]); err == nil { + if value := int64(binary.LittleEndian.Uint64(raw[:])); value != 0 { + return value + } + } + value := privateNoForwardsRandomFallback.Add(1) + return int64(value | 1<<62) +} diff --git a/internal/rpc/messages_no_forwards_rpc_test.go b/internal/rpc/messages_no_forwards_rpc_test.go new file mode 100644 index 00000000..877c3939 --- /dev/null +++ b/internal/rpc/messages_no_forwards_rpc_test.go @@ -0,0 +1,167 @@ +package rpc + +import ( + "context" + "testing" + "time" + + "github.com/iamxvbaba/td/clock" + "github.com/iamxvbaba/td/tg" + "github.com/iamxvbaba/td/tgerr" + "go.uber.org/zap/zaptest" + + appdialogs "telesrv/internal/app/dialogs" + appmessages "telesrv/internal/app/messages" + appusers "telesrv/internal/app/users" + "telesrv/internal/domain" + "telesrv/internal/store/memory" +) + +func TestMessagesToggleNoForwardsPrivateFullFlow(t *testing.T) { + ctx := context.Background() + usersStore := memory.NewUserStore() + alice, _ := usersStore.Create(ctx, domain.User{AccessHash: 5101, Phone: "15550005101", FirstName: "Alice"}) + bob, _ := usersStore.Create(ctx, domain.User{AccessHash: 5102, Phone: "15550005102", FirstName: "Bob"}) + if _, err := usersStore.SetPremiumUntil(ctx, alice.ID, int(time.Now().Add(time.Hour).Unix())); err != nil { + t.Fatalf("grant alice premium: %v", err) + } + dialogsStore := memory.NewDialogStore() + messagesStore := memory.NewMessageStore(dialogsStore) + router := New(Config{}, Deps{ + Users: appusers.NewService(usersStore), + Dialogs: appdialogs.NewService(dialogsStore), + Messages: appmessages.NewService(messagesStore, dialogsStore), + }, zaptest.NewLogger(t), clock.System) + + enable, err := router.onMessagesToggleNoForwards(WithUserID(ctx, alice.ID), &tg.MessagesToggleNoForwardsRequest{ + Peer: &tg.InputPeerUser{UserID: bob.ID, AccessHash: bob.AccessHash}, Enabled: true, + }) + if err != nil { + t.Fatalf("enable private noforwards: %v", err) + } + enableMessage := noForwardsServiceMessage(t, enable) + if _, ok := enableMessage.Action.(*tg.MessageActionNoForwardsToggle); !ok { + t.Fatalf("enable action = %T", enableMessage.Action) + } + assertNoForwardsFullFlags(t, router, ctx, alice, bob, true, false) + assertNoForwardsFullFlags(t, router, ctx, bob, alice, false, true) + source, err := messagesStore.SendPrivateText(ctx, domain.SendPrivateTextRequest{ + SenderUserID: alice.ID, RecipientUserID: bob.ID, RandomID: 5199, Message: "protected source", Date: int(time.Now().Unix()), + }) + if err != nil { + t.Fatalf("send protected source: %v", err) + } + if _, err := router.onMessagesForwardMessages(WithUserID(ctx, alice.ID), &tg.MessagesForwardMessagesRequest{ + FromPeer: &tg.InputPeerUser{UserID: bob.ID, AccessHash: bob.AccessHash}, + ID: []int{source.SenderMessage.ID}, + RandomID: []int64{5200}, + ToPeer: &tg.InputPeerSelf{}, + }); !tgerr.Is(err, "CHAT_FORWARDS_RESTRICTED") { + t.Fatalf("forward protected private chat err=%v, want CHAT_FORWARDS_RESTRICTED", err) + } + + // The other party cannot steal ownership by setting enabled=true. This is a + // no-op and does not require that party to be premium. + noOp, err := router.onMessagesToggleNoForwards(WithUserID(ctx, bob.ID), &tg.MessagesToggleNoForwardsRequest{ + Peer: &tg.InputPeerUser{UserID: alice.ID, AccessHash: alice.AccessHash}, Enabled: true, + }) + if err != nil || len(noOp.(*tg.Updates).Updates) != 0 { + t.Fatalf("peer repeat enable = %#v err=%v, want empty no-op", noOp, err) + } + + requestUpdates, err := router.onMessagesToggleNoForwards(WithUserID(ctx, bob.ID), &tg.MessagesToggleNoForwardsRequest{ + Peer: &tg.InputPeerUser{UserID: alice.ID, AccessHash: alice.AccessHash}, Enabled: false, + }) + if err != nil { + t.Fatalf("request sharing: %v", err) + } + requestMessage := noForwardsServiceMessage(t, requestUpdates) + requestAction, ok := requestMessage.Action.(*tg.MessageActionNoForwardsRequest) + if !ok || requestAction.Expired || !requestAction.PrevValue || requestAction.NewValue { + t.Fatalf("request action = %#v", requestMessage.Action) + } + aliceHistory, err := messagesStore.ListByUser(ctx, alice.ID, domain.MessageFilter{ + HasPeer: true, Peer: domain.Peer{Type: domain.PeerTypeUser, ID: bob.ID}, Limit: 20, + }) + if err != nil { + t.Fatal(err) + } + var aliceRequestID int + for _, msg := range aliceHistory.Messages { + if msg.Media != nil && msg.Media.ServiceAction != nil && + msg.Media.ServiceAction.Kind == domain.MessageServiceActionNoForwardsRequest { + aliceRequestID = msg.ID + } + } + if aliceRequestID == 0 { + t.Fatal("alice request box not found") + } + + answerReq := &tg.MessagesToggleNoForwardsRequest{ + Peer: &tg.InputPeerUser{UserID: bob.ID, AccessHash: bob.AccessHash}, + Enabled: false, + } + answerReq.SetRequestMsgID(aliceRequestID) + answerUpdates, err := router.onMessagesToggleNoForwards(WithUserID(ctx, alice.ID), answerReq) + if err != nil { + t.Fatalf("accept sharing request: %v", err) + } + answerMessage := noForwardsServiceMessage(t, answerUpdates) + answerAction, ok := answerMessage.Action.(*tg.MessageActionNoForwardsToggle) + if !ok || !answerAction.PrevValue || answerAction.NewValue { + t.Fatalf("answer action = %#v", answerMessage.Action) + } + if answerMessage.ReplyTo == nil { + t.Fatal("answer service message has no reply_to") + } + assertNoForwardsFullFlags(t, router, ctx, alice, bob, false, false) + assertNoForwardsFullFlags(t, router, ctx, bob, alice, false, false) + + if _, err := router.onMessagesToggleNoForwards(WithUserID(ctx, alice.ID), answerReq); !tgerr.Is(err, "REQUEST_MSG_EXPIRED") { + t.Fatalf("repeat request answer err=%v, want REQUEST_MSG_EXPIRED", err) + } + if _, err := router.onMessagesToggleNoForwards(WithUserID(ctx, bob.ID), &tg.MessagesToggleNoForwardsRequest{ + Peer: &tg.InputPeerUser{UserID: alice.ID, AccessHash: alice.AccessHash}, Enabled: true, + }); !tgerr.Is(err, "PREMIUM_ACCOUNT_REQUIRED") { + t.Fatalf("non-premium fresh enable err=%v, want PREMIUM_ACCOUNT_REQUIRED", err) + } + if _, err := router.onMessagesToggleNoForwards(WithUserID(ctx, alice.ID), &tg.MessagesToggleNoForwardsRequest{ + Peer: &tg.InputPeerUser{UserID: bob.ID, AccessHash: bob.AccessHash + 1}, Enabled: true, + }); !tgerr.Is(err, "PEER_ID_INVALID") { + t.Fatalf("wrong access hash err=%v, want PEER_ID_INVALID", err) + } +} + +func noForwardsServiceMessage(t *testing.T, updates tg.UpdatesClass) *tg.MessageService { + t.Helper() + full, ok := updates.(*tg.Updates) + if !ok || len(full.Updates) != 1 { + t.Fatalf("updates = %#v, want one updateNewMessage", updates) + } + newMessage, ok := full.Updates[0].(*tg.UpdateNewMessage) + if !ok || newMessage.Pts <= 0 || newMessage.PtsCount != 1 { + t.Fatalf("update = %#v, want updateNewMessage pts_count=1", full.Updates[0]) + } + service, ok := newMessage.Message.(*tg.MessageService) + if !ok { + t.Fatalf("message = %#v, want messageService (which has no message.noforwards field)", newMessage.Message) + } + return service +} + +func assertNoForwardsFullFlags(t *testing.T, router *Router, ctx context.Context, viewer, target domain.User, wantMy, wantPeer bool) { + t.Helper() + full, err := router.onUsersGetFullUser(WithUserID(ctx, viewer.ID), &tg.InputUser{ + UserID: target.ID, AccessHash: target.AccessHash, + }) + if err != nil { + t.Fatalf("get full user %d->%d: %v", viewer.ID, target.ID, err) + } + if full.FullUser.GetNoforwardsMyEnabled() != wantMy || + full.FullUser.GetNoforwardsPeerEnabled() != wantPeer { + t.Fatalf("full flags %d->%d my=%v peer=%v, want %v/%v", + viewer.ID, target.ID, + full.FullUser.GetNoforwardsMyEnabled(), full.FullUser.GetNoforwardsPeerEnabled(), + wantMy, wantPeer) + } +} diff --git a/internal/rpc/users.go b/internal/rpc/users.go index d0e58866..412d09d7 100644 --- a/internal/rpc/users.go +++ b/internal/rpc/users.go @@ -265,6 +265,17 @@ func (r *Router) buildUserFullProjection(ctx context.Context, currentUserID int6 Settings: tg.PeerSettings{}, NotifySettings: *tdesktop.NotifySettings(), } + if u.ID != currentUserID { + if svc, ok := r.deps.Messages.(PrivateNoForwardsService); ok { + state, err := svc.GetPrivateNoForwards(ctx, currentUserID, u.ID) + if err != nil { + return tg.UserFull{}, internalErr() + } + myEnabled, peerEnabled := state.ForViewer(currentUserID) + full.SetNoforwardsMyEnabled(myEnabled) + full.SetNoforwardsPeerEnabled(peerEnabled) + } + } // 通话入口:客户端不见 phone_calls_available=true 不显示通话按钮(P1 前置项)。 // phone_calls_private 标记对端禁 P2P(p2p_allowed 真值在通话确认时另行计算)。 if !u.Bot && u.ID != currentUserID { diff --git a/internal/store/memory/message_forward.go b/internal/store/memory/message_forward.go index 52baf693..cf72ef5b 100644 --- a/internal/store/memory/message_forward.go +++ b/internal/store/memory/message_forward.go @@ -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 { diff --git a/internal/store/memory/message_helpers.go b/internal/store/memory/message_helpers.go index ab5fd796..4ecdfe5a 100644 --- a/internal/store/memory/message_helpers.go +++ b/internal/store/memory/message_helpers.go @@ -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...) diff --git a/internal/store/memory/message_no_forwards.go b/internal/store/memory/message_no_forwards.go new file mode 100644 index 00000000..49192a0d --- /dev/null +++ b/internal/store/memory/message_no_forwards.go @@ -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() +} diff --git a/internal/store/memory/message_no_forwards_test.go b/internal/store/memory/message_no_forwards_test.go new file mode 100644 index 00000000..b076a2b2 --- /dev/null +++ b/internal/store/memory/message_no_forwards_test.go @@ -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) + } +} diff --git a/internal/store/memory/message_store.go b/internal/store/memory/message_store.go index 2d3677bf..830b56cc 100644 --- a/internal/store/memory/message_store.go +++ b/internal/store/memory/message_store.go @@ -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] diff --git a/internal/store/postgres/message_forward.go b/internal/store/postgres/message_forward.go index dfa4445e..89f841a7 100644 --- a/internal/store/postgres/message_forward.go +++ b/internal/store/postgres/message_forward.go @@ -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 { diff --git a/internal/store/postgres/message_no_forwards.go b/internal/store/postgres/message_no_forwards.go new file mode 100644 index 00000000..b76fc31b --- /dev/null +++ b/internal/store/postgres/message_no_forwards.go @@ -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 +} diff --git a/internal/store/postgres/message_no_forwards_integration_test.go b/internal/store/postgres/message_no_forwards_integration_test.go new file mode 100644 index 00000000..d58a70a9 --- /dev/null +++ b/internal/store/postgres/message_no_forwards_integration_test.go @@ -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) + } + } +} diff --git a/internal/store/private_no_forwards.go b/internal/store/private_no_forwards.go new file mode 100644 index 00000000..cf648904 --- /dev/null +++ b/internal/store/private_no_forwards.go @@ -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) +}