package load import ( "encoding/json" "fmt" "sync" "sync/atomic" "time" "git.asio.asia/nixevol/NixMsg/internal/protocol" "git.asio.asia/nixevol/NixMsg/test/harness" "github.com/mochi-mqtt/server/v2/packets" ) // Client 是压测用 MQTT 5 端会话:CONNECT → 订阅 down → hello,之后收发应用帧。 type Client struct { ID string mc harness.MQTTClient pktID atomic.Uint32 rid atomic.Uint64 mu sync.Mutex inbox []map[string]any closed bool disc atomic.Bool done chan struct{} } // Resp 是 type=resp 的解析结果。 type Resp struct { OK bool Error map[string]any Data any Raw map[string]any } func (r Resp) ErrorCode() string { if r.OK || r.Error == nil { return "" } code, _ := r.Error["code"].(string) if code == "" { return "error" } return code } // DialAndHello 连接(WebSocket /mqtt 或裸 TCP)、MQTT 5 认证、订阅、hello。 func DialAndHello(httpBase, tcpAddr, endpointID, password string, timeout time.Duration) (*Client, error) { if timeout <= 0 { timeout = 15 * time.Second } var mc harness.MQTTClient var err error if tcpAddr != "" { mc, err = harness.DialMQTTTCP(tcpAddr, timeout) } else { if httpBase == "" { return nil, fmt.Errorf("需要 -http 或 -tcp") } mc, err = harness.DialMQTTWebSocket(httpBase, timeout) } if err != nil { return nil, fmt.Errorf("dial %s: %w", endpointID, err) } c := &Client{ID: endpointID, mc: mc, done: make(chan struct{})} c.pktID.Store(10) if err := c.connectSubscribeHello(password, timeout); err != nil { _ = mc.Close() return nil, err } go c.readLoop() go c.pingLoop() return c, nil } func (c *Client) connectSubscribeHello(password string, timeout time.Duration) error { pkt, err := encodeConnect(c.ID, password) if err != nil { return err } if sendErr := c.mc.Send(pkt); sendErr != nil { return fmt.Errorf("%s CONNECT: %w", c.ID, sendErr) } ack, err := c.mc.Recv() if err != nil { return fmt.Errorf("%s CONNACK: %w", c.ID, err) } reason, err := connackReason(ack) if err != nil { return fmt.Errorf("%s CONNACK: %w", c.ID, err) } if reason != 0 { return fmt.Errorf("%s CONNACK reason=0x%02x", c.ID, reason) } sub, err := encodeSubscribe(c.nextPkt(), c.ID) if err != nil { return err } if sendErr := c.mc.Send(sub); sendErr != nil { return fmt.Errorf("%s SUBSCRIBE: %w", c.ID, sendErr) } if _, recvErr := c.mc.Recv(); recvErr != nil { return fmt.Errorf("%s SUBACK: %w", c.ID, recvErr) } hello := protocol.Hello{V: protocol.Version, Type: protocol.TypeHello, RID: "h0", Client: "mqttbench/0.1"} payload, err := protocol.Marshal(hello) if err != nil { return err } if err := c.publishRaw(payload); err != nil { return fmt.Errorf("%s hello publish: %w", c.ID, err) } deadline := time.Now().Add(timeout) for time.Now().Before(deadline) { raw, err := c.mc.Recv() if err != nil { return fmt.Errorf("%s hello recv: %w", c.ID, err) } m := c.handlePacket(raw) if m == nil { continue } if m["type"] == protocol.TypeResp && m["rid"] == "h0" { if m["ok"] == true { return nil } return fmt.Errorf("%s hello 失败: %v", c.ID, m) } c.push(m) } return fmt.Errorf("%s hello 超时", c.ID) } func (c *Client) nextPkt() uint16 { for { v := c.pktID.Add(1) id := uint16(v) if id != 0 { return id } } } func (c *Client) NextRID() string { return fmt.Sprintf("r%d", c.rid.Add(1)) } func (c *Client) publishRaw(payload []byte) error { pkt, err := encodePublish(c.nextPkt(), c.ID, payload) if err != nil { return err } return c.mc.Send(pkt) } func (c *Client) pingLoop() { t := time.NewTicker(10 * time.Second) defer t.Stop() for { select { case <-c.done: return case <-t.C: if err := c.mc.Send(encodePingreq()); err != nil { c.disc.Store(true) return } } } } func (c *Client) readLoop() { defer close(c.done) for { raw, err := c.mc.Recv() if err != nil { c.disc.Store(true) return } m := c.handlePacket(raw) if m != nil { c.push(m) } } } func (c *Client) handlePacket(raw []byte) map[string]any { if len(raw) < 2 { return nil } typ := raw[0] >> 4 switch typ { case packets.Puback, packets.Pingresp, packets.Suback: return nil case packets.Disconnect: c.disc.Store(true) return nil case packets.Publish: payload, packetID, qos, err := decodePublish(raw) if err != nil { return nil } if qos == 1 && packetID != 0 { if ack, encErr := encodePuback(packetID); encErr == nil { _ = c.mc.Send(ack) } } var m map[string]any if json.Unmarshal(payload, &m) != nil { return nil } return m default: return nil } } func (c *Client) push(m map[string]any) { c.mu.Lock() c.inbox = append(c.inbox, m) c.mu.Unlock() } // Request 发上行帧并等待对应 rid 的 resp。 func (c *Client) Request(frame map[string]any, timeout time.Duration) (Resp, error) { if timeout <= 0 { timeout = 15 * time.Second } rid, _ := frame["rid"].(string) if rid == "" { rid = c.NextRID() frame["rid"] = rid } payload, err := protocol.Marshal(frame) if err != nil { return Resp{}, err } if err := c.publishRaw(payload); err != nil { c.disc.Store(true) return Resp{}, err } deadline := time.Now().Add(timeout) for time.Now().Before(deadline) { if c.Disconnected() { return Resp{}, fmt.Errorf("%s 已断开,等待 resp rid=%s", c.ID, rid) } m := c.takeMatching(func(x map[string]any) bool { return x["type"] == protocol.TypeResp && x["rid"] == rid }) if m != nil { r := Resp{OK: m["ok"] == true, Raw: m, Data: m["data"]} if e, ok := m["error"].(map[string]any); ok { r.Error = e } return r, nil } time.Sleep(2 * time.Millisecond) } return Resp{}, fmt.Errorf("%s 等待 resp rid=%s 超时", c.ID, rid) } // WaitType 等到指定 type 的下行帧。 func (c *Client) WaitType(typ string, timeout time.Duration) (map[string]any, error) { deadline := time.Now().Add(timeout) for time.Now().Before(deadline) { m := c.takeMatching(func(x map[string]any) bool { return x["type"] == typ }) if m != nil { return m, nil } if c.Disconnected() { return nil, fmt.Errorf("%s 已断开,等待 type=%s", c.ID, typ) } time.Sleep(2 * time.Millisecond) } return nil, fmt.Errorf("%s 等待 type=%s 超时", c.ID, typ) } // WaitMsg 等到指定消息号的 msg 帧。 func (c *Client) WaitMsg(id string, timeout time.Duration) (map[string]any, error) { deadline := time.Now().Add(timeout) for time.Now().Before(deadline) { m := c.takeMatching(func(x map[string]any) bool { return x["type"] == protocol.TypeMsg && x["id"] == id }) if m != nil { return m, nil } if c.Disconnected() { return nil, fmt.Errorf("%s 已断开,等待 msg id=%s", c.ID, id) } time.Sleep(2 * time.Millisecond) } return nil, fmt.Errorf("%s 等待 msg id=%s 超时", c.ID, id) } // DrainEvents 丢掉 group_event / presence,避免干扰收发统计。 func (c *Client) DrainEvents() { c.mu.Lock() defer c.mu.Unlock() kept := c.inbox[:0] for _, x := range c.inbox { typ, _ := x["type"].(string) if typ == protocol.TypeGroupEvent || typ == protocol.TypePresence { continue } kept = append(kept, x) } c.inbox = kept } func (c *Client) takeMatching(pred func(map[string]any) bool) map[string]any { c.mu.Lock() defer c.mu.Unlock() for i, m := range c.inbox { if pred(m) { c.inbox = append(c.inbox[:i], c.inbox[i+1:]...) return m } } return nil } // Disconnected 连接是否已掉。 func (c *Client) Disconnected() bool { return c.disc.Load() } // Close 关闭底层连接。 func (c *Client) Close() { c.mu.Lock() if c.closed { c.mu.Unlock() return } c.closed = true c.mu.Unlock() _ = c.mc.Close() select { case <-c.done: case <-time.After(3 * time.Second): } }