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.
 
 

334 lines
7.2 KiB

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
}