339 lines
7.7 KiB
Go
339 lines
7.7 KiB
Go
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
|
|
}
|