fix: sync support private chat content protection
This commit is contained in:
parent
74c9249091
commit
037ce017d4
25 changed files with 1395 additions and 36 deletions
2
deploy/migrations/0159_private_no_forwards.down.sql
Normal file
2
deploy/migrations/0159_private_no_forwards.down.sql
Normal file
|
|
@ -0,0 +1,2 @@
|
||||||
|
DROP TABLE IF EXISTS private_no_forwards_requests;
|
||||||
|
DROP TABLE IF EXISTS private_no_forwards_chats;
|
||||||
36
deploy/migrations/0159_private_no_forwards.up.sql
Normal file
36
deploy/migrations/0159_private_no_forwards.up.sql
Normal file
|
|
@ -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;
|
||||||
|
|
@ -56,8 +56,9 @@ const tdesktopClient = "tdesktop"
|
||||||
// WebK directly calls Array.some on fragment_prefixes while rendering user profiles,
|
// 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.
|
// 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 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 提供客户端启动配置与国家区号目录。
|
// Service 提供客户端启动配置与国家区号目录。
|
||||||
//
|
//
|
||||||
|
|
@ -117,14 +118,14 @@ func defaultAppConfig(mapboxToken string) domain.AppConfig {
|
||||||
|
|
||||||
func defaultAppConfigJSON(mapboxToken string) []byte {
|
func defaultAppConfigJSON(mapboxToken string) []byte {
|
||||||
if mapboxToken == "" {
|
if mapboxToken == "" {
|
||||||
return []byte(tdesktopDefaultAppConfigBase + `}`)
|
return []byte(tdesktopDefaultAppConfigBase + tdesktopNoForwardsAppConfig + `}`)
|
||||||
}
|
}
|
||||||
token, err := json.Marshal(mapboxToken)
|
token, err := json.Marshal(mapboxToken)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return []byte(tdesktopDefaultAppConfigBase + `}`)
|
return []byte(tdesktopDefaultAppConfigBase + tdesktopNoForwardsAppConfig + `}`)
|
||||||
}
|
}
|
||||||
tokenJSON := string(token)
|
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 {
|
func defaultAppConfigHashFor(mapboxToken string) int {
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,8 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"telesrv/internal/domain"
|
||||||
)
|
)
|
||||||
|
|
||||||
// TestAppConfigPremiumKeys 断言 premium / Stars 相关 key 完整下发且 hash 已递增:
|
// TestAppConfigPremiumKeys 断言 premium / Stars 相关 key 完整下发且 hash 已递增:
|
||||||
|
|
@ -26,6 +28,9 @@ func TestAppConfigPremiumKeys(t *testing.T) {
|
||||||
if err := json.Unmarshal(cfg.JSON, &decoded); err != nil {
|
if err := json.Unmarshal(cfg.JSON, &decoded); err != nil {
|
||||||
t.Fatalf("app config json invalid: %v", err)
|
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 {
|
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"])
|
t.Fatalf("premium_purchase_blocked = %v, want false (star gift 送礼入口耦合此 flag)", decoded["premium_purchase_blocked"])
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -222,6 +222,42 @@ func (s *Service) SetChatTheme(ctx context.Context, userID int64, req domain.Set
|
||||||
return out, nil
|
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 {
|
func chatThemeServiceMedia(emoticon string) *domain.MessageMedia {
|
||||||
return &domain.MessageMedia{
|
return &domain.MessageMedia{
|
||||||
Kind: domain.MessageMediaKindService,
|
Kind: domain.MessageMediaKindService,
|
||||||
|
|
|
||||||
|
|
@ -558,6 +558,10 @@ const (
|
||||||
// MessageServiceActionSetChatTheme 映射 messageActionSetChatTheme:
|
// MessageServiceActionSetChatTheme 映射 messageActionSetChatTheme:
|
||||||
// 私聊双方共享的 chat theme token 变更。
|
// 私聊双方共享的 chat theme token 变更。
|
||||||
MessageServiceActionSetChatTheme MessageServiceActionKind = "set_chat_theme"
|
MessageServiceActionSetChatTheme MessageServiceActionKind = "set_chat_theme"
|
||||||
|
// MessageServiceActionNoForwardsToggle / Request 映射私聊内容保护的
|
||||||
|
// 状态切换与关闭请求。会话级保护不能写入普通消息的 NoForwards 字段。
|
||||||
|
MessageServiceActionNoForwardsToggle MessageServiceActionKind = "no_forwards_toggle"
|
||||||
|
MessageServiceActionNoForwardsRequest MessageServiceActionKind = "no_forwards_request"
|
||||||
// MessageServiceActionStarGift 映射 messageActionStarGift:收到一份 Star 礼物。
|
// MessageServiceActionStarGift 映射 messageActionStarGift:收到一份 Star 礼物。
|
||||||
// 礼物快照(贴纸/星价)内嵌在 action 里,收礼人无需额外拉取即可渲染气泡。
|
// 礼物快照(贴纸/星价)内嵌在 action 里,收礼人无需额外拉取即可渲染气泡。
|
||||||
MessageServiceActionStarGift MessageServiceActionKind = "star_gift"
|
MessageServiceActionStarGift MessageServiceActionKind = "star_gift"
|
||||||
|
|
@ -629,6 +633,15 @@ type MessageRequestedPeerDetails struct {
|
||||||
Photo *Photo `json:"photo,omitempty"`
|
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 是私聊服务消息动作的协议中立表示。
|
// MessageServiceAction 是私聊服务消息动作的协议中立表示。
|
||||||
type MessageServiceAction struct {
|
type MessageServiceAction struct {
|
||||||
Kind MessageServiceActionKind `json:"kind"`
|
Kind MessageServiceActionKind `json:"kind"`
|
||||||
|
|
@ -639,6 +652,7 @@ type MessageServiceAction struct {
|
||||||
WebViewData *MessageWebViewDataAction `json:"web_view_data,omitempty"`
|
WebViewData *MessageWebViewDataAction `json:"web_view_data,omitempty"`
|
||||||
RequestedPeer *MessageRequestedPeerAction `json:"requested_peer,omitempty"`
|
RequestedPeer *MessageRequestedPeerAction `json:"requested_peer,omitempty"`
|
||||||
ChatThemeEmoticon string `json:"chat_theme_emoticon,omitempty"`
|
ChatThemeEmoticon string `json:"chat_theme_emoticon,omitempty"`
|
||||||
|
NoForwards *MessageNoForwardsAction `json:"no_forwards,omitempty"`
|
||||||
StarGift *MessageStarGiftAction `json:"star_gift,omitempty"`
|
StarGift *MessageStarGiftAction `json:"star_gift,omitempty"`
|
||||||
StarGiftUnique *MessageStarGiftUniqueAction `json:"star_gift_unique,omitempty"`
|
StarGiftUnique *MessageStarGiftUniqueAction `json:"star_gift_unique,omitempty"`
|
||||||
StarGiftOffer *MessageStarGiftOfferAction `json:"star_gift_offer,omitempty"`
|
StarGiftOffer *MessageStarGiftOfferAction `json:"star_gift_offer,omitempty"`
|
||||||
|
|
|
||||||
|
|
@ -407,6 +407,47 @@ type ForwardPrivateMessagesResult struct {
|
||||||
ReplayDeleteEvents []*UpdateEvent
|
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 命令。
|
// ReadHistoryRequest 是账号视角的 messages.readHistory 命令。
|
||||||
type ReadHistoryRequest struct {
|
type ReadHistoryRequest struct {
|
||||||
OwnerUserID int64
|
OwnerUserID int64
|
||||||
|
|
|
||||||
|
|
@ -27,6 +27,7 @@ var (
|
||||||
ErrLoginCodeDeliveryCommitAmbiguous = errors.New("login code delivery commit ambiguous")
|
ErrLoginCodeDeliveryCommitAmbiguous = errors.New("login code delivery commit ambiguous")
|
||||||
ErrReplyMessageIDInvalid = errors.New("reply message id invalid")
|
ErrReplyMessageIDInvalid = errors.New("reply message id invalid")
|
||||||
ErrChatForwardsRestricted = errors.New("chat forwards restricted")
|
ErrChatForwardsRestricted = errors.New("chat forwards restricted")
|
||||||
|
ErrNoForwardsRequestExpired = errors.New("no forwards request expired")
|
||||||
// ErrPinnedSavedDialogsTooMuch 映射 PINNED_TOO_MUCH:收藏夹子会话置顶
|
// ErrPinnedSavedDialogsTooMuch 映射 PINNED_TOO_MUCH:收藏夹子会话置顶
|
||||||
// 数量达到 MaxPinnedSavedDialogs 上限。
|
// 数量达到 MaxPinnedSavedDialogs 上限。
|
||||||
ErrPinnedSavedDialogsTooMuch = errors.New("pinned saved dialogs too much")
|
ErrPinnedSavedDialogsTooMuch = errors.New("pinned saved dialogs too much")
|
||||||
|
|
|
||||||
|
|
@ -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) {
|
func (r *Router) onMessagesSetChatAvailableReactions(ctx context.Context, req *tg.MessagesSetChatAvailableReactionsRequest) (tg.UpdatesClass, error) {
|
||||||
if r.deps.Channels == nil {
|
if r.deps.Channels == nil {
|
||||||
return nil, notImplementedErr()
|
return nil, notImplementedErr()
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,10 @@
|
||||||
package rpc
|
package rpc
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/iamxvbaba/td/tg"
|
"github.com/iamxvbaba/td/tg"
|
||||||
|
|
||||||
"telesrv/internal/domain"
|
"telesrv/internal/domain"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -145,6 +148,25 @@ func tgMessageServiceAction(msg domain.Message) tg.MessageActionClass {
|
||||||
return &tg.MessageActionSetChatTheme{
|
return &tg.MessageActionSetChatTheme{
|
||||||
Theme: &tg.ChatTheme{Emoticon: m.ServiceAction.ChatThemeEmoticon},
|
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:
|
case domain.MessageServiceActionPhoneCall:
|
||||||
if m.ServiceAction.Call == nil {
|
if m.ServiceAction.Call == nil {
|
||||||
return &tg.MessageActionEmpty{}
|
return &tg.MessageActionEmpty{}
|
||||||
|
|
|
||||||
|
|
@ -604,6 +604,13 @@ type MessagesService interface {
|
||||||
DeleteSavedHistory(ctx context.Context, userID int64, req domain.DeleteSavedHistoryRequest) (domain.DeleteSavedHistoryResult, error)
|
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
|
// TranslationService owns read-only translation and the durable per-account
|
||||||
// peer preference. It only exposes domain values to the RPC edge.
|
// peer preference. It only exposes domain values to the RPC edge.
|
||||||
type TranslationService interface {
|
type TranslationService interface {
|
||||||
|
|
|
||||||
|
|
@ -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 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 inputRequestInvalidErr() error { return tgerr.New(400, "INPUT_REQUEST_INVALID") }
|
||||||
|
|
||||||
func inputRequestTooLongErr() error { return tgerr.New(400, "INPUT_REQUEST_TOO_LONG") }
|
func inputRequestTooLongErr() error { return tgerr.New(400, "INPUT_REQUEST_TOO_LONG") }
|
||||||
|
|
|
||||||
|
|
@ -515,6 +515,15 @@ func (r *Router) forwardSourcesFromPrivateMessages(ctx context.Context, userID i
|
||||||
if fromPeer.Type != domain.PeerTypeUser || fromPeer.ID == 0 {
|
if fromPeer.Type != domain.PeerTypeUser || fromPeer.ID == 0 {
|
||||||
return nil, domain.ErrMessageIDInvalid
|
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))
|
byID := make(map[int]domain.Message, len(messages))
|
||||||
for _, msg := range messages {
|
for _, msg := range messages {
|
||||||
byID[msg.ID] = msg
|
byID[msg.ID] = msg
|
||||||
|
|
|
||||||
144
internal/rpc/messages_no_forwards.go
Normal file
144
internal/rpc/messages_no_forwards.go
Normal file
|
|
@ -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)
|
||||||
|
}
|
||||||
167
internal/rpc/messages_no_forwards_rpc_test.go
Normal file
167
internal/rpc/messages_no_forwards_rpc_test.go
Normal file
|
|
@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -265,6 +265,17 @@ func (r *Router) buildUserFullProjection(ctx context.Context, currentUserID int6
|
||||||
Settings: tg.PeerSettings{},
|
Settings: tg.PeerSettings{},
|
||||||
NotifySettings: *tdesktop.NotifySettings(),
|
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_available=true 不显示通话按钮(P1 前置项)。
|
||||||
// phone_calls_private 标记对端禁 P2P(p2p_allowed 真值在通话确认时另行计算)。
|
// phone_calls_private 标记对端禁 P2P(p2p_allowed 真值在通话确认时另行计算)。
|
||||||
if !u.Bot && u.ID != currentUserID {
|
if !u.Bot && u.ID != currentUserID {
|
||||||
|
|
|
||||||
|
|
@ -20,6 +20,9 @@ func (s *MessageStore) ForwardPrivateMessages(ctx context.Context, req domain.Fo
|
||||||
if req.Date == 0 {
|
if req.Date == 0 {
|
||||||
req.Date = int(time.Now().Unix())
|
req.Date = int(time.Now().Unix())
|
||||||
}
|
}
|
||||||
|
if s.privateNoForwardsEnabled(req.OwnerUserID, req.FromPeer.ID) {
|
||||||
|
return res, domain.ErrChatForwardsRestricted
|
||||||
|
}
|
||||||
s.mu.RLock()
|
s.mu.RLock()
|
||||||
sources := make([]domain.Message, 0, len(req.MessageIDs))
|
sources := make([]domain.Message, 0, len(req.MessageIDs))
|
||||||
for _, id := range req.MessageIDs {
|
for _, id := range req.MessageIDs {
|
||||||
|
|
|
||||||
|
|
@ -114,10 +114,18 @@ func cloneRequestedPeerMedia(media *domain.MessageMedia) *domain.MessageMedia {
|
||||||
video.Attributes = append([]domain.DocumentAttribute(nil), media.LivePhotoVideo.Attributes...)
|
video.Attributes = append([]domain.DocumentAttribute(nil), media.LivePhotoVideo.Attributes...)
|
||||||
clone.LivePhotoVideo = &video
|
clone.LivePhotoVideo = &video
|
||||||
}
|
}
|
||||||
if media.ServiceAction == nil || media.ServiceAction.RequestedPeer == nil {
|
if media.ServiceAction == nil {
|
||||||
return &clone
|
return &clone
|
||||||
}
|
}
|
||||||
action := *media.ServiceAction
|
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 := *media.ServiceAction.RequestedPeer
|
||||||
requested.Peers = append([]domain.Peer(nil), requested.Peers...)
|
requested.Peers = append([]domain.Peer(nil), requested.Peers...)
|
||||||
requested.Details = append([]domain.MessageRequestedPeerDetails(nil), requested.Details...)
|
requested.Details = append([]domain.MessageRequestedPeerDetails(nil), requested.Details...)
|
||||||
|
|
|
||||||
200
internal/store/memory/message_no_forwards.go
Normal file
200
internal/store/memory/message_no_forwards.go
Normal file
|
|
@ -0,0 +1,200 @@
|
||||||
|
package memory
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"telesrv/internal/domain"
|
||||||
|
)
|
||||||
|
|
||||||
|
type privateNoForwardsPair struct {
|
||||||
|
low int64
|
||||||
|
high int64
|
||||||
|
}
|
||||||
|
|
||||||
|
type memoryNoForwardsRequest struct {
|
||||||
|
privateMessageID int64
|
||||||
|
requesterUserID int64
|
||||||
|
responderUserID int64
|
||||||
|
expiresAt int
|
||||||
|
handled bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func noForwardsPair(a, b int64) (privateNoForwardsPair, bool) {
|
||||||
|
if a <= 0 || b <= 0 || a == b {
|
||||||
|
return privateNoForwardsPair{}, false
|
||||||
|
}
|
||||||
|
if a > b {
|
||||||
|
a, b = b, a
|
||||||
|
}
|
||||||
|
return privateNoForwardsPair{low: a, high: b}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *MessageStore) GetPrivateNoForwards(_ context.Context, viewerUserID, peerUserID int64) (domain.PrivateNoForwardsState, error) {
|
||||||
|
pair, ok := noForwardsPair(viewerUserID, peerUserID)
|
||||||
|
if !ok {
|
||||||
|
return domain.PrivateNoForwardsState{}, domain.ErrMessageIDInvalid
|
||||||
|
}
|
||||||
|
s.noForwardsMu.Lock()
|
||||||
|
defer s.noForwardsMu.Unlock()
|
||||||
|
state := s.privateNoForwards[pair]
|
||||||
|
state.UserLowID, state.UserHighID = pair.low, pair.high
|
||||||
|
return state, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *MessageStore) TogglePrivateNoForwards(ctx context.Context, req domain.TogglePrivateNoForwardsRequest) (domain.TogglePrivateNoForwardsResult, error) {
|
||||||
|
pair, ok := noForwardsPair(req.ActorUserID, req.PeerUserID)
|
||||||
|
if !ok || req.RequestMsgID < 0 || req.RequestMsgID > domain.MaxMessageBoxID {
|
||||||
|
return domain.TogglePrivateNoForwardsResult{}, domain.ErrMessageIDInvalid
|
||||||
|
}
|
||||||
|
if req.Date == 0 {
|
||||||
|
req.Date = int(time.Now().Unix())
|
||||||
|
}
|
||||||
|
if req.RandomID == 0 {
|
||||||
|
req.RandomID = time.Now().UnixNano()
|
||||||
|
if req.RandomID == 0 {
|
||||||
|
req.RandomID = 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
s.noForwardsMu.Lock()
|
||||||
|
defer s.noForwardsMu.Unlock()
|
||||||
|
|
||||||
|
state := s.privateNoForwards[pair]
|
||||||
|
state.UserLowID, state.UserHighID = pair.low, pair.high
|
||||||
|
previousEnabled := state.Enabled()
|
||||||
|
var (
|
||||||
|
kind domain.MessageServiceActionKind
|
||||||
|
action domain.MessageNoForwardsAction
|
||||||
|
requestRecord *memoryNoForwardsRequest
|
||||||
|
requestUID int64
|
||||||
|
)
|
||||||
|
|
||||||
|
if req.RequestMsgID != 0 {
|
||||||
|
s.mu.RLock()
|
||||||
|
var source domain.Message
|
||||||
|
for _, msg := range s.m[req.ActorUserID] {
|
||||||
|
if msg.ID == req.RequestMsgID && msg.Peer == (domain.Peer{Type: domain.PeerTypeUser, ID: req.PeerUserID}) {
|
||||||
|
source = msg
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if source.ID != 0 {
|
||||||
|
record := s.privateNoForwardsRequests[source.UID]
|
||||||
|
requestRecord = &record
|
||||||
|
requestUID = source.UID
|
||||||
|
}
|
||||||
|
s.mu.RUnlock()
|
||||||
|
if source.ID == 0 || requestRecord == nil || requestRecord.privateMessageID != source.UID ||
|
||||||
|
requestRecord.requesterUserID != req.PeerUserID || requestRecord.responderUserID != req.ActorUserID ||
|
||||||
|
requestRecord.handled || requestRecord.expiresAt <= req.Date {
|
||||||
|
return domain.TogglePrivateNoForwardsResult{}, domain.ErrNoForwardsRequestExpired
|
||||||
|
}
|
||||||
|
kind = domain.MessageServiceActionNoForwardsToggle
|
||||||
|
action = domain.MessageNoForwardsAction{PrevValue: previousEnabled, NewValue: req.Enabled}
|
||||||
|
if req.Enabled {
|
||||||
|
state.EnabledByUserID = req.ActorUserID
|
||||||
|
} else {
|
||||||
|
state.EnabledByUserID = 0
|
||||||
|
}
|
||||||
|
} else if req.Enabled {
|
||||||
|
if state.EnabledByUserID != 0 {
|
||||||
|
return domain.TogglePrivateNoForwardsResult{State: state}, nil
|
||||||
|
}
|
||||||
|
kind = domain.MessageServiceActionNoForwardsToggle
|
||||||
|
action = domain.MessageNoForwardsAction{PrevValue: false, NewValue: true}
|
||||||
|
state.EnabledByUserID = req.ActorUserID
|
||||||
|
} else {
|
||||||
|
switch state.EnabledByUserID {
|
||||||
|
case 0:
|
||||||
|
return domain.TogglePrivateNoForwardsResult{State: state}, nil
|
||||||
|
case req.ActorUserID:
|
||||||
|
kind = domain.MessageServiceActionNoForwardsToggle
|
||||||
|
action = domain.MessageNoForwardsAction{PrevValue: true, NewValue: false}
|
||||||
|
state.EnabledByUserID = 0
|
||||||
|
default:
|
||||||
|
kind = domain.MessageServiceActionNoForwardsRequest
|
||||||
|
action = domain.MessageNoForwardsAction{
|
||||||
|
PrevValue: true,
|
||||||
|
NewValue: false,
|
||||||
|
ExpiresAt: req.Date + domain.PrivateNoForwardsRequestExpirePeriod,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
reply := (*domain.MessageReply)(nil)
|
||||||
|
if req.RequestMsgID != 0 {
|
||||||
|
reply = &domain.MessageReply{
|
||||||
|
MessageID: req.RequestMsgID,
|
||||||
|
Peer: domain.Peer{Type: domain.PeerTypeUser, ID: req.PeerUserID},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
send, err := s.SendPrivateText(ctx, domain.SendPrivateTextRequest{
|
||||||
|
SenderUserID: req.ActorUserID,
|
||||||
|
RecipientUserID: req.PeerUserID,
|
||||||
|
RandomID: req.RandomID,
|
||||||
|
Silent: true,
|
||||||
|
Date: req.Date,
|
||||||
|
OriginAuthKeyID: req.OriginAuthKeyID,
|
||||||
|
OriginSessionID: req.OriginSessionID,
|
||||||
|
ReplyTo: reply,
|
||||||
|
Media: &domain.MessageMedia{
|
||||||
|
Kind: domain.MessageMediaKindService,
|
||||||
|
ServiceAction: &domain.MessageServiceAction{
|
||||||
|
Kind: kind,
|
||||||
|
NoForwards: &action,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
if req.RequestMsgID != 0 && err == domain.ErrReplyMessageIDInvalid {
|
||||||
|
return domain.TogglePrivateNoForwardsResult{}, domain.ErrNoForwardsRequestExpired
|
||||||
|
}
|
||||||
|
return domain.TogglePrivateNoForwardsResult{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
s.privateNoForwards[pair] = state
|
||||||
|
if kind == domain.MessageServiceActionNoForwardsRequest {
|
||||||
|
s.privateNoForwardsRequests[send.SenderMessage.UID] = memoryNoForwardsRequest{
|
||||||
|
privateMessageID: send.SenderMessage.UID,
|
||||||
|
requesterUserID: req.ActorUserID,
|
||||||
|
responderUserID: req.PeerUserID,
|
||||||
|
expiresAt: action.ExpiresAt,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if requestUID != 0 {
|
||||||
|
record := s.privateNoForwardsRequests[requestUID]
|
||||||
|
record.handled = true
|
||||||
|
s.privateNoForwardsRequests[requestUID] = record
|
||||||
|
s.markNoForwardsRequestExpired(requestUID)
|
||||||
|
}
|
||||||
|
return domain.TogglePrivateNoForwardsResult{State: state, Changed: true, Send: send}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *MessageStore) markNoForwardsRequestExpired(privateMessageID int64) {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
for ownerID, messages := range s.m {
|
||||||
|
for i := range messages {
|
||||||
|
action := messages[i].Media
|
||||||
|
if messages[i].UID != privateMessageID || action == nil || action.ServiceAction == nil ||
|
||||||
|
action.ServiceAction.Kind != domain.MessageServiceActionNoForwardsRequest ||
|
||||||
|
action.ServiceAction.NoForwards == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
messages[i].Media = cloneRequestedPeerMedia(messages[i].Media)
|
||||||
|
messages[i].Media.ServiceAction.NoForwards.Expired = true
|
||||||
|
}
|
||||||
|
s.m[ownerID] = messages
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *MessageStore) privateNoForwardsEnabled(a, b int64) bool {
|
||||||
|
pair, ok := noForwardsPair(a, b)
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
s.noForwardsMu.Lock()
|
||||||
|
defer s.noForwardsMu.Unlock()
|
||||||
|
return s.privateNoForwards[pair].Enabled()
|
||||||
|
}
|
||||||
162
internal/store/memory/message_no_forwards_test.go
Normal file
162
internal/store/memory/message_no_forwards_test.go
Normal file
|
|
@ -0,0 +1,162 @@
|
||||||
|
package memory
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"telesrv/internal/domain"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestPrivateNoForwardsStateMachineAndForwardGate(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
messages := NewMessageStore()
|
||||||
|
const alice, bob int64 = 1001, 1002
|
||||||
|
|
||||||
|
enable, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: alice, PeerUserID: bob, Enabled: true, RandomID: 11, Date: 100,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("enable: %v", err)
|
||||||
|
}
|
||||||
|
if !enable.Changed || enable.State.EnabledByUserID != alice ||
|
||||||
|
enable.Send.SenderMessage.Pts != 1 || enable.Send.RecipientMessage.Pts != 1 ||
|
||||||
|
enable.Send.SenderMessage.NoForwards {
|
||||||
|
t.Fatalf("enable result = %+v", enable)
|
||||||
|
}
|
||||||
|
assertMemoryNoForwardsAction(t, enable.Send.SenderMessage, domain.MessageServiceActionNoForwardsToggle, false, true, false)
|
||||||
|
|
||||||
|
repeat, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: alice, PeerUserID: bob, Enabled: true, RandomID: 12, Date: 101,
|
||||||
|
})
|
||||||
|
if err != nil || repeat.Changed || repeat.State.EnabledByUserID != alice {
|
||||||
|
t.Fatalf("repeat enable = %+v err=%v, want no-op", repeat, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
request, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: bob, PeerUserID: alice, Enabled: false, RandomID: 13, Date: 102,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("request disable: %v", err)
|
||||||
|
}
|
||||||
|
if request.State.EnabledByUserID != alice || request.Send.SenderMessage.Pts != 2 ||
|
||||||
|
request.Send.RecipientMessage.Pts != 2 {
|
||||||
|
t.Fatalf("request result = %+v", request)
|
||||||
|
}
|
||||||
|
assertMemoryNoForwardsAction(t, request.Send.SenderMessage, domain.MessageServiceActionNoForwardsRequest, true, false, false)
|
||||||
|
|
||||||
|
answer, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: alice,
|
||||||
|
PeerUserID: bob,
|
||||||
|
Enabled: false,
|
||||||
|
RequestMsgID: request.Send.RecipientMessage.ID,
|
||||||
|
RandomID: 14,
|
||||||
|
Date: 103,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("accept request: %v", err)
|
||||||
|
}
|
||||||
|
if answer.State.Enabled() || answer.Send.SenderMessage.Pts != 3 || answer.Send.RecipientMessage.Pts != 3 {
|
||||||
|
t.Fatalf("answer result = %+v", answer)
|
||||||
|
}
|
||||||
|
if answer.Send.SenderMessage.ReplyTo == nil ||
|
||||||
|
answer.Send.SenderMessage.ReplyTo.MessageID != request.Send.RecipientMessage.ID ||
|
||||||
|
answer.Send.RecipientMessage.ReplyTo == nil ||
|
||||||
|
answer.Send.RecipientMessage.ReplyTo.MessageID != request.Send.SenderMessage.ID {
|
||||||
|
t.Fatalf("answer reply mapping sender=%+v recipient=%+v", answer.Send.SenderMessage.ReplyTo, answer.Send.RecipientMessage.ReplyTo)
|
||||||
|
}
|
||||||
|
assertMemoryNoForwardsAction(t, answer.Send.SenderMessage, domain.MessageServiceActionNoForwardsToggle, true, false, false)
|
||||||
|
|
||||||
|
aliceHistory, err := messages.ListByUser(ctx, alice, domain.MessageFilter{
|
||||||
|
HasPeer: true, Peer: domain.Peer{Type: domain.PeerTypeUser, ID: bob}, Limit: 20,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("alice history: %v", err)
|
||||||
|
}
|
||||||
|
var expired bool
|
||||||
|
for _, msg := range aliceHistory.Messages {
|
||||||
|
if msg.ID == request.Send.RecipientMessage.ID && msg.Media != nil && msg.Media.ServiceAction != nil &&
|
||||||
|
msg.Media.ServiceAction.NoForwards != nil {
|
||||||
|
expired = msg.Media.ServiceAction.NoForwards.Expired
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !expired {
|
||||||
|
t.Fatal("handled request was not projected expired")
|
||||||
|
}
|
||||||
|
if _, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: alice, PeerUserID: bob, RequestMsgID: request.Send.RecipientMessage.ID,
|
||||||
|
RandomID: 15, Date: 104,
|
||||||
|
}); !errors.Is(err, domain.ErrNoForwardsRequestExpired) {
|
||||||
|
t.Fatalf("repeat answer err=%v, want ErrNoForwardsRequestExpired", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
source, err := messages.SendPrivateText(ctx, domain.SendPrivateTextRequest{
|
||||||
|
SenderUserID: alice, RecipientUserID: bob, RandomID: 20, Message: "source", Date: 105,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("send source: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: alice, PeerUserID: bob, Enabled: true, RandomID: 21, Date: 106,
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("re-enable: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := messages.ForwardPrivateMessages(ctx, domain.ForwardPrivateMessagesRequest{
|
||||||
|
OwnerUserID: alice,
|
||||||
|
FromPeer: domain.Peer{Type: domain.PeerTypeUser, ID: bob},
|
||||||
|
ToUserID: alice,
|
||||||
|
MessageIDs: []int{source.SenderMessage.ID},
|
||||||
|
RandomIDs: []int64{22},
|
||||||
|
Date: 107,
|
||||||
|
}); !errors.Is(err, domain.ErrChatForwardsRestricted) {
|
||||||
|
t.Fatalf("forward protected chat err=%v, want ErrChatForwardsRestricted", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPrivateNoForwardsRequestExpiresWithoutPTS(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
messages := NewMessageStore()
|
||||||
|
const alice, bob int64 = 2001, 2002
|
||||||
|
if _, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: alice, PeerUserID: bob, Enabled: true, RandomID: 31, Date: 200,
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
request, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: bob, PeerUserID: alice, RandomID: 32, Date: 201,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: alice,
|
||||||
|
PeerUserID: bob,
|
||||||
|
RequestMsgID: request.Send.RecipientMessage.ID,
|
||||||
|
RandomID: 33,
|
||||||
|
Date: 201 + domain.PrivateNoForwardsRequestExpirePeriod,
|
||||||
|
}); !errors.Is(err, domain.ErrNoForwardsRequestExpired) {
|
||||||
|
t.Fatalf("expired answer err=%v", err)
|
||||||
|
}
|
||||||
|
state, _ := messages.GetPrivateNoForwards(ctx, alice, bob)
|
||||||
|
if state.EnabledByUserID != alice {
|
||||||
|
t.Fatalf("expired answer changed state = %+v", state)
|
||||||
|
}
|
||||||
|
history, _ := messages.ListByUser(ctx, alice, domain.MessageFilter{
|
||||||
|
HasPeer: true, Peer: domain.Peer{Type: domain.PeerTypeUser, ID: bob}, Limit: 20,
|
||||||
|
})
|
||||||
|
if len(history.Messages) != 2 || history.Messages[0].Pts != 2 {
|
||||||
|
t.Fatalf("expired answer allocated message/pts: %+v", history.Messages)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func assertMemoryNoForwardsAction(t *testing.T, msg domain.Message, kind domain.MessageServiceActionKind, prev, next, expired bool) {
|
||||||
|
t.Helper()
|
||||||
|
if msg.Media == nil || msg.Media.ServiceAction == nil || msg.Media.ServiceAction.Kind != kind ||
|
||||||
|
msg.Media.ServiceAction.NoForwards == nil {
|
||||||
|
t.Fatalf("message action = %+v, want %s", msg.Media, kind)
|
||||||
|
}
|
||||||
|
action := msg.Media.ServiceAction.NoForwards
|
||||||
|
if action.PrevValue != prev || action.NewValue != next || action.Expired != expired {
|
||||||
|
t.Fatalf("action = %+v, want prev=%v new=%v expired=%v", action, prev, next, expired)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -8,6 +8,7 @@ import (
|
||||||
// MessageStore 是 store.MessageStore 的内存实现。
|
// MessageStore 是 store.MessageStore 的内存实现。
|
||||||
type MessageStore struct {
|
type MessageStore struct {
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
|
noForwardsMu sync.Mutex
|
||||||
m map[int64][]domain.Message
|
m map[int64][]domain.Message
|
||||||
nextUID int64
|
nextUID int64
|
||||||
nextBox map[int64]int
|
nextBox map[int64]int
|
||||||
|
|
@ -24,6 +25,11 @@ type MessageStore struct {
|
||||||
polls *PollStore
|
polls *PollStore
|
||||||
// savedPins 是收藏夹子会话置顶顺序(下标即 pinned_order,越小越前)。
|
// savedPins 是收藏夹子会话置顶顺序(下标即 pinned_order,越小越前)。
|
||||||
savedPins map[int64][]domain.Peer
|
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 共用同一实例)。
|
// AttachPollStore 注入共享 poll 权威(与 ChannelStore 共用同一实例)。
|
||||||
|
|
@ -40,18 +46,20 @@ type readOutboxDateKey struct {
|
||||||
// NewMessageStore 创建内存 MessageStore。
|
// NewMessageStore 创建内存 MessageStore。
|
||||||
func NewMessageStore(dialogs ...*DialogStore) *MessageStore {
|
func NewMessageStore(dialogs ...*DialogStore) *MessageStore {
|
||||||
s := &MessageStore{
|
s := &MessageStore{
|
||||||
m: make(map[int64][]domain.Message),
|
m: make(map[int64][]domain.Message),
|
||||||
nextUID: 1,
|
nextUID: 1,
|
||||||
nextBox: make(map[int64]int),
|
nextBox: make(map[int64]int),
|
||||||
nextPts: make(map[int64]int),
|
nextPts: make(map[int64]int),
|
||||||
readOutboxDates: make(map[readOutboxDateKey]int),
|
readOutboxDates: make(map[readOutboxDateKey]int),
|
||||||
privateReactions: make(map[int64]map[int64][]domain.ChannelMessagePeerReaction),
|
privateReactions: make(map[int64]map[int64][]domain.ChannelMessagePeerReaction),
|
||||||
savedMessageTags: make(map[int64]map[int][]domain.MessageReaction),
|
savedMessageTags: make(map[int64]map[int][]domain.MessageReaction),
|
||||||
savedTagTitles: make(map[int64]map[string]string),
|
savedTagTitles: make(map[int64]map[string]string),
|
||||||
privateSendDedup: make(map[privateSendDedupKey]privateSendDedupRecord),
|
privateSendDedup: make(map[privateSendDedupKey]privateSendDedupRecord),
|
||||||
loginCodeDeliveries: make(map[[32]byte]loginCodeDeliveryRecord),
|
loginCodeDeliveries: make(map[[32]byte]loginCodeDeliveryRecord),
|
||||||
albumGroups: make(map[albumGroupKey]albumGroupRecord),
|
albumGroups: make(map[albumGroupKey]albumGroupRecord),
|
||||||
savedPins: make(map[int64][]domain.Peer),
|
savedPins: make(map[int64][]domain.Peer),
|
||||||
|
privateNoForwards: make(map[privateNoForwardsPair]domain.PrivateNoForwardsState),
|
||||||
|
privateNoForwardsRequests: make(map[int64]memoryNoForwardsRequest),
|
||||||
}
|
}
|
||||||
if len(dialogs) > 0 {
|
if len(dialogs) > 0 {
|
||||||
s.dialogs = dialogs[0]
|
s.dialogs = dialogs[0]
|
||||||
|
|
|
||||||
|
|
@ -33,6 +33,13 @@ func (s *MessageStore) ForwardPrivateMessages(ctx context.Context, req domain.Fo
|
||||||
if req.Date == 0 {
|
if req.Date == 0 {
|
||||||
req.Date = int(time.Now().Unix())
|
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))
|
boxIDs := make([]int32, 0, len(req.MessageIDs))
|
||||||
for i, id := range req.MessageIDs {
|
for i, id := range req.MessageIDs {
|
||||||
if id <= 0 || id > domain.MaxMessageBoxID || req.RandomIDs[i] == 0 {
|
if id <= 0 || id > domain.MaxMessageBoxID || req.RandomIDs[i] == 0 {
|
||||||
|
|
|
||||||
248
internal/store/postgres/message_no_forwards.go
Normal file
248
internal/store/postgres/message_no_forwards.go
Normal file
|
|
@ -0,0 +1,248 @@
|
||||||
|
package postgres
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
|
||||||
|
"telesrv/internal/domain"
|
||||||
|
)
|
||||||
|
|
||||||
|
var errPrivateNoForwardsNoop = errors.New("private no forwards no-op")
|
||||||
|
|
||||||
|
func pgNoForwardsPair(a, b int64) (low, high int64, ok bool) {
|
||||||
|
if a <= 0 || b <= 0 || a == b {
|
||||||
|
return 0, 0, false
|
||||||
|
}
|
||||||
|
if a > b {
|
||||||
|
a, b = b, a
|
||||||
|
}
|
||||||
|
return a, b, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *MessageStore) GetPrivateNoForwards(ctx context.Context, viewerUserID, peerUserID int64) (domain.PrivateNoForwardsState, error) {
|
||||||
|
low, high, ok := pgNoForwardsPair(viewerUserID, peerUserID)
|
||||||
|
if !ok {
|
||||||
|
return domain.PrivateNoForwardsState{}, domain.ErrMessageIDInvalid
|
||||||
|
}
|
||||||
|
state := domain.PrivateNoForwardsState{UserLowID: low, UserHighID: high}
|
||||||
|
err := s.db.QueryRow(ctx, `
|
||||||
|
SELECT COALESCE(enabled_by_user_id, 0)
|
||||||
|
FROM private_no_forwards_chats
|
||||||
|
WHERE user_low_id = $1 AND user_high_id = $2`, low, high).Scan(&state.EnabledByUserID)
|
||||||
|
if errors.Is(err, pgx.ErrNoRows) {
|
||||||
|
return state, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return domain.PrivateNoForwardsState{}, fmt.Errorf("get private no forwards: %w", err)
|
||||||
|
}
|
||||||
|
return state, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *MessageStore) TogglePrivateNoForwards(ctx context.Context, req domain.TogglePrivateNoForwardsRequest) (domain.TogglePrivateNoForwardsResult, error) {
|
||||||
|
low, high, ok := pgNoForwardsPair(req.ActorUserID, req.PeerUserID)
|
||||||
|
if !ok || req.RequestMsgID < 0 || req.RequestMsgID > domain.MaxMessageBoxID {
|
||||||
|
return domain.TogglePrivateNoForwardsResult{}, domain.ErrMessageIDInvalid
|
||||||
|
}
|
||||||
|
if req.Date == 0 {
|
||||||
|
req.Date = int(time.Now().Unix())
|
||||||
|
}
|
||||||
|
if req.RandomID == 0 {
|
||||||
|
req.RandomID = time.Now().UnixNano()
|
||||||
|
if req.RandomID == 0 {
|
||||||
|
req.RandomID = 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
state := domain.PrivateNoForwardsState{UserLowID: low, UserHighID: high}
|
||||||
|
actionKind := domain.MessageServiceActionNoForwardsToggle
|
||||||
|
action := domain.MessageNoForwardsAction{}
|
||||||
|
var answeredRequestSenderID, answeredRequestMessageID int64
|
||||||
|
|
||||||
|
sendReq := domain.SendPrivateTextRequest{
|
||||||
|
SenderUserID: req.ActorUserID,
|
||||||
|
RecipientUserID: req.PeerUserID,
|
||||||
|
RandomID: req.RandomID,
|
||||||
|
Silent: true,
|
||||||
|
Date: req.Date,
|
||||||
|
OriginAuthKeyID: req.OriginAuthKeyID,
|
||||||
|
OriginSessionID: req.OriginSessionID,
|
||||||
|
// A non-empty placeholder is required before the send transaction starts.
|
||||||
|
// The pair-locked before-hook replaces it with the authoritative action.
|
||||||
|
Media: noForwardsServiceMedia(actionKind, action),
|
||||||
|
}
|
||||||
|
if req.RequestMsgID != 0 {
|
||||||
|
sendReq.ReplyTo = &domain.MessageReply{
|
||||||
|
MessageID: req.RequestMsgID,
|
||||||
|
Peer: domain.Peer{Type: domain.PeerTypeUser, ID: req.PeerUserID},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
hooks := privateSendTxHooks{
|
||||||
|
before: func(ctx context.Context, tx pgx.Tx, send *domain.SendPrivateTextRequest) error {
|
||||||
|
if _, err := tx.Exec(ctx, `
|
||||||
|
INSERT INTO private_no_forwards_chats (user_low_id, user_high_id)
|
||||||
|
VALUES ($1, $2)
|
||||||
|
ON CONFLICT (user_low_id, user_high_id) DO NOTHING`, low, high); err != nil {
|
||||||
|
return fmt.Errorf("ensure private no forwards state: %w", err)
|
||||||
|
}
|
||||||
|
if err := tx.QueryRow(ctx, `
|
||||||
|
SELECT COALESCE(enabled_by_user_id, 0)
|
||||||
|
FROM private_no_forwards_chats
|
||||||
|
WHERE user_low_id = $1 AND user_high_id = $2
|
||||||
|
FOR UPDATE`, low, high).Scan(&state.EnabledByUserID); err != nil {
|
||||||
|
return fmt.Errorf("lock private no forwards state: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
previousEnabled := state.Enabled()
|
||||||
|
actionKind = domain.MessageServiceActionNoForwardsToggle
|
||||||
|
action = domain.MessageNoForwardsAction{}
|
||||||
|
if req.RequestMsgID != 0 {
|
||||||
|
var expiresAt, handledAt int
|
||||||
|
err := tx.QueryRow(ctx, `
|
||||||
|
SELECT r.private_message_sender_user_id, r.private_message_id, r.expires_at, r.handled_at
|
||||||
|
FROM message_boxes AS b
|
||||||
|
JOIN private_no_forwards_requests AS r
|
||||||
|
ON r.private_message_sender_user_id = b.message_sender_id
|
||||||
|
AND r.private_message_id = b.private_message_id
|
||||||
|
WHERE b.owner_user_id = $1
|
||||||
|
AND b.box_id = $2
|
||||||
|
AND b.peer_type = 'user'
|
||||||
|
AND b.peer_id = $3
|
||||||
|
AND r.requester_user_id = $3
|
||||||
|
AND r.responder_user_id = $1
|
||||||
|
FOR UPDATE OF r`,
|
||||||
|
req.ActorUserID, req.RequestMsgID, req.PeerUserID,
|
||||||
|
).Scan(&answeredRequestSenderID, &answeredRequestMessageID, &expiresAt, &handledAt)
|
||||||
|
if errors.Is(err, pgx.ErrNoRows) {
|
||||||
|
return domain.ErrNoForwardsRequestExpired
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("lock private no forwards request: %w", err)
|
||||||
|
}
|
||||||
|
if handledAt != 0 || expiresAt <= req.Date {
|
||||||
|
return domain.ErrNoForwardsRequestExpired
|
||||||
|
}
|
||||||
|
action = domain.MessageNoForwardsAction{PrevValue: previousEnabled, NewValue: req.Enabled}
|
||||||
|
if req.Enabled {
|
||||||
|
state.EnabledByUserID = req.ActorUserID
|
||||||
|
} else {
|
||||||
|
state.EnabledByUserID = 0
|
||||||
|
}
|
||||||
|
if _, err := tx.Exec(ctx, `
|
||||||
|
UPDATE private_no_forwards_requests
|
||||||
|
SET handled_at = $3
|
||||||
|
WHERE private_message_sender_user_id = $1
|
||||||
|
AND private_message_id = $2
|
||||||
|
AND handled_at = 0`,
|
||||||
|
answeredRequestSenderID, answeredRequestMessageID, req.Date,
|
||||||
|
); err != nil {
|
||||||
|
return fmt.Errorf("handle private no forwards request: %w", err)
|
||||||
|
}
|
||||||
|
if err := expirePGNoForwardsRequest(ctx, tx, answeredRequestSenderID, answeredRequestMessageID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
} else if req.Enabled {
|
||||||
|
if state.EnabledByUserID != 0 {
|
||||||
|
return errPrivateNoForwardsNoop
|
||||||
|
}
|
||||||
|
action = domain.MessageNoForwardsAction{PrevValue: false, NewValue: true}
|
||||||
|
state.EnabledByUserID = req.ActorUserID
|
||||||
|
} else {
|
||||||
|
switch state.EnabledByUserID {
|
||||||
|
case 0:
|
||||||
|
return errPrivateNoForwardsNoop
|
||||||
|
case req.ActorUserID:
|
||||||
|
action = domain.MessageNoForwardsAction{PrevValue: true, NewValue: false}
|
||||||
|
state.EnabledByUserID = 0
|
||||||
|
default:
|
||||||
|
actionKind = domain.MessageServiceActionNoForwardsRequest
|
||||||
|
action = domain.MessageNoForwardsAction{
|
||||||
|
PrevValue: true,
|
||||||
|
NewValue: false,
|
||||||
|
ExpiresAt: req.Date + domain.PrivateNoForwardsRequestExpirePeriod,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var enabledBy any
|
||||||
|
if state.EnabledByUserID != 0 {
|
||||||
|
enabledBy = state.EnabledByUserID
|
||||||
|
}
|
||||||
|
if _, err := tx.Exec(ctx, `
|
||||||
|
UPDATE private_no_forwards_chats
|
||||||
|
SET enabled_by_user_id = $3, updated_at = now()
|
||||||
|
WHERE user_low_id = $1 AND user_high_id = $2`, low, high, enabledBy); err != nil {
|
||||||
|
return fmt.Errorf("update private no forwards state: %w", err)
|
||||||
|
}
|
||||||
|
send.Media = noForwardsServiceMedia(actionKind, action)
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
after: func(ctx context.Context, tx pgx.Tx, result domain.SendPrivateTextResult) error {
|
||||||
|
if actionKind != domain.MessageServiceActionNoForwardsRequest {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if _, err := tx.Exec(ctx, `
|
||||||
|
INSERT INTO private_no_forwards_requests (
|
||||||
|
private_message_sender_user_id,
|
||||||
|
private_message_id,
|
||||||
|
requester_user_id,
|
||||||
|
responder_user_id,
|
||||||
|
expires_at
|
||||||
|
) VALUES ($1, $2, $3, $4, $5)`,
|
||||||
|
req.ActorUserID, result.SenderMessage.UID, req.ActorUserID, req.PeerUserID, action.ExpiresAt,
|
||||||
|
); err != nil {
|
||||||
|
return fmt.Errorf("create private no forwards request: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
send, err := s.sendPrivateTextWithHooks(ctx, sendReq, hooks)
|
||||||
|
if errors.Is(err, errPrivateNoForwardsNoop) {
|
||||||
|
return domain.TogglePrivateNoForwardsResult{State: state}, nil
|
||||||
|
}
|
||||||
|
if errors.Is(err, domain.ErrReplyMessageIDInvalid) {
|
||||||
|
return domain.TogglePrivateNoForwardsResult{}, domain.ErrNoForwardsRequestExpired
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return domain.TogglePrivateNoForwardsResult{}, err
|
||||||
|
}
|
||||||
|
return domain.TogglePrivateNoForwardsResult{State: state, Changed: true, Send: send}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func noForwardsServiceMedia(kind domain.MessageServiceActionKind, action domain.MessageNoForwardsAction) *domain.MessageMedia {
|
||||||
|
return &domain.MessageMedia{
|
||||||
|
Kind: domain.MessageMediaKindService,
|
||||||
|
ServiceAction: &domain.MessageServiceAction{
|
||||||
|
Kind: kind,
|
||||||
|
NoForwards: &action,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func expirePGNoForwardsRequest(ctx context.Context, tx pgx.Tx, senderUserID, privateMessageID int64) error {
|
||||||
|
for _, statement := range []string{
|
||||||
|
`UPDATE private_messages
|
||||||
|
SET media = jsonb_set(media, '{service_action,no_forwards,expired}', 'true'::jsonb, true)
|
||||||
|
WHERE sender_user_id = $1 AND id = $2`,
|
||||||
|
`UPDATE message_boxes
|
||||||
|
SET media = jsonb_set(media, '{service_action,no_forwards,expired}', 'true'::jsonb, true)
|
||||||
|
WHERE message_sender_id = $1 AND private_message_id = $2`,
|
||||||
|
} {
|
||||||
|
tag, err := tx.Exec(ctx, statement, senderUserID, privateMessageID)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("expire private no forwards request: %w", err)
|
||||||
|
}
|
||||||
|
if tag.RowsAffected() == 0 {
|
||||||
|
return fmt.Errorf("expire private no forwards request: message disappeared")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *MessageStore) privateNoForwardsEnabled(ctx context.Context, a, b int64) (bool, error) {
|
||||||
|
state, err := s.GetPrivateNoForwards(ctx, a, b)
|
||||||
|
return state.Enabled(), err
|
||||||
|
}
|
||||||
229
internal/store/postgres/message_no_forwards_integration_test.go
Normal file
229
internal/store/postgres/message_no_forwards_integration_test.go
Normal file
|
|
@ -0,0 +1,229 @@
|
||||||
|
package postgres
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"telesrv/internal/domain"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestPostgresPrivateNoForwardsAtomicStateAndDifference(t *testing.T) {
|
||||||
|
pool := testPool(t)
|
||||||
|
ctx := context.Background()
|
||||||
|
suffix := randomSuffix(t)
|
||||||
|
users := NewUserStore(pool)
|
||||||
|
alice, err := users.Create(ctx, domain.User{AccessHash: 6101, Phone: "+1668" + suffix + "01", FirstName: "Alice"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
bob, err := users.Create(ctx, domain.User{AccessHash: 6102, Phone: "+1668" + suffix + "02", FirstName: "Bob"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
_, _ = pool.Exec(ctx, "DELETE FROM users WHERE id = ANY($1::bigint[])", []int64{alice.ID, bob.ID})
|
||||||
|
})
|
||||||
|
|
||||||
|
messages := NewMessageStore(pool)
|
||||||
|
baseRandom := time.Now().UnixNano()
|
||||||
|
enable, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: alice.ID, PeerUserID: bob.ID, Enabled: true, RandomID: baseRandom, Date: 1700100000,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("enable: %v", err)
|
||||||
|
}
|
||||||
|
request, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: bob.ID, PeerUserID: alice.ID, RandomID: baseRandom + 1, Date: 1700100001,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("request: %v", err)
|
||||||
|
}
|
||||||
|
answer, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: alice.ID, PeerUserID: bob.ID, RequestMsgID: request.Send.RecipientMessage.ID,
|
||||||
|
RandomID: baseRandom + 2, Date: 1700100002,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("answer: %v", err)
|
||||||
|
}
|
||||||
|
if enable.Send.SenderMessage.Pts != 1 || request.Send.SenderMessage.Pts != 2 ||
|
||||||
|
answer.Send.SenderMessage.Pts != 3 || answer.Send.SenderMessage.ReplyTo == nil ||
|
||||||
|
answer.Send.SenderMessage.ReplyTo.MessageID != request.Send.RecipientMessage.ID ||
|
||||||
|
answer.Send.RecipientMessage.ReplyTo == nil ||
|
||||||
|
answer.Send.RecipientMessage.ReplyTo.MessageID != request.Send.SenderMessage.ID {
|
||||||
|
t.Fatalf("pts/reply mapping enable=%+v request=%+v answer=%+v", enable.Send, request.Send, answer.Send)
|
||||||
|
}
|
||||||
|
state, err := messages.GetPrivateNoForwards(ctx, alice.ID, bob.ID)
|
||||||
|
if err != nil || state.Enabled() {
|
||||||
|
t.Fatalf("final state=%+v err=%v, want disabled", state, err)
|
||||||
|
}
|
||||||
|
if _, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: alice.ID, PeerUserID: bob.ID, RequestMsgID: request.Send.RecipientMessage.ID,
|
||||||
|
RandomID: baseRandom + 3, Date: 1700100003,
|
||||||
|
}); !errors.Is(err, domain.ErrNoForwardsRequestExpired) {
|
||||||
|
t.Fatalf("repeat answer err=%v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, userID := range []int64{alice.ID, bob.ID} {
|
||||||
|
events, err := NewUpdateEventStore(pool).ListAfter(ctx, userID, 0, 10)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("events user %d: %v", userID, err)
|
||||||
|
}
|
||||||
|
if len(events) != 3 || events[0].Pts != 1 || events[1].Pts != 2 || events[2].Pts != 3 {
|
||||||
|
t.Fatalf("events user %d = %+v, want continuous 1..3", userID, events)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var eventCount, outboxCount int
|
||||||
|
if err := pool.QueryRow(ctx, `
|
||||||
|
SELECT
|
||||||
|
(SELECT count(*) FROM user_update_events WHERE user_id = ANY($1::bigint[])),
|
||||||
|
(SELECT count(*) FROM dispatch_outbox WHERE target_user_id = ANY($1::bigint[]))`,
|
||||||
|
[]int64{alice.ID, bob.ID},
|
||||||
|
).Scan(&eventCount, &outboxCount); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if eventCount != 6 || outboxCount != 6 {
|
||||||
|
t.Fatalf("event/outbox count=%d/%d, want 6/6", eventCount, outboxCount)
|
||||||
|
}
|
||||||
|
var handledAt int
|
||||||
|
var logicalExpired bool
|
||||||
|
var expiredBoxes int
|
||||||
|
if err := pool.QueryRow(ctx, `
|
||||||
|
SELECT r.handled_at,
|
||||||
|
COALESCE((pm.media #>> '{service_action,no_forwards,expired}')::boolean, false),
|
||||||
|
(SELECT count(*)
|
||||||
|
FROM message_boxes b
|
||||||
|
WHERE b.message_sender_id = r.private_message_sender_user_id
|
||||||
|
AND b.private_message_id = r.private_message_id
|
||||||
|
AND COALESCE((b.media #>> '{service_action,no_forwards,expired}')::boolean, false))
|
||||||
|
FROM private_no_forwards_requests r
|
||||||
|
JOIN private_messages pm
|
||||||
|
ON pm.sender_user_id = r.private_message_sender_user_id
|
||||||
|
AND pm.id = r.private_message_id
|
||||||
|
WHERE r.private_message_sender_user_id = $1
|
||||||
|
AND r.private_message_id = $2`, bob.ID, request.Send.SenderMessage.UID,
|
||||||
|
).Scan(&handledAt, &logicalExpired, &expiredBoxes); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if handledAt != 1700100002 || !logicalExpired || expiredBoxes != 2 {
|
||||||
|
t.Fatalf("handled request handled_at=%d logical_expired=%v boxes=%d", handledAt, logicalExpired, expiredBoxes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPostgresPrivateNoForwardsConcurrentOwnershipAndOneShotAnswer(t *testing.T) {
|
||||||
|
pool := testPool(t)
|
||||||
|
ctx := context.Background()
|
||||||
|
suffix := randomSuffix(t)
|
||||||
|
users := NewUserStore(pool)
|
||||||
|
alice, err := users.Create(ctx, domain.User{AccessHash: 6201, Phone: "+1669" + suffix + "01", FirstName: "Alice"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
bob, err := users.Create(ctx, domain.User{AccessHash: 6202, Phone: "+1669" + suffix + "02", FirstName: "Bob"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
_, _ = pool.Exec(ctx, "DELETE FROM users WHERE id = ANY($1::bigint[])", []int64{alice.ID, bob.ID})
|
||||||
|
})
|
||||||
|
|
||||||
|
messages := NewMessageStore(pool)
|
||||||
|
baseRandom := time.Now().UnixNano()
|
||||||
|
enableResults := make([]domain.TogglePrivateNoForwardsResult, 2)
|
||||||
|
enableErrors := make([]error, 2)
|
||||||
|
actors := []int64{alice.ID, bob.ID}
|
||||||
|
peers := []int64{bob.ID, alice.ID}
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for i := range actors {
|
||||||
|
wg.Add(1)
|
||||||
|
go func(i int) {
|
||||||
|
defer wg.Done()
|
||||||
|
enableResults[i], enableErrors[i] = messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: actors[i],
|
||||||
|
PeerUserID: peers[i],
|
||||||
|
Enabled: true,
|
||||||
|
RandomID: baseRandom + int64(i),
|
||||||
|
Date: 1700200000,
|
||||||
|
})
|
||||||
|
}(i)
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
changed := 0
|
||||||
|
for i, err := range enableErrors {
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("concurrent enable %d: %v", i, err)
|
||||||
|
}
|
||||||
|
if enableResults[i].Changed {
|
||||||
|
changed++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if changed != 1 {
|
||||||
|
t.Fatalf("concurrent enable changed=%d, want exactly one service message", changed)
|
||||||
|
}
|
||||||
|
state, err := messages.GetPrivateNoForwards(ctx, alice.ID, bob.ID)
|
||||||
|
if err != nil || (state.EnabledByUserID != alice.ID && state.EnabledByUserID != bob.ID) {
|
||||||
|
t.Fatalf("concurrent enable state=%+v err=%v", state, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ownerID := state.EnabledByUserID
|
||||||
|
requesterID := alice.ID
|
||||||
|
if ownerID == alice.ID {
|
||||||
|
requesterID = bob.ID
|
||||||
|
}
|
||||||
|
request, err := messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: requesterID,
|
||||||
|
PeerUserID: ownerID,
|
||||||
|
RandomID: baseRandom + 10,
|
||||||
|
Date: 1700200001,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("create disable request: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
answerResults := make([]domain.TogglePrivateNoForwardsResult, 2)
|
||||||
|
answerErrors := make([]error, 2)
|
||||||
|
for i := range answerResults {
|
||||||
|
wg.Add(1)
|
||||||
|
go func(i int) {
|
||||||
|
defer wg.Done()
|
||||||
|
answerResults[i], answerErrors[i] = messages.TogglePrivateNoForwards(ctx, domain.TogglePrivateNoForwardsRequest{
|
||||||
|
ActorUserID: ownerID,
|
||||||
|
PeerUserID: requesterID,
|
||||||
|
Enabled: false,
|
||||||
|
RequestMsgID: request.Send.RecipientMessage.ID,
|
||||||
|
RandomID: baseRandom + 20 + int64(i),
|
||||||
|
Date: 1700200002,
|
||||||
|
})
|
||||||
|
}(i)
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
successes, expired := 0, 0
|
||||||
|
for i, err := range answerErrors {
|
||||||
|
switch {
|
||||||
|
case err == nil && answerResults[i].Changed:
|
||||||
|
successes++
|
||||||
|
case errors.Is(err, domain.ErrNoForwardsRequestExpired):
|
||||||
|
expired++
|
||||||
|
default:
|
||||||
|
t.Fatalf("concurrent answer %d result=%+v err=%v", i, answerResults[i], err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if successes != 1 || expired != 1 {
|
||||||
|
t.Fatalf("concurrent answers successes=%d expired=%d, want 1/1", successes, expired)
|
||||||
|
}
|
||||||
|
state, err = messages.GetPrivateNoForwards(ctx, alice.ID, bob.ID)
|
||||||
|
if err != nil || state.Enabled() {
|
||||||
|
t.Fatalf("state after concurrent answer=%+v err=%v, want disabled", state, err)
|
||||||
|
}
|
||||||
|
for _, userID := range []int64{alice.ID, bob.ID} {
|
||||||
|
events, err := NewUpdateEventStore(pool).ListAfter(ctx, userID, 0, 10)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("events user %d: %v", userID, err)
|
||||||
|
}
|
||||||
|
if len(events) != 3 || events[0].Pts != 1 || events[1].Pts != 2 || events[2].Pts != 3 {
|
||||||
|
t.Fatalf("events user %d = %+v, want one enable/request/answer sequence", userID, events)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
15
internal/store/private_no_forwards.go
Normal file
15
internal/store/private_no_forwards.go
Normal file
|
|
@ -0,0 +1,15 @@
|
||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"telesrv/internal/domain"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PrivateNoForwardsStore keeps the canonical pair state and its service-message
|
||||||
|
// transition in one store transaction. It is optional so unrelated lightweight
|
||||||
|
// MessageStore test doubles do not need to implement this capability.
|
||||||
|
type PrivateNoForwardsStore interface {
|
||||||
|
GetPrivateNoForwards(ctx context.Context, viewerUserID, peerUserID int64) (domain.PrivateNoForwardsState, error)
|
||||||
|
TogglePrivateNoForwards(ctx context.Context, req domain.TogglePrivateNoForwardsRequest) (domain.TogglePrivateNoForwardsResult, error)
|
||||||
|
}
|
||||||
Loading…
Add table
Add a link
Reference in a new issue