Browse Source

grpc etcd discovery

master
tangmingyou 3 years ago
parent
commit
fbdf3ab365
  1. 8
      api/postal.proto
  2. 38
      benchmark/main.go
  3. 6
      cmd/auth/main.go
  4. 19
      cmd/chat/main.go
  5. 20
      cmd/gateway_http/main.go
  6. 2
      cmd/gateway_ws/config.toml
  7. 27
      cmd/gateway_ws/main.go
  8. 16
      cmd/mahjong/main.go
  9. 10
      internal/auth/logic/auth_server.go
  10. 10
      internal/chat/logic/chat_server.go
  11. 10
      internal/mahjong/logic/mahjong_server.go
  12. 10
      internal/postal/logic/postal_server.go
  13. 37
      pkg/grpc/discovery/discovery.go
  14. 39
      pkg/grpc/discovery/etcd/instance.go
  15. 4
      pkg/grpc/discovery/etcd/register.go
  16. 3
      pkg/grpc/discovery/etcd/resolver.go
  17. 183
      pkg/grpc/discovery/etcd_naming.go
  18. 43
      pkg/grpc/discovery/etcd_naming_test.go
  19. 80
      pkg/grpc/discovery/etcd_registry.go
  20. 38
      pkg/protocol/deliver/deliver.go
  21. 133
      pkg/protocol/deliver/group_deliver.go
  22. 1
      pkg/protocol/deliver/group_loader.go
  23. 55
      pkg/utils/collect/concurrent_map.go
  24. 18
      pkg/utils/collect/concurrent_map_test.go
  25. 9
      pkg/utils/logger/logger.go
  26. 29
      pkg/utils/shutdown/option.go
  27. 51
      pkg/utils/shutdown/signal.go

8
api/postal.proto

