Files
NixMsg/test/load/client.go
T

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 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):
}
}