owpengram-server/internal/rpc/messages_forward.go

549 lines
17 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package rpc
import (
"context"
"errors"
"strings"
"github.com/gotd/td/tg"
"telesrv/internal/domain"
)
func (r *Router) onMessagesForwardMessages(ctx context.Context, req *tg.MessagesForwardMessagesRequest) (tg.UpdatesClass, error) {
ids, randomIDs, ok := normalizeForwardMessageVectors(req.ID, req.RandomID)
if !ok {
return nil, inputRequestInvalidErr()
}
req.ID = ids
req.RandomID = randomIDs
if len(req.ID) > domain.MaxForwardMessageIDs {
return nil, limitInvalidErr()
}
if req.ScheduleRepeatPeriod != 0 {
return nil, scheduleDateInvalidErr()
}
if err := forwardMessagesUnsupportedOptionErr(req); err != nil {
return nil, err
}
topMsgID, topMsgIDSet := req.GetTopMsgID()
if !topMsgIDSet && req.TopMsgID != 0 {
topMsgID, topMsgIDSet = req.TopMsgID, true
}
if topMsgIDSet && topMsgID == -1 {
topMsgID, topMsgIDSet = 0, false
}
if topMsgIDSet && (topMsgID < 0 || topMsgID > domain.MaxMessageBoxID) {
return nil, replyMessageIDInvalidErr()
}
userID, _, err := r.currentUserID(ctx)
if err != nil {
return nil, internalErr()
}
if userID == 0 {
return nil, peerIDInvalidErr()
}
fromPeer, preloadedSources, err := r.forwardFromPeerAndSources(ctx, userID, req.FromPeer, req.ID, req.RandomID)
if err != nil {
return nil, err
}
toPeer, err := r.checkedDomainPeerFromInputPeer(ctx, userID, req.ToPeer)
if err != nil {
return nil, err
}
sendAs, err := r.resolveSendAsPeer(ctx, userID, toPeer, req.SendAs)
if err != nil {
return nil, err
}
replyTo, err := r.messageReplyFromInput(ctx, userID, toPeer, req.ReplyTo)
if err != nil {
return nil, err
}
replyTo, err = mergeForwardTopMsgID(toPeer, replyTo, topMsgID, topMsgIDSet)
if err != nil {
return nil, err
}
if r.deps.Users != nil {
for _, peer := range []domain.Peer{fromPeer, toPeer} {
if peer.Type != domain.PeerTypeUser || peer.ID == userID {
continue
}
if _, found, err := r.deps.Users.ByID(ctx, userID, peer.ID); err != nil {
return nil, internalErr()
} else if !found {
return nil, peerIDInvalidErr()
}
}
}
if !forwardMessageIDsValid(req.ID, req.RandomID) {
return nil, messageIDInvalidErr()
}
if err := r.checkSendRateLimit(ctx, userID, len(req.ID)); err != nil {
return nil, err
}
if req.ScheduleDate != 0 && !scheduleDateIsImmediate(req.ScheduleDate, int(r.clock.Now().Unix())) {
return r.scheduleForwardMessages(ctx, userID, fromPeer, toPeer, req, replyTo, sendAs, preloadedSources)
}
if toPeer.Type == domain.PeerTypeChannel {
if r.deps.Channels == nil {
return nil, peerIDInvalidErr()
}
sources, err := r.forwardSourcesForRequest(ctx, userID, fromPeer, req.ID, preloadedSources)
if err != nil {
return nil, messageForwardErr(err)
}
recipients := make([]int64, 0)
results := make([]domain.SendChannelMessageResult, 0, len(sources))
extraUserIDs := make([]int64, 0, len(sources))
for i, source := range sources {
forward := source.forward
if req.DropAuthor {
forward = nil
}
mentionUserIDs := r.mentionUserIDsFromDomain(ctx, userID, source.body, source.entities)
res, err := r.deps.Channels.SendMessage(ctx, userID, domain.SendChannelMessageRequest{
UserID: userID,
ChannelID: toPeer.ID,
RandomID: req.RandomID[i],
Message: source.body,
Entities: source.entities,
Media: source.media,
MentionUserIDs: mentionUserIDs,
Silent: req.Silent,
NoForwards: req.Noforwards,
ReplyTo: replyTo,
Forward: forward,
SendAs: sendAs,
Date: int(r.clock.Now().Unix()),
})
if err != nil {
return nil, channelInvalidErr(err)
}
results = append(results, res)
// 收件人是该频道的活跃成员集,对本次转发的每条源都相同;只取一次,
// 避免一次转发 ≤100 条到 N 成员大群时把 recipients 累积成 ~100×N 条目
// 的巨大临时切片N=10^5 时约千万级)。
if len(recipients) == 0 {
recipients = res.Recipients
}
if sourceUserID := source.userID(); sourceUserID != 0 {
extraUserIDs = append(extraUserIDs, sourceUserID)
}
}
// echo 与 fan-out 用各自独立 cacheRPC vs worker goroutine 防竞态)。多条转发汇成
// 一个 fan-out jobchannelMessagesUpdatesWithPeerCache 内含多条 UpdateNewChannelMessage
// 由同 channel 分片 FIFO 原子投递。
echoCache := newViewerPeerCache(r)
updates := r.channelMessagesUpdatesWithPeerCache(ctx, userID, results, req.RandomID, true, extraUserIDs, echoCache)
fanoutPts := 0
if n := len(results); n > 0 {
fanoutPts = results[n-1].Event.Pts
}
r.enqueueChannelMessagesFanout(ctx, userID, toPeer.ID, fanoutPts, recipients, results, extraUserIDs)
for _, res := range results {
r.pushChannelDiscussionUpdate(ctx, userID, res.Discussion)
}
return updates, nil
}
if toPeer.Type == domain.PeerTypeUser {
if r.deps.Messages == nil {
return nil, peerIDInvalidErr()
}
recipientBlocked, err := r.peerBlocksUser(ctx, userID, toPeer.ID)
if err != nil {
return nil, err
}
// 私聊源与频道源统一经 forwardSources 取源:首次生成的 forward header 在
// forwardSources 内已按原作者 PrivacyKeyForwards 降级(不允许链接回账号时仅保
// 留 from_name避免私聊→私聊路径泄漏原作者可点击账号media 也随 source 透传。
sources, err := r.forwardSourcesForRequest(ctx, userID, fromPeer, req.ID, preloadedSources)
if err != nil {
return nil, messageForwardErr(err)
}
sessionID, _ := SessionIDFrom(ctx)
authKeyID, _ := AuthKeyIDFrom(ctx)
res := domain.ForwardPrivateMessagesResult{OwnerUserID: userID}
for i, source := range sources {
forward := source.forward
if req.DropAuthor {
forward = nil
}
if forward != nil && toPeer.ID == userID {
saved := *forward
saved.SavedFrom = fromPeer
saved.SavedFromMsgID = req.ID[i]
forward = &saved
}
sent, err := r.deps.Messages.SendPrivateText(ctx, userID, domain.SendPrivateTextRequest{
SenderUserID: userID,
RecipientUserID: toPeer.ID,
RandomID: req.RandomID[i],
Message: source.body,
Entities: source.entities,
Media: source.media,
Silent: req.Silent,
NoForwards: req.Noforwards,
ReplyTo: replyTo,
Forward: forward,
Date: int(r.clock.Now().Unix()),
OriginAuthKeyID: authKeyID,
OriginSessionID: sessionID,
RecipientBlocked: recipientBlocked,
})
if err != nil {
return nil, messageForwardErr(err)
}
if !sent.Duplicate {
r.enqueueBotAPIPrivateMessageUpdateAsync(ctx, sent)
}
res.SenderMessages = append(res.SenderMessages, sent.SenderMessage)
res.RecipientMessages = append(res.RecipientMessages, sent.RecipientMessage)
res.SenderEvents = append(res.SenderEvents, sent.SenderEvent)
res.RecipientEvents = append(res.RecipientEvents, sent.RecipientEvent)
res.Duplicates = append(res.Duplicates, sent.Duplicate)
}
return tgForwardMessagesUpdates(res, req.RandomID, r.usersForMessageUpdates(ctx, userID, res.SenderMessages), r.chatsForMessageUpdates(ctx, userID, res.SenderMessages)), nil
}
return nil, peerIDInvalidErr()
}
func normalizeForwardMessageVectors(ids []int, randomIDs []int64) ([]int, []int64, bool) {
if len(ids) == 0 || len(randomIDs) == 0 {
return nil, nil, false
}
if len(ids) == len(randomIDs) {
return ids, randomIDs, true
}
if len(ids) < len(randomIDs) {
return nil, nil, false
}
compact := make([]int, 0, len(randomIDs))
runLength := 0
for i, id := range ids {
if i == 0 || id != ids[i-1] {
compact = append(compact, id)
runLength = 1
continue
}
runLength++
if runLength > 2 {
return nil, nil, false
}
}
if len(compact) != len(randomIDs) {
return nil, nil, false
}
return compact, randomIDs, true
}
func (r *Router) forwardFromPeerAndSources(ctx context.Context, userID int64, input tg.InputPeerClass, ids []int, randomIDs []int64) (domain.Peer, []forwardSource, error) {
if forwardFromPeerIsEmpty(input) {
if !forwardMessageIDsValid(ids, randomIDs) {
return domain.Peer{}, nil, messageIDInvalidErr()
}
fromPeer, sources, err := r.forwardSourcesFromEmptyPeer(ctx, userID, ids)
if err != nil {
return domain.Peer{}, nil, messageForwardErr(err)
}
return fromPeer, sources, nil
}
fromPeer, err := r.checkedDomainPeerFromInputPeer(ctx, userID, input)
return fromPeer, nil, err
}
func forwardMessageIDsValid(ids []int, randomIDs []int64) bool {
if len(ids) == 0 || len(ids) != len(randomIDs) {
return false
}
for i, id := range ids {
if id <= 0 || id > domain.MaxMessageBoxID || randomIDs[i] == 0 {
return false
}
}
return true
}
func forwardFromPeerIsEmpty(peer tg.InputPeerClass) bool {
if inputPeerClassNil(peer) {
return false
}
_, ok := peer.(*tg.InputPeerEmpty)
return ok
}
func (r *Router) forwardSourcesFromEmptyPeer(ctx context.Context, userID int64, ids []int) (domain.Peer, []forwardSource, error) {
if r.deps.Messages == nil {
return domain.Peer{}, nil, domain.ErrMessageIDInvalid
}
list, err := r.deps.Messages.GetMessages(ctx, userID, ids)
if err != nil {
return domain.Peer{}, nil, domain.ErrMessageIDInvalid
}
var fromPeer domain.Peer
byID := make(map[int]domain.Message, len(list.Messages))
for _, msg := range list.Messages {
byID[msg.ID] = msg
}
for _, id := range ids {
msg, ok := byID[id]
if !ok || msg.Peer.Type != domain.PeerTypeUser || msg.Peer.ID == 0 {
return domain.Peer{}, nil, domain.ErrMessageIDInvalid
}
if fromPeer.ID == 0 {
fromPeer = msg.Peer
continue
}
if msg.Peer != fromPeer {
return domain.Peer{}, nil, domain.ErrMessageIDInvalid
}
}
sources, err := r.forwardSourcesFromPrivateMessages(ctx, userID, fromPeer, ids, list.Messages)
if err != nil {
return domain.Peer{}, nil, err
}
return fromPeer, sources, nil
}
func (r *Router) forwardSourcesForRequest(ctx context.Context, userID int64, fromPeer domain.Peer, ids []int, preloaded []forwardSource) ([]forwardSource, error) {
if preloaded != nil {
return preloaded, nil
}
return r.forwardSources(ctx, userID, fromPeer, ids)
}
func mergeForwardTopMsgID(toPeer domain.Peer, replyTo *domain.MessageReply, topMsgID int, topMsgIDSet bool) (*domain.MessageReply, error) {
if !topMsgIDSet || topMsgID == 0 {
return replyTo, nil
}
if topMsgID < 0 || topMsgID > domain.MaxMessageBoxID || toPeer.Type != domain.PeerTypeChannel {
return nil, replyMessageIDInvalidErr()
}
if replyTo == nil {
return &domain.MessageReply{
Peer: toPeer,
TopMessageID: topMsgID,
ForumTopic: true,
}, nil
}
if replyTo.Peer.ID != 0 && replyTo.Peer != toPeer {
return nil, replyMessageIDInvalidErr()
}
if replyTo.TopMessageID != 0 && replyTo.TopMessageID != topMsgID {
return nil, replyMessageIDInvalidErr()
}
merged := *replyTo
merged.Peer = toPeer
merged.TopMessageID = topMsgID
merged.QuoteEntities = append([]domain.MessageEntity(nil), replyTo.QuoteEntities...)
if merged.MessageID == 0 {
merged.ForumTopic = true
}
return &merged, nil
}
func (r *Router) forwardSources(ctx context.Context, userID int64, fromPeer domain.Peer, ids []int) ([]forwardSource, error) {
out := make([]forwardSource, 0, len(ids))
switch fromPeer.Type {
case domain.PeerTypeUser:
if r.deps.Messages == nil {
return nil, domain.ErrMessageIDInvalid
}
list, err := r.deps.Messages.GetMessages(ctx, userID, ids)
if err != nil {
return nil, domain.ErrMessageIDInvalid
}
sources, err := r.forwardSourcesFromPrivateMessages(ctx, userID, fromPeer, ids, list.Messages)
if err != nil {
return nil, err
}
out = append(out, sources...)
case domain.PeerTypeChannel:
if r.deps.Channels == nil {
return nil, domain.ErrMessageIDInvalid
}
history, err := r.deps.Channels.GetMessages(ctx, userID, fromPeer.ID, ids)
if err != nil {
return nil, domain.ErrMessageIDInvalid
}
byID := make(map[int]domain.ChannelMessage, len(history.Messages))
for _, msg := range history.Messages {
byID[msg.ID] = msg
}
for _, id := range ids {
msg, ok := byID[id]
if !ok {
return nil, domain.ErrMessageIDInvalid
}
if msg.NoForwards || history.Channel.NoForwards {
return nil, domain.ErrChatForwardsRestricted
}
if msg.Action != nil || (msg.Body == "" && msg.Media.IsZero()) {
return nil, domain.ErrMessageIDInvalid
}
forward := cloneDomainMessageForward(msg.Forward)
from := msg.From
if from.ID == 0 && msg.SenderUserID != 0 {
from = domain.Peer{Type: domain.PeerTypeUser, ID: msg.SenderUserID}
}
if msg.Post {
from = domain.Peer{Type: domain.PeerTypeChannel, ID: msg.ChannelID}
}
if forward == nil {
forward = &domain.MessageForward{From: from, Date: msg.Date}
if from.Type == domain.PeerTypeChannel {
forward.ChannelPost = msg.ID
}
r.applyForwardAuthorPrivacy(ctx, userID, forward)
}
out = append(out, forwardSource{
body: msg.Body,
entities: append([]domain.MessageEntity(nil),
msg.Entities...),
media: msg.Media,
forward: forward,
from: from,
date: msg.Date,
})
}
default:
return nil, domain.ErrMessageIDInvalid
}
return out, nil
}
func (r *Router) forwardSourcesFromPrivateMessages(ctx context.Context, userID int64, fromPeer domain.Peer, ids []int, messages []domain.Message) ([]forwardSource, error) {
if fromPeer.Type != domain.PeerTypeUser || fromPeer.ID == 0 {
return nil, domain.ErrMessageIDInvalid
}
byID := make(map[int]domain.Message, len(messages))
for _, msg := range messages {
byID[msg.ID] = msg
}
out := make([]forwardSource, 0, len(ids))
for _, id := range ids {
msg, ok := byID[id]
if !ok {
return nil, domain.ErrMessageIDInvalid
}
if msg.Peer != fromPeer {
return nil, domain.ErrMessageIDInvalid
}
if msg.NoForwards {
return nil, domain.ErrChatForwardsRestricted
}
forward := cloneDomainMessageForward(msg.Forward)
if forward == nil {
forward = &domain.MessageForward{From: msg.From, Date: msg.Date}
r.applyForwardAuthorPrivacy(ctx, userID, forward)
}
out = append(out, forwardSource{
body: msg.Body,
entities: append([]domain.MessageEntity(nil),
msg.Entities...),
media: msg.Media,
forward: forward,
from: msg.From,
date: msg.Date,
})
}
return out, nil
}
func cloneDomainMessageForward(in *domain.MessageForward) *domain.MessageForward {
if in == nil {
return nil
}
out := *in
return &out
}
func (r *Router) forwardAuthorDisplayName(ctx context.Context, viewerUserID, authorUserID int64) string {
name := ""
if r.deps.Users != nil {
if author, found, err := r.deps.Users.ByID(ctx, viewerUserID, authorUserID); err == nil && found {
name = strings.TrimSpace(strings.TrimSpace(author.FirstName) + " " + strings.TrimSpace(author.LastName))
}
}
if name == "" {
name = "Deleted Account"
}
return name
}
// applyForwardAuthorPrivacy 在首次生成 forward header 时按原作者的
// forwards 隐私规则降级:不允许链接回账号时只保留展示名 from_name。
// 已有 header 的再转发沿用原 header不重新评估。
func (r *Router) applyForwardAuthorPrivacy(ctx context.Context, forwarderUserID int64, forward *domain.MessageForward) {
if forward == nil || forward.From.Type != domain.PeerTypeUser || forward.From.ID == 0 || forward.FromName != "" {
return
}
if forward.From.ID == forwarderUserID || r.deps.Privacy == nil {
return
}
allowed, err := r.deps.Privacy.CanSee(ctx, forward.From.ID, forwarderUserID, domain.PrivacyKeyForwards)
if err != nil || allowed {
return
}
name := r.forwardAuthorDisplayName(ctx, forwarderUserID, forward.From.ID)
forward.From = domain.Peer{}
forward.FromName = name
}
func messageForwardErr(err error) error {
switch {
case errors.Is(err, domain.ErrMessageIDInvalid):
return messageIDInvalidErr()
case errors.Is(err, domain.ErrChatForwardsRestricted):
return chatForwardsRestrictedErr()
case errors.Is(err, domain.ErrReplyMessageIDInvalid):
return replyMessageIDInvalidErr()
default:
return internalErr()
}
}
func tgForwardMessagesUpdates(res domain.ForwardPrivateMessagesResult, randomIDs []int64, users []tg.UserClass, chats []tg.ChatClass) *tg.Updates {
updates := make([]tg.UpdateClass, 0, len(res.SenderMessages)*2)
date := 0
for i, msg := range res.SenderMessages {
randomID := int64(0)
if i < len(randomIDs) {
randomID = randomIDs[i]
}
updates = append(updates, &tg.UpdateMessageID{ID: msg.ID, RandomID: randomID})
event := domain.UpdateEvent{}
if i < len(res.SenderEvents) {
event = res.SenderEvents[i]
}
item := tgMessage(msg)
if item == nil {
item = &tg.MessageEmpty{ID: msg.ID}
}
pts := event.Pts
if pts == 0 {
pts = msg.Pts
}
ptsCount := event.PtsCount
if ptsCount == 0 {
ptsCount = 1
}
updates = append(updates, &tg.UpdateNewMessage{
Message: item,
Pts: pts,
PtsCount: ptsCount,
})
if date == 0 {
date = event.Date
}
if date == 0 {
date = msg.Date
}
}
return &tg.Updates{
Updates: updates,
Users: users,
Chats: chats,
Date: date,
Seq: 0,
}
}