package main import ( "context" "encoding/base64" "encoding/json" "errors" "flag" "fmt" "github.com/bytedance/sonic" "github.com/gorilla/websocket" clientv3 "go.etcd.io/etcd/client/v3" "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" "google.golang.org/protobuf/proto" "io" "math/rand" "net" "net/http" "reflect" "runtime" "sonet/api/gen/auth" "sonet/api/gen/chat" "sonet/pkg/config" "sonet/pkg/grpc/discovery" "sonet/pkg/protocol" "sonet/pkg/protocol/deliver" "sonet/pkg/utils/collect" "sonet/pkg/utils/logger" "sonet/pkg/utils/security" "sonet/pkg/utils/shutdown" "strconv" "strings" "sync" "sync/atomic" "time" ) var ( gatewayHttp = "http://192.168.110.41:7000" benchmarkMode = "deliver" // deliver/group mockUsers = 2000 eachUserSend = 100 mockNetUsers []*NetUser sendCounter int64 = 0 receiverCounter int64 = 0 httpClient *http.Client etcdClient *clientv3.Client seqId int32 callbacks *collect.ConcurrentMap[int32, func(res any)] useMsSum int64 useMsAvg int64 ) func init() { config.InitLogger(false) //callbacks = make(map[int32]func(res any), 128) //callbackMutex = &sync.Mutex{} callbacks = collect.NewConcurrentMap[int32, func(res any)](16, func(k int32) string { return strconv.Itoa(int(k)) }) httpClient = &http.Client{Timeout: 10 * time.Second} var err error etcdClient, err = clientv3.New(clientv3.Config{ Endpoints: []string{"124.222.131.236:3279"}, Username: "root", Password: "sopod@etcd", }) if err != nil { panic(err) } } // main -mode=deliver -users=10 -send=10 func main() { runtime.GOMAXPROCS(runtime.NumCPU()) flag.StringVar(&gatewayHttp, "gateway", "http://192.168.110.41:7000", "http gateway address") // 124.222.131.236:30830 flag.StringVar(&benchmarkMode, "mode", "group", "benchmark mode: deliver/group") flag.IntVar(&mockUsers, "users", 100, "mock users") flag.IntVar(&eachUserSend, "send", 10, "each user send msg count") flag.Parse() if benchmarkMode == "group" { benchmarkGroup() } else { benchmark() } } func records(ctx context.Context) { ticker := time.NewTicker(time.Second) defer ticker.Stop() var prevSendCounter, prevReceiverCounter int64 for { select { case <-ctx.Done(): return case <-ticker.C: // fmt.Printf("record: %d/%d, total: %d/%d \n", receiverCounter-prevReceiverCounter, sendCounter-prevSendCounter, receiverCounter, sendCounter) logger.Infof("record: %d/%d, total: %d/%d, avg %dms\n", receiverCounter-prevReceiverCounter, sendCounter-prevSendCounter, receiverCounter, sendCounter, useMsAvg) prevReceiverCounter = receiverCounter prevSendCounter = sendCounter } } } type NetUser struct { Uid string Conn *websocket.Conn } func benchmark() { runtime.GOMAXPROCS(runtime.NumCPU()) // go prof.StartPprof(":8888") ctx, cancel := context.WithCancel(context.Background()) go records(ctx) mockUsersConnect(ctx) for i := 0; i < mockUsers; i++ { netUser := mockNetUsers[i] go sendBatchChatMessage(ctx, netUser.Uid, netUser.Conn, eachUserSend) } shutdown.Await() cancel() logger.Infof("total receiver:%d, send:%d\n", receiverCounter, sendCounter) } func benchmarkGroup() { ctx, cancel := context.WithCancel(context.Background()) shutdown.AddHook(cancel) // postal group deliver dis := discovery.NewEtcdDiscovery(etcdClient) picker := deliver.NewPostalPicker(dis, grpc.WithTransportCredentials(insecure.NewCredentials())) if err := picker.Init(ctx); err != nil { panic(err) } groupDeliver := deliver.NewGroupDeliver(chat.Chat_ServiceDesc.ServiceName, picker, nil, nil) if err := groupDeliver.Init(ctx); err != nil { panic(err) } groupId := "9527" mockUsersConnect(ctx) // create room channel, err := send(mockNetUsers[0].Conn, chat.Chat_ServiceDesc.ServiceName, "RoomCreate", &chat.ReqRoomCreate{Rname: "benchmark"}) if err != nil { panic(err) return } // wait result res := <-channel if err, failed := res.(error); failed { panic(err) } else { groupId = res.(*chat.ResRoomCreate).Room.Rid } for _, user := range mockNetUsers { channel, err := send(user.Conn, chat.Chat_ServiceDesc.ServiceName, "RoomJoin", &chat.ReqRoomJoin{Rid: groupId}) if err != nil { panic(err) } // wait result if err, failed := (<-channel).(error); failed { panic(err) } } shutdown.AddHook(func() { // delete room channel, err := send(mockNetUsers[0].Conn, chat.Chat_ServiceDesc.ServiceName, "RoomDissolve", &chat.ReqRoomDissolve{Rid: groupId}) if err != nil { panic(err) return } // wait result if err, failed := (<-channel).(error); failed { panic(err) } logger.Info("test group dissolved") }) go records(ctx) // send group message args := &chat.ReqRoomSend{Rid: groupId, Message: "hi"} for i := 0; i < eachUserSend; i++ { channel, err := send(mockNetUsers[0].Conn, chat.Chat_ServiceDesc.ServiceName, "RoomSend", args) if err != nil { logger.Error("room send error: ", err) continue } // wait result if err, failed := (<-channel).(error); failed { logger.Error("room send failed: ", err) } } shutdown.Await() } // mockUsersConnect 并行快速创建链接 func mockUsersConnect(ctx context.Context) { // initial mock users, join to postal group mockNetUsers = make([]*NetUser, mockUsers) concurrent := 100 wg := &sync.WaitGroup{} wg.Add(concurrent) each := (mockUsers / concurrent) + 1 for i := 0; i < concurrent; i++ { //begin, end := i*each, (i+1)*each //fmt.Printf("%d: %d~%d \n", i, begin, end) go func(segment int) { defer wg.Done() begin, end := segment*each, (segment+1)*each for i := begin; i < end && i < mockUsers; i++ { uid := strconv.Itoa(110000 + i) token, err := getToken(uid) if err != nil { panic(err) } conn, err := getConn(token) if err != nil { panic(err) } err = handleConn(ctx, uid, token, conn) if err != nil { panic(err) } mockNetUsers[i] = &NetUser{Uid: uid, Conn: conn} } }(i) } wg.Wait() } func sendBatchChatMessage(ctx context.Context, uid string, conn *websocket.Conn, count int) { r := rand.New(rand.NewSource(time.Now().UnixMilli())) for i := 0; i < count; i++ { receiverUid := mockNetUsers[r.Intn(mockUsers)].Uid args := &chat.ReqSend{ Receiver: receiverUid, Content: "hi", } channel, err := send(conn, chat.Chat_ServiceDesc.ServiceName, "Send", args) if err != nil { fmt.Println("send failed:", err) return } // wait result if err, failed := (<-channel).(error); failed { fmt.Println("send failed:", err) } } } func getConn(token string) (conn *websocket.Conn, err error) { req, err := http.NewRequest("GET", gatewayHttp+"/api/lb/ws", io.LimitReader(nil, 0)) if err != nil { return } req.Header.Set("Authorization", token) resp, err := httpClient.Do(req) if err != nil { logger.Error("get ws endpoint error: ", err) return } defer resp.Body.Close() bytes, err := io.ReadAll(resp.Body) if err != nil { return } body := make(map[string]any, 2) err = json.Unmarshal(bytes, &body) if err != nil { return } wsUrl := body["data"].(map[string]any)["ws"].(string) // wsUrl := wsUrls[rand.Intn(len(wsUrls))] conn, _, err = websocket.DefaultDialer.Dial("ws://"+wsUrl+"/ws", nil) if err != nil { return } return } func getToken(uid string) (token string, err error) { // generate token subject := &auth.Subject{Uid: uid, Username: uid, Time: time.Now().UnixMilli()} bytes, err := sonic.Marshal(subject) if err != nil { return } keyBytes, err := base64.StdEncoding.DecodeString("9Nz3Y6DES3msAFndz4QsJAECUIFkf+KKaRa+jRnNALk=") if err != nil { return } encode, err := security.EncryptAesCBC(bytes, keyBytes) if err != nil { return } token = base64.URLEncoding.EncodeToString(encode) return } // handleConn listen and auth verify func handleConn(ctx context.Context, uid string, token string, conn *websocket.Conn) (err error) { go listen(ctx, conn) // handshake args := &auth.ReqVerify{Token: token} channel, err := send(conn, auth.Auth_ServiceDesc.ServiceName, "Verify", args) if err != nil { return } ch := time.After(time.Second * 5) select { case <-ch: err = errors.New("req verify timeout") case res := <-channel: if e, failed := res.(error); failed { err = e } } return } func listen(ctx context.Context, conn *websocket.Conn) { var e error defer func() { if e != nil && !strings.Contains(e.Error(), "close") { fmt.Println("conn error: ", e) } }() go func() { select { case <-ctx.Done(): // finished close connection _ = conn.Close() } }() for { t, message, err := conn.ReadMessage() if err != nil { if opErr, ok := err.(*net.OpError); ok && opErr.Err.Error() == "use of closed network connection" { return } e = err _ = conn.Close() break } if t != websocket.BinaryMessage { e = fmt.Errorf("unknown message type: %d\n", t) continue } payload, err := protocol.DecodeSo(message) if err != nil { logger.Error("decode message error: ", err) continue } header := payload.Header if benchmarkMode == "group" { atomic.AddInt64(&receiverCounter, 1) } // notice message if header.Type == protocol.TypeNotice { switch header.Target { case "ChatMessage": msg := &chat.ChatMessage{} if err := proto.Unmarshal(payload.Body, msg); err != nil { logger.Error("unmarshal proto message error: ", err) return } //atomic.AddInt64(&receiverCounter, 1) // fmt.Printf("reciever: %v\n", msg) } continue } callback, ok := callbacks.LoadAndDelete(header.SeqId) if !ok { logger.Errorf("callback %d not found", header.SeqId) continue } if header.Type == protocol.TypeError { err = errors.New(string(payload.Body)) callback(err) continue } // rpc response var msg proto.Message if svc, ok := protoStructs[header.Svc]; ok { if s, ok := svc[header.Target]; ok { val := reflect.New(reflect.TypeOf(s).Elem()) msg = val.Interface().(proto.Message) } } if msg == nil { // protobuf.Empty // logger.Warning("unknown rpc response: ", len(payload.Body), header) callback(nil) continue } err = proto.Unmarshal(payload.Body, msg) if err != nil { logger.Error("unmarshal rcp res body error: ", payload, err) callback(fmt.Errorf("unmarshal rcp res body error: %v", err)) return } callback(msg) } } var protoStructs = map[string]map[string]proto.Message{ auth.Auth_ServiceDesc.ServiceName: { "Subject": &auth.Subject{}, }, chat.Chat_ServiceDesc.ServiceName: { "ResSend": &chat.ResSend{}, "ResRoomCreate": &chat.ResRoomCreate{}, "ResRoomInfo": &chat.ResRoomInfo{}, "ResRoomList": &chat.ResRoomList{}, }, } func send(conn *websocket.Conn, svc, method string, reqArgs proto.Message) (channel chan any, err error) { seq := atomic.AddInt32(&seqId, 1) header := &protocol.Header{ Magic: protocol.Magic, Type: 1, UrlType: 1, SerializeType: 1, SeqId: seq, Svc: svc, Target: method, } payload := &protocol.Payload{Header: header} payload.Body, err = proto.Marshal(reqArgs) if err != nil { panic(err) } bytes, err := protocol.EncodeSo(payload) if err != nil { panic(err) } begin := time.Now().UnixMilli() // seq, await channel = make(chan any) // ready callback callbacks.Store(seq, func(res any) { ms := time.Now().UnixMilli() - begin msSum := atomic.AddInt64(&useMsSum, ms) counter := atomic.AddInt64(&receiverCounter, 1) atomic.StoreInt64(&useMsAvg, msSum/counter) channel <- res }) // send message err = conn.WriteMessage(websocket.BinaryMessage, bytes) if err != nil { return } atomic.AddInt64(&sendCounter, 1) return }