fix: sync session membership routing fixes
This commit is contained in:
parent
7e64d9c30e
commit
56d995474c
20 changed files with 639 additions and 102 deletions
73
internal/mtprotoedge/layer_seed_test.go
Normal file
73
internal/mtprotoedge/layer_seed_test.go
Normal file
|
|
@ -0,0 +1,73 @@
|
|||
package mtprotoedge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap/zaptest"
|
||||
|
||||
"github.com/gotd/td/bin"
|
||||
"github.com/gotd/td/mt"
|
||||
"github.com/gotd/td/proto"
|
||||
"github.com/gotd/td/tg"
|
||||
)
|
||||
|
||||
func TestSessionManagerSetClientLayerForAuthKey(t *testing.T) {
|
||||
sm := NewSessionManager(zaptest.NewLogger(t))
|
||||
c := &Conn{sessionID: 42, authKeyID: [8]byte{1, 2, 3}}
|
||||
sm.Register(c)
|
||||
defer sm.Unregister(c)
|
||||
|
||||
sm.SetClientLayerForAuthKey([8]byte{1, 2, 3}, 42, 225)
|
||||
if got := c.ClientLayer(); got != 225 {
|
||||
t.Fatalf("ClientLayer = %d, want 225", got)
|
||||
}
|
||||
// 未注册 session:no-op,不 panic。
|
||||
sm.SetClientLayerForAuthKey([8]byte{9}, 1, 220)
|
||||
// 非法 layer:忽略,保留已有值。
|
||||
sm.SetClientLayerForAuthKey([8]byte{1, 2, 3}, 42, 0)
|
||||
if got := c.ClientLayer(); got != 225 {
|
||||
t.Fatalf("ClientLayer after layer=0 = %d, want 225", got)
|
||||
}
|
||||
}
|
||||
|
||||
type seededLayerRPC struct{}
|
||||
|
||||
func (seededLayerRPC) Dispatch(context.Context, [8]byte, int64, *bin.Buffer) (bin.Encoder, error) {
|
||||
return &tg.Config{ThisDC: 2}, nil
|
||||
}
|
||||
|
||||
func (seededLayerRPC) NegotiatedLayer([8]byte, int64) (int, bool) { return 225, true }
|
||||
|
||||
// TestRegisterSeedsNegotiatedLayerBeforeFirstRPC 验证连接注册即从 rpc 层播种协商 layer:
|
||||
// 只发一条 ping(服务消息,不经 RPC Dispatch),连接的 ClientLayer 就必须是协商值,
|
||||
// 而不是等首条 RPC 的 Dispatch 返回后才刷新——否则重连老客户端在首条 RPC handler
|
||||
// 执行期间收到的 pending flush / 并发 push 会按 canonical 227 漏降级。
|
||||
func TestRegisterSeedsNegotiatedLayerBeforeFirstRPC(t *testing.T) {
|
||||
addr, pub, srv := startTestServer(t, Options{DC: 2, RPC: seededLayerRPC{}})
|
||||
conn, auth, cipher := dialHandshake(t, addr, 2, pub)
|
||||
|
||||
clientMsgID := proto.NewMessageIDGen(time.Now)
|
||||
sendEncrypted(t, conn, cipher, auth, clientMsgID.New(proto.MessageFromClient), &mt.PingRequest{PingID: 7})
|
||||
|
||||
// 等 pong 回来,确保携带注册动作的那一帧已处理完成。
|
||||
gotPong := false
|
||||
for i := 0; i < 12 && !gotPong; i++ {
|
||||
_, id, _ := readServerMessage(t, conn, cipher, auth.AuthKey)
|
||||
gotPong = id == mt.PongTypeID
|
||||
}
|
||||
if !gotPong {
|
||||
t.Fatal("missing pong for ping")
|
||||
}
|
||||
|
||||
srv.conns.mu.RLock()
|
||||
c := srv.conns.bySession[sessionKey{authKeyID: auth.AuthKey.ID, sessionID: auth.SessionID}]
|
||||
srv.conns.mu.RUnlock()
|
||||
if c == nil {
|
||||
t.Fatal("connection not registered")
|
||||
}
|
||||
if got := c.ClientLayer(); got != 225 {
|
||||
t.Fatalf("ClientLayer after registration = %d, want seeded 225", got)
|
||||
}
|
||||
}
|
||||
|
|
@ -27,7 +27,7 @@ func TestSetReceivesUpdatesFlushesPendingBeforeActivation(t *testing.T) {
|
|||
|
||||
// 完全就绪还要求 membership 路由建立(ReceivesUpdatesForAuthKey 的另一半条件)。
|
||||
srv.Conns().BindUserForAuthKey(raw, auth.SessionID, 100)
|
||||
srv.Conns().SetSessionChannelMemberships(raw, auth.SessionID, 100, nil)
|
||||
srv.Conns().SetSessionChannelMemberships(raw, auth.SessionID, 100, nil, srv.Conns().ChannelMembershipGeneration(raw, auth.SessionID))
|
||||
|
||||
// 未就绪:推送进 pending 而非直发。
|
||||
for i := 0; i < 2; i++ {
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ import (
|
|||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zapcore"
|
||||
|
||||
"github.com/gotd/td/bin"
|
||||
"github.com/gotd/td/proto"
|
||||
|
|
@ -128,14 +129,16 @@ func (m *SessionManager) Register(c *Conn) {
|
|||
replaced = old
|
||||
m.removeLocked(old, false)
|
||||
} else if existing := m.byAuthKey[c.authKeyID]; len(existing) >= maxSessionsPerAuthKey {
|
||||
// 同 raw auth_key 的 session 数达上限且本次是新 session:驱逐一个现有 session 让位,
|
||||
// 同 raw auth_key 的 session 数达上限且本次是新 session:驱逐建连最早的 session 让位,
|
||||
// 防对抗客户端用海量 session_id 撑爆索引。驱逐对象与新连接同属一个设备凭据,
|
||||
// 触顶基本是该凭据自身异常。被驱逐连接的 serveConn 会在下一帧因 actor 已关而退出。
|
||||
// 触顶基本是该凭据自身异常;选最旧而非 map 随机,避免误杀刚建立的活跃下载/主连接。
|
||||
// 被驱逐连接的 serveConn 会在下一帧因 actor 已关而退出。O(cap) 扫描仅在触顶时发生。
|
||||
for _, ec := range existing {
|
||||
evicted = ec
|
||||
m.removeLocked(ec, true)
|
||||
break
|
||||
if evicted == nil || ec.createdAt.Before(evicted.createdAt) {
|
||||
evicted = ec
|
||||
}
|
||||
}
|
||||
m.removeLocked(evicted, true)
|
||||
m.log.Debug("Evicted oldest session for auth key at cap",
|
||||
zap.String("auth_key_id", sessionKeyLog(c.authKeyID)),
|
||||
zap.Int("cap", maxSessionsPerAuthKey),
|
||||
|
|
@ -280,7 +283,7 @@ func (m *SessionManager) bindUserLocked(c *Conn, key sessionKey, userID int64) {
|
|||
removeUserIndex(m.byUser, old, key)
|
||||
if old != userID {
|
||||
m.clearChannelInterestsLocked(key)
|
||||
m.clearChannelMembershipsLocked(key)
|
||||
m.clearChannelMembershipsLocked(c, key)
|
||||
c.membershipsSynced.Store(false)
|
||||
// 身份变化即丢弃暂存推送:它们属于前一个账号,flush 给新账号是跨账号泄露。
|
||||
// 同时取消进行中的排空(runFlush 还另有 owner 校验做批内兜底)。
|
||||
|
|
@ -293,7 +296,7 @@ func (m *SessionManager) bindUserLocked(c *Conn, key sessionKey, userID int64) {
|
|||
addUserIndex(m.byUser, userID, key, c)
|
||||
} else {
|
||||
m.clearChannelInterestsLocked(key)
|
||||
m.clearChannelMembershipsLocked(key)
|
||||
m.clearChannelMembershipsLocked(c, key)
|
||||
c.membershipsSynced.Store(false)
|
||||
delete(m.pending, key)
|
||||
delete(m.flushing, key)
|
||||
|
|
@ -394,7 +397,7 @@ func (m *SessionManager) bindAuthKeyLocked(c *Conn, key sessionKey, authKeyID [8
|
|||
removeUserIndex(m.byUser, oldUserID, key)
|
||||
}
|
||||
m.clearChannelInterestsLocked(key)
|
||||
m.clearChannelMembershipsLocked(key)
|
||||
m.clearChannelMembershipsLocked(c, key)
|
||||
c.membershipsSynced.Store(false)
|
||||
delete(m.pending, key)
|
||||
delete(m.flushing, key)
|
||||
|
|
@ -514,7 +517,7 @@ func (m *SessionManager) UnbindAuthKey(authKeyID [8]byte) int {
|
|||
removeUserIndex(m.byUser, old, key)
|
||||
}
|
||||
m.clearChannelInterestsLocked(key)
|
||||
m.clearChannelMembershipsLocked(key)
|
||||
m.clearChannelMembershipsLocked(c, key)
|
||||
c.membershipsSynced.Store(false)
|
||||
// 授权解除后暂存推送属于已登出的账号,不能等下一个登录者置位时 flush 出去。
|
||||
delete(m.pending, key)
|
||||
|
|
@ -556,7 +559,7 @@ func (m *SessionManager) setReceivesUpdatesLocked(c *Conn, key sessionKey, recei
|
|||
if !receives {
|
||||
c.receivesUpdates.Store(false)
|
||||
m.clearChannelInterestsLocked(key)
|
||||
m.clearChannelMembershipsLocked(key)
|
||||
m.clearChannelMembershipsLocked(c, key)
|
||||
c.membershipsSynced.Store(false)
|
||||
// 取消进行中的排空激活:runFlush 在置位前会复查该标志,标志已删则放弃置位,
|
||||
// 避免把刚置 false 的开关翻回 true。
|
||||
|
|
@ -679,6 +682,22 @@ func (m *SessionManager) ReceivesUpdatesForAuthKey(authKeyID [8]byte, sessionID
|
|||
return ok && c.receivesUpdates.Load() && c.membershipsSynced.Load()
|
||||
}
|
||||
|
||||
// SetClientLayerForAuthKey 把协商的 TL layer 即时写到指定连接。由 rpc 层在
|
||||
// invokeWithLayer 观测到新 layer 时(Dispatch 入口,早于鉴权门与 updates 就绪
|
||||
// 置位)调用,使同一条请求触发的 pending flush 与并发 push 立即按正确 layer
|
||||
// 降级,不必等连接层在 Dispatch 返回后的兜底刷新。
|
||||
func (m *SessionManager) SetClientLayerForAuthKey(authKeyID [8]byte, sessionID int64, layer int) {
|
||||
if m == nil || layer <= 0 {
|
||||
return
|
||||
}
|
||||
m.mu.RLock()
|
||||
c, ok := m.bySession[sessionKey{authKeyID: authKeyID, sessionID: sessionID}]
|
||||
m.mu.RUnlock()
|
||||
if ok {
|
||||
c.SetClientLayer(layer)
|
||||
}
|
||||
}
|
||||
|
||||
// SetReceivesUpdatesForAuthKey 标记指定 raw auth_key_id + session_id 是否接收主动 updates。
|
||||
func (m *SessionManager) SetReceivesUpdatesForAuthKey(authKeyID [8]byte, sessionID int64, receives bool) {
|
||||
m.mu.Lock()
|
||||
|
|
@ -905,14 +924,21 @@ func onceEncodedOutbound(msg bin.Encoder) func() (*encodedOutboundMessage, error
|
|||
}
|
||||
|
||||
func (m *SessionManager) pushToUserWithSender(ctx context.Context, userID int64, excludeAuthKeyID *[8]byte, excludeSessionID int64, t proto.MessageType, msg bin.Encoder, queueWhenNotReady bool, send func(*Conn) error) (int, error) {
|
||||
m.mu.Lock()
|
||||
// push fan-out 是连接层最热路径之一:debug 日志的字段构造(含 auth_key hex 格式化)
|
||||
// 在关闭 debug 时也会求值,先查级别一次、按需记日志。
|
||||
debug := m.log.Core().Enabled(zapcore.DebugLevel)
|
||||
// 快路径:稳态下目标连接全部就绪(或 transient 直接跳过未就绪者),收集连接
|
||||
// 只读不写,全程共享读锁即可,避免每次 push 都拿独占写锁串行整个注册表。
|
||||
// 仅当 durable 推送遇到未就绪连接(需写 pending)时才回落写锁重扫。
|
||||
m.mu.RLock()
|
||||
total := len(m.byUser[userID])
|
||||
conns := make([]*Conn, 0, total)
|
||||
queued := 0
|
||||
dropped := 0
|
||||
excluded := 0
|
||||
skipped := 0
|
||||
for key, c := range m.byUser[userID] {
|
||||
needQueue := false
|
||||
for _, c := range m.byUser[userID] {
|
||||
if shouldExcludeSession(c, excludeAuthKeyID, excludeSessionID) {
|
||||
excluded++
|
||||
continue
|
||||
|
|
@ -924,26 +950,53 @@ func (m *SessionManager) pushToUserWithSender(ctx context.Context, userID int64,
|
|||
skipped++
|
||||
continue
|
||||
}
|
||||
if m.queueLocked(key, t, msg) {
|
||||
queued++
|
||||
m.log.Debug("Push queued (session not updates-ready)",
|
||||
zap.Int64("user_id", userID),
|
||||
zap.String("auth_key_id", sessionKeyLog(key.authKeyID)),
|
||||
zap.Int64("session_id", key.sessionID),
|
||||
)
|
||||
} else {
|
||||
dropped++
|
||||
m.log.Debug("Push dropped (stale pending; durable log covers)",
|
||||
zap.Int64("user_id", userID),
|
||||
zap.String("auth_key_id", sessionKeyLog(key.authKeyID)),
|
||||
zap.Int64("session_id", key.sessionID),
|
||||
)
|
||||
}
|
||||
continue
|
||||
needQueue = true
|
||||
break
|
||||
}
|
||||
conns = append(conns, c)
|
||||
}
|
||||
m.mu.Unlock()
|
||||
m.mu.RUnlock()
|
||||
if needQueue {
|
||||
// 写锁下完整重扫(读锁释放到此之间状态可能变化,以重扫结果为准)。
|
||||
conns = conns[:0]
|
||||
queued, dropped, excluded, skipped = 0, 0, 0, 0
|
||||
m.mu.Lock()
|
||||
total = len(m.byUser[userID])
|
||||
for key, c := range m.byUser[userID] {
|
||||
if shouldExcludeSession(c, excludeAuthKeyID, excludeSessionID) {
|
||||
excluded++
|
||||
continue
|
||||
}
|
||||
if !c.receivesUpdates.Load() {
|
||||
if !queueWhenNotReady {
|
||||
skipped++
|
||||
continue
|
||||
}
|
||||
if m.queueLocked(key, t, msg) {
|
||||
queued++
|
||||
if debug {
|
||||
m.log.Debug("Push queued (session not updates-ready)",
|
||||
zap.Int64("user_id", userID),
|
||||
zap.String("auth_key_id", sessionKeyLog(key.authKeyID)),
|
||||
zap.Int64("session_id", key.sessionID),
|
||||
)
|
||||
}
|
||||
} else {
|
||||
dropped++
|
||||
if debug {
|
||||
m.log.Debug("Push dropped (stale pending; durable log covers)",
|
||||
zap.Int64("user_id", userID),
|
||||
zap.String("auth_key_id", sessionKeyLog(key.authKeyID)),
|
||||
zap.Int64("session_id", key.sessionID),
|
||||
)
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
conns = append(conns, c)
|
||||
}
|
||||
m.mu.Unlock()
|
||||
}
|
||||
|
||||
var firstErr error
|
||||
sent := 0
|
||||
|
|
@ -958,33 +1011,39 @@ func (m *SessionManager) pushToUserWithSender(ctx context.Context, userID int64,
|
|||
if firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
m.log.Debug("Push to conn failed",
|
||||
zap.Int64("user_id", userID),
|
||||
zap.String("auth_key_id", sessionKeyLog(c.authKeyID)),
|
||||
zap.Int64("session_id", c.sessionID),
|
||||
zap.Error(err),
|
||||
)
|
||||
if debug {
|
||||
m.log.Debug("Push to conn failed",
|
||||
zap.Int64("user_id", userID),
|
||||
zap.String("auth_key_id", sessionKeyLog(c.authKeyID)),
|
||||
zap.Int64("session_id", c.sessionID),
|
||||
zap.Error(err),
|
||||
)
|
||||
}
|
||||
continue
|
||||
}
|
||||
sent++
|
||||
m.log.Debug("Push to conn ok",
|
||||
zap.Int64("user_id", userID),
|
||||
zap.String("auth_key_id", sessionKeyLog(c.authKeyID)),
|
||||
zap.Int64("session_id", c.sessionID),
|
||||
)
|
||||
if debug {
|
||||
m.log.Debug("Push to conn ok",
|
||||
zap.Int64("user_id", userID),
|
||||
zap.String("auth_key_id", sessionKeyLog(c.authKeyID)),
|
||||
zap.Int64("session_id", c.sessionID),
|
||||
)
|
||||
}
|
||||
}
|
||||
if total == 0 {
|
||||
m.log.Debug("Push to user: no active conns", zap.Int64("user_id", userID))
|
||||
} else if excluded > 0 || queued > 0 || dropped > 0 || skipped > 0 || sent < len(conns) {
|
||||
m.log.Debug("Push to user summary",
|
||||
zap.Int64("user_id", userID),
|
||||
zap.Int("conns", total),
|
||||
zap.Int("sent", sent),
|
||||
zap.Int("queued", queued),
|
||||
zap.Int("dropped", dropped),
|
||||
zap.Int("skipped_transient", skipped),
|
||||
zap.Int("excluded", excluded),
|
||||
)
|
||||
if debug {
|
||||
if total == 0 {
|
||||
m.log.Debug("Push to user: no active conns", zap.Int64("user_id", userID))
|
||||
} else if excluded > 0 || queued > 0 || dropped > 0 || skipped > 0 || sent < len(conns) {
|
||||
m.log.Debug("Push to user summary",
|
||||
zap.Int64("user_id", userID),
|
||||
zap.Int("conns", total),
|
||||
zap.Int("sent", sent),
|
||||
zap.Int("queued", queued),
|
||||
zap.Int("dropped", dropped),
|
||||
zap.Int("skipped_transient", skipped),
|
||||
zap.Int("excluded", excluded),
|
||||
)
|
||||
}
|
||||
}
|
||||
return sent + queued, firstErr
|
||||
}
|
||||
|
|
@ -1014,7 +1073,7 @@ func (m *SessionManager) OnlineUserIDsForCandidates(candidateUserIDs []int64, li
|
|||
}
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
out := make([]int64, 0, minInt(len(candidateUserIDs), positiveLimitOrLen(limit, len(candidateUserIDs))))
|
||||
out := make([]int64, 0, min(len(candidateUserIDs), positiveLimitOrLen(limit, len(candidateUserIDs))))
|
||||
seen := make(map[int64]struct{}, len(candidateUserIDs))
|
||||
for _, userID := range candidateUserIDs {
|
||||
if userID == 0 {
|
||||
|
|
@ -1078,10 +1137,29 @@ func (m *SessionManager) OnlineChannelUserIDs(channelID int64, limit int) []int6
|
|||
return m.onlineChannelUsers(m.byChannel, channelID, limit)
|
||||
}
|
||||
|
||||
// ChannelMembershipGeneration 返回该 session 的 membership 索引修订号。
|
||||
// 全量同步方必须在读取持久成员列表【之前】采样,并经 SetSessionChannelMemberships
|
||||
// 带回比对;session 不在线时返回 0(后续 Set 也会因查不到连接而放弃)。
|
||||
func (m *SessionManager) ChannelMembershipGeneration(rawAuthKeyID [8]byte, sessionID int64) int64 {
|
||||
m.mu.RLock()
|
||||
c, ok := m.bySession[sessionKey{authKeyID: rawAuthKeyID, sessionID: sessionID}]
|
||||
m.mu.RUnlock()
|
||||
if !ok {
|
||||
return 0
|
||||
}
|
||||
return c.membershipGen.Load()
|
||||
}
|
||||
|
||||
// SetSessionChannelMemberships replaces the joined-channel index for one
|
||||
// updates-ready session. This index is broader than TrackChannelInterest and is
|
||||
// used for durable channel updates such as new/edit/delete message.
|
||||
func (m *SessionManager) SetSessionChannelMemberships(rawAuthKeyID [8]byte, sessionID, userID int64, channelIDs []int64) {
|
||||
//
|
||||
// expectedGen 是调用方在读取持久成员列表前经 ChannelMembershipGeneration 采样的
|
||||
// 修订号。若落地时修订号已变(读取窗口内发生了 join/leave/kick 的增量修订或整体
|
||||
// 清除),全量替换会覆盖掉窗口内的增量——此时改走并集合并(保留增量 Add;合并回的
|
||||
// stale 条目由 fan-out 前的 PG active 复核兜底),并保持 membershipsSynced=false,
|
||||
// 让下一条 RPC 重新走全量同步收敛。
|
||||
func (m *SessionManager) SetSessionChannelMemberships(rawAuthKeyID [8]byte, sessionID, userID int64, channelIDs []int64, expectedGen int64) {
|
||||
key := sessionKey{authKeyID: rawAuthKeyID, sessionID: sessionID}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
|
@ -1089,11 +1167,22 @@ func (m *SessionManager) SetSessionChannelMemberships(rawAuthKeyID [8]byte, sess
|
|||
if !ok {
|
||||
return
|
||||
}
|
||||
m.clearChannelMembershipsLocked(key)
|
||||
c.membershipsSynced.Store(false)
|
||||
if userID == 0 || c.userID.Load() != userID {
|
||||
m.clearChannelMembershipsLocked(c, key)
|
||||
c.membershipsSynced.Store(false)
|
||||
return
|
||||
}
|
||||
if c.membershipGen.Load() != expectedGen {
|
||||
c.membershipsSynced.Store(false)
|
||||
m.trackChannelIndexLocked(m.byMemberChannel, m.bySessionMembers, key, userID, channelIDs)
|
||||
m.log.Debug("Channel membership sync raced with incremental updates; merged and kept unsynced",
|
||||
zap.String("auth_key_id", sessionKeyLog(rawAuthKeyID)),
|
||||
zap.Int64("session_id", sessionID),
|
||||
)
|
||||
return
|
||||
}
|
||||
m.clearChannelMembershipsLocked(c, key)
|
||||
c.membershipsSynced.Store(false)
|
||||
m.trackChannelIndexLocked(m.byMemberChannel, m.bySessionMembers, key, userID, channelIDs)
|
||||
c.membershipsSynced.Store(true)
|
||||
}
|
||||
|
|
@ -1110,6 +1199,7 @@ func (m *SessionManager) AddUserChannelMembership(userID, channelID int64) {
|
|||
if c == nil || c.userID.Load() != userID {
|
||||
continue
|
||||
}
|
||||
c.membershipGen.Add(1)
|
||||
m.trackChannelIndexLocked(m.byMemberChannel, m.bySessionMembers, key, userID, []int64{channelID})
|
||||
}
|
||||
}
|
||||
|
|
@ -1122,7 +1212,10 @@ func (m *SessionManager) RemoveUserChannelMembership(userID, channelID int64) {
|
|||
}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
for key := range m.byUser[userID] {
|
||||
for key, c := range m.byUser[userID] {
|
||||
if c != nil {
|
||||
c.membershipGen.Add(1)
|
||||
}
|
||||
m.removeChannelIndexLocked(m.byMemberChannel, m.bySessionMembers, key, channelID)
|
||||
}
|
||||
}
|
||||
|
|
@ -1151,7 +1244,7 @@ func (m *SessionManager) OnlineChannelMemberUserIDsExcluding(channelID int64, ex
|
|||
if len(sessions) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make([]int64, 0)
|
||||
out := make([]int64, 0, positiveLimitOrLen(limit, len(sessions)))
|
||||
seen := make(map[int64]struct{}, len(sessions))
|
||||
for key, userID := range sessions {
|
||||
if userID == 0 {
|
||||
|
|
@ -1219,7 +1312,7 @@ func (m *SessionManager) removeLocked(c *Conn, dropPending bool) int64 {
|
|||
removeUserIndex(m.byUser, uid, key)
|
||||
}
|
||||
m.clearChannelInterestsLocked(key)
|
||||
m.clearChannelMembershipsLocked(key)
|
||||
m.clearChannelMembershipsLocked(c, key)
|
||||
if dropPending {
|
||||
delete(m.pending, key)
|
||||
}
|
||||
|
|
@ -1247,7 +1340,10 @@ func (m *SessionManager) clearChannelInterestsLocked(key sessionKey) {
|
|||
m.clearChannelIndexLocked(m.byChannel, m.bySessionChannels, key)
|
||||
}
|
||||
|
||||
func (m *SessionManager) clearChannelMembershipsLocked(key sessionKey) {
|
||||
// clearChannelMembershipsLocked 整体清除某连接的 membership 索引并递增其修订号,
|
||||
// 使在飞的全量同步(SetSessionChannelMemberships)能检测到清除并放弃过期替换。
|
||||
func (m *SessionManager) clearChannelMembershipsLocked(c *Conn, key sessionKey) {
|
||||
c.membershipGen.Add(1)
|
||||
m.clearChannelIndexLocked(m.byMemberChannel, m.bySessionMembers, key)
|
||||
}
|
||||
|
||||
|
|
@ -1322,13 +1418,6 @@ func positiveLimitOrLen(limit, length int) int {
|
|||
return length
|
||||
}
|
||||
|
||||
func minInt(a, b int) int {
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func (m *SessionManager) takePendingLocked(key sessionKey, ready bool) []queuedPush {
|
||||
if !ready || len(m.pending[key]) == 0 {
|
||||
return nil
|
||||
|
|
|
|||
|
|
@ -260,7 +260,7 @@ func TestSessionManagerChannelInterestIndex(t *testing.T) {
|
|||
if got := sm.OnlineChannelMemberUserIDs(10, 10); len(got) != 0 {
|
||||
t.Fatalf("channel 10 online members before membership sync = %v, want empty", got)
|
||||
}
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{10, 30})
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{10, 30}, sm.ChannelMembershipGeneration(raw, 42))
|
||||
if got := sm.OnlineChannelMemberUserIDs(10, 10); len(got) != 1 || got[0] != 100 {
|
||||
t.Fatalf("channel 10 online members = %v, want [100]", got)
|
||||
}
|
||||
|
|
@ -315,7 +315,7 @@ func TestSessionManagerClearsChannelIndexesOnAuthAndReadinessChanges(t *testing.
|
|||
|
||||
track := func() {
|
||||
sm.TrackChannelInterest(raw, 42, 100, []int64{10})
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{10})
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{10}, sm.ChannelMembershipGeneration(raw, 42))
|
||||
if got := sm.OnlineChannelUserIDs(10, 10); len(got) != 1 || got[0] != 100 {
|
||||
t.Fatalf("channel viewers before cleanup = %v, want [100]", got)
|
||||
}
|
||||
|
|
|
|||
95
internal/mtprotoedge/session_membership_gen_test.go
Normal file
95
internal/mtprotoedge/session_membership_gen_test.go
Normal file
|
|
@ -0,0 +1,95 @@
|
|||
package mtprotoedge
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap/zaptest"
|
||||
)
|
||||
|
||||
// TestSetSessionChannelMembershipsDetectsConcurrentIncrementalUpdates 验证全量
|
||||
// membership 同步的丢失更新防护:同步方在读持久成员列表前采样修订号,读取窗口内
|
||||
// 若发生增量 join/leave(另一设备操作经 Add/RemoveUserChannelMembership 落索引),
|
||||
// 携带过期修订号的全量替换必须改走并集合并(不得覆盖增量),并保持
|
||||
// membershipsSynced=false 促使下一条 RPC 重试全量同步收敛。
|
||||
func TestSetSessionChannelMembershipsDetectsConcurrentIncrementalUpdates(t *testing.T) {
|
||||
sm := NewSessionManager(zaptest.NewLogger(t))
|
||||
raw := [8]byte{1, 2, 3}
|
||||
c := &Conn{sessionID: 42, authKeyID: raw}
|
||||
sm.Register(c)
|
||||
sm.BindUserForAuthKey(raw, 42, 100)
|
||||
sm.SetReceivesUpdatesForAuthKey(raw, 42, true)
|
||||
|
||||
// 同步方采样修订号后、全量列表落地前,用户在另一台设备加入了频道 7。
|
||||
gen := sm.ChannelMembershipGeneration(raw, 42)
|
||||
sm.AddUserChannelMembership(100, 7)
|
||||
// 基于旧快照的全量列表(只有频道 5,不含 7)携带过期修订号落地。
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{5}, gen)
|
||||
|
||||
if got := sm.OnlineChannelMemberUserIDs(7, 10); len(got) != 1 || got[0] != 100 {
|
||||
t.Fatalf("channel 7 members = %v, want [100]: full replace overwrote the in-window incremental join", got)
|
||||
}
|
||||
if got := sm.OnlineChannelMemberUserIDs(5, 10); len(got) != 1 || got[0] != 100 {
|
||||
t.Fatalf("channel 5 members = %v, want [100]: merge path must still apply the full list", got)
|
||||
}
|
||||
if sm.ReceivesUpdatesForAuthKey(raw, 42) {
|
||||
t.Fatal("session fully ready despite raced membership sync; retry would never happen")
|
||||
}
|
||||
|
||||
// 重试:新修订号下的全量同步正常替换并置就绪。
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{5, 7}, sm.ChannelMembershipGeneration(raw, 42))
|
||||
if !sm.ReceivesUpdatesForAuthKey(raw, 42) {
|
||||
t.Fatal("session not ready after clean resync")
|
||||
}
|
||||
|
||||
// 反方向:窗口内被移出频道 5,stale 全量含 5 → 合并会短暂保留 stale 条目
|
||||
// (fan-out 前的 PG active 复核兜底),但必须保持未就绪等待重试。
|
||||
gen = sm.ChannelMembershipGeneration(raw, 42)
|
||||
sm.RemoveUserChannelMembership(100, 5)
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{5, 7}, gen)
|
||||
if sm.ReceivesUpdatesForAuthKey(raw, 42) {
|
||||
t.Fatal("session ready despite raced removal during sync")
|
||||
}
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{7}, sm.ChannelMembershipGeneration(raw, 42))
|
||||
if got := sm.OnlineChannelMemberUserIDs(5, 10); len(got) != 0 {
|
||||
t.Fatalf("channel 5 members after resync = %v, want empty", got)
|
||||
}
|
||||
if !sm.ReceivesUpdatesForAuthKey(raw, 42) {
|
||||
t.Fatal("session not ready after final resync")
|
||||
}
|
||||
}
|
||||
|
||||
// TestRegisterEvictsOldestSessionAtCap 验证同 raw auth_key session 数触顶时驱逐的是
|
||||
// 建连最早的连接,而不是 map 迭代顺序下的随机一个(随机可能误杀刚建立的活跃连接)。
|
||||
func TestRegisterEvictsOldestSessionAtCap(t *testing.T) {
|
||||
sm := NewSessionManager(zaptest.NewLogger(t))
|
||||
raw := [8]byte{9}
|
||||
base := time.Unix(1_700_000_000, 0)
|
||||
|
||||
const oldestSession = int64(100)
|
||||
for i := 0; i < maxSessionsPerAuthKey; i++ {
|
||||
sid := int64(i + 1)
|
||||
created := base.Add(time.Duration(i+1) * time.Second)
|
||||
if sid == oldestSession {
|
||||
created = base // 唯一早于所有其它连接的时间戳,且故意不在注册顺序首位。
|
||||
}
|
||||
sm.Register(&Conn{sessionID: sid, authKeyID: raw, createdAt: created})
|
||||
}
|
||||
|
||||
sm.Register(&Conn{sessionID: 9999, authKeyID: raw, createdAt: base.Add(time.Hour)})
|
||||
|
||||
sm.mu.RLock()
|
||||
_, oldestAlive := sm.bySession[sessionKey{authKeyID: raw, sessionID: oldestSession}]
|
||||
_, newestAlive := sm.bySession[sessionKey{authKeyID: raw, sessionID: 9999}]
|
||||
total := len(sm.byAuthKey[raw])
|
||||
sm.mu.RUnlock()
|
||||
if oldestAlive {
|
||||
t.Fatal("oldest session survived eviction at cap")
|
||||
}
|
||||
if !newestAlive {
|
||||
t.Fatal("newly registered session missing after eviction")
|
||||
}
|
||||
if total != maxSessionsPerAuthKey {
|
||||
t.Fatalf("sessions for auth key = %d, want cap %d", total, maxSessionsPerAuthKey)
|
||||
}
|
||||
}
|
||||
|
|
@ -23,23 +23,23 @@ func TestReceivesUpdatesForAuthKeyRequiresMembershipSync(t *testing.T) {
|
|||
t.Fatal("ready before membership sync — a failed sync would never be retried")
|
||||
}
|
||||
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{7})
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{7}, sm.ChannelMembershipGeneration(raw, 42))
|
||||
if !sm.ReceivesUpdatesForAuthKey(raw, 42) {
|
||||
t.Fatal("not ready after successful membership sync")
|
||||
}
|
||||
|
||||
// 没有任何频道的账号:空列表的成功同步同样算就绪。
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, nil)
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, nil, sm.ChannelMembershipGeneration(raw, 42))
|
||||
if !sm.ReceivesUpdatesForAuthKey(raw, 42) {
|
||||
t.Fatal("not ready after successful empty membership sync")
|
||||
}
|
||||
|
||||
// userID 与连接当前绑定不一致(换号竞态)时不算就绪,等正确身份重试。
|
||||
sm.SetSessionChannelMemberships(raw, 42, 999, []int64{7})
|
||||
sm.SetSessionChannelMemberships(raw, 42, 999, []int64{7}, sm.ChannelMembershipGeneration(raw, 42))
|
||||
if sm.ReceivesUpdatesForAuthKey(raw, 42) {
|
||||
t.Fatal("ready after membership sync for mismatched user")
|
||||
}
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{7})
|
||||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{7}, sm.ChannelMembershipGeneration(raw, 42))
|
||||
|
||||
// 登出清除就绪标志。
|
||||
sm.BindUserForAuthKey(raw, 42, 0)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue