owpengram-server/internal/rpc/messages_send_rpc_test.go

389 lines
14 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"
"github.com/gotd/td/bin"
"github.com/gotd/td/clock"
"github.com/gotd/td/tg"
"go.uber.org/zap/zaptest"
"strconv"
"strings"
"telesrv/internal/domain"
"testing"
"time"
)
func TestMessagesSendMessageReturnsUpdateAndRecordsOwnerContext(t *testing.T) {
sender := domain.User{ID: 1000000001, AccessHash: 11, FirstName: "Sender"}
recipient := domain.User{ID: 1000000002, AccessHash: 22, FirstName: "Recipient"}
messages := &captureMessages{}
metrics := &captureRPCMetrics{}
r := New(Config{}, Deps{
Messages: messages,
Users: mapUsersService{users: map[int64]domain.User{sender.ID: sender, recipient.ID: recipient}},
Metrics: metrics,
}, zaptest.NewLogger(t), clock.System)
req := &tg.MessagesSendMessageRequest{
Peer: &tg.InputPeerUser{UserID: recipient.ID, AccessHash: recipient.AccessHash},
Message: "hello",
RandomID: 123456,
Entities: []tg.MessageEntityClass{
&tg.MessageEntityBold{Offset: 0, Length: 5},
},
}
var in bin.Buffer
if err := req.Encode(&in); err != nil {
t.Fatalf("encode request: %v", err)
}
enc, err := r.Dispatch(WithUserID(context.Background(), sender.ID), [8]byte{}, 77, &in)
if err != nil {
t.Fatalf("dispatch: %v", err)
}
box, ok := enc.(*tg.UpdatesBox)
if !ok {
t.Fatalf("response = %T, want *tg.UpdatesBox", enc)
}
got, ok := box.Updates.(*tg.Updates)
if !ok {
t.Fatalf("boxed response = %T, want *tg.Updates", box.Updates)
}
if messages.sendUserID != sender.ID || messages.sendReq.SenderUserID != sender.ID || messages.sendReq.RecipientUserID != recipient.ID || messages.sendReq.OriginSessionID != 77 {
t.Fatalf("send context = user %d req %+v, want sender/recipient/session", messages.sendUserID, messages.sendReq)
}
if len(messages.sendReq.Entities) != 1 || messages.sendReq.Entities[0].Type != domain.MessageEntityBold {
t.Fatalf("entities = %+v, want bold entity converted to domain", messages.sendReq.Entities)
}
if len(got.Updates) != 2 {
t.Fatalf("updates = %+v, want message id + new message", got.Updates)
}
if id, ok := got.Updates[0].(*tg.UpdateMessageID); !ok || id.ID != 1 || id.RandomID != req.RandomID {
t.Fatalf("update id = %#v, want id=1 random_id=%d", got.Updates[0], req.RandomID)
}
newMsg, ok := got.Updates[1].(*tg.UpdateNewMessage)
if !ok || newMsg.Pts != 1 || newMsg.PtsCount != 1 {
t.Fatalf("new message update = %#v, want pts=1 pts_count=1", got.Updates[1])
}
msg, ok := newMsg.Message.(*tg.Message)
if !ok || !msg.Out || msg.PeerID.(*tg.PeerUser).UserID != recipient.ID || msg.Message != req.Message {
t.Fatalf("message = %#v, want outgoing private text to recipient", newMsg.Message)
}
if metrics.messageSend != 1 || metrics.messageSendErr != nil {
t.Fatalf("metrics send=%d err=%v, want one successful send", metrics.messageSend, metrics.messageSendErr)
}
}
func TestMessagesSendMessageRateLimitReturnsFloodWait(t *testing.T) {
const userID = int64(1000000001)
limiter := &captureRateLimiter{block: true, retryAfter: 9}
metrics := &captureRPCMetrics{}
r := New(Config{SendRateLimit: 1, SendRateWindow: time.Minute}, Deps{
Limiter: limiter,
Metrics: metrics,
}, zaptest.NewLogger(t), clock.System)
_, err := r.onMessagesSendMessage(WithUserID(context.Background(), userID), &tg.MessagesSendMessageRequest{
Peer: &tg.InputPeerSelf{},
Message: "too fast",
RandomID: 123456,
})
if err == nil || !strings.Contains(err.Error(), "FLOOD_WAIT") || !strings.Contains(err.Error(), "(9)") {
t.Fatalf("sendMessage rate err = %v, want FLOOD_WAIT 9", err)
}
if len(limiter.calls) != 1 {
t.Fatalf("limiter calls = %d, want 1", len(limiter.calls))
}
call := limiter.calls[0]
if call.key != sendRateLimitKeyPrefix+strconv.FormatInt(userID, 10) || call.cost != 1 || call.limit != 1 || call.window != time.Minute {
t.Fatalf("limiter call = %+v, want send key cost=1 limit=1 window=1m", call)
}
if metrics.rateLimited != 9 {
t.Fatalf("rate limited metric = %d, want 9", metrics.rateLimited)
}
}
func TestMessagesForwardMessagesRateLimitCountsIDs(t *testing.T) {
const userID = int64(1000000001)
limiter := &captureRateLimiter{block: true, retryAfter: 13}
r := New(Config{SendRateLimit: 3, SendRateWindow: 4 * time.Second}, Deps{Limiter: limiter}, zaptest.NewLogger(t), clock.System)
_, err := r.onMessagesForwardMessages(WithUserID(context.Background(), userID), &tg.MessagesForwardMessagesRequest{
FromPeer: &tg.InputPeerSelf{},
ToPeer: &tg.InputPeerSelf{},
ID: []int{1, 2, 3},
RandomID: []int64{11, 22, 33},
})
if err == nil || !strings.Contains(err.Error(), "FLOOD_WAIT") || !strings.Contains(err.Error(), "(13)") {
t.Fatalf("forwardMessages rate err = %v, want FLOOD_WAIT 13", err)
}
if len(limiter.calls) != 1 {
t.Fatalf("limiter calls = %d, want 1", len(limiter.calls))
}
call := limiter.calls[0]
if call.key != sendRateLimitKeyPrefix+strconv.FormatInt(userID, 10) || call.cost != 3 || call.limit != 3 || call.window != 4*time.Second {
t.Fatalf("limiter call = %+v, want shared send key cost=3 limit=3 window=4s", call)
}
}
func TestMessagesSendMessageSupportsReplyAndFlags(t *testing.T) {
const (
senderID = int64(1000000001)
recipientID = int64(1000000002)
)
messages := &captureMessages{}
r := New(Config{}, Deps{
Messages: messages,
Users: mapUsersService{users: map[int64]domain.User{senderID: {ID: senderID, FirstName: "Sender"}, recipientID: {ID: recipientID, FirstName: "Recipient"}}},
}, zaptest.NewLogger(t), clock.System)
reply := &tg.InputReplyToMessage{ReplyToMsgID: 7}
reply.SetQuoteText("hello")
reply.SetQuoteOffset(1)
req := &tg.MessagesSendMessageRequest{
Peer: &tg.InputPeerUser{UserID: recipientID},
Message: "reply",
RandomID: 456,
Silent: true,
}
req.SetNoforwards(true)
req.SetReplyTo(reply)
var in bin.Buffer
if err := req.Encode(&in); err != nil {
t.Fatalf("encode request: %v", err)
}
enc, err := r.Dispatch(WithUserID(context.Background(), senderID), [8]byte{}, 88, &in)
if err != nil {
t.Fatalf("dispatch: %v", err)
}
if messages.sendReq.ReplyTo == nil || messages.sendReq.ReplyTo.MessageID != 7 || messages.sendReq.ReplyTo.Peer.ID != recipientID || messages.sendReq.ReplyTo.QuoteText != "hello" {
t.Fatalf("reply request = %+v, want reply metadata", messages.sendReq.ReplyTo)
}
if !messages.sendReq.Silent || !messages.sendReq.NoForwards {
t.Fatalf("send flags silent=%v noforwards=%v, want true/true", messages.sendReq.Silent, messages.sendReq.NoForwards)
}
box, ok := enc.(*tg.UpdatesBox)
if !ok {
t.Fatalf("response = %T, want *tg.UpdatesBox", enc)
}
got := box.Updates.(*tg.Updates)
newMsg := got.Updates[1].(*tg.UpdateNewMessage)
msg := newMsg.Message.(*tg.Message)
if !msg.Silent || !msg.Noforwards {
t.Fatalf("message flags silent=%v noforwards=%v, want true/true", msg.Silent, msg.Noforwards)
}
header, ok := msg.ReplyTo.(*tg.MessageReplyHeader)
if !ok || header.ReplyToMsgID != 7 {
t.Fatalf("reply header = %#v, want msg id 7", msg.ReplyTo)
}
}
func TestMessagesSendMessageRejectsHugeReplyQuoteOffset(t *testing.T) {
const (
senderID = int64(1000000001)
recipientID = int64(1000000002)
)
messages := &captureMessages{}
r := New(Config{}, Deps{
Messages: messages,
Users: mapUsersService{users: map[int64]domain.User{recipientID: {ID: recipientID, FirstName: "Recipient"}}},
}, zaptest.NewLogger(t), clock.System)
reply := &tg.InputReplyToMessage{ReplyToMsgID: 7}
reply.SetQuoteText("hello")
reply.SetQuoteOffset(domain.MaxMessageReplyQuoteOffset + 1)
req := &tg.MessagesSendMessageRequest{
Peer: &tg.InputPeerUser{UserID: recipientID},
Message: "reply",
RandomID: 457,
}
req.SetReplyTo(reply)
if _, err := r.onMessagesSendMessage(WithUserID(context.Background(), senderID), req); err == nil || !strings.Contains(err.Error(), "REPLY_MESSAGE_ID_INVALID") {
t.Fatalf("huge quote offset err = %v, want REPLY_MESSAGE_ID_INVALID", err)
}
if messages.sendReq.RandomID != 0 {
t.Fatalf("send request reached service: %+v", messages.sendReq)
}
}
// story 回复评论reply_to=inputReplyToStory 必须被接受并投影为 messageReplyStoryHeader
// 而非旧的 STORY_ID_INVALID 拒绝真机暴露Alice 回复 Bob story 时评论消息发送失败)。
func TestMessageReplyFromInputStorySucceedsAndProjectsStoryHeader(t *testing.T) {
sender := domain.User{ID: 1000000001, AccessHash: 11, FirstName: "Sender"}
recipient := domain.User{ID: 1000000002, AccessHash: 22, FirstName: "Recipient"}
r := New(Config{}, Deps{
Users: mapUsersService{users: map[int64]domain.User{sender.ID: sender, recipient.ID: recipient}},
}, zaptest.NewLogger(t), clock.System)
ctx := WithUserID(context.Background(), sender.ID)
peer := domain.Peer{Type: domain.PeerTypeUser, ID: recipient.ID}
input := &tg.InputReplyToStory{Peer: &tg.InputPeerUser{UserID: recipient.ID, AccessHash: recipient.AccessHash}, StoryID: 7}
reply, err := r.messageReplyFromInput(ctx, sender.ID, peer, input)
if err != nil {
t.Fatalf("story reply err = %v, want nilstory 回复应被接受)", err)
}
if reply == nil || reply.StoryID != 7 || reply.Peer != peer {
t.Fatalf("reply = %+v, want StoryID=7 peer=recipient", reply)
}
header := tgMessageReplyHeader(domain.Message{Peer: peer, ReplyTo: reply})
sh, ok := header.(*tg.MessageReplyStoryHeader)
if !ok {
t.Fatalf("header = %T, want *tg.MessageReplyStoryHeader", header)
}
if sh.StoryID != 7 {
t.Fatalf("header story id = %d, want 7", sh.StoryID)
}
if pu, ok := sh.Peer.(*tg.PeerUser); !ok || pu.UserID != recipient.ID {
t.Fatalf("header peer = %#v, want recipient story owner", sh.Peer)
}
// 回复非会话对端的 storypeer 不匹配)仍被拒。
wrongPeer := domain.Peer{Type: domain.PeerTypeUser, ID: sender.ID}
if _, err := r.messageReplyFromInput(ctx, sender.ID, wrongPeer, input); err == nil || !strings.Contains(err.Error(), "STORY_ID_INVALID") {
t.Fatalf("mismatched story owner err = %v, want STORY_ID_INVALID", err)
}
}
func TestMessageReplyFromInputUnsupportedShapesReturnExplicitErrors(t *testing.T) {
const userID = int64(1000000001)
ctx := WithUserID(context.Background(), userID)
r := New(Config{}, Deps{}, zaptest.NewLogger(t), clock.System)
peer := domain.Peer{Type: domain.PeerTypeChannel, ID: 1000000002}
withReplyMsg := func(update func(*tg.InputReplyToMessage)) *tg.InputReplyToMessage {
reply := &tg.InputReplyToMessage{ReplyToMsgID: 7}
update(reply)
return reply
}
cases := []struct {
name string
input tg.InputReplyToClass
want string
}{
{
name: "story",
input: &tg.InputReplyToStory{Peer: &tg.InputPeerUser{UserID: 1000000003}, StoryID: 1},
want: "STORY_ID_INVALID",
},
{
name: "monoforum constructor",
input: &tg.InputReplyToMonoForum{MonoforumPeerID: &tg.InputPeerChannel{ChannelID: 1000000004}},
want: "REPLY_TO_MONOFORUM_PEER_INVALID",
},
{
name: "monoforum field",
input: withReplyMsg(func(reply *tg.InputReplyToMessage) {
reply.SetMonoforumPeerID(&tg.InputPeerChannel{ChannelID: 1000000004})
}),
want: "REPLY_TO_MONOFORUM_PEER_INVALID",
},
{
name: "todo item",
input: withReplyMsg(func(reply *tg.InputReplyToMessage) {
reply.SetTodoItemID(1)
}),
want: "REPLY_MESSAGE_ID_INVALID",
},
{
name: "poll option",
input: withReplyMsg(func(reply *tg.InputReplyToMessage) {
reply.SetPollOption([]byte{1})
}),
want: "POLL_OPTION_INVALID",
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
if _, err := r.messageReplyFromInput(ctx, userID, peer, tc.input); err == nil || !strings.Contains(err.Error(), tc.want) {
t.Fatalf("reply err = %v, want %s", err, tc.want)
}
})
}
}
func TestMessagesSendMessageUnsupportedOptionErrors(t *testing.T) {
const (
senderID = int64(1000000001)
recipientID = int64(1000000002)
)
ctx := WithUserID(context.Background(), senderID)
r := New(Config{}, Deps{}, zaptest.NewLogger(t), clock.System)
base := func() *tg.MessagesSendMessageRequest {
return &tg.MessagesSendMessageRequest{
Peer: &tg.InputPeerUser{UserID: recipientID},
Message: "hello",
RandomID: 456,
}
}
suggested := func() tg.SuggestedPost {
post := tg.SuggestedPost{}
post.SetAccepted(true)
return post
}
cases := []struct {
name string
req *tg.MessagesSendMessageRequest
want string
}{
{
name: "quick reply",
req: func() *tg.MessagesSendMessageRequest {
req := base()
req.SetQuickReplyShortcut(&tg.InputQuickReplyShortcut{Shortcut: "hello"})
return req
}(),
want: "SHORTCUT_INVALID",
},
{
name: "effect",
req: func() *tg.MessagesSendMessageRequest {
req := base()
req.SetEffect(1)
return req
}(),
want: "EFFECT_ID_INVALID",
},
{
name: "negative paid stars",
req: func() *tg.MessagesSendMessageRequest {
req := base()
req.SetAllowPaidStars(-1)
return req
}(),
want: "STARS_AMOUNT_INVALID",
},
{
name: "paid stars",
req: func() *tg.MessagesSendMessageRequest {
req := base()
req.SetAllowPaidStars(1)
return req
}(),
want: "PAYMENT_UNSUPPORTED",
},
{
name: "paid floodskip",
req: func() *tg.MessagesSendMessageRequest {
req := base()
req.SetAllowPaidFloodskip(true)
return req
}(),
want: "PAYMENT_UNSUPPORTED",
},
{
name: "suggested post",
req: func() *tg.MessagesSendMessageRequest {
req := base()
req.SetSuggestedPost(suggested())
return req
}(),
want: "SUGGESTED_POST_PEER_INVALID",
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
if _, err := r.onMessagesSendMessage(ctx, tc.req); err == nil || !strings.Contains(err.Error(), tc.want) {
t.Fatalf("send err = %v, want %s", err, tc.want)
}
})
}
}