27 changed files with 685 additions and 214 deletions
@ -0,0 +1,37 @@
|
||||
package discovery |
||||
|
||||
import ( |
||||
"context" |
||||
"google.golang.org/grpc/resolver" |
||||
) |
||||
|
||||
type Registry interface { |
||||
// Registry server instance
|
||||
Registry(ctx context.Context, server Server) (err error) |
||||
} |
||||
|
||||
type GrpcResolver interface { |
||||
DialUrl(serviceName string) string |
||||
// Resolver get grpc dial resolver
|
||||
Resolver() (builder resolver.Builder, err error) |
||||
} |
||||
|
||||
type Resolver interface { |
||||
// ResolveAll get service all instance
|
||||
ResolveAll(ctx context.Context, serviceName string) (servers []Server, err error) |
||||
// Watch when service instance change, send current all instances to channel
|
||||
Watch(ctx context.Context, serviceName string) (ch chan []Server, err error) |
||||
} |
||||
|
||||
type Discovery interface { |
||||
Registry |
||||
Resolver |
||||
GrpcResolver |
||||
} |
||||
|
||||
// Server registry format
|
||||
type Server struct { |
||||
Name string `json:"name"` |
||||
Addr string `json:"addr"` // 地址
|
||||
Attrs map[string]string `json:"attrs"` // attributes
|
||||
} |
||||
@ -0,0 +1,183 @@
|
||||
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 |
||||
} |
||||
@ -0,0 +1,43 @@
|
||||
package discovery |
||||
|
||||
import ( |
||||
"context" |
||||
clientv3 "go.etcd.io/etcd/client/v3" |
||||
"testing" |
||||
"time" |
||||
) |
||||
|
||||
func TestResolveAll(t *testing.T) { |
||||
client, err := clientv3.New(clientv3.Config{ |
||||
Endpoints: []string{"127.0.0.1:2379"}, |
||||
}) |
||||
if err != nil { |
||||
t.Error(err) |
||||
} |
||||
dis := NewEtcdDiscovery(client) |
||||
|
||||
serviceName := "TestService" |
||||
// registry
|
||||
s1 := Server{ |
||||
Addr: "127.0.0.1:1234", |
||||
Name: serviceName, |
||||
Attrs: map[string]string{"weight": "10"}, |
||||
} |
||||
ctx, cancel := context.WithCancel(context.Background()) |
||||
err = dis.Registry(ctx, s1) |
||||
if err != nil { |
||||
t.Error(err) |
||||
} |
||||
|
||||
// resolve
|
||||
servers, err := dis.ResolveAll(context.Background(), serviceName) |
||||
if err != nil { |
||||
t.Error(err) |
||||
} |
||||
if len(servers) == 0 || servers[0].Addr != s1.Addr { |
||||
t.Error("resolveAll server addr error") |
||||
} |
||||
|
||||
cancel() |
||||
time.Sleep(time.Second) |
||||
} |
||||
@ -1,80 +0,0 @@
|
||||
package discovery |
||||
|
||||
import ( |
||||
"context" |
||||
"encoding/json" |
||||
"fmt" |
||||
clientv3 "go.etcd.io/etcd/client/v3" |
||||
"go.etcd.io/etcd/client/v3/naming/endpoints" |
||||
"sonet/pkg/utils/logger" |
||||
"sonet/pkg/utils/shutdown" |
||||
"time" |
||||
) |
||||
|
||||
func EtcdDialUrl(serviceName string) string { |
||||
return fmt.Sprintf("etcd:///%s", serviceName) |
||||
} |
||||
|
||||
func EtcdRegistry(client *clientv3.Client, server *Server) (err error) { |
||||
em, err := endpoints.NewManager(client, server.Name) |
||||
if err != nil { |
||||
return |
||||
} |
||||
ip, port, err := RegisterIpPort(server.Addr) |
||||
if err != nil { |
||||
return |
||||
} |
||||
|
||||
server.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) |
||||
} |
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second*10) |
||||
defer cancel() |
||||
|
||||
lease, err := client.Grant(ctx, DefaultRegisterTTL) // 使用租约注册端点确保如果主机无法维持保活心跳, 从服务中删除
|
||||
if err != nil { |
||||
return |
||||
} |
||||
endpointKey := fmt.Sprintf("%s/%s", server.Name, ip) |
||||
err = em.AddEndpoint(ctx, |
||||
endpointKey, |
||||
endpoints.Endpoint{ |
||||
Addr: server.Addr, |
||||
Metadata: meta, |
||||
}, |
||||
clientv3.WithLease(lease.ID), |
||||
) |
||||
|
||||
// keepalive lease
|
||||
keepAliveCh, err := client.KeepAlive(context.Background(), lease.ID) |
||||
doneCh := make(chan bool) |
||||
c := func() { |
||||
doneCh <- true |
||||
<-doneCh |
||||
} |
||||
shutdown.AddShutdownHook(c) |
||||
go func() { |
||||
for { |
||||
select { |
||||
case <-doneCh: // async... not done
|
||||
logger.Info("keepalive done") |
||||
ctx, c := context.WithTimeout(context.Background(), time.Second*5) |
||||
_, _ = client.Revoke(ctx, lease.ID) |
||||
c() |
||||
doneCh <- true |
||||
return |
||||
case _ = <-keepAliveCh: |
||||
} |
||||
} |
||||
}() |
||||
return |
||||
} |
||||
@ -0,0 +1,133 @@
|
||||
package deliver |
||||
|
||||
import ( |
||||
"context" |
||||
"google.golang.org/grpc" |
||||
"google.golang.org/protobuf/proto" |
||||
"sonet/api/gen/postal" |
||||
"sonet/pkg/grpc/discovery" |
||||
"sonet/pkg/utils/logger" |
||||
"sync" |
||||
) |
||||
|
||||
type GroupLoader interface { |
||||
Load(uid string) (groupIds []string) |
||||
} |
||||
|
||||
type Postal struct { |
||||
conn *grpc.ClientConn |
||||
Client postal.PostalClient |
||||
} |
||||
|
||||
type GroupDeliver struct { |
||||
svcName string |
||||
groupLoader GroupLoader |
||||
resolver discovery.Resolver |
||||
postalDialOptions []grpc.DialOption |
||||
postals map[string]*Postal |
||||
lock *sync.RWMutex |
||||
} |
||||
|
||||
func NewGroupDeliver( |
||||
msgInServiceName string, |
||||
groupLoader GroupLoader, |
||||
resolver discovery.Resolver, |
||||
postalDialOptions []grpc.DialOption, |
||||
) *GroupDeliver { |
||||
return &GroupDeliver{ |
||||
svcName: msgInServiceName, |
||||
groupLoader: groupLoader, |
||||
resolver: resolver, |
||||
postalDialOptions: postalDialOptions, |
||||
} |
||||
} |
||||
|
||||
func (d *GroupDeliver) Init(ctx context.Context) (err error) { |
||||
// initial all postal clients
|
||||
servers, err := d.resolver.ResolveAll(ctx, postal.Postal_ServiceDesc.ServiceName) |
||||
if err != nil { |
||||
return |
||||
} |
||||
d.buildPostals(servers) |
||||
|
||||
// watch postal server instance
|
||||
err = d.watchPostal(ctx) |
||||
if err != nil { |
||||
return |
||||
} |
||||
return |
||||
} |
||||
|
||||
// watchPostal watch postal service list
|
||||
func (d *GroupDeliver) watchPostal(ctx context.Context) (err error) { |
||||
ch, err := d.resolver.Watch(ctx, postal.Postal_ServiceDesc.ServiceName) |
||||
if err != nil { |
||||
return |
||||
} |
||||
go func() { |
||||
for { |
||||
select { |
||||
case <-ctx.Done(): |
||||
return |
||||
case servers := <-ch: |
||||
d.buildPostals(servers) |
||||
} |
||||
} |
||||
}() |
||||
return |
||||
} |
||||
|
||||
func (d *GroupDeliver) buildPostals(servers []discovery.Server) { |
||||
postals := make(map[string]*Postal) |
||||
for _, server := range servers { |
||||
if p, ok := d.postals[server.Addr]; ok { |
||||
postals[server.Addr] = p |
||||
continue |
||||
} |
||||
// new connection
|
||||
conn, err := grpc.DialContext(context.Background(), server.Addr, d.postalDialOptions...) |
||||
if err != nil { |
||||
logger.Errorf("dial postal server %+v error: %v", server.Addr, err) |
||||
continue |
||||
} |
||||
postals[server.Addr] = &Postal{ |
||||
conn: conn, |
||||
Client: postal.NewPostalClient(conn), |
||||
} |
||||
} |
||||
|
||||
// close old connection
|
||||
oldPostals := d.postals |
||||
d.postals = postals |
||||
for addr, p := range oldPostals { |
||||
if _, ok := d.postals[addr]; !ok { |
||||
if err := p.conn.Close(); err != nil { |
||||
logger.Errorf("close old postal conn %s error: %v", addr, err) |
||||
} |
||||
} |
||||
} |
||||
} |
||||
|
||||
func (d *GroupDeliver) DeliverGroup(ctx context.Context, gid string, msg proto.Message, options ...Option) (err error) { |
||||
// deliver to all postal
|
||||
message, err := protoMessage2Deliver(d.svcName, msg) |
||||
if err != nil { |
||||
return |
||||
} |
||||
req := &postal.ReqDeliverGroup{ |
||||
Gid: gid, |
||||
Msg: message, |
||||
} |
||||
for _, p := range d.postals { |
||||
res, err := p.Client.DeliverGroup(ctx, req) |
||||
if err != nil || !res.Ok { |
||||
logger.Errorf("deliver to postal failed: %+v, %v", res, err) |
||||
} |
||||
} |
||||
return |
||||
} |
||||
|
||||
func (d *GroupDeliver) GroupJoin(ctx context.Context, uid string, gid []string) (err error) { |
||||
|
||||
return |
||||
} |
||||
@ -0,0 +1,29 @@
|
||||
package shutdown |
||||
|
||||
var defaultOptions = Options{ |
||||
Order: 1, |
||||
} |
||||
|
||||
type Options struct { |
||||
Order int // 0头部, 1中间, 2尾部
|
||||
} |
||||
|
||||
type Option func(opts *Options) |
||||
|
||||
func WithOrderFront() Option { |
||||
return func(opts *Options) { |
||||
opts.Order = 0 |
||||
} |
||||
} |
||||
|
||||
func WithOrderMiddle() Option { |
||||
return func(opts *Options) { |
||||
opts.Order = 1 |
||||
} |
||||
} |
||||
|
||||
func WithOrderBack() Option { |
||||
return func(opts *Options) { |
||||
opts.Order = 2 |
||||
} |
||||
} |
||||
Loading…
Reference in new issue