From 6f40a94ec8e0b9a410e33ddb8f670a73d48e4977 Mon Sep 17 00:00:00 2001 From: tangmingyou Date: Thu, 8 Feb 2024 17:57:16 +0800 Subject: [PATCH] postal cluster dispatch --- benchmark/main.go | 158 +++++++++++++++--- build-all.bat | 3 + cmd/chat/main.go | 2 +- cmd/gateway_ws/main.go | 2 + .../gateway_ws/gws_server/conn_handler.go | 45 ++++- internal/postal/group/group.go | 15 +- .../postal/logic/postal_cluster_server.go | 10 +- internal/postal/logic/postal_server.go | 34 ++-- pkg/grpc/discovery/discovery.go | 8 +- pkg/grpc/discovery/etcd_naming.go | 59 ++++++- pkg/protocol/deliver/group_deliver.go | 12 +- .../session/{session.go => channel.go} | 23 ++- 12 files changed, 307 insertions(+), 64 deletions(-) create mode 100644 build-all.bat rename pkg/protocol/session/{session.go => channel.go} (68%) diff --git a/benchmark/main.go b/benchmark/main.go index 3df4abd..eef6bb9 100644 --- a/benchmark/main.go +++ b/benchmark/main.go @@ -3,38 +3,51 @@ 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 ( - mockUsers = 100 - eachUserSend = 100000 + 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"} - seqId int32 - callbacks map[int32]func(res any) - callbackMutex *sync.Mutex - useMsSum int64 - useMsAvg int64 + //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() { @@ -42,10 +55,21 @@ func init() { 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() + // benchmark() + benchmarkGroup() } func records(ctx context.Context) { @@ -83,11 +107,15 @@ func benchmark() { // initial uids for i := 0; i < mockUsers; i++ { uid := strconv.Itoa(110000 + i) - conn, err := getConn() + token, err := getToken(uid) if err != nil { panic(err) } - err = handleConn(ctx, uid, conn) + conn, err := getConn(token) + if err != nil { + panic(err) + } + err = handleConn(ctx, uid, token, conn) if err != nil { panic(err) } @@ -108,6 +136,75 @@ func benchmark() { 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++ { @@ -132,9 +229,30 @@ func sendBatchChatMessage(ctx context.Context, uid string, conn *websocket.Conn, } } -func getConn() (conn *websocket.Conn, err error) { - wsUrl := wsUrls[rand.Intn(len(wsUrls))] - conn, _, err = websocket.DefaultDialer.Dial(wsUrl+"/ws", nil) +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 } @@ -161,14 +279,10 @@ func getToken(uid string) (token string, err error) { } // handleConn listen and auth verify -func handleConn(ctx context.Context, uid string, conn *websocket.Conn) (err error) { +func handleConn(ctx context.Context, uid string, token 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) @@ -186,7 +300,7 @@ func handleConn(ctx context.Context, uid string, conn *websocket.Conn) (err erro func listen(ctx context.Context, conn *websocket.Conn) { var e error defer func() { - if e != nil { + if e != nil && !strings.Contains(e.Error(), "close") { fmt.Println("conn error: ", e) } }() @@ -221,6 +335,10 @@ func listen(ctx context.Context, conn *websocket.Conn) { } header := payload.Header + if benchmarkMode == "groupDeliver" { + atomic.AddInt64(&receiverCounter, 1) + } + if header.Type == 4 { callbackMutex.Lock() callback, ok := callbacks[header.SeqId] diff --git a/build-all.bat b/build-all.bat new file mode 100644 index 0000000..c358464 --- /dev/null +++ b/build-all.bat @@ -0,0 +1,3 @@ +@title build all sonet + +build.bat gateway_ws && build.bat gateway_http && build.bat auth && build.bat chat && build.bat mahjong diff --git a/cmd/chat/main.go b/cmd/chat/main.go index 87b6be7..fab84c5 100644 --- a/cmd/chat/main.go +++ b/cmd/chat/main.go @@ -42,7 +42,7 @@ func main() { if err != nil { panic(err) } - groupDeli := deliver.NewGroupDeliver(chat.Chat_ServiceDesc.ServiceName, logic.NewRoomLoader(), consumer) + groupDeli := deliver.NewGroupDeliver(chat.Chat_ServiceDesc.ServiceName, picker, logic.NewRoomLoader(), consumer) chatServer := logic.NewChatServer(deli, groupDeli) go func() { diff --git a/cmd/gateway_ws/main.go b/cmd/gateway_ws/main.go index 391b832..bd12434 100644 --- a/cmd/gateway_ws/main.go +++ b/cmd/gateway_ws/main.go @@ -5,6 +5,7 @@ import ( clientv3 "go.etcd.io/etcd/client/v3" "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" + "runtime" "sonet/internal/gateway_ws/gws_server" "sonet/internal/postal/group" "sonet/internal/postal/logic" @@ -31,6 +32,7 @@ type GatewayWsConfig struct { // websocket server with postalService func main() { + runtime.GOMAXPROCS(runtime.NumCPU()) appConf := &GatewayWsConfig{} conf := config.LoadConfig(appConf, "cmd/gateway_ws") diff --git a/internal/gateway_ws/gws_server/conn_handler.go b/internal/gateway_ws/gws_server/conn_handler.go index c1ef18b..4f8320d 100644 --- a/internal/gateway_ws/gws_server/conn_handler.go +++ b/internal/gateway_ws/gws_server/conn_handler.go @@ -7,6 +7,7 @@ import ( "google.golang.org/protobuf/proto" "regexp" "sonet/api/gen/auth" + "sonet/api/gen/postal" "sonet/pkg/grpc/generic" "sonet/pkg/protocol" "sonet/pkg/protocol/session" @@ -140,14 +141,16 @@ func (c *GwsHandler) OnMessage(socket *gws.Conn, message *gws.Message) { // auth verify success if service == "Auth" && method == "Verify" { - sessionUid, err = c.extractAuthVerify(conn, payload.Body) + var channel *session.Channel + sessionUid, channel, err = c.extractAuthVerify(conn, payload.Body) if err != nil { logger.Error("extractAuthVerify error: ", err) writeError(conn, header.SeqId, "server error") return } socket.Session().Store(SessionUidKey, sessionUid) - // todo online event + // channel + go c.dispatch(channel) } header.Type = protocol.TypeResponse @@ -163,7 +166,40 @@ func (c *GwsHandler) OnMessage(socket *gws.Conn, message *gws.Message) { } } -func (c *GwsHandler) extractAuthVerify(conn *gwsConn, resp []byte) (sessionUid string, err error) { +func encodeDeliverMessage(msg *postal.Message) (bytes []byte, err error) { + header := &protocol.Header{ + Magic: protocol.Magic, + Type: protocol.TypeNotice, + UrlType: 1, + SerializeType: 1, + Svc: msg.Svc, + Target: msg.Msg, + } + payload := &protocol.Payload{Header: header, Body: msg.Body} + bytes, err = protocol.EncodeSo(payload) + if err != nil { + logger.Error("encode deliver message error: ", err) + } + return +} + +// todo close +func (c *GwsHandler) dispatch(channel *session.Channel) { + for { + msg := channel.Ready() + bytes, err := encodeDeliverMessage(msg) + if err != nil { + logger.Error("encode dispatch message error:", err) + continue + } + err = channel.Conn.Write(bytes) + if err != nil { + logger.Error("write dispatch message error:", err) + } + } +} + +func (c *GwsHandler) extractAuthVerify(conn *gwsConn, resp []byte) (sessionUid string, channel *session.Channel, err error) { authSubject := &auth.Subject{} err = proto.Unmarshal(resp, authSubject) if err != nil { @@ -180,7 +216,8 @@ func (c *GwsHandler) extractAuthVerify(conn *gwsConn, resp []byte) (sessionUid s logger.Infof("close subject old conn: %s\n", sessionUid) } // 存储session到内存 - c.sessionStore.Store(sessionUid, session.NewChannel(sessionUid, conn)) + channel = session.NewChannel(sessionUid, conn) + c.sessionStore.Store(sessionUid, channel) logger.Infof("subject online: %s\n", sessionUid) return } diff --git a/internal/postal/group/group.go b/internal/postal/group/group.go index 3e880d0..e7d26c6 100644 --- a/internal/postal/group/group.go +++ b/internal/postal/group/group.go @@ -3,6 +3,7 @@ package group import ( "errors" "github.com/redis/go-redis/v9" + "sonet/api/gen/postal" "sonet/pkg/protocol/session" "sonet/pkg/utils/logger" "sync" @@ -36,27 +37,27 @@ func (d *RedisPostalDao) LoadGroupIdsByUid(uid string) (gids []string, err error type Group struct { sync.RWMutex Gid string - uids map[string]session.NetConn + uids map[string]*session.Channel } func NewGroup(gid string) *Group { return &Group{ Gid: gid, - uids: make(map[string]session.NetConn), + uids: make(map[string]*session.Channel, 6), } } -func (g *Group) Write(data []byte) { +func (g *Group) Write(msg *postal.Message) { g.RLock() defer g.RUnlock() - for uid, conn := range g.uids { - if err := conn.Write(data); err != nil { + for uid, channel := range g.uids { + if err := channel.Push(msg); err != nil { logger.Errorf("group send %s.%s error: ", g.Gid, uid, err) } } } -func (g *Group) Join(uid string, conn session.NetConn) { +func (g *Group) Join(uid string, conn *session.Channel) { g.Lock() defer g.Unlock() g.uids[uid] = conn @@ -74,7 +75,7 @@ func (g *Group) Dismiss() { } -func (g *Group) Load(uid string) (conn session.NetConn, ok bool) { +func (g *Group) Load(uid string) (conn *session.Channel, ok bool) { g.RLock() defer g.RUnlock() conn, ok = g.uids[uid] diff --git a/internal/postal/logic/postal_cluster_server.go b/internal/postal/logic/postal_cluster_server.go index efebeaa..c82205a 100644 --- a/internal/postal/logic/postal_cluster_server.go +++ b/internal/postal/logic/postal_cluster_server.go @@ -69,10 +69,10 @@ func (s *PostalClusterServer) Run(postalAddr string, opts ...grpc.ServerOption) } func (s *PostalClusterServer) Redeliver(ctx context.Context, req *postal.ReqRedeliver) (res *postal.ResDeliver, err error) { - bytes, err := encodeDeliverMessage(req.Msg) - if err != nil { - return - } + //bytes, err := encodeDeliverMessage(req.Msg) + //if err != nil { + // return + //} var redelivers []string for _, receiver := range req.Receivers { channel, ok := s.sessionStore.Load(receiver) @@ -80,7 +80,7 @@ func (s *PostalClusterServer) Redeliver(ctx context.Context, req *postal.ReqRede redelivers = append(redelivers, receiver) continue } - err = channel.Conn.Write(bytes) + err = channel.Push(req.Msg) if err != nil { logger.Errorf("redeliver channel %s write error: %v", receiver, err) continue diff --git a/internal/postal/logic/postal_server.go b/internal/postal/logic/postal_server.go index ba0cc77..323cb47 100644 --- a/internal/postal/logic/postal_server.go +++ b/internal/postal/logic/postal_server.go @@ -131,13 +131,13 @@ func (s *PostalServer) processSessionStoreEvent(ctx context.Context) { func (s *PostalServer) Deliver(ctx context.Context, req *postal.ReqDeliver) (res *postal.ResDeliver, err error) { receiver, ok := s.sessionStore.Load(req.Receiver) if ok { - var bytes []byte - bytes, err = encodeDeliverMessage(req.Msg) - if err != nil { - return - } + //var bytes []byte + //bytes, err = encodeDeliverMessage(req.Msg) + //if err != nil { + // return + //} // write msg - err = receiver.Conn.Write(bytes) + err = receiver.Push(req.Msg) if err != nil { return } @@ -161,10 +161,10 @@ func (s *PostalServer) Deliver(ctx context.Context, req *postal.ReqDeliver) (res } func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverBatch) (res *postal.ResDeliver, err error) { - bytes, err := encodeDeliverMessage(req.Msg) - if err != nil { - return - } + //bytes, err := encodeDeliverMessage(req.Msg) + //if err != nil { + // return + //} var redeliverReceivers []string for _, receiverId := range req.Receivers { @@ -173,7 +173,7 @@ func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverB redeliverReceivers = append(redeliverReceivers, receiverId) continue } - err = receiver.Conn.Write(bytes) + err = receiver.Push(req.Msg) if err != nil { logger.Errorf("deliver to %s error: ", receiver, err) } @@ -200,17 +200,17 @@ func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverB // DeliverGroup postal broadcast func (s *PostalServer) DeliverGroup(ctx context.Context, req *postal.ReqDeliverGroup) (res *postal.ResDeliver, err error) { - bytes, err := encodeDeliverMessage(req.Msg) - if err != nil { - return - } + //bytes, err := encodeDeliverMessage(req.Msg) + //if err != nil { + // return + //} res = &postal.ResDeliver{Ok: true} g, ok := s.groupStore.Load(req.Gid) if !ok { return } - g.Write(bytes) + g.Write(req.Msg) return } @@ -223,7 +223,7 @@ func (s *PostalServer) GroupJoin(ctx context.Context, req *postal.ReqGroupJoin) } for _, gid := range req.Gids { g, _ := s.groupStore.ComputeIfAbsent(gid, func(gid string) *group.Group { return group.NewGroup(gid) }) - g.Join(req.Uid, channel.Conn) + g.Join(req.Uid, channel) channel.GroupJoin(gid) } return diff --git a/pkg/grpc/discovery/discovery.go b/pkg/grpc/discovery/discovery.go index 2e8f9a3..bb2a6fb 100644 --- a/pkg/grpc/discovery/discovery.go +++ b/pkg/grpc/discovery/discovery.go @@ -3,6 +3,7 @@ package discovery import ( "context" "google.golang.org/grpc/resolver" + "sonet/pkg/utils/logger" "strconv" ) @@ -47,8 +48,11 @@ func (s Server) GetWeight() (weight int) { if !ok { return } - if w, err := strconv.Atoi(v); err != nil { - weight = w + w, err := strconv.Atoi(v) + if err != nil { + logger.Warning("failed parse discovery server attr weight: ", v) + return } + weight = w return } diff --git a/pkg/grpc/discovery/etcd_naming.go b/pkg/grpc/discovery/etcd_naming.go index ad402e2..6a90eee 100644 --- a/pkg/grpc/discovery/etcd_naming.go +++ b/pkg/grpc/discovery/etcd_naming.go @@ -75,17 +75,70 @@ func (r *EtcdDiscovery) Registry(ctx context.Context, server Server) (err error) ) // keepalive lease - keepAliveCh, err := r.client.KeepAlive(context.Background(), lease.ID) + keepCtx, keepCancel := context.WithCancel(context.Background()) + keepAliveCh, err := r.client.KeepAlive(keepCtx, lease.ID) + if err != nil { + keepCancel() + logger.Error("registry keepalive error: ", err) + return + } + + // ticker := time.NewTicker(time.Second * time.Duration(DefaultRegisterTTL)) + w := r.client.Watch(ctx, endpointKey) go func() { for { select { case <-ctx.Done(): logger.Info("registry keepalive done") - ctx, c := context.WithTimeout(context.Background(), time.Second*2) + ctx, c := context.WithTimeout(context.Background(), time.Second*3) _, _ = r.client.Revoke(ctx, lease.ID) c() + keepCancel() + //ticker.Stop() return - case _ = <-keepAliveCh: + case <-keepAliveCh: + //logger.Infof("keepalive: lease id %d, %+v", lease.ID, res) + case res := <-w: + if err := res.Err(); err != nil { + logger.Errorf("registry watch endpoint %s error: %v", endpointKey, err) + continue + } + deleted := false + for _, event := range res.Events { + if event.Type == clientv3.EventTypeDelete { + deleted = true + } + } + // endpoint key被删除, 重放 + if deleted { + logger.Infof("registry endpoint %s deleted ", endpointKey) + // 删除旧的 lease + keepCancel() + _, _ = r.client.Revoke(context.Background(), lease.ID) + lease, err = r.client.Grant(ctx, DefaultRegisterTTL) + if err != nil { + logger.Error("registry grant lease again error: ", err) + } + err = em.AddEndpoint(ctx, + endpointKey, + endpoints.Endpoint{ + Addr: addr, + Metadata: meta, + }, + clientv3.WithLease(lease.ID), + ) + if err != nil { + logger.Error("refresh endpoint error: ", err) + continue + } + // keepalive lease + keepCtx, keepCancel = context.WithCancel(context.Background()) + keepAliveCh, err = r.client.KeepAlive(keepCtx, lease.ID) + if err != nil { + logger.Error("registry keepalive error: ", err) + } + } + //case <-ticker.C: // 定时重放防止etcd中key被删除 } } }() diff --git a/pkg/protocol/deliver/group_deliver.go b/pkg/protocol/deliver/group_deliver.go index 4c7fa47..f0761bf 100644 --- a/pkg/protocol/deliver/group_deliver.go +++ b/pkg/protocol/deliver/group_deliver.go @@ -25,18 +25,22 @@ type GroupDeliver struct { func NewGroupDeliver( msgInServiceName string, + postalPicker *PostalPicker, groupLoader GroupLoader, consumer mq.Consumer, ) *GroupDeliver { return &GroupDeliver{ - svcName: msgInServiceName, - groupLoader: groupLoader, - consumer: consumer, + svcName: msgInServiceName, + postalPicker: postalPicker, + groupLoader: groupLoader, + consumer: consumer, } } func (d *GroupDeliver) Init(ctx context.Context) (err error) { - err = d.subscribe() + if d.consumer != nil { + err = d.subscribe() + } return } diff --git a/pkg/protocol/session/session.go b/pkg/protocol/session/channel.go similarity index 68% rename from pkg/protocol/session/session.go rename to pkg/protocol/session/channel.go index db7aecb..2b650cd 100644 --- a/pkg/protocol/session/session.go +++ b/pkg/protocol/session/channel.go @@ -1,6 +1,12 @@ package session -import "sync" +import ( + "errors" + "sonet/api/gen/postal" + "sync" +) + +var ErrChannelFullMsgDropped = errors.New("channel full, msg dropped") // NetConn 各类型连接的 write 接口 type NetConn interface { @@ -10,6 +16,7 @@ type NetConn interface { type Channel struct { Uid string Conn NetConn + ch chan *postal.Message groups []string // 记录 uid 对应的群组列表 lock *sync.Mutex } @@ -18,10 +25,24 @@ func NewChannel(uid string, conn NetConn) *Channel { return &Channel{ Uid: uid, Conn: conn, + ch: make(chan *postal.Message, 16), lock: &sync.Mutex{}, } } +func (c *Channel) Push(msg *postal.Message) (err error) { + select { + case c.ch <- msg: + default: + err = ErrChannelFullMsgDropped + } + return +} + +func (c *Channel) Ready() *postal.Message { + return <-c.ch +} + func (c *Channel) GroupJoin(gid string) { c.lock.Lock() defer c.lock.Unlock()