package logic import ( "context" "google.golang.org/grpc" "google.golang.org/grpc/reflection" "google.golang.org/protobuf/types/known/emptypb" "net" "sonet/api/gen/postal" "sonet/internal/gateway_ws/session" "sonet/pkg/config" "sonet/pkg/grpc/client" "sonet/pkg/grpc/discovery" "sonet/pkg/grpc/interceptor" "sonet/pkg/plugins/cache" "sonet/pkg/protocol" "sonet/pkg/utils/logger" "sync" ) type PostalServer struct { postal.UnimplementedPostalServer endpointAddress string // websocket 前端连接地址 broadcastAddress string // postal server 在注册中心注册的地址, 其他服务可直连访问 sessionStore *sync.Map // 在线用户conn存储 subjectStore cache.MultiLevelCache // online subject address cache, ws gateway集群所有在线用户存储 clientFactory *client.GrpcDirectClientFactory } func NewPostalServer( endpointAddress string, sessionStore *sync.Map, subjectStore cache.MultiLevelCache, clientFactory *client.GrpcDirectClientFactory) *PostalServer { return &PostalServer{ endpointAddress: endpointAddress, sessionStore: sessionStore, subjectStore: subjectStore, clientFactory: clientFactory, } } func (s *PostalServer) Run(conf config.GrpcConfig, postalRegister *discovery.Register) (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 reg := conf.Register if reg.Name == "" { reg.Name = postal.Postal_ServiceDesc.ServiceName } if reg.Addr == "" { reg.Addr, err = discovery.RegisterAddress(conf.Address) if err != nil { return } } if err = postalRegister.Register(reg); err != nil { return } // 其他服务直连地址 s.broadcastAddress = reg.Addr // run serve logger.Infof("%s grpc server running %s\n", reg.Name, listen.Addr().String()) err = server.Serve(listen) return } func (s *PostalServer) deliverMessage(msg *postal.Message, netSubject *session.NetSubject) 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("req deliver encode notice error: ", err) return err } return netSubject.Client.Write(bytes) } // receiverGates 找receiver在集群内哪些其他节点 func (s *PostalServer) clusterGateReceivers(ctx context.Context, receivers []string) (gateReceivers map[string][]string, offline []string) { gateReceivers = make(map[string][]string) for _, receiver := range receivers { subject := &session.Subject{} err := s.subjectStore.Load(ctx, receiver, subject) if err != nil { offline = append(offline, receiver) if err != cache.NotExists { logger.Error("load from subject store error: ", err) offline = append(offline, receiver) } continue } if subject.Gate == s.broadcastAddress { offline = append(offline, receiver) // 清除失效缓存 err := s.subjectStore.Del(ctx, receiver) if err != nil { logger.Error("del subject store error: ", receiver, err) } continue } // put receiver gate addr gateReceivers[subject.Gate] = append(gateReceivers[subject.Gate], receiver) } return } func (s *PostalServer) Deliver(ctx context.Context, req *postal.ReqDeliver) (*postal.ResDeliver, error) { val, ok := s.sessionStore.Load(req.Receiver) if ok { receiver := val.(*session.NetSubject) err := s.deliverMessage(req.Msg, receiver) if err != nil { return nil, err } return &postal.ResDeliver{Ok: true}, nil } // 用户连接不在当前gateway gateReceivers, offline := s.clusterGateReceivers(ctx, []string{req.Receiver}) if offline != nil && len(offline) > 0 { return &postal.ResDeliver{Code: int32(postal.DeliverResult_ReceiverOffline.Number())}, nil } for postalAddr := range gateReceivers { // gateway集群中转发消息 clusterAddr, err := PostalAddr2Cluster(postalAddr) if err != nil { logger.Errorf("parse postal server addr error: %s", postalAddr, err) continue } conn, err := s.clientFactory.GetConn(context.Background(), clusterAddr) if err != nil { logger.Error("get postal cluster conn error: ", err) return nil, err } clusterClient := postal.NewPostalClusterClient(conn) reqRedirect := &postal.ReqRedirect{ Ttl: 3, // TODO 转发n次就丢弃 RedirectMethod: methodRedirectDeliver, Deliver: req, } resp, err := clusterClient.Redirect(ctx, reqRedirect) return resp, err } return nil, nil } func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverBatch) (*postal.ResDeliver, error) { var redirectReceivers []string for _, receiverId := range req.Receivers { val, ok := s.sessionStore.Load(receiverId) if !ok { redirectReceivers = append(redirectReceivers, receiverId) continue } receiver := val.(*session.NetSubject) err := s.deliverMessage(req.Msg, receiver) if err != nil { logger.Errorf("deliver to %s error: ", receiver, err) } } if len(redirectReceivers) == 0 { return &postal.ResDeliver{Ok: true}, nil } gateReceivers, offline := s.clusterGateReceivers(ctx, redirectReceivers) if len(offline) > 0 { logger.Warning("offline redirect receivers: ", offline) } if len(gateReceivers) > 0 { for postalAddr, receivers := range gateReceivers { // gateway集群中转发消息 clusterAddr, err := PostalAddr2Cluster(postalAddr) if err != nil { logger.Errorf("parse postal server addr error: %s", postalAddr, err) continue } conn, err := s.clientFactory.GetConn(ctx, clusterAddr) if err != nil { logger.Error("get postal cluster conn error: ", err) continue } clusterClient := postal.NewPostalClusterClient(conn) req.Receivers = receivers reqRedirect := &postal.ReqRedirect{ Ttl: 3, // TODO 转发n次就丢弃 RedirectMethod: methodRedirectDeliverBatch, DeliverBatch: req, } _, err = clusterClient.Redirect(ctx, reqRedirect) if err != nil { logger.Error("redirect batch error: ", err) } } } return &postal.ResDeliver{Ok: true}, nil } func (s *PostalServer) DeliverGroup(ctx context.Context, group *postal.ReqDeliverGroup) (*postal.ResDeliver, error) { return nil, nil } func (s *PostalServer) GroupCreate(ctx context.Context, create *postal.ReqGroupCreate) (*postal.ResGroupCreate, error) { return nil, nil } func (s *PostalServer) GroupJoin(ctx context.Context, join *postal.ReqGroupJoin) (*emptypb.Empty, error) { return nil, nil } func (s *PostalServer) GroupLeave(ctx context.Context, leave *postal.ReqGroupLeave) (*emptypb.Empty, error) { return nil, nil } func (s *PostalServer) GroupDissolve(ctx context.Context, dissolve *postal.ReqGroupDissolve) (*emptypb.Empty, error) { return nil, nil } func (s *PostalServer) Endpoint(context.Context, *emptypb.Empty) (*postal.ResEndpoint, error) { res := &postal.ResEndpoint{ Endpoint: s.endpointAddress, } return res, nil }