469 lines
16 KiB
Go
469 lines
16 KiB
Go
package mtprotoedge
|
||
|
||
import (
|
||
"context"
|
||
"errors"
|
||
"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"
|
||
)
|
||
|
||
type countingOutboundEncoder struct {
|
||
count *int
|
||
}
|
||
|
||
func (e *countingOutboundEncoder) Encode(b *bin.Buffer) error {
|
||
*e.count++
|
||
return (&tg.UpdatesTooLong{}).Encode(b)
|
||
}
|
||
|
||
type closeCountingTransport struct {
|
||
closes int
|
||
}
|
||
|
||
func (t *closeCountingTransport) Send(context.Context, *bin.Buffer) error {
|
||
return errors.New("test transport send")
|
||
}
|
||
|
||
func (t *closeCountingTransport) Recv(context.Context, *bin.Buffer) error {
|
||
return errors.New("test transport recv")
|
||
}
|
||
|
||
func (t *closeCountingTransport) Close() error {
|
||
t.closes++
|
||
return nil
|
||
}
|
||
|
||
// TestSessionManagerRegistry 验证注册表的注册/注销/查找语义(不涉及网络发送)。
|
||
func TestSessionManagerRegistry(t *testing.T) {
|
||
sm := NewSessionManager(zaptest.NewLogger(t))
|
||
c := &Conn{sessionID: 42, authKeyID: [8]byte{1, 2, 3}}
|
||
c.receivesUpdates.Store(true)
|
||
|
||
sm.Register(c)
|
||
if got := sm.Online(); got != 1 {
|
||
t.Fatalf("online = %d, want 1", got)
|
||
}
|
||
sm.BindAuthKey(42, [8]byte{1, 2, 3})
|
||
sm.BindUser(42, 100)
|
||
if userID, ok := sm.UserID(42); !ok || userID != 100 {
|
||
t.Fatalf("cached user = %d ok %v, want 100/true", userID, ok)
|
||
}
|
||
sm.BindAuthKey(42, [8]byte{9})
|
||
if userID, ok := sm.UserID(42); ok || userID != 0 {
|
||
t.Fatalf("cached user after auth key switch = %d ok %v, want 0/false", userID, ok)
|
||
}
|
||
if userID, resolved := sm.UserIDResolved(42); resolved || userID != 0 {
|
||
t.Fatalf("resolved user after auth key switch = %d resolved %v, want unresolved", userID, resolved)
|
||
}
|
||
sm.BindUser(42, 0)
|
||
if userID, resolved := sm.UserIDResolved(42); !resolved || userID != 0 {
|
||
t.Fatalf("negative user cache = %d resolved %v, want 0/true", userID, resolved)
|
||
}
|
||
|
||
sm.Unregister(c)
|
||
if got := sm.Online(); got != 0 {
|
||
t.Fatalf("online after unregister = %d, want 0", got)
|
||
}
|
||
|
||
err := sm.PushToSession(context.Background(), 42, proto.MessageFromServer, &tg.UpdatesTooLong{})
|
||
if !errors.Is(err, ErrSessionNotFound) {
|
||
t.Fatalf("push to missing session err = %v, want ErrSessionNotFound", err)
|
||
}
|
||
}
|
||
|
||
func TestSessionManagerBestEffortFanoutPreencodesOnce(t *testing.T) {
|
||
sm := NewSessionManager(zaptest.NewLogger(t))
|
||
const userID = int64(100)
|
||
for i := 0; i < 2; i++ {
|
||
c := &Conn{
|
||
sessionID: int64(i + 1),
|
||
authKeyID: [8]byte{byte(i + 1)},
|
||
outbound: make(chan outboundOp, 1),
|
||
outboundControl: make(chan outboundOp, 1),
|
||
outboundStop: make(chan struct{}),
|
||
}
|
||
c.userID.Store(userID)
|
||
c.userIDResolved.Store(true)
|
||
c.receivesUpdates.Store(true)
|
||
sm.Register(c)
|
||
}
|
||
|
||
encodes := 0
|
||
sent, err := sm.PushToUserExceptSessionBestEffort(
|
||
context.Background(),
|
||
userID,
|
||
0,
|
||
proto.MessageFromServer,
|
||
&countingOutboundEncoder{count: &encodes},
|
||
0,
|
||
)
|
||
if err != nil {
|
||
t.Fatalf("push: %v", err)
|
||
}
|
||
if sent != 2 {
|
||
t.Fatalf("sent = %d, want 2", sent)
|
||
}
|
||
if encodes != 1 {
|
||
t.Fatalf("encoded %d times, want 1", encodes)
|
||
}
|
||
}
|
||
|
||
func TestSessionManagerScopesSameSessionIDByAuthKey(t *testing.T) {
|
||
sm := NewSessionManager(zaptest.NewLogger(t))
|
||
raw1 := [8]byte{1}
|
||
raw2 := [8]byte{2}
|
||
perm1 := [8]byte{9}
|
||
c1 := &Conn{sessionID: 42, authKeyID: raw1}
|
||
c2 := &Conn{sessionID: 42, authKeyID: raw2}
|
||
|
||
sm.Register(c1)
|
||
sm.Register(c2)
|
||
if got := sm.Online(); got != 2 {
|
||
t.Fatalf("online = %d, want 2", got)
|
||
}
|
||
|
||
sm.BindAuthKeyForSession(raw1, 42, perm1)
|
||
sm.BindUserForAuthKey(raw1, 42, 100)
|
||
sm.BindUserForAuthKey(raw2, 42, 200)
|
||
|
||
if userID, ok := sm.UserIDForAuthKey(raw1, 42); !ok || userID != 100 {
|
||
t.Fatalf("scoped user raw1 = %d ok %v, want 100/true", userID, ok)
|
||
}
|
||
if userID, ok := sm.UserIDForAuthKey(raw2, 42); !ok || userID != 200 {
|
||
t.Fatalf("scoped user raw2 = %d ok %v, want 200/true", userID, ok)
|
||
}
|
||
if _, ok := sm.UserID(42); ok {
|
||
t.Fatal("legacy UserID unexpectedly resolved ambiguous session_id")
|
||
}
|
||
if err := sm.PushToSession(context.Background(), 42, proto.MessageFromServer, &tg.UpdatesTooLong{}); !errors.Is(err, ErrSessionAmbiguous) {
|
||
t.Fatalf("ambiguous push err = %v, want ErrSessionAmbiguous", err)
|
||
}
|
||
|
||
sm.BindUserForAuthKey(raw1, 42, 300)
|
||
sm.BindUserForAuthKey(raw2, 42, 300)
|
||
sent, err := sm.PushToUserExceptAuthKeySession(context.Background(), 300, perm1, 42, proto.MessageFromServer, &tg.UpdatesTooLong{})
|
||
if err != nil {
|
||
t.Fatalf("push except scoped session: %v", err)
|
||
}
|
||
if sent != 1 {
|
||
t.Fatalf("pushed to %d sessions, want 1", sent)
|
||
}
|
||
if _, ok := sm.pending[sessionKey{authKeyID: raw1, sessionID: 42}]; ok {
|
||
t.Fatal("excluded session received pending push")
|
||
}
|
||
if got := len(sm.pending[sessionKey{authKeyID: raw2, sessionID: 42}]); got != 1 {
|
||
t.Fatalf("raw2 pending pushes = %d, want 1", got)
|
||
}
|
||
}
|
||
|
||
func TestSessionManagerCloseSessionsForBusinessAuthKeyClosesBoundTempAndRaw(t *testing.T) {
|
||
sm := NewSessionManager(zaptest.NewLogger(t))
|
||
rawTemp := [8]byte{1}
|
||
perm := [8]byte{9}
|
||
otherRaw := [8]byte{2}
|
||
otherPerm := [8]byte{8}
|
||
tempTransport := &closeCountingTransport{}
|
||
permTransport := &closeCountingTransport{}
|
||
otherTransport := &closeCountingTransport{}
|
||
cTemp := &Conn{sessionID: 11, authKeyID: rawTemp, transport: tempTransport}
|
||
cPerm := &Conn{sessionID: 12, authKeyID: perm, transport: permTransport}
|
||
cOther := &Conn{sessionID: 13, authKeyID: otherRaw}
|
||
cOther.transport = otherTransport
|
||
|
||
sm.Register(cTemp)
|
||
sm.Register(cPerm)
|
||
sm.Register(cOther)
|
||
sm.BindAuthKeyForSession(rawTemp, 11, perm)
|
||
sm.BindAuthKeyForSession(perm, 12, perm)
|
||
sm.BindAuthKeyForSession(otherRaw, 13, otherPerm)
|
||
sm.BindUserForAuthKey(rawTemp, 11, 100)
|
||
sm.BindUserForAuthKey(perm, 12, 100)
|
||
sm.BindUserForAuthKey(otherRaw, 13, 200)
|
||
|
||
if closed := sm.CloseSessionsForBusinessAuthKey(perm); closed != 2 {
|
||
t.Fatalf("closed sessions = %d, want 2", closed)
|
||
}
|
||
if tempTransport.closes != 1 || permTransport.closes != 1 {
|
||
t.Fatalf("transport closes temp=%d perm=%d, want 1/1", tempTransport.closes, permTransport.closes)
|
||
}
|
||
if otherTransport.closes != 0 {
|
||
t.Fatalf("other transport closes = %d, want 0", otherTransport.closes)
|
||
}
|
||
if got := sm.Online(); got != 1 {
|
||
t.Fatalf("online after close = %d, want 1", got)
|
||
}
|
||
if _, ok := sm.AuthKeyIDForSession(rawTemp, 11); ok {
|
||
t.Fatal("temp session still indexed after business auth key close")
|
||
}
|
||
if _, ok := sm.AuthKeyIDForSession(perm, 12); ok {
|
||
t.Fatal("raw perm session still indexed after business auth key close")
|
||
}
|
||
if userID, ok := sm.UserIDForAuthKey(otherRaw, 13); !ok || userID != 200 {
|
||
t.Fatalf("other session user = %d ok %v, want 200/true", userID, ok)
|
||
}
|
||
if closed := sm.CloseSessionsForBusinessAuthKey(perm); closed != 0 {
|
||
t.Fatalf("second close = %d, want 0", closed)
|
||
}
|
||
}
|
||
|
||
func TestSessionManagerBusinessAuthKeyIndexTracksRebind(t *testing.T) {
|
||
sm := NewSessionManager(zaptest.NewLogger(t))
|
||
raw := [8]byte{1}
|
||
oldPerm := [8]byte{7}
|
||
newPerm := [8]byte{8}
|
||
c := &Conn{sessionID: 21, authKeyID: raw}
|
||
|
||
sm.Register(c)
|
||
sm.BindAuthKeyForSession(raw, 21, oldPerm)
|
||
sm.BindAuthKeyForSession(raw, 21, newPerm)
|
||
|
||
if closed := sm.CloseSessionsForBusinessAuthKey(oldPerm); closed != 0 {
|
||
t.Fatalf("close old business auth key = %d, want 0", closed)
|
||
}
|
||
if got := sm.Online(); got != 1 {
|
||
t.Fatalf("online after closing old key = %d, want 1", got)
|
||
}
|
||
if closed := sm.CloseSessionsForBusinessAuthKey(newPerm); closed != 1 {
|
||
t.Fatalf("close new business auth key = %d, want 1", closed)
|
||
}
|
||
if got := sm.Online(); got != 0 {
|
||
t.Fatalf("online after closing new key = %d, want 0", got)
|
||
}
|
||
}
|
||
|
||
func TestSessionManagerChannelInterestIndex(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.TrackChannelInterest(raw, 42, 100, []int64{10, 10, 20})
|
||
if got := sm.OnlineChannelUserIDs(10, 10); len(got) != 1 || got[0] != 100 {
|
||
t.Fatalf("channel 10 online users = %v, want [100]", got)
|
||
}
|
||
sm.TrackChannelInterest(raw, 42, 100, []int64{20})
|
||
if got := sm.OnlineChannelUserIDs(10, 10); len(got) != 0 {
|
||
t.Fatalf("channel 10 after viewer switch = %v, want empty", got)
|
||
}
|
||
if got := sm.OnlineChannelUserIDs(20, 10); len(got) != 1 || got[0] != 100 {
|
||
t.Fatalf("channel 20 after viewer switch = %v, want [100]", got)
|
||
}
|
||
sm.TrackChannelInterest(raw, 42, 100, []int64{10})
|
||
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})
|
||
if got := sm.OnlineChannelMemberUserIDs(10, 10); len(got) != 1 || got[0] != 100 {
|
||
t.Fatalf("channel 10 online members = %v, want [100]", got)
|
||
}
|
||
if got := sm.OnlineChannelUserIDs(30, 10); len(got) != 0 {
|
||
t.Fatalf("channel 30 viewers = %v, want empty", got)
|
||
}
|
||
if got := sm.OnlineUserIDsForCandidates([]int64{0, 200, 100, 100}, 10); len(got) != 1 || got[0] != 100 {
|
||
t.Fatalf("candidate online users = %v, want [100]", got)
|
||
}
|
||
|
||
sm.BindUserForAuthKey(raw, 42, 200)
|
||
if got := sm.OnlineChannelUserIDs(10, 10); len(got) != 0 {
|
||
t.Fatalf("channel interest after user switch = %v, want empty", got)
|
||
}
|
||
if got := sm.OnlineChannelMemberUserIDs(10, 10); len(got) != 0 {
|
||
t.Fatalf("channel membership after user switch = %v, want empty", got)
|
||
}
|
||
sm.TrackChannelInterest(raw, 42, 200, []int64{10})
|
||
if got := sm.OnlineChannelUserIDs(10, 10); len(got) != 1 || got[0] != 200 {
|
||
t.Fatalf("channel 10 after re-track = %v, want [200]", got)
|
||
}
|
||
sm.AddUserChannelMembership(200, 10)
|
||
if got := sm.OnlineChannelMemberUserIDs(10, 10); len(got) != 1 || got[0] != 200 {
|
||
t.Fatalf("channel 10 membership after add = %v, want [200]", got)
|
||
}
|
||
sm.RemoveUserChannelMembership(200, 10)
|
||
if got := sm.OnlineChannelMemberUserIDs(10, 10); len(got) != 0 {
|
||
t.Fatalf("channel membership after remove = %v, want empty", got)
|
||
}
|
||
sm.ClearChannelInterest(raw, 42, 200)
|
||
if got := sm.OnlineChannelUserIDs(10, 10); len(got) != 0 {
|
||
t.Fatalf("channel interest after explicit clear = %v, want empty", got)
|
||
}
|
||
|
||
sm.Unregister(c)
|
||
if got := sm.OnlineChannelUserIDs(10, 10); len(got) != 0 {
|
||
t.Fatalf("channel interest after unregister = %v, want empty", got)
|
||
}
|
||
if got := sm.OnlineChannelMemberUserIDs(10, 10); len(got) != 0 {
|
||
t.Fatalf("channel membership after unregister = %v, want empty", got)
|
||
}
|
||
}
|
||
|
||
func TestSessionManagerClearsChannelIndexesOnAuthAndReadinessChanges(t *testing.T) {
|
||
sm := NewSessionManager(zaptest.NewLogger(t))
|
||
raw := [8]byte{1, 2, 3}
|
||
business := [8]byte{8}
|
||
c := &Conn{sessionID: 42, authKeyID: raw}
|
||
sm.Register(c)
|
||
sm.BindAuthKeyForSession(raw, 42, business)
|
||
sm.BindUserForAuthKey(raw, 42, 100)
|
||
|
||
track := func() {
|
||
sm.TrackChannelInterest(raw, 42, 100, []int64{10})
|
||
sm.SetSessionChannelMemberships(raw, 42, 100, []int64{10})
|
||
if got := sm.OnlineChannelUserIDs(10, 10); len(got) != 1 || got[0] != 100 {
|
||
t.Fatalf("channel viewers before cleanup = %v, want [100]", got)
|
||
}
|
||
if got := sm.OnlineChannelMemberUserIDs(10, 10); len(got) != 1 || got[0] != 100 {
|
||
t.Fatalf("channel members before cleanup = %v, want [100]", got)
|
||
}
|
||
}
|
||
assertCleared := func(label string) {
|
||
if got := sm.OnlineChannelUserIDs(10, 10); len(got) != 0 {
|
||
t.Fatalf("%s viewers = %v, want empty", label, got)
|
||
}
|
||
if got := sm.OnlineChannelMemberUserIDs(10, 10); len(got) != 0 {
|
||
t.Fatalf("%s members = %v, want empty", label, got)
|
||
}
|
||
}
|
||
|
||
track()
|
||
sm.SetReceivesUpdatesForAuthKey(raw, 42, false)
|
||
assertCleared("after receivesUpdates=false")
|
||
|
||
track()
|
||
sm.BindAuthKeyForSession(raw, 42, [8]byte{9})
|
||
assertCleared("after business auth key change")
|
||
|
||
sm.BindAuthKeyForSession(raw, 42, business)
|
||
sm.BindUserForAuthKey(raw, 42, 100)
|
||
track()
|
||
if n := sm.UnbindAuthKey(business); n != 1 {
|
||
t.Fatalf("UnbindAuthKey count = %d, want 1", n)
|
||
}
|
||
assertCleared("after unbind auth key")
|
||
}
|
||
|
||
func TestPushToSessionForAuthKeyImmediateBypassesReadinessQueue(t *testing.T) {
|
||
sm := NewSessionManager(zaptest.NewLogger(t))
|
||
raw := [8]byte{1, 2, 3}
|
||
c := &Conn{
|
||
sessionID: 42,
|
||
authKeyID: raw,
|
||
outbound: make(chan outboundOp, 1),
|
||
outboundControl: make(chan outboundOp, 1),
|
||
outboundStop: make(chan struct{}),
|
||
}
|
||
sm.Register(c)
|
||
|
||
msg := &tg.UpdateShort{Update: &tg.UpdateLoginToken{}, Date: 1700000000}
|
||
if err := sm.PushToSessionForAuthKeyImmediate(context.Background(), raw, 42, proto.MessageFromServer, msg); err != nil {
|
||
t.Fatalf("immediate push: %v", err)
|
||
}
|
||
|
||
select {
|
||
case op := <-c.outbound:
|
||
if op.msg != msg {
|
||
t.Fatalf("enqueued msg = %T, want original update", op.msg)
|
||
}
|
||
case <-time.After(time.Second):
|
||
t.Fatal("immediate push was not enqueued")
|
||
}
|
||
|
||
sm.mu.RLock()
|
||
pending := len(sm.pending[sessionKey{authKeyID: raw, sessionID: 42}])
|
||
sm.mu.RUnlock()
|
||
if pending != 0 {
|
||
t.Fatalf("pending pushes = %d, want 0", pending)
|
||
}
|
||
}
|
||
|
||
// TestSessionManagerPush 验证主动推送端到端:两个 client 连接握手并建立 session 后,
|
||
// server 经 PushToSession / PushToUser 主动向其推送,client 收到。
|
||
func TestSessionManagerPush(t *testing.T) {
|
||
const dc = 2
|
||
addr, pub, srv := startTestServer(t, Options{DC: dc})
|
||
|
||
conn1, auth1, cipher1 := dialHandshake(t, addr, dc, pub)
|
||
conn2, auth2, cipher2 := dialHandshake(t, addr, dc, pub)
|
||
|
||
// 各发一个 ping 建立 session,触发注册(并清掉 new_session_created/pong/ack)。
|
||
msgGen := proto.NewMessageIDGen(time.Now)
|
||
sendEncrypted(t, conn1, cipher1, auth1, msgGen.New(proto.MessageFromClient), &mt.PingRequest{PingID: 1})
|
||
collectReplies(t, conn1, cipher1, auth1.AuthKey, mt.PongTypeID)
|
||
sendEncrypted(t, conn2, cipher2, auth2, msgGen.New(proto.MessageFromClient), &mt.PingRequest{PingID: 2})
|
||
collectReplies(t, conn2, cipher2, auth2.AuthKey, mt.PongTypeID)
|
||
|
||
if got := srv.Conns().Online(); got != 2 {
|
||
t.Fatalf("online = %d, want 2", got)
|
||
}
|
||
|
||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||
defer cancel()
|
||
|
||
// 1) PushToSession:session2 尚未进入 updates 同步入口时先暂存,ready 后下发。
|
||
if err := srv.Conns().PushToSession(ctx, auth2.SessionID, proto.MessageFromServer, &tg.UpdatesTooLong{}); err != nil {
|
||
t.Fatalf("push to session: %v", err)
|
||
}
|
||
srv.Conns().SetReceivesUpdates(auth2.SessionID, true)
|
||
r2 := collectReplies(t, conn2, cipher2, auth2.AuthKey, tg.UpdatesTooLongTypeID)
|
||
mustHave(t, r2, tg.UpdatesTooLongTypeID, "pushed updates on conn2")
|
||
|
||
// 2) BindUser + PushToUser:按 user 维度推送给 conn1。
|
||
srv.Conns().BindUser(auth1.SessionID, 100)
|
||
srv.Conns().SetReceivesUpdates(auth1.SessionID, true)
|
||
sent, err := srv.Conns().PushToUser(ctx, 100, proto.MessageFromServer, &tg.UpdatesTooLong{})
|
||
if err != nil {
|
||
t.Fatalf("push to user: %v", err)
|
||
}
|
||
if sent != 1 {
|
||
t.Fatalf("pushed to %d conns, want 1", sent)
|
||
}
|
||
r1 := collectReplies(t, conn1, cipher1, auth1.AuthKey, tg.UpdatesTooLongTypeID)
|
||
mustHave(t, r1, tg.UpdatesTooLongTypeID, "pushed updates on conn1")
|
||
|
||
// 3) PushToUserExceptSession:模拟 SyncUpdatesNotMe,跳过当前 session。
|
||
srv.Conns().BindUser(auth1.SessionID, 200)
|
||
srv.Conns().BindUser(auth2.SessionID, 200)
|
||
sent, err = srv.Conns().PushToUserExceptSession(ctx, 200, auth2.SessionID, proto.MessageFromServer, &tg.UpdatesTooLong{})
|
||
if err != nil {
|
||
t.Fatalf("push to user except session: %v", err)
|
||
}
|
||
if sent != 1 {
|
||
t.Fatalf("pushed to %d conns, want 1 after excluding current session", sent)
|
||
}
|
||
r1 = collectReplies(t, conn1, cipher1, auth1.AuthKey, tg.UpdatesTooLongTypeID)
|
||
mustHave(t, r1, tg.UpdatesTooLongTypeID, "pushed not-me updates on conn1")
|
||
}
|
||
|
||
func BenchmarkSessionManagerOnlineCandidateFilter(b *testing.B) {
|
||
sm := NewSessionManager(zaptest.NewLogger(b))
|
||
const online = 200_000
|
||
rawPrefix := [8]byte{9}
|
||
for i := 1; i <= online; i++ {
|
||
raw := rawPrefix
|
||
raw[1] = byte(i)
|
||
raw[2] = byte(i >> 8)
|
||
raw[3] = byte(i >> 16)
|
||
raw[4] = byte(i >> 24)
|
||
c := &Conn{sessionID: int64(i), authKeyID: raw}
|
||
sm.Register(c)
|
||
sm.BindUserForAuthKey(raw, int64(i), int64(i))
|
||
}
|
||
candidates := make([]int64, 0, 500)
|
||
for i := 0; i < 500; i++ {
|
||
candidates = append(candidates, int64(i*97+1))
|
||
}
|
||
b.ResetTimer()
|
||
for i := 0; i < b.N; i++ {
|
||
got := sm.OnlineUserIDsForCandidates(candidates, 500)
|
||
if len(got) == 0 {
|
||
b.Fatal("no candidates matched")
|
||
}
|
||
}
|
||
}
|