diff --git a/api/postal.proto b/api/postal.proto index 460a1f1..644fa6f 100644 --- a/api/postal.proto +++ b/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 { diff --git a/benchmark/main.go b/benchmark/main.go index 4fea675..d37bdcd 100644 --- a/benchmark/main.go +++ b/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") diff --git a/cmd/auth/main.go b/cmd/auth/main.go index 0b7237e..c465c58 100644 --- a/cmd/auth/main.go +++ b/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) } diff --git a/cmd/chat/main.go b/cmd/chat/main.go index a303aa3..35cb915 100644 --- a/cmd/chat/main.go +++ b/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) } diff --git a/cmd/gateway_http/main.go b/cmd/gateway_http/main.go index b5025d5..bbd802e 100644 --- a/cmd/gateway_http/main.go +++ b/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() { diff --git a/cmd/gateway_ws/config.toml b/cmd/gateway_ws/config.toml index 3e0bed5..7a97676 100644 --- a/cmd/gateway_ws/config.toml +++ b/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" diff --git a/cmd/gateway_ws/main.go b/cmd/gateway_ws/main.go index 86262bc..36e7471 100644 --- a/cmd/gateway_ws/main.go +++ b/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{ diff --git a/cmd/mahjong/main.go b/cmd/mahjong/main.go index 1dae610..7c1fb3e 100644 --- a/cmd/mahjong/main.go +++ b/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) } diff --git a/internal/auth/logic/auth_server.go b/internal/auth/logic/auth_server.go index 3d30d87..1679579 100644 --- a/internal/auth/logic/auth_server.go +++ b/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()) diff --git a/internal/chat/logic/chat_server.go b/internal/chat/logic/chat_server.go index 6e23468..6ff2342 100644 --- a/internal/chat/logic/chat_server.go +++ b/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()) diff --git a/internal/mahjong/logic/mahjong_server.go b/internal/mahjong/logic/mahjong_server.go index 145f3ab..6278cca 100644 --- a/internal/mahjong/logic/mahjong_server.go +++ b/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()) diff --git a/internal/postal/logic/postal_server.go b/internal/postal/logic/postal_server.go index 4830010..b221ecb 100644 --- a/internal/postal/logic/postal_server.go +++ b/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 diff --git a/pkg/grpc/discovery/discovery.go b/pkg/grpc/discovery/discovery.go new file mode 100644 index 0000000..8ea4794 --- /dev/null +++ b/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 +} diff --git a/pkg/grpc/discovery/instance.go b/pkg/grpc/discovery/etcd/instance.go similarity index 95% rename from pkg/grpc/discovery/instance.go rename to pkg/grpc/discovery/etcd/instance.go index 2409ba6..3aa3b56 100644 --- a/pkg/grpc/discovery/instance.go +++ b/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 } diff --git a/pkg/grpc/discovery/register.go b/pkg/grpc/discovery/etcd/register.go similarity index 98% rename from pkg/grpc/discovery/register.go rename to pkg/grpc/discovery/etcd/register.go index 6eefa55..3c007d1 100644 --- a/pkg/grpc/discovery/register.go +++ b/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 diff --git a/pkg/grpc/discovery/resolver.go b/pkg/grpc/discovery/etcd/resolver.go similarity index 99% rename from pkg/grpc/discovery/resolver.go rename to pkg/grpc/discovery/etcd/resolver.go index 1e68d16..a4b0fa7 100644 --- a/pkg/grpc/discovery/resolver.go +++ b/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 diff --git a/pkg/grpc/discovery/etcd_naming.go b/pkg/grpc/discovery/etcd_naming.go new file mode 100644 index 0000000..8d3f99c --- /dev/null +++ b/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 +} diff --git a/pkg/grpc/discovery/etcd_naming_test.go b/pkg/grpc/discovery/etcd_naming_test.go new file mode 100644 index 0000000..be9b3aa --- /dev/null +++ b/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) +} diff --git a/pkg/grpc/discovery/etcd_registry.go b/pkg/grpc/discovery/etcd_registry.go deleted file mode 100644 index 71b779f..0000000 --- a/pkg/grpc/discovery/etcd_registry.go +++ /dev/null @@ -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 -} diff --git a/pkg/protocol/deliver/deliver.go b/pkg/protocol/deliver/deliver.go index 667ff50..09e1ab1 100644 --- a/pkg/protocol/deliver/deliver.go +++ b/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 diff --git a/pkg/protocol/deliver/group_deliver.go b/pkg/protocol/deliver/group_deliver.go new file mode 100644 index 0000000..053b25b --- /dev/null +++ b/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 +} diff --git a/pkg/protocol/deliver/group_loader.go b/pkg/protocol/deliver/group_loader.go new file mode 100644 index 0000000..9eae15d --- /dev/null +++ b/pkg/protocol/deliver/group_loader.go @@ -0,0 +1 @@ +package deliver diff --git a/internal/postal/group/concurrent_map.go b/pkg/utils/collect/concurrent_map.go similarity index 62% rename from internal/postal/group/concurrent_map.go rename to pkg/utils/collect/concurrent_map.go index e5f3213..54518dd 100644 --- a/internal/postal/group/concurrent_map.go +++ b/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)) diff --git a/internal/postal/group/concurrent_map_test.go b/pkg/utils/collect/concurrent_map_test.go similarity index 70% rename from internal/postal/group/concurrent_map_test.go rename to pkg/utils/collect/concurrent_map_test.go index 841adb2..c5fe944 100644 --- a/internal/postal/group/concurrent_map_test.go +++ b/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() } diff --git a/pkg/utils/logger/logger.go b/pkg/utils/logger/logger.go index c9660e4..ca40eee 100644 --- a/pkg/utils/logger/logger.go +++ b/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 } diff --git a/pkg/utils/shutdown/option.go b/pkg/utils/shutdown/option.go new file mode 100644 index 0000000..723bdd8 --- /dev/null +++ b/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 + } +} diff --git a/pkg/utils/shutdown/signal.go b/pkg/utils/shutdown/signal.go index 3e75cfa..31e87da 100644 --- a/pkg/utils/shutdown/signal.go +++ b/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() {