146 lines
3.2 KiB
Go
146 lines
3.2 KiB
Go
// Package load 提供 MQTT 压测客户端骨架:先连上并统计连接数。
|
|
package load
|
|
|
|
import (
|
|
"encoding/binary"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"sync"
|
|
"sync/atomic"
|
|
"time"
|
|
)
|
|
|
|
// MQTTBench 保持多个 MQTT 3.1.1 TCP 连接(仅 CONNECT/CONNACK)。
|
|
type MQTTBench struct {
|
|
Addr string
|
|
ClientIDPrefix string
|
|
Timeout time.Duration
|
|
|
|
mu sync.Mutex
|
|
conns []net.Conn
|
|
alive atomic.Int64
|
|
}
|
|
|
|
// ConnectN 建立 n 条连接;成功一条 alive+1,并打印当前连接数到 w(可为 nil)。
|
|
func (b *MQTTBench) ConnectN(n int, w io.Writer) error {
|
|
if b.Addr == "" {
|
|
return fmt.Errorf("addr required")
|
|
}
|
|
if n <= 0 {
|
|
return fmt.Errorf("n must be > 0")
|
|
}
|
|
timeout := b.Timeout
|
|
if timeout <= 0 {
|
|
timeout = 5 * time.Second
|
|
}
|
|
prefix := b.ClientIDPrefix
|
|
if prefix == "" {
|
|
prefix = "q-bench-"
|
|
}
|
|
|
|
for i := 0; i < n; i++ {
|
|
conn, err := net.DialTimeout("tcp", b.Addr, timeout)
|
|
if err != nil {
|
|
return fmt.Errorf("dial %d: %w", i, err)
|
|
}
|
|
_ = conn.SetDeadline(time.Now().Add(timeout))
|
|
cid := fmt.Sprintf("%s%d", prefix, i)
|
|
if err := mqttConnect(conn, cid); err != nil {
|
|
_ = conn.Close()
|
|
return fmt.Errorf("connect %d: %w", i, err)
|
|
}
|
|
_ = conn.SetDeadline(time.Time{})
|
|
b.mu.Lock()
|
|
b.conns = append(b.conns, conn)
|
|
b.mu.Unlock()
|
|
cur := b.alive.Add(1)
|
|
if w != nil {
|
|
_, _ = fmt.Fprintf(w, "mqtt connections: %d\n", cur)
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// Alive 当前仍打开的连接数。
|
|
func (b *MQTTBench) Alive() int64 {
|
|
return b.alive.Load()
|
|
}
|
|
|
|
// Close 关闭全部连接。
|
|
func (b *MQTTBench) Close() {
|
|
b.mu.Lock()
|
|
defer b.mu.Unlock()
|
|
for _, c := range b.conns {
|
|
_ = c.Close()
|
|
}
|
|
b.conns = nil
|
|
b.alive.Store(0)
|
|
}
|
|
|
|
func mqttConnect(conn net.Conn, clientID string) error {
|
|
pkt := buildConnect(clientID)
|
|
if _, err := conn.Write(pkt); err != nil {
|
|
return err
|
|
}
|
|
header := make([]byte, 4)
|
|
if _, err := io.ReadFull(conn, header[:2]); err != nil {
|
|
return err
|
|
}
|
|
if header[0] != 0x20 {
|
|
return fmt.Errorf("unexpected packet type 0x%02x", header[0])
|
|
}
|
|
// remaining length 对 CONNACK 固定为 2
|
|
if header[1] != 2 {
|
|
return fmt.Errorf("unexpected remaining length %d", header[1])
|
|
}
|
|
if _, err := io.ReadFull(conn, header[2:4]); err != nil {
|
|
return err
|
|
}
|
|
if header[3] != 0 {
|
|
return fmt.Errorf("CONNACK rc=%d", header[3])
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func buildConnect(clientID string) []byte {
|
|
// Variable header: protocol name MQTT, level 4, flags 0, keepalive 60
|
|
vh := []byte{
|
|
0x00, 0x04, 'M', 'Q', 'T', 'T',
|
|
0x04,
|
|
0x00, // clean session=0 flags for skeleton; brokers may still accept
|
|
0x00, 0x3c,
|
|
}
|
|
// Actually clean session bit should be set for simple benches
|
|
vh[7] = 0x02 // Clean Session
|
|
|
|
id := []byte(clientID)
|
|
payload := make([]byte, 2+len(id))
|
|
binary.BigEndian.PutUint16(payload[0:2], uint16(len(id)))
|
|
copy(payload[2:], id)
|
|
|
|
remaining := len(vh) + len(payload)
|
|
pkt := make([]byte, 0, 2+remaining)
|
|
pkt = append(pkt, 0x10)
|
|
pkt = append(pkt, encodeRemainingLength(remaining)...)
|
|
pkt = append(pkt, vh...)
|
|
pkt = append(pkt, payload...)
|
|
return pkt
|
|
}
|
|
|
|
func encodeRemainingLength(n int) []byte {
|
|
var out []byte
|
|
for {
|
|
encoded := byte(n % 128)
|
|
n /= 128
|
|
if n > 0 {
|
|
encoded |= 0x80
|
|
}
|
|
out = append(out, encoded)
|
|
if n == 0 {
|
|
break
|
|
}
|
|
}
|
|
return out
|
|
}
|