package deliver import ( "context" "fmt" "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" "google.golang.org/protobuf/proto" "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, err error) } type GroupDeliver struct { 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, directPostalDialOptions []grpc.DialOption, ) *GroupDeliver { directPostals := &atomic.Value{} directPostals.Store(make(map[string]*Postal, 3)) return &GroupDeliver{ svcName: msgInServiceName, groupLoader: groupLoader, consumer: consumer, resolver: resolver, directPostalDialOptions: directPostalDialOptions, directPostals: directPostals, lock: &sync.RWMutex{}, } } func (d *GroupDeliver) Init(ctx context.Context, grpcResolver discovery.GrpcResolver, opts ...grpc.DialOption) (err error) { err = d.initPostal(ctx, grpcResolver, opts...) if err != nil { return } // initial all postal direct clients servers, err := d.resolver.ResolveAll(ctx, postal.Postal_ServiceDesc.ServiceName) if err != nil { return } d.buildPostals(servers) // watch postal server instance err = d.watchPostal(ctx) 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() if err != nil { return } postalUrl := resolver.DialUrl(postal.Postal_ServiceDesc.ServiceName) var options []grpc.DialOption // consistent hash lb options = append(options, grpc.WithDefaultServiceConfig(fmt.Sprintf(`{"loadBalancingPolicy":"%s"}`, balancer.ConsistentHash))) options = append(options, grpc.WithResolvers(rb)) options = append(options, grpc.WithTransportCredentials(insecure.NewCredentials())) options = append(options, opts...) conn, err := grpc.DialContext(ctx, postalUrl, options...) if err != nil { return } d.postal = postal.NewPostalClient(conn) return } // watchPostal watch postal service list func (d *GroupDeliver) watchPostal(ctx context.Context) (err error) { ch, err := d.resolver.Watch(ctx, postal.Postal_ServiceDesc.ServiceName) if err != nil { return } go func() { for { select { case <-ctx.Done(): return case servers := <-ch: d.buildPostals(servers) } } }() return } 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 := directPostals[server.Addr]; ok { postals[server.Addr] = p continue } // new connection 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 } postals[server.Addr] = NewPostal(conn) } // close old connection oldPostals := directPostals d.directPostals.Store(postals) for addr, p := range oldPostals { 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.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.directPostals.Load().(map[string]*Postal) { p.GroupDissolve(&req) } } func (d *GroupDeliver) GroupJoin(ctx context.Context, uid string, gids []string) (err error) { 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 } func (d *GroupDeliver) GroupLeave(ctx context.Context, uid string, gids []string) (err error) { req := &postal.ReqGroupLeave{Uid: uid, Gids: gids} ctx = context.WithValue(ctx, balancer.ConsistentHashKey, uid) _, err = d.postal.GroupLeave(ctx, req) return }