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.
 
 

166 lines
3.3 KiB

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
}