package main import ( "context" "encoding/base64" "fmt" "github.com/bytedance/sonic" "github.com/gorilla/websocket" "google.golang.org/protobuf/proto" "math/rand" "net" "runtime" "sonet/api/gen/auth" "sonet/api/gen/chat" "sonet/pkg/config" "sonet/pkg/protocol" "sonet/pkg/utils/logger" "sonet/pkg/utils/security" "sonet/pkg/utils/shutdown" "strconv" "sync" "sync/atomic" "time" ) var ( mockUsers = 100 eachUserSend = 100000 mockNetUsers []*NetUser sendCounter int64 = 0 receiverCounter int64 = 0 wsUrls = []string{"ws://192.168.110.36:7001", "ws://192.168.110.36:7003"} 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{} } 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 main() { 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) conn, err := getConn() if err != nil { panic(err) } err = handleConn(ctx, uid, 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 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() (conn *websocket.Conn, err error) { wsUrl := wsUrls[rand.Intn(len(wsUrls))] conn, _, err = websocket.DefaultDialer.Dial(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, conn *websocket.Conn) (err error) { go listen(ctx, conn) // handshake token, err := getToken(uid) if err != nil { return } 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 { 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 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 }