package discovery import ( "context" "encoding/json" "errors" "fmt" "google.golang.org/grpc/grpclog" "net" "sonet/pkg/utils/nets" "strings" "time" clientv3 "go.etcd.io/etcd/client/v3" ) var DefaultRegisterTTL int64 = 10 func RegisterAddress(listener net.Listener) (addr string, err error) { port := listener.Addr().(*net.TCPAddr).Port ipv4, err := nets.GetHostIpv4() if err != nil { return } addr = fmt.Sprintf("%s:%d", ipv4, port) return } func MustGetRegisterAddr(listener net.Listener) (addr string) { port := listener.Addr().(*net.TCPAddr).Port ipv4, err := nets.GetHostIpv4() if err != nil { panic(err) } addr = fmt.Sprintf("%s:%s", ipv4, port) return } type Register struct { DialTimeout int closeCh chan struct{} leasesID clientv3.LeaseID keepAliveCh <-chan *clientv3.LeaseKeepAliveResponse srvInfo Server srvTTL int64 cli *clientv3.Client } // NewRegister create a register based on etcd func NewRegister(client *clientv3.Client) *Register { return &Register{ cli: client, DialTimeout: 3, } } // Register a user func (r *Register) Register(srvInfo Server) (err error) { if strings.Split(srvInfo.Addr, ":")[0] == "" { return errors.New("invalid ip address") } r.srvInfo = srvInfo r.srvTTL = DefaultRegisterTTL if err = r.register(); err != nil { return err } if r.closeCh == nil { r.closeCh = make(chan struct{}) } go r.keepAlive() return nil } func (r *Register) register() error { ctx, cancel := context.WithTimeout(context.Background(), time.Duration(r.DialTimeout)*time.Second) defer cancel() leaseResp, err := r.cli.Grant(ctx, r.srvTTL) if err != nil { return err } r.leasesID = leaseResp.ID if r.keepAliveCh, err = r.cli.KeepAlive(context.Background(), r.leasesID); err != nil { return err } data, err := json.Marshal(r.srvInfo) if err != nil { return err } _, err = r.cli.Put(context.Background(), BuildRegisterPath(r.srvInfo), string(data), clientv3.WithLease(r.leasesID)) return err } // Stop stop register func (r *Register) Stop() { if r.closeCh == nil { return } r.closeCh <- struct{}{} <-r.closeCh // 阻塞到关闭 close(r.closeCh) } // unregister 删除节点 func (r *Register) unregister() error { _, err := r.cli.Delete(context.Background(), BuildRegisterPath(r.srvInfo)) return err } func (r *Register) keepAlive() { ticker := time.NewTicker(time.Duration(r.srvTTL) * time.Second) defer ticker.Stop() for { select { case <-r.closeCh: if err := r.unregister(); err != nil { grpclog.Error("unregister failed, error: ", err) } if _, err := r.cli.Revoke(context.Background(), r.leasesID); err != nil { grpclog.Error("revoke failed, error: ", err) } r.closeCh <- struct{}{} return case res := <-r.keepAliveCh: if res == nil { if err := r.register(); err != nil { grpclog.Error("register failed, error: ", err) } } case <-ticker.C: if r.keepAliveCh == nil { if err := r.register(); err != nil { grpclog.Error("register failed, error: ", err) } } } } } func (r *Register) GetServerInfo() (Server, error) { resp, err := r.cli.Get(context.Background(), BuildRegisterPath(r.srvInfo)) if err != nil { return r.srvInfo, err } server := Server{} if resp.Count >= 1 { if err := json.Unmarshal(resp.Kvs[0].Value, &server); err != nil { return server, err } } return server, err }