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 }