347 lines
7.6 KiB
Go
347 lines
7.6 KiB
Go
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):
|
|
}
|
|
}
|