package main import ( "context" "encoding/base64" "encoding/json" "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" "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/logger" "sonet/pkg/utils/security" "sonet/pkg/utils/shutdown" "strconv" "strings" "sync" "sync/atomic" "time" ) var ( benchmarkMode = "deliver" // deliver / groupDeliver mockUsers = 2000 eachUserSend = 100 mockNetUsers []*NetUser sendCounter int64 = 0 receiverCounter int64 = 0 //wsUrls = []string{"ws://192.168.110.36:7001", "ws://192.168.110.36:7003"} gatewayHttp = "http://192.168.110.36:7000" httpClient *http.Client etcdClient *clientv3.Client seqId int32 callbacks map[int32]func(res any) callbackMutex *sync.Mutex useMsSum int64 useMsAvg int64 ) func init() { config.InitLogger() callbacks = make(map[int32]func(res any), 128) callbackMutex = &sync.Mutex{} 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) } } func main() { // benchmark() benchmarkGroup() } 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) mockNetUsers = make([]*NetUser, mockUsers) // initial uids for i := 0; 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, } } 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() { benchmarkMode = "groupDeliver" runtime.GOMAXPROCS(runtime.NumCPU()) 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" // initial mock users, join to postal group mockNetUsers = make([]*NetUser, mockUsers) for i := 0; 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, } // join group err = groupDeliver.GroupJoin(context.Background(), uid, []string{groupId}) if err != nil { panic(err) } } shutdown.AddHook(func() { if err := groupDeliver.GroupDissolve(context.Background(), groupId); err != nil { logger.Error("dissolve group error: ", err) return } logger.Info("test group dissolved") }) go records(ctx) // send group message message := &chat.ChatMessage{Sender: "100001", Content: "hello"} for i := 0; i < eachUserSend; i++ { err := groupDeliver.DeliverGroup(context.Background(), groupId, message) if err != nil { logger.Error("deliver group error:", err) } atomic.AddInt64(&sendCounter, 1) } shutdown.Await() } 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: "hello", } channel, err := send(conn, chat.Chat_ServiceDesc.ServiceName, "send", args) if err != nil { fmt.Println("send failed:", err) return } res := <-channel // wait result err, failed := res.(error) if failed { fmt.Println("send failed:", err) } else { // fmt.Println("send ok:", res) } } } 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 } res := <-channel // fmt.Printf("%v\n", res) 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 == "groupDeliver" { atomic.AddInt64(&receiverCounter, 1) } if header.Type == 4 { callbackMutex.Lock() callback, ok := callbacks[header.SeqId] if ok { delete(callbacks, header.SeqId) } callbackMutex.Unlock() if !ok { e = fmt.Errorf("callback %d not found", header.SeqId) return } callback(err) continue } // 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 } // rpc response var msg proto.Message switch header.Svc { case auth.Auth_ServiceDesc.ServiceName: switch header.Target { case "Subject": // fmt.Println("res verify...") msg = &auth.Subject{} } case chat.Chat_ServiceDesc.ServiceName: switch header.Target { case "ResSend": // fmt.Println("res send...") msg = &chat.ResSend{} } } if msg == nil { logger.Warning("unknown rpc response: ", len(payload.Body), header) continue } err = proto.Unmarshal(payload.Body, msg) if e != nil { logger.Error("unmarshal rcp res body error: ", payload, err) return } callbackMutex.Lock() callback, ok := callbacks[header.SeqId] if ok { delete(callbacks, header.SeqId) } callbackMutex.Unlock() if !ok { e = fmt.Errorf("callback %d not found", header.SeqId) return } callback(msg) } } 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 callbackMutex.Lock() callbacks[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 } callbackMutex.Unlock() // send message err = conn.WriteMessage(websocket.BinaryMessage, bytes) if err != nil { return } atomic.AddInt64(&sendCounter, 1) return }