package logic import ( "context" "errors" "google.golang.org/grpc" "google.golang.org/grpc/reflection" "google.golang.org/protobuf/types/known/emptypb" "net" "sonet/api/gen/postal" "sonet/internal/postal/group" "sonet/pkg/config" "sonet/pkg/grpc/discovery" "sonet/pkg/grpc/interceptor" "sonet/pkg/plugins/mq" "sonet/pkg/protocol/event" "sonet/pkg/protocol/session" "sonet/pkg/utils/collect" "sonet/pkg/utils/logger" "sonet/pkg/utils/shutdown" ) type PostalServer struct { postal.UnimplementedPostalServer endpointAddress string // websocket 前端连接地址 broadcastAddress string // postal server 在注册中心注册的地址, 其他服务可直连访问 sessionStore session.Store // k:uid 在线用户conn存储 groupStore *collect.ConcurrentMap[string, *group.Group] // k:gid 在线用户 groups 存储 postalClusterServer *PostalClusterServer producer mq.Producer } func NewPostalServer( endpointAddress string, sessionStore session.Store, groupStore *collect.ConcurrentMap[string, *group.Group], postalClusterServer *PostalClusterServer, producer mq.Producer) *PostalServer { return &PostalServer{ endpointAddress: endpointAddress, sessionStore: sessionStore, groupStore: groupStore, postalClusterServer: postalClusterServer, producer: producer, } } func (s *PostalServer) Run(conf config.GrpcConfig, registry discovery.Registry) (err error) { server := grpc.NewServer( config.GetGrpcOptions( conf, grpc.UnaryInterceptor(interceptor.RecoverInterceptor), )..., ) if !conf.NoReflection { // 注册反射服务 reflection.Register(server) } postal.RegisterPostalServer(server, s) listen, err := net.Listen("tcp", conf.Address) if err != nil { return } // registry discovery register := conf.Register if register.Name == "" { register.Name = postal.Postal_ServiceDesc.ServiceName } if register.Addr == "" { register.Addr = conf.Address } ctx, cancel := context.WithCancel(context.Background()) err = registry.Registry(ctx, register) if err != nil { panic(err) } shutdown.AddHook(cancel) go s.processSessionStoreEvent(ctx) // 其他服务直连地址 s.broadcastAddress = register.Addr // run serve logger.Infof("%s grpc server running %s\n", register.Name, listen.Addr().String()) err = server.Serve(listen) return } func (s *PostalServer) processSessionStoreEvent(ctx context.Context) { for { select { case <-ctx.Done(): return case uid := <-s.sessionStore.OnStore(): logger.Info("uid %s online", uid) // 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) for _, gid := range channel.Groups() { if g, ok := s.groupStore.Load(gid); ok { g.Leave(channel.Uid) } } // 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) } } } } 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 } // write msg err = receiver.Conn.Write(bytes) if err != nil { return } res = &postal.ResDeliver{Ok: true} return } // 向一致性 hash 下一个节点传递 // picker reqCluster offset +3: for _, offset := range []int32{3, -3} { res, err = s.postalClusterServer.Redeliver(ctx, &postal.ReqRedeliver{ Receivers: []string{req.Receiver}, Msg: req.Msg, Offset: offset, }) if err == nil { res = &postal.ResDeliver{Ok: true} return } } return } func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverBatch) (res *postal.ResDeliver, err error) { bytes, err := encodeDeliverMessage(req.Msg) if err != nil { return } var redeliverReceivers []string for _, receiverId := range req.Receivers { receiver, ok := s.sessionStore.Load(receiverId) if !ok { redeliverReceivers = append(redeliverReceivers, receiverId) continue } err = receiver.Conn.Write(bytes) if err != nil { logger.Errorf("deliver to %s error: ", receiver, err) } } if len(redeliverReceivers) == 0 { return &postal.ResDeliver{Ok: true}, nil } // 向一致性 hash 下一个节点传递 for _, offset := range []int32{3, -3} { res, err = s.postalClusterServer.Redeliver(ctx, &postal.ReqRedeliver{ Receivers: redeliverReceivers, Msg: req.Msg, Offset: offset, }) if err == nil { res = &postal.ResDeliver{Ok: true} return } } return } // 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 } res = &postal.ResDeliver{Ok: true} g, ok := s.groupStore.Load(req.Gid) if !ok { return } g.Write(bytes) return } // GroupJoin uid consistent hash func (s *PostalServer) GroupJoin(ctx context.Context, req *postal.ReqGroupJoin) (emp *emptypb.Empty, err error) { channel, ok := s.sessionStore.Load(req.Uid) if !ok { err = errors.New("uid not online") return } 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) channel.GroupJoin(gid) } return } // GroupLeave uid consistent hash func (s *PostalServer) GroupLeave(ctx context.Context, req *postal.ReqGroupLeave) (emp *emptypb.Empty, err error) { channel, ok := s.sessionStore.Load(req.Uid) if !ok { err = errors.New("uid not online") return } for _, gid := range req.Gids { g, ok := s.groupStore.Load(gid) if !ok { continue } g.Leave(req.Uid) channel.GroupLeave(gid) } return } // GroupDissolve postal broadcast func (s *PostalServer) GroupDissolve(ctx context.Context, req *postal.ReqGroupDissolve) (emp *emptypb.Empty, err error) { s.groupStore.Delete(req.Gid) return } func (s *PostalServer) Endpoint(context.Context, *emptypb.Empty) (*postal.ResEndpoint, error) { res := &postal.ResEndpoint{ Endpoint: s.endpointAddress, } return res, nil }