merged with fixes

This commit is contained in:
onysd 2026-09-09 02:49:30 +03:00
parent a9e758b712
commit 2f1818d656
176 changed files with 9000 additions and 907 deletions

View file

@ -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,

View file

@ -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)

View file

@ -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)
}

View file

@ -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
}

View file

@ -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)

View file

@ -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 {

View file

@ -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)

View file

@ -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{

View 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")
}
}

View file

@ -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):

View file

@ -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.

View file

@ -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,

View file

@ -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",

View 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)
}
}

View 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)
}
}

View file

@ -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 {

View 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)
}
}
}

View file

@ -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

View file

@ -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 移除再执行 firefire 内部会查在线态做最终去抖),全程不持 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)

View file

@ -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))
}

View file

@ -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

View 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")
}
}

View 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()
}

View 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())
}
}

View file

@ -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

View file

@ -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)

View file

@ -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
}

View file

@ -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

View file

@ -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}

View file

@ -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{}{}

View file

@ -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)
}