package logic import ( "context" "fmt" "google.golang.org/grpc" "net" "sonet/api/gen/postal" "sonet/pkg/grpc/client" "sonet/pkg/protocol/deliver" "sonet/pkg/protocol/session" "sonet/pkg/utils/logger" "sonet/pkg/utils/nets" ) // PostalClusterPortOffset 相对与 postal grpc server 接口偏移量 const PostalClusterPortOffset = 1000 // PostalAddr2Cluster 根据 postal server 地址得到 postal cluster server 地址 func PostalAddr2Cluster(postalServerAddr string) (string, error) { addr, err := net.ResolveTCPAddr("tcp", postalServerAddr) if err != nil { return "", err } if addr.IP == nil { return fmt.Sprintf(":%d", addr.Port+PostalClusterPortOffset), nil } return fmt.Sprintf("%s:%d", addr.IP.String(), addr.Port+PostalClusterPortOffset), nil } // PostalClusterServer 用于postal server集群间, 消息转发通信 type PostalClusterServer struct { postal.UnimplementedPostalClusterServer sessionStore session.Store postalPicker *deliver.PostalPicker clientFactory *client.GrpcDirectClientFactory // todo 清理过期连接 } func NewPostalClusterServer( sessionStore session.Store, postalPicker *deliver.PostalPicker, clientFactory *client.GrpcDirectClientFactory, ) *PostalClusterServer { return &PostalClusterServer{ sessionStore: sessionStore, postalPicker: postalPicker, clientFactory: clientFactory, } } // Run 运行在同进程 postalServer 端口+1000 func (s *PostalClusterServer) Run(postalAddr string, opts ...grpc.ServerOption) (err error) { server := grpc.NewServer(opts...) postal.RegisterPostalClusterServer(server, s) // 解析端口 address, err := PostalAddr2Cluster(postalAddr) if err != nil { return } listen, err := net.Listen("tcp", address) if err != nil { return } // run serve logger.Infof("%s grpc server running %s\n", postal.PostalCluster_ServiceDesc.ServiceName, listen.Addr().String()) err = server.Serve(listen) return } func (s *PostalClusterServer) Redeliver(ctx context.Context, req *postal.ReqRedeliver) (res *postal.ResDeliver, err error) { bytes, err := encodeDeliverMessage(req.Msg) if err != nil { return } var redelivers []string for _, receiver := range req.Receivers { channel, ok := s.sessionStore.Load(receiver) if !ok { redelivers = append(redelivers, receiver) continue } err = channel.Conn.Write(bytes) if err != nil { logger.Errorf("redeliver channel %s write error: %v", receiver, err) continue } } if len(redelivers) == 0 { res = &postal.ResDeliver{Ok: true} return } req.Receivers = redelivers // 重新投递 var next int next, req.Offset = convergeOffset(req.Offset) if next == 0 || req.Offset == 0 { // 重投次数结束 err = errRedeliverTtlEnd return } node, _, err := s.postalPicker.PickOffsetNode(redelivers[0], next) if err != nil { return } addr := node.GetAddr() addrI64, err := nets.Address2i64(addr) if err != nil { return } for _, nodeI64 := range req.Nodes { if addrI64 == nodeI64 { // 重复投递节点, 结束投递 err = errRedeliverOffsetNodeRepeat return } } // 再次投递 req.Nodes = append(req.Nodes, addrI64) // postal集群中转发消息 clusterAddr, err := PostalAddr2Cluster(addr) if err != nil { logger.Errorf("parse postal server addr error: %s", addr, err) return } conn, err := s.clientFactory.GetConn(ctx, clusterAddr) if err != nil { logger.Error("get postal cluster conn error: ", err) return nil, err } clusterClient := postal.NewPostalClusterClient(conn) res, err = clusterClient.Redeliver(ctx, req) return } // convergeOffset 向0收敛offset func convergeOffset(offset int32) (int, int32) { if offset < 0 { return -1, offset + 1 } if offset > 0 { return 1, offset - 1 } return 0, 0 }