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 provErr := Provision(cfg, ids); provErr != nil { return nil, provErr } 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 }