merged with fixes
This commit is contained in:
parent
a9e758b712
commit
2f1818d656
176 changed files with 9000 additions and 907 deletions
|
|
@ -43,7 +43,7 @@ func (r *Router) resolveAIComposeStyleWebPage(ctx context.Context, rawURL string
|
|||
Hash: aiComposeToneWebPageHash(tone),
|
||||
Date: int(now.Unix()),
|
||||
Type: aiComposeToneWebPageType,
|
||||
SiteName: branding.ProductName,
|
||||
SiteName: branding.ProductName(),
|
||||
Title: tone.Title,
|
||||
Description: tone.Prompt,
|
||||
ComposeToneEmojiID: tone.EmojiID,
|
||||
|
|
|
|||
|
|
@ -41,8 +41,18 @@ func (r *Router) onUpdatesGetChannelDifference(ctx context.Context, req *tg.Upda
|
|||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, domain.ErrPersistentTimestamp) {
|
||||
r.log.Debug("channel difference cursor rejected",
|
||||
zap.Int64("viewer_user_id", userID),
|
||||
zap.Int64("channel_id", channelID),
|
||||
zap.Int("request_pts", req.Pts))
|
||||
return nil, persistentTimestampInvalidErr()
|
||||
}
|
||||
r.log.Warn("load channel difference failed",
|
||||
zap.Int64("viewer_user_id", userID),
|
||||
zap.Int64("channel_id", channelID),
|
||||
zap.Int("request_pts", req.Pts),
|
||||
zap.Int("limit", req.Limit),
|
||||
zap.Error(err))
|
||||
return nil, channelInvalidErr(err)
|
||||
}
|
||||
diff, err = r.enrichChannelDifferenceStrict(ctx, userID, diff)
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import (
|
|||
"github.com/iamxvbaba/td/tg"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/transport"
|
||||
)
|
||||
|
||||
type ctxKey int
|
||||
|
|
@ -303,3 +304,15 @@ func invokeWithoutUpdatesFrom(ctx context.Context) bool {
|
|||
v, _ := ctx.Value(invokeWithoutUpdatesKey).(bool)
|
||||
return v
|
||||
}
|
||||
|
||||
// WithClientIP 在 ctx 注入客户端连接的对端 IP(来自 MTProto 连接的 RemoteAddr)。
|
||||
// 仅在需要时(绑定设备授权)由 edge 写入,其余 RPC 不依赖它。edge 通过中立
|
||||
// internal/transport 载体写入,这里仅做别名以便 rpc 业务层读取,避免反向依赖。
|
||||
func WithClientIP(ctx context.Context, ip string) context.Context {
|
||||
return transport.WithClientIP(ctx, ip)
|
||||
}
|
||||
|
||||
// ClientIPFrom 返回 ctx 中的客户端对端 IP,未设置时 ok=false。
|
||||
func ClientIPFrom(ctx context.Context) (string, bool) {
|
||||
return transport.ClientIPFrom(ctx)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -16,6 +16,9 @@ func (r *Router) authzFromCtx(ctx context.Context) domain.Authorization {
|
|||
a.AppVersion = ci.AppVersion
|
||||
a.APIID = ci.APIID
|
||||
}
|
||||
if ip, ok := ClientIPFrom(ctx); ok {
|
||||
a.IP = ip
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -346,7 +346,7 @@ func tgMessageReplyHeader(m domain.Message) tg.MessageReplyHeaderClass {
|
|||
}
|
||||
return &tg.MessageReplyStoryHeader{Peer: peer, StoryID: m.ReplyTo.StoryID}
|
||||
}
|
||||
if m.ReplyTo.MessageID <= 0 && m.ReplyTo.TopMessageID <= 0 {
|
||||
if m.ReplyTo.MessageID <= 0 && m.ReplyTo.TopMessageID <= 0 && m.ReplyTo.External == nil {
|
||||
return nil
|
||||
}
|
||||
header := &tg.MessageReplyHeader{}
|
||||
|
|
@ -364,6 +364,18 @@ func tgMessageReplyHeader(m domain.Message) tg.MessageReplyHeaderClass {
|
|||
header.SetReplyToPeerID(peer)
|
||||
}
|
||||
}
|
||||
if external := m.ReplyTo.External; external != nil {
|
||||
if from := tgMessageFwdHeader(&external.From); from != nil {
|
||||
header.SetReplyFrom(*from)
|
||||
}
|
||||
if !external.Media.IsZero() {
|
||||
header.SetReplyMedia(tgMessageMedia(external.Media))
|
||||
}
|
||||
if m.ReplyTo.QuoteText == "" && external.Text != "" {
|
||||
header.SetQuoteText(external.Text)
|
||||
header.SetQuoteEntities(tgMessageEntities(external.Entities))
|
||||
}
|
||||
}
|
||||
if m.ReplyTo.QuoteText != "" {
|
||||
header.SetQuote(true)
|
||||
header.SetQuoteText(m.ReplyTo.QuoteText)
|
||||
|
|
|
|||
|
|
@ -62,6 +62,8 @@ func peersListEmptyErr() error { return tgerr.New(400, "PEERS_LIST_EMPTY") }
|
|||
// peerIDInvalidErr 表示目标 peer 不存在或当前阶段不支持。
|
||||
func peerIDInvalidErr() error { return tgerr.New(400, "PEER_ID_INVALID") }
|
||||
|
||||
func fromPeerInvalidErr() error { return tgerr.New(400, "FROM_PEER_INVALID") }
|
||||
|
||||
func parentPeerInvalidErr() error { return tgerr.New(400, "PARENT_PEER_INVALID") }
|
||||
|
||||
func sendAsPeerInvalidErr() error { return tgerr.New(400, "SEND_AS_PEER_INVALID") }
|
||||
|
|
@ -442,6 +444,8 @@ func dhGAInvalidErr() error { return tgerr.New(400, "DH_G_A_INVALI
|
|||
func maxDateInvalidErr() error { return tgerr.New(400, "MAX_DATE_INVALID") }
|
||||
func fileEmptyErr() error { return tgerr.New(400, "FILE_EMPTY") }
|
||||
|
||||
func quoteTextInvalidErr() error { return tgerr.New(400, "QUOTE_TEXT_INVALID") }
|
||||
|
||||
// signInErr 把登录业务错误映射为客户端可识别的 rpc_error。
|
||||
func signInErr(err error) error {
|
||||
switch {
|
||||
|
|
|
|||
|
|
@ -25,7 +25,7 @@ func (r *Router) registerHelp(d *tlprofile.Dispatcher) {
|
|||
}, nil
|
||||
})
|
||||
registerRPC[*tg.HelpGetInviteTextRequest](d, tlprofile.SemanticMethodHelpGetInviteText, func(ctx context.Context, layerRequest *tg.HelpGetInviteTextRequest) (any, error) {
|
||||
return &tg.HelpInviteText{Message: "Join me on " + branding.ProductName + "."}, nil
|
||||
return &tg.HelpInviteText{Message: "Join me on " + branding.ProductName() + "."}, nil
|
||||
})
|
||||
registerRPC[*tg.HelpSaveAppLogRequest](d, tlprofile.SemanticMethodHelpSaveAppLog, func(ctx context.Context, _ *tg.HelpSaveAppLogRequest) (any, error) {
|
||||
return r.onHelpSaveAppLog(ctx)
|
||||
|
|
@ -178,7 +178,7 @@ func (r *Router) onHelpDismissSuggestion(ctx context.Context, req *tg.HelpDismis
|
|||
// dead payment URLs. All six TL fields are mandatory.
|
||||
func (r *Router) onHelpGetPremiumPromo(ctx context.Context) (*tg.HelpPremiumPromo, error) {
|
||||
promo := &tg.HelpPremiumPromo{
|
||||
StatusText: branding.PremiumName + " is not active on this account.",
|
||||
StatusText: branding.PremiumName() + " is not active on this account.",
|
||||
StatusEntities: []tg.MessageEntityClass{},
|
||||
VideoSections: []string{},
|
||||
Videos: []tg.DocumentClass{},
|
||||
|
|
@ -207,7 +207,7 @@ func (r *Router) onHelpGetPremiumPromo(ctx context.Context) (*tg.HelpPremiumProm
|
|||
}
|
||||
if u.PremiumActiveAt(r.clock.Now().Unix()) {
|
||||
until := time.Unix(int64(u.PremiumUntil), 0)
|
||||
promo.StatusText = branding.PremiumName + " is active until " + until.Format("2006-01-02") + "."
|
||||
promo.StatusText = branding.PremiumName() + " is active until " + until.Format("2006-01-02") + "."
|
||||
}
|
||||
if r.deps.PremiumPromo != nil {
|
||||
catalog, found, err := r.deps.PremiumPromo.PremiumPromo(ctx)
|
||||
|
|
|
|||
|
|
@ -12,12 +12,48 @@ import (
|
|||
"github.com/iamxvbaba/td/tlprofile"
|
||||
compatandroid "telesrv/internal/compat/android"
|
||||
"telesrv/internal/observability/dbtrace"
|
||||
"telesrv/internal/rpcresult"
|
||||
)
|
||||
|
||||
type layerWrappersAppliedKey struct{}
|
||||
type layerAdmissionSequenceKey struct{}
|
||||
type layerRPCProfileEvidenceFreshKey struct{}
|
||||
|
||||
type immutableLayerRPCReplaySource struct {
|
||||
call tlprofile.Call
|
||||
source rpcresult.ValueSource
|
||||
}
|
||||
|
||||
func (s *immutableLayerRPCReplaySource) EncodeInner(ctx context.Context, out *bin.Buffer) error {
|
||||
if s == nil || s.source == nil {
|
||||
return fmt.Errorf("immutable layer RPC replay source is unavailable")
|
||||
}
|
||||
value, err := s.source.Value(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.call.EncodeResult(value, out)
|
||||
}
|
||||
|
||||
func (s *immutableLayerRPCReplaySource) RetainedBytes() int {
|
||||
if s == nil || s.source == nil {
|
||||
return 0
|
||||
}
|
||||
return s.source.RetainedBytes() + 128
|
||||
}
|
||||
|
||||
type immutableLayerRPCResult struct {
|
||||
tlprofile.Result
|
||||
source rpcresult.ReplaySource
|
||||
}
|
||||
|
||||
func (r *immutableLayerRPCResult) ExactReplaySource() rpcresult.ReplaySource {
|
||||
if r == nil {
|
||||
return nil
|
||||
}
|
||||
return r.source
|
||||
}
|
||||
|
||||
// WithLayerRPCProfileEvidenceFresh records whether an admitted request's
|
||||
// explicit selector is inside MTProto's mutable msg_id freshness window. A
|
||||
// stale selector still owns its immutable request/result codec, but wrapper
|
||||
|
|
@ -278,7 +314,16 @@ func (r *Router) DispatchAdmitted(
|
|||
}
|
||||
dbBefore := dbtrace.SnapshotFromContext(ctx)
|
||||
start := time.Now()
|
||||
result, err := r.dispatchGeneratedSafely(ctx, method, request)
|
||||
dispatchCtx, replayCapture := rpcresult.WithCapture(ctx)
|
||||
result, err := r.dispatchGeneratedSafely(dispatchCtx, method, request)
|
||||
if err == nil && result != nil {
|
||||
if valueSource := replayCapture.Take(); valueSource != nil {
|
||||
result = &immutableLayerRPCResult{
|
||||
Result: result,
|
||||
source: &immutableLayerRPCReplaySource{call: call, source: valueSource},
|
||||
}
|
||||
}
|
||||
}
|
||||
dur := time.Since(start)
|
||||
dbDelta := dbtrace.SnapshotFromContext(ctx).Sub(dbBefore)
|
||||
fields := append([]zap.Field{
|
||||
|
|
|
|||
87
internal/rpc/messages_external_reply_test.go
Normal file
87
internal/rpc/messages_external_reply_test.go
Normal file
|
|
@ -0,0 +1,87 @@
|
|||
package rpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/iamxvbaba/td/tg"
|
||||
"github.com/iamxvbaba/td/tgerr"
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
func TestExternalPrivateReplySnapshotWireAndProtection(t *testing.T) {
|
||||
r, store, events, a, b := savedForwardFixture(t)
|
||||
ctx := WithUserID(context.Background(), a.ID)
|
||||
seed, err := store.SendPrivateText(ctx, domain.SendPrivateTextRequest{SenderUserID: a.ID, RecipientUserID: a.ID, RandomID: 1, Message: "a🌕 quote", Date: 1700000000, Media: &domain.MessageMedia{Kind: domain.MessageMediaKindContact, Contact: &domain.MessageContact{FirstName: "snapshot", PhoneNumber: "123"}}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, manual := range []bool{false, true} {
|
||||
reply := &tg.InputReplyToMessage{ReplyToMsgID: seed.SenderMessage.ID}
|
||||
reply.SetReplyToPeerID(&tg.InputPeerSelf{})
|
||||
random := int64(2)
|
||||
if manual {
|
||||
random = 3
|
||||
reply.SetQuoteText("quote")
|
||||
reply.SetQuoteOffset(4)
|
||||
}
|
||||
req := &tg.MessagesSendMessageRequest{Peer: &tg.InputPeerUser{UserID: b.ID, AccessHash: b.AccessHash}, Message: "external", RandomID: random}
|
||||
req.SetReplyTo(reply)
|
||||
out, err := r.onMessagesSendMessage(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
m := out.(*tg.Updates).Updates[1].(*tg.UpdateNewMessage).Message.(*tg.Message)
|
||||
h := m.ReplyTo.(*tg.MessageReplyHeader)
|
||||
if h.ReplyFrom.FromID.(*tg.PeerUser).UserID != a.ID || h.ReplyFrom.Date != seed.SenderMessage.Date || h.ReplyToMsgID != seed.SenderMessage.ID || h.Quote != manual {
|
||||
t.Fatalf("sender reply=%+v", h)
|
||||
}
|
||||
if manual {
|
||||
if h.QuoteText != "quote" || h.QuoteOffset != 4 {
|
||||
t.Fatal("manual quote")
|
||||
}
|
||||
} else if h.QuoteText != seed.SenderMessage.Body {
|
||||
t.Fatal("external preview text missing")
|
||||
}
|
||||
if h.ReplyMedia.(*tg.MessageMediaContact).FirstName != "snapshot" {
|
||||
t.Fatal("media snapshot")
|
||||
}
|
||||
ownerEvents := savedForwardEvents(t, events, b.ID)[0]
|
||||
recipient := ownerEvents[len(ownerEvents)-1].Message
|
||||
rh := tgMessageReplyHeader(recipient).(*tg.MessageReplyHeader)
|
||||
if _, set := rh.GetReplyToMsgID(); set {
|
||||
t.Fatal("recipient received sender-owned source ID")
|
||||
}
|
||||
if !reflect.DeepEqual(rh.ReplyFrom, h.ReplyFrom) {
|
||||
t.Fatal("external author differs by owner")
|
||||
}
|
||||
if manual {
|
||||
if _, err := store.DeleteMessages(ctx, domain.DeleteMessagesRequest{OwnerUserID: a.ID, IDs: []int{seed.SenderMessage.ID}}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if manual {
|
||||
before := savedForwardEvents(t, events, a.ID, b.ID)
|
||||
replayed, err := r.onMessagesSendMessage(ctx, req)
|
||||
if err != nil || !reflect.DeepEqual(out.(*tg.Updates).Updates, replayed.(*tg.Updates).Updates) || !reflect.DeepEqual(before, savedForwardEvents(t, events, a.ID, b.ID)) {
|
||||
t.Fatalf("exact replay=%v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
protected, err := store.SendPrivateText(ctx, domain.SendPrivateTextRequest{SenderUserID: a.ID, RecipientUserID: a.ID, RandomID: 10, Message: "protected", NoForwards: true, Date: 1700000001})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
reply := &tg.InputReplyToMessage{ReplyToMsgID: protected.SenderMessage.ID}
|
||||
reply.SetReplyToPeerID(&tg.InputPeerSelf{})
|
||||
req := &tg.MessagesSendMessageRequest{Peer: &tg.InputPeerUser{UserID: b.ID, AccessHash: b.AccessHash}, Message: "forbidden", RandomID: 11}
|
||||
req.SetReplyTo(reply)
|
||||
before := savedForwardEvents(t, events, a.ID, b.ID)
|
||||
if _, err := r.onMessagesSendMessage(ctx, req); !tgerr.Is(err, "CHAT_FORWARDS_RESTRICTED") {
|
||||
t.Fatalf("protected source=%v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(before, savedForwardEvents(t, events, a.ID, b.ID)) {
|
||||
t.Fatal("rejection wrote event")
|
||||
}
|
||||
}
|
||||
|
|
@ -514,7 +514,7 @@ func (r *Router) forwardFromPeerAndSources(ctx context.Context, userID int64, in
|
|||
}
|
||||
return fromPeer, sources, nil
|
||||
}
|
||||
fromPeer, err := r.checkedDomainPeerFromInputPeer(ctx, userID, input)
|
||||
fromPeer, err := r.checkedMessageReadPeer(ctx, userID, input, false)
|
||||
return fromPeer, nil, err
|
||||
}
|
||||
|
||||
|
|
@ -682,7 +682,9 @@ func (r *Router) forwardSourcesFromPrivateMessages(ctx context.Context, userID i
|
|||
if fromPeer.Type != domain.PeerTypeUser || fromPeer.ID == 0 {
|
||||
return nil, domain.ErrMessageIDInvalid
|
||||
}
|
||||
if svc, ok := r.deps.Messages.(PrivateNoForwardsService); ok {
|
||||
// Saved Messages has no two-user protection state. Its individual source
|
||||
// messages still pass the noforwards check below, including inferred peers.
|
||||
if svc, ok := r.deps.Messages.(PrivateNoForwardsService); ok && fromPeer.ID != userID {
|
||||
state, err := svc.GetPrivateNoForwards(ctx, userID, fromPeer.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
|
@ -771,6 +773,8 @@ func messageForwardErr(err error) error {
|
|||
return messageIDInvalidErr()
|
||||
case errors.Is(err, domain.ErrChatForwardsRestricted):
|
||||
return chatForwardsRestrictedErr()
|
||||
case errors.Is(err, domain.ErrQuoteTextInvalid):
|
||||
return quoteTextInvalidErr()
|
||||
case errors.Is(err, domain.ErrReplyMessageIDInvalid):
|
||||
return replyMessageIDInvalidErr()
|
||||
case errors.Is(err, domain.ErrMessageRandomIDDuplicate):
|
||||
|
|
|
|||
|
|
@ -806,10 +806,13 @@ func tgGlobalSearchMessages(viewerUserID int64, limit int, private domain.Messag
|
|||
return &tg.MessagesMessages{Messages: messages, Chats: chats, Users: users}
|
||||
}
|
||||
|
||||
func (r *Router) messageFilterFromHistoryRequest(userID int64, req *tg.MessagesGetHistoryRequest) (domain.MessageFilter, bool) {
|
||||
peer, ok := r.domainPeerFromInputPeer(userID, req.Peer)
|
||||
if !ok {
|
||||
return domain.MessageFilter{}, false
|
||||
func (r *Router) messageFilterFromHistoryRequest(ctx context.Context, userID int64, req *tg.MessagesGetHistoryRequest) (domain.MessageFilter, error) {
|
||||
if err := validateMessageReadBounds(req.Limit, req.OffsetID, req.MaxID, req.MinID); err != nil {
|
||||
return domain.MessageFilter{}, err
|
||||
}
|
||||
peer, err := r.checkedMessageReadPeer(ctx, userID, req.Peer, false)
|
||||
if err != nil {
|
||||
return domain.MessageFilter{}, err
|
||||
}
|
||||
limit := req.Limit
|
||||
if limit > 50 {
|
||||
|
|
@ -825,10 +828,13 @@ func (r *Router) messageFilterFromHistoryRequest(userID int64, req *tg.MessagesG
|
|||
MaxID: req.MaxID,
|
||||
MinID: req.MinID,
|
||||
Hash: req.Hash,
|
||||
}, true
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *Router) messageFilterFromSearchRequest(ctx context.Context, userID int64, req *tg.MessagesSearchRequest) (domain.MessageFilter, error) {
|
||||
if err := validateMessageReadBounds(req.Limit, req.OffsetID, req.MaxID, req.MinID); err != nil {
|
||||
return domain.MessageFilter{}, err
|
||||
}
|
||||
limit := req.Limit
|
||||
if limit > 500 {
|
||||
limit = 500
|
||||
|
|
@ -840,6 +846,7 @@ func (r *Router) messageFilterFromSearchRequest(ctx context.Context, userID int6
|
|||
MaxDate: req.MaxDate,
|
||||
AddOffset: domain.ClampMessageHistoryAddOffset(req.AddOffset),
|
||||
Limit: limit,
|
||||
CountOnly: req.Limit == 0,
|
||||
MaxID: req.MaxID,
|
||||
MinID: req.MinID,
|
||||
Hash: req.Hash,
|
||||
|
|
@ -850,10 +857,21 @@ func (r *Router) messageFilterFromSearchRequest(ctx context.Context, userID int6
|
|||
filter.PhoneCallsOnly = true
|
||||
filter.MissedPhoneCallsOnly = phoneCalls.Missed
|
||||
}
|
||||
if peer, ok := r.domainPeerFromInputPeer(userID, req.Peer); ok {
|
||||
if empty, ok := req.Peer.(*tg.InputPeerEmpty); !ok || empty == nil {
|
||||
peer, err := r.checkedMessageReadPeer(ctx, userID, req.Peer, false)
|
||||
if err != nil {
|
||||
return domain.MessageFilter{}, err
|
||||
}
|
||||
filter.HasPeer = true
|
||||
filter.Peer = peer
|
||||
}
|
||||
if req.FromID != nil {
|
||||
from, err := r.checkedMessageReadPeer(ctx, userID, req.FromID, true)
|
||||
if err != nil {
|
||||
return domain.MessageFilter{}, err
|
||||
}
|
||||
filter.SenderUserID = from.ID
|
||||
}
|
||||
savedReactions, hasSavedReactions := req.GetSavedReaction()
|
||||
// An empty optional vector carries no reaction-filtering semantics. Some TL
|
||||
// clients emit flags.3 with a zero-length vector on ordinary peer searches.
|
||||
|
|
|
|||
|
|
@ -580,7 +580,7 @@ func TestMessagesGetHistoryReturnsStoredMessages(t *testing.T) {
|
|||
Hash: 99,
|
||||
},
|
||||
}
|
||||
r := New(Config{}, Deps{Messages: messages}, zaptest.NewLogger(t), clock.System)
|
||||
r := New(Config{}, Deps{Messages: messages, Users: mapUsersService{users: map[int64]domain.User{domain.OfficialSystemUserID: domain.OfficialSystemUser()}}}, zaptest.NewLogger(t), clock.System)
|
||||
req := &tg.MessagesGetHistoryRequest{
|
||||
Peer: &tg.InputPeerUser{UserID: domain.OfficialSystemUserID, AccessHash: domain.OfficialSystemUser().AccessHash},
|
||||
Limit: 20,
|
||||
|
|
|
|||
|
|
@ -206,7 +206,7 @@ func TestSavedReactionTagHashMatchesClientShape(t *testing.T) {
|
|||
|
||||
func TestMessageFilterFromSearchRequestParsesSavedTagsAndPeer(t *testing.T) {
|
||||
const userID = int64(1000000001)
|
||||
r := New(Config{}, Deps{}, zaptest.NewLogger(t), clock.System)
|
||||
r := New(Config{}, Deps{Users: mapUsersService{users: map[int64]domain.User{userID + 1: {ID: userID + 1, AccessHash: 1}}}}, zaptest.NewLogger(t), clock.System)
|
||||
req := &tg.MessagesSearchRequest{
|
||||
Peer: &tg.InputPeerSelf{},
|
||||
Q: "needle",
|
||||
|
|
|
|||
83
internal/rpc/messages_read_filter.go
Normal file
83
internal/rpc/messages_read_filter.go
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
package rpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/iamxvbaba/td/tg"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
type privateSendBaseUserResolver interface {
|
||||
BaseUsersByIDs(ctx context.Context, userIDs []int64) ([]domain.User, error)
|
||||
}
|
||||
|
||||
// Read endpoints retain their existing positive page caps and signed
|
||||
// add_offset semantics. Invalid IDs must not become an unbounded query.
|
||||
func validateMessageReadBounds(limit, offsetID, maxID, minID int) error {
|
||||
if limit < 0 {
|
||||
return limitInvalidErr()
|
||||
}
|
||||
for _, id := range [...]int{offsetID, maxID, minID} {
|
||||
if id < 0 || id > domain.MaxMessageBoxID {
|
||||
return msgIDInvalidErr()
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// checkedMessageReadPeer never treats an invalid reference as global scope.
|
||||
// Only the search parser admits InputPeerEmpty, explicitly. Base identity
|
||||
// lookup avoids loading viewer-specific user projections for count requests.
|
||||
func (r *Router) checkedMessageReadPeer(ctx context.Context, userID int64, input tg.InputPeerClass, sender bool) (domain.Peer, error) {
|
||||
invalid := peerIDInvalidErr
|
||||
if sender {
|
||||
invalid = fromPeerInvalidErr
|
||||
}
|
||||
if inputPeerClassNil(input) {
|
||||
return domain.Peer{}, invalid()
|
||||
}
|
||||
switch peer := input.(type) {
|
||||
case *tg.InputPeerSelf:
|
||||
if userID <= 0 {
|
||||
return domain.Peer{}, invalid()
|
||||
}
|
||||
return domain.Peer{Type: domain.PeerTypeUser, ID: userID}, nil
|
||||
case *tg.InputPeerUser:
|
||||
if peer.UserID <= 0 {
|
||||
return domain.Peer{}, invalid()
|
||||
}
|
||||
if r.deps.Users == nil {
|
||||
return domain.Peer{}, internalErr()
|
||||
}
|
||||
if resolver, ok := r.deps.Users.(privateSendBaseUserResolver); ok {
|
||||
users, err := resolver.BaseUsersByIDs(ctx, []int64{peer.UserID})
|
||||
if err != nil {
|
||||
return domain.Peer{}, internalErr()
|
||||
}
|
||||
for _, user := range users {
|
||||
if user.ID == peer.UserID && user.AccessHash == peer.AccessHash {
|
||||
return domain.Peer{Type: domain.PeerTypeUser, ID: peer.UserID}, nil
|
||||
}
|
||||
}
|
||||
return domain.Peer{}, invalid()
|
||||
}
|
||||
user, found, err := r.deps.Users.ByID(ctx, userID, peer.UserID)
|
||||
if err != nil {
|
||||
return domain.Peer{}, internalErr()
|
||||
}
|
||||
if found && user.AccessHash == peer.AccessHash {
|
||||
return domain.Peer{Type: domain.PeerTypeUser, ID: peer.UserID}, nil
|
||||
}
|
||||
return domain.Peer{}, invalid()
|
||||
default:
|
||||
if sender {
|
||||
return domain.Peer{}, invalid()
|
||||
}
|
||||
resolved, ok := r.domainPeerFromInputPeer(userID, input)
|
||||
if !ok || resolved.ID <= 0 {
|
||||
return domain.Peer{}, invalid()
|
||||
}
|
||||
return r.checkedDomainPeerFromInputPeer(ctx, userID, input)
|
||||
}
|
||||
}
|
||||
83
internal/rpc/messages_read_filter_test.go
Normal file
83
internal/rpc/messages_read_filter_test.go
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
package rpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/iamxvbaba/td/clock"
|
||||
"github.com/iamxvbaba/td/tg"
|
||||
"github.com/iamxvbaba/td/tgerr"
|
||||
"go.uber.org/zap/zaptest"
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
type failingReadUsers struct{ mapUsersService }
|
||||
|
||||
func (failingReadUsers) BaseUsersByIDs(context.Context, []int64) ([]domain.User, error) {
|
||||
return nil, errors.New("identity unavailable")
|
||||
}
|
||||
|
||||
func TestMessageReadInputScope(t *testing.T) {
|
||||
users := &countingMapUsersService{mapUsersService: mapUsersService{users: map[int64]domain.User{22: {ID: 22, AccessHash: 77}}}}
|
||||
r := New(Config{}, Deps{Users: users}, zaptest.NewLogger(t), clock.System)
|
||||
ctx := context.Background()
|
||||
invalid := []tg.InputPeerClass{nil, (*tg.InputPeerUser)(nil), &tg.InputPeerEmpty{}, &tg.InputPeerUser{UserID: 0}, &tg.InputPeerUser{UserID: -22}, &tg.InputPeerUser{UserID: 22, AccessHash: 0}, &tg.InputPeerUser{UserID: 22, AccessHash: 78}, &tg.InputPeerUser{UserID: 23, AccessHash: 77}, &tg.InputPeerUserFromMessage{Peer: &tg.InputPeerSelf{}, MsgID: 1, UserID: 22}}
|
||||
for i, p := range invalid {
|
||||
for _, sender := range []bool{false, true} {
|
||||
_, err := r.checkedMessageReadPeer(ctx, 11, p, sender)
|
||||
want := "PEER_ID_INVALID"
|
||||
if sender {
|
||||
want = "FROM_PEER_INVALID"
|
||||
}
|
||||
if !tgerr.Is(err, want) {
|
||||
t.Fatalf("invalid[%d], sender=%t err=%v", i, sender, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, p := range []tg.InputPeerClass{&tg.InputPeerEmpty{}, &tg.InputPeerSelf{}, &tg.InputPeerUser{UserID: 22, AccessHash: 77}} {
|
||||
req := &tg.MessagesSearchRequest{Peer: p, FromID: &tg.InputPeerUser{UserID: 22, AccessHash: 77}, Filter: &tg.InputMessagesFilterEmpty{}}
|
||||
f, err := r.messageFilterFromSearchRequest(ctx, 11, req)
|
||||
_, global := p.(*tg.InputPeerEmpty)
|
||||
if err != nil || f.HasPeer == global || f.SenderUserID != 22 || !f.CountOnly {
|
||||
t.Fatalf("scope=%T filter=%+v err=%v", p, f, err)
|
||||
}
|
||||
}
|
||||
if users.byIDCalls != 0 || users.byIDsCalls != 0 || users.selfCalls != 0 || users.baseByIDsCalls == 0 {
|
||||
t.Fatalf("read validation projected users: %+v", users)
|
||||
}
|
||||
for _, deps := range []Deps{{}, {Users: failingReadUsers{}}} {
|
||||
broken := New(Config{}, deps, zaptest.NewLogger(t), clock.System)
|
||||
_, err := broken.checkedMessageReadPeer(ctx, 11, &tg.InputPeerUser{UserID: 22, AccessHash: 77}, false)
|
||||
if !tgerr.Is(err, "INTERNAL_SERVER_ERROR") {
|
||||
t.Fatalf("missing identity boundary err=%v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessageReadBoundsAndCaps(t *testing.T) {
|
||||
r := New(Config{}, Deps{}, zaptest.NewLogger(t), clock.System)
|
||||
ctx := context.Background()
|
||||
for _, tt := range []struct {
|
||||
limit, offset, max, min int
|
||||
want string
|
||||
}{{-1, 0, 0, 0, "LIMIT_INVALID"}, {0, -1, 0, 0, "MSG_ID_INVALID"}, {0, 0, -1, 0, "MSG_ID_INVALID"}, {0, 0, 0, -1, "MSG_ID_INVALID"}, {0, domain.MaxMessageBoxID + 1, 0, 0, "MSG_ID_INVALID"}, {0, 0, domain.MaxMessageBoxID + 1, 0, "MSG_ID_INVALID"}} {
|
||||
_, se := r.messageFilterFromSearchRequest(ctx, 11, &tg.MessagesSearchRequest{Peer: &tg.InputPeerSelf{}, Limit: tt.limit, OffsetID: tt.offset, MaxID: tt.max, MinID: tt.min})
|
||||
_, he := r.messageFilterFromHistoryRequest(ctx, 11, &tg.MessagesGetHistoryRequest{Peer: &tg.InputPeerSelf{}, Limit: tt.limit, OffsetID: tt.offset, MaxID: tt.max, MinID: tt.min})
|
||||
if !tgerr.Is(se, tt.want) || !tgerr.Is(he, tt.want) {
|
||||
t.Fatalf("bounds %+v search=%v history=%v", tt, se, he)
|
||||
}
|
||||
}
|
||||
s, err := r.messageFilterFromSearchRequest(ctx, 11, &tg.MessagesSearchRequest{Peer: &tg.InputPeerSelf{}, Limit: 501, AddOffset: -2})
|
||||
if err != nil || s.Limit != 500 || s.AddOffset != -2 || s.CountOnly {
|
||||
t.Fatalf("search %+v %v", s, err)
|
||||
}
|
||||
h, err := r.messageFilterFromHistoryRequest(ctx, 11, &tg.MessagesGetHistoryRequest{Peer: &tg.InputPeerSelf{}, Limit: 501, AddOffset: -2})
|
||||
if err != nil || h.Limit != 50 || h.AddOffset != -2 || h.CountOnly {
|
||||
t.Fatalf("history %+v %v", h, err)
|
||||
}
|
||||
h, err = r.messageFilterFromHistoryRequest(ctx, 11, &tg.MessagesGetHistoryRequest{Peer: &tg.InputPeerSelf{}})
|
||||
if err != nil || h.CountOnly {
|
||||
t.Fatalf("default history %+v %v", h, err)
|
||||
}
|
||||
}
|
||||
|
|
@ -611,9 +611,9 @@ func (r *Router) registerMessages(d *tlprofile.Dispatcher) {
|
|||
if err != nil {
|
||||
return nil, internalErr()
|
||||
}
|
||||
filter, ok := r.messageFilterFromHistoryRequest(userID, req)
|
||||
if !ok {
|
||||
return messagesNotModifiedOrEmpty(req.Hash), nil
|
||||
filter, err := r.messageFilterFromHistoryRequest(ctx, userID, req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if filter.Peer.Type == domain.PeerTypeChannel {
|
||||
if r.deps.Channels == nil {
|
||||
|
|
|
|||
259
internal/rpc/messages_saved_forward_boundary_test.go
Normal file
259
internal/rpc/messages_saved_forward_boundary_test.go
Normal file
|
|
@ -0,0 +1,259 @@
|
|||
package rpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/iamxvbaba/td/clock"
|
||||
"github.com/iamxvbaba/td/tg"
|
||||
"github.com/iamxvbaba/td/tgerr"
|
||||
"go.uber.org/zap/zaptest"
|
||||
appdialogs "telesrv/internal/app/dialogs"
|
||||
appmessages "telesrv/internal/app/messages"
|
||||
appusers "telesrv/internal/app/users"
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/store/memory"
|
||||
)
|
||||
|
||||
func savedForwardFixture(t *testing.T) (*Router, *memory.MessageStore, *memory.UpdateEventStore, domain.User, domain.User) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
a, err := users.Create(ctx, domain.User{AccessHash: 51, Phone: "15550009501", FirstName: "A"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, err := users.Create(ctx, domain.User{AccessHash: 52, Phone: "15550009502", FirstName: "B"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dialogs := memory.NewDialogStore()
|
||||
messages := memory.NewMessageStore(dialogs)
|
||||
events := memory.NewUpdateEventStore()
|
||||
messages.AttachUpdateEventStore(events)
|
||||
r := New(Config{}, Deps{Users: appusers.NewService(users), Dialogs: appdialogs.NewService(dialogs), Messages: appmessages.NewService(messages, dialogs)}, zaptest.NewLogger(t), clock.System)
|
||||
return r, messages, events, a, b
|
||||
}
|
||||
|
||||
func savedForwardEvents(t *testing.T, events *memory.UpdateEventStore, owners ...int64) [][]domain.UpdateEvent {
|
||||
t.Helper()
|
||||
result := make([][]domain.UpdateEvent, len(owners))
|
||||
for i, owner := range owners {
|
||||
var err error
|
||||
result[i], err = events.ListAfter(context.Background(), owner, 0, 100)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func TestSavedForwardSourceScopeProtectionAndDeletedReplay(t *testing.T) {
|
||||
for _, scope := range []string{"self", "explicit-user", "inferred"} {
|
||||
for _, saved := range []bool{false, true} {
|
||||
for _, protected := range []bool{false, true} {
|
||||
t.Run(fmt.Sprintf("%s/saved=%v/protected=%v", scope, saved, protected), func(t *testing.T) {
|
||||
r, store, events, a, b := savedForwardFixture(t)
|
||||
ctx := WithUserID(context.Background(), a.ID)
|
||||
seed, err := store.SendPrivateText(ctx, domain.SendPrivateTextRequest{SenderUserID: a.ID, RecipientUserID: a.ID, RandomID: 1, Message: "saved 🌕 quote", NoForwards: protected, Date: 1700000000})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
to := tg.InputPeerClass(&tg.InputPeerSelf{})
|
||||
targetSender := a.ID
|
||||
if !saved {
|
||||
to = &tg.InputPeerUser{UserID: b.ID, AccessHash: b.AccessHash}
|
||||
targetSender = b.ID
|
||||
}
|
||||
target, err := store.SendPrivateText(ctx, domain.SendPrivateTextRequest{SenderUserID: targetSender, RecipientUserID: a.ID, RandomID: 2, Message: "target", Date: 1700000001})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
from := tg.InputPeerClass(&tg.InputPeerSelf{})
|
||||
if scope == "explicit-user" {
|
||||
from = &tg.InputPeerUser{UserID: a.ID, AccessHash: a.AccessHash}
|
||||
}
|
||||
if scope == "inferred" {
|
||||
from = &tg.InputPeerEmpty{}
|
||||
}
|
||||
req := &tg.MessagesForwardMessagesRequest{FromPeer: from, ToPeer: to, ID: []int{seed.SenderMessage.ID}, RandomID: []int64{3}}
|
||||
reply := &tg.InputReplyToMessage{ReplyToMsgID: target.RecipientMessage.ID}
|
||||
reply.SetQuoteText("target")
|
||||
req.SetReplyTo(reply)
|
||||
before := savedForwardEvents(t, events, a.ID, b.ID)
|
||||
out, err := r.onMessagesForwardMessages(ctx, req)
|
||||
if protected {
|
||||
if !tgerr.Is(err, "CHAT_FORWARDS_RESTRICTED") || !reflect.DeepEqual(before, savedForwardEvents(t, events, a.ID, b.ID)) {
|
||||
t.Fatalf("protected Saved source err=%v or wrote events", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
full := out.(*tg.Updates)
|
||||
msg := full.Updates[1].(*tg.UpdateNewMessage).Message.(*tg.Message)
|
||||
if msg.Message != seed.SenderMessage.Body || msg.FwdFrom.FromID.(*tg.PeerUser).UserID != a.ID || msg.FwdFrom.Date != seed.SenderMessage.Date || msg.ReplyTo.(*tg.MessageReplyHeader).ReplyToMsgID != target.RecipientMessage.ID {
|
||||
t.Fatalf("forward metadata: %+v", msg)
|
||||
}
|
||||
after := savedForwardEvents(t, events, a.ID, b.ID)
|
||||
for i := range after {
|
||||
want := 1
|
||||
if saved && i == 1 {
|
||||
want = 0
|
||||
}
|
||||
if len(after[i])-len(before[i]) != want {
|
||||
t.Fatal("new forward event cardinality")
|
||||
}
|
||||
if want == 1 && after[i][len(after[i])-1].PtsCount != 1 {
|
||||
t.Fatal("new forward PTS count")
|
||||
}
|
||||
}
|
||||
if !saved {
|
||||
received := after[1][len(after[1])-1].Message
|
||||
if received.ReplyTo == nil || received.ReplyTo.MessageID != target.SenderMessage.ID {
|
||||
t.Fatal("recipient reply not mapped")
|
||||
}
|
||||
}
|
||||
if _, err := store.DeleteMessages(ctx, domain.DeleteMessagesRequest{OwnerUserID: a.ID, IDs: []int{seed.SenderMessage.ID}, Date: 1700000002}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
beforeReplay := savedForwardEvents(t, events, a.ID, b.ID)
|
||||
replay, err := r.onMessagesForwardMessages(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !reflect.DeepEqual(full.Updates, replay.(*tg.Updates).Updates) || !reflect.DeepEqual(beforeReplay, savedForwardEvents(t, events, a.ID, b.ID)) {
|
||||
t.Fatal("replay changed message or appended event")
|
||||
}
|
||||
fresh := *req
|
||||
fresh.RandomID = []int64{4}
|
||||
if _, err := r.onMessagesForwardMessages(ctx, &fresh); !tgerr.Is(err, "MESSAGE_ID_INVALID") {
|
||||
t.Fatalf("fresh deleted source err=%v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(beforeReplay, savedForwardEvents(t, events, a.ID, b.ID)) {
|
||||
t.Fatal("rejected fresh request appended event")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The injected boundary fails before the second send enters the real store.
|
||||
// All successful sends and replay lookups retain the normal service/store path.
|
||||
type failSecondForward struct {
|
||||
*appmessages.Service
|
||||
failRandom int64
|
||||
calls []int64
|
||||
}
|
||||
|
||||
func (s *failSecondForward) SendPrivateText(ctx context.Context, user int64, req domain.SendPrivateTextRequest) (domain.SendPrivateTextResult, error) {
|
||||
s.calls = append(s.calls, req.RandomID)
|
||||
if req.RandomID == s.failRandom {
|
||||
return domain.SendPrivateTextResult{}, errors.New("injected second-send failure")
|
||||
}
|
||||
return s.Service.SendPrivateText(ctx, user, req)
|
||||
}
|
||||
|
||||
func TestSavedForwardPartialCommitRetryAfterCommittedSourceDeleted(t *testing.T) {
|
||||
r, store, events, a, b := savedForwardFixture(t)
|
||||
ctx := WithUserID(context.Background(), a.ID)
|
||||
ids := []int{}
|
||||
for i := int64(1); i <= 2; i++ {
|
||||
source, err := store.SendPrivateText(ctx, domain.SendPrivateTextRequest{SenderUserID: a.ID, RecipientUserID: a.ID, RandomID: i, Message: fmt.Sprintf("source-%d", i), Date: 1700000000})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ids = append(ids, source.SenderMessage.ID)
|
||||
}
|
||||
fault := &failSecondForward{Service: r.deps.Messages.(*appmessages.Service), failRandom: 12}
|
||||
r.deps.Messages = fault
|
||||
req := &tg.MessagesForwardMessagesRequest{FromPeer: &tg.InputPeerSelf{}, ToPeer: &tg.InputPeerUser{UserID: b.ID, AccessHash: b.AccessHash}, ID: ids, RandomID: []int64{11, 12}}
|
||||
before := savedForwardEvents(t, events, a.ID, b.ID)
|
||||
if _, err := r.onMessagesForwardMessages(ctx, req); !tgerr.Is(err, "INTERNAL_SERVER_ERROR") {
|
||||
t.Fatalf("partial send err=%v", err)
|
||||
}
|
||||
after := savedForwardEvents(t, events, a.ID, b.ID)
|
||||
for i := range after {
|
||||
if len(after[i])-len(before[i]) != 1 {
|
||||
t.Fatal("first item must commit before second fails")
|
||||
}
|
||||
}
|
||||
first := after[0][len(after[0])-1].Message
|
||||
if _, err := store.DeleteMessages(ctx, domain.DeleteMessagesRequest{OwnerUserID: a.ID, IDs: ids[:1], Date: 1700000001}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
before = savedForwardEvents(t, events, a.ID, b.ID)
|
||||
fault.failRandom = 0
|
||||
out, err := r.onMessagesForwardMessages(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !reflect.DeepEqual(fault.calls, []int64{11, 12, 12}) {
|
||||
t.Fatalf("send calls=%v; committed item must not be re-sent", fault.calls)
|
||||
}
|
||||
full := out.(*tg.Updates)
|
||||
if full.Updates[0].(*tg.UpdateMessageID).ID != first.ID || len(full.Updates) != 4 {
|
||||
t.Fatal("partial retry lost first committed ID")
|
||||
}
|
||||
after = savedForwardEvents(t, events, a.ID, b.ID)
|
||||
for i := range after {
|
||||
if len(after[i])-len(before[i]) != 1 {
|
||||
t.Fatal("retry must only append second item")
|
||||
}
|
||||
}
|
||||
if _, err := r.onMessagesForwardMessages(ctx, req); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !reflect.DeepEqual(after, savedForwardEvents(t, events, a.ID, b.ID)) {
|
||||
t.Fatal("full replay appended events")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReplyAndForwardExplicitSourceCredentials(t *testing.T) {
|
||||
for _, method := range []string{"reply", "forward"} {
|
||||
for _, wrong := range []tg.InputPeerClass{&tg.InputPeerUser{UserID: 22, AccessHash: 78}, &tg.InputPeerUser{UserID: 23, AccessHash: 77}, &tg.InputPeerEmpty{}} {
|
||||
t.Run(fmt.Sprintf("%s/%#v", method, wrong), func(t *testing.T) {
|
||||
messages := &captureMessages{}
|
||||
r := New(Config{}, Deps{Messages: messages, Users: mapUsersService{users: map[int64]domain.User{22: {ID: 22, AccessHash: 77}}}}, zaptest.NewLogger(t), clock.System)
|
||||
ctx := WithUserID(context.Background(), 11)
|
||||
var err error
|
||||
if method == "reply" {
|
||||
req := &tg.MessagesSendMessageRequest{Peer: &tg.InputPeerSelf{}, Message: "reply", RandomID: 91}
|
||||
reply := &tg.InputReplyToMessage{ReplyToMsgID: 7}
|
||||
reply.SetReplyToPeerID(wrong)
|
||||
req.SetReplyTo(reply)
|
||||
_, err = r.onMessagesSendMessage(ctx, req)
|
||||
if !tgerr.Is(err, "REPLY_MESSAGE_ID_INVALID") {
|
||||
t.Fatalf("reply err=%v", err)
|
||||
}
|
||||
} else {
|
||||
_, err = r.onMessagesForwardMessages(ctx, &tg.MessagesForwardMessagesRequest{FromPeer: wrong, ToPeer: &tg.InputPeerSelf{}, ID: []int{7}, RandomID: []int64{92}})
|
||||
// Empty from_peer is permitted only when an owned source can be inferred.
|
||||
want := "PEER_ID_INVALID"
|
||||
if _, ok := wrong.(*tg.InputPeerEmpty); ok {
|
||||
want = "MESSAGE_ID_INVALID"
|
||||
}
|
||||
if !tgerr.Is(err, want) {
|
||||
t.Fatalf("forward err=%v want=%s", err, want)
|
||||
}
|
||||
}
|
||||
if messages.sendReq.RandomID != 0 {
|
||||
t.Fatal("invalid credentials reached write service")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
for _, deps := range []Deps{{}, {Users: failingReadUsers{}}} {
|
||||
r := New(Config{}, deps, zaptest.NewLogger(t), clock.System)
|
||||
reply := &tg.InputReplyToMessage{ReplyToMsgID: 1}
|
||||
reply.SetReplyToPeerID(&tg.InputPeerUser{UserID: 22, AccessHash: 77})
|
||||
if _, err := r.messageReplyFromInput(context.Background(), 11, domain.Peer{Type: domain.PeerTypeUser, ID: 11}, reply); !tgerr.Is(err, "INTERNAL_SERVER_ERROR") {
|
||||
t.Fatalf("identity unavailable err=%v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -8,6 +8,7 @@ import (
|
|||
"unicode/utf8"
|
||||
|
||||
"github.com/iamxvbaba/td/tg"
|
||||
"github.com/iamxvbaba/td/tgerr"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
|
@ -273,6 +274,10 @@ func messageSendErr(err error) error {
|
|||
return randomIDDuplicateErr()
|
||||
case errors.Is(err, domain.ErrMessageEmpty):
|
||||
return messageEmptyErr()
|
||||
case errors.Is(err, domain.ErrChatForwardsRestricted):
|
||||
return chatForwardsRestrictedErr()
|
||||
case errors.Is(err, domain.ErrQuoteTextInvalid):
|
||||
return quoteTextInvalidErr()
|
||||
default:
|
||||
return internalErr()
|
||||
}
|
||||
|
|
@ -330,8 +335,11 @@ func (r *Router) messageReplyFromInput(ctx context.Context, userID int64, peer d
|
|||
}
|
||||
replyPeer := peer
|
||||
if inputPeer, ok := reply.GetReplyToPeerID(); ok {
|
||||
parsed, err := r.checkedDomainPeerFromInputPeer(ctx, userID, inputPeer)
|
||||
parsed, err := r.checkedMessageReadPeer(ctx, userID, inputPeer, false)
|
||||
if err != nil {
|
||||
if tgerr.Is(err, "INTERNAL_SERVER_ERROR") {
|
||||
return nil, err
|
||||
}
|
||||
return nil, replyMessageIDInvalidErr()
|
||||
}
|
||||
replyPeer = parsed
|
||||
|
|
|
|||
|
|
@ -51,7 +51,9 @@ type presenceTracker struct {
|
|||
byUser map[int64]map[presenceSessionKey]domain.UserStatus
|
||||
// offlineTimers 跟踪每个 user 挂起的 offline 广播去抖定时器,使重新上线能取消它,
|
||||
// 避免断连风暴下 O(N) 个裸 time.AfterFunc 在 runtime timer heap 长期堆积。
|
||||
offlineTimers map[int64]*time.Timer
|
||||
offlineTimers map[int64]*time.Timer
|
||||
offlineRunning int64
|
||||
lastSeenDirectRunning int64
|
||||
}
|
||||
|
||||
func newPresenceTracker() *presenceTracker {
|
||||
|
|
@ -62,8 +64,8 @@ func newPresenceTracker() *presenceTracker {
|
|||
}
|
||||
}
|
||||
|
||||
// armOfflineTimer 安排(或重置)某 user 的 offline 广播去抖定时器。定时器触发时先把自己
|
||||
// 从 map 移除再执行 fire(fire 内部会查在线态做最终去抖),全程不持 p.mu 调用 fire。
|
||||
// armOfflineTimer replaces the user's pending grace timer. Stop cannot join
|
||||
// an already-started AfterFunc; the callback must still prove timer ownership.
|
||||
func (p *presenceTracker) armOfflineTimer(userID int64, d time.Duration, fire func()) {
|
||||
if p == nil || userID == 0 {
|
||||
return
|
||||
|
|
@ -75,12 +77,24 @@ func (p *presenceTracker) armOfflineTimer(userID int64, d time.Duration, fire fu
|
|||
if old := p.offlineTimers[userID]; old != nil {
|
||||
old.Stop()
|
||||
}
|
||||
p.offlineTimers[userID] = time.AfterFunc(d, func() {
|
||||
var timer *time.Timer
|
||||
timer = time.AfterFunc(d, func() {
|
||||
p.mu.Lock()
|
||||
if p.offlineTimers[userID] != timer {
|
||||
p.mu.Unlock()
|
||||
return
|
||||
}
|
||||
delete(p.offlineTimers, userID)
|
||||
p.offlineRunning++
|
||||
p.mu.Unlock()
|
||||
defer func() {
|
||||
p.mu.Lock()
|
||||
p.offlineRunning--
|
||||
p.mu.Unlock()
|
||||
}()
|
||||
fire()
|
||||
})
|
||||
p.offlineTimers[userID] = timer
|
||||
p.mu.Unlock()
|
||||
}
|
||||
|
||||
|
|
@ -419,8 +433,13 @@ func (r *Router) persistReservedLastSeenAsync(
|
|||
lastSeenAt int,
|
||||
) {
|
||||
bgCtx, cancel := r.presenceBackgroundContext(ctx, 10*time.Second)
|
||||
// Keep the direct overflow write visible after its parent callback returns.
|
||||
// The production batch normally handles this work; admission errors retain
|
||||
// their existing explicit log and authoritative write behavior.
|
||||
r.presence.changeDirectLastSeenRunning(1)
|
||||
go func() {
|
||||
defer cancel()
|
||||
defer r.presence.changeDirectLastSeenRunning(-1)
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
r.log.Error("Update user last seen panicked", zap.Int64("user_id", userID), zap.Any("panic", rec))
|
||||
|
|
@ -488,6 +507,12 @@ func (r *Router) presenceBackgroundContext(ctx context.Context, timeout time.Dur
|
|||
// 去抖宽限后在后台带超时执行——连接关闭发生在 serveConn 退出路径上,同步做
|
||||
// DB 写与逐 peer 查询会在断连风暴时放大 DB 压力并拖住 goroutine 退出。
|
||||
func (r *Router) SessionOffline(rawAuthKeyID [8]byte, sessionID, userID int64, lastForUser bool) {
|
||||
r.SessionOfflineAt(rawAuthKeyID, sessionID, userID, lastForUser, int(r.clock.Now().Unix()))
|
||||
}
|
||||
|
||||
// SessionOfflineAt preserves the observed disconnect time and is idempotent
|
||||
// for repeated physical-session departure reports.
|
||||
func (r *Router) SessionOfflineAt(rawAuthKeyID [8]byte, sessionID, userID int64, lastForUser bool, disconnectedAt int) {
|
||||
r.forgetClientSessionInfo(rawAuthKeyID, sessionID)
|
||||
if userID == 0 {
|
||||
return
|
||||
|
|
@ -497,7 +522,9 @@ func (r *Router) SessionOffline(rawAuthKeyID [8]byte, sessionID, userID int64, l
|
|||
if !lastForUser {
|
||||
return
|
||||
}
|
||||
disconnectedAt := int(r.clock.Now().Unix())
|
||||
if disconnectedAt <= 0 {
|
||||
disconnectedAt = int(r.clock.Now().Unix())
|
||||
}
|
||||
// 用可跟踪的定时器,重新上线时能取消(见 cancelOfflineTimer),避免断连风暴堆积。
|
||||
r.presence.armOfflineTimer(userID, offlineAnnounceGrace, func() {
|
||||
r.announceUserOfflineIfStillGone(rawAuthKeyID, sessionID, userID, disconnectedAt)
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import (
|
|||
"errors"
|
||||
"sort"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
|
@ -71,6 +72,7 @@ type presenceLastSeenBatchDispatcher struct {
|
|||
|
||||
gate sync.RWMutex
|
||||
accepting bool
|
||||
pending atomic.Int64
|
||||
}
|
||||
|
||||
func newPresenceLastSeenBatchDispatcher(
|
||||
|
|
@ -89,6 +91,7 @@ func newPresenceLastSeenBatchDispatcher(
|
|||
if metrics == nil {
|
||||
metrics = NopMetrics{}
|
||||
}
|
||||
metrics.PresenceLastSeenPending(0)
|
||||
return &presenceLastSeenBatchDispatcher{
|
||||
updater: updater,
|
||||
cfg: cfg,
|
||||
|
|
@ -108,18 +111,23 @@ func (d *presenceLastSeenBatchDispatcher) submit(update store.UserLastSeenUpdate
|
|||
if !d.accepting {
|
||||
return errPresenceLastSeenBatchStopped
|
||||
}
|
||||
d.metrics.PresenceLastSeenPending(1)
|
||||
d.changePending(1)
|
||||
select {
|
||||
case d.queue <- update:
|
||||
d.metrics.PresenceLastSeenSubmitted()
|
||||
return nil
|
||||
default:
|
||||
d.metrics.PresenceLastSeenPending(-1)
|
||||
d.changePending(-1)
|
||||
d.metrics.PresenceLastSeenOverflow()
|
||||
return errPresenceLastSeenBatchFull
|
||||
}
|
||||
}
|
||||
|
||||
func (d *presenceLastSeenBatchDispatcher) changePending(delta int) {
|
||||
d.pending.Add(int64(delta))
|
||||
d.metrics.PresenceLastSeenPending(delta)
|
||||
}
|
||||
|
||||
func (d *presenceLastSeenBatchDispatcher) stopAccepting() {
|
||||
if d == nil {
|
||||
return
|
||||
|
|
@ -223,7 +231,7 @@ func (d *presenceLastSeenBatchDispatcher) executeWithRetry(
|
|||
cancel()
|
||||
d.metrics.PresenceLastSeenBatch(len(updates), time.Since(started), err)
|
||||
if err == nil {
|
||||
d.metrics.PresenceLastSeenPending(-rawCount)
|
||||
d.changePending(-rawCount)
|
||||
return true
|
||||
}
|
||||
if attempt == 1 || attempt&(attempt-1) == 0 {
|
||||
|
|
@ -286,6 +294,6 @@ func (d *presenceLastSeenBatchDispatcher) reportDrainDropped(count int) {
|
|||
return
|
||||
}
|
||||
d.metrics.PresenceLastSeenDrainDropped(count)
|
||||
d.metrics.PresenceLastSeenPending(-count)
|
||||
d.changePending(-count)
|
||||
d.log.Error("presence last-seen shutdown drain expired", zap.Int("updates", count))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ package rpc
|
|||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
|
@ -10,9 +12,27 @@ import (
|
|||
|
||||
"go.uber.org/zap/zaptest"
|
||||
|
||||
obsmetrics "telesrv/internal/observability/metrics"
|
||||
"telesrv/internal/store"
|
||||
)
|
||||
|
||||
func TestPresenceLastSeenBatchExportsIdlePending(t *testing.T) {
|
||||
registry := obsmetrics.New()
|
||||
updater := &capturePresenceLastSeenUpdater{}
|
||||
d := newPresenceLastSeenBatchDispatcher(updater, presenceLastSeenBatchConfig{}, zaptest.NewLogger(t), registry)
|
||||
if d == nil {
|
||||
t.Fatal("presence owner was not constructed")
|
||||
}
|
||||
recorder := httptest.NewRecorder()
|
||||
registry.ServeHTTP(recorder, httptest.NewRequest("GET", "/metrics", nil))
|
||||
if !strings.Contains(recorder.Body.String(), "telesrv_presence_last_seen_pending 0\n") || len(updater.snapshot()) != 0 {
|
||||
t.Fatalf("idle owner must expose zero pending without a last-seen write: %s", recorder.Body.String())
|
||||
}
|
||||
if strings.Contains(recorder.Body.String(), "telesrv_presence_last_seen_submitted_total") {
|
||||
t.Fatal("idle registration fabricated a last-seen submission")
|
||||
}
|
||||
}
|
||||
|
||||
type capturePresenceLastSeenUpdater struct {
|
||||
mu sync.Mutex
|
||||
calls [][]store.UserLastSeenUpdate
|
||||
|
|
|
|||
39
internal/rpc/presence_timer_generation_test.go
Normal file
39
internal/rpc/presence_timer_generation_test.go
Normal file
|
|
@ -0,0 +1,39 @@
|
|||
package rpc
|
||||
|
||||
import (
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestPresenceOfflineExpiredCallbackCannotEraseReplacement(t *testing.T) {
|
||||
p := newPresenceTracker()
|
||||
var fired atomic.Int64
|
||||
p.armOfflineTimer(7, 10*time.Millisecond, func() { fired.Add(1) })
|
||||
p.mu.Lock()
|
||||
old := p.offlineTimers[7]
|
||||
// The real AfterFunc has expired and is waiting for the tracker lock.
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
if old.Stop() {
|
||||
p.mu.Unlock()
|
||||
t.Fatal("old callback has not started")
|
||||
}
|
||||
// Reproduce the replacement critical section before the old callback
|
||||
// reacquires the lock. No scheduler timing or production test hook.
|
||||
replacement := time.AfterFunc(time.Hour, func() { fired.Add(100) })
|
||||
defer replacement.Stop()
|
||||
p.offlineTimers[7] = replacement
|
||||
p.mu.Unlock()
|
||||
// Synchronize on the callback lock acquisition without holding it.
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
p.mu.RLock()
|
||||
current := p.offlineTimers[7]
|
||||
p.mu.RUnlock()
|
||||
if current != replacement || fired.Load() != 0 {
|
||||
t.Fatalf("stale callback: replacement retained=%v, business callbacks=%d", current == replacement, fired.Load())
|
||||
}
|
||||
p.cancelOfflineTimer(7)
|
||||
if replacement.Stop() {
|
||||
t.Fatal("replacement was no longer cancelable through the tracker")
|
||||
}
|
||||
}
|
||||
41
internal/rpc/presence_work.go
Normal file
41
internal/rpc/presence_work.go
Normal file
|
|
@ -0,0 +1,41 @@
|
|||
package rpc
|
||||
|
||||
// PresenceWorkSnapshot measures pending grace callbacks and their last-seen
|
||||
// writers, including online last-seen work in the shared batch. It is not a
|
||||
// durability receipt or an inventory of all router background work.
|
||||
type PresenceWorkSnapshot struct {
|
||||
WaitingTimers int64
|
||||
RunningCallbacks int64
|
||||
DirectWrites int64
|
||||
PendingLastSeen int64
|
||||
}
|
||||
|
||||
// PresenceWorkSnapshot is bounded and identity-free. Holding the tracker lock
|
||||
// across the batch read prevents a false zero at callback -> batch handoff:
|
||||
// submission precedes the callback's locked running-count decrement.
|
||||
func (r *Router) PresenceWorkSnapshot() PresenceWorkSnapshot {
|
||||
if r == nil || r.presence == nil {
|
||||
return PresenceWorkSnapshot{}
|
||||
}
|
||||
p := r.presence
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
s := PresenceWorkSnapshot{
|
||||
WaitingTimers: int64(len(p.offlineTimers)),
|
||||
RunningCallbacks: p.offlineRunning,
|
||||
DirectWrites: p.lastSeenDirectRunning,
|
||||
}
|
||||
if r.lastSeenBatch != nil {
|
||||
s.PendingLastSeen = r.lastSeenBatch.pending.Load()
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func (p *presenceTracker) changeDirectLastSeenRunning(delta int64) {
|
||||
if p == nil {
|
||||
return
|
||||
}
|
||||
p.mu.Lock()
|
||||
p.lastSeenDirectRunning += delta
|
||||
p.mu.Unlock()
|
||||
}
|
||||
202
internal/rpc/presence_work_test.go
Normal file
202
internal/rpc/presence_work_test.go
Normal file
|
|
@ -0,0 +1,202 @@
|
|||
package rpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/store"
|
||||
)
|
||||
|
||||
func waitPresenceWork(t *testing.T, r *Router, want PresenceWorkSnapshot) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(time.Second)
|
||||
for {
|
||||
got := r.PresenceWorkSnapshot()
|
||||
if got == want {
|
||||
return
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatalf("presence work = %+v, want %+v", got, want)
|
||||
}
|
||||
time.Sleep(time.Millisecond)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPresenceWorkRunningCallbackSurvivesReplacementAndCancel(t *testing.T) {
|
||||
r := &Router{presence: newPresenceTracker()}
|
||||
started, release := make(chan struct{}), make(chan struct{})
|
||||
r.presence.armOfflineTimer(7, 0, func() { close(started); <-release })
|
||||
select {
|
||||
case <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("callback did not start")
|
||||
}
|
||||
defer func() {
|
||||
close(release)
|
||||
waitPresenceWork(t, r, PresenceWorkSnapshot{})
|
||||
}()
|
||||
waitPresenceWork(t, r, PresenceWorkSnapshot{RunningCallbacks: 1})
|
||||
r.presence.armOfflineTimer(7, time.Hour, func() { t.Error("canceled replacement fired") })
|
||||
waitPresenceWork(t, r, PresenceWorkSnapshot{WaitingTimers: 1, RunningCallbacks: 1})
|
||||
r.presence.cancelOfflineTimer(7)
|
||||
waitPresenceWork(t, r, PresenceWorkSnapshot{RunningCallbacks: 1})
|
||||
}
|
||||
|
||||
func TestPresenceWorkWaitingCancellationAndOtherDevice(t *testing.T) {
|
||||
r := &Router{presence: newPresenceTracker()}
|
||||
key := presenceSessionKey{sessionID: 2}
|
||||
r.presence.setSessionStatus(key, 7, domain.UserStatus{Kind: domain.UserStatusOnline, Expires: 100})
|
||||
r.SessionOfflineAt([8]byte{}, 1, 7, false, 50)
|
||||
waitPresenceWork(t, r, PresenceWorkSnapshot{})
|
||||
if _, ok := r.presence.statusFor(7, 60); !ok {
|
||||
t.Fatal("departure erased the other device")
|
||||
}
|
||||
r.presence.armOfflineTimer(7, time.Hour, func() { t.Error("canceled timer fired") })
|
||||
waitPresenceWork(t, r, PresenceWorkSnapshot{WaitingTimers: 1})
|
||||
r.presence.cancelOfflineTimer(7)
|
||||
waitPresenceWork(t, r, PresenceWorkSnapshot{})
|
||||
}
|
||||
|
||||
type heldPresenceBatch struct {
|
||||
started chan struct{}
|
||||
release chan struct{}
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func (s *heldPresenceBatch) UpdateLastSeenBatch(ctx context.Context, _ []store.UserLastSeenUpdate) error {
|
||||
s.once.Do(func() { close(s.started) })
|
||||
select {
|
||||
case <-s.release:
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func TestPresenceWorkCallbackToBatchHasNoUnobservedHandoff(t *testing.T) {
|
||||
u := &heldPresenceBatch{started: make(chan struct{}), release: make(chan struct{})}
|
||||
metrics := &capturePresenceLastSeenMetrics{}
|
||||
d := newPresenceLastSeenBatchDispatcher(u, presenceLastSeenBatchConfig{MaxSize: 1}, zap.NewNop(), metrics)
|
||||
r := &Router{presence: newPresenceTracker(), lastSeenBatch: d}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
done := make(chan struct{})
|
||||
go func() { d.Run(ctx); close(done) }()
|
||||
defer func() { close(u.release); cancel(); <-done }()
|
||||
callbackStarted, submit := make(chan struct{}), make(chan struct{})
|
||||
r.presence.armOfflineTimer(7, 0, func() {
|
||||
close(callbackStarted)
|
||||
<-submit
|
||||
if err := d.submit(store.UserLastSeenUpdate{UserID: 7, LastSeenAt: 99}); err != nil {
|
||||
t.Error(err)
|
||||
}
|
||||
})
|
||||
<-callbackStarted
|
||||
var observing atomic.Bool
|
||||
observing.Store(true)
|
||||
snapshotsDone := make(chan struct{})
|
||||
go func() {
|
||||
defer close(snapshotsDone)
|
||||
for observing.Load() {
|
||||
s := r.PresenceWorkSnapshot()
|
||||
if s.RunningCallbacks+s.PendingLastSeen != 1 && s.RunningCallbacks+s.PendingLastSeen != 2 {
|
||||
t.Errorf("unobserved callback handoff: %+v", s)
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
close(submit)
|
||||
select {
|
||||
case <-u.started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("batch did not begin")
|
||||
}
|
||||
waitPresenceWork(t, r, PresenceWorkSnapshot{PendingLastSeen: 1})
|
||||
observing.Store(false)
|
||||
<-snapshotsDone
|
||||
if d.pending.Load() != metrics.pending.Load() || metrics.pending.Load() != 1 {
|
||||
t.Fatal("snapshot and exported batch pending disagree")
|
||||
}
|
||||
}
|
||||
|
||||
type heldDirectLastSeen struct{ started, release chan struct{} }
|
||||
|
||||
func (u heldDirectLastSeen) UpdateLastSeen(ctx context.Context, _ int64, _ int) error {
|
||||
close(u.started)
|
||||
select {
|
||||
case <-u.release:
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func TestPresenceWorkDirectWriteRemainsVisibleAfterParentReturns(t *testing.T) {
|
||||
u := heldDirectLastSeen{make(chan struct{}), make(chan struct{})}
|
||||
r := &Router{presence: newPresenceTracker(), log: zap.NewNop()}
|
||||
r.persistReservedLastSeenAsync(context.Background(), u, 7, 99)
|
||||
<-u.started
|
||||
waitPresenceWork(t, r, PresenceWorkSnapshot{DirectWrites: 1})
|
||||
close(u.release)
|
||||
waitPresenceWork(t, r, PresenceWorkSnapshot{})
|
||||
}
|
||||
|
||||
func TestPresenceWorkConcurrentArmCancel(t *testing.T) {
|
||||
p := newPresenceTracker()
|
||||
var workers sync.WaitGroup
|
||||
for i := 0; i < 8; i++ {
|
||||
workers.Go(func() {
|
||||
for j := 0; j < 100; j++ {
|
||||
p.armOfflineTimer(int64(j%4+1), time.Hour, func() { t.Error("canceled timer fired") })
|
||||
p.cancelOfflineTimer(int64(j%4 + 1))
|
||||
}
|
||||
})
|
||||
}
|
||||
workers.Wait()
|
||||
for user := int64(1); user <= 4; user++ {
|
||||
p.cancelOfflineTimer(user)
|
||||
}
|
||||
waitPresenceWork(t, &Router{presence: p}, PresenceWorkSnapshot{})
|
||||
}
|
||||
|
||||
func TestPresenceWorkBatchRetryDrainAndOverflowAccounting(t *testing.T) {
|
||||
u := &capturePresenceLastSeenUpdater{failFirst: 1, called: make(chan struct{}, 8)}
|
||||
m := &capturePresenceLastSeenMetrics{}
|
||||
d := newPresenceLastSeenBatchDispatcher(u, presenceLastSeenBatchConfig{MaxSize: 1, QueueSize: 1, DrainTimeout: time.Second}, zap.NewNop(), m)
|
||||
if err := d.submit(store.UserLastSeenUpdate{UserID: 7, LastSeenAt: 99}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := d.submit(store.UserLastSeenUpdate{UserID: 8, LastSeenAt: 100}); !errors.Is(err, errPresenceLastSeenBatchFull) {
|
||||
t.Fatalf("overflow = %v", err)
|
||||
}
|
||||
if d.pending.Load() != 1 || m.pending.Load() != 1 {
|
||||
t.Fatal("rejected work changed pending")
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
d.Run(ctx)
|
||||
if d.pending.Load() != 0 || m.pending.Load() != 0 || m.failures.Load() != 1 || m.dropped.Load() != 0 {
|
||||
t.Fatalf("drain accounting: pending=%d metric=%d failures=%d dropped=%d", d.pending.Load(), m.pending.Load(), m.failures.Load(), m.dropped.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestPresenceWorkFailedDrainIsCountedNotSuccessful(t *testing.T) {
|
||||
u := &capturePresenceLastSeenUpdater{failFirst: 1000, called: make(chan struct{}, 8)}
|
||||
m := &capturePresenceLastSeenMetrics{}
|
||||
d := newPresenceLastSeenBatchDispatcher(u, presenceLastSeenBatchConfig{MaxSize: 1, DrainTimeout: 10 * time.Millisecond}, zap.NewNop(), m)
|
||||
if err := d.submit(store.UserLastSeenUpdate{UserID: 7, LastSeenAt: 99}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
d.Run(ctx)
|
||||
if d.pending.Load() != 0 || m.pending.Load() != 0 || m.dropped.Load() != 1 || m.failures.Load() < 1 {
|
||||
t.Fatalf("failed drain accounting: pending=%d metric=%d failures=%d dropped=%d", d.pending.Load(), m.pending.Load(), m.failures.Load(), m.dropped.Load())
|
||||
}
|
||||
}
|
||||
|
|
@ -1078,6 +1078,9 @@ func (r *Router) persistAuthKeyClientInfo(ctx context.Context, info clientSessio
|
|||
return
|
||||
}
|
||||
domainInfo := domainAuthKeyClientInfo(info)
|
||||
if ip, ok := ClientIPFrom(ctx); ok {
|
||||
domainInfo.IP = ip
|
||||
}
|
||||
if domainInfo.Layer == 0 && domainInfo.DeviceModel == "" && domainInfo.Platform == "" &&
|
||||
domainInfo.SystemVersion == "" && domainInfo.APIID == 0 && domainInfo.AppVersion == "" {
|
||||
return
|
||||
|
|
|
|||
|
|
@ -508,6 +508,31 @@ func TestObservedClientLayerNeverLeaksAcrossAuthKeySessions(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestPersistAuthKeyClientInfoCarriesClientIP(t *testing.T) {
|
||||
authKeyID := [8]byte{0x68, 0x25, 0x7a, 0x09}
|
||||
auth := &captureAuthService{}
|
||||
r := New(Config{DC: 2, IP: "127.0.0.1", Port: 2398}, Deps{
|
||||
Auth: auth,
|
||||
}, zaptest.NewLogger(t), clock.System)
|
||||
ctx := WithAuthKeyID(
|
||||
WithRawAuthKeyID(WithClientIP(context.Background(), "203.0.113.17"), authKeyID),
|
||||
authKeyID,
|
||||
)
|
||||
|
||||
r.persistAuthKeyClientInfo(ctx, clientSessionInfo{
|
||||
layer: currentClientLayer,
|
||||
hasClientInfo: true,
|
||||
clientInfo: ClientInfo{
|
||||
APIID: 2040, DeviceModel: "Pixel 9", SystemVersion: "SDK 36", AppVersion: "12.8.7",
|
||||
},
|
||||
})
|
||||
|
||||
got := auth.authKeyClientInfos[authKeyID].IP
|
||||
if got != "203.0.113.17" {
|
||||
t.Fatalf("persisted client IP = %q, want %q", got, "203.0.113.17")
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvokeWithLayerPersistsClientLayerUpgrade(t *testing.T) {
|
||||
authKeyID := [8]byte{0x68, 0x25, 0x7a, 0x02}
|
||||
userID := int64(1780269504)
|
||||
|
|
|
|||
|
|
@ -315,6 +315,9 @@ func (s *captureAuthService) UpdateAuthKeyClientInfo(_ context.Context, authKeyI
|
|||
if info.AppVersion != "" {
|
||||
current.AppVersion = info.AppVersion
|
||||
}
|
||||
if info.IP != "" {
|
||||
current.IP = info.IP
|
||||
}
|
||||
s.authKeyClientInfos[authKeyID] = current
|
||||
for i := range s.authorizations {
|
||||
if s.authorizations[i].AuthKeyID != authKeyID {
|
||||
|
|
@ -335,6 +338,9 @@ func (s *captureAuthService) UpdateAuthKeyClientInfo(_ context.Context, authKeyI
|
|||
if info.AppVersion != "" {
|
||||
s.authorizations[i].AppVersion = info.AppVersion
|
||||
}
|
||||
if info.IP != "" {
|
||||
s.authorizations[i].IP = info.IP
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
|
|
|||
|
|
@ -15,11 +15,12 @@ type mapUsersService struct {
|
|||
|
||||
type countingMapUsersService struct {
|
||||
mapUsersService
|
||||
selfCalls int
|
||||
byIDCalls int
|
||||
byIDsCalls int
|
||||
lastByIDs []int64
|
||||
byIDsBatches [][]int64
|
||||
selfCalls int
|
||||
byIDCalls int
|
||||
byIDsCalls int
|
||||
baseByIDsCalls int
|
||||
lastByIDs []int64
|
||||
byIDsBatches [][]int64
|
||||
}
|
||||
|
||||
func (s staticUsersService) Self(context.Context, int64) (domain.User, error) {
|
||||
|
|
@ -104,6 +105,11 @@ func (s *countingMapUsersService) ByIDs(ctx context.Context, currentUserID int64
|
|||
return s.mapUsersService.ByIDs(ctx, currentUserID, userIDs)
|
||||
}
|
||||
|
||||
func (s *countingMapUsersService) BaseUsersByIDs(ctx context.Context, userIDs []int64) ([]domain.User, error) {
|
||||
s.baseByIDsCalls++
|
||||
return s.mapUsersService.ByIDs(ctx, 1, userIDs)
|
||||
}
|
||||
|
||||
type captureUsersService struct {
|
||||
user domain.User
|
||||
userID int64
|
||||
|
|
|
|||
|
|
@ -782,6 +782,7 @@ func TestMessagesHistoryAndSearchProjectStoriesMaxID(t *testing.T) {
|
|||
}}
|
||||
r := New(Config{}, Deps{
|
||||
Messages: messages,
|
||||
Users: mapUsersService{users: map[int64]domain.User{owner.ID: owner, viewer.ID: viewer}},
|
||||
Stories: appstories.NewService(storyStore),
|
||||
}, zaptest.NewLogger(t), fixedClock{now: time.Unix(1700000100, 0)})
|
||||
peer := &tg.InputPeerUser{UserID: owner.ID, AccessHash: owner.AccessHash}
|
||||
|
|
|
|||
|
|
@ -343,6 +343,10 @@ func collectMessagePeerRefs(msg domain.Message, currentChannelID int64, userIDs,
|
|||
if msg.ReplyTo != nil {
|
||||
addDomainPeerRef(msg.ReplyTo.Peer, currentChannelID, userIDs, channelIDs)
|
||||
collectMessageEntityUserRefs(msg.ReplyTo.QuoteEntities, userIDs)
|
||||
if external := msg.ReplyTo.External; external != nil {
|
||||
addDomainPeerRef(external.From.From, currentChannelID, userIDs, channelIDs)
|
||||
collectMessagePeerRefs(domain.Message{Media: external.Media, Entities: external.Entities}, currentChannelID, userIDs, channelIDs)
|
||||
}
|
||||
}
|
||||
if msg.Media != nil && msg.Media.Contact != nil && msg.Media.Contact.UserID != 0 {
|
||||
userIDs[msg.Media.Contact.UserID] = struct{}{}
|
||||
|
|
@ -429,6 +433,10 @@ func collectChannelMessagePeerRefs(msg domain.ChannelMessage, currentChannelID i
|
|||
if msg.ReplyTo != nil {
|
||||
addDomainPeerRef(msg.ReplyTo.Peer, currentChannelID, userIDs, channelIDs)
|
||||
collectMessageEntityUserRefs(msg.ReplyTo.QuoteEntities, userIDs)
|
||||
if external := msg.ReplyTo.External; external != nil {
|
||||
addDomainPeerRef(external.From.From, currentChannelID, userIDs, channelIDs)
|
||||
collectMessagePeerRefs(domain.Message{Media: external.Media, Entities: external.Entities}, currentChannelID, userIDs, channelIDs)
|
||||
}
|
||||
}
|
||||
if msg.Media != nil && msg.Media.Contact != nil && msg.Media.Contact.UserID != 0 {
|
||||
userIDs[msg.Media.Contact.UserID] = struct{}{}
|
||||
|
|
|
|||
|
|
@ -159,11 +159,6 @@ func clonePeerPtr(in *domain.Peer) *domain.Peer {
|
|||
return &out
|
||||
}
|
||||
|
||||
func cloneMessageReply(in *domain.MessageReply) *domain.MessageReply {
|
||||
if in == nil {
|
||||
return nil
|
||||
}
|
||||
out := *in
|
||||
out.QuoteEntities = append([]domain.MessageEntity(nil), in.QuoteEntities...)
|
||||
return &out
|
||||
func cloneMessageReply(reply *domain.MessageReply) *domain.MessageReply {
|
||||
return domain.CloneMessageReply(reply)
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue