fix: 压测客户端改为 MQTT 5 并完成 hello 与收发统计
This commit is contained in:
@@ -0,0 +1,346 @@
|
||||
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 err := c.mc.Send(pkt); err != nil {
|
||||
return fmt.Errorf("%s CONNECT: %w", c.ID, err)
|
||||
}
|
||||
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 err := c.mc.Send(sub); err != nil {
|
||||
return fmt.Errorf("%s SUBSCRIBE: %w", c.ID, err)
|
||||
}
|
||||
if _, err := c.mc.Recv(); err != nil {
|
||||
return fmt.Errorf("%s SUBACK: %w", c.ID, err)
|
||||
}
|
||||
|
||||
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):
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user