diff --git a/deploy/migrations/0001_init.up.sql b/deploy/migrations/0001_init.up.sql index 372df81e..bb0d8541 100644 --- a/deploy/migrations/0001_init.up.sql +++ b/deploy/migrations/0001_init.up.sql @@ -2892,7 +2892,7 @@ CREATE TABLE public.user_saved_reaction_tags ( created_at timestamp with time zone DEFAULT now() NOT NULL, updated_at timestamp with time zone DEFAULT now() NOT NULL, CONSTRAINT user_saved_reaction_tags_reaction_count_check CHECK ((reaction_count >= 0)), - CONSTRAINT user_saved_reaction_tags_reaction_type_check CHECK (((reaction_type)::text = 'emoji'::text)), + CONSTRAINT user_saved_reaction_tags_reaction_type_check CHECK (((reaction_type)::text = ANY (ARRAY['emoji'::text, 'custom_emoji'::text]))), CONSTRAINT user_saved_reaction_tags_reaction_value_check CHECK ((reaction_value <> ''::text)), CONSTRAINT user_saved_reaction_tags_title_check CHECK ((char_length(title) <= 12)) ); @@ -2948,7 +2948,7 @@ CREATE TABLE public.user_update_events ( story_payload jsonb DEFAULT '{}'::jsonb NOT NULL, reaction_payload jsonb DEFAULT '{}'::jsonb NOT NULL, CONSTRAINT user_update_events_peer_type_check CHECK (((peer_type IS NULL) OR ((peer_type)::text = ANY (ARRAY[('user'::character varying)::text, ('channel'::character varying)::text])))), - CONSTRAINT user_update_events_type_check CHECK (((event_type)::text = ANY (ARRAY[('new_message'::character varying)::text, ('read_history_inbox'::character varying)::text, ('read_history_outbox'::character varying)::text, ('read_message_contents'::character varying)::text, ('edit_message'::character varying)::text, ('message_reactions'::character varying)::text, ('message_poll'::character varying)::text, ('draft_message'::character varying)::text, ('quick_replies'::character varying)::text, ('new_quick_reply'::character varying)::text, ('delete_quick_reply'::character varying)::text, ('quick_reply_message'::character varying)::text, ('delete_quick_reply_messages'::character varying)::text, ('contacts_reset'::character varying)::text, ('dialog_pinned'::character varying)::text, ('pinned_dialogs'::character varying)::text, ('pinned_messages'::character varying)::text, ('dialog_unread_mark'::character varying)::text, ('peer_settings'::character varying)::text, ('peer_story_blocked'::character varying)::text, ('delete_messages'::character varying)::text, ('dialog_filter'::character varying)::text, ('dialog_filter_order'::character varying)::text, ('dialog_filters'::character varying)::text, ('folder_peers'::character varying)::text, ('channel_available_messages'::character varying)::text, ('channel_view_forum_as_messages'::character varying)::text, ('channel_state'::character varying)::text, ('saved_dialog_pinned'::character varying)::text, ('pinned_saved_dialogs'::character varying)::text, ('story'::character varying)::text, ('read_stories'::character varying)::text, ('sent_story_reaction'::character varying)::text, ('new_story_reaction'::character varying)::text, ('noop'::character varying)::text]))) + CONSTRAINT user_update_events_type_check CHECK (((event_type)::text = ANY (ARRAY[('new_message'::character varying)::text, ('read_history_inbox'::character varying)::text, ('read_history_outbox'::character varying)::text, ('read_message_contents'::character varying)::text, ('edit_message'::character varying)::text, ('message_poll'::character varying)::text, ('draft_message'::character varying)::text, ('quick_replies'::character varying)::text, ('new_quick_reply'::character varying)::text, ('delete_quick_reply'::character varying)::text, ('quick_reply_message'::character varying)::text, ('delete_quick_reply_messages'::character varying)::text, ('contacts_reset'::character varying)::text, ('dialog_pinned'::character varying)::text, ('pinned_dialogs'::character varying)::text, ('pinned_messages'::character varying)::text, ('dialog_unread_mark'::character varying)::text, ('peer_settings'::character varying)::text, ('peer_story_blocked'::character varying)::text, ('delete_messages'::character varying)::text, ('dialog_filter'::character varying)::text, ('dialog_filter_order'::character varying)::text, ('dialog_filters'::character varying)::text, ('folder_peers'::character varying)::text, ('channel_available_messages'::character varying)::text, ('channel_view_forum_as_messages'::character varying)::text, ('channel_state'::character varying)::text, ('saved_dialog_pinned'::character varying)::text, ('pinned_saved_dialogs'::character varying)::text, ('story'::character varying)::text, ('read_stories'::character varying)::text, ('sent_story_reaction'::character varying)::text, ('new_story_reaction'::character varying)::text, ('noop'::character varying)::text]))) ); diff --git a/deploy/migrations/0147_user_moderation_profile_events.up.sql b/deploy/migrations/0147_user_moderation_profile_events.up.sql index a04a5d98..d141788e 100644 --- a/deploy/migrations/0147_user_moderation_profile_events.up.sql +++ b/deploy/migrations/0147_user_moderation_profile_events.up.sql @@ -7,7 +7,7 @@ ALTER TABLE public.user_update_events DROP CONSTRAINT IF EXISTS user_update_even ALTER TABLE public.user_update_events ADD CONSTRAINT user_update_events_type_check CHECK ( (event_type)::text = ANY (ARRAY[ 'new_message', 'read_history_inbox', 'read_history_outbox', 'read_message_contents', - 'edit_message', 'web_page', 'message_reactions', 'message_poll', 'draft_message', 'quick_replies', + 'edit_message', 'web_page', 'message_poll', 'draft_message', 'quick_replies', 'new_quick_reply', 'delete_quick_reply', 'quick_reply_message', 'delete_quick_reply_messages', 'contacts_reset', 'dialog_pinned', 'pinned_dialogs', 'pinned_messages', 'dialog_unread_mark', 'peer_settings', 'peer_story_blocked', 'user_phone', 'user_emoji_status', diff --git a/deploy/migrations/0148_saved_message_reaction_tags.down.sql b/deploy/migrations/0148_saved_message_reaction_tags.down.sql new file mode 100644 index 00000000..083fa567 --- /dev/null +++ b/deploy/migrations/0148_saved_message_reaction_tags.down.sql @@ -0,0 +1,9 @@ +DROP TABLE IF EXISTS public.saved_message_reaction_tags; + +ALTER TABLE public.user_saved_reaction_tags + DROP CONSTRAINT IF EXISTS user_saved_reaction_tags_reaction_type_check; +DELETE FROM public.user_saved_reaction_tags +WHERE (reaction_type)::text <> 'emoji'::text; +ALTER TABLE public.user_saved_reaction_tags + ADD CONSTRAINT user_saved_reaction_tags_reaction_type_check + CHECK ((reaction_type)::text = 'emoji'::text); diff --git a/deploy/migrations/0148_saved_message_reaction_tags.up.sql b/deploy/migrations/0148_saved_message_reaction_tags.up.sql new file mode 100644 index 00000000..8b9f7807 --- /dev/null +++ b/deploy/migrations/0148_saved_message_reaction_tags.up.sql @@ -0,0 +1,33 @@ +CREATE TABLE public.saved_message_reaction_tags ( + user_id bigint NOT NULL, + message_box_id integer NOT NULL, + reaction_type character varying(16) NOT NULL, + reaction_value text NOT NULL, + chosen_order integer NOT NULL, + created_at timestamp with time zone DEFAULT now() NOT NULL, + updated_at timestamp with time zone DEFAULT now() NOT NULL, + CONSTRAINT saved_message_reaction_tags_pkey + PRIMARY KEY (user_id, message_box_id, reaction_type, reaction_value), + CONSTRAINT saved_message_reaction_tags_order_check CHECK (chosen_order > 0), + CONSTRAINT saved_message_reaction_tags_type_check + CHECK ((reaction_type)::text = ANY (ARRAY['emoji'::text, 'custom_emoji'::text])), + CONSTRAINT saved_message_reaction_tags_value_check CHECK (reaction_value <> ''), + CONSTRAINT saved_message_reaction_tags_user_id_fkey + FOREIGN KEY (user_id) REFERENCES public.users(id) ON DELETE CASCADE, + CONSTRAINT saved_message_reaction_tags_message_box_fkey + FOREIGN KEY (user_id, message_box_id) + REFERENCES public.message_boxes(owner_user_id, box_id) ON DELETE CASCADE +); + +CREATE INDEX saved_message_reaction_tags_reaction_message_idx + ON public.saved_message_reaction_tags + (user_id, ((reaction_type)::text || ':' || reaction_value), message_box_id DESC); + +ALTER TABLE public.user_saved_reaction_tags + DROP CONSTRAINT user_saved_reaction_tags_reaction_type_check; +ALTER TABLE public.user_saved_reaction_tags + ADD CONSTRAINT user_saved_reaction_tags_reaction_type_check + CHECK ((reaction_type)::text = ANY (ARRAY['emoji'::text, 'custom_emoji'::text])); + +COMMENT ON COLUMN public.user_saved_reaction_tags.reaction_count IS + 'Legacy unused column; visible counts are aggregated from saved_message_reaction_tags.'; diff --git a/internal/app/channels/service.go b/internal/app/channels/service.go index 27822a4a..48d86681 100644 --- a/internal/app/channels/service.go +++ b/internal/app/channels/service.go @@ -1165,34 +1165,6 @@ func (s *Service) ClearRecentReactions(ctx context.Context, userID int64) error return s.channels.ClearRecentMessageReactions(ctx, userID) } -// SavedReactionTags returns account-level saved-message reaction tag titles. -func (s *Service) SavedReactionTags(ctx context.Context, userID int64, limit int) ([]domain.SavedReactionTag, error) { - if s == nil || s.channels == nil || userID == 0 { - return nil, domain.ErrChannelInvalid - } - if limit <= 0 { - return []domain.SavedReactionTag{}, nil - } - if limit > domain.MaxSavedReactionTags { - limit = domain.MaxSavedReactionTags - } - return s.channels.ListSavedReactionTags(ctx, userID, limit) -} - -// UpdateSavedReactionTag stores the account-level custom title for one saved-message reaction tag. -func (s *Service) UpdateSavedReactionTag(ctx context.Context, userID int64, tag domain.SavedReactionTag) error { - if s == nil || s.channels == nil || userID == 0 { - return domain.ErrChannelInvalid - } - if tag.UserID == 0 { - tag.UserID = userID - } - if tag.UserID != userID || tag.Reaction.Type != domain.MessageReactionEmoji || tag.Reaction.Emoticon == "" { - return domain.ErrChannelInvalid - } - return s.channels.UpsertSavedReactionTag(ctx, tag) -} - // ReadMessageContents returns visible channel messages whose content-read state can be synced. func (s *Service) ReadMessageContents(ctx context.Context, userID int64, req domain.ReadChannelMessageContentsRequest) (domain.ReadChannelMessageContentsResult, error) { if s == nil || s.channels == nil || userID == 0 { diff --git a/internal/app/messages/service.go b/internal/app/messages/service.go index 82ee42fb..bf1e1ab2 100644 --- a/internal/app/messages/service.go +++ b/internal/app/messages/service.go @@ -433,6 +433,31 @@ func (s *Service) GetMessageReactions(ctx context.Context, userID int64, req dom return s.messages.GetMessageReactions(ctx, req) } +// SavedReactionTags returns the global or one-sub-dialog Saved Messages tag list. +func (s *Service) SavedReactionTags(ctx context.Context, userID int64, savedPeer domain.Peer, limit int) ([]domain.SavedReactionTag, error) { + if s == nil || s.messages == nil || userID == 0 { + return nil, nil + } + if limit <= 0 || limit > domain.MaxSavedReactionTags { + limit = domain.MaxSavedReactionTags + } + return s.messages.ListSavedReactionTags(ctx, domain.SavedReactionTagsRequest{ + UserID: userID, + SavedPeer: savedPeer, + Limit: limit, + }) +} + +// UpdateSavedReactionTag stores or removes the optional global title for one +// tag that is currently assigned to at least one visible Saved Message. +func (s *Service) UpdateSavedReactionTag(ctx context.Context, userID int64, tag domain.SavedReactionTag) error { + if s == nil || s.messages == nil || userID == 0 || !tag.Reaction.Valid() { + return domain.ErrReactionInvalid + } + tag.UserID = userID + return s.messages.UpsertSavedReactionTag(ctx, tag) +} + // EditMessage 编辑当前账号发出的私聊文本消息。 func (s *Service) EditMessage(ctx context.Context, userID int64, req domain.EditMessageRequest) (domain.EditMessageResult, error) { if s == nil || s.messages == nil || userID == 0 { diff --git a/internal/app/messages/service_test.go b/internal/app/messages/service_test.go index f7f49377..feb5dc6b 100644 --- a/internal/app/messages/service_test.go +++ b/internal/app/messages/service_test.go @@ -682,6 +682,14 @@ func (s projectionMessageStore) GetMessageReactions(context.Context, domain.Priv return domain.PrivateMessageReactionsResult{}, nil } +func (s projectionMessageStore) ListSavedReactionTags(context.Context, domain.SavedReactionTagsRequest) ([]domain.SavedReactionTag, error) { + return nil, nil +} + +func (s projectionMessageStore) UpsertSavedReactionTag(context.Context, domain.SavedReactionTag) error { + return nil +} + func (s projectionMessageStore) VoteMessagePoll(context.Context, domain.VotePrivateMessagePollRequest) (domain.PrivateMessagePollResult, error) { return domain.PrivateMessagePollResult{}, nil } diff --git a/internal/app/updates/service.go b/internal/app/updates/service.go index bbae8f17..3817a34b 100644 --- a/internal/app/updates/service.go +++ b/internal/app/updates/service.go @@ -400,31 +400,9 @@ func (s *Service) PublishNewMessage(ctx context.Context, userID int64, msg domai }, true, 0, false) } -// RecordMessageReactions records a durable marker for message reaction changes. -// -// updateMessageReactions has no pts fields in Layer 225, but TDesktop still -// needs getDifference to advance account pts and carry the latest reaction -// aggregate for offline devices. -func (s *Service) RecordMessageReactions(ctx context.Context, authKeyID [8]byte, userID int64, msg domain.Message) (domain.UpdateEvent, domain.UpdateState, error) { - if userID == 0 { - userID = msg.OwnerUserID - } - date := msg.Date - if date == 0 { - date = int(time.Now().Unix()) - } - return s.recordEventWithoutState(ctx, userID, domain.UpdateEvent{ - Type: domain.UpdateEventMessageReactions, - Date: date, - Message: msg, - Peer: msg.Peer, - PtsCount: 1, - }) -} - // RecordMessagePoll records a durable marker for message poll state changes // (vote / close). updateMessagePoll has no pts fields in Layer 225 — same -// bookkeeping shape as RecordMessageReactions. +// historical bookkeeping shape pending its own audit. func (s *Service) RecordMessagePoll(ctx context.Context, authKeyID [8]byte, userID int64, msg domain.Message) (domain.UpdateEvent, domain.UpdateState, error) { if userID == 0 { userID = msg.OwnerUserID diff --git a/internal/domain/channel.go b/internal/domain/channel.go index cddefe56..a43c2971 100644 --- a/internal/domain/channel.go +++ b/internal/domain/channel.go @@ -802,8 +802,11 @@ type ChannelMessagePeerReaction struct { // ChannelMessageReactions is the read model carried by channel messages and reaction updates. type ChannelMessageReactions struct { CanSeeList bool - Results []ChannelMessageReactionCount - Recent []ChannelMessagePeerReaction + // AsTags marks reactions on Saved Messages as private message tags. + // It is never set for ordinary private/channel reactions. + AsTags bool + Results []ChannelMessageReactionCount + Recent []ChannelMessagePeerReaction // Paid 是付费 reaction(Stars)聚合(nil = 无);读路径从 channel_message_paid_reactions // 填充、tg 转换注入 ReactionPaid 计数 + top reactors。与普通 reaction 分表存储。 Paid *ChannelMessagePaidReactions @@ -923,6 +926,14 @@ type SavedReactionTag struct { Count int } +// SavedReactionTagsRequest lists account-level Saved Messages tags. SavedPeer +// is zero for the global list and non-zero for one Saved Messages sub-dialog. +type SavedReactionTagsRequest struct { + UserID int64 + SavedPeer Peer + Limit int +} + // ChannelDiscussionRef links a broadcast post to its discussion megagroup root message. type ChannelDiscussionRef struct { ChannelID int64 diff --git a/internal/domain/message.go b/internal/domain/message.go index c46fc9e0..5334e2d0 100644 --- a/internal/domain/message.go +++ b/internal/domain/message.go @@ -232,6 +232,8 @@ type MessageFilter struct { Query string OffsetID int OffsetDate int + MinDate int + MaxDate int AddOffset int Limit int MaxID int @@ -245,6 +247,9 @@ type MessageFilter struct { // SavedPeer 非零时仅返回 self-chat 中该 saved 子会话的消息 // (messages.getSavedHistory);Peer 必须同时是 self。 SavedPeer Peer + // SavedReactions 非空时仅返回至少带其中一个 tag 的 Saved Messages。 + // 仅 messages.search(peer=self) 使用;普通私聊 reaction 不参与匹配。 + SavedReactions []MessageReaction // PeerIDs restricts a global private search to these user peers. Empty is a // valid restricted set, so RestrictPeerIDs carries presence separately. PeerIDs []int64 diff --git a/internal/domain/update_event.go b/internal/domain/update_event.go index 0c1aa7b8..9dab6286 100644 --- a/internal/domain/update_event.go +++ b/internal/domain/update_event.go @@ -18,11 +18,10 @@ const ( // UpdateEventWebPage 映射 updateWebPage:异步解析完成后把消息里的 pending 链接预览 // 占位就地替换为已解析卡片。携带账号 pts(非 LacksWirePts),消息快照经 box JOIN 重建, // 故 difference/dispatch 与 edit_message 同走通用消息事件路径,仅 tg 投影构造器不同。 - UpdateEventWebPage UpdateEventType = "web_page" - UpdateEventMessageReactions UpdateEventType = "message_reactions" + UpdateEventWebPage UpdateEventType = "web_page" // UpdateEventMessagePoll 映射 updateMessagePoll(投票/关闭后 poll 状态变化; // Message 为该 owner 视角消息,media 在 difference 重放时按 viewer 重新 enrich)。 - // 与 reaction 同款:占账号 pts 但 TL 构造器无 pts,见 LacksWirePts。 + // 当前历史实现仍将 poll 作为待复核的 LacksWirePts 事件。 UpdateEventMessagePoll UpdateEventType = "message_poll" UpdateEventContactsReset UpdateEventType = "contacts_reset" UpdateEventDialogPinned UpdateEventType = "dialog_pinned" @@ -126,8 +125,7 @@ type UpdateEvent struct { // 下一条真正带 pts 的更新会被判为空洞。 func (e UpdateEvent) LacksWirePts() bool { switch e.Type { - case UpdateEventMessageReactions, - UpdateEventMessagePoll, + case UpdateEventMessagePoll, UpdateEventDraftMessage, UpdateEventChannelState, UpdateEventContactsReset, diff --git a/internal/rpc/channels_passive_stubs_rpc_test.go b/internal/rpc/channels_passive_stubs_rpc_test.go index 77214d85..26e6a145 100644 --- a/internal/rpc/channels_passive_stubs_rpc_test.go +++ b/internal/rpc/channels_passive_stubs_rpc_test.go @@ -4,6 +4,7 @@ import ( "github.com/iamxvbaba/td/tg" "strings" apptelemetry "telesrv/internal/app/clienttelemetry" + appmessages "telesrv/internal/app/messages" appmoderation "telesrv/internal/app/moderation" "telesrv/internal/domain" "telesrv/internal/store/memory" @@ -865,28 +866,22 @@ func TestTDesktopPassiveChannelStubs(t *testing.T) { if !ok || staleEmptyPage.Hash != 0 || len(staleEmptyPage.Tags) != 0 { t.Fatalf("messages.getSavedReactionTags stale empty hash = %#v, want empty page hash 0", staleEmptyTags) } + r.deps.Messages = appmessages.NewService(memory.NewMessageStore(), nil) + if _, err := f.users.SetPremiumUntil(ownerCtx, owner.ID, int(time.Now().Add(time.Hour).Unix())); err != nil { + t.Fatalf("grant owner premium for saved tag rename: %v", err) + } updateTagReq := &tg.MessagesUpdateSavedReactionTagRequest{Reaction: &tg.ReactionEmoji{Emoticon: "ok"}} updateTagReq.SetTitle("Work") - if ok, err := r.onMessagesUpdateSavedReactionTag(ownerCtx, updateTagReq); err != nil || !ok { - t.Fatalf("messages.updateSavedReactionTag = ok %v err %v, want true nil", ok, err) + if _, err := r.onMessagesUpdateSavedReactionTag(ownerCtx, updateTagReq); err == nil || !strings.Contains(err.Error(), "REACTION_INVALID") { + t.Fatalf("messages.updateSavedReactionTag unassigned err = %v, want REACTION_INVALID", err) } globalTags, err := r.onMessagesGetSavedReactionTags(ownerCtx, &tg.MessagesGetSavedReactionTagsRequest{}) if err != nil { t.Fatalf("messages.getSavedReactionTags global: %v", err) } globalPage, ok := globalTags.(*tg.MessagesSavedReactionTags) - if !ok || globalPage.Hash == 0 || len(globalPage.Tags) != 1 { - t.Fatalf("messages.getSavedReactionTags global = %#v, want one hashable tag", globalTags) - } - if emoji, ok := globalPage.Tags[0].Reaction.(*tg.ReactionEmoji); !ok || emoji.Emoticon != "ok" || globalPage.Tags[0].Title != "Work" || globalPage.Tags[0].Count != 0 { - t.Fatalf("messages.getSavedReactionTags tag = %+v, want ok/Work/count0", globalPage.Tags[0]) - } - globalNotModified, err := r.onMessagesGetSavedReactionTags(ownerCtx, &tg.MessagesGetSavedReactionTagsRequest{Hash: globalPage.Hash}) - if err != nil { - t.Fatalf("messages.getSavedReactionTags hash: %v", err) - } - if _, ok := globalNotModified.(*tg.MessagesSavedReactionTagsNotModified); !ok { - t.Fatalf("messages.getSavedReactionTags hash = %#v, want notModified", globalNotModified) + if !ok || globalPage.Hash != 0 || len(globalPage.Tags) != 0 { + t.Fatalf("messages.getSavedReactionTags global = %#v, want empty", globalTags) } peerTagsAfterUpdate, err := r.onMessagesGetSavedReactionTags(ownerCtx, savedTagsReq) if err != nil { @@ -909,9 +904,11 @@ func TestTDesktopPassiveChannelStubs(t *testing.T) { if err != nil { t.Fatalf("messages.getDefaultTagReactions: %v", err) } - if got := tagReactions.(*tg.MessagesReactions).Reactions; len(got) != 0 { - t.Fatalf("messages.getDefaultTagReactions = %+v, want empty", got) + defaultPage, ok := tagReactions.(*tg.MessagesReactions) + if !ok || defaultPage.Hash == 0 || len(defaultPage.Reactions) == 0 { + t.Fatalf("messages.getDefaultTagReactions = %#v, want non-empty hashable catalog", tagReactions) } + r.deps.Messages = nil // poll 链路已是真实现:对非 poll 消息一律 MESSAGE_ID_INVALID(与官方一致)。 if _, err := r.onMessagesSendVote(ownerCtx, &tg.MessagesSendVoteRequest{ Peer: inputPeerChannel(channel), diff --git a/internal/rpc/convert_messages.go b/internal/rpc/convert_messages.go index 96f88f54..af793457 100644 --- a/internal/rpc/convert_messages.go +++ b/internal/rpc/convert_messages.go @@ -444,6 +444,9 @@ func tgMessageReactions(viewerUserID int64, in *domain.ChannelMessageReactions) if in.CanSeeList { out.SetCanSeeList(true) } + if in.AsTags { + out.SetReactionsAsTags(true) + } for _, item := range in.Results { reaction := tgMessageReaction(item.Reaction) if reaction == nil || item.Count <= 0 { diff --git a/internal/rpc/convert_updates.go b/internal/rpc/convert_updates.go index efdb6bea..7d3f7bd8 100644 --- a/internal/rpc/convert_updates.go +++ b/internal/rpc/convert_updates.go @@ -35,7 +35,7 @@ func tgUpdatesDifference(viewerUserID int64, diff domain.UpdateDifference) tg.Up if update := tgReadHistoryOutboxUpdate(event); update != nil { out.OtherUpdates = append(out.OtherUpdates, update) } - case domain.UpdateEventMessageReactions, domain.UpdateEventMessagePoll: + case domain.UpdateEventMessagePoll: // 同时下发消息快照(含最新聚合)与对应通知 update;事件无 TL pts, // pts 推进靠 difference state 本身。 if msg := tgMessage(event.Message); msg != nil { @@ -464,32 +464,6 @@ func tgOtherUpdateFromEvent(event domain.UpdateEvent) tg.UpdateClass { return nil } return tgUpdateMessagePoll(pollPeer, event.Message.ID, media.Poll) - case domain.UpdateEventMessageReactions: - if event.Message.ID <= 0 || event.Message.ID > domain.MaxMessageBoxID { - return nil - } - peer := event.Message.Peer - if peer.Type == "" || peer.ID == 0 { - peer = event.Peer - } - outPeer := tgPeer(peer) - if outPeer == nil { - return nil - } - reactions := event.Message.Reactions - if reactions == nil { - empty := domain.ChannelMessageReactions{CanSeeList: true, Results: []domain.ChannelMessageReactionCount{}, Recent: []domain.ChannelMessagePeerReaction{}} - reactions = &empty - } - converted := tgMessageReactions(event.UserID, reactions) - if converted == nil { - converted = &tg.MessageReactions{Results: []tg.ReactionCount{}} - } - return &tg.UpdateMessageReactions{ - Peer: outPeer, - MsgID: event.Message.ID, - Reactions: *converted, - } case domain.UpdateEventDialogFilter: update := &tg.UpdateDialogFilter{ID: event.FilterID} if event.DialogFilter != nil { diff --git a/internal/rpc/deps.go b/internal/rpc/deps.go index 495a0d3b..431799ee 100644 --- a/internal/rpc/deps.go +++ b/internal/rpc/deps.go @@ -564,6 +564,8 @@ type MessagesService interface { GetOutboxReadDate(ctx context.Context, userID int64, req domain.OutboxReadDateRequest) (int, error) SetMessageReactions(ctx context.Context, userID int64, req domain.SetPrivateMessageReactionsRequest) (domain.PrivateMessageReactionsResult, error) GetMessageReactions(ctx context.Context, userID int64, req domain.PrivateMessageReactionsRequest) (domain.PrivateMessageReactionsResult, error) + SavedReactionTags(ctx context.Context, userID int64, savedPeer domain.Peer, limit int) ([]domain.SavedReactionTag, error) + UpdateSavedReactionTag(ctx context.Context, userID int64, tag domain.SavedReactionTag) error VoteMessagePoll(ctx context.Context, userID int64, req domain.VotePrivateMessagePollRequest) (domain.PrivateMessagePollResult, error) CloseMessagePoll(ctx context.Context, userID int64, req domain.ClosePrivateMessagePollRequest) (domain.PrivateMessagePollResult, error) ListUnreadReactionMessages(ctx context.Context, userID int64, peer domain.Peer, limit int) ([]domain.Message, error) @@ -691,8 +693,6 @@ type ChannelsService interface { TopReactions(ctx context.Context, userID int64, limit int) ([]domain.MessageReaction, error) RecentReactions(ctx context.Context, userID int64, limit int) ([]domain.MessageReaction, error) ClearRecentReactions(ctx context.Context, userID int64) error - SavedReactionTags(ctx context.Context, userID int64, limit int) ([]domain.SavedReactionTag, error) - UpdateSavedReactionTag(ctx context.Context, userID int64, tag domain.SavedReactionTag) error GetPremiumBoostStatus(ctx context.Context, userID, channelID int64, now int) (domain.PremiumBoostStatus, error) ListPremiumBoosts(ctx context.Context, userID, channelID int64, gifts bool, offset string, limit, now int) (domain.PremiumBoostList, error) GetPremiumMyBoosts(ctx context.Context, userID int64, now, premiumUntil int) (domain.PremiumMyBoosts, error) diff --git a/internal/rpc/messages_contracts.go b/internal/rpc/messages_contracts.go index 22ccbefd..5844f699 100644 --- a/internal/rpc/messages_contracts.go +++ b/internal/rpc/messages_contracts.go @@ -19,10 +19,6 @@ type accountPaidReactionPrivacyService interface { SetPaidReactionPrivacy(ctx context.Context, userID int64, privacy domain.PaidReactionPrivacy) (domain.AccountReactionSettings, error) } -type messageReactionUpdateRecorder interface { - RecordMessageReactions(ctx context.Context, authKeyID [8]byte, userID int64, msg domain.Message) (domain.UpdateEvent, domain.UpdateState, error) -} - type messagePollUpdateRecorder interface { RecordMessagePoll(ctx context.Context, authKeyID [8]byte, userID int64, msg domain.Message) (domain.UpdateEvent, domain.UpdateState, error) } diff --git a/internal/rpc/messages_history.go b/internal/rpc/messages_history.go index 5ba0e8e8..eb5234e4 100644 --- a/internal/rpc/messages_history.go +++ b/internal/rpc/messages_history.go @@ -766,7 +766,7 @@ func (r *Router) messageFilterFromHistoryRequest(userID int64, req *tg.MessagesG }, true } -func (r *Router) messageFilterFromSearchRequest(userID int64, req *tg.MessagesSearchRequest) domain.MessageFilter { +func (r *Router) messageFilterFromSearchRequest(ctx context.Context, userID int64, req *tg.MessagesSearchRequest) (domain.MessageFilter, error) { limit := req.Limit if limit > 500 { limit = 500 @@ -774,6 +774,8 @@ func (r *Router) messageFilterFromSearchRequest(userID int64, req *tg.MessagesSe filter := domain.MessageFilter{ Query: req.Q, OffsetID: req.OffsetID, + MinDate: req.MinDate, + MaxDate: req.MaxDate, AddOffset: domain.ClampMessageHistoryAddOffset(req.AddOffset), Limit: limit, MaxID: req.MaxID, @@ -786,7 +788,42 @@ func (r *Router) messageFilterFromSearchRequest(userID int64, req *tg.MessagesSe filter.HasPeer = true filter.Peer = peer } - return filter + savedReactions, hasSavedReactions := req.GetSavedReaction() + savedPeerInput, hasSavedPeer := req.GetSavedPeerID() + if hasSavedReactions || hasSavedPeer { + if !filter.HasPeer || + filter.Peer != (domain.Peer{Type: domain.PeerTypeUser, ID: userID}) { + return domain.MessageFilter{}, peerIDInvalidErr() + } + } + if hasSavedPeer { + if savedPeerInput == nil { + return domain.MessageFilter{}, peerIDInvalidErr() + } + savedPeer, err := r.checkedDomainPeerFromInputPeer(ctx, userID, savedPeerInput) + if err != nil { + return domain.MessageFilter{}, err + } + filter.SavedPeer = savedPeer + } + if hasSavedReactions { + if len(savedReactions) == 0 || len(savedReactions) > maxReactionVector { + return domain.MessageFilter{}, reactionInvalidErr() + } + seen := make(map[string]struct{}, len(savedReactions)) + for _, item := range savedReactions { + reaction, err := domainMessageReactionFromTL(item) + if err != nil { + return domain.MessageFilter{}, err + } + if _, ok := seen[reaction.Key()]; ok { + continue + } + seen[reaction.Key()] = struct{}{} + filter.SavedReactions = append(filter.SavedReactions, reaction) + } + } + return filter, nil } func (r *Router) channelHistoryFilterFromSearchRequest(userID int64, req *tg.MessagesSearchRequest, channelID int64) (domain.ChannelHistoryFilter, bool) { diff --git a/internal/rpc/messages_reactions_catalog.go b/internal/rpc/messages_reactions_catalog.go index f4cad6ea..31232795 100644 --- a/internal/rpc/messages_reactions_catalog.go +++ b/internal/rpc/messages_reactions_catalog.go @@ -2,10 +2,13 @@ package rpc import ( "context" - "github.com/iamxvbaba/td/tg" - "hash/fnv" - "strconv" + "crypto/md5" + "encoding/binary" + "sort" "strings" + + "github.com/iamxvbaba/td/tg" + "telesrv/internal/compat/tdesktop" "telesrv/internal/domain" ) @@ -69,33 +72,34 @@ func (r *Router) onMessagesGetSavedReactionTags(ctx context.Context, req *tg.Mes if err != nil { return nil, internalErr() } + var savedPeer domain.Peer if peer, ok := req.GetPeer(); ok && peer != nil { - if _, err := r.checkedDomainPeerFromInputPeer(ctx, userID, peer); err != nil { + savedPeer, err = r.checkedDomainPeerFromInputPeer(ctx, userID, peer) + if err != nil { return nil, err } + } + if r.deps.Messages == nil { return savedReactionTagsEmpty(req.Hash), nil } - if r.deps.Channels == nil { - return savedReactionTagsEmpty(req.Hash), nil - } - tags, err := r.deps.Channels.SavedReactionTags(ctx, userID, domain.MaxSavedReactionTags) + tags, err := r.deps.Messages.SavedReactionTags(ctx, userID, savedPeer, domain.MaxSavedReactionTags) if err != nil { - return nil, channelInvalidErr(err) + return nil, messageReactionErr(err) } - return savedReactionTagsFromDomain(tags, req.Hash), nil + return savedReactionTagsFromDomain(tags, req.Hash, savedPeer.ID == 0), nil } func (r *Router) onMessagesGetDefaultTagReactions(ctx context.Context, hash int64) (tg.MessagesReactionsClass, error) { if _, _, err := r.currentUserID(ctx); err != nil { return nil, internalErr() } - return messagesReactionsEmpty(hash), nil + return messagesReactionsFromDomain( + r.reactionsWithCatalogFallback(ctx, nil, domain.MaxChannelMessageReactionsPerUser), + hash, + ), nil } -func messagesReactionsEmpty(hash int64) tg.MessagesReactionsClass { - if hash != 0 { - return &tg.MessagesReactionsNotModified{} - } +func messagesReactionsEmpty(_ int64) tg.MessagesReactionsClass { return &tg.MessagesReactions{ Hash: 0, Reactions: []tg.ReactionClass{}, @@ -127,8 +131,15 @@ func savedReactionTagsEmpty(_ int64) tg.MessagesSavedReactionTagsClass { } } -func savedReactionTagsFromDomain(tags []domain.SavedReactionTag, requestHash int64) tg.MessagesSavedReactionTagsClass { - hash := savedReactionTagListHash(tags) +func savedReactionTagsFromDomain(tags []domain.SavedReactionTag, requestHash int64, includeTitles bool) tg.MessagesSavedReactionTagsClass { + tags = append([]domain.SavedReactionTag(nil), tags...) + sort.SliceStable(tags, func(i, j int) bool { + if tags[i].Count != tags[j].Count { + return tags[i].Count > tags[j].Count + } + return savedReactionTagLongID(tags[i].Reaction) > savedReactionTagLongID(tags[j].Reaction) + }) + hash := savedReactionTagListHash(tags, includeTitles) if hash != 0 && requestHash == hash { return &tg.MessagesSavedReactionTagsNotModified{} } @@ -142,7 +153,7 @@ func savedReactionTagsFromDomain(tags []domain.SavedReactionTag, requestHash int Reaction: reaction, Count: tag.Count, } - if tag.Title != "" { + if includeTitles && tag.Title != "" { item.SetTitle(tag.Title) } out = append(out, item) @@ -239,38 +250,47 @@ func messageReactionListHash(reactions []domain.MessageReaction) int64 { if len(reactions) == 0 { return 0 } - h := fnv.New64a() + var hash uint64 for _, reaction := range reactions { - _, _ = h.Write([]byte(reaction.Type)) - _, _ = h.Write([]byte{0}) - _, _ = h.Write([]byte(reaction.Value())) - _, _ = h.Write([]byte{0xff}) + hash = telegramListHashNext(hash, savedReactionTagLongID(reaction)) } - sum := int64(h.Sum64() & 0x7fffffffffffffff) - if sum == 0 { - return 1 - } - return sum + return int64(hash) } -func savedReactionTagListHash(tags []domain.SavedReactionTag) int64 { +func savedReactionTagListHash(tags []domain.SavedReactionTag, includeTitles bool) int64 { if len(tags) == 0 { return 0 } - h := fnv.New64a() + var hash uint64 for _, tag := range tags { - _, _ = h.Write([]byte(tag.Reaction.Type)) - _, _ = h.Write([]byte{0}) - _, _ = h.Write([]byte(tag.Reaction.Value())) - _, _ = h.Write([]byte{0}) - _, _ = h.Write([]byte(tag.Title)) - _, _ = h.Write([]byte{0}) - _, _ = h.Write([]byte(strconv.Itoa(tag.Count))) - _, _ = h.Write([]byte{0xff}) + hash = telegramListHashNext(hash, savedReactionTagLongID(tag.Reaction)) + if includeTitles && tag.Title != "" { + hash = telegramListHashNext(hash, md5LongID(tag.Title)) + } + hash = telegramListHashNext(hash, uint64(tag.Count)) } - sum := int64(h.Sum64() & 0x7fffffffffffffff) - if sum == 0 { - return 1 - } - return sum + return int64(hash) +} + +func savedReactionTagLongID(reaction domain.MessageReaction) uint64 { + switch reaction.Type { + case domain.MessageReactionEmoji: + return md5LongID(strings.ReplaceAll(reaction.Emoticon, "\ufe0f", "")) + case domain.MessageReactionCustomEmoji: + return uint64(reaction.DocumentID) + default: + return 0 + } +} + +func md5LongID(value string) uint64 { + sum := md5.Sum([]byte(value)) + return binary.BigEndian.Uint64(sum[:8]) +} + +func telegramListHashNext(hash, id uint64) uint64 { + hash ^= hash >> 21 + hash ^= hash << 35 + hash ^= hash >> 4 + return hash + id } diff --git a/internal/rpc/messages_reactions_helpers.go b/internal/rpc/messages_reactions_helpers.go index aace5861..c23c0d69 100644 --- a/internal/rpc/messages_reactions_helpers.go +++ b/internal/rpc/messages_reactions_helpers.go @@ -3,10 +3,12 @@ package rpc import ( "context" "errors" - "github.com/iamxvbaba/td/tg" "strings" - "telesrv/internal/domain" "unicode/utf8" + + "github.com/iamxvbaba/td/tg" + + "telesrv/internal/domain" ) func (r *Router) onMessagesUpdateSavedReactionTag(ctx context.Context, req *tg.MessagesUpdateSavedReactionTagRequest) (bool, error) { @@ -18,8 +20,8 @@ func (r *Router) onMessagesUpdateSavedReactionTag(ctx context.Context, req *tg.M if err != nil { return false, err } - if reaction.Type != domain.MessageReactionEmoji { - return false, reactionInvalidErr() + if !r.viewerPremium(ctx, userID) { + return false, premiumAccountRequiredErr() } title, ok := req.GetTitle() if !ok { @@ -28,13 +30,13 @@ func (r *Router) onMessagesUpdateSavedReactionTag(ctx context.Context, req *tg.M if utf8.RuneCountInString(title) > maxSavedReactionTagTitle { return false, limitInvalidErr() } - if r.deps.Channels != nil { - if err := r.deps.Channels.UpdateSavedReactionTag(ctx, userID, domain.SavedReactionTag{ + if r.deps.Messages != nil { + if err := r.deps.Messages.UpdateSavedReactionTag(ctx, userID, domain.SavedReactionTag{ UserID: userID, Reaction: reaction, Title: title, }); err != nil { - return false, channelInvalidErr(err) + return false, messageReactionErr(err) } } r.pushUserUpdates(ctx, userID, &tg.Updates{ @@ -166,6 +168,8 @@ func messageReactionErr(err error) error { switch { case errors.Is(err, domain.ErrMessageIDInvalid): return messageIDInvalidErr() + case errors.Is(err, domain.ErrReactionInvalid): + return reactionInvalidErr() default: return internalErr() } diff --git a/internal/rpc/messages_reactions_send.go b/internal/rpc/messages_reactions_send.go index e81f948f..a1c93c8b 100644 --- a/internal/rpc/messages_reactions_send.go +++ b/internal/rpc/messages_reactions_send.go @@ -29,6 +29,10 @@ func (r *Router) onMessagesSendReaction(ctx context.Context, req *tg.MessagesSen // reactions_user_max_premium=3),否则客户端允许的多 reaction 会被静默裁剪。 perUserMax := domain.MessageReactionsUserMax(r.viewerPremium(ctx, userID)) reactions = domain.TrimMessageReactionsToUserMax(reactions, perUserMax) + if peer.Type == domain.PeerTypeUser && peer.ID == userID && + len(reactions) > 0 && !r.viewerPremium(ctx, userID) { + return nil, premiumAccountRequiredErr() + } date := int(r.clock.Now().Unix()) if peer.Type == domain.PeerTypeChannel && r.deps.Channels != nil { res, err := r.deps.Channels.SetMessageReactions(ctx, userID, domain.SetChannelMessageReactionsRequest{ @@ -54,7 +58,7 @@ func (r *Router) onMessagesSendReaction(ctx context.Context, req *tg.MessagesSen return updates, nil } if peer.Type == domain.PeerTypeUser && r.deps.Messages != nil { - if len(reactions) == 0 && r.shouldSuppressTransientPrivateReactionClear(userID, peer, req.MsgID, date) { + if peer.ID != userID && len(reactions) == 0 && r.shouldSuppressTransientPrivateReactionClear(userID, peer, req.MsgID, date) { res, err := r.deps.Messages.GetMessageReactions(ctx, userID, domain.PrivateMessageReactionsRequest{ OwnerUserID: userID, Peer: peer, @@ -65,7 +69,7 @@ func (r *Router) onMessagesSendReaction(ctx context.Context, req *tg.MessagesSen } return r.privateMessagesReactionsUpdates(ctx, userID, peer, res, []int{req.MsgID}), nil } - if req.Big && len(reactions) > 0 { + if peer.ID != userID && req.Big && len(reactions) > 0 { r.rememberTransientPrivateBigReaction(userID, peer, req.MsgID, date) } res, err := r.deps.Messages.SetMessageReactions(ctx, userID, domain.SetPrivateMessageReactionsRequest{ @@ -84,19 +88,13 @@ func (r *Router) onMessagesSendReaction(ctx context.Context, req *tg.MessagesSen if len(reactions) == 0 { r.forgetTransientPrivateBigReaction(userID, peer, req.MsgID) } - if err := r.recordMessageReactionUse(ctx, userID, reactions, req.GetAddToRecent(), date); err != nil { - return nil, internalErr() + // Saved Messages reactions are private tags, not ordinary reaction usage. + if peer.ID != userID { + if err := r.recordMessageReactionUse(ctx, userID, reactions, req.GetAddToRecent(), date); err != nil { + return nil, internalErr() + } } - recordedEvents, err := r.recordPrivateMessageReactionEvents(ctx, userID, res) - if err != nil { - return nil, internalErr() - } - // reaction 事件占双方账号 pts 但 updateMessageReactions 不带 pts; - // 在线直推必须附 pts 簿记,否则双方下一条带 pts 的更新被判空洞。 updates := r.privateMessageReactionsUpdates(ctx, userID, peer, res) - if updates != nil { - updates.Updates = appendAuxPtsBookkeeping(updates.Updates, recordedEvents[userID]) - } r.pushUserUpdates(ctx, userID, updates) for _, msg := range res.Messages { if msg.OwnerUserID == 0 || msg.OwnerUserID == userID { @@ -104,9 +102,6 @@ func (r *Router) onMessagesSendReaction(ctx context.Context, req *tg.MessagesSen } viewerPeer := msg.Peer viewerUpdates := r.privateMessageReactionsUpdates(ctx, msg.OwnerUserID, viewerPeer, res) - if viewerUpdates != nil { - viewerUpdates.Updates = appendAuxPtsBookkeeping(viewerUpdates.Updates, recordedEvents[msg.OwnerUserID]) - } r.pushUserUpdates(ctx, msg.OwnerUserID, viewerUpdates) } return updates, nil @@ -223,33 +218,6 @@ func (r *Router) recordMessageReactionUse(ctx context.Context, userID int64, rea return recorder.RecordMessageReactionUse(ctx, userID, reactions, addToRecent, date) } -func (r *Router) recordPrivateMessageReactionEvents(ctx context.Context, requestUserID int64, res domain.PrivateMessageReactionsResult) (map[int64]domain.UpdateEvent, error) { - if r.deps.Updates == nil { - return nil, nil - } - recorder, ok := r.deps.Updates.(messageReactionUpdateRecorder) - if !ok { - return nil, nil - } - authKeyID, _ := AuthKeyIDFrom(ctx) - events := make(map[int64]domain.UpdateEvent, len(res.Messages)) - for _, msg := range res.Messages { - if msg.OwnerUserID == 0 || msg.ID == 0 { - continue - } - eventAuthKeyID := [8]byte{} - if msg.OwnerUserID == requestUserID { - eventAuthKeyID = authKeyID - } - event, _, err := recorder.RecordMessageReactions(ctx, eventAuthKeyID, msg.OwnerUserID, msg) - if err != nil { - return nil, err - } - events[msg.OwnerUserID] = event - } - return events, nil -} - func (r *Router) channelMessageReactionsUpdates(ctx context.Context, viewerUserID int64, res domain.ChannelMessageReactionsResult) *tg.Updates { ids := []int{res.Message.ID} if res.Message.ID <= 0 && len(res.Messages) > 0 { @@ -311,6 +279,7 @@ func minifyChannelReactionsResult(res domain.ChannelMessageReactionsResult) doma } out := domain.ChannelMessageReactions{ CanSeeList: in.CanSeeList, + AsTags: in.AsTags, Results: make([]domain.ChannelMessageReactionCount, 0, len(in.Results)), Recent: make([]domain.ChannelMessagePeerReaction, 0, len(in.Recent)), } diff --git a/internal/rpc/messages_reactions_test.go b/internal/rpc/messages_reactions_test.go index d0cdeedc..d102a582 100644 --- a/internal/rpc/messages_reactions_test.go +++ b/internal/rpc/messages_reactions_test.go @@ -5,82 +5,25 @@ import ( "github.com/iamxvbaba/td/clock" "github.com/iamxvbaba/td/proto" "github.com/iamxvbaba/td/tg" + "github.com/iamxvbaba/td/tgerr" "go.uber.org/zap/zaptest" - appchannels "telesrv/internal/app/channels" + appusers "telesrv/internal/app/users" "telesrv/internal/domain" "telesrv/internal/store/memory" "testing" "time" ) -func TestUpdatesDifferenceIncludesReactionMessageAndUpdate(t *testing.T) { - const ( - aliceID = int64(1000000001) - bobID = int64(1000000002) - ) - reaction := domain.MessageReaction{Type: domain.MessageReactionEmoji, Emoticon: "\U0001f44d"} - reactions := domain.ChannelMessageReactions{ - CanSeeList: true, - Results: []domain.ChannelMessageReactionCount{{ - Reaction: reaction, - Count: 1, - ChosenOrder: 1, - }}, - Recent: []domain.ChannelMessagePeerReaction{{ - UserID: bobID, - Reaction: reaction, - My: true, - ChosenOrder: 1, - Date: 1700000310, - }}, - } - msg := domain.Message{ - ID: 68, - UID: 7001, - OwnerUserID: aliceID, - Peer: domain.Peer{Type: domain.PeerTypeUser, ID: bobID}, - From: domain.Peer{Type: domain.PeerTypeUser, ID: aliceID}, - Date: 1700000300, - Body: "rx", - Reactions: &reactions, - } - got, ok := tgUpdatesDifference(0, domain.UpdateDifference{ - State: domain.UpdateState{Pts: 9, Date: 1700000310}, - Events: []domain.UpdateEvent{{ - UserID: aliceID, - Type: domain.UpdateEventMessageReactions, - Pts: 9, - PtsCount: 1, - Date: 1700000310, - Peer: domain.Peer{Type: domain.PeerTypeUser, ID: bobID}, - Message: msg, - }}, - }).(*tg.UpdatesDifference) - if !ok { - t.Fatalf("difference = %T, want *tg.UpdatesDifference", got) - } - if len(got.NewMessages) != 1 || len(got.OtherUpdates) != 1 { - t.Fatalf("difference messages/updates = %d/%d, want 1/1", len(got.NewMessages), len(got.OtherUpdates)) - } - wireMsg, ok := got.NewMessages[0].(*tg.Message) - if !ok || wireMsg.ID != msg.ID { - t.Fatalf("message = %T %+v, want message %d", got.NewMessages[0], got.NewMessages[0], msg.ID) - } - msgReactions, ok := wireMsg.GetReactions() - if !ok || len(msgReactions.Results) != 1 || msgReactions.Results[0].Count != 1 || msgReactions.Results[0].ChosenOrder != 1 { - t.Fatalf("message reactions = %+v set=%v, want chosen reaction", msgReactions, ok) - } - update, ok := got.OtherUpdates[0].(*tg.UpdateMessageReactions) - if !ok || update.MsgID != msg.ID || len(update.Reactions.Results) != 1 || update.Reactions.Results[0].ChosenOrder != 1 { - t.Fatalf("reaction update = %T %+v, want update for msg %d", got.OtherUpdates[0], got.OtherUpdates[0], msg.ID) - } -} - func TestMessagesUpdateSavedReactionTagPersistsAndPushesRefresh(t *testing.T) { - const userID = int64(1000000001) + userID, users := newReactionTestUsers(t, true) sessions := &captureSessions{} + reaction := domain.MessageReaction{Type: domain.MessageReactionEmoji, Emoticon: "\U0001f44d"} + messages := &captureMessages{savedTags: []domain.SavedReactionTag{{ + UserID: userID, Reaction: reaction, Count: 1, + }}} r := New(Config{}, Deps{ - Channels: appchannels.NewService(memory.NewChannelStore()), + Messages: messages, + Users: users, Sessions: sessions, }, zaptest.NewLogger(t), clock.System) @@ -121,6 +64,198 @@ func TestMessagesUpdateSavedReactionTagPersistsAndPushesRefresh(t *testing.T) { } } +func TestMessagesUpdateSavedReactionTagAcceptsCustomEmoji(t *testing.T) { + userID, users := newReactionTestUsers(t, true) + custom := domain.MessageReaction{Type: domain.MessageReactionCustomEmoji, DocumentID: 90001} + messages := &captureMessages{savedTags: []domain.SavedReactionTag{{ + UserID: userID, Reaction: custom, Count: 1, + }}} + r := New(Config{}, Deps{Messages: messages, Users: users}, zaptest.NewLogger(t), clock.System) + req := &tg.MessagesUpdateSavedReactionTagRequest{ + Reaction: &tg.ReactionCustomEmoji{DocumentID: custom.DocumentID}, + } + req.SetTitle("Work") + ok, err := r.onMessagesUpdateSavedReactionTag(WithUserID(context.Background(), userID), req) + if err != nil || !ok { + t.Fatalf("rename custom saved tag = %v, %v", ok, err) + } + if messages.updatedSavedTag.Reaction.Key() != custom.Key() || messages.updatedSavedTag.Title != "Work" { + t.Fatalf("updated custom tag = %+v", messages.updatedSavedTag) + } +} + +func newReactionTestUsers(t *testing.T, premium bool) (int64, UsersService) { + t.Helper() + users := memory.NewUserStore() + user, err := users.Create(context.Background(), domain.User{ + Phone: "+15550000001", + FirstName: "Reaction", + }) + if err != nil { + t.Fatalf("create reaction test user: %v", err) + } + if premium { + if _, err := users.SetPremiumUntil(context.Background(), user.ID, int(time.Now().Add(time.Hour).Unix())); err != nil { + t.Fatalf("set reaction test premium: %v", err) + } + } + return user.ID, appusers.NewService(users) +} + +func TestMessagesSendReactionSavedMessageUsesTagsWithoutPTSBookkeeping(t *testing.T) { + userID, users := newReactionTestUsers(t, true) + messages := &captureMessages{} + sessions := &captureSessions{} + r := New(Config{}, Deps{ + Messages: messages, + Users: users, + Sessions: sessions, + }, zaptest.NewLogger(t), fixedClock{now: time.Unix(1_700_000_200, 0)}) + req := &tg.MessagesSendReactionRequest{ + Peer: &tg.InputPeerSelf{}, + MsgID: 7, + Reaction: []tg.ReactionClass{&tg.ReactionCustomEmoji{DocumentID: 90001}}, + } + req.SetReaction(req.Reaction) + got, err := r.onMessagesSendReaction( + WithSessionID(WithUserID(context.Background(), userID), 72), + req, + ) + if err != nil { + t.Fatalf("send saved tag: %v", err) + } + updates, ok := got.(*tg.Updates) + if !ok || len(updates.Updates) != 1 { + t.Fatalf("saved tag result = %T %+v, want one update", got, got) + } + update, ok := updates.Updates[0].(*tg.UpdateMessageReactions) + if !ok || !update.Reactions.ReactionsAsTags || len(update.Reactions.Results) != 1 { + t.Fatalf("saved tag update = %T %+v, want reactions_as_tags", updates.Updates[0], updates.Updates[0]) + } + for _, item := range updates.Updates { + if deleted, ok := item.(*tg.UpdateDeleteMessages); ok { + t.Fatalf("saved tag emitted fake delete pts bookkeeping: %+v", deleted) + } + } + if messages.setReactionReq.Peer != (domain.Peer{Type: domain.PeerTypeUser, ID: userID}) { + t.Fatalf("saved tag peer = %+v, want self", messages.setReactionReq.Peer) + } + push := sessions.snapshot() + pushed, ok := push.message.(*tg.Updates) + if push.userID != userID || push.sessionID != 72 || !ok || len(pushed.Updates) != 1 { + t.Fatalf("saved tag push = user %d exclude %d %T %+v", push.userID, push.sessionID, push.message, push.message) + } + pushedReaction, ok := pushed.Updates[0].(*tg.UpdateMessageReactions) + if !ok || !pushedReaction.Reactions.ReactionsAsTags { + t.Fatalf("saved tag pushed update = %T %+v, want reactions_as_tags", pushed.Updates[0], pushed.Updates[0]) + } +} + +func TestMessagesSendReactionSavedMessageRequiresPremiumButAllowsClear(t *testing.T) { + userID, users := newReactionTestUsers(t, false) + messages := &captureMessages{} + r := New(Config{}, Deps{ + Messages: messages, + Users: users, + }, zaptest.NewLogger(t), clock.System) + add := &tg.MessagesSendReactionRequest{ + Peer: &tg.InputPeerSelf{}, + MsgID: 8, + Reaction: []tg.ReactionClass{&tg.ReactionEmoji{Emoticon: "👍"}}, + } + add.SetReaction(add.Reaction) + if _, err := r.onMessagesSendReaction(WithUserID(context.Background(), userID), add); !tgerr.Is(err, "PREMIUM_ACCOUNT_REQUIRED") { + t.Fatalf("non-premium add err = %v, want PREMIUM_ACCOUNT_REQUIRED", err) + } + clear := &tg.MessagesSendReactionRequest{Peer: &tg.InputPeerSelf{}, MsgID: 8} + if _, err := r.onMessagesSendReaction(WithUserID(context.Background(), userID), clear); err != nil { + t.Fatalf("non-premium clear: %v", err) + } +} + +func TestSavedReactionTagHashMatchesClientShape(t *testing.T) { + plain := domain.MessageReaction{Type: domain.MessageReactionEmoji, Emoticon: "❤️"} + withoutVariation := domain.MessageReaction{Type: domain.MessageReactionEmoji, Emoticon: "❤"} + if got, want := messageReactionListHash([]domain.MessageReaction{plain}), messageReactionListHash([]domain.MessageReaction{withoutVariation}); got != want { + t.Fatalf("emoji variation-selector hash = %d, want normalized %d", got, want) + } + tags := []domain.SavedReactionTag{{ + Reaction: plain, + Title: "Love", + Count: 3, + }} + full := savedReactionTagsFromDomain(tags, 0, true) + page, ok := full.(*tg.MessagesSavedReactionTags) + if !ok || page.Hash == 0 || len(page.Tags) != 1 || page.Tags[0].Title != "Love" { + t.Fatalf("saved tag page = %T %+v", full, full) + } + if page.Hash != -4770309592622053821 { + t.Fatalf("saved tag client hash = %d, want -4770309592622053821", page.Hash) + } + if cached := savedReactionTagsFromDomain(tags, page.Hash, true); cached == nil { + t.Fatal("cached saved tag result is nil") + } else if _, ok := cached.(*tg.MessagesSavedReactionTagsNotModified); !ok { + t.Fatalf("cached saved tag result = %T, want not modified", cached) + } + perPeer := savedReactionTagsFromDomain(tags, 0, false) + peerPage, ok := perPeer.(*tg.MessagesSavedReactionTags) + if !ok || peerPage.Tags[0].Title != "" || peerPage.Hash == page.Hash { + t.Fatalf("per-peer saved tags = %T %+v, want title omitted and scope hash", perPeer, perPeer) + } +} + +func TestMessageFilterFromSearchRequestParsesSavedTagsAndPeer(t *testing.T) { + const userID = int64(1000000001) + r := New(Config{}, Deps{}, zaptest.NewLogger(t), clock.System) + req := &tg.MessagesSearchRequest{ + Peer: &tg.InputPeerSelf{}, + Q: "needle", + MinDate: 100, + MaxDate: 200, + Limit: 50, + Filter: &tg.InputMessagesFilterEmpty{}, + } + req.SetSavedPeerID(&tg.InputPeerSelf{}) + req.SetSavedReaction([]tg.ReactionClass{ + &tg.ReactionEmoji{Emoticon: "👍"}, + &tg.ReactionCustomEmoji{DocumentID: 90001}, + }) + filter, err := r.messageFilterFromSearchRequest(WithUserID(context.Background(), userID), userID, req) + if err != nil { + t.Fatalf("parse saved search filter: %v", err) + } + if filter.Peer != (domain.Peer{Type: domain.PeerTypeUser, ID: userID}) || + filter.SavedPeer != filter.Peer || len(filter.SavedReactions) != 2 || + filter.MinDate != 100 || filter.MaxDate != 200 { + t.Fatalf("saved search filter = %+v", filter) + } + + req.Peer = &tg.InputPeerUser{UserID: userID + 1, AccessHash: 1} + if _, err := r.messageFilterFromSearchRequest(WithUserID(context.Background(), userID), userID, req); !tgerr.Is(err, "PEER_ID_INVALID") { + t.Fatalf("non-self saved search err = %v, want PEER_ID_INVALID", err) + } +} + +func TestMessagesGetDefaultTagReactionsReturnsHashableCatalog(t *testing.T) { + r := New(Config{}, Deps{}, zaptest.NewLogger(t), clock.System) + ctx := WithUserID(context.Background(), 1000000001) + got, err := r.onMessagesGetDefaultTagReactions(ctx, 0) + if err != nil { + t.Fatalf("get default tag reactions: %v", err) + } + page, ok := got.(*tg.MessagesReactions) + if !ok || page.Hash == 0 || len(page.Reactions) == 0 { + t.Fatalf("default tags = %T %+v, want non-empty hashable catalog", got, got) + } + cached, err := r.onMessagesGetDefaultTagReactions(ctx, page.Hash) + if err != nil { + t.Fatalf("get cached default tags: %v", err) + } + if _, ok := cached.(*tg.MessagesReactionsNotModified); !ok { + t.Fatalf("cached default tags = %T, want not modified", cached) + } +} + func TestMessagesSendReactionPrivatePeerReturnsReactionUpdate(t *testing.T) { const ( userID = int64(1000000001) diff --git a/internal/rpc/messages_register.go b/internal/rpc/messages_register.go index 8e025d05..67c6e510 100644 --- a/internal/rpc/messages_register.go +++ b/internal/rpc/messages_register.go @@ -702,7 +702,10 @@ func (r *Router) registerMessages(d *tlprofile.Dispatcher) { if err != nil { return nil, internalErr() } - filter := r.messageFilterFromSearchRequest(userID, req) + filter, err := r.messageFilterFromSearchRequest(ctx, userID, req) + if err != nil { + return nil, err + } if filter.HasPeer && filter.Peer.Type == domain.PeerTypeChannel { if r.deps.Channels == nil { return messagesNotModifiedOrEmpty(req.Hash), nil diff --git a/internal/rpc/rpc_testkit_messages_test.go b/internal/rpc/rpc_testkit_messages_test.go index e871e205..b72c80f9 100644 --- a/internal/rpc/rpc_testkit_messages_test.go +++ b/internal/rpc/rpc_testkit_messages_test.go @@ -27,6 +27,10 @@ type captureMessages struct { setReactionRes domain.PrivateMessageReactionsResult getReactionReq domain.PrivateMessageReactionsRequest getReactionRes domain.PrivateMessageReactionsResult + savedTagPeer domain.Peer + savedTags []domain.SavedReactionTag + updatedSavedTag domain.SavedReactionTag + savedTagErr error getMessagesCalls int getMessagesIDs [][]int getMessagesListed bool @@ -432,7 +436,12 @@ func (s *captureMessages) SetMessageReactions(_ context.Context, userID int64, r s.setReactionReq = req if len(s.setReactionRes.Messages) == 0 { if len(req.Reactions) == 0 { - reactions := domain.ChannelMessageReactions{CanSeeList: true, Results: []domain.ChannelMessageReactionCount{}, Recent: []domain.ChannelMessagePeerReaction{}} + reactions := domain.ChannelMessageReactions{ + CanSeeList: req.Peer.ID != userID, + AsTags: req.Peer.ID == userID, + Results: []domain.ChannelMessageReactionCount{}, + Recent: []domain.ChannelMessagePeerReaction{}, + } s.setReactionRes = domain.PrivateMessageReactionsResult{ Messages: []domain.Message{{ ID: req.MessageID, @@ -447,20 +456,23 @@ func (s *captureMessages) SetMessageReactions(_ context.Context, userID int64, r return s.setReactionRes, nil } reactions := domain.ChannelMessageReactions{ - CanSeeList: true, + CanSeeList: req.Peer.ID != userID, + AsTags: req.Peer.ID == userID, Results: []domain.ChannelMessageReactionCount{{ Reaction: req.Reactions[0], Count: 1, ChosenOrder: 1, }}, - Recent: []domain.ChannelMessagePeerReaction{{ + } + if req.Peer.ID != userID { + reactions.Recent = []domain.ChannelMessagePeerReaction{{ UserID: userID, Reaction: req.Reactions[0], My: true, Big: req.Big, ChosenOrder: 1, Date: req.Date, - }}, + }} } s.setReactionRes = domain.PrivateMessageReactionsResult{ Messages: []domain.Message{{ @@ -503,6 +515,22 @@ func (s *captureMessages) GetMessageReactions(_ context.Context, userID int64, r return s.getReactionRes, nil } +func (s *captureMessages) SavedReactionTags(_ context.Context, _ int64, savedPeer domain.Peer, _ int) ([]domain.SavedReactionTag, error) { + s.savedTagPeer = savedPeer + return append([]domain.SavedReactionTag(nil), s.savedTags...), s.savedTagErr +} + +func (s *captureMessages) UpdateSavedReactionTag(_ context.Context, _ int64, tag domain.SavedReactionTag) error { + s.updatedSavedTag = tag + for i := range s.savedTags { + if s.savedTags[i].Reaction.Key() == tag.Reaction.Key() { + s.savedTags[i].Title = tag.Title + return s.savedTagErr + } + } + return s.savedTagErr +} + func (s *captureMessages) EditMessage(_ context.Context, userID int64, req domain.EditMessageRequest) (domain.EditMessageResult, error) { s.editReq = req if s.editRes.OwnerUserID == 0 { diff --git a/internal/rpc/update_peer_refs.go b/internal/rpc/update_peer_refs.go index 98916767..6e594e8e 100644 --- a/internal/rpc/update_peer_refs.go +++ b/internal/rpc/update_peer_refs.go @@ -45,9 +45,6 @@ func (r *Router) enrichUpdateEventsWithPeerCache(ctx context.Context, viewerUser } } } - if out[i].Type == domain.UpdateEventMessageReactions { - out[i] = r.enrichMessageReactionEvent(ctx, viewerUserID, out[i]) - } if out[i].Type == domain.UpdateEventMessagePoll { out[i] = r.enrichMessagePollEvent(ctx, viewerUserID, out[i]) } @@ -115,36 +112,6 @@ type updateEventPeerRefs struct { channelIDs map[int64]struct{} } -func (r *Router) enrichMessageReactionEvent(ctx context.Context, viewerUserID int64, event domain.UpdateEvent) domain.UpdateEvent { - if r.deps.Messages == nil || event.Message.ID <= 0 { - return event - } - peer := event.Message.Peer - if peer.Type == "" || peer.ID == 0 { - peer = event.Peer - } - if peer.Type != domain.PeerTypeUser || peer.ID == 0 { - return event - } - res, err := r.deps.Messages.GetMessageReactions(ctx, viewerUserID, domain.PrivateMessageReactionsRequest{ - OwnerUserID: viewerUserID, - Peer: peer, - IDs: []int{event.Message.ID}, - }) - if err != nil { - return event - } - for _, msg := range res.Messages { - if msg.OwnerUserID == viewerUserID && msg.ID == event.Message.ID { - msg.Pts = event.Pts - event.Message = msg - event.Peer = msg.Peer - return event - } - } - return event -} - // enrichMessagePollEvent 在 difference 重放时按 viewer 重载消息(media 含最新 poll 权威态与 // viewer 门控),与 reaction 事件 enrich 同构。 func (r *Router) enrichMessagePollEvent(ctx context.Context, viewerUserID int64, event domain.UpdateEvent) domain.UpdateEvent { diff --git a/internal/store/channel.go b/internal/store/channel.go index 519af694..80e3ca2e 100644 --- a/internal/store/channel.go +++ b/internal/store/channel.go @@ -77,8 +77,6 @@ type ChannelStore interface { ListTopMessageReactions(ctx context.Context, userID int64, limit int) ([]domain.MessageReaction, error) ListRecentMessageReactions(ctx context.Context, userID int64, limit int) ([]domain.MessageReaction, error) ClearRecentMessageReactions(ctx context.Context, userID int64) error - ListSavedReactionTags(ctx context.Context, userID int64, limit int) ([]domain.SavedReactionTag, error) - UpsertSavedReactionTag(ctx context.Context, tag domain.SavedReactionTag) error GetPremiumBoostStatus(ctx context.Context, viewerUserID, channelID int64, now int) (domain.PremiumBoostStatus, error) ListPremiumBoosts(ctx context.Context, viewerUserID, channelID int64, gifts bool, offset string, limit, now int) (domain.PremiumBoostList, error) GetPremiumMyBoosts(ctx context.Context, userID int64, now, premiumUntil int) (domain.PremiumMyBoosts, error) diff --git a/internal/store/memory/channel_reactions.go b/internal/store/memory/channel_reactions.go index 75f34584..37d76a4e 100644 --- a/internal/store/memory/channel_reactions.go +++ b/internal/store/memory/channel_reactions.go @@ -632,54 +632,6 @@ func (s *ChannelStore) ClearRecentMessageReactions(_ context.Context, userID int return nil } -func (s *ChannelStore) ListSavedReactionTags(_ context.Context, userID int64, limit int) ([]domain.SavedReactionTag, error) { - if userID == 0 { - return nil, domain.ErrChannelInvalid - } - if limit <= 0 { - return []domain.SavedReactionTag{}, nil - } - if limit > domain.MaxSavedReactionTags { - limit = domain.MaxSavedReactionTags - } - s.mu.RLock() - defer s.mu.RUnlock() - rows := make([]domain.SavedReactionTag, 0, len(s.savedTags[userID])) - for _, row := range s.savedTags[userID] { - rows = append(rows, row) - } - sort.Slice(rows, func(i, j int) bool { - if rows[i].Count != rows[j].Count { - return rows[i].Count > rows[j].Count - } - if rows[i].Reaction.Type != rows[j].Reaction.Type { - return rows[i].Reaction.Type < rows[j].Reaction.Type - } - return rows[i].Reaction.Value() < rows[j].Reaction.Value() - }) - if len(rows) > limit { - rows = rows[:limit] - } - return rows, nil -} - -func (s *ChannelStore) UpsertSavedReactionTag(_ context.Context, tag domain.SavedReactionTag) error { - if tag.UserID == 0 || tag.Reaction.Type != domain.MessageReactionEmoji || strings.TrimSpace(tag.Reaction.Emoticon) == "" { - return domain.ErrChannelInvalid - } - s.mu.Lock() - defer s.mu.Unlock() - if s.savedTags[tag.UserID] == nil { - s.savedTags[tag.UserID] = make(map[string]domain.SavedReactionTag) - } - tag.Reaction.Emoticon = strings.TrimSpace(tag.Reaction.Emoticon) - if tag.Count < 0 { - tag.Count = 0 - } - s.savedTags[tag.UserID][messageReactionKey(tag.Reaction)] = tag - return nil -} - func (s *ChannelStore) ListChannelUnreadReactions(_ context.Context, viewerUserID int64, filter domain.ChannelUnreadReactionsFilter) (domain.ChannelHistory, error) { s.mu.RLock() defer s.mu.RUnlock() diff --git a/internal/store/memory/channel_store.go b/internal/store/memory/channel_store.go index 77788026..c9a38cf5 100644 --- a/internal/store/memory/channel_store.go +++ b/internal/store/memory/channel_store.go @@ -76,7 +76,6 @@ type ChannelStore struct { paidReactions map[int64]map[int]map[int64]memoryPaidReaction top map[int64]map[string]domain.TopMessageReaction recent map[int64]map[string]domain.RecentMessageReaction - savedTags map[int64]map[string]domain.SavedReactionTag mentions map[int64]map[int64]map[int]memoryMention msgViews map[int64]map[int]int msgViewers map[int64]map[int]map[int64]struct{} @@ -124,7 +123,6 @@ func NewChannelStore() *ChannelStore { paidReactions: make(map[int64]map[int]map[int64]memoryPaidReaction), top: make(map[int64]map[string]domain.TopMessageReaction), recent: make(map[int64]map[string]domain.RecentMessageReaction), - savedTags: make(map[int64]map[string]domain.SavedReactionTag), mentions: make(map[int64]map[int64]map[int]memoryMention), msgViews: make(map[int64]map[int]int), msgViewers: make(map[int64]map[int]map[int64]struct{}), diff --git a/internal/store/memory/message_delete.go b/internal/store/memory/message_delete.go index 0bd00140..076204ae 100644 --- a/internal/store/memory/message_delete.go +++ b/internal/store/memory/message_delete.go @@ -52,6 +52,12 @@ func (s *MessageStore) finishMemoryDeleteLocked(res domain.DeleteMessagesResult, idsByOwner := make(map[int64][]int) peersByOwner := make(map[int64]map[domain.Peer]struct{}) for _, row := range deleted { + if byMessage := s.savedMessageTags[row.userID]; byMessage != nil { + delete(byMessage, row.id) + if len(byMessage) == 0 { + delete(s.savedMessageTags, row.userID) + } + } idsByOwner[row.userID] = append(idsByOwner[row.userID], row.id) if peersByOwner[row.userID] == nil { peersByOwner[row.userID] = make(map[domain.Peer]struct{}) diff --git a/internal/store/memory/message_history.go b/internal/store/memory/message_history.go index 53bec30a..5a7d21fa 100644 --- a/internal/store/memory/message_history.go +++ b/internal/store/memory/message_history.go @@ -4,8 +4,9 @@ import ( "context" "sort" "strings" - "telesrv/internal/domain" "time" + + "telesrv/internal/domain" ) func (s *MessageStore) GetByIDs(_ context.Context, userID int64, ids []int) (domain.MessageList, error) { @@ -282,6 +283,12 @@ func filterMessageList(messages []domain.Message, filter domain.MessageFilter) d if query != "" && !strings.Contains(strings.ToLower(msg.Body), query) { continue } + if filter.MinDate > 0 && msg.Date <= filter.MinDate { + continue + } + if filter.MaxDate > 0 && msg.Date >= filter.MaxDate { + continue + } if filter.MaxID > 0 && msg.ID >= filter.MaxID { continue } @@ -297,6 +304,9 @@ func filterMessageList(messages []domain.Message, filter domain.MessageFilter) d if filter.SavedPeer.ID != 0 && msg.SavedPeer != filter.SavedPeer { continue } + if len(filter.SavedReactions) > 0 && !messageHasAnySavedTag(msg, filter.SavedReactions) { + continue + } base = append(base, msg) } @@ -316,6 +326,22 @@ func filterMessageList(messages []domain.Message, filter domain.MessageFilter) d } } +func messageHasAnySavedTag(msg domain.Message, wanted []domain.MessageReaction) bool { + if msg.Reactions == nil || !msg.Reactions.AsTags { + return false + } + have := make(map[string]struct{}, len(msg.Reactions.Results)) + for _, result := range msg.Reactions.Results { + have[result.Reaction.Key()] = struct{}{} + } + for _, reaction := range wanted { + if _, ok := have[reaction.Key()]; ok { + return true + } + } + return false +} + func pageMessageHistory(base []domain.Message, filter domain.MessageFilter, limit int) []domain.Message { if limit <= 0 || len(base) == 0 { return nil diff --git a/internal/store/memory/message_reactions.go b/internal/store/memory/message_reactions.go index 98528780..f2328df7 100644 --- a/internal/store/memory/message_reactions.go +++ b/internal/store/memory/message_reactions.go @@ -22,6 +22,9 @@ func (s *MessageStore) SetMessageReactions(_ context.Context, req domain.SetPriv } s.mu.Lock() defer s.mu.Unlock() + if req.Peer.ID == req.UserID { + return s.setSavedMessageTagsLocked(req) + } var target domain.Message for _, msg := range s.m[req.UserID] { if msg.ID == req.MessageID && msg.Peer == req.Peer { @@ -119,6 +122,10 @@ func (s *MessageStore) privateReactionResultLocked(uid int64) domain.PrivateMess } func (s *MessageStore) privateMessageReactionsForMessageLocked(msg domain.Message) domain.ChannelMessageReactions { + if msg.OwnerUserID != 0 && + msg.Peer == (domain.Peer{Type: domain.PeerTypeUser, ID: msg.OwnerUserID}) { + return s.savedMessageTagsForMessageLocked(msg) + } reactions := s.privateMessageReactionsLocked(msg.UID, msg.OwnerUserID) if len(reactions.Recent) == 0 || msg.From.ID == 0 { return reactions @@ -232,6 +239,11 @@ func writeMessageReactionsHash(h hash.Hash64, reactions *domain.ChannelMessageRe return } var buf [16]byte + if reactions.AsTags { + _, _ = h.Write([]byte{1}) + } else { + _, _ = h.Write([]byte{0}) + } for _, item := range reactions.Results { _, _ = h.Write([]byte(item.Reaction.Type)) _, _ = h.Write([]byte{0}) diff --git a/internal/store/memory/message_saved_reactions.go b/internal/store/memory/message_saved_reactions.go new file mode 100644 index 00000000..98743362 --- /dev/null +++ b/internal/store/memory/message_saved_reactions.go @@ -0,0 +1,159 @@ +package memory + +import ( + "context" + "sort" + + "telesrv/internal/domain" +) + +func (s *MessageStore) setSavedMessageTagsLocked(req domain.SetPrivateMessageReactionsRequest) (domain.PrivateMessageReactionsResult, error) { + var target domain.Message + for _, msg := range s.m[req.UserID] { + if msg.ID == req.MessageID && + msg.Peer == (domain.Peer{Type: domain.PeerTypeUser, ID: req.UserID}) { + target = msg + break + } + } + if target.ID == 0 { + return domain.PrivateMessageReactionsResult{}, domain.ErrMessageIDInvalid + } + for _, reaction := range req.Reactions { + if !reaction.Valid() { + return domain.PrivateMessageReactionsResult{}, domain.ErrReactionInvalid + } + } + if len(req.Reactions) == 0 { + if byMessage := s.savedMessageTags[req.UserID]; byMessage != nil { + delete(byMessage, target.ID) + if len(byMessage) == 0 { + delete(s.savedMessageTags, req.UserID) + } + } + } else { + if s.savedMessageTags[req.UserID] == nil { + s.savedMessageTags[req.UserID] = make(map[int][]domain.MessageReaction) + } + s.savedMessageTags[req.UserID][target.ID] = append([]domain.MessageReaction(nil), req.Reactions...) + } + item := cloneMessage(target) + reactions := s.savedMessageTagsForMessageLocked(item) + item.Reactions = cloneChannelMessageReactionsPtr(&reactions) + return domain.PrivateMessageReactionsResult{ + Messages: []domain.Message{item}, + Reactions: reactions, + }, nil +} + +func (s *MessageStore) savedMessageTagsForMessageLocked(msg domain.Message) domain.ChannelMessageReactions { + out := domain.ChannelMessageReactions{ + AsTags: true, + Results: []domain.ChannelMessageReactionCount{}, + Recent: []domain.ChannelMessagePeerReaction{}, + } + for i, reaction := range s.savedMessageTags[msg.OwnerUserID][msg.ID] { + out.Results = append(out.Results, domain.ChannelMessageReactionCount{ + Reaction: reaction, + Count: 1, + ChosenOrder: i + 1, + }) + } + return out +} + +func (s *MessageStore) ListSavedReactionTags(_ context.Context, req domain.SavedReactionTagsRequest) ([]domain.SavedReactionTag, error) { + if req.UserID == 0 { + return nil, domain.ErrReactionInvalid + } + if req.Limit <= 0 || req.Limit > domain.MaxSavedReactionTags { + req.Limit = domain.MaxSavedReactionTags + } + s.mu.RLock() + defer s.mu.RUnlock() + + visible := make(map[int]domain.Message, len(s.m[req.UserID])) + for _, msg := range s.m[req.UserID] { + if msg.Peer == (domain.Peer{Type: domain.PeerTypeUser, ID: req.UserID}) { + visible[msg.ID] = msg + } + } + byKey := make(map[string]domain.SavedReactionTag) + for messageID, reactions := range s.savedMessageTags[req.UserID] { + msg, ok := visible[messageID] + if !ok || (req.SavedPeer.ID != 0 && msg.SavedPeer != req.SavedPeer) { + continue + } + for _, reaction := range reactions { + key := reaction.Key() + tag := byKey[key] + tag.UserID = req.UserID + tag.Reaction = reaction + tag.Count++ + if req.SavedPeer.ID == 0 { + tag.Title = s.savedTagTitles[req.UserID][key] + } + byKey[key] = tag + } + } + out := make([]domain.SavedReactionTag, 0, len(byKey)) + for _, tag := range byKey { + out = append(out, tag) + } + sort.Slice(out, func(i, j int) bool { + if out[i].Count != out[j].Count { + return out[i].Count > out[j].Count + } + return out[i].Reaction.Key() > out[j].Reaction.Key() + }) + if len(out) > req.Limit { + out = out[:req.Limit] + } + return out, nil +} + +func (s *MessageStore) UpsertSavedReactionTag(_ context.Context, tag domain.SavedReactionTag) error { + if tag.UserID == 0 || !tag.Reaction.Valid() { + return domain.ErrReactionInvalid + } + key := tag.Reaction.Key() + s.mu.Lock() + defer s.mu.Unlock() + found := false + for messageID, reactions := range s.savedMessageTags[tag.UserID] { + alive := false + for _, msg := range s.m[tag.UserID] { + if msg.ID == messageID && + msg.Peer == (domain.Peer{Type: domain.PeerTypeUser, ID: tag.UserID}) { + alive = true + break + } + } + if !alive { + continue + } + for _, reaction := range reactions { + if reaction.Key() == key { + found = true + break + } + } + if found { + break + } + } + if !found { + return domain.ErrReactionInvalid + } + if tag.Title == "" { + if titles := s.savedTagTitles[tag.UserID]; titles != nil { + delete(titles, key) + } + return nil + } + if s.savedTagTitles[tag.UserID] == nil { + s.savedTagTitles[tag.UserID] = make(map[string]string) + } + s.savedTagTitles[tag.UserID][key] = tag.Title + return nil +} diff --git a/internal/store/memory/message_saved_reactions_test.go b/internal/store/memory/message_saved_reactions_test.go new file mode 100644 index 00000000..c1f9ad37 --- /dev/null +++ b/internal/store/memory/message_saved_reactions_test.go @@ -0,0 +1,150 @@ +package memory + +import ( + "context" + "testing" + + "telesrv/internal/domain" +) + +func TestSavedMessageTagsAssignmentCountsSearchAndDelete(t *testing.T) { + ctx := context.Background() + const userID int64 = 1001 + self := domain.Peer{Type: domain.PeerTypeUser, ID: userID} + peerA := domain.Peer{Type: domain.PeerTypeUser, ID: 2001} + peerB := domain.Peer{Type: domain.PeerTypeChannel, ID: 3001} + thumb := domain.MessageReaction{Type: domain.MessageReactionEmoji, Emoticon: "👍"} + custom := domain.MessageReaction{Type: domain.MessageReactionCustomEmoji, DocumentID: 90001} + + store := NewMessageStore() + create := func(body string, savedPeer domain.Peer) domain.Message { + msg, err := store.Create(ctx, domain.Message{ + OwnerUserID: userID, + Peer: self, + From: self, + SavedPeer: savedPeer, + Date: 1_700_000_000, + Body: body, + }) + if err != nil { + t.Fatalf("create saved message: %v", err) + } + return msg + } + first := create("first", peerA) + second := create("second", peerA) + third := create("third", peerB) + + set := func(msg domain.Message, reactions ...domain.MessageReaction) { + t.Helper() + result, err := store.SetMessageReactions(ctx, domain.SetPrivateMessageReactionsRequest{ + UserID: userID, + Peer: self, + MessageID: msg.ID, + Reactions: reactions, + ReactionsPerUserMax: 3, + }) + if err != nil { + t.Fatalf("set saved tags for %d: %v", msg.ID, err) + } + if len(result.Messages) != 1 || result.Messages[0].Reactions == nil || + !result.Messages[0].Reactions.AsTags { + t.Fatalf("saved tag result = %+v, want one reactions_as_tags message", result) + } + } + set(first, thumb) + set(second, thumb, custom) + set(third, custom) + if got := store.nextPts[userID]; got != 0 { + t.Fatalf("tag mutations pts = %d, want 0", got) + } + + if err := store.UpsertSavedReactionTag(ctx, domain.SavedReactionTag{ + UserID: userID, Reaction: thumb, Title: "Fav", + }); err != nil { + t.Fatalf("rename saved tag: %v", err) + } + global, err := store.ListSavedReactionTags(ctx, domain.SavedReactionTagsRequest{UserID: userID, Limit: 100}) + if err != nil { + t.Fatalf("list global saved tags: %v", err) + } + assertMemorySavedTag(t, global, thumb, 2, "Fav") + assertMemorySavedTag(t, global, custom, 2, "") + + perPeer, err := store.ListSavedReactionTags(ctx, domain.SavedReactionTagsRequest{ + UserID: userID, SavedPeer: peerA, Limit: 100, + }) + if err != nil { + t.Fatalf("list per-peer saved tags: %v", err) + } + assertMemorySavedTag(t, perPeer, thumb, 2, "") + assertMemorySavedTag(t, perPeer, custom, 1, "") + + found, err := store.ListByUser(ctx, userID, domain.MessageFilter{ + HasPeer: true, + Peer: self, + SavedPeer: peerA, + SavedReactions: []domain.MessageReaction{custom}, + Limit: 10, + }) + if err != nil { + t.Fatalf("search saved tag: %v", err) + } + if len(found.Messages) != 1 || found.Messages[0].ID != second.ID || + found.Messages[0].Reactions == nil || !found.Messages[0].Reactions.AsTags { + t.Fatalf("saved tag search = %+v, want second message", found.Messages) + } + foundAny, err := store.ListByUser(ctx, userID, domain.MessageFilter{ + HasPeer: true, + Peer: self, + SavedReactions: []domain.MessageReaction{thumb, custom}, + Limit: 10, + }) + if err != nil { + t.Fatalf("search any saved tag: %v", err) + } + if len(foundAny.Messages) != 3 { + t.Fatalf("saved tag OR search = %+v, want all three messages", foundAny.Messages) + } + + if _, err := store.DeleteMessages(ctx, domain.DeleteMessagesRequest{ + OwnerUserID: userID, + IDs: []int{second.ID}, + Date: 1_700_000_100, + }); err != nil { + t.Fatalf("delete tagged message: %v", err) + } + global, err = store.ListSavedReactionTags(ctx, domain.SavedReactionTagsRequest{UserID: userID, Limit: 100}) + if err != nil { + t.Fatalf("list tags after delete: %v", err) + } + assertMemorySavedTag(t, global, thumb, 1, "Fav") + assertMemorySavedTag(t, global, custom, 1, "") + + set(first) + global, err = store.ListSavedReactionTags(ctx, domain.SavedReactionTagsRequest{UserID: userID, Limit: 100}) + if err != nil { + t.Fatalf("list tags after clear: %v", err) + } + if len(global) != 1 || global[0].Reaction.Key() != custom.Key() { + t.Fatalf("tags after clear = %+v, want only custom", global) + } + if err := store.UpsertSavedReactionTag(ctx, domain.SavedReactionTag{ + UserID: userID, Reaction: thumb, Title: "ghost", + }); err != domain.ErrReactionInvalid { + t.Fatalf("rename unassigned tag err = %v, want ErrReactionInvalid", err) + } +} + +func assertMemorySavedTag(t *testing.T, tags []domain.SavedReactionTag, reaction domain.MessageReaction, count int, title string) { + t.Helper() + for _, tag := range tags { + if tag.Reaction.Key() == reaction.Key() { + if tag.Count != count || tag.Title != title { + t.Fatalf("tag %s = %+v, want count=%d title=%q", reaction.Key(), tag, count, title) + } + return + } + } + t.Fatalf("tag %s not found in %+v", reaction.Key(), tags) +} diff --git a/internal/store/memory/message_store.go b/internal/store/memory/message_store.go index 152bc2a7..2d3677bf 100644 --- a/internal/store/memory/message_store.go +++ b/internal/store/memory/message_store.go @@ -14,6 +14,8 @@ type MessageStore struct { nextPts map[int64]int readOutboxDates map[readOutboxDateKey]int privateReactions map[int64]map[int64][]domain.ChannelMessagePeerReaction + savedMessageTags map[int64]map[int][]domain.MessageReaction + savedTagTitles map[int64]map[string]string privateSendDedup map[privateSendDedupKey]privateSendDedupRecord loginCodeDeliveries map[[32]byte]loginCodeDeliveryRecord albumGroups map[albumGroupKey]albumGroupRecord @@ -44,6 +46,8 @@ func NewMessageStore(dialogs ...*DialogStore) *MessageStore { 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), diff --git a/internal/store/message.go b/internal/store/message.go index 7709c15b..945825a4 100644 --- a/internal/store/message.go +++ b/internal/store/message.go @@ -16,6 +16,8 @@ type MessageStore interface { GetOutboxReadDate(ctx context.Context, req domain.OutboxReadDateRequest) (int, error) SetMessageReactions(ctx context.Context, req domain.SetPrivateMessageReactionsRequest) (domain.PrivateMessageReactionsResult, error) GetMessageReactions(ctx context.Context, req domain.PrivateMessageReactionsRequest) (domain.PrivateMessageReactionsResult, error) + ListSavedReactionTags(ctx context.Context, req domain.SavedReactionTagsRequest) ([]domain.SavedReactionTag, error) + UpsertSavedReactionTag(ctx context.Context, tag domain.SavedReactionTag) error VoteMessagePoll(ctx context.Context, req domain.VotePrivateMessagePollRequest) (domain.PrivateMessagePollResult, error) CloseMessagePoll(ctx context.Context, req domain.ClosePrivateMessagePollRequest) (domain.PrivateMessagePollResult, error) EditMessage(ctx context.Context, req domain.EditMessageRequest) (domain.EditMessageResult, error) diff --git a/internal/store/postgres/channel_reaction_catalog.go b/internal/store/postgres/channel_reaction_catalog.go index 92f33d04..365f6e98 100644 --- a/internal/store/postgres/channel_reaction_catalog.go +++ b/internal/store/postgres/channel_reaction_catalog.go @@ -121,70 +121,3 @@ func (s *ChannelStore) ClearRecentMessageReactions(ctx context.Context, userID i } return nil } - -func (s *ChannelStore) ListSavedReactionTags(ctx context.Context, userID int64, limit int) ([]domain.SavedReactionTag, error) { - if userID == 0 { - return nil, domain.ErrChannelInvalid - } - if limit <= 0 { - return []domain.SavedReactionTag{}, nil - } - if limit > domain.MaxSavedReactionTags { - limit = domain.MaxSavedReactionTags - } - rows, err := s.db.Query(ctx, ` -SELECT reaction_type, reaction_value, title, reaction_count -FROM user_saved_reaction_tags -WHERE user_id = $1 -ORDER BY reaction_count DESC, updated_at DESC, reaction_type ASC, reaction_value ASC -LIMIT $2`, userID, limit) - if err != nil { - return nil, fmt.Errorf("list saved reaction tags: %w", err) - } - defer rows.Close() - out := make([]domain.SavedReactionTag, 0, limit) - for rows.Next() { - var reactionType, reactionValue, title string - var count int - if err := rows.Scan(&reactionType, &reactionValue, &title, &count); err != nil { - return nil, err - } - reaction, ok := domain.MessageReactionFromValue(domain.MessageReactionType(reactionType), reactionValue) - if !ok { - continue - } - out = append(out, domain.SavedReactionTag{ - UserID: userID, - Reaction: reaction, - Title: title, - Count: count, - }) - } - if err := rows.Err(); err != nil { - return nil, err - } - return out, nil -} - -func (s *ChannelStore) UpsertSavedReactionTag(ctx context.Context, tag domain.SavedReactionTag) error { - if tag.UserID == 0 || tag.Reaction.Type != domain.MessageReactionEmoji { - return domain.ErrChannelInvalid - } - reactionValue := strings.TrimSpace(tag.Reaction.Emoticon) - if reactionValue == "" { - return domain.ErrChannelInvalid - } - if tag.Count < 0 { - tag.Count = 0 - } - if _, err := s.db.Exec(ctx, ` -INSERT INTO user_saved_reaction_tags (user_id, reaction_type, reaction_value, title, reaction_count) -VALUES ($1, $2, $3, $4, $5) -ON CONFLICT (user_id, reaction_type, reaction_value) DO UPDATE SET - title = EXCLUDED.title, - reaction_count = GREATEST(user_saved_reaction_tags.reaction_count, EXCLUDED.reaction_count), - updated_at = now()`, tag.UserID, string(tag.Reaction.Type), reactionValue, tag.Title, tag.Count); err != nil { - return fmt.Errorf("upsert saved reaction tag: %w", err) - } - return nil -} diff --git a/internal/store/postgres/message_history.go b/internal/store/postgres/message_history.go index e9a8151c..92372430 100644 --- a/internal/store/postgres/message_history.go +++ b/internal/store/postgres/message_history.go @@ -4,10 +4,12 @@ import ( "context" "errors" "fmt" + "time" + "github.com/jackc/pgx/v5" + "telesrv/internal/domain" "telesrv/internal/store/postgres/sqlcgen" - "time" ) func (s *MessageStore) GetByIDs(ctx context.Context, userID int64, ids []int) (domain.MessageList, error) { @@ -92,6 +94,7 @@ func (s *MessageStore) ListByUser(ctx context.Context, userID int64, filter doma savedPeerType = string(filter.SavedPeer.Type) savedPeerID = filter.SavedPeer.ID } + savedReactionKeys := postgresSavedReactionKeys(filter.SavedReactions) // add_offset>=0 是 backward 热路径(初始加载/上滑翻页,占 getHistory 绝大多数)。 // 走扁平静态查询 ListMessagesBackward:规划仅单 index scan + 2 LEFT JOIN,避免 // ListMessagesByUser 大 CTE 把 4 个分支+total 全树规划(6.7ms→~1ms)。与 CTE @@ -100,23 +103,26 @@ func (s *MessageStore) ListByUser(ctx context.Context, userID int64, filter doma var rows []sqlcgen.ListMessagesByUserRow if addOffset >= 0 { bw, err := s.q.ListMessagesBackward(ctx, sqlcgen.ListMessagesBackwardParams{ - OwnerUserID: userID, - HasPeer: filter.HasPeer, - PeerType: string(filter.Peer.Type), - PeerID: filter.Peer.ID, - RestrictPeerIds: filter.RestrictPeerIDs, - PeerIds: filter.PeerIDs, - Query: filter.Query, - MaxID: pgInt32NonNegative(filter.MaxID), - MinID: pgInt32NonNegative(filter.MinID), - PinnedOnly: filter.PinnedOnly, - MusicOnly: filter.MusicOnly, - SavedPeerType: savedPeerType, - SavedPeerID: savedPeerID, - OffsetDate: pgInt32NonNegative(filter.OffsetDate), - OffsetID: pgInt32NonNegative(filter.OffsetID), - RowOffset: pgInt32Bounded(addOffset), - LimitCount: int32(queryLimit), + OwnerUserID: userID, + HasPeer: filter.HasPeer, + PeerType: string(filter.Peer.Type), + PeerID: filter.Peer.ID, + RestrictPeerIds: filter.RestrictPeerIDs, + PeerIds: filter.PeerIDs, + Query: filter.Query, + MinDate: pgInt32NonNegative(filter.MinDate), + MaxDate: pgInt32NonNegative(filter.MaxDate), + MaxID: pgInt32NonNegative(filter.MaxID), + MinID: pgInt32NonNegative(filter.MinID), + PinnedOnly: filter.PinnedOnly, + MusicOnly: filter.MusicOnly, + SavedPeerType: savedPeerType, + SavedPeerID: savedPeerID, + SavedReactionKeys: savedReactionKeys, + OffsetDate: pgInt32NonNegative(filter.OffsetDate), + OffsetID: pgInt32NonNegative(filter.OffsetID), + RowOffset: pgInt32Bounded(addOffset), + LimitCount: int32(queryLimit), }) if err != nil { return domain.MessageList{}, fmt.Errorf("list messages (backward): %w", err) @@ -127,19 +133,22 @@ func (s *MessageStore) ListByUser(ctx context.Context, userID int64, filter doma } if filter.NeedTotalCount { total, err := s.q.CountMessagesByUser(ctx, sqlcgen.CountMessagesByUserParams{ - OwnerUserID: userID, - HasPeer: filter.HasPeer, - PeerType: string(filter.Peer.Type), - PeerID: filter.Peer.ID, - RestrictPeerIds: filter.RestrictPeerIDs, - PeerIds: filter.PeerIDs, - Query: filter.Query, - MaxID: pgInt32NonNegative(filter.MaxID), - MinID: pgInt32NonNegative(filter.MinID), - PinnedOnly: filter.PinnedOnly, - MusicOnly: filter.MusicOnly, - SavedPeerType: savedPeerType, - SavedPeerID: savedPeerID, + OwnerUserID: userID, + HasPeer: filter.HasPeer, + PeerType: string(filter.Peer.Type), + PeerID: filter.Peer.ID, + RestrictPeerIds: filter.RestrictPeerIDs, + PeerIds: filter.PeerIDs, + Query: filter.Query, + MinDate: pgInt32NonNegative(filter.MinDate), + MaxDate: pgInt32NonNegative(filter.MaxDate), + MaxID: pgInt32NonNegative(filter.MaxID), + MinID: pgInt32NonNegative(filter.MinID), + PinnedOnly: filter.PinnedOnly, + MusicOnly: filter.MusicOnly, + SavedPeerType: savedPeerType, + SavedPeerID: savedPeerID, + SavedReactionKeys: savedReactionKeys, }) if err != nil { return domain.MessageList{}, fmt.Errorf("count messages: %w", err) @@ -153,24 +162,27 @@ func (s *MessageStore) ListByUser(ctx context.Context, userID int64, filter doma } else { var err error rows, err = s.q.ListMessagesByUser(ctx, sqlcgen.ListMessagesByUserParams{ - OwnerUserID: userID, - HasPeer: filter.HasPeer, - PeerType: string(filter.Peer.Type), - PeerID: filter.Peer.ID, - RestrictPeerIds: filter.RestrictPeerIDs, - PeerIds: filter.PeerIDs, - Query: filter.Query, - OffsetID: pgInt32NonNegative(filter.OffsetID), - OffsetDate: pgInt32NonNegative(filter.OffsetDate), - MaxID: pgInt32NonNegative(filter.MaxID), - MinID: pgInt32NonNegative(filter.MinID), - AddOffset: pgInt32Bounded(addOffset), - LimitCount: int32(queryLimit), - PinnedOnly: filter.PinnedOnly, - MusicOnly: filter.MusicOnly, - NeedTotalCount: filter.NeedTotalCount, - SavedPeerType: savedPeerType, - SavedPeerID: savedPeerID, + OwnerUserID: userID, + HasPeer: filter.HasPeer, + PeerType: string(filter.Peer.Type), + PeerID: filter.Peer.ID, + RestrictPeerIds: filter.RestrictPeerIDs, + PeerIds: filter.PeerIDs, + Query: filter.Query, + MinDate: pgInt32NonNegative(filter.MinDate), + MaxDate: pgInt32NonNegative(filter.MaxDate), + OffsetID: pgInt32NonNegative(filter.OffsetID), + OffsetDate: pgInt32NonNegative(filter.OffsetDate), + MaxID: pgInt32NonNegative(filter.MaxID), + MinID: pgInt32NonNegative(filter.MinID), + AddOffset: pgInt32Bounded(addOffset), + LimitCount: int32(queryLimit), + PinnedOnly: filter.PinnedOnly, + MusicOnly: filter.MusicOnly, + NeedTotalCount: filter.NeedTotalCount, + SavedPeerType: savedPeerType, + SavedPeerID: savedPeerID, + SavedReactionKeys: savedReactionKeys, }) if err != nil { return domain.MessageList{}, fmt.Errorf("list messages: %w", err) @@ -273,6 +285,16 @@ func (s *MessageStore) ListByUser(ctx context.Context, userID int64, filter doma return out, nil } +func postgresSavedReactionKeys(reactions []domain.MessageReaction) []string { + out := make([]string, 0, len(reactions)) + for _, reaction := range reactions { + if reaction.Valid() { + out = append(out, string(reaction.Type)+":"+reaction.Value()) + } + } + return out +} + func (s *MessageStore) ReadHistory(ctx context.Context, req domain.ReadHistoryRequest) (res domain.ReadHistoryResult, err error) { res = domain.ReadHistoryResult{OwnerUserID: req.OwnerUserID, Peer: req.Peer, MaxID: req.MaxID} if req.OwnerUserID == 0 { diff --git a/internal/store/postgres/message_reactions.go b/internal/store/postgres/message_reactions.go index 47abf083..6d8c452e 100644 --- a/internal/store/postgres/message_reactions.go +++ b/internal/store/postgres/message_reactions.go @@ -29,6 +29,9 @@ func (s *MessageStore) SetMessageReactions(ctx context.Context, req domain.SetPr return domain.PrivateMessageReactionsResult{}, domain.ErrMessageIDInvalid } } + if req.Peer.ID == req.UserID { + return s.setSavedMessageTags(ctx, req) + } beginner, ok := s.db.(txBeginner) if !ok { return domain.PrivateMessageReactionsResult{}, fmt.Errorf("set message reactions: db does not support transactions") @@ -240,10 +243,17 @@ func (s *MessageStore) enrichPrivateMessageReactions(ctx context.Context, db sql if err := s.enrichPrivateMessagePolls(ctx, db, viewerUserID, messages); err != nil { return err } + if err := s.enrichSavedMessageTags(ctx, db, messages); err != nil { + return err + } keySet := make(map[privateMessageReactionKey]struct{}, len(messages)) senderIDs := make([]int64, 0, len(messages)) privateIDs := make([]int64, 0, len(messages)) for _, msg := range messages { + if msg.OwnerUserID != 0 && + msg.Peer == (domain.Peer{Type: domain.PeerTypeUser, ID: msg.OwnerUserID}) { + continue + } if msg.UID == 0 || msg.From.ID == 0 { continue } @@ -420,6 +430,11 @@ func writeMessageReactionsHash(h hash.Hash64, reactions *domain.ChannelMessageRe return } var buf [16]byte + if reactions.AsTags { + _, _ = h.Write([]byte{1}) + } else { + _, _ = h.Write([]byte{0}) + } for _, item := range reactions.Results { _, _ = h.Write([]byte(item.Reaction.Type)) _, _ = h.Write([]byte{0}) diff --git a/internal/store/postgres/message_saved_reactions.go b/internal/store/postgres/message_saved_reactions.go new file mode 100644 index 00000000..81585400 --- /dev/null +++ b/internal/store/postgres/message_saved_reactions.go @@ -0,0 +1,308 @@ +package postgres + +import ( + "context" + "errors" + "fmt" + "sort" + "unicode/utf8" + + "github.com/jackc/pgx/v5" + + "telesrv/internal/domain" + "telesrv/internal/store/postgres/sqlcgen" +) + +func (s *MessageStore) setSavedMessageTags(ctx context.Context, req domain.SetPrivateMessageReactionsRequest) (domain.PrivateMessageReactionsResult, error) { + beginner, ok := s.db.(txBeginner) + if !ok { + return domain.PrivateMessageReactionsResult{}, fmt.Errorf("set saved message tags: db does not support transactions") + } + tx, err := beginner.Begin(ctx) + if err != nil { + return domain.PrivateMessageReactionsResult{}, fmt.Errorf("begin set saved message tags tx: %w", err) + } + committed := false + defer func() { + if !committed { + _ = tx.Rollback(ctx) + } + }() + if err := lockUsersForUpdate(ctx, tx, req.UserID); err != nil { + return domain.PrivateMessageReactionsResult{}, fmt.Errorf("lock saved message tag owner: %w", err) + } + + var boxID int32 + if err := tx.QueryRow(ctx, ` +SELECT box_id +FROM message_boxes +WHERE owner_user_id = $1 + AND box_id = $2 + AND peer_type = 'user' + AND peer_id = $1 + AND NOT deleted +LIMIT 1 +FOR UPDATE`, req.UserID, int32(req.MessageID)).Scan(&boxID); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return domain.PrivateMessageReactionsResult{}, domain.ErrMessageIDInvalid + } + return domain.PrivateMessageReactionsResult{}, fmt.Errorf("get saved message for tags: %w", err) + } + if _, err := tx.Exec(ctx, ` +DELETE FROM saved_message_reaction_tags +WHERE user_id = $1 AND message_box_id = $2`, req.UserID, boxID); err != nil { + return domain.PrivateMessageReactionsResult{}, fmt.Errorf("delete old saved message tags: %w", err) + } + for i, reaction := range req.Reactions { + if !reaction.Valid() { + return domain.PrivateMessageReactionsResult{}, domain.ErrReactionInvalid + } + if _, err := tx.Exec(ctx, ` +INSERT INTO saved_message_reaction_tags ( + user_id, message_box_id, reaction_type, reaction_value, chosen_order +) VALUES ($1, $2, $3, $4, $5)`, + req.UserID, boxID, string(reaction.Type), reaction.Value(), int32(i+1)); err != nil { + return domain.PrivateMessageReactionsResult{}, fmt.Errorf("insert saved message tag: %w", err) + } + } + rows, err := sqlcgen.New(tx).GetMessageBoxesByIDs(ctx, sqlcgen.GetMessageBoxesByIDsParams{ + OwnerUserID: req.UserID, + BoxIds: []int32{boxID}, + }) + if err != nil { + return domain.PrivateMessageReactionsResult{}, fmt.Errorf("reload saved message tags box: %w", err) + } + if len(rows) != 1 { + return domain.PrivateMessageReactionsResult{}, domain.ErrMessageIDInvalid + } + msg, err := messageFromIDRow(rows[0]) + if err != nil { + return domain.PrivateMessageReactionsResult{}, err + } + messages := []domain.Message{msg} + if err := s.enrichPrivateMessageReactions(ctx, tx, req.UserID, messages); err != nil { + return domain.PrivateMessageReactionsResult{}, err + } + if err := tx.Commit(ctx); err != nil { + return domain.PrivateMessageReactionsResult{}, fmt.Errorf("commit saved message tags tx: %w", err) + } + committed = true + reactions := domain.ChannelMessageReactions{AsTags: true} + if messages[0].Reactions != nil { + reactions = *messages[0].Reactions + } + return domain.PrivateMessageReactionsResult{ + Messages: messages, + Reactions: reactions, + }, nil +} + +func (s *MessageStore) enrichSavedMessageTags(ctx context.Context, db sqlcgen.DBTX, messages []domain.Message) error { + ownerIDs := make([]int64, 0, len(messages)) + boxIDs := make([]int32, 0, len(messages)) + indexes := make(map[[2]int64]int, len(messages)) + for i := range messages { + msg := messages[i] + if msg.OwnerUserID == 0 || + msg.Peer != (domain.Peer{Type: domain.PeerTypeUser, ID: msg.OwnerUserID}) { + continue + } + ownerIDs = append(ownerIDs, msg.OwnerUserID) + boxIDs = append(boxIDs, int32(msg.ID)) + indexes[[2]int64{msg.OwnerUserID, int64(msg.ID)}] = i + } + if len(ownerIDs) == 0 { + return nil + } + rows, err := db.Query(ctx, ` +WITH wanted AS ( + SELECT user_id, message_box_id + FROM unnest($1::bigint[], $2::int[]) AS w(user_id, message_box_id) +) +SELECT t.user_id, t.message_box_id, t.reaction_type, t.reaction_value, t.chosen_order +FROM saved_message_reaction_tags t +JOIN wanted w + ON w.user_id = t.user_id + AND w.message_box_id = t.message_box_id +ORDER BY t.user_id, t.message_box_id, t.chosen_order, t.reaction_type, t.reaction_value`, + ownerIDs, boxIDs) + if err != nil { + return fmt.Errorf("load saved message tags: %w", err) + } + defer rows.Close() + for rows.Next() { + var ( + userID int64 + messageBoxID int32 + reactionType string + reactionValue string + chosenOrder int32 + ) + if err := rows.Scan(&userID, &messageBoxID, &reactionType, &reactionValue, &chosenOrder); err != nil { + return fmt.Errorf("scan saved message tag: %w", err) + } + reaction, ok := domain.MessageReactionFromValue(domain.MessageReactionType(reactionType), reactionValue) + if !ok { + continue + } + index, ok := indexes[[2]int64{userID, int64(messageBoxID)}] + if !ok { + continue + } + if messages[index].Reactions == nil { + messages[index].Reactions = &domain.ChannelMessageReactions{ + AsTags: true, + Results: []domain.ChannelMessageReactionCount{}, + Recent: []domain.ChannelMessagePeerReaction{}, + } + } + messages[index].Reactions.Results = append(messages[index].Reactions.Results, domain.ChannelMessageReactionCount{ + Reaction: reaction, + Count: 1, + ChosenOrder: int(chosenOrder), + }) + } + if err := rows.Err(); err != nil { + return fmt.Errorf("saved message tag rows: %w", err) + } + return nil +} + +func (s *MessageStore) ListSavedReactionTags(ctx context.Context, req domain.SavedReactionTagsRequest) ([]domain.SavedReactionTag, error) { + if req.UserID == 0 { + return nil, domain.ErrReactionInvalid + } + if req.Limit <= 0 || req.Limit > domain.MaxSavedReactionTags { + req.Limit = domain.MaxSavedReactionTags + } + savedPeerType := "" + var savedPeerID int64 + if req.SavedPeer.ID != 0 { + savedPeerType = string(req.SavedPeer.Type) + savedPeerID = req.SavedPeer.ID + } + rows, err := s.db.Query(ctx, ` +SELECT + a.reaction_type, + a.reaction_value, + CASE WHEN $2 = '' THEN COALESCE(t.title, '') ELSE '' END AS title, + COUNT(*)::int AS reaction_count +FROM saved_message_reaction_tags a +JOIN message_boxes m + ON m.owner_user_id = a.user_id + AND m.box_id = a.message_box_id + AND NOT m.deleted + AND m.peer_type = 'user' + AND m.peer_id = a.user_id +LEFT JOIN user_saved_reaction_tags t + ON t.user_id = a.user_id + AND t.reaction_type = a.reaction_type + AND t.reaction_value = a.reaction_value +WHERE a.user_id = $1 + AND ($2 = '' OR (m.saved_peer_type = $2 AND m.saved_peer_id = $3)) +GROUP BY a.reaction_type, a.reaction_value, title +ORDER BY + reaction_count DESC, + CASE + WHEN a.reaction_type = 'custom_emoji' + THEN lpad(to_hex(a.reaction_value::bigint), 16, '0') + ELSE substr(md5(replace(a.reaction_value, U&'\FE0F', '')), 1, 16) + END DESC +LIMIT $4`, req.UserID, savedPeerType, savedPeerID, int32(req.Limit)) + if err != nil { + return nil, fmt.Errorf("list saved reaction tags: %w", err) + } + defer rows.Close() + out := make([]domain.SavedReactionTag, 0, req.Limit) + for rows.Next() { + var reactionType, reactionValue, title string + var count int32 + if err := rows.Scan(&reactionType, &reactionValue, &title, &count); err != nil { + return nil, fmt.Errorf("scan saved reaction tag: %w", err) + } + reaction, ok := domain.MessageReactionFromValue(domain.MessageReactionType(reactionType), reactionValue) + if !ok || count <= 0 { + continue + } + out = append(out, domain.SavedReactionTag{ + UserID: req.UserID, + Reaction: reaction, + Title: title, + Count: int(count), + }) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("saved reaction tag rows: %w", err) + } + sort.SliceStable(out, func(i, j int) bool { + if out[i].Count != out[j].Count { + return out[i].Count > out[j].Count + } + return out[i].Reaction.Key() > out[j].Reaction.Key() + }) + return out, nil +} + +func (s *MessageStore) UpsertSavedReactionTag(ctx context.Context, tag domain.SavedReactionTag) error { + if tag.UserID == 0 || !tag.Reaction.Valid() || utf8.RuneCountInString(tag.Title) > 12 { + return domain.ErrReactionInvalid + } + beginner, ok := s.db.(txBeginner) + if !ok { + return fmt.Errorf("update saved reaction tag title: db does not support transactions") + } + tx, err := beginner.Begin(ctx) + if err != nil { + return fmt.Errorf("begin update saved reaction tag title tx: %w", err) + } + committed := false + defer func() { + if !committed { + _ = tx.Rollback(ctx) + } + }() + if err := lockUsersForUpdate(ctx, tx, tag.UserID); err != nil { + return fmt.Errorf("lock saved reaction tag owner: %w", err) + } + var exists bool + if err := tx.QueryRow(ctx, ` +SELECT EXISTS ( + SELECT 1 + FROM saved_message_reaction_tags a + JOIN message_boxes m + ON m.owner_user_id = a.user_id + AND m.box_id = a.message_box_id + AND NOT m.deleted + AND m.peer_type = 'user' + AND m.peer_id = a.user_id + WHERE a.user_id = $1 + AND a.reaction_type = $2 + AND a.reaction_value = $3 +)`, tag.UserID, string(tag.Reaction.Type), tag.Reaction.Value()).Scan(&exists); err != nil { + return fmt.Errorf("check saved reaction tag assignment: %w", err) + } + if !exists { + return domain.ErrReactionInvalid + } + if tag.Title == "" { + if _, err := tx.Exec(ctx, ` +DELETE FROM user_saved_reaction_tags +WHERE user_id = $1 AND reaction_type = $2 AND reaction_value = $3`, + tag.UserID, string(tag.Reaction.Type), tag.Reaction.Value()); err != nil { + return fmt.Errorf("delete saved reaction tag title: %w", err) + } + } else if _, err := tx.Exec(ctx, ` +INSERT INTO user_saved_reaction_tags ( + user_id, reaction_type, reaction_value, title, reaction_count +) VALUES ($1, $2, $3, $4, 0) +ON CONFLICT (user_id, reaction_type, reaction_value) +DO UPDATE SET title = EXCLUDED.title, reaction_count = 0, updated_at = now()`, + tag.UserID, string(tag.Reaction.Type), tag.Reaction.Value(), tag.Title); err != nil { + return fmt.Errorf("upsert saved reaction tag title: %w", err) + } + if err := tx.Commit(ctx); err != nil { + return fmt.Errorf("commit saved reaction tag title tx: %w", err) + } + committed = true + return nil +} diff --git a/internal/store/postgres/message_saved_reactions_integration_test.go b/internal/store/postgres/message_saved_reactions_integration_test.go new file mode 100644 index 00000000..68a198b4 --- /dev/null +++ b/internal/store/postgres/message_saved_reactions_integration_test.go @@ -0,0 +1,177 @@ +package postgres + +import ( + "context" + "testing" + "time" + + "telesrv/internal/domain" +) + +func TestSavedMessageTagsPostgresAssignmentCountsSearchAndDelete(t *testing.T) { + pool := testPool(t) + ctx := context.Background() + suffix := randomSuffix(t) + users := NewUserStore(pool) + user, err := users.Create(ctx, domain.User{ + AccessHash: 1, + Phone: "+1777" + suffix + "01", + FirstName: "SavedTags", + }) + if err != nil { + t.Fatalf("create saved-tag user: %v", err) + } + t.Cleanup(func() { + _, _ = pool.Exec(ctx, "DELETE FROM saved_message_reaction_tags WHERE user_id = $1", user.ID) + _, _ = pool.Exec(ctx, "DELETE FROM user_saved_reaction_tags WHERE user_id = $1", user.ID) + _, _ = pool.Exec(ctx, "DELETE FROM message_boxes WHERE owner_user_id = $1", user.ID) + _, _ = pool.Exec(ctx, "DELETE FROM private_messages WHERE sender_user_id = $1 OR recipient_user_id = $1", user.ID) + _, _ = pool.Exec(ctx, "DELETE FROM user_update_events WHERE user_id = $1", user.ID) + _, _ = pool.Exec(ctx, "DELETE FROM dialogs WHERE user_id = $1", user.ID) + _, _ = pool.Exec(ctx, "DELETE FROM users WHERE id = $1", user.ID) + }) + + messages := NewMessageStore(pool) + self := domain.Peer{Type: domain.PeerTypeUser, ID: user.ID} + peerA := domain.Peer{Type: domain.PeerTypeUser, ID: user.ID} + peerB := domain.Peer{Type: domain.PeerTypeChannel, ID: 90001} + create := func(body string, savedPeer domain.Peer) domain.Message { + msg, err := messages.Create(ctx, domain.Message{ + OwnerUserID: user.ID, + Peer: self, + From: self, + Date: int(time.Now().Unix()), + Body: body, + }) + if err != nil { + t.Fatalf("create saved message: %v", err) + } + if _, err := pool.Exec(ctx, ` +UPDATE message_boxes +SET saved_peer_type = $3, saved_peer_id = $4 +WHERE owner_user_id = $1 AND box_id = $2`, + user.ID, msg.ID, string(savedPeer.Type), savedPeer.ID); err != nil { + t.Fatalf("set saved peer: %v", err) + } + msg.SavedPeer = savedPeer + return msg + } + first := create("first", peerA) + second := create("second", peerA) + third := create("third", peerB) + thumb := domain.MessageReaction{Type: domain.MessageReactionEmoji, Emoticon: "👍"} + custom := domain.MessageReaction{Type: domain.MessageReactionCustomEmoji, DocumentID: 70001} + + set := func(msg domain.Message, reactions ...domain.MessageReaction) { + t.Helper() + result, err := messages.SetMessageReactions(ctx, domain.SetPrivateMessageReactionsRequest{ + UserID: user.ID, + Peer: self, + MessageID: msg.ID, + Reactions: reactions, + ReactionsPerUserMax: 3, + }) + if err != nil { + t.Fatalf("set saved tags on %d: %v", msg.ID, err) + } + if len(result.Messages) != 1 || result.Messages[0].Reactions == nil || + !result.Messages[0].Reactions.AsTags { + t.Fatalf("set saved tags result = %+v", result) + } + } + set(first, thumb) + set(second, thumb, custom) + set(third, custom) + if err := messages.UpsertSavedReactionTag(ctx, domain.SavedReactionTag{ + UserID: user.ID, Reaction: custom, Title: "Custom", + }); err != nil { + t.Fatalf("rename custom saved tag: %v", err) + } + + global, err := messages.ListSavedReactionTags(ctx, domain.SavedReactionTagsRequest{ + UserID: user.ID, Limit: 100, + }) + if err != nil { + t.Fatalf("list global saved tags: %v", err) + } + assertPostgresSavedTag(t, global, thumb, 2, "") + assertPostgresSavedTag(t, global, custom, 2, "Custom") + + perPeer, err := messages.ListSavedReactionTags(ctx, domain.SavedReactionTagsRequest{ + UserID: user.ID, SavedPeer: peerA, Limit: 100, + }) + if err != nil { + t.Fatalf("list per-peer saved tags: %v", err) + } + assertPostgresSavedTag(t, perPeer, thumb, 2, "") + assertPostgresSavedTag(t, perPeer, custom, 1, "") + + search, err := messages.ListByUser(ctx, user.ID, domain.MessageFilter{ + HasPeer: true, + Peer: self, + SavedPeer: peerA, + SavedReactions: []domain.MessageReaction{custom}, + NeedTotalCount: true, + Limit: 10, + }) + if err != nil { + t.Fatalf("search saved tag: %v", err) + } + if len(search.Messages) != 1 || search.Messages[0].ID != second.ID || + search.Messages[0].Reactions == nil || !search.Messages[0].Reactions.AsTags { + t.Fatalf("saved tag search = %+v, want second message", search.Messages) + } + searchAny, err := messages.ListByUser(ctx, user.ID, domain.MessageFilter{ + HasPeer: true, + Peer: self, + SavedReactions: []domain.MessageReaction{thumb, custom}, + NeedTotalCount: true, + Limit: 10, + }) + if err != nil { + t.Fatalf("search any saved tag: %v", err) + } + if len(searchAny.Messages) != 3 || searchAny.Count != 3 { + t.Fatalf("saved tag OR search = count %d messages %+v, want all three", searchAny.Count, searchAny.Messages) + } + + var reactionEvents int + if err := pool.QueryRow(ctx, ` +SELECT COUNT(*)::int +FROM user_update_events +WHERE user_id = $1 AND event_type = 'message_reactions'`, user.ID).Scan(&reactionEvents); err != nil { + t.Fatalf("count reaction events: %v", err) + } + if reactionEvents != 0 { + t.Fatalf("reaction durable events = %d, want 0", reactionEvents) + } + + if _, err := messages.DeleteMessages(ctx, domain.DeleteMessagesRequest{ + OwnerUserID: user.ID, + IDs: []int{second.ID}, + Date: int(time.Now().Unix()), + }); err != nil { + t.Fatalf("delete tagged saved message: %v", err) + } + global, err = messages.ListSavedReactionTags(ctx, domain.SavedReactionTagsRequest{ + UserID: user.ID, Limit: 100, + }) + if err != nil { + t.Fatalf("list tags after delete: %v", err) + } + assertPostgresSavedTag(t, global, thumb, 1, "") + assertPostgresSavedTag(t, global, custom, 1, "Custom") +} + +func assertPostgresSavedTag(t *testing.T, tags []domain.SavedReactionTag, reaction domain.MessageReaction, count int, title string) { + t.Helper() + for _, tag := range tags { + if tag.Reaction.Key() == reaction.Key() { + if tag.Count != count || tag.Title != title { + t.Fatalf("tag %s = %+v, want count=%d title=%q", reaction.Key(), tag, count, title) + } + return + } + } + t.Fatalf("tag %s not found in %+v", reaction.Key(), tags) +} diff --git a/internal/store/postgres/queries/message.sql b/internal/store/postgres/queries/message.sql index 65c30d09..920d1af6 100644 --- a/internal/store/postgres/queries/message.sql +++ b/internal/store/postgres/queries/message.sql @@ -545,6 +545,8 @@ base AS NOT MATERIALIZED ( sqlc.arg(query)::text = '' OR m.body ILIKE ('%' || sqlc.arg(query)::text || '%') ) + AND (sqlc.arg(min_date)::int <= 0 OR m.message_date > sqlc.arg(min_date)::int) + AND (sqlc.arg(max_date)::int <= 0 OR m.message_date < sqlc.arg(max_date)::int) AND (sqlc.arg(max_id)::int <= 0 OR m.box_id < sqlc.arg(max_id)::int) AND (sqlc.arg(min_id)::int <= 0 OR m.box_id > sqlc.arg(min_id)::int) AND (NOT sqlc.arg(pinned_only)::boolean OR m.pinned) @@ -564,6 +566,17 @@ base AS NOT MATERIALIZED ( sqlc.arg(saved_peer_type)::text = '' OR (m.saved_peer_type = sqlc.arg(saved_peer_type)::text AND m.saved_peer_id = sqlc.arg(saved_peer_id)::bigint) ) + AND ( + cardinality(sqlc.arg(saved_reaction_keys)::text[]) = 0 + OR EXISTS ( + SELECT 1 + FROM saved_message_reaction_tags tag + WHERE tag.user_id = m.owner_user_id + AND tag.message_box_id = m.box_id + AND (tag.reaction_type || ':' || tag.reaction_value) + = ANY(sqlc.arg(saved_reaction_keys)::text[]) + ) + ) ), total AS ( SELECT count(*)::int AS total_count @@ -810,6 +823,8 @@ WHERE m.owner_user_id = sqlc.arg(owner_user_id)::bigint sqlc.arg(query)::text = '' OR m.body ILIKE ('%' || sqlc.arg(query)::text || '%') ) + AND (sqlc.arg(min_date)::int <= 0 OR m.message_date > sqlc.arg(min_date)::int) + AND (sqlc.arg(max_date)::int <= 0 OR m.message_date < sqlc.arg(max_date)::int) AND (sqlc.arg(max_id)::int <= 0 OR m.box_id < sqlc.arg(max_id)::int) AND (sqlc.arg(min_id)::int <= 0 OR m.box_id > sqlc.arg(min_id)::int) AND (NOT sqlc.arg(pinned_only)::boolean OR m.pinned) @@ -829,6 +844,17 @@ WHERE m.owner_user_id = sqlc.arg(owner_user_id)::bigint sqlc.arg(saved_peer_type)::text = '' OR (m.saved_peer_type = sqlc.arg(saved_peer_type)::text AND m.saved_peer_id = sqlc.arg(saved_peer_id)::bigint) ) + AND ( + cardinality(sqlc.arg(saved_reaction_keys)::text[]) = 0 + OR EXISTS ( + SELECT 1 + FROM saved_message_reaction_tags tag + WHERE tag.user_id = m.owner_user_id + AND tag.message_box_id = m.box_id + AND (tag.reaction_type || ':' || tag.reaction_value) + = ANY(sqlc.arg(saved_reaction_keys)::text[]) + ) + ) AND ( (sqlc.arg(offset_date)::int > 0 AND m.message_date < sqlc.arg(offset_date)::int) OR (sqlc.arg(offset_date)::int <= 0 AND (sqlc.arg(offset_id)::int <= 0 OR m.box_id < sqlc.arg(offset_id)::int)) @@ -856,6 +882,8 @@ WHERE m.owner_user_id = sqlc.arg(owner_user_id)::bigint sqlc.arg(query)::text = '' OR m.body ILIKE ('%' || sqlc.arg(query)::text || '%') ) + AND (sqlc.arg(min_date)::int <= 0 OR m.message_date > sqlc.arg(min_date)::int) + AND (sqlc.arg(max_date)::int <= 0 OR m.message_date < sqlc.arg(max_date)::int) AND (sqlc.arg(max_id)::int <= 0 OR m.box_id < sqlc.arg(max_id)::int) AND (sqlc.arg(min_id)::int <= 0 OR m.box_id > sqlc.arg(min_id)::int) AND (NOT sqlc.arg(pinned_only)::boolean OR m.pinned) @@ -874,6 +902,17 @@ WHERE m.owner_user_id = sqlc.arg(owner_user_id)::bigint AND ( sqlc.arg(saved_peer_type)::text = '' OR (m.saved_peer_type = sqlc.arg(saved_peer_type)::text AND m.saved_peer_id = sqlc.arg(saved_peer_id)::bigint) + ) + AND ( + cardinality(sqlc.arg(saved_reaction_keys)::text[]) = 0 + OR EXISTS ( + SELECT 1 + FROM saved_message_reaction_tags tag + WHERE tag.user_id = m.owner_user_id + AND tag.message_box_id = m.box_id + AND (tag.reaction_type || ':' || tag.reaction_value) + = ANY(sqlc.arg(saved_reaction_keys)::text[]) + ) ); -- name: GetMessageBoxesByIDs :many diff --git a/internal/store/postgres/sqlcgen/message.sql.go b/internal/store/postgres/sqlcgen/message.sql.go index ed0eb1c9..b86bfec9 100644 --- a/internal/store/postgres/sqlcgen/message.sql.go +++ b/internal/store/postgres/sqlcgen/message.sql.go @@ -26,11 +26,13 @@ WHERE m.owner_user_id = $1::bigint $7::text = '' OR m.body ILIKE ('%' || $7::text || '%') ) - AND ($8::int <= 0 OR m.box_id < $8::int) - AND ($9::int <= 0 OR m.box_id > $9::int) - AND (NOT $10::boolean OR m.pinned) + AND ($8::int <= 0 OR m.message_date > $8::int) + AND ($9::int <= 0 OR m.message_date < $9::int) + AND ($10::int <= 0 OR m.box_id < $10::int) + AND ($11::int <= 0 OR m.box_id > $11::int) + AND (NOT $12::boolean OR m.pinned) AND ( - NOT $11::boolean + NOT $13::boolean OR ( m.media->>'kind' = 'document' AND EXISTS ( @@ -42,25 +44,39 @@ WHERE m.owner_user_id = $1::bigint ) ) AND ( - $12::text = '' - OR (m.saved_peer_type = $12::text AND m.saved_peer_id = $13::bigint) + $14::text = '' + OR (m.saved_peer_type = $14::text AND m.saved_peer_id = $15::bigint) + ) + AND ( + cardinality($16::text[]) = 0 + OR EXISTS ( + SELECT 1 + FROM saved_message_reaction_tags tag + WHERE tag.user_id = m.owner_user_id + AND tag.message_box_id = m.box_id + AND (tag.reaction_type || ':' || tag.reaction_value) + = ANY($16::text[]) + ) ) ` type CountMessagesByUserParams struct { - OwnerUserID int64 - HasPeer bool - PeerType string - PeerID int64 - RestrictPeerIds bool - PeerIds []int64 - Query string - MaxID int32 - MinID int32 - PinnedOnly bool - MusicOnly bool - SavedPeerType string - SavedPeerID int64 + OwnerUserID int64 + HasPeer bool + PeerType string + PeerID int64 + RestrictPeerIds bool + PeerIds []int64 + Query string + MinDate int32 + MaxDate int32 + MaxID int32 + MinID int32 + PinnedOnly bool + MusicOnly bool + SavedPeerType string + SavedPeerID int64 + SavedReactionKeys []string } // ListMessagesByUser total CTE 的独立化:相同 base 过滤(不含分页 anchor), @@ -74,12 +90,15 @@ func (q *Queries) CountMessagesByUser(ctx context.Context, arg CountMessagesByUs arg.RestrictPeerIds, arg.PeerIds, arg.Query, + arg.MinDate, + arg.MaxDate, arg.MaxID, arg.MinID, arg.PinnedOnly, arg.MusicOnly, arg.SavedPeerType, arg.SavedPeerID, + arg.SavedReactionKeys, ) var total_count int32 err := row.Scan(&total_count) @@ -2295,11 +2314,13 @@ WHERE m.owner_user_id = $1::bigint $7::text = '' OR m.body ILIKE ('%' || $7::text || '%') ) - AND ($8::int <= 0 OR m.box_id < $8::int) - AND ($9::int <= 0 OR m.box_id > $9::int) - AND (NOT $10::boolean OR m.pinned) + AND ($8::int <= 0 OR m.message_date > $8::int) + AND ($9::int <= 0 OR m.message_date < $9::int) + AND ($10::int <= 0 OR m.box_id < $10::int) + AND ($11::int <= 0 OR m.box_id > $11::int) + AND (NOT $12::boolean OR m.pinned) AND ( - NOT $11::boolean + NOT $13::boolean OR ( m.media->>'kind' = 'document' AND EXISTS ( @@ -2311,36 +2332,50 @@ WHERE m.owner_user_id = $1::bigint ) ) AND ( - $12::text = '' - OR (m.saved_peer_type = $12::text AND m.saved_peer_id = $13::bigint) + $14::text = '' + OR (m.saved_peer_type = $14::text AND m.saved_peer_id = $15::bigint) ) AND ( - ($14::int > 0 AND m.message_date < $14::int) - OR ($14::int <= 0 AND ($15::int <= 0 OR m.box_id < $15::int)) + cardinality($16::text[]) = 0 + OR EXISTS ( + SELECT 1 + FROM saved_message_reaction_tags tag + WHERE tag.user_id = m.owner_user_id + AND tag.message_box_id = m.box_id + AND (tag.reaction_type || ':' || tag.reaction_value) + = ANY($16::text[]) + ) + ) + AND ( + ($17::int > 0 AND m.message_date < $17::int) + OR ($17::int <= 0 AND ($18::int <= 0 OR m.box_id < $18::int)) ) ORDER BY m.box_id DESC -OFFSET GREATEST($16::int, 0) -LIMIT $17::int +OFFSET GREATEST($19::int, 0) +LIMIT $20::int ` type ListMessagesBackwardParams struct { - OwnerUserID int64 - HasPeer bool - PeerType string - PeerID int64 - RestrictPeerIds bool - PeerIds []int64 - Query string - MaxID int32 - MinID int32 - PinnedOnly bool - MusicOnly bool - SavedPeerType string - SavedPeerID int64 - OffsetDate int32 - OffsetID int32 - RowOffset int32 - LimitCount int32 + OwnerUserID int64 + HasPeer bool + PeerType string + PeerID int64 + RestrictPeerIds bool + PeerIds []int64 + Query string + MinDate int32 + MaxDate int32 + MaxID int32 + MinID int32 + PinnedOnly bool + MusicOnly bool + SavedPeerType string + SavedPeerID int64 + SavedReactionKeys []string + OffsetDate int32 + OffsetID int32 + RowOffset int32 + LimitCount int32 } type ListMessagesBackwardRow struct { @@ -2433,12 +2468,15 @@ func (q *Queries) ListMessagesBackward(ctx context.Context, arg ListMessagesBack arg.RestrictPeerIds, arg.PeerIds, arg.Query, + arg.MinDate, + arg.MaxDate, arg.MaxID, arg.MinID, arg.PinnedOnly, arg.MusicOnly, arg.SavedPeerType, arg.SavedPeerID, + arg.SavedReactionKeys, arg.OffsetDate, arg.OffsetID, arg.RowOffset, @@ -2641,11 +2679,13 @@ base AS NOT MATERIALIZED ( $11::text = '' OR m.body ILIKE ('%' || $11::text || '%') ) - AND ($12::int <= 0 OR m.box_id < $12::int) - AND ($13::int <= 0 OR m.box_id > $13::int) - AND (NOT $14::boolean OR m.pinned) + AND ($12::int <= 0 OR m.message_date > $12::int) + AND ($13::int <= 0 OR m.message_date < $13::int) + AND ($14::int <= 0 OR m.box_id < $14::int) + AND ($15::int <= 0 OR m.box_id > $15::int) + AND (NOT $16::boolean OR m.pinned) AND ( - NOT $15::boolean + NOT $17::boolean OR ( m.media->>'kind' = 'document' AND EXISTS ( @@ -2657,14 +2697,25 @@ base AS NOT MATERIALIZED ( ) ) AND ( - $16::text = '' - OR (m.saved_peer_type = $16::text AND m.saved_peer_id = $17::bigint) + $18::text = '' + OR (m.saved_peer_type = $18::text AND m.saved_peer_id = $19::bigint) + ) + AND ( + cardinality($20::text[]) = 0 + OR EXISTS ( + SELECT 1 + FROM saved_message_reaction_tags tag + WHERE tag.user_id = m.owner_user_id + AND tag.message_box_id = m.box_id + AND (tag.reaction_type || ':' || tag.reaction_value) + = ANY($20::text[]) + ) ) ), total AS ( SELECT count(*)::int AS total_count FROM base - WHERE $18::boolean + WHERE $21::boolean ), backward AS ( SELECT b.box_id, b.private_message_id, b.owner_user_id, b.peer_type, b.peer_id, b.from_user_id, b.message_date, b.ttl_period, b.expires_at, b.edit_date, b.hide_edited, b.outgoing, b.body, b.entities_json, b.silent, b.noforwards, b.reply_to_msg_id, b.reply_to_peer_type, b.reply_to_peer_id, b.reply_to_top_id, b.reply_to_story_id, b.quote_text, b.quote_entities_json, b.quote_offset, b.fwd_from_peer_type, b.fwd_from_peer_id, b.fwd_from_name, b.fwd_date, b.fwd_saved_from_peer_type, b.fwd_saved_from_peer_id, b.fwd_saved_from_msg_id, b.saved_peer_type, b.saved_peer_id, b.pts, b.media_json, b.media_unread, b.reaction_unread, b.pinned, b.via_bot_id, b.grouped_id, b.effect, b.reply_markup_json, b.rich_message_json, b.peer_user_id, b.peer_access_hash, b.peer_phone, b.peer_first_name, b.peer_last_name, b.peer_username, b.peer_country_code, b.peer_verified, b.peer_support, b.peer_is_bot, b.peer_bot_info_version, b.peer_premium_until, b.peer_emoji_status_document_id, b.peer_emoji_status_until, b.peer_last_seen_at, b.from_user_user_id, b.from_user_access_hash, b.from_user_phone, b.from_user_first_name, b.from_user_last_name, b.from_user_username, b.from_user_country_code, b.from_user_verified, b.from_user_support, b.from_user_is_bot, b.from_user_bot_info_version, b.from_user_premium_until, b.from_user_emoji_status_document_id, b.from_user_emoji_status_until, b.from_user_last_seen_at @@ -2811,24 +2862,27 @@ ORDER BY box_id DESC ` type ListMessagesByUserParams struct { - OwnerUserID int64 - OffsetID int32 - OffsetDate int32 - AddOffset int32 - LimitCount int32 - HasPeer bool - PeerType string - PeerID int64 - RestrictPeerIds bool - PeerIds []int64 - Query string - MaxID int32 - MinID int32 - PinnedOnly bool - MusicOnly bool - SavedPeerType string - SavedPeerID int64 - NeedTotalCount bool + OwnerUserID int64 + OffsetID int32 + OffsetDate int32 + AddOffset int32 + LimitCount int32 + HasPeer bool + PeerType string + PeerID int64 + RestrictPeerIds bool + PeerIds []int64 + Query string + MinDate int32 + MaxDate int32 + MaxID int32 + MinID int32 + PinnedOnly bool + MusicOnly bool + SavedPeerType string + SavedPeerID int64 + SavedReactionKeys []string + NeedTotalCount bool } type ListMessagesByUserRow struct { @@ -2921,12 +2975,15 @@ func (q *Queries) ListMessagesByUser(ctx context.Context, arg ListMessagesByUser arg.RestrictPeerIds, arg.PeerIds, arg.Query, + arg.MinDate, + arg.MaxDate, arg.MaxID, arg.MinID, arg.PinnedOnly, arg.MusicOnly, arg.SavedPeerType, arg.SavedPeerID, + arg.SavedReactionKeys, arg.NeedTotalCount, ) if err != nil { diff --git a/internal/store/postgres/star_gift_lifecycle_migration_integration_test.go b/internal/store/postgres/star_gift_lifecycle_migration_integration_test.go index 05576a5d..9436fae2 100644 --- a/internal/store/postgres/star_gift_lifecycle_migration_integration_test.go +++ b/internal/store/postgres/star_gift_lifecycle_migration_integration_test.go @@ -14,7 +14,7 @@ func TestStarGiftLifecycleMigrationsApply(t *testing.T) { if err != nil { t.Fatalf("migrate star gift lifecycle schema: %v", err) } - if status.Dirty || status.Empty || status.Version != 147 { - t.Fatalf("migration status = %+v, want clean version 147", status) + if status.Dirty || status.Empty || status.Version != 148 { + t.Fatalf("migration status = %+v, want clean version 148", status) } }