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/utils/logger" "sync" ) type GroupLoader interface { Load(uid string) (groupIds []string) } 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 } func NewGroupDeliver( msgInServiceName string, groupLoader GroupLoader, resolver discovery.Resolver, postalDialOptions []grpc.DialOption, ) *GroupDeliver { return &GroupDeliver{ svcName: msgInServiceName, groupLoader: groupLoader, resolver: resolver, postalDialOptions: postalDialOptions, } } 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 } 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) for _, server := range servers { if p, ok := d.postals[server.Addr]; ok { postals[server.Addr] = p continue } // new connection conn, err := grpc.DialContext(context.Background(), server.Addr, d.postalDialOptions...) if err != nil { logger.Errorf("dial postal server %+v error: %v", server.Addr, err) continue } postals[server.Addr] = NewPostal(conn) } // close old connection oldPostals := d.postals d.postals = postals for addr, p := range oldPostals { if _, ok := d.postals[addr]; !ok { p.Close() } } } func (d *GroupDeliver) DeliverGroup(ctx context.Context, gid string, msg proto.Message, options ...Option) (err error) { // 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 { p.DeliverGroup(&req) } return } func (d *GroupDeliver) GroupDissolve(gid string) { req := postal.ReqGroupDissolve{Gid: gid} for _, p := range d.postals { p.GroupDissolve(&req) } } func (d *GroupDeliver) GroupJoin(ctx context.Context, uid string, gids []string) (err error) { req := &postal.ReqGroupJoin{Uid: uid, Gids: gids} ctx = context.WithValue(ctx, balancer.ConsistentHashKey, uid) _, 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 }