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
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 |
|
}
|
|
|