// 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 }