209 lines
6.4 KiB
Go
209 lines
6.4 KiB
Go
package mtprotoedge
|
||
|
||
import (
|
||
"context"
|
||
"crypto/rand"
|
||
"crypto/rsa"
|
||
"net"
|
||
"testing"
|
||
"time"
|
||
|
||
"go.uber.org/zap/zaptest"
|
||
|
||
"github.com/gotd/td/bin"
|
||
"github.com/gotd/td/crypto"
|
||
"github.com/gotd/td/exchange"
|
||
"github.com/gotd/td/proto"
|
||
"github.com/gotd/td/transport"
|
||
)
|
||
|
||
// startTestServer 生成 RSA key、监听随机端口并启动 Server,返回监听地址与公钥。
|
||
// 通过 t.Cleanup 自动取消并校验优雅退出。opts 的 RSAKey/Logger/DC 会被补默认。
|
||
func startTestServer(t *testing.T, opts Options) (addr string, pub exchange.PublicKey, srv *Server) {
|
||
t.Helper()
|
||
rsaKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||
if err != nil {
|
||
t.Fatalf("gen rsa: %v", err)
|
||
}
|
||
opts.RSAKey = rsaKey
|
||
if opts.Logger == nil {
|
||
opts.Logger = zaptest.NewLogger(t)
|
||
}
|
||
if opts.DC == 0 {
|
||
opts.DC = 2
|
||
}
|
||
|
||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||
if err != nil {
|
||
t.Fatalf("listen: %v", err)
|
||
}
|
||
|
||
srv = New(opts)
|
||
ctx, cancel := context.WithCancel(context.Background())
|
||
serveErr := make(chan error, 1)
|
||
go func() { serveErr <- srv.Serve(ctx, ln) }()
|
||
t.Cleanup(func() {
|
||
cancel()
|
||
select {
|
||
case err := <-serveErr:
|
||
if err != nil {
|
||
t.Errorf("serve: %v", err)
|
||
}
|
||
case <-time.After(5 * time.Second):
|
||
t.Error("server did not stop after ctx cancel")
|
||
}
|
||
})
|
||
|
||
return ln.Addr().String(), exchange.PublicKey{RSA: &rsaKey.PublicKey}, srv
|
||
}
|
||
|
||
// dialHandshake 建立 TCP 连接、完成 intermediate 协商与 MTProto 密钥交换,
|
||
// 返回连接、握手结果与 client 端 cipher。连接通过 t.Cleanup 自动关闭。
|
||
func dialHandshake(t *testing.T, addr string, dc int, pub exchange.PublicKey) (transport.Conn, exchange.ClientExchangeResult, crypto.Cipher) {
|
||
t.Helper()
|
||
raw, err := net.Dial("tcp", addr)
|
||
if err != nil {
|
||
t.Fatalf("dial: %v", err)
|
||
}
|
||
conn, err := transport.Intermediate.Handshake(raw)
|
||
if err != nil {
|
||
t.Fatalf("transport handshake: %v", err)
|
||
}
|
||
t.Cleanup(func() { _ = conn.Close() })
|
||
|
||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||
defer cancel()
|
||
auth, err := exchange.NewExchanger(conn, dc).
|
||
WithRand(rand.Reader).
|
||
WithLogger(zaptest.NewLogger(t).Named("client")).
|
||
Client([]exchange.PublicKey{pub}).
|
||
Run(ctx)
|
||
if err != nil {
|
||
t.Fatalf("client exchange: %v", err)
|
||
}
|
||
return conn, auth, crypto.NewClientCipher(rand.Reader)
|
||
}
|
||
|
||
// sendEncrypted 用 client cipher 加密并发送一条带 msgID 的消息。
|
||
func sendEncrypted(t *testing.T, conn transport.Conn, cipher crypto.Cipher, auth exchange.ClientExchangeResult, msgID int64, msg bin.Encoder) {
|
||
t.Helper()
|
||
sendEncryptedWithSalt(t, conn, cipher, auth, auth.ServerSalt, msgID, msg)
|
||
}
|
||
|
||
// sendEncryptedWithSalt 用指定 salt 加密并发送一条消息。
|
||
func sendEncryptedWithSalt(t *testing.T, conn transport.Conn, cipher crypto.Cipher, auth exchange.ClientExchangeResult, salt, msgID int64, msg bin.Encoder) {
|
||
t.Helper()
|
||
body, seqNo := encodeClientMessageForTest(t, msg)
|
||
sendEncryptedWithSaltAndSeq(t, conn, cipher, auth, salt, msgID, seqNo, body)
|
||
}
|
||
|
||
func sendEncryptedWithSeq(t *testing.T, conn transport.Conn, cipher crypto.Cipher, auth exchange.ClientExchangeResult, msgID int64, seqNo int32, msg bin.Encoder) {
|
||
t.Helper()
|
||
body := encodeClientMessageBodyForTest(t, msg)
|
||
sendEncryptedWithSaltAndSeq(t, conn, cipher, auth, auth.ServerSalt, msgID, seqNo, body)
|
||
}
|
||
|
||
func sendEncryptedWithSaltAndSeq(t *testing.T, conn transport.Conn, cipher crypto.Cipher, auth exchange.ClientExchangeResult, salt, msgID int64, seqNo int32, body []byte) {
|
||
t.Helper()
|
||
var buf bin.Buffer
|
||
if err := cipher.Encrypt(auth.AuthKey, crypto.EncryptedMessageData{
|
||
Salt: salt,
|
||
SessionID: auth.SessionID,
|
||
MessageID: msgID,
|
||
SeqNo: seqNo,
|
||
MessageDataLen: int32(len(body)),
|
||
MessageDataWithPadding: body,
|
||
}, &buf); err != nil {
|
||
t.Fatalf("encrypt: %v", err)
|
||
}
|
||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||
defer cancel()
|
||
if err := conn.Send(ctx, &buf); err != nil {
|
||
t.Fatalf("send: %v", err)
|
||
}
|
||
}
|
||
|
||
func encodeClientMessageForTest(t *testing.T, msg bin.Encoder) ([]byte, int32) {
|
||
t.Helper()
|
||
raw := encodeClientMessageBodyForTest(t, msg)
|
||
typeID, err := (&bin.Buffer{Buf: raw}).PeekID()
|
||
if err != nil {
|
||
t.Fatalf("peek encrypted message type: %v", err)
|
||
}
|
||
if container, ok := msg.(*proto.MessageContainer); ok {
|
||
return raw, clientContainerSeqNoForTest(container)
|
||
}
|
||
if clientMessageNeedsAck(typeID) {
|
||
return raw, 1
|
||
}
|
||
return raw, 0
|
||
}
|
||
|
||
func encodeClientMessageBodyForTest(t *testing.T, msg bin.Encoder) []byte {
|
||
t.Helper()
|
||
var body bin.Buffer
|
||
if err := msg.Encode(&body); err != nil {
|
||
t.Fatalf("encode encrypted message: %v", err)
|
||
}
|
||
return body.Copy()
|
||
}
|
||
|
||
func clientContainerSeqNoForTest(container *proto.MessageContainer) int32 {
|
||
var maxSeq int32
|
||
for _, msg := range container.Messages {
|
||
if seq := int32(msg.SeqNo); seq > maxSeq {
|
||
maxSeq = seq
|
||
}
|
||
}
|
||
if maxSeq%2 != 0 {
|
||
maxSeq++
|
||
}
|
||
return maxSeq
|
||
}
|
||
|
||
// collectReplies 读取并解密 server 回发的消息,按 TypeID 收集明文 buffer,
|
||
// 直到见到 wantID(含)或达到上限。用于断言一次请求触发的多条响应
|
||
// (new_session_created / 业务响应 / msgs_ack)。
|
||
func collectReplies(t *testing.T, conn transport.Conn, cipher crypto.Cipher, key crypto.AuthKey, wantID uint32) map[uint32]*bin.Buffer {
|
||
t.Helper()
|
||
got := make(map[uint32]*bin.Buffer)
|
||
for i := 0; i < 8; i++ {
|
||
_, id, plain := readServerMessage(t, conn, cipher, key)
|
||
got[id] = plain
|
||
if id == wantID {
|
||
break
|
||
}
|
||
}
|
||
return got
|
||
}
|
||
|
||
func readServerMessage(t *testing.T, conn transport.Conn, cipher crypto.Cipher, key crypto.AuthKey) (*crypto.EncryptedMessageData, uint32, *bin.Buffer) {
|
||
t.Helper()
|
||
var buf bin.Buffer
|
||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||
err := conn.Recv(ctx, &buf)
|
||
cancel()
|
||
if err != nil {
|
||
t.Fatalf("recv server message: %v", err)
|
||
}
|
||
data, err := cipher.DecryptFromBuffer(key, &buf)
|
||
if err != nil {
|
||
t.Fatalf("decrypt server message: %v", err)
|
||
}
|
||
plain := append([]byte(nil), data.Data()...)
|
||
id, err := (&bin.Buffer{Buf: plain}).PeekID()
|
||
if err != nil {
|
||
t.Fatalf("peek server message: %v", err)
|
||
}
|
||
return data, id, &bin.Buffer{Buf: plain}
|
||
}
|
||
|
||
// mustHave 断言 replies 含指定 TypeID 的消息并返回其 buffer。
|
||
func mustHave(t *testing.T, replies map[uint32]*bin.Buffer, id uint32, name string) *bin.Buffer {
|
||
t.Helper()
|
||
b, ok := replies[id]
|
||
if !ok {
|
||
t.Fatalf("missing %s (%#x)", name, id)
|
||
}
|
||
return b
|
||
}
|