diff --git a/api/postal.proto b/api/postal.proto index 644fa6f..25aed7e 100644 --- a/api/postal.proto +++ b/api/postal.proto @@ -11,6 +11,7 @@ enum DeliverResult { service Postal { rpc Deliver(ReqDeliver) returns (ResDeliver); rpc DeliverBatch(ReqDeliverBatch) returns(ResDeliver); + // broadcast request to postal cluster rpc DeliverGroup(ReqDeliverGroup) returns(ResDeliver); // rpc GroupCreate(ReqGroupCreate) returns(ResGroupCreate); @@ -67,12 +68,12 @@ message ResGroupCreate { message ReqGroupJoin { string uid = 1; - repeated string gid = 2; + repeated string gids = 2; } message ReqGroupLeave { string uid = 1; - repeated string gid = 2; + repeated string gids = 2; } message ReqGroupDissolve { diff --git a/pkg/protocol/deliver/deliver.go b/pkg/protocol/deliver/deliver.go index 09e1ab1..b72597c 100644 --- a/pkg/protocol/deliver/deliver.go +++ b/pkg/protocol/deliver/deliver.go @@ -34,7 +34,7 @@ func NewDeliver(msgInServiceName string) *Deliver { } } -func (d *Deliver) InitWithResolver(ctx context.Context, resolver discovery.GrpcResolver) (err error) { +func (d *Deliver) InitWithResolver(ctx context.Context, resolver discovery.GrpcResolver, opts ...grpc.DialOption) (err error) { balancer.InitConsistentHashBuilder() rb, err := resolver.Resolver() if err != nil { @@ -42,12 +42,13 @@ func (d *Deliver) InitWithResolver(ctx context.Context, resolver discovery.GrpcR } postalUrl := resolver.DialUrl(postal.Postal_ServiceDesc.ServiceName) - conn, err := grpc.DialContext(ctx, postalUrl, - grpc.WithTransportCredentials(insecure.NewCredentials()), - // consistent hash lb - grpc.WithDefaultServiceConfig(fmt.Sprintf(`{"loadBalancingPolicy":"%s"}`, balancer.ConsistentHash)), - grpc.WithResolvers(rb), - ) + 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 } @@ -55,9 +56,12 @@ func (d *Deliver) InitWithResolver(ctx context.Context, resolver discovery.GrpcR return } -func (d *Deliver) InitWithAddr(postalAddr string) (err error) { +func (d *Deliver) InitWithAddr(postalAddr string, opts ...grpc.DialOption) (err error) { // Conn *grpc.ClientConn - conn, err := grpc.Dial(postalAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) + var options []grpc.DialOption + options = append(options, grpc.WithTransportCredentials(insecure.NewCredentials())) + options = append(options, opts...) + conn, err := grpc.Dial(postalAddr, options...) if err != nil { return } diff --git a/pkg/protocol/deliver/group_deliver.go b/pkg/protocol/deliver/deliver_group.go similarity index 62% rename from pkg/protocol/deliver/group_deliver.go rename to pkg/protocol/deliver/deliver_group.go index 053b25b..b8d5399 100644 --- a/pkg/protocol/deliver/group_deliver.go +++ b/pkg/protocol/deliver/deliver_group.go @@ -2,9 +2,12 @@ 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" @@ -14,15 +17,11 @@ type GroupLoader interface { Load(uid string) (groupIds []string) } -type Postal struct { - conn *grpc.ClientConn - Client postal.PostalClient -} - 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 @@ -42,8 +41,13 @@ func NewGroupDeliver( } } -func (d *GroupDeliver) Init(ctx context.Context) (err error) { - // initial all postal clients +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 @@ -58,6 +62,28 @@ func (d *GroupDeliver) Init(ctx context.Context) (err error) { 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) @@ -90,10 +116,7 @@ func (d *GroupDeliver) buildPostals(servers []discovery.Server) { logger.Errorf("dial postal server %+v error: %v", server.Addr, err) continue } - postals[server.Addr] = &Postal{ - conn: conn, - Client: postal.NewPostalClient(conn), - } + postals[server.Addr] = NewPostal(conn) } // close old connection @@ -101,9 +124,7 @@ func (d *GroupDeliver) buildPostals(servers []discovery.Server) { d.postals = postals for addr, p := range oldPostals { if _, ok := d.postals[addr]; !ok { - if err := p.conn.Close(); err != nil { - logger.Errorf("close old postal conn %s error: %v", addr, err) - } + p.Close() } } } @@ -114,20 +135,21 @@ func (d *GroupDeliver) DeliverGroup(ctx context.Context, gid string, msg proto.M if err != nil { return } - req := &postal.ReqDeliverGroup{ + req := postal.ReqDeliverGroup{ Gid: gid, Msg: message, } for _, p := range d.postals { - res, err := p.Client.DeliverGroup(ctx, req) - if err != nil || !res.Ok { - logger.Errorf("deliver to postal failed: %+v, %v", res, err) - } + p.DeliverGroup(&req) } return } -func (d *GroupDeliver) GroupJoin(ctx context.Context, uid string, gid []string) (err error) { - +func (d *GroupDeliver) GroupJoin(ctx context.Context, uid string, gids []string) (err error) { + req := &postal.ReqGroupJoin{ + Uid: uid, + Gids: gids, + } + _, err = d.postal.GroupJoin(ctx, req) return } diff --git a/pkg/protocol/deliver/group_loader.go b/pkg/protocol/deliver/group_loader.go deleted file mode 100644 index 9eae15d..0000000 --- a/pkg/protocol/deliver/group_loader.go +++ /dev/null @@ -1 +0,0 @@ -package deliver diff --git a/pkg/protocol/deliver/postal.go b/pkg/protocol/deliver/postal.go new file mode 100644 index 0000000..05f4993 --- /dev/null +++ b/pkg/protocol/deliver/postal.go @@ -0,0 +1,35 @@ +package deliver + +import ( + "context" + "google.golang.org/grpc" + "sonet/api/gen/postal" + "sonet/pkg/utils/logger" +) + +type Postal struct { + conn *grpc.ClientConn + client postal.PostalClient +} + +func NewPostal(conn *grpc.ClientConn) *Postal { + return &Postal{ + conn: conn, + client: postal.NewPostalClient(conn), + } +} + +func (p *Postal) Close() { + err := p.conn.Close() + if err != nil { + logger.Error("postal close error: ", err) + } +} + +func (p *Postal) DeliverGroup(req *postal.ReqDeliverGroup) { + // todo put in chan + res, err := p.client.DeliverGroup(context.Background(), req) + if err != nil || !res.Ok { + logger.Errorf("deliver to postal failed: %+v, %v", res, err) + } +}