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.
 
 

172 lines
4.3 KiB

package deliver
import (
"context"
"errors"
"fmt"
"google.golang.org/grpc"
"sonet/api/gen/postal"
"sonet/pkg/grpc/discovery"
"sonet/pkg/utils/collect"
"sonet/pkg/utils/logger"
"sync/atomic"
)
type PostalPicker struct {
resolver discovery.Resolver
directPostalDialOptions []grpc.DialOption
directPostalNodes atomic.Value // map[string]*PostalNode , postal server 直连客户端
picker atomic.Value // *ConsistentHashPicker
}
func NewPostalPicker(resolver discovery.Resolver, directPostalDialOptions ...grpc.DialOption) *PostalPicker {
return &PostalPicker{
resolver: resolver,
directPostalDialOptions: directPostalDialOptions,
}
}
func (p *PostalPicker) Init(ctx context.Context) (err error) {
// initial all postal direct clients
servers, err := p.resolver.ResolveAll(ctx, postal.Postal_ServiceDesc.ServiceName)
if err != nil {
return
}
p.newPostalServers(servers)
// watch postal server instance
err = p.watchPostal(ctx)
if err != nil {
return
}
return
}
// watchPostal watch postal service list
func (p *PostalPicker) watchPostal(ctx context.Context) (err error) {
ch, err := p.resolver.Watch(ctx, postal.Postal_ServiceDesc.ServiceName)
if err != nil {
return
}
go func() {
for {
select {
case <-ctx.Done():
return
case servers := <-ch:
p.newPostalServers(servers)
}
}
}()
return
}
// newPostalServers postal集群节点增减
func (p *PostalPicker) newPostalServers(servers []discovery.Server) {
newPostal := make(map[string]*PostalNode)
var pickServers []discovery.Server
oldPostal, ok := p.directPostalNodes.Load().(map[string]*PostalNode)
if !ok {
oldPostal = make(map[string]*PostalNode)
}
for _, server := range servers {
if p, ok := oldPostal[server.Addr]; ok {
newPostal[server.Addr] = p
continue
}
// new connection
conn, err := grpc.DialContext(context.Background(), server.Addr, p.directPostalDialOptions...)
if err != nil {
logger.Errorf("dial postal server %+v error: %v", server, err)
continue
}
newPostal[server.Addr] = NewPostalNode(server.Addr, conn)
pickServers = append(pickServers, server)
}
// 重新构建 consistent hash picker
p.newConsistentHash(pickServers)
// close old connection
p.directPostalNodes.Store(newPostal)
for addr, p := range oldPostal {
if _, ok := newPostal[addr]; !ok {
p.Close()
}
}
}
func (p *PostalPicker) newConsistentHash(servers []discovery.Server) {
nodes := collect.Mapping(servers, func(server discovery.Server) *PickNode {
return &PickNode{
Key: server.Addr,
Weight: server.GetWeight(),
}
})
picker := NewConsistentHashPicker(nodes, DefaultReplicas, DefaultSalt)
picker.Init()
p.picker.Store(picker)
}
func (p *PostalPicker) Pick(key string) (client postal.PostalClient, err error) {
postalNode, err := p.PickNode(key)
if err != nil {
return
}
client = postalNode.client
return
}
func (p *PostalPicker) PickNode(key string) (postalNode *PostalNode, err error) {
picker := p.picker.Load().(*ConsistentHashPicker)
node, ok := picker.Pick(key)
if !ok {
err = errors.New("no available postal server")
return
}
postalNode, ok = p.directPostalNodes.Load().(map[string]*PostalNode)[node.Key]
if !ok {
err = fmt.Errorf("no conn for postal server %s", node.Key)
return
}
return
}
// PickOffset
// return same: client 是否与 offset=0 时相同
func (p *PostalPicker) PickOffset(key string, offset int) (client postal.PostalClient, same bool, err error) {
postalNode, same, err := p.PickOffsetNode(key, offset)
if err != nil {
return
}
client = postalNode.client
return
}
func (p *PostalPicker) PickOffsetNode(key string, offset int) (postalNode *PostalNode, same bool, err error) {
picker := p.picker.Load().(*ConsistentHashPicker)
node, same, ok := picker.PickOffset(key, offset)
if !ok {
err = errors.New("no available postal server")
return
}
postalNode, ok = p.directPostalNodes.Load().(map[string]*PostalNode)[node.Key]
if !ok {
err = fmt.Errorf("no conn for postal server %s", node.Key)
return
}
return
}
func (p *PostalPicker) PickAll() (clients []postal.PostalClient) {
nodes, ok := p.directPostalNodes.Load().(map[string]*PostalNode)
if !ok {
return
}
clients = make([]postal.PostalClient, 0, len(nodes))
for _, node := range nodes {
clients = append(clients, node.client)
}
return
}