@ -66,13 +66,13 @@ message ResGroupCreate {
}
message ReqGroupJoin {
string gid = 1;
repeated string uid = 2;
string uid = 1;
repeated string gid = 2;
}
message ReqGroupLeave {
string gid = 1;
repeated string uid = 2;
string uid = 1;
repeated string gid = 2;
}
message ReqGroupDissolve {

38
benchmark/main.go

@ -6,13 +6,16 @@ import (
"fmt"
"github.com/bytedance/sonic"
"github.com/gorilla/websocket"
clientv3 "go.etcd.io/etcd/client/v3"
"google.golang.org/protobuf/proto"
"math/rand"
"net"
"runtime"
"sonet/api/gen/auth"
"sonet/api/gen/chat"
"sonet/api/gen/postal"
"sonet/pkg/config"
"sonet/pkg/grpc/discovery"
"sonet/pkg/protocol"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/security"
@ -69,6 +72,41 @@ type NetUser struct {
}
func main() {
client, err := clientv3.New(clientv3.Config{
Endpoints: []string{"124.222.131.236:3279"},
Username: "root",
Password: "sopod@etcd",
})
if err != nil {
panic(err)
}
servers, err := discovery.ResolveAll(context.Background(), client, postal.Postal_ServiceDesc.ServiceName)
if err != nil {
panic(err)
}
fmt.Printf("%+v\n", servers)
ctx, cancel := context.WithCancel(context.Background())
ch := discovery.Watch(ctx, client, postal.Postal_ServiceDesc.ServiceName)
go func() {
for {
select {
case <-ctx.Done():
fmt.Println("done2")
return
case servers := <-ch:
fmt.Printf("watch services: %+v\n", servers)
}
}
}()
shutdown.AddHook(cancel)
shutdown.Await()
}
func main2() {
runtime.GOMAXPROCS(runtime.NumCPU())
// go prof.StartPprof(":8888")

6
cmd/auth/main.go

@ -5,6 +5,7 @@ import (
"sonet/internal/auth/data"
"sonet/internal/auth/logic"
"sonet/pkg/config"
"sonet/pkg/grpc/discovery"
"sonet/pkg/utils/shutdown"
)
@ -21,12 +22,15 @@ func main() {
if err != nil {
panic(err)
}
shutdown.AddHook(func() { _ = etcdClient.Close() }, shutdown.WithOrderBack())
// user dao
userDao := data.NewUserDao(config.NewGorm(conf.Gorm))
authServer := logic.NewAuthServer(appConf.AesTokenKey, userDao)
dis := discovery.NewEtcdDiscovery(etcdClient)
go func() {
err = authServer.Run(conf.Grpc, etcdClient)
err = authServer.Run(conf.Grpc, dis)
if err != nil {
panic(err)
}

19
cmd/chat/main.go

@ -3,10 +3,10 @@ package main
import (
"context"
clientv3 "go.etcd.io/etcd/client/v3"
"go.etcd.io/etcd/client/v3/naming/resolver"
"sonet/api/gen/chat"
"sonet/internal/chat/logic"
"sonet/pkg/config"
"sonet/pkg/grpc/discovery"
"sonet/pkg/protocol/deliver"
"sonet/pkg/utils/shutdown"
)
@ -19,18 +19,23 @@ func main() {
if err != nil {
panic(err)
}
shutdown.AddHook(func() {
if err := etcdClient.Close(); err != nil {
panic(err)
}
}, shutdown.WithOrderBack())
// etcd discovery
dis := discovery.NewEtcdDiscovery(etcdClient)
etcdResolver, err := resolver.NewBuilder(etcdClient)
if err != nil {
panic(err)
}
deli := deliver.NewDeliver(chat.Chat_ServiceDesc.ServiceName)
if err := deli.InitWithResolver(context.Background(), etcdResolver); err != nil {
if err := deli.InitWithResolver(context.Background(), dis); err != nil {
panic(err)
}
chatServer := logic.NewChatServer(deli)
go func() {
err = chatServer.Run(conf.Grpc, etcdClient)
err = chatServer.Run(conf.Grpc, dis)
if err != nil {
panic(err)
}

20
cmd/gateway_http/main.go

@ -5,7 +5,6 @@ import (
"fmt"
"github.com/gin-gonic/gin"
clientv3 "go.etcd.io/etcd/client/v3"
"go.etcd.io/etcd/client/v3/naming/resolver"
"google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure"
"net/http"
@ -13,6 +12,7 @@ import (
config2 "sonet/internal/gateway_http/config"
"sonet/internal/gateway_http/logic"
"sonet/pkg/config"
"sonet/pkg/grpc/discovery"
"sonet/pkg/grpc/generic"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/resp"
@ -34,9 +34,11 @@ func main() {
if err != nil {
panic(err)
}
//etcdResolver := discovery.NewResolver(etcdClient)
//shutdown.AddShutdownHook(etcdResolver.Close)
//resolver.Register(etcdResolver)
shutdown.AddHook(func() {
if err := etcdClient.Close(); err != nil {
panic(err)
}
}, shutdown.WithOrderBack())
// gin http server
authFilter, err := config2.NewAuthFilter(appConf.AesTokenKey, appConf.IgnoreUrls)
@ -62,13 +64,15 @@ func main() {
// grpc services
grpcGroup := server.Group("/api/svc")
etcdResolver, err := resolver.NewBuilder(etcdClient)
dis := discovery.NewEtcdDiscovery(etcdClient)
resolver, err := dis.Resolver()
if err != nil {
panic(err)
}
grpcFactory := generic.NewGpcGenericClientFactory("etcd",
grpcFactory := generic.NewGpcGenericClientFactory(
discovery.EtcdSchema,
grpc.WithTransportCredentials(insecure.NewCredentials()),
grpc.WithResolvers(etcdResolver),
grpc.WithResolvers(resolver),
)
grpcFactory.Init()
grpcGenericHandler := logic.NewGrpcGenericHandler(grpcFactory)
@ -76,7 +80,7 @@ func main() {
// postal loadBalancer handler
postalBalancer := logic.NewPostalBalancer()
postalBalancer.Init(context.Background(), etcdResolver)
postalBalancer.Init(context.Background(), resolver)
server.GET("/api/lb/ws", postalBalancer.Endpoint)
go func() {

2
cmd/gateway_ws/config.toml

@ -6,7 +6,7 @@ subjectLrcExpiration = "10m"
subjectLrcCleanupInterval = "5m"
[grpc]
address = ":7011"
address = ":7011" # postal cluster offset port +1000=8011
maxSendMsgSize = "8Mi"
maxRecvMsgSize = "8Mi"
readBufferSize = "8Ki"

27
cmd/gateway_ws/main.go

@ -4,7 +4,6 @@ import (
"github.com/nats-io/nats.go"
"github.com/redis/go-redis/v9"
clientv3 "go.etcd.io/etcd/client/v3"
"go.etcd.io/etcd/client/v3/naming/resolver"
"google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure"
"sonet/internal/gateway_ws/gws_server"
@ -12,12 +11,12 @@ import (
"sonet/pkg/config"
"sonet/pkg/grpc/client"
"sonet/pkg/grpc/discovery"
"sonet/pkg/grpc/discovery/etcd"
"sonet/pkg/grpc/generic"
"sonet/pkg/plugins/cache"
"sonet/pkg/plugins/mq"
"sonet/pkg/protocol/session"
"sonet/pkg/utils/conver"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/shutdown"
)
@ -40,16 +39,20 @@ func main() {
if err != nil {
panic(err)
}
shutdown.AddHook(func() { _ = etcdClient.Close() }, shutdown.WithOrderBack())
subjectStore, err := getSubjectMultiCache(appConf, conf.Redis, conf.Nats)
if err != nil {
panic(err)
}
etcdResolver, err := resolver.NewBuilder(etcdClient)
dis := discovery.NewEtcdDiscovery(etcdClient)
etcdResolver, err := dis.Resolver()
if err != nil {
panic(err)
}
grpcFactory := generic.NewGpcGenericClientFactory("etcd",
grpcFactory := generic.NewGpcGenericClientFactory(
discovery.EtcdSchema,
grpc.WithTransportCredentials(insecure.NewCredentials()),
grpc.WithResolvers(etcdResolver),
)
@ -57,7 +60,8 @@ func main() {
sessionStore := session.NewMapStore()
postalAddr := discovery.MustRegisterAddress(conf.Grpc.Address)
// run websocket server
postalAddr := etcd.MustRegisterAddress(conf.Grpc.Address)
//connHandler := server.NewConnHandler(postalAddr, grpcFactory, sessionStore, subjectStore)
//httpServer := server.NewHttpServer(connHandler)
gwsHandler := gws_server.NewGwsHandler(postalAddr, grpcFactory, sessionStore, subjectStore)
@ -72,15 +76,10 @@ func main() {
}()
// run postal server
shutdown.AddShutdownHook(func() {
if err := etcdClient.Close(); err != nil {
logger.Error("etcd close error: ", err)
}
})
clientFactory := client.NewGrpcDirectClientFactory(grpc.WithTransportCredentials(insecure.NewCredentials()))
postalServer := logic.NewPostalServer(appConf.EndpointAddress, sessionStore, subjectStore, clientFactory)
go func() {
err = postalServer.Run(conf.Grpc, etcdClient)
err = postalServer.Run(conf.Grpc, dis)
if err != nil {
panic(err)
}
@ -102,19 +101,19 @@ func getSubjectMultiCache(appConf *GatewayWsConfig, redisOptions redis.Options,
// initial cache...
rdb := redis.NewClient(&redisOptions)
subjectRedisCache := cache.NewRedisCache(appConf.SubjectCacheTopic, rdb)
shutdown.AddShutdownHook(func() { _ = rdb.Close() })
shutdown.AddHook(func() { _ = rdb.Close() })
// nats mq
producer, err := mq.NewNatsProducer(natsOptions)
if err != nil {
return nil, err
}
shutdown.AddShutdownHook(func() { producer.Stop() })
shutdown.AddHook(func() { producer.Stop() })
consumer, err := mq.NewNatsConsumer(natsOptions)
if err != nil {
return nil, err
}
shutdown.AddShutdownHook(func() { consumer.Stop() })
shutdown.AddHook(func() { consumer.Stop() })
// 多级缓存
subjectLrcOpts := cache.LocalRemoteCacheOptions{

16
cmd/mahjong/main.go

@ -3,7 +3,6 @@ package main
import (
"context"
clientv3 "go.etcd.io/etcd/client/v3"
"go.etcd.io/etcd/client/v3/naming/resolver"
"google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure"
"sonet/api/gen/auth"
@ -12,6 +11,7 @@ import (
"sonet/internal/mahjong/store"
"sonet/pkg/config"
"sonet/pkg/grpc/discovery"
"sonet/pkg/grpc/discovery/etcd"
"sonet/pkg/protocol/deliver"
"sonet/pkg/utils/shutdown"
)
@ -25,17 +25,17 @@ func main() {
panic(err)
}
registry := discovery.NewRegister(etcdClient)
shutdown.AddShutdownHook(registry.Stop)
registry := etcd.NewRegister(etcdClient)
shutdown.AddHook(registry.Stop)
// init deliver
// etcdResolver := discovery.NewResolver(etcdClient)
etcdResolver, err := resolver.NewBuilder(etcdClient)
dis := discovery.NewEtcdDiscovery(etcdClient)
resolver, err := dis.Resolver()
if err != nil {
panic(err)
}
deli := deliver.NewDeliver(mahjong.Mahjong_ServiceDesc.ServiceName)
if err := deli.InitWithResolver(context.Background(), etcdResolver); err != nil {
if err := deli.InitWithResolver(context.Background(), dis); err != nil {
panic(err)
}
mjStore := store.NewStore(deli)
@ -44,14 +44,14 @@ func main() {
url := discovery.EtcdDialUrl(auth.Auth_ServiceDesc.ServiceName)
authConn, err := grpc.DialContext(context.Background(), url,
grpc.WithTransportCredentials(insecure.NewCredentials()),
grpc.WithResolvers(etcdResolver),
grpc.WithResolvers(resolver),
)
if err != nil {
panic(err)
}
mahjongServer := logic.NewMahjongServer(mjStore, deli, auth.NewAuthClient(authConn))
go func() {
err = mahjongServer.Run(conf.Grpc, etcdClient)
err = mahjongServer.Run(conf.Grpc, dis)
if err != nil {
panic(err)
}

10
internal/auth/logic/auth_server.go

@ -5,7 +5,6 @@ import (
"encoding/base64"
"errors"
"github.com/bytedance/sonic"
clientv3 "go.etcd.io/etcd/client/v3"
"google.golang.org/grpc"
"google.golang.org/grpc/reflection"
"net"
@ -16,6 +15,7 @@ import (
"sonet/pkg/grpc/interceptor"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/security"
"sonet/pkg/utils/shutdown"
"time"
)
@ -36,7 +36,7 @@ func NewAuthServer(aesTokenKey string, userDao *data.AuthUserDao) *AuthServer {
}
}
func (s *AuthServer) Run(conf config.GrpcConfig, etcd *clientv3.Client) (err error) {
func (s *AuthServer) Run(conf config.GrpcConfig, registry discovery.Registry) (err error) {
server := grpc.NewServer(
config.GetGrpcOptions(
conf,
@ -54,17 +54,19 @@ func (s *AuthServer) Run(conf config.GrpcConfig, etcd *clientv3.Client) (err err
}
// registry discovery
register := &(conf.Register)
register := conf.Register
if register.Name == "" {
register.Name = auth.Auth_ServiceDesc.ServiceName
}
if register.Addr == "" {
register.Addr = conf.Address
}
err = discovery.EtcdRegistry(etcd, register)
ctx, cancel := context.WithCancel(context.Background())
err = registry.Registry(ctx, register)
if err != nil {
panic(err)
}
shutdown.AddHook(cancel, shutdown.WithOrderFront())
// run serve
logger.Infof("%s grpc server running %s\n", register.Name, listen.Addr().String())

10
internal/chat/logic/chat_server.go

@ -2,7 +2,6 @@ package logic
import (
"context"
clientv3 "go.etcd.io/etcd/client/v3"
"google.golang.org/grpc"
"google.golang.org/grpc/reflection"
"google.golang.org/protobuf/types/known/emptypb"
@ -14,6 +13,7 @@ import (
"sonet/pkg/protocol/deliver"
"sonet/pkg/protocol/session"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/shutdown"
)
type ChatServer struct {
@ -27,7 +27,7 @@ func NewChatServer(deliver *deliver.Deliver) *ChatServer {
}
}
func (s *ChatServer) Run(conf config.GrpcConfig, etcd *clientv3.Client) (err error) {
func (s *ChatServer) Run(conf config.GrpcConfig, registry discovery.Registry) (err error) {
server := grpc.NewServer(
config.GetGrpcOptions(
conf,
@ -45,17 +45,19 @@ func (s *ChatServer) Run(conf config.GrpcConfig, etcd *clientv3.Client) (err err
}
// registry discovery
register := &(conf.Register)
register := conf.Register
if register.Name == "" {
register.Name = chat.Chat_ServiceDesc.ServiceName
}
if register.Addr == "" {
register.Addr = conf.Address
}
err = discovery.EtcdRegistry(etcd, register)
ctx, cancel := context.WithCancel(context.Background())
err = registry.Registry(ctx, register)
if err != nil {
panic(err)
}
shutdown.AddHook(cancel)
// run serve
logger.Infof("%s grpc server running %s\n", register.Name, listen.Addr().String())

10
internal/mahjong/logic/mahjong_server.go

@ -3,7 +3,6 @@ package logic
import (
"context"
"errors"
clientv3 "go.etcd.io/etcd/client/v3"
"google.golang.org/grpc"
"google.golang.org/grpc/reflection"
"google.golang.org/protobuf/types/known/emptypb"
@ -20,6 +19,7 @@ import (
"sonet/pkg/protocol/session"
"sonet/pkg/utils/collect"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/shutdown"
"sonet/pkg/utils/state"
)
@ -38,7 +38,7 @@ func NewMahjongServer(store *store.Store, deliver *deliver.Deliver, authClient a
}
}
func (mj *MahjongServer) Run(conf config.GrpcConfig, etcd *clientv3.Client) (err error) {
func (mj *MahjongServer) Run(conf config.GrpcConfig, registry discovery.Registry) (err error) {
server := grpc.NewServer(
config.GetGrpcOptions(
conf,
@ -56,17 +56,19 @@ func (mj *MahjongServer) Run(conf config.GrpcConfig, etcd *clientv3.Client) (err
}
// registry discovery
register := &(conf.Register)
register := conf.Register
if register.Name == "" {
register.Name = mahjong.Mahjong_ServiceDesc.ServiceName
}
if register.Addr == "" {
register.Addr = conf.Address
}
err = discovery.EtcdRegistry(etcd, register)
ctx, cancel := context.WithCancel(context.Background())
err = registry.Registry(ctx, register)
if err != nil {
panic(err)
}
shutdown.AddHook(cancel)
// run serve
logger.Infof("%s grpc server running %s\n", register.Name, listen.Addr().String())

10
internal/postal/logic/postal_server.go

@ -2,7 +2,6 @@ package logic
import (
"context"
clientv3 "go.etcd.io/etcd/client/v3"
"google.golang.org/grpc"
"google.golang.org/grpc/reflection"
"google.golang.org/protobuf/types/known/emptypb"
@ -16,6 +15,7 @@ import (
"sonet/pkg/protocol"
"sonet/pkg/protocol/session"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/shutdown"
)
type PostalServer struct {
@ -40,7 +40,7 @@ func NewPostalServer(
}
}
func (s *PostalServer) Run(conf config.GrpcConfig, etcd *clientv3.Client) (err error) {
func (s *PostalServer) Run(conf config.GrpcConfig, registry discovery.Registry) (err error) {
server := grpc.NewServer(
config.GetGrpcOptions(
conf,
@ -58,17 +58,19 @@ func (s *PostalServer) Run(conf config.GrpcConfig, etcd *clientv3.Client) (err e
}
// registry discovery
register := &(conf.Register)
register := conf.Register
if register.Name == "" {
register.Name = postal.Postal_ServiceDesc.ServiceName
}
if register.Addr == "" {
register.Addr = conf.Address
}
err = discovery.EtcdRegistry(etcd, register)
ctx, cancel := context.WithCancel(context.Background())
err = registry.Registry(ctx, register)
if err != nil {
panic(err)
}
shutdown.AddHook(cancel)
// 其他服务直连地址
s.broadcastAddress = register.Addr

37
pkg/grpc/discovery/discovery.go

@ -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
}

39
pkg/grpc/discovery/instance.go → pkg/grpc/discovery/etcd/instance.go

@ -1,4 +1,4 @@
package discovery
package etcd
import (
"encoding/json"
@ -9,29 +9,13 @@ import (
"google.golang.org/grpc/resolver"
)
// Server registry format
type Server struct {
Name string `json:"name"`
Addr string `json:"addr"` // 地址
Attrs map[string]string `json:"attrs"` // attributes
}
func BuildPrefix(server Server) string {
return fmt.Sprintf("/%s/", server.Name)
}
func BuildRegisterPath(server Server) string {
return fmt.Sprintf("%s%s", BuildPrefix(server), server.Addr)
}
func ParseValue(value []byte) (Server, error) {
server := Server{}
if err := json.Unmarshal(value, &server); err != nil {
return server, err
}
return server, nil
}
func SplitPath(path string) (Server, error) {
server := Server{}
strs := strings.Split(path, "/")
@ -67,5 +51,22 @@ func Remove(s []resolver.Address, addr resolver.Address) ([]resolver.Address, bo
}
func BuildResolverUrl(app string) string {
return schema + ":///" + app
return "etcd:///" + app
}
func BuildPrefix(server Server) string {
return fmt.Sprintf("/%s/", server.Name)
}
func BuildRegisterPath(server Server) string {
return fmt.Sprintf("%s%s", BuildPrefix(server), server.Addr)
}
func ParseValue(value []byte) (Server, error) {
server := Server{}
if err := json.Unmarshal(value, &server); err != nil {
return server, err
}
return server, nil
}

4
pkg/grpc/discovery/register.go → pkg/grpc/discovery/etcd/register.go

@ -1,4 +1,4 @@
package discovery
package etcd
import (
"context"
@ -47,6 +47,8 @@ func MustRegisterAddress(addr string) (fullAddr string) {
return
}
// Register
// Deprecated
type Register struct {
DialTimeout int

3
pkg/grpc/discovery/resolver.go → pkg/grpc/discovery/etcd/resolver.go

@ -1,4 +1,4 @@
package discovery
package etcd
import (
"context"
@ -14,6 +14,7 @@ const (
)
// Resolver for grpc client
// Deprecated
type Resolver struct {
schema string
DialTimeout int

183
pkg/grpc/discovery/etcd_naming.go

@ -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
}

43
pkg/grpc/discovery/etcd_naming_test.go

@ -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)
}

80
pkg/grpc/discovery/etcd_registry.go

@ -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
}

38
pkg/protocol/deliver/deliver.go

@ -3,10 +3,9 @@ package deliver
import (
"context"
"fmt"
"github.com/golang/protobuf/proto"
"google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure"
"google.golang.org/grpc/resolver"
"google.golang.org/protobuf/proto"
"reflect"
"sonet/api/gen/postal"
"sonet/pkg/grpc/balancer"
@ -35,15 +34,19 @@ func NewDeliver(msgInServiceName string) *Deliver {
}
}
func (d *Deliver) InitWithResolver(ctx context.Context, builder resolver.Builder) (err error) {
func (d *Deliver) InitWithResolver(ctx context.Context, resolver discovery.GrpcResolver) (err error) {
balancer.InitConsistentHashBuilder()
rb, err := resolver.Resolver()
if err != nil {
return
}
postalUrl := discovery.EtcdDialUrl(postal.Postal_ServiceDesc.ServiceName)
postalUrl := resolver.DialUrl(postal.Postal_ServiceDesc.ServiceName)
conn, err := grpc.DialContext(ctx, postalUrl,
grpc.WithTransportCredentials(insecure.NewCredentials()),
// consistent hash lb
grpc.WithDefaultServiceConfig(fmt.Sprintf(`{"loadBalancingPolicy":"%s"}`, balancer.ConsistentHash)),
grpc.WithResolvers(builder),
grpc.WithResolvers(rb),
)
if err != nil {
return
@ -70,6 +73,21 @@ func (d *Deliver) DeliverBatch(ctx context.Context, msg proto.Message, receivers
return d.deliver0(ctx, msg, receivers, options...)
}
func protoMessage2Deliver(svcName string, msg proto.Message) (message *postal.Message, err error) {
body, err := proto.Marshal(msg)
if err != nil {
return
}
msgName := reflect.TypeOf(msg).Elem().Name()
message = &postal.Message{
Time: time.Now().UnixMilli(),
Svc: svcName,
Msg: msgName,
Body: body,
}
return
}
func (d *Deliver) deliver0(ctx context.Context, msg proto.Message, receivers []string, options ...Option) (status Status, err error) {
opts := defaultOptions
if options != nil {
@ -83,17 +101,11 @@ func (d *Deliver) deliver0(ctx context.Context, msg proto.Message, receivers []s
}
// encode msg
body, err := proto.Marshal(msg)
message, err := protoMessage2Deliver(d.svcName, msg)
if err != nil {
return StatusError, err
}
msgName := reflect.TypeOf(msg).Elem().Name()
message := &postal.Message{
Time: time.Now().UnixMilli(),
Svc: d.svcName,
Msg: msgName,
Body: body,
}
// deliver to gateway
if len(receivers) == 1 {
// deliver one receiver

133
pkg/protocol/deliver/group_deliver.go

@ -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
}

1
pkg/protocol/deliver/group_loader.go

@ -0,0 +1 @@
package deliver

55
internal/postal/group/concurrent_map.go → pkg/utils/collect/concurrent_map.go

@ -1,4 +1,4 @@
package group
package collect
import (
"hash/fnv"
@ -18,8 +18,12 @@ type ConcurrentMap[K comparable, V any] struct {
// NewConcurrentMap 分段锁并发 map
// segments: 分段数
// hashKeyFunc: key转string函数
// equalsFunc: value比较函数, Put时新旧值相同则不返回旧值
func NewConcurrentMap[K comparable, V any](segments int, hashKeyFunc func(K) string) *ConcurrentMap[K, V] {
func NewConcurrentMap[K comparable, V any](concurrencyLevel int, hashKeyFunc func(K) string) *ConcurrentMap[K, V] {
segments := 1 // segments = 2^n
for segments < concurrencyLevel {
segments <<= 1
}
m := &ConcurrentMap[K, V]{
hashKeyFunc: hashKeyFunc,
segments: segments,
@ -37,21 +41,22 @@ func NewConcurrentMap[K comparable, V any](segments int, hashKeyFunc func(K) str
func (m *ConcurrentMap[K, V]) segment(k K) int {
hashK := m.hashKeyFunc(k)
hash := fnv32Hash(hashK)
return int(hash) % m.segments
return int(hash) & (m.segments - 1)
}
// Store 放置新值
func (m *ConcurrentMap[K, V]) Store(k K, v V) { // (old V, hasOld bool) // 返回旧值
func (m *ConcurrentMap[K, V]) update(k K, update func(map[K]V)) {
segment := m.segment(k)
lock := m.segmentsLock[segment]
lock.Lock()
defer lock.Unlock()
//if prev, ok := m.segmentsMap[segment][k]; ok {
// if !m.equalsFunc(v, prev) { // 两值不同返回旧值
// old, hasOld = prev, true
// }
//}
m.segmentsMap[segment][k] = v
update(m.segmentsMap[segment])
}
// Store 放置新值
func (m *ConcurrentMap[K, V]) Store(k K, v V) {
m.update(k, func(segment map[K]V) {
segment[k] = v
})
return
}
@ -65,14 +70,12 @@ func (m *ConcurrentMap[K, V]) Load(k K) (v V, ok bool) {
}
func (m *ConcurrentMap[K, V]) Delete(k K) {
segment := m.segment(k)
lock := m.segmentsLock[segment]
lock.Lock()
defer lock.Unlock()
delete(m.segmentsMap[segment], k)
m.update(k, func(segment map[K]V) {
delete(segment, k)
})
}
func (m *ConcurrentMap[K, V]) Range(f func(key, value any) bool) {
func (m *ConcurrentMap[K, V]) Range(f func(key K, value V) bool) {
for i := 0; i < m.segments; i++ {
lock := m.segmentsLock[i]
func() {
@ -89,6 +92,22 @@ func (m *ConcurrentMap[K, V]) Range(f func(key, value any) bool) {
}
}
func (m *ConcurrentMap[K, V]) LoadAndDelete(k K) (prev V, loaded bool) {
m.update(k, func(segment map[K]V) {
prev, loaded = segment[k]
delete(segment, k)
})
return
}
func (m *ConcurrentMap[K, V]) Swap(k K, v V) (prev V, loaded bool) {
m.update(k, func(segment map[K]V) {
prev, loaded = segment[k]
segment[k] = v
})
return
}
func fnv32Hash(k string) uint32 {
f := fnv.New32()
_, err := f.Write([]byte(k))

18
internal/postal/group/concurrent_map_test.go → pkg/utils/collect/concurrent_map_test.go

@ -1,4 +1,4 @@
package group
package collect
import (
"fmt"
@ -9,7 +9,7 @@ import (
)
func TestConcurrentMap(t *testing.T) {
cm := NewConcurrentMap[string, string](16, func(k string) string { return k })
cm := NewConcurrentMap[string, string](19, func(k string) string { return k })
concurrent := 1000
wg := sync.WaitGroup{}
@ -26,8 +26,19 @@ func TestConcurrentMap(t *testing.T) {
}
wg.Wait()
wg.Add(concurrent)
var sum int
cm.Range(func(key, value string) bool {
sum++
if key != value {
t.Errorf("value load error: %s:%s", key, value)
}
return true
})
if sum != (concurrent * concurrent) {
t.Errorf("count error: %d", sum)
}
wg.Add(concurrent)
for i := 0; i < concurrent; i++ {
go func() {
r := rand.New(rand.NewSource(time.Now().UnixMilli()))
@ -41,6 +52,5 @@ func TestConcurrentMap(t *testing.T) {
wg.Done()
}()
}
wg.Wait()
}

9
pkg/utils/logger/logger.go

@ -71,8 +71,9 @@ var (
// grpclog:info=2,warn=3...
// logrus: warn=3,info=4...
func (l *SoLogger) V(level int) bool {
if level >= logrusLevelsLen {
return true
}
return logrusLevels[level] <= l.logger.Level
return level == 1
//if level >= logrusLevelsLen {
// return true
//}
//return logrusLevels[level] <= l.logger.Level
}

29
pkg/utils/shutdown/option.go

@ -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
}
}

51
pkg/utils/shutdown/signal.go

@ -8,18 +8,25 @@ import (
"sync"
"sync/atomic"
"syscall"
"time"
)
var (
shutdownHooks []func()
AwaitSeconds = 3
shutdownHooks [3][]func() // front,middle,back hooks
lock = &sync.Mutex{}
sigChan = make(chan os.Signal, 1)
)
func AddShutdownHook(hook func()) {
func AddHook(hook func(), options ...Option) {
opts := defaultOptions
for _, opt := range options {
opt(&opts)
}
lock.Lock()
defer lock.Unlock()
shutdownHooks = append(shutdownHooks, hook)
shutdownHooks[opts.Order] = append(shutdownHooks[opts.Order], hook)
}
func Await() {
@ -29,23 +36,37 @@ func Await() {
// 监听到关闭信号
logger.Info("catch exit signal: ", s)
awaitSeconds := AwaitSeconds
var success, fail int32
for _, hook := range shutdownHooks {
func() {
defer func() {
if err := recover(); err != nil {
log.Println("exec shutdown hook panic: ", err)
atomic.AddInt32(&fail, 1)
return
}
atomic.AddInt32(&success, 1)
for i, hooks := range shutdownHooks {
for _, hook := range hooks {
func() {
defer func() {
if err := recover(); err != nil {
log.Println("exec shutdown hook panic: ", err)
atomic.AddInt32(&fail, 1)
return
}
atomic.AddInt32(&success, 1)
}()
hook()
}()
}
// 间隔1s再执行
if i < 2 && len(hooks) > 0 {
time.Sleep(time.Second)
awaitSeconds -= 1
}
}
hook()
}()
if awaitSeconds < 1 {
awaitSeconds = 1
}
logger.Infof("execute %d shutdown hook %d ok, %d failed\n", success+fail, success, fail)
logger.Infof("execute %d shutdown hook %d ok, %d failed, exit in %d seconds...\n", success+fail, success, fail, awaitSeconds)
time.Sleep(time.Second * time.Duration(awaitSeconds))
}
//func Shutdown() {

Loading…
Cancel
Save