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