You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 

493 lines
12 KiB

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
}