You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 

218 lines
5.9 KiB

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
}