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.
 
 

133 lines
2.9 KiB

package deliver
import (
"context"
"google.golang.org/grpc"
"google.golang.org/protobuf/proto"
"sonet/api/gen/postal"
"sonet/pkg/grpc/discovery"
"sonet/pkg/utils/logger"
"sync"
)
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
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) (err error) {
// initial all postal 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
}
// 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] = &Postal{
conn: conn,
Client: postal.NewPostalClient(conn),
}
}
// close old connection
oldPostals := d.postals
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)
}
}
}
}
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 {
res, err := p.Client.DeliverGroup(ctx, req)
if err != nil || !res.Ok {
logger.Errorf("deliver to postal failed: %+v, %v", res, err)
}
}
return
}
func (d *GroupDeliver) GroupJoin(ctx context.Context, uid string, gid []string) (err error) {
return
}