owpengram-server/internal/mtprotoedge/helpers_test.go
2026-06-04 01:37:39 +08:00

209 lines
6.4 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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
}