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.
 
 

183 lines
3.9 KiB

package discovery
import (
"context"
"encoding/json"
"fmt"
"github.com/bytedance/sonic"
"go.etcd.io/etcd/api/v3/mvccpb"
clientv3 "go.etcd.io/etcd/client/v3"
"go.etcd.io/etcd/client/v3/naming/endpoints"
etcdResolver "go.etcd.io/etcd/client/v3/naming/resolver"
"google.golang.org/grpc/resolver"
"net"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/nets"
"time"
)
var (
DefaultRegisterTTL int64 = 30
)
const (
EtcdSchema = "etcd"
)
func EtcdDialUrl(serviceName string) string {
return fmt.Sprintf("%s:///%s", EtcdSchema, serviceName)
}
type EtcdDiscovery struct {
client *clientv3.Client
}
func NewEtcdDiscovery(client *clientv3.Client) *EtcdDiscovery {
return &EtcdDiscovery{
client: client,
}
}
func (r *EtcdDiscovery) Registry(ctx context.Context, server Server) (err error) {
em, err := endpoints.NewManager(r.client, server.Name)
if err != nil {
return
}
ip, port, err := RegisterIpPort(server.Addr)
if err != nil {
return
}
addr := fmt.Sprintf("%s:%d", ip, port)
// 序列化 metadata 信息
meta := "{}"
if server.Attrs != nil {
bytes, e := json.Marshal(server.Attrs)
if e != nil {
err = e
return
}
meta = string(bytes)
}
lease, err := r.client.Grant(ctx, DefaultRegisterTTL) // 使用租约注册端点确保如果主机无法维持保活心跳, 从服务中删除
if err != nil {
return
}
endpointKey := fmt.Sprintf("%s/%s", server.Name, addr)
err = em.AddEndpoint(ctx,
endpointKey,
endpoints.Endpoint{
Addr: addr,
Metadata: meta,
},
clientv3.WithLease(lease.ID),
)
// keepalive lease
keepAliveCh, err := r.client.KeepAlive(context.Background(), lease.ID)
go func() {
for {
select {
case <-ctx.Done():
logger.Info("registry keepalive done")
ctx, c := context.WithTimeout(context.Background(), time.Second*2)
defer c()
_, _ = r.client.Revoke(ctx, lease.ID)
return
case _ = <-keepAliveCh:
}
}
}()
return
}
func (r *EtcdDiscovery) DialUrl(serviceName string) string {
return fmt.Sprintf("%s:///%s", EtcdSchema, serviceName)
}
func (r *EtcdDiscovery) Resolver() (builder resolver.Builder, err error) {
builder, err = etcdResolver.NewBuilder(r.client)
return
}
func (r *EtcdDiscovery) ResolveAll(ctx context.Context, serviceName string) (servers []Server, err error) {
res, err := r.client.Get(ctx, serviceName+"/", clientv3.WithPrefix())
if err != nil {
return
}
for _, kv := range res.Kvs {
endpoint := endpoints.Endpoint{}
if err = sonic.Unmarshal(kv.Value, &endpoint); err != nil {
logger.Errorf("resolve service %s error: ", string(kv.Value))
return
}
server := Server{
Name: serviceName,
Addr: endpoint.Addr,
}
if endpoint.Metadata != nil {
if strMeta, ok := endpoint.Metadata.(string); ok {
err = json.Unmarshal([]byte(strMeta), &server.Attrs)
if err != nil {
return
}
}
}
servers = append(servers, server)
}
return
}
func (r *EtcdDiscovery) Watch(ctx context.Context, serviceName string) (ch chan []Server, err error) {
w := r.client.Watch(ctx, serviceName+"/", clientv3.WithPrefix())
ch = make(chan []Server, 1)
go func() {
for {
select {
case <-ctx.Done():
close(ch)
return
case res := <-w:
if err := res.Err(); err != nil {
logger.Errorf("watch service %s error: %v", serviceName, err)
continue
}
for _, event := range res.Events {
switch event.Type {
case mvccpb.DELETE:
fallthrough
case mvccpb.PUT:
servers, err := r.ResolveAll(context.Background(), serviceName)
if err != nil {
logger.Errorf("watch event %v for service %s error: %v", event.Type, serviceName, err)
continue
}
ch <- servers
}
}
}
}
}()
return
}
func RegisterIpPort(addr string) (ip string, port int, err error) {
tcpAddr, err := net.ResolveTCPAddr("tcp", addr)
if err != nil {
return
}
port = tcpAddr.Port
if tcpAddr.IP != nil {
ip = tcpAddr.IP.String()
} else {
ip, err = nets.GetHostIpv4()
if err != nil {
return
}
}
return
}