fix: 压测客户端改为 MQTT 5 并完成 hello 与收发统计
This commit is contained in:
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user