fix: sync admit TDLib handshake message id sentinels

This commit is contained in:
iamxvbaba 2026-07-31 15:35:10 +08:00
parent 0d917ef87e
commit 9cb76b7c5a
2 changed files with 296 additions and 1 deletions

View file

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

View file

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