fix(mtproto): sync accept TDLib reconnect service parity

This commit is contained in:
iamxvbaba 2026-08-02 01:51:03 +08:00
parent c854fc7b94
commit c13fbf8885
9 changed files with 176 additions and 41 deletions

View file

@ -15,7 +15,7 @@ func TestEncryptedConnectionSwitchesAuthKeyEvenWhenSessionIDIsReused(t *testing.
_, authB, cipherB := dialHandshake(t, addr, dc, pub)
msgID := proto.NewMessageIDGen(time.Now)
sendEncrypted(t, connA, cipherA, authA, msgID.New(proto.MessageFromClient), &mt.PingRequest{PingID: 1})
sendEncryptedWithSeq(t, connA, cipherA, authA, msgID.New(proto.MessageFromClient), 1, &mt.PingRequest{PingID: 1})
for range 3 { // new_session_created + pong + msgs_ack; leave no A-key frame on the socket.
readServerMessage(t, connA, cipherA, authA.AuthKey)
}

View file

@ -1481,15 +1481,15 @@ func validateClientEnvelope(now time.Time, msgID int64, seqNo int32, typeID uint
if msgTime.After(now.Add(30 * time.Second)) {
return badMsgIDTooHigh
}
if clientMessageAllowsEitherSeqParity(typeID) {
return 0
}
if clientMessageNeedsAck(typeID) {
switch clientMessageContentPolicyFor(typeID) {
case clientMessageContentRequired:
if seqNo%2 == 0 {
return badMsgSeqNotOdd
}
} else if seqNo%2 != 0 {
return badMsgSeqNotEven
case clientMessageContentForbidden:
if seqNo%2 != 0 {
return badMsgSeqNotEven
}
}
return 0
}
@ -1498,48 +1498,68 @@ func validateClientContainerEnvelope(msgID int64, seqNo int32, typeID uint32) in
if !validClientMessageIDBits(msgID) {
return badMsgIDInvalidBits
}
if clientMessageAllowsEitherSeqParity(typeID) {
return 0
}
if clientMessageNeedsAck(typeID) {
switch clientMessageContentPolicyFor(typeID) {
case clientMessageContentRequired:
if seqNo%2 == 0 {
return badMsgSeqNotOdd
}
} else if seqNo%2 != 0 {
return badMsgSeqNotEven
case clientMessageContentForbidden:
if seqNo%2 != 0 {
return badMsgSeqNotEven
}
}
return 0
}
func clientMessageAllowsEitherSeqParity(typeID uint32) bool {
switch typeID {
case mt.PingDelayDisconnectRequestTypeID,
// get_future_salts 的 seqno 奇偶在客户端间不一致:部分客户端按内容消息发奇数,
// gotd 按服务消息发偶数。两者都合法(官方服务器都接受),故不在此卡奇偶,避免
// 误判 bad_msg 触发客户端重连风暴。ack/content 行为仍由 clientMessageNeedsAck 决定。
mt.GetFutureSaltsRequestTypeID:
return true
default:
return false
}
}
type clientMessageContentPolicy uint8
func clientMessageNeedsAck(typeID uint32) bool {
const (
clientMessageContentRequired clientMessageContentPolicy = iota + 1
clientMessageContentForbidden
clientMessageContentOptional
)
// clientMessageContentPolicyFor classifies the client envelope, not merely the
// constructor's usual sending convention. MTProto requires API RPCs to be
// content-related and requires containers/acknowledgements to be irrelevant,
// but clients may mark the other service constructors as either. TDLib uses
// even sequence numbers for its reconnect state/resend/cancel service batch,
// while gotd and DrKLO use odd sequence numbers for some of the same requests.
func clientMessageContentPolicyFor(typeID uint32) clientMessageContentPolicy {
switch typeID {
case proto.MessageContainerTypeID,
mt.MsgsAckTypeID,
mt.MsgCopyTypeID:
return clientMessageContentForbidden
case mt.PingRequestTypeID,
mt.PingDelayDisconnectRequestTypeID,
mt.DestroySessionRequestTypeID,
mt.HTTPWaitRequestTypeID,
mt.BadMsgNotificationTypeID,
mt.BadServerSaltTypeID,
mt.GetFutureSaltsRequestTypeID,
mt.MsgsStateReqTypeID,
mt.MsgResendReqTypeID,
mt.MsgsAllInfoTypeID,
mt.MsgsStateInfoTypeID,
mt.DestroySessionRequestTypeID,
mt.HTTPWaitRequestTypeID,
mt.RPCDropAnswerRequestTypeID,
mt.BadMsgNotificationTypeID,
mt.BadServerSaltTypeID,
mt.MsgDetailedInfoTypeID,
mt.MsgNewDetailedInfoTypeID:
return false
mt.MsgNewDetailedInfoTypeID,
destroyAuthKeyRequestTypeID:
return clientMessageContentOptional
default:
return clientMessageContentRequired
}
}
func clientMessageIsContentRelated(typeID uint32, seqNo int32) bool {
switch clientMessageContentPolicyFor(typeID) {
case clientMessageContentRequired:
return true
case clientMessageContentOptional:
return seqNo%2 != 0
default:
return false
}
}

View file

@ -69,6 +69,72 @@ func (h *durableDestroyLayerRPC) deletion() ([8]byte, int64) {
return h.authKeyID, h.sessionID
}
func TestClientMessageContentPolicy(t *testing.T) {
tests := []struct {
name string
typeID uint32
want clientMessageContentPolicy
}{
{name: "api_rpc", typeID: tg.HelpGetConfigRequestTypeID, want: clientMessageContentRequired},
{name: "container", typeID: proto.MessageContainerTypeID, want: clientMessageContentForbidden},
{name: "msgs_ack", typeID: mt.MsgsAckTypeID, want: clientMessageContentForbidden},
{name: "msg_copy", typeID: mt.MsgCopyTypeID, want: clientMessageContentForbidden},
{name: "bad_msg_notification", typeID: mt.BadMsgNotificationTypeID, want: clientMessageContentOptional},
{name: "bad_server_salt", typeID: mt.BadServerSaltTypeID, want: clientMessageContentOptional},
{name: "msg_detailed_info", typeID: mt.MsgDetailedInfoTypeID, want: clientMessageContentOptional},
{name: "msg_new_detailed_info", typeID: mt.MsgNewDetailedInfoTypeID, want: clientMessageContentOptional},
{name: "ping", typeID: mt.PingRequestTypeID, want: clientMessageContentOptional},
{name: "ping_delay_disconnect", typeID: mt.PingDelayDisconnectRequestTypeID, want: clientMessageContentOptional},
{name: "get_future_salts", typeID: mt.GetFutureSaltsRequestTypeID, want: clientMessageContentOptional},
{name: "msgs_state_req", typeID: mt.MsgsStateReqTypeID, want: clientMessageContentOptional},
{name: "msg_resend_req", typeID: mt.MsgResendReqTypeID, want: clientMessageContentOptional},
{name: "msgs_all_info", typeID: mt.MsgsAllInfoTypeID, want: clientMessageContentOptional},
{name: "msgs_state_info", typeID: mt.MsgsStateInfoTypeID, want: clientMessageContentOptional},
{name: "destroy_session", typeID: mt.DestroySessionRequestTypeID, want: clientMessageContentOptional},
{name: "http_wait", typeID: mt.HTTPWaitRequestTypeID, want: clientMessageContentOptional},
{name: "rpc_drop_answer", typeID: mt.RPCDropAnswerRequestTypeID, want: clientMessageContentOptional},
{name: "destroy_auth_key", typeID: destroyAuthKeyRequestTypeID, want: clientMessageContentOptional},
}
now := time.Now()
msgID := proto.NewMessageIDGen(func() time.Time { return now }).New(proto.MessageFromClient)
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if got := clientMessageContentPolicyFor(test.typeID); got != test.want {
t.Fatalf("content policy = %d, want %d", got, test.want)
}
evenCode := validateClientContainerEnvelope(msgID, 8, test.typeID)
oddCode := validateClientContainerEnvelope(msgID, 9, test.typeID)
directEvenCode := validateClientEnvelope(now, msgID, 8, test.typeID)
directOddCode := validateClientEnvelope(now, msgID, 9, test.typeID)
if directEvenCode != evenCode || directOddCode != oddCode {
t.Fatalf(
"top-level/container parity mismatch = top(%d,%d) container(%d,%d)",
directEvenCode, directOddCode, evenCode, oddCode,
)
}
switch test.want {
case clientMessageContentRequired:
if evenCode != badMsgSeqNotOdd || oddCode != 0 {
t.Fatalf("required content parity codes = even:%d odd:%d", evenCode, oddCode)
}
case clientMessageContentForbidden:
if evenCode != 0 || oddCode != badMsgSeqNotEven {
t.Fatalf("forbidden content parity codes = even:%d odd:%d", evenCode, oddCode)
}
case clientMessageContentOptional:
if evenCode != 0 || oddCode != 0 {
t.Fatalf("optional content parity codes = even:%d odd:%d", evenCode, oddCode)
}
if clientMessageIsContentRelated(test.typeID, 8) || !clientMessageIsContentRelated(test.typeID, 9) {
t.Fatal("optional service content bit was not derived from seq_no parity")
}
}
})
}
}
// TestEncryptedPingPong 验证 M2/M4握手后 client 加密 ping
// server 回 new_session_created + pong + msgs_ack。
func TestEncryptedPingPong(t *testing.T) {
@ -79,7 +145,7 @@ func TestEncryptedPingPong(t *testing.T) {
clientMsgID := proto.NewMessageIDGen(time.Now)
const pingID int64 = 0x1234beef
pingMsgID := clientMsgID.New(proto.MessageFromClient)
sendEncrypted(t, conn, cipher, auth, pingMsgID, &mt.PingRequest{PingID: pingID})
sendEncryptedWithSeq(t, conn, cipher, auth, pingMsgID, 1, &mt.PingRequest{PingID: pingID})
replies := collectReplies(t, conn, cipher, auth.AuthKey, mt.PongTypeID)
mustHave(t, replies, mt.NewSessionCreatedTypeID, "new_session_created")
@ -230,6 +296,55 @@ func TestMsgsStateReq(t *testing.T) {
}
}
// TestTDLibReconnectRecoveryContainerAcceptsEvenServiceMessages reproduces the
// first container TDLib emits after reopening an authenticated session with
// unknown queries. All service entries and the outer container use the current
// even sequence number. Rejecting msgs_state_req as a content-only constructor
// turns the valid inner message into bad_msg_notification(code=64), after which
// TDLib closes the session and retries the same container forever.
func TestTDLibReconnectRecoveryContainerAcceptsEvenServiceMessages(t *testing.T) {
const dc = 2
addr, pub, _ := startTestServer(t, Options{DC: dc})
conn, auth, cipher := dialHandshake(t, addr, dc, pub)
ids := proto.NewMessageIDGen(time.Now)
pendingMsgID := ids.New(proto.MessageFromClient)
ackMsgID := ids.New(proto.MessageFromClient)
stateMsgID := ids.New(proto.MessageFromClient)
pingMsgID := ids.New(proto.MessageFromClient)
outerMsgID := ids.New(proto.MessageFromClient)
const (
pendingSeqNo int32 = 9
serviceSeqNo int32 = 10
)
pendingBody := mustEncodeTL(t, &mt.PingRequest{PingID: 0xc01d})
ackBody := mustEncodeTL(t, &mt.MsgsAck{MsgIDs: []int64{stateMsgID - 4}})
stateBody := mustEncodeTL(t, &mt.MsgsStateReq{MsgIDs: []int64{pendingMsgID, stateMsgID - 4}})
pingBody := mustEncodeTL(t, &mt.PingDelayDisconnectRequest{PingID: 0x5eed, DisconnectDelay: 60})
container := &proto.MessageContainer{Messages: []proto.Message{
{ID: pendingMsgID, SeqNo: int(pendingSeqNo), Bytes: len(pendingBody), Body: pendingBody},
{ID: ackMsgID, SeqNo: int(serviceSeqNo), Bytes: len(ackBody), Body: ackBody},
{ID: stateMsgID, SeqNo: int(serviceSeqNo), Bytes: len(stateBody), Body: stateBody},
{ID: pingMsgID, SeqNo: int(serviceSeqNo), Bytes: len(pingBody), Body: pingBody},
}}
sendEncryptedWithSeq(t, conn, cipher, auth, outerMsgID, serviceSeqNo, container)
frames := collectReplyFrames(t, conn, cipher, auth.AuthKey, map[uint32]int{
mt.MsgsStateInfoTypeID: 1,
mt.PongTypeID: 2,
})
for _, frame := range frames {
if frame.TypeID == mt.BadMsgNotificationTypeID {
var bad mt.BadMsgNotification
if err := bad.Decode(frame.Plain); err != nil {
t.Fatalf("decode bad_msg_notification: %v", err)
}
t.Fatalf("TDLib reconnect recovery container was rejected: %+v", bad)
}
}
}
// TestMsgResendReq 验证 MTProto msg_resend_req 由连接层按状态查询兜底响应,
// 不会落入业务 RPC fallback。
func TestMsgResendReq(t *testing.T) {

View file

@ -1044,7 +1044,7 @@ func TestReconnectFakeReqPQThenEncryptedFrame(t *testing.T) {
cancel()
msgGen := tgproto.NewMessageIDGen(time.Now)
sendEncrypted(t, conn, cipher, auth, msgGen.New(tgproto.MessageFromClient), &mt.PingRequest{PingID: 7})
sendEncryptedWithSeq(t, conn, cipher, auth, msgGen.New(tgproto.MessageFromClient), 1, &mt.PingRequest{PingID: 7})
var resPQFrame bin.Buffer
ctx, cancel = context.WithTimeout(context.Background(), 5*time.Second)

View file

@ -250,7 +250,7 @@ func encodeClientMessageForTest(t *testing.T, msg bin.Encoder) ([]byte, int32) {
if container, ok := msg.(*proto.MessageContainer); ok {
return raw, clientContainerSeqNoForTest(container)
}
if clientMessageNeedsAck(typeID) {
if clientMessageContentPolicyFor(typeID) == clientMessageContentRequired {
return raw, 1
}
return raw, 0

View file

@ -454,7 +454,7 @@ func (s *Server) walkInbound(
return &dispatchBadMsgError{msgID: msgID, seqNo: seqNo, code: code}
}
content := clientMessageNeedsAck(typeID)
content := clientMessageIsContentRelated(typeID, seqNo)
if record, seen := overlay.seenRecord(msgID); seen {
if record.seqNo != seqNo || record.content != content {
return &dispatchBadMsgError{msgID: msgID, seqNo: seqNo, code: badMsgContainer}

View file

@ -28,7 +28,7 @@ func TestRegisterSeedsNegotiatedLayerBeforeFirstRPC(t *testing.T) {
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})
sendEncryptedWithSeq(t, conn, cipher, auth, clientMsgID.New(proto.MessageFromClient), 1, &mt.PingRequest{PingID: 7})
// 等 pong 回来,确保携带注册动作的那一帧已处理完成。
gotPong := false

View file

@ -544,7 +544,7 @@ func TestOutboundActorSerializesConcurrentSends(t *testing.T) {
conn, auth, cipher := dialHandshake(t, addr, dc, pub)
clientMsgID := proto.NewMessageIDGen(time.Now)
sendEncrypted(t, conn, cipher, auth, clientMsgID.New(proto.MessageFromClient), &mt.PingRequest{PingID: 1})
sendEncryptedWithSeq(t, conn, cipher, auth, clientMsgID.New(proto.MessageFromClient), 1, &mt.PingRequest{PingID: 1})
collectReplies(t, conn, cipher, auth.AuthKey, mt.MsgsAckTypeID)
freezeActiveTestSessionProfile(t, srv.Conns(), auth.AuthKey.ID, auth.SessionID, tlprofile.ProfileCanonical)
srv.Conns().SetReceivesUpdates(auth.SessionID, true)
@ -1034,7 +1034,7 @@ func TestOutboundResendAndAckState(t *testing.T) {
conn, auth, cipher := dialHandshake(t, addr, dc, pub)
clientMsgID := proto.NewMessageIDGen(time.Now)
sendEncrypted(t, conn, cipher, auth, clientMsgID.New(proto.MessageFromClient), &mt.PingRequest{PingID: 1})
sendEncryptedWithSeq(t, conn, cipher, auth, clientMsgID.New(proto.MessageFromClient), 1, &mt.PingRequest{PingID: 1})
collectReplies(t, conn, cipher, auth.AuthKey, mt.MsgsAckTypeID)
freezeActiveTestSessionProfile(t, srv.Conns(), auth.AuthKey.ID, auth.SessionID, tlprofile.ProfileCanonical)
srv.Conns().SetReceivesUpdates(auth.SessionID, true)

View file

@ -1205,9 +1205,9 @@ func TestSessionManagerPush(t *testing.T) {
// 各发一个 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})
sendEncryptedWithSeq(t, conn1, cipher1, auth1, msgGen.New(proto.MessageFromClient), 1, &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})
sendEncryptedWithSeq(t, conn2, cipher2, auth2, msgGen.New(proto.MessageFromClient), 1, &mt.PingRequest{PingID: 2})
collectReplies(t, conn2, cipher2, auth2.AuthKey, mt.PongTypeID)
if got := srv.Conns().Online(); got != 2 {