From 669a8eb1515d5a0c6515acd37f76e560f970f341 Mon Sep 17 00:00:00 2001 From: tangmingyou Date: Sat, 3 Feb 2024 17:36:45 +0800 Subject: [PATCH] deliver group --- api/chat.proto | 7 +- cmd/chat/main.go | 2 +- cmd/gateway_ws/config.toml | 2 +- cmd/mahjong/main.go | 3 +- internal/chat/logic/chat_server.go | 122 +++++++++++++++--- internal/gateway_http/logic/http_server.go | 10 +- internal/gateway_http/logic/postal_monitor.go | 80 ------------ .../gateway_ws/gws_server/conn_handler.go | 9 +- internal/gateway_ws/server/ws_server.go | 14 +- internal/mahjong/logic/mahjong_server.go | 9 +- internal/postal/logic/postal_server.go | 24 +++- pkg/grpc/generic/generic_client.go | 4 +- pkg/plugins/mq/nats_jet_stream.go | 1 + pkg/protocol/deliver/deliver_group.go | 94 +++++++++++--- pkg/protocol/event/event.go | 50 +++++-- 15 files changed, 279 insertions(+), 152 deletions(-) delete mode 100644 internal/gateway_http/logic/postal_monitor.go diff --git a/api/chat.proto b/api/chat.proto index dea8929..bb26a03 100644 --- a/api/chat.proto +++ b/api/chat.proto @@ -23,9 +23,10 @@ service Chat { } message ChatMessage { - string sender = 1; - string content = 7; -// Message message = 8; + int32 type = 1; // 0默认单聊消息, 1群消息 + string sender = 2; + string gid = 3; + string content = 15; } message ReqSend { diff --git a/cmd/chat/main.go b/cmd/chat/main.go index 35cb915..51d91ee 100644 --- a/cmd/chat/main.go +++ b/cmd/chat/main.go @@ -32,7 +32,7 @@ func main() { if err := deli.InitWithResolver(context.Background(), dis); err != nil { panic(err) } - chatServer := logic.NewChatServer(deli) + chatServer := logic.NewChatServer(deli, nil) go func() { err = chatServer.Run(conf.Grpc, dis) diff --git a/cmd/gateway_ws/config.toml b/cmd/gateway_ws/config.toml index 7a97676..5cfdfd1 100644 --- a/cmd/gateway_ws/config.toml +++ b/cmd/gateway_ws/config.toml @@ -1,6 +1,6 @@ [app] httpPort = 7001 -endpointAddress = "127.0.0.1:7001" +endpointAddress = "192.168.110.41:7001" subjectCacheTopic = "wsgate:subject:" subjectLrcExpiration = "10m" subjectLrcCleanupInterval = "5m" diff --git a/cmd/mahjong/main.go b/cmd/mahjong/main.go index 7c1fb3e..29e48ca 100644 --- a/cmd/mahjong/main.go +++ b/cmd/mahjong/main.go @@ -34,6 +34,7 @@ func main() { if err != nil { panic(err) } + deli := deliver.NewDeliver(mahjong.Mahjong_ServiceDesc.ServiceName) if err := deli.InitWithResolver(context.Background(), dis); err != nil { panic(err) @@ -49,7 +50,7 @@ func main() { if err != nil { panic(err) } - mahjongServer := logic.NewMahjongServer(mjStore, deli, auth.NewAuthClient(authConn)) + mahjongServer := logic.NewMahjongServer(mjStore, auth.NewAuthClient(authConn), deli) go func() { err = mahjongServer.Run(conf.Grpc, dis) if err != nil { diff --git a/internal/chat/logic/chat_server.go b/internal/chat/logic/chat_server.go index 6ff2342..dd684c4 100644 --- a/internal/chat/logic/chat_server.go +++ b/internal/chat/logic/chat_server.go @@ -2,6 +2,7 @@ package logic import ( "context" + "fmt" "google.golang.org/grpc" "google.golang.org/grpc/reflection" "google.golang.org/protobuf/types/known/emptypb" @@ -12,18 +13,29 @@ import ( "sonet/pkg/grpc/interceptor" "sonet/pkg/protocol/deliver" "sonet/pkg/protocol/session" + "sonet/pkg/utils/collect" "sonet/pkg/utils/logger" "sonet/pkg/utils/shutdown" + "strconv" + "sync/atomic" + "time" ) type ChatServer struct { chat.UnimplementedChatServer - deliver *deliver.Deliver + deliver *deliver.Deliver + deliverGroup *deliver.GroupDeliver + roomId int64 + rooms *collect.ConcurrentMap[string, *chat.Group] } -func NewChatServer(deliver *deliver.Deliver) *ChatServer { +func NewChatServer(deliver *deliver.Deliver, deliverGroup *deliver.GroupDeliver) *ChatServer { + rooms := collect.NewConcurrentMap[string, *chat.Group](32, func(k string) string { return k }) return &ChatServer{ - deliver: deliver, + deliver: deliver, + deliverGroup: deliverGroup, + roomId: 90000, + rooms: rooms, } } @@ -71,7 +83,7 @@ func (s *ChatServer) Send(ctx context.Context, req *chat.ReqSend) (*chat.ResSend return nil, err } - // TODO 存储消息,ack 队列重发(可异步发送),asc超时第二次发送时同步判断是否不在线? + // TODO 消息队列消峰发送,ack 队列重发(可异步发送),asc超时第二次发送时同步判断是否不在线? // 投递消息 message := &chat.ChatMessage{Sender: subject.Uid, Content: req.Content} @@ -83,30 +95,104 @@ func (s *ChatServer) Send(ctx context.Context, req *chat.ReqSend) (*chat.ResSend return &chat.ResSend{}, nil } -func (s *ChatServer) RoomSend(ctx context.Context, send *chat.ReqRoomSend) (*chat.ResSend, error) { - return nil, nil +func (s *ChatServer) RoomSend(ctx context.Context, req *chat.ReqRoomSend) (res *chat.ResSend, err error) { + subject, err := session.GetSubject(ctx) + if err != nil { + return nil, err + } + gMessage := &chat.ChatMessage{ + Type: 1, + Sender: subject.Uid, + Gid: req.Gid, + Content: req.Message, + } + err = s.deliverGroup.DeliverGroup(ctx, req.Gid, gMessage) + return } -func (s *ChatServer) RoomCreate(ctx context.Context, create *chat.ReqRoomCreate) (*chat.ResRoomCreate, error) { - return nil, nil +func (s *ChatServer) RoomCreate(ctx context.Context, req *chat.ReqRoomCreate) (res *chat.ResRoomCreate, err error) { + subject, err := session.GetSubject(ctx) + if err != nil { + return nil, err + } + logger.Infof("create room: %s", req.Gname) + + group := &chat.Group{ + Gid: strconv.Itoa(int(atomic.AddInt64(&s.roomId, 1))), + MasterUid: subject.Uid, + UserLimit: 10000, + UserNum: 1, + Gname: req.Gname, + CreateUid: subject.Uid, + CreateAt: time.Now().UnixMilli(), + } + + // join postal group + err = s.deliverGroup.GroupJoin(ctx, subject.Uid, []string{group.Gid}) + if err != nil { + return + } + + s.rooms.Store(group.Gid, group) + res = &chat.ResRoomCreate{Group: group} + return } -func (s *ChatServer) RoomJoin(ctx context.Context, join *chat.ReqRoomJoin) (*emptypb.Empty, error) { - return nil, nil +func (s *ChatServer) RoomJoin(ctx context.Context, req *chat.ReqRoomJoin) (_ *emptypb.Empty, err error) { + subject, err := session.GetSubject(ctx) + if err != nil { + return nil, err + } + + room, ok := s.rooms.Load(req.Gid) + if !ok { + err = fmt.Errorf("room %s not exists", req.Gid) + return + } + atomic.AddInt32(&room.UserNum, 1) + + err = s.deliverGroup.GroupJoin(ctx, subject.Uid, []string{req.Gid}) + return } -func (s *ChatServer) RoomLeave(ctx context.Context, leave *chat.ReqRoomLeave) (*emptypb.Empty, error) { - return nil, nil +func (s *ChatServer) RoomLeave(ctx context.Context, req *chat.ReqRoomLeave) (_ *emptypb.Empty, err error) { + subject, err := session.GetSubject(ctx) + if err != nil { + return nil, err + } + + room, ok := s.rooms.Load(req.Gid) + if !ok { + err = fmt.Errorf("room %s not exists", req.Gid) + return + } + atomic.AddInt32(&room.UserNum, -1) + + err = s.deliverGroup.GroupLeave(ctx, subject.Uid, []string{req.Gid}) + return } -func (s *ChatServer) RoomKickOut(ctx context.Context, out *chat.ReqRoomKickOut) (*emptypb.Empty, error) { - return nil, nil +func (s *ChatServer) RoomKickOut(ctx context.Context, req *chat.ReqRoomKickOut) (_ *emptypb.Empty, err error) { + return } -func (s *ChatServer) RoomInfo(ctx context.Context, info *chat.ReqRoomInfo) (*chat.ResRoomInfo, error) { - return nil, nil +func (s *ChatServer) RoomInfo(ctx context.Context, req *chat.ReqRoomInfo) (res *chat.ResRoomInfo, err error) { + room, ok := s.rooms.Load(req.Gid) + if !ok { + err = fmt.Errorf("room %s not exists", req.Gid) + return + } + + res = &chat.ResRoomInfo{Group: room} + return } -func (s *ChatServer) RoomList(ctx context.Context, list *chat.ReqRoomList) (*chat.ResRoomList, error) { - return nil, nil +func (s *ChatServer) RoomList(ctx context.Context, req *chat.ReqRoomList) (res *chat.ResRoomList, err error) { + rooms := make([]*chat.Group, 16) + s.rooms.Range(func(_ string, room *chat.Group) bool { + rooms = append(rooms, room) + return true + }) + res = &chat.ResRoomList{Groups: rooms} + return } diff --git a/internal/gateway_http/logic/http_server.go b/internal/gateway_http/logic/http_server.go index 035b294..bc9e027 100644 --- a/internal/gateway_http/logic/http_server.go +++ b/internal/gateway_http/logic/http_server.go @@ -9,6 +9,7 @@ import ( "sonet/pkg/protocol/session" "sonet/pkg/utils/resp" "sonet/pkg/utils/strs" + "time" ) type GrpcGenericHandler struct { @@ -39,11 +40,14 @@ func (h *GrpcGenericHandler) handler(c *gin.Context) { return } - grpcClient, err := h.grpcFactory.GetClient(ctx, svc) + ctx2, cancel := context.WithTimeout(ctx, time.Second*3) + defer cancel() + grpcClient, err := h.grpcFactory.GetClient(ctx2, svc) if err != nil { c.JSON(http.StatusForbidden, resp.Error(err.Error())) return } + body := make(map[string]interface{}) err = c.BindJSON(&body) if err != nil { @@ -51,7 +55,9 @@ func (h *GrpcGenericHandler) handler(c *gin.Context) { return } - res, err := grpcClient.InvokeUnaryJson(ctx, method, body) + ctx3, cancel := context.WithTimeout(ctx, time.Second*10) + defer cancel() + res, err := grpcClient.InvokeUnaryJson(ctx3, method, body) if err != nil { c.JSON(http.StatusInternalServerError, resp.Error(err.Error())) return diff --git a/internal/gateway_http/logic/postal_monitor.go b/internal/gateway_http/logic/postal_monitor.go deleted file mode 100644 index c70323b..0000000 --- a/internal/gateway_http/logic/postal_monitor.go +++ /dev/null @@ -1,80 +0,0 @@ -package logic - -import ( - "context" - "go.etcd.io/etcd/api/v3/mvccpb" - clientv3 "go.etcd.io/etcd/client/v3" - "sonet/api/gen/postal" - "sonet/pkg/grpc/discovery" - "sonet/pkg/utils/logger" -) - -type PostalMonitor struct { - client *clientv3.Client - keyPrefix string -} - -func NewPostalMonitor(client *clientv3.Client) *PostalMonitor { - return &PostalMonitor{ - client: client, - } -} - -func (m *PostalMonitor) Init(ctx context.Context) (err error) { - m.keyPrefix = discovery.BuildPrefix(discovery.Server{Name: postal.Postal_ServiceDesc.ServiceName}) - go m.watch(ctx) - m.build(ctx) - return -} - -func (m *PostalMonitor) Next() { - -} - -func (m *PostalMonitor) build(ctx context.Context) { - res, err := m.client.Get(ctx, m.keyPrefix, clientv3.WithPrefix()) - if err != nil { - return - } - //for _, kv := range res.Kvs { - // - //} - m.update(res.Kvs) - -} - -func (m *PostalMonitor) watch(ctx context.Context) { - w := m.client.Watch(ctx, m.keyPrefix, clientv3.WithPrefix()) - cancelCh := ctx.Done() - - for { - select { - case <-cancelCh: - return - case res := <-w: - if err := res.Err(); err != nil { - logger.Errorf("watch etcd instance error: %v\n", err) - continue - } - rebuild := false - eLoop: - for _, event := range res.Events { - switch event.Type { - case clientv3.EventTypePut: - fallthrough - case clientv3.EventTypeDelete: - rebuild = true - break eLoop - } - } - if rebuild { - go m.build(ctx) - } - } - } - -} - -func (m *PostalMonitor) update(value []*mvccpb.KeyValue) { - -} diff --git a/internal/gateway_ws/gws_server/conn_handler.go b/internal/gateway_ws/gws_server/conn_handler.go index fa61961..3b4a941 100644 --- a/internal/gateway_ws/gws_server/conn_handler.go +++ b/internal/gateway_ws/gws_server/conn_handler.go @@ -110,7 +110,9 @@ func (c *GwsHandler) OnMessage(socket *gws.Conn, message *gws.Message) { // grpc generic call ctx := context.Background() - grpcClient, err := c.grpcFactory.GetClient(ctx, service) + ctx2, cancel2 := context.WithTimeout(ctx, time.Second*3) + defer cancel2() + grpcClient, err := c.grpcFactory.GetClient(ctx2, service) if err != nil { logger.Error("get grpc generic client error: ", err) writeError(conn, header.SeqId, fmt.Sprintf("service %s not avaliable: %s", header.Svc, err.Error())) @@ -120,7 +122,10 @@ func (c *GwsHandler) OnMessage(socket *gws.Conn, message *gws.Message) { if authorized { ctx = session.PutSubject(ctx, session.NewRpcSubject(sessionUid)) } - resp, err := grpcClient.InvokeUnary(ctx, header.Target, payload.Body) + + ctx3, cancel3 := context.WithTimeout(ctx, time.Second*10) + defer cancel3() + resp, err := grpcClient.InvokeUnary(ctx3, header.Target, payload.Body) if err != nil { writeError(conn, header.SeqId, unwrapRpcError(err.Error())) logger.Error("grpc generic call error: ", err) diff --git a/internal/gateway_ws/server/ws_server.go b/internal/gateway_ws/server/ws_server.go index f5e51ac..f884f4f 100644 --- a/internal/gateway_ws/server/ws_server.go +++ b/internal/gateway_ws/server/ws_server.go @@ -97,7 +97,9 @@ func (c *ConnHandler) handleConn(conn *websocket.Conn) { // grpc generic call ctx := context.Background() - grpcClient, err := c.grpcFactory.GetClient(ctx, service) + ctx2, cancel2 := context.WithTimeout(ctx, time.Second*3) + grpcClient, err := c.grpcFactory.GetClient(ctx2, service) + cancel2() if err != nil { logger.Error("get grpc generic client error: ", err) writeError(wsConn, header.SeqId, fmt.Sprintf("service %s not avaliable: %s", header.Svc, err.Error())) @@ -107,7 +109,9 @@ func (c *ConnHandler) handleConn(conn *websocket.Conn) { if sessionUid != "" { ctx = session.PutSubject(ctx, session.NewRpcSubject(sessionUid)) } - resp, err := grpcClient.InvokeUnary(ctx, header.Target, payload.Body) + ctx3, cancel3 := context.WithTimeout(ctx, time.Second*10) + resp, err := grpcClient.InvokeUnary(ctx3, header.Target, payload.Body) + cancel3() if err != nil { writeError(wsConn, header.SeqId, unwrapRpcError(err.Error())) logger.Error("grpc generic call error: ", err) @@ -172,16 +176,16 @@ func (c *ConnHandler) extractAuthVerify(conn *wsConn, resp []byte) (sessionUid s } // 关闭旧的链接 - oldConn, ok := c.sessionStore.Load(sessionUid) + oldChannel, ok := c.sessionStore.Load(sessionUid) if ok { // TODO nats offline / force load subject target cluster call offline - if err = oldConn.(*wsConn).conn.Close(); err != nil { + if err = oldChannel.Conn.(*wsConn).conn.Close(); err != nil { logger.Error("old conn close error: ", err) } logger.Infof("close subject old conn: %s\n", sessionUid) } // 存储session到内存 - c.sessionStore.Store(sessionUid, conn) + c.sessionStore.Store(sessionUid, session.NewChannel(sessionUid, conn)) logger.Infof("subject online: %s\n", sessionUid) return } diff --git a/internal/mahjong/logic/mahjong_server.go b/internal/mahjong/logic/mahjong_server.go index 6278cca..f9e299f 100644 --- a/internal/mahjong/logic/mahjong_server.go +++ b/internal/mahjong/logic/mahjong_server.go @@ -26,15 +26,18 @@ import ( type MahjongServer struct { mahjong.UnimplementedMahjongServer store *store.Store - deliver *deliver.Deliver authClient auth.AuthClient + deliver *deliver.Deliver } -func NewMahjongServer(store *store.Store, deliver *deliver.Deliver, authClient auth.AuthClient) *MahjongServer { +func NewMahjongServer( + store *store.Store, + authClient auth.AuthClient, + deliver *deliver.Deliver) *MahjongServer { return &MahjongServer{ store: store, - deliver: deliver, authClient: authClient, + deliver: deliver, } } diff --git a/internal/postal/logic/postal_server.go b/internal/postal/logic/postal_server.go index c295fd2..151999a 100644 --- a/internal/postal/logic/postal_server.go +++ b/internal/postal/logic/postal_server.go @@ -16,6 +16,7 @@ import ( "sonet/pkg/plugins/cache" "sonet/pkg/plugins/mq" "sonet/pkg/protocol" + "sonet/pkg/protocol/event" "sonet/pkg/protocol/session" "sonet/pkg/utils/collect" "sonet/pkg/utils/logger" @@ -101,8 +102,16 @@ func (s *PostalServer) processSessionStoreEvent(ctx context.Context) { case uid := <-s.sessionStore.OnStore(): logger.Info("uid %s online", uid) - // todo mq event..., delay 5s - // s.producer.Publish() + // publish mq online event todo delay 5s + online := &event.Online{Uid: uid} + data, err := online.MarshalBinary() + if err != nil { + logger.Errorf("marshal %s online event error: %v", uid, err) + continue + } + if err = s.producer.Publish(online.Topic(), data); err != nil { + logger.Errorf("publish %s online event error: %v", uid, err) + } case channel := <-s.sessionStore.OnDelete(): logger.Info("uid %s offline", channel.Uid) @@ -111,7 +120,16 @@ func (s *PostalServer) processSessionStoreEvent(ctx context.Context) { g.Leave(channel.Uid) } } - // todo mq event... + // publish mq offline event todo delay 5s + offline := event.Offline{Uid: channel.Uid} + data, err := offline.MarshalBinary() + if err != nil { + logger.Errorf("marshal %s offline event error: %v", channel.Uid, err) + continue + } + if err = s.producer.Publish(offline.Topic(), data); err != nil { + logger.Errorf("publish %s offline event error: %v", channel.Uid, err) + } } } } diff --git a/pkg/grpc/generic/generic_client.go b/pkg/grpc/generic/generic_client.go index 902ff87..a8640d9 100644 --- a/pkg/grpc/generic/generic_client.go +++ b/pkg/grpc/generic/generic_client.go @@ -117,8 +117,8 @@ func (c *GrpcGenericClient) invokeUnary0(ctx context.Context, method string, req if err != nil { return } - ctx, cancel := context.WithTimeout(ctx, 10*time.Second) - defer cancel() + // ctx, cancel := context.WithTimeout(ctx, 10*time.Second) + // defer cancel() res, err := caller.Stub.InvokeRpc(ctx, caller.Mtd, request, opts...) if err != nil { return diff --git a/pkg/plugins/mq/nats_jet_stream.go b/pkg/plugins/mq/nats_jet_stream.go index 16fab60..bf5bab6 100644 --- a/pkg/plugins/mq/nats_jet_stream.go +++ b/pkg/plugins/mq/nats_jet_stream.go @@ -144,6 +144,7 @@ func (mq *NatsJetStreamConsumer) Subscribe(filterSubject string, channel string, } consumeContext, err := consumer.Consume(func(msg jetstream.Msg) { + // todo recover error message := &Message{ Body: msg.Data(), Time: time.Now().UnixMilli(), diff --git a/pkg/protocol/deliver/deliver_group.go b/pkg/protocol/deliver/deliver_group.go index e8d447c..dd0b4f5 100644 --- a/pkg/protocol/deliver/deliver_group.go +++ b/pkg/protocol/deliver/deliver_group.go @@ -9,35 +9,46 @@ import ( "sonet/api/gen/postal" "sonet/pkg/grpc/balancer" "sonet/pkg/grpc/discovery" + "sonet/pkg/plugins/mq" + "sonet/pkg/protocol/event" "sonet/pkg/utils/logger" "sync" + "sync/atomic" ) type GroupLoader interface { - Load(uid string) (groupIds []string) + Load(uid string) (groupIds []string, err error) } type GroupDeliver struct { - svcName string - groupLoader GroupLoader - resolver discovery.Resolver - postal postal.PostalClient // postal consistent hash client - postalDialOptions []grpc.DialOption - postals map[string]*Postal - lock *sync.RWMutex + svcName string + groupLoader GroupLoader + consumer mq.Consumer + resolver discovery.Resolver + postal postal.PostalClient // postal consistent hash client + directPostalDialOptions []grpc.DialOption + directPostals *atomic.Value //map[string]*Postal , postal server 直连客户端 + lock *sync.RWMutex } func NewGroupDeliver( msgInServiceName string, groupLoader GroupLoader, + consumer mq.Consumer, resolver discovery.Resolver, - postalDialOptions []grpc.DialOption, + directPostalDialOptions []grpc.DialOption, ) *GroupDeliver { + directPostals := &atomic.Value{} + directPostals.Store(make(map[string]*Postal, 3)) + return &GroupDeliver{ - svcName: msgInServiceName, - groupLoader: groupLoader, - resolver: resolver, - postalDialOptions: postalDialOptions, + svcName: msgInServiceName, + groupLoader: groupLoader, + consumer: consumer, + resolver: resolver, + directPostalDialOptions: directPostalDialOptions, + directPostals: directPostals, + lock: &sync.RWMutex{}, } } @@ -59,9 +70,39 @@ func (d *GroupDeliver) Init(ctx context.Context, grpcResolver discovery.GrpcReso if err != nil { return } + + err = d.subscribe() return } +// toPostalGid 加上 svc name 前缀避免和其他服务群组冲突 +func (d *GroupDeliver) toPostalGid(gid string) string { + return d.svcName + "." + gid +} + +// subscribe postal online events +func (d *GroupDeliver) subscribe() error { + consumerChannel := fmt.Sprintf("%s:%s", "deliver", d.svcName) + return d.consumer.Subscribe(event.TopicOnline, consumerChannel, func(message *mq.Message) (err error) { + online := &event.Online{} + if e := online.UnmarshalBinary(message.Body); e != nil { + logger.Error("deliver unmarshal online event payload error: ", e) + return + } + // load uid groups join to postal + groupIds, err := d.groupLoader.Load(online.Uid) + if err != nil { + logger.Error("deliver load groups error:", err) + return + } + if err = d.GroupJoin(context.Background(), online.Uid, groupIds); err != nil { + logger.Errorf("deliver uid %s group join error: %v", online.Uid, err) + return + } + return + }) +} + func (d *GroupDeliver) initPostal(ctx context.Context, resolver discovery.GrpcResolver, opts ...grpc.DialOption) (err error) { balancer.InitConsistentHashBuilder() rb, err := resolver.Resolver() @@ -105,13 +146,14 @@ func (d *GroupDeliver) watchPostal(ctx context.Context) (err error) { func (d *GroupDeliver) buildPostals(servers []discovery.Server) { postals := make(map[string]*Postal) + directPostals := d.directPostals.Load().(map[string]*Postal) for _, server := range servers { - if p, ok := d.postals[server.Addr]; ok { + if p, ok := directPostals[server.Addr]; ok { postals[server.Addr] = p continue } // new connection - conn, err := grpc.DialContext(context.Background(), server.Addr, d.postalDialOptions...) + conn, err := grpc.DialContext(context.Background(), server.Addr, d.directPostalDialOptions...) if err != nil { logger.Errorf("dial postal server %+v error: %v", server.Addr, err) continue @@ -120,39 +162,49 @@ func (d *GroupDeliver) buildPostals(servers []discovery.Server) { } // close old connection - oldPostals := d.postals - d.postals = postals + oldPostals := directPostals + d.directPostals.Store(postals) for addr, p := range oldPostals { - if _, ok := d.postals[addr]; !ok { + if _, ok := postals[addr]; !ok { p.Close() } } } func (d *GroupDeliver) DeliverGroup(ctx context.Context, gid string, msg proto.Message, options ...Option) (err error) { + gid = d.toPostalGid(gid) + // deliver to all postal message, err := protoMessage2Deliver(d.svcName, msg) if err != nil { return } req := postal.ReqDeliverGroup{Gid: gid, Msg: message} - for _, p := range d.postals { + for _, p := range d.directPostals.Load().(map[string]*Postal) { p.DeliverGroup(&req) } return } func (d *GroupDeliver) GroupDissolve(gid string) { + gid = d.toPostalGid(gid) + req := postal.ReqGroupDissolve{Gid: gid} - for _, p := range d.postals { + for _, p := range d.directPostals.Load().(map[string]*Postal) { p.GroupDissolve(&req) } } func (d *GroupDeliver) GroupJoin(ctx context.Context, uid string, gids []string) (err error) { - req := &postal.ReqGroupJoin{Uid: uid, Gids: gids} + if len(gids) == 0 { + return + } + for i := 0; i < len(gids); i++ { + gids[i] = d.toPostalGid(gids[i]) + } ctx = context.WithValue(ctx, balancer.ConsistentHashKey, uid) + req := &postal.ReqGroupJoin{Uid: uid, Gids: gids} _, err = d.postal.GroupJoin(ctx, req) return } diff --git a/pkg/protocol/event/event.go b/pkg/protocol/event/event.go index 3fc7258..a7d8fb4 100644 --- a/pkg/protocol/event/event.go +++ b/pkg/protocol/event/event.go @@ -1,22 +1,52 @@ package event -type Event[T any] interface { +import ( + "encoding" + "github.com/bytedance/sonic" +) + +const ( + TopicOnline = "postal:event:online" + TopicOffline = "postal:event:offline" +) + +type Event interface { Topic() string - Payload2Bytes(T) []byte - Bytes2Payload([]byte) T + encoding.BinaryMarshaler + encoding.BinaryUnmarshaler +} + +// Online ==================================================== Online +type Online struct { + Uid string +} + +func (o *Online) Topic() string { + return TopicOnline +} +func (o *Online) MarshalBinary() (data []byte, err error) { + data, err = sonic.Marshal(o) + return +} + +func (o *Online) UnmarshalBinary(data []byte) error { + return sonic.Unmarshal(data, o) } -type online struct { +// Offline ==================================================== Offline +type Offline struct { + Uid string } -func (o *online) Topic() string { - return "postal_event_online" +func (o *Offline) Topic() string { + return TopicOffline } -func (o *online) Payload2Bytes(t string) []byte { - return []byte(t) +func (o *Offline) MarshalBinary() (data []byte, err error) { + data, err = sonic.Marshal(o) + return } -func (o *online) Bytes2Payload(bytes []byte) string { - return string(bytes) +func (o *Offline) UnmarshalBinary(data []byte) error { + return sonic.Unmarshal(data, o) }