From 9cb76b7c5a6c257bbe540c24fe77b438d97b8a6d Mon Sep 17 00:00:00 2001 From: iamxvbaba <28732408+iamxvbaba@users.noreply.github.com> Date: Fri, 31 Jul 2026 15:35:10 +0800 Subject: [PATCH] fix: sync admit TDLib handshake message id sentinels --- internal/mtprotoedge/exchange_compat.go | 36 +++- internal/mtprotoedge/exchange_test.go | 261 ++++++++++++++++++++++++ 2 files changed, 296 insertions(+), 1 deletion(-) diff --git a/internal/mtprotoedge/exchange_compat.go b/internal/mtprotoedge/exchange_compat.go index 56db3218..51c63331 100644 --- a/internal/mtprotoedge/exchange_compat.go +++ b/internal/mtprotoedge/exchange_compat.go @@ -434,7 +434,7 @@ func (s serverExchangeCompat) readUnencrypted(ctx context.Context, b *bin.Buffer if err := msg.Decode(b); err != nil { return err } - if !validClientMessageIDBits(msg.MessageID) { + if !validUnencryptedHandshakeMessageID(msg.MessageID, msg.MessageData) { return gofaster.New("bad msg type") } b.ResetTo(msg.MessageData) @@ -442,6 +442,40 @@ func (s serverExchangeCompat) readUnencrypted(ctx context.Context, b *bin.Buffer return data.Decode(b) } +// validUnencryptedHandshakeMessageID preserves the normal client message-id +// rules while admitting the two sentinel ids emitted by official TDLib: +// +// - PingConnectionReqPQ sends req_pq_multi with message_id=1. +// - HandshakeConnection sends every auth-key exchange request with message_id=0. +// +// The exception is deliberately constructor-scoped and is only called after an +// auth_key_id=0 envelope has been decoded. Encrypted traffic continues through +// validClientMessageIDBits and the full inbound preflight without this carve-out. +func validUnencryptedHandshakeMessageID(messageID int64, messageData []byte) bool { + if validClientMessageIDBits(messageID) { + return true + } + + payload := &bin.Buffer{Buf: messageData} + typeID, err := payload.PeekID() + if err != nil { + return false + } + + switch messageID { + case 1: + return typeID == mt.ReqPqMultiRequestTypeID + case 0: + switch typeID { + case mt.ReqPqMultiRequestTypeID, + mt.ReqDHParamsRequestTypeID, + mt.SetClientDHParamsRequestTypeID: + return true + } + } + return false +} + type compatReqPQ struct { Type uint32 Nonce bin.Int128 diff --git a/internal/mtprotoedge/exchange_test.go b/internal/mtprotoedge/exchange_test.go index d98a96db..ff2044ea 100644 --- a/internal/mtprotoedge/exchange_test.go +++ b/internal/mtprotoedge/exchange_test.go @@ -110,6 +110,221 @@ func TestKeyExchange(t *testing.T) { } } +type tdlibZeroHandshakeMessageIDConn struct { + transport.Conn +} + +func (c *tdlibZeroHandshakeMessageIDConn) Send(ctx context.Context, frame *bin.Buffer) error { + candidate := &bin.Buffer{Buf: frame.Copy()} + var message tgproto.UnencryptedMessage + if err := message.Decode(candidate); err != nil { + return c.Conn.Send(ctx, frame) + } + + payload := &bin.Buffer{Buf: message.MessageData} + typeID, err := payload.PeekID() + if err != nil { + return c.Conn.Send(ctx, frame) + } + switch typeID { + case mt.ReqPqMultiRequestTypeID, + mt.ReqDHParamsRequestTypeID, + mt.SetClientDHParamsRequestTypeID: + message.MessageID = 0 + default: + return c.Conn.Send(ctx, frame) + } + + // TDLib NoCryptoImpl includes 0-255 random bytes in message_data_length. + // Use its maximum legal alignment-plus-15-block shape so the complete + // permanent and temporary exchanges exercise padded bodies at every stage. + paddingSize := (-len(message.MessageData)) & 15 + paddingSize += 16 * 15 + message.MessageData = append( + message.MessageData, + bytes.Repeat([]byte{0xa5}, paddingSize)..., + ) + + var rewritten bin.Buffer + if err := message.Encode(&rewritten); err != nil { + return err + } + return c.Conn.Send(ctx, &rewritten) +} + +func TestKeyExchangeAcceptsTDLibZeroMessageIDs(t *testing.T) { + const ( + dc = 2 + expiresIn = 60 + ) + tests := []struct { + name string + clientDC int + temporary bool + wantExpiry bool + }{ + {name: "permanent key", clientDC: dc}, + { + name: "media temporary key", + clientDC: -dc, + temporary: true, + wantExpiry: true, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + addr, pub, srv := startTestServer(t, Options{DC: dc}) + conn := &tdlibZeroHandshakeMessageIDConn{Conn: dialTransportOnly(t, addr)} + t.Cleanup(func() { _ = conn.Close() }) + + exchanger := exchange.NewExchanger(conn, test.clientDC). + WithRand(rand.Reader). + WithLogger(logzap.New(zaptest.NewLogger(t).Named("tdlib-zero-msg-id-client"))) + if test.temporary { + exchanger = exchanger.WithTempMode(expiresIn) + } + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + result, err := exchanger. + Client([]exchange.PublicKey{pub}). + Run(ctx) + if err != nil { + t.Fatalf("TDLib-shaped client exchange: %v", err) + } + if result.AuthKey.ID == ([8]byte{}) { + t.Fatal("TDLib-shaped client exchange returned an empty auth key id") + } + if got := result.ExpiresAt > 0; got != test.wantExpiry { + t.Fatalf("client expiry present = %v, want %v", got, test.wantExpiry) + } + + var ( + saved store.AuthKeyData + found bool + ) + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + saved, found, _ = srv.authKeys.Get(context.Background(), result.AuthKey.ID) + if found { + break + } + time.Sleep(20 * time.Millisecond) + } + if !found { + t.Fatalf("server did not store TDLib-shaped auth key %x", result.AuthKey.ID) + } + if got := saved.ExpiresAt > 0; got != test.wantExpiry { + t.Fatalf("server expiry present = %v, want %v", got, test.wantExpiry) + } + }) + } +} + +func TestValidUnencryptedHandshakeMessageID(t *testing.T) { + encode := func(value bin.Encoder) []byte { + t.Helper() + var payload bin.Buffer + if err := value.Encode(&payload); err != nil { + t.Fatalf("encode payload: %v", err) + } + return payload.Copy() + } + + const ordinaryClientMessageID = int64(1<<32 | 4) + tests := []struct { + name string + messageID int64 + payload []byte + want bool + }{ + { + name: "ordinary client id retains existing admission", + messageID: ordinaryClientMessageID, + payload: encode(&mt.PingRequest{}), + want: true, + }, + { + name: "probe sentinel rejects legacy req pq", + messageID: 1, + payload: encode(&mt.ReqPqRequest{}), + want: false, + }, + { + name: "TDLib probe req pq multi", + messageID: 1, + payload: encode(&mt.ReqPqMultiRequest{}), + want: true, + }, + { + name: "probe sentinel cannot carry req DH", + messageID: 1, + payload: encode(&mt.ReqDHParamsRequest{}), + want: false, + }, + { + name: "zero sentinel rejects legacy req pq", + messageID: 0, + payload: encode(&mt.ReqPqRequest{}), + want: false, + }, + { + name: "TDLib handshake req pq multi", + messageID: 0, + payload: encode(&mt.ReqPqMultiRequest{}), + want: true, + }, + { + name: "TDLib handshake req DH", + messageID: 0, + payload: encode(&mt.ReqDHParamsRequest{}), + want: true, + }, + { + name: "TDLib handshake set client DH", + messageID: 0, + payload: encode(&mt.SetClientDHParamsRequest{}), + want: true, + }, + { + name: "zero sentinel cannot carry ack", + messageID: 0, + payload: encode(&mt.MsgsAck{}), + want: false, + }, + { + name: "other invalid nonzero id remains rejected", + messageID: 2, + payload: encode(&mt.ReqPqMultiRequest{}), + want: false, + }, + { + name: "negative id remains rejected", + messageID: -4, + payload: encode(&mt.ReqPqMultiRequest{}), + want: false, + }, + { + name: "sentinel requires a complete constructor id", + messageID: 0, + payload: []byte{1, 2, 3}, + want: false, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + if got := validUnencryptedHandshakeMessageID(test.messageID, test.payload); got != test.want { + t.Fatalf( + "validUnencryptedHandshakeMessageID(%d) = %v, want %v", + test.messageID, + got, + test.want, + ) + } + }) + } +} + type authKeySaveContextObservation struct { hasDeadline bool deadline time.Time @@ -146,6 +361,52 @@ func (c *ownershipFrameConn) Recv(_ context.Context, b *bin.Buffer) error { return nil } +func TestReadUnencryptedAcceptsTDLibProbeMessageID(t *testing.T) { + nonce := bin.Int128{1, 2, 3, 4} + var payload bin.Buffer + if err := (&mt.ReqPqMultiRequest{Nonce: nonce}).Encode(&payload); err != nil { + t.Fatalf("encode req_pq_multi: %v", err) + } + + tests := []struct { + name string + padding []byte + }{ + {name: "without TDLib no-crypto padding"}, + {name: "with minimum TDLib no-crypto padding", padding: bytes.Repeat([]byte{0xa5}, 12)}, + {name: "with maximum TDLib no-crypto padding", padding: bytes.Repeat([]byte{0x5a}, 252)}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + messageData := append(payload.Copy(), test.padding...) + var frame bin.Buffer + if err := (tgproto.UnencryptedMessage{ + MessageID: 1, + MessageData: messageData, + }).Encode(&frame); err != nil { + t.Fatalf("encode unencrypted probe: %v", err) + } + ex := serverExchangeCompat{ + conn: &ownershipFrameConn{frame: frame.Copy()}, + timeout: time.Second, + } + var decoded compatReqPQ + var scratch bin.Buffer + if err := ex.readUnencrypted(context.Background(), &scratch, &decoded); err != nil { + t.Fatalf("read TDLib probe: %v", err) + } + if decoded.Type != mt.ReqPqMultiRequestTypeID || decoded.Nonce != nonce { + t.Fatalf( + "decoded probe = {type:%#x nonce:%x}, want req_pq_multi nonce %x", + decoded.Type, + decoded.Nonce, + nonce, + ) + } + }) + } +} + func TestExchangeEncryptedReplayTransfersFrameOwnership(t *testing.T) { backing := make([]byte, 64) copy(backing[:8], []byte{1, 2, 3, 4, 5, 6, 7, 8})