diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 233ea8e..da97e46 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -1698,3 +1698,13 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是 - 原因:`resp` 回不来是测试客户端丢掉合帧里的后续 MQTT 包,以及并发写把发往服务端的流写乱;不是服务端死锁。 - 备选方案:合入 fix3 的 broker `OnPublish`/延后队列(否决,基于错误假设,还会改 PUBACK 语义)。 - 影响:跨线改了总控目录 `test/harness` 与身份线 `emit`;不改 `internal/broker/`。所有经 `DialMQTTWebSocket` 的测试客户端一并受益。 + +### 复审修复 T-03 + +1. **压测工具改为 MQTT 5 认证握手并覆盖收发场景** + - 日期:2026-09-30 + - 原条款:PRD 第 8 节规模/吞吐/延迟;DEVELOPMENT 第 5 节 CONNECT(ClientID=Username=端编号、CleanStart、心跳 30)与 hello;issue #64。 + - 实际做法:重写 `test/load`。`mqttbench` 用 MQTT 5 连接(WebSocket `/mqtt` 或裸 TCP),完成订阅与 hello;端可通过管理接口 CSV 批量开通、自助注册,或 `-reuse` 复用已开通编号。场景:保持 N 条在线;按目标速率单聊 send/ack;建群发一条并统计从提交成功到成员收到 `msg` 的耗时。输出提交/送达 P50/P95/P99、错误码分布、掉线数。用法见 `test/load/README.md`。不进 CI;`go test ./test/load/` 只跑编码/分位单测、假 broker CONNACK 和短时冒烟。长时间 1000 连接 / 10 分钟本轮不跑。 + - 原因:原工具只发 MQTT 3.1.1 CONNECT、无用户名密码、无 hello/收发,连不上 NixMsg,也测不出 PRD 第 8 节任何指标。 + - 备选方案:直接用四套 SDK 压测(否决:工具需覆盖裸设备路径,且不改 sdk)。 + - 影响:仅 `test/load/*` 与本文本节;不改 broker/message/sdk。本地验收为 50 连接、每秒 20 条、1 分钟出报告。 diff --git a/test/load/README.md b/test/load/README.md new file mode 100644 index 0000000..caff513 --- /dev/null +++ b/test/load/README.md @@ -0,0 +1,59 @@ +# NixMsg 压测工具(`test/load`) + +本目录提供可执行 PRD 第 8 节指标的压测客户端。**不进 CI**;`go test ./test/load/` 只跑工具自身的最小单测和短时冒烟。 + +长时间 1000 连接 / 每秒 200 条 / 10 分钟本轮按负责人要求不跑,但工具已能执行。本地验收口径:50 个连接、每秒 20 条、跑 1 分钟能打出报告。 + +## 连接规则(DEVELOPMENT 第 5 节) + +- MQTT 5;`CleanStart = true`;心跳 30 秒。 +- ClientID、Username 都等于端编号;Password 为登录密码。 +- 只订阅 `nix/c/{编号}/down`,只向 `nix/c/{编号}/up` 发布。 +- 订阅完成后发 `hello`,收到成功 `resp` 才算在线,才开始收发。 + +默认走 WebSocket `ws:///mqtt`(子协议 `mqtt`)。也可用 `-tcp` 走裸 MQTT。 + +## 准备端 + +任选其一: + +1. **管理接口批量开通**:`-admin-pass` 或 `-admin-token`。工具用 `POST /api/admin/endpoints/import` 一次最多 1000 行。 +2. **自助注册**:管理员已开启注册时传 `-register-code`。 +3. **复用已开通的端**:`-reuse`。编号冲突(`409 id_taken`)视为已存在并继续登录。 + +端编号为 `{prefix}{序号}`,默认 `qb0000` 起。登录密码默认 `password1234`。 + +## 场景 + +| `-scenario` | 行为 | +|---|---| +| `hold` | 只保持 N 条已握手连接(配合 `-hold`) | +| `dm` | 按 `-rate` 做单聊 send → 等对方 msg → ack | +| `group` | 建一个 `-group-size` 人的群,发一条,统计从提交成功到成员都收到 `msg` 的耗时(推送排队的上界) | +| `all`(默认) | 先 hold(若设置),再单聊,再群发 | + +输出:提交与送达的 P50/P95/P99(毫秒)、错误码分布、连接掉线数、群推送排队耗时。 + +## 用法 + +在仓库根目录: + +```text +go run ./test/load/cmd/mqttbench -http http://127.0.0.1:<端口> -admin-pass <管理员密码> -n 50 -rate 20 -duration 1m +``` + +常用参数: + +| 参数 | 默认 | 说明 | +|---|---|---| +| `-http` | (必填,除非只测裸 TCP 且 `-reuse`) | 服务 HTTP 根 | +| `-tcp` | 空 | 裸 MQTT `host:port` | +| `-n` | 50 | 连接数 | +| `-rate` | 20 | 单聊每秒条数 | +| `-duration` | 1m | 单聊时长;`0` 跳过单聊 | +| `-group-size` | 与 `-n` 相同 | `0` 跳过群发 | +| `-reuse` | false | 复用已开通编号 | +| `-register-code` | 空 | 自助注册 | +| `-prefix` | qb | 端编号前缀 | + +不要把本命令加进 `task check`。完整规模(1000 在线、每秒 200、10 分钟、1000 人群发 5 秒)在专用机器上把 `-n`/`-rate`/`-duration`/`-group-size` 调到 PRD 第 8 节即可。 diff --git a/test/load/bench.go b/test/load/bench.go new file mode 100644 index 0000000..9e2ed79 --- /dev/null +++ b/test/load/bench.go @@ -0,0 +1,338 @@ +package load + +import ( + "fmt" + "strconv" + "strings" + "sync" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/protocol" +) + +// Config 控制压测工具。 +type Config struct { + HTTPBase string + AdminHTTP string + TCPAddr string + AdminUser string + AdminPass string + APIToken string + RegisterCode string + Reuse bool + Password string + Prefix string + N int + Rate float64 + Duration time.Duration + Hold time.Duration + GroupSize int + Scenario string // hold / dm / group / all + DialTimeout time.Duration +} + +func (c *Config) defaults() error { + if c.N <= 0 { + return fmt.Errorf("n 必须 > 0") + } + if c.Password == "" { + c.Password = "password1234" + } + if c.Prefix == "" { + c.Prefix = "qb" + } + if c.DialTimeout <= 0 { + c.DialTimeout = 15 * time.Second + } + if c.Scenario == "" { + c.Scenario = "all" + } + if c.GroupSize < 0 { + c.GroupSize = 0 + } + if c.GroupSize > c.N { + c.GroupSize = c.N + } + if c.HTTPBase == "" && c.TCPAddr == "" { + return fmt.Errorf("需要 -http 或 -tcp") + } + if !protocol.ValidLoginPassword(c.Password) || c.Password == "" { + return fmt.Errorf("登录密码不合法") + } + return nil +} + +// EndpointIDs 生成 n 个合法端编号。 +func EndpointIDs(prefix string, n int) ([]string, error) { + if prefix == "" { + prefix = "qb" + } + width := len(strconv.Itoa(n - 1)) + if width < 4 { + width = 4 + } + ids := make([]string, n) + for i := 0; i < n; i++ { + id := fmt.Sprintf("%s%0*d", prefix, width, i) + if !protocol.ValidEndpointID(id) { + return nil, fmt.Errorf("生成的端编号不合法: %s", id) + } + ids[i] = id + } + return ids, nil +} + +// Run 开通(或复用)端、保持 N 条在线连接,按场景跑单聊速率和群发排队。 +func Run(cfg Config) (*Report, error) { + if err := cfg.defaults(); err != nil { + return nil, err + } + ids, err := EndpointIDs(cfg.Prefix, cfg.N) + if err != nil { + return nil, err + } + if err := Provision(cfg, ids); err != nil { + return nil, err + } + clients, err := ConnectAll(cfg, ids) + if err != nil { + return nil, err + } + defer closeAll(clients) + + rep := &Report{Connections: len(clients), Errors: map[string]int{}} + sc := strings.ToLower(cfg.Scenario) + runHold := sc == "hold" || sc == "all" + runDM := sc == "dm" || sc == "all" + runGroup := sc == "group" || sc == "all" + + if runHold && cfg.Hold > 0 { + time.Sleep(cfg.Hold) + } + if runDM && cfg.Duration > 0 { + if cfg.Rate <= 0 { + return nil, fmt.Errorf("单聊场景需要 rate > 0") + } + if err := runDirectMessages(clients, cfg, rep); err != nil { + return nil, err + } + } + if runGroup && cfg.GroupSize >= 2 { + if err := runGroupFanout(clients, cfg.GroupSize, rep); err != nil { + return nil, err + } + } + + alive := 0 + disc := 0 + for _, c := range clients { + if c.Disconnected() { + disc++ + } else { + alive++ + } + } + rep.Alive = alive + rep.Disconnects = disc + return rep, nil +} + +// ConnectAll 为每个编号建立 MQTT 5 会话并完成 hello。 +func ConnectAll(cfg Config, ids []string) ([]*Client, error) { + clients := make([]*Client, len(ids)) + type result struct { + i int + c *Client + err error + } + ch := make(chan result, len(ids)) + sem := make(chan struct{}, 8) + var wg sync.WaitGroup + for i, id := range ids { + wg.Add(1) + go func(i int, id string) { + defer wg.Done() + sem <- struct{}{} + defer func() { <-sem }() + c, err := DialAndHello(cfg.HTTPBase, cfg.TCPAddr, id, cfg.Password, cfg.DialTimeout) + ch <- result{i: i, c: c, err: err} + }(i, id) + } + go func() { + wg.Wait() + close(ch) + }() + var first error + for r := range ch { + if r.err != nil && first == nil { + first = r.err + continue + } + if r.c != nil { + clients[r.i] = r.c + } + } + if first != nil { + closeAll(clients) + return nil, first + } + return clients, nil +} + +func closeAll(clients []*Client) { + for _, c := range clients { + if c != nil { + c.Close() + } + } +} + +func runDirectMessages(clients []*Client, cfg Config, rep *Report) error { + n := len(clients) + if n < 2 { + return fmt.Errorf("单聊至少需要 2 个连接") + } + var submit []float64 + var deliver []float64 + start := time.Now() + end := start.Add(cfg.Duration) + seq := 0 + for time.Now().Before(end) { + tickStart := start.Add(time.Duration(float64(seq) / cfg.Rate * float64(time.Second))) + if wait := time.Until(tickStart); wait > 0 { + time.Sleep(wait) + } + a := clients[seq%n] + b := clients[(seq+1)%n] + if a.Disconnected() || b.Disconnected() { + addError(rep.Errors, "disconnected") + seq++ + continue + } + msgID := fmt.Sprintf("dm-%d", seq) + t0 := time.Now() + resp, err := a.Request(map[string]any{ + "v": 1, "type": protocol.TypeSend, "rid": a.NextRID(), "id": msgID, + "to": map[string]any{"kind": protocol.TargetEndpoint, "id": b.ID}, + "body": map[string]any{"enc": protocol.EncUTF8, "data": "load"}, + "delay_ms": int64(0), + }, 15*time.Second) + if err != nil { + addError(rep.Errors, "timeout") + seq++ + continue + } + submit = append(submit, float64(time.Since(t0).Milliseconds())) + if !resp.OK { + addError(rep.Errors, resp.ErrorCode()) + seq++ + continue + } + tSubmit := time.Now() + msg, err := b.WaitMsg(msgID, 15*time.Second) + if err != nil { + addError(rep.Errors, "deliver_timeout") + seq++ + continue + } + deliver = append(deliver, float64(time.Since(tSubmit).Milliseconds())) + from, _ := msg["from"].(string) + if from == "" { + from = a.ID + } + ack, err := b.Request(map[string]any{ + "v": 1, "type": protocol.TypeAck, "rid": b.NextRID(), + "from": from, "id": msgID, + }, 15*time.Second) + if err != nil { + addError(rep.Errors, "ack_timeout") + seq++ + continue + } + if !ack.OK { + addError(rep.Errors, ack.ErrorCode()) + } + seq++ + } + elapsed := time.Since(start).Seconds() + rep.Sends = seq + if elapsed > 0 { + rep.ActualRate = float64(seq) / elapsed + } + rep.SubmitMs = calcPercentiles(submit) + rep.DeliverMs = calcPercentiles(deliver) + return nil +} + +func runGroupFanout(clients []*Client, groupSize int, rep *Report) error { + if groupSize > len(clients) { + groupSize = len(clients) + } + members := clients[:groupSize] + owner := members[0] + gID := fmt.Sprintf("lg%d", time.Now().UnixMilli()) + if !protocol.ValidEndpointID(gID) { + gID = "lgload1" + } + memberIn := make([]map[string]any, 0, groupSize-1) + for _, c := range members[1:] { + memberIn = append(memberIn, map[string]any{"id": c.ID}) + } + created, err := owner.Request(map[string]any{ + "v": 1, "type": protocol.TypeGroupCreate, "rid": owner.NextRID(), + "id": gID, "name": "load", "members": memberIn, + }, 30*time.Second) + if err != nil { + return fmt.Errorf("建群: %w", err) + } + if !created.OK { + addError(rep.Errors, created.ErrorCode()) + return fmt.Errorf("建群失败: %v", created.Error) + } + time.Sleep(200 * time.Millisecond) + for _, c := range members { + c.DrainEvents() + } + + msgID := "gload-1" + expect := groupSize - 1 // 发送者自己不收群消息 + sent, err := owner.Request(map[string]any{ + "v": 1, "type": protocol.TypeSend, "rid": owner.NextRID(), "id": msgID, + "to": map[string]any{"kind": protocol.TargetGroup, "id": gID}, + "body": map[string]any{"enc": protocol.EncUTF8, "data": "g"}, + "delay_ms": int64(0), + }, 30*time.Second) + if err != nil { + return fmt.Errorf("群发: %w", err) + } + if !sent.OK { + addError(rep.Errors, sent.ErrorCode()) + return fmt.Errorf("群发失败: %v", sent.Error) + } + tQueued := time.Now() + var wg sync.WaitGroup + errCh := make(chan error, expect) + for _, c := range members[1:] { + wg.Add(1) + go func(c *Client) { + defer wg.Done() + if _, werr := c.WaitMsg(msgID, 15*time.Second); werr != nil { + errCh <- werr + } + }(c) + } + wg.Wait() + close(errCh) + fail := 0 + for e := range errCh { + if e != nil { + fail++ + addError(rep.Errors, "group_deliver_timeout") + } + } + if fail == 0 { + rep.GroupQueueMs = float64(time.Since(tQueued).Milliseconds()) + } + rep.GroupMembers = groupSize + return nil +} diff --git a/test/load/client.go b/test/load/client.go new file mode 100644 index 0000000..05ee3f2 --- /dev/null +++ b/test/load/client.go @@ -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): + } +} diff --git a/test/load/cmd/mqttbench/main.go b/test/load/cmd/mqttbench/main.go index 66b2516..f48c1c2 100644 --- a/test/load/cmd/mqttbench/main.go +++ b/test/load/cmd/mqttbench/main.go @@ -11,30 +11,72 @@ import ( "git.asio.asia/nixevol/NixMsg/test/load" ) -// 压测客户端骨架:连上 N 个 MQTT 连接并打印连接数。完整 1000 连接 / 10 分钟压测留给 Q3。 func main() { - addr := flag.String("addr", "", "MQTT broker host:port(必填,不要写死业务端口)") - n := flag.Int("n", 10, "连接数") - prefix := flag.String("prefix", "q-bench-", "client id 前缀") - hold := flag.Duration("hold", 3*time.Second, "保持连接时长") + httpBase := flag.String("http", "", "服务 HTTP 根(WebSocket /mqtt 与管理/注册接口),如 http://127.0.0.1:12345") + adminHTTP := flag.String("admin-http", "", "管理接口根;默认与 -http 相同") + tcp := flag.String("tcp", "", "裸 MQTT TCP host:port;非空则 MQTT 走 TCP,HTTP 仍用于开通端") + n := flag.Int("n", 50, "在线连接数") + rate := flag.Float64("rate", 20, "单聊提交每秒条数") + duration := flag.Duration("duration", time.Minute, "单聊持续时长;0 表示跳过单聊") + hold := flag.Duration("hold", 0, "连上后先保持这么久再跑收发") + groupSize := flag.Int("group-size", -1, "群发成员数;默认与 -n 相同,0 表示跳过群发") + prefix := flag.String("prefix", "qb", "端编号前缀(小写)") + password := flag.String("password", "password1234", "端登录密码") + adminUser := flag.String("admin-user", "admin", "管理员用户名") + adminPass := flag.String("admin-pass", "", "管理员密码(批量开通)") + adminToken := flag.String("admin-token", "", "管理 API 令牌,与密码二选一") + regCode := flag.String("register-code", "", "自助注册安全码(不走管理开通时使用)") + reuse := flag.Bool("reuse", false, "编号已存在则跳过开通/注册") + scenario := flag.String("scenario", "all", "hold | dm | group | all") flag.Parse() - if *addr == "" { - fmt.Fprintln(os.Stderr, "usage: mqttbench -addr host:port [-n 10]") + + if *httpBase == "" && *tcp == "" { + fmt.Fprintln(os.Stderr, "usage: mqttbench -http http://host:port [options]") + flag.PrintDefaults() os.Exit(2) } + gs := *groupSize + if gs < 0 { + gs = *n + } - b := &load.MQTTBench{Addr: *addr, ClientIDPrefix: *prefix, Timeout: 5 * time.Second} - defer b.Close() - if err := b.ConnectN(*n, os.Stdout); err != nil { - fmt.Fprintln(os.Stderr, err) + cfg := load.Config{ + HTTPBase: *httpBase, + AdminHTTP: *adminHTTP, + TCPAddr: *tcp, + AdminUser: *adminUser, + AdminPass: *adminPass, + APIToken: *adminToken, + RegisterCode: *regCode, + Reuse: *reuse, + Password: *password, + Prefix: *prefix, + N: *n, + Rate: *rate, + Duration: *duration, + Hold: *hold, + GroupSize: gs, + Scenario: *scenario, + } + + stop := make(chan os.Signal, 1) + signal.Notify(stop, os.Interrupt, syscall.SIGTERM) + done := make(chan struct{}) + var runErr error + var report *load.Report + go func() { + defer close(done) + report, runErr = load.Run(cfg) + }() + select { + case <-done: + case <-stop: + fmt.Fprintln(os.Stderr, "interrupted") + os.Exit(130) + } + if runErr != nil { + fmt.Fprintln(os.Stderr, runErr) os.Exit(1) } - fmt.Printf("held %d connections for %s\n", b.Alive(), hold.String()) - - ch := make(chan os.Signal, 1) - signal.Notify(ch, os.Interrupt, syscall.SIGTERM) - select { - case <-time.After(*hold): - case <-ch: - } + fmt.Print(report.Format()) } diff --git a/test/load/connect.go b/test/load/connect.go new file mode 100644 index 0000000..564f6a9 --- /dev/null +++ b/test/load/connect.go @@ -0,0 +1,153 @@ +package load + +import ( + "bytes" + "fmt" + "io" + + "github.com/mochi-mqtt/server/v2/packets" +) + +const ( + mqttProtocolLevel5 = 5 + mqttKeepaliveSec = 30 +) + +// encodeConnect 按 DEVELOPMENT §5 生成 MQTT 5 CONNECT: +// ClientID=Username=端编号、CleanStart、心跳 30、会话过期间隔缺省为 0(Clean Start 时即为 0)。 +// 不声明 Receive Maximum,避免踩 mochi 发送配额路径。 +func encodeConnect(endpointID, password string) ([]byte, error) { + pk := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Connect}, + ProtocolVersion: mqttProtocolLevel5, + Connect: packets.ConnectParams{ + ProtocolName: []byte("MQTT"), + Clean: true, + ClientIdentifier: endpointID, + Keepalive: mqttKeepaliveSec, + UsernameFlag: true, + Username: []byte(endpointID), + PasswordFlag: true, + Password: []byte(password), + }, + } + var buf bytes.Buffer + if err := pk.ConnectEncode(&buf); err != nil { + return nil, err + } + return buf.Bytes(), nil +} + +func encodeSubscribe(packetID uint16, endpointID string) ([]byte, error) { + pk := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Subscribe, Qos: 1}, + ProtocolVersion: mqttProtocolLevel5, + PacketID: packetID, + Filters: packets.Subscriptions{ + {Filter: downTopic(endpointID), Qos: 1}, + }, + } + var buf bytes.Buffer + if err := pk.SubscribeEncode(&buf); err != nil { + return nil, err + } + return buf.Bytes(), nil +} + +func encodePublish(packetID uint16, endpointID string, payload []byte) ([]byte, error) { + pk := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Publish, Qos: 1}, + ProtocolVersion: mqttProtocolLevel5, + TopicName: upTopic(endpointID), + PacketID: packetID, + Payload: payload, + } + var buf bytes.Buffer + if err := pk.PublishEncode(&buf); err != nil { + return nil, err + } + return buf.Bytes(), nil +} + +func encodePuback(packetID uint16) ([]byte, error) { + pk := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Puback}, + ProtocolVersion: mqttProtocolLevel5, + PacketID: packetID, + } + var buf bytes.Buffer + if err := pk.PubackEncode(&buf); err != nil { + return nil, err + } + return buf.Bytes(), nil +} + +func encodePingreq() []byte { + return []byte{0xC0, 0x00} +} + +func downTopic(id string) string { return "nix/c/" + id + "/down" } +func upTopic(id string) string { return "nix/c/" + id + "/up" } + +func connackReason(raw []byte) (byte, error) { + if len(raw) < 2 || raw[0]>>4 != packets.Connack { + return 0, fmt.Errorf("不是 CONNACK:type=0x%02x len=%d", byteAt(raw, 0), len(raw)) + } + rem, n, err := decodeRemainingLength(raw[1:]) + if err != nil { + return 0, err + } + body := raw[1+n:] + if rem < 2 || len(body) < 2 { + return 0, fmt.Errorf("CONNACK 过短 remaining=%d body=%d", rem, len(body)) + } + return body[1], nil +} + +func decodePublish(raw []byte) (payload []byte, packetID uint16, qos byte, err error) { + if len(raw) < 2 { + return nil, 0, 0, io.ErrUnexpectedEOF + } + qos = (raw[0] >> 1) & 0x3 + rem, n, err := decodeRemainingLength(raw[1:]) + if err != nil { + return nil, 0, 0, err + } + body := raw[1+n:] + if len(body) != rem { + return nil, 0, 0, io.ErrUnexpectedEOF + } + pk := packets.Packet{ + ProtocolVersion: mqttProtocolLevel5, + FixedHeader: packets.FixedHeader{ + Type: packets.Publish, + Remaining: rem, + Qos: qos, + }, + } + if err := pk.PublishDecode(body); err != nil { + return nil, 0, qos, err + } + return pk.Payload, pk.PacketID, qos, nil +} + +func decodeRemainingLength(b []byte) (value int, n int, err error) { + var mul uint32 = 1 + var v uint32 + for i := 0; i < len(b) && i < 4; i++ { + v += uint32(b[i]&127) * mul + n++ + if b[i]&128 == 0 { + return int(v), n, nil + } + mul *= 128 + } + return 0, 0, io.ErrUnexpectedEOF +} + +func byteAt(b []byte, i int) byte { + if i < 0 || i >= len(b) { + return 0 + } + return b[i] +} diff --git a/test/load/mqtt_bench.go b/test/load/mqtt_bench.go deleted file mode 100644 index 3d3df16..0000000 --- a/test/load/mqtt_bench.go +++ /dev/null @@ -1,145 +0,0 @@ -// 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 -} diff --git a/test/load/mqtt_bench_test.go b/test/load/mqtt_bench_test.go index 8e90a65..0dd5fc7 100644 --- a/test/load/mqtt_bench_test.go +++ b/test/load/mqtt_bench_test.go @@ -4,72 +4,182 @@ import ( "bytes" "io" "net" + "os" + "strings" "testing" "time" + + "git.asio.asia/nixevol/NixMsg/test/harness" + "github.com/mochi-mqtt/server/v2/packets" ) -// 极简 MQTT broker:读 CONNECT,回 CONNACK accepted,保持连接。 -func startFakeMQTTBroker(t *testing.T) (addr string, closeFn func()) { +func TestEncodeConnectMQTT5(t *testing.T) { + raw, err := encodeConnect("ep-1", "password1234") + if err != nil { + t.Fatal(err) + } + if raw[0]>>4 != packets.Connect { + t.Fatalf("type=0x%02x", raw[0]) + } + rem, n, err := decodeRemainingLength(raw[1:]) + if err != nil { + t.Fatal(err) + } + body := raw[1+n:] + pk := packets.Packet{ + ProtocolVersion: mqttProtocolLevel5, + FixedHeader: packets.FixedHeader{Type: packets.Connect, Remaining: rem}, + } + if err := pk.ConnectDecode(body); err != nil { + t.Fatal(err) + } + if pk.ProtocolVersion != mqttProtocolLevel5 { + t.Fatalf("protocol version %d", pk.ProtocolVersion) + } + if !pk.Connect.Clean { + t.Fatal("CleanStart=false") + } + if pk.Connect.Keepalive != mqttKeepaliveSec { + t.Fatalf("keepalive=%d", pk.Connect.Keepalive) + } + if pk.Connect.ClientIdentifier != "ep-1" { + t.Fatalf("client id %q", pk.Connect.ClientIdentifier) + } + if string(pk.Connect.Username) != "ep-1" { + t.Fatalf("username %q", pk.Connect.Username) + } + if string(pk.Connect.Password) != "password1234" { + t.Fatalf("password mismatch") + } + if !pk.Connect.UsernameFlag || !pk.Connect.PasswordFlag { + t.Fatal("username/password flag") + } + if !bytes.Contains(raw, []byte{0, 4, 'M', 'Q', 'T', 'T', mqttProtocolLevel5}) { + t.Fatal("CONNECT 不是 MQTT 5") + } +} + +func TestConnackReasonMQTT5(t *testing.T) { + reason, err := connackReason([]byte{0x20, 0x03, 0x00, 0x00, 0x00}) + if err != nil || reason != 0 { + t.Fatalf("mqtt5 connack: %d %v", reason, err) + } + reason, err = connackReason([]byte{0x20, 0x02, 0x00, 0x87}) + if err != nil || reason != 0x87 { + t.Fatalf("mqtt311-style: %d %v", reason, err) + } +} + +func TestPercentiles(t *testing.T) { + empty := calcPercentiles(nil) + if empty.N != 0 { + t.Fatalf("empty n=%d", empty.N) + } + one := calcPercentiles([]float64{7}) + if one.P50 != 7 || one.P95 != 7 || one.P99 != 7 { + t.Fatalf("%+v", one) + } + vals := []float64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10} + p := calcPercentiles(vals) + if p.P50 != 5 || p.P95 != 10 || p.P99 != 10 { + t.Fatalf("%+v", p) + } +} + +func TestEndpointIDs(t *testing.T) { + ids, err := EndpointIDs("qb", 3) + if err != nil { + t.Fatal(err) + } + if strings.Join(ids, ",") != "qb0000,qb0001,qb0002" { + t.Fatalf("%v", ids) + } +} + +func TestMQTT5ConnectFakeBroker(t *testing.T) { + addr, stop := startFakeMQTT5(t) + defer stop() + mc, err := harness.DialMQTTTCP(addr, 2*time.Second) + if err != nil { + t.Fatal(err) + } + defer func() { _ = mc.Close() }() + pkt, err := encodeConnect("ep1", "password1234") + if err != nil { + t.Fatal(err) + } + if err := mc.Send(pkt); err != nil { + t.Fatal(err) + } + ack, err := mc.Recv() + if err != nil { + t.Fatal(err) + } + reason, err := connackReason(ack) + if err != nil || reason != 0 { + t.Fatalf("reason=%d err=%v ack=%x", reason, err, ack) + } +} + +func startFakeMQTT5(t *testing.T) (addr string, closeFn func()) { t.Helper() ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } - done := make(chan struct{}) go func() { for { conn, err := ln.Accept() if err != nil { - select { - case <-done: - return - default: - return - } + return } go func(c net.Conn) { defer func() { _ = c.Close() }() _ = c.SetDeadline(time.Now().Add(5 * time.Second)) - buf := make([]byte, 256) + buf := make([]byte, 2048) n, err := c.Read(buf) - if err != nil || n < 2 || buf[0] != 0x10 { + if err != nil || n < 10 || buf[0] != 0x10 { return } - // CONNACK: type 0x20, remaining 2, flags 0, rc 0 - _, _ = c.Write([]byte{0x20, 0x02, 0x00, 0x00}) + if !bytes.Contains(buf[:n], []byte{0, 4, 'M', 'Q', 'T', 'T', mqttProtocolLevel5}) { + return + } + _, _ = c.Write([]byte{0x20, 0x03, 0x00, 0x00, 0x00}) _ = c.SetDeadline(time.Time{}) _, _ = io.Copy(io.Discard, c) }(conn) } }() - return ln.Addr().String(), func() { - close(done) - _ = ln.Close() - } + return ln.Addr().String(), func() { _ = ln.Close() } } -func TestMQTTBenchConnectN(t *testing.T) { - addr, stop := startFakeMQTTBroker(t) - defer stop() +func TestLocalLoadReport(t *testing.T) { + if os.Getenv("NIXMSG_LOAD_VERIFY") == "" { + t.Skip("本地 50×20/s×1m:设置 NIXMSG_LOAD_VERIFY=1") + } + srv, err := harness.Start(harness.Options{}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = srv.Stop() }() - var out bytes.Buffer - b := &MQTTBench{Addr: addr, ClientIDPrefix: "q-t-", Timeout: 2 * time.Second} - defer b.Close() - if err := b.ConnectN(3, &out); err != nil { - t.Fatalf("ConnectN: %v", err) + cfg := Config{ + HTTPBase: srv.HTTPBase, + AdminPass: srv.AdminPassword, + Password: "password1234", + Prefix: "vr", + N: 50, + Rate: 20, + Duration: time.Minute, + GroupSize: 50, + Scenario: "all", } - if b.Alive() != 3 { - t.Fatalf("alive=%d", b.Alive()) + rep, err := Run(cfg) + if err != nil { + t.Fatal(err) } - got := out.String() - if !bytes.Contains(out.Bytes(), []byte("mqtt connections: 3")) { - t.Fatalf("output missing count: %q", got) - } -} - -func TestMQTTBenchRequiresAddr(t *testing.T) { - b := &MQTTBench{} - if err := b.ConnectN(1, nil); err == nil { - t.Fatal("expected error") + if rep.SubmitMs.N == 0 { + t.Fatalf("empty report:\n%s", rep.Format()) } + t.Log(rep.Format()) } diff --git a/test/load/provision.go b/test/load/provision.go new file mode 100644 index 0000000..5d85900 --- /dev/null +++ b/test/load/provision.go @@ -0,0 +1,170 @@ +package load + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" + + "git.asio.asia/nixevol/NixMsg/test/harness" +) + +const maxImportRows = 1000 + +type apiEnvelope struct { + OK bool `json:"ok"` + Data map[string]any `json:"data"` + Error *struct { + Code string `json:"code"` + Message string `json:"message"` + } `json:"error"` +} + +// Provision 按配置准备端:管理 API 批量开通、自助注册,或复用已开通编号。 +func Provision(cfg Config, ids []string) error { + if cfg.Reuse && cfg.AdminUser == "" && cfg.AdminPass == "" && cfg.APIToken == "" && cfg.RegisterCode == "" { + return nil + } + if cfg.RegisterCode != "" && cfg.AdminPass == "" && cfg.APIToken == "" { + return registerAll(cfg, ids) + } + if cfg.AdminPass == "" && cfg.APIToken == "" { + return fmt.Errorf("开通端需要 -admin-pass / -admin-token,或 -register-code;复用已开通端请加 -reuse") + } + ac, err := adminLogin(cfg) + if err != nil { + return err + } + if cfg.Reuse { + return createOneByOne(ac, ids, cfg.Password, true) + } + return importCSV(ac, ids, cfg.Password) +} + +func adminLogin(cfg Config) (*harness.AdminClient, error) { + base := cfg.adminBase() + ac, err := harness.NewAdminClient(base) + if err != nil { + return nil, err + } + ac.HTTP.Timeout = 10 * time.Minute + if cfg.APIToken != "" { + ac.APIToken = cfg.APIToken + return ac, nil + } + user := cfg.AdminUser + if user == "" { + user = "admin" + } + body, _ := json.Marshal(map[string]string{ + "username": user, + "password": cfg.AdminPass, + }) + resp, err := ac.PostJSON("/api/admin/login", body) + if err != nil { + return nil, fmt.Errorf("管理员登录: %w", err) + } + raw, _ := io.ReadAll(resp.Body) + _ = resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("管理员登录失败: %d %s", resp.StatusCode, raw) + } + return ac, nil +} + +func importCSV(ac *harness.AdminClient, ids []string, password string) error { + for start := 0; start < len(ids); start += maxImportRows { + end := start + maxImportRows + if end > len(ids) { + end = len(ids) + } + var b strings.Builder + b.WriteString("id,name,login_password,talk_password,default_delay_seconds,remark\n") + for _, id := range ids[start:end] { + fmt.Fprintf(&b, "%s,%s,%s,,,\n", id, id, password) + } + resp, err := ac.Do(http.MethodPost, "/api/admin/endpoints/import", []byte(b.String()), "text/csv") + if err != nil { + return fmt.Errorf("批量开通: %w", err) + } + raw, _ := io.ReadAll(resp.Body) + _ = resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("批量开通失败: %d %s", resp.StatusCode, raw) + } + } + return nil +} + +func createOneByOne(ac *harness.AdminClient, ids []string, password string, reuse bool) error { + for _, id := range ids { + body, _ := json.Marshal(map[string]string{ + "id": id, + "login_password": password, + }) + resp, err := ac.PostJSON("/api/admin/endpoints", body) + if err != nil { + return fmt.Errorf("开通 %s: %w", id, err) + } + raw, _ := io.ReadAll(resp.Body) + _ = resp.Body.Close() + if resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusCreated { + continue + } + if reuse && isTaken(resp.StatusCode, raw) { + continue + } + return fmt.Errorf("开通 %s 失败: %d %s", id, resp.StatusCode, raw) + } + return nil +} + +func registerAll(cfg Config, ids []string) error { + base := strings.TrimRight(cfg.HTTPBase, "/") + client := &http.Client{Timeout: 30 * time.Second} + for _, id := range ids { + payload, _ := json.Marshal(map[string]string{ + "registration_code": cfg.RegisterCode, + "id": id, + "login_password": cfg.Password, + }) + resp, err := client.Post(base+"/api/client/register", "application/json", bytes.NewReader(payload)) + if err != nil { + return fmt.Errorf("注册 %s: %w", id, err) + } + raw, _ := io.ReadAll(resp.Body) + _ = resp.Body.Close() + if resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusCreated { + continue + } + if cfg.Reuse && isTaken(resp.StatusCode, raw) { + continue + } + return fmt.Errorf("注册 %s 失败: %d %s", id, resp.StatusCode, raw) + } + return nil +} + +func isTaken(status int, raw []byte) bool { + if status == http.StatusConflict { + return true + } + var env apiEnvelope + if json.Unmarshal(raw, &env) != nil { + return false + } + if env.Error != nil && (env.Error.Code == "id_taken" || env.Error.Code == "conflict") { + return true + } + return false +} + +func (c Config) adminBase() string { + if c.AdminHTTP != "" { + return strings.TrimRight(c.AdminHTTP, "/") + } + return strings.TrimRight(c.HTTPBase, "/") +} diff --git a/test/load/q3_live_test.go b/test/load/q3_live_test.go index e39ce5e..7c968b4 100644 --- a/test/load/q3_live_test.go +++ b/test/load/q3_live_test.go @@ -1,21 +1,15 @@ -package load_test +package load import ( - "fmt" "testing" "time" - "git.asio.asia/nixevol/NixMsg/test/accept" "git.asio.asia/nixevol/NixMsg/test/harness" ) -const ( - epPassword = "password1234" - // 本机短时压测目标:几十连接收发。不做 1000 / 10 分钟(见 DEVIATIONS)。 - q3LivePairs = 16 // 32 个端、16 对收发 -) +const epPassword = "password1234" -// TestQ3LiveShortBurst 短时几十连接登录并完成单聊收发。 +// TestQ3LiveShortBurst 短时真实登录+单聊,覆盖原先 Q3 冒烟;完整 50×1m 见 TestLocalLoadReport。 func TestQ3LiveShortBurst(t *testing.T) { srv, err := harness.Start(harness.Options{}) if err != nil { @@ -23,56 +17,23 @@ func TestQ3LiveShortBurst(t *testing.T) { } defer func() { _ = srv.Stop() }() - ac := accept.AdminLogin(t, srv) - type pair struct { - a, b string + cfg := Config{ + HTTPBase: srv.HTTPBase, + AdminPass: srv.AdminPassword, + Password: epPassword, + Prefix: "q3", + N: 8, + Rate: 10, + Duration: 3 * time.Second, + GroupSize: 8, + Scenario: "all", } - pairs := make([]pair, 0, q3LivePairs) - for i := 0; i < q3LivePairs; i++ { - a := fmt.Sprintf("q3la%04d", i) - b := fmt.Sprintf("q3lb%04d", i) - accept.CreateEndpoint(t, ac, a, epPassword) - accept.CreateEndpoint(t, ac, b, epPassword) - pairs = append(pairs, pair{a: a, b: b}) + rep, err := Run(cfg) + if err != nil { + t.Fatal(err) } - - sessions := make([]*accept.MQTTSession, 0, q3LivePairs*2) - defer func() { - for _, s := range sessions { - s.Close() - } - }() - - for _, p := range pairs { - sa := accept.MQTTLogin(t, srv.HTTPBase, p.a, epPassword) - sb := accept.MQTTLogin(t, srv.HTTPBase, p.b, epPassword) - sessions = append(sessions, sa, sb) + if rep.Disconnects != 0 || rep.SubmitMs.N == 0 || len(rep.Errors) != 0 || rep.GroupMembers != 8 { + t.Fatalf("%s", rep.Format()) } - - for i, p := range pairs { - sa := sessions[i*2] - sb := sessions[i*2+1] - msgID := fmt.Sprintf("q3-load-%d", i) - resp := sa.Request(t, map[string]any{ - "v": 1, "type": "send", "rid": fmt.Sprintf("ls%d", i), "id": msgID, - "to": map[string]any{"kind": "endpoint", "id": p.b}, - "body": map[string]any{"enc": "utf8", "data": "burst"}, - "delay_ms": int64(0), - }) - if !resp.OK { - t.Fatalf("pair %d send: %+v", i, resp) - } - msg := sb.WaitType(t, "msg", 15*time.Second) - if msg["id"] != msgID { - t.Fatalf("pair %d want %s got %v", i, msgID, msg) - } - ack := sb.Request(t, map[string]any{ - "v": 1, "type": "ack", "rid": fmt.Sprintf("la%d", i), - "from": p.a, "id": msgID, - }) - if !ack.OK { - t.Fatalf("pair %d ack: %+v", i, ack) - } - } - t.Logf("短时压测通过:%d 连接、%d 对单聊收发成功(未做 1000 连接 / 10 分钟)", q3LivePairs*2, q3LivePairs) + t.Logf("短时压测通过:\n%s", rep.Format()) } diff --git a/test/load/stats.go b/test/load/stats.go new file mode 100644 index 0000000..ba789fd --- /dev/null +++ b/test/load/stats.go @@ -0,0 +1,95 @@ +package load + +import ( + "fmt" + "math" + "sort" + "strings" +) + +// Percentiles 是毫秒延迟分位。 +type Percentiles struct { + P50 float64 + P95 float64 + P99 float64 + N int +} + +func calcPercentiles(ms []float64) Percentiles { + out := Percentiles{N: len(ms)} + if len(ms) == 0 { + return out + } + cp := append([]float64(nil), ms...) + sort.Float64s(cp) + out.P50 = percentileAt(cp, 0.50) + out.P95 = percentileAt(cp, 0.95) + out.P99 = percentileAt(cp, 0.99) + return out +} + +func percentileAt(sorted []float64, p float64) float64 { + n := len(sorted) + if n == 1 { + return sorted[0] + } + idx := int(math.Ceil(p*float64(n))) - 1 + if idx < 0 { + idx = 0 + } + if idx >= n { + idx = n - 1 + } + return sorted[idx] +} + +// Report 是一次压测输出:提交/送达分位、错误码分布、掉线数。 +type Report struct { + Connections int + Alive int + Disconnects int + Sends int + ActualRate float64 + SubmitMs Percentiles + DeliverMs Percentiles + Errors map[string]int + GroupMembers int + GroupQueueMs float64 +} + +func (r Report) Format() string { + var b strings.Builder + fmt.Fprintf(&b, "connections: %d alive: %d disconnects: %d\n", r.Connections, r.Alive, r.Disconnects) + fmt.Fprintf(&b, "sends: %d", r.Sends) + if r.ActualRate > 0 { + fmt.Fprintf(&b, " actual_rate: %.2f/s", r.ActualRate) + } + b.WriteByte('\n') + fmt.Fprintf(&b, "submit_ms: n=%d p50=%.1f p95=%.1f p99=%.1f\n", r.SubmitMs.N, r.SubmitMs.P50, r.SubmitMs.P95, r.SubmitMs.P99) + fmt.Fprintf(&b, "deliver_ms: n=%d p50=%.1f p95=%.1f p99=%.1f\n", r.DeliverMs.N, r.DeliverMs.P50, r.DeliverMs.P95, r.DeliverMs.P99) + if r.GroupMembers > 0 { + fmt.Fprintf(&b, "group: members=%d queue_ms=%.1f\n", r.GroupMembers, r.GroupQueueMs) + } + if len(r.Errors) == 0 { + b.WriteString("errors: (none)\n") + } else { + b.WriteString("errors:") + keys := make([]string, 0, len(r.Errors)) + for k := range r.Errors { + keys = append(keys, k) + } + sort.Strings(keys) + for _, k := range keys { + fmt.Fprintf(&b, " %s=%d", k, r.Errors[k]) + } + b.WriteByte('\n') + } + return b.String() +} + +func addError(m map[string]int, code string) { + if code == "" { + code = "error" + } + m[code]++ +}