Files
NixMsg/test/load/bench.go
T

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
}