From f8b57c3e4b0d7e6294b94ad8cce81f3d3b41c378 Mon Sep 17 00:00:00 2001 From: tangmingyou Date: Thu, 11 Jan 2024 15:17:42 +0800 Subject: [PATCH] postal consistent hash balancer --- api/postal.proto | 8 ++ cmd/auth/main.go | 3 - cmd/chat/main.go | 5 +- cmd/gateway_http/config.toml | 1 + cmd/gateway_http/main.go | 50 ++----- cmd/gateway_ws/config.toml | 7 +- cmd/gateway_ws/main.go | 9 +- cmd/mahjong/main.go | 25 ---- internal/gateway_http/config/auth_filter.go | 11 ++ internal/gateway_http/logic/http_server.go | 61 ++++++++ .../gateway_http/logic/postal_balancer.go | 54 +++++++ internal/gateway_http/logic/postal_monitor.go | 80 +++++++++++ internal/postal/logic/postal_server.go | 19 ++- pkg/config/loader.go | 4 +- pkg/config/logger.go | 16 +++ pkg/grpc/balancer/consistent_hash.go | 98 +++++++++++++ pkg/grpc/balancer/consistent_ketama.go | 136 ++++++++++++++++++ pkg/grpc/balancer/properties.go | 6 + pkg/grpc/discovery/register.go | 2 +- pkg/grpc/discovery/resolver.go | 8 +- pkg/grpc/meta/meta.go | 18 +++ pkg/protocol/deliver/deliver.go | 16 ++- pkg/protocol/protocol_test.go | 2 +- pkg/utils/logger/logger.go | 2 +- 24 files changed, 544 insertions(+), 97 deletions(-) create mode 100644 internal/gateway_http/logic/postal_balancer.go create mode 100644 internal/gateway_http/logic/postal_monitor.go create mode 100644 pkg/config/logger.go create mode 100644 pkg/grpc/balancer/consistent_hash.go create mode 100644 pkg/grpc/balancer/consistent_ketama.go create mode 100644 pkg/grpc/balancer/properties.go create mode 100644 pkg/grpc/meta/meta.go diff --git a/api/postal.proto b/api/postal.proto index 79232e5..548c066 100644 --- a/api/postal.proto +++ b/api/postal.proto @@ -17,6 +17,9 @@ service Postal { rpc GroupJoin(ReqGroupJoin) returns(google.protobuf.Empty); rpc GroupLeave(ReqGroupLeave) returns(google.protobuf.Empty); rpc GroupDissolve(ReqGroupDissolve) returns(google.protobuf.Empty); + + // 返回当前节点外部连接断点, 如 websocket addr + rpc Endpoint(google.protobuf.Empty) returns(ResEndpoint); } message Message { @@ -90,3 +93,8 @@ message ReqRedirect { ReqDeliverBatch deliverBatch = 11; ReqDeliverGroup deliverGroup = 12; } + +message ResEndpoint { + string endpoint = 1; + map extra = 2; +} diff --git a/cmd/auth/main.go b/cmd/auth/main.go index 8761e3f..9553d12 100644 --- a/cmd/auth/main.go +++ b/cmd/auth/main.go @@ -2,12 +2,10 @@ package main import ( clientv3 "go.etcd.io/etcd/client/v3" - "google.golang.org/grpc/grpclog" "sonet/internal/auth/data" "sonet/internal/auth/logic" "sonet/pkg/config" "sonet/pkg/grpc/discovery" - "sonet/pkg/utils/logger" "sonet/pkg/utils/shutdown" ) @@ -16,7 +14,6 @@ type AuthConfig struct { } func main() { - grpclog.SetLoggerV2(logger.Logger) appConf := &AuthConfig{} conf := config.LoadConfig(appConf, "cmd/auth") diff --git a/cmd/chat/main.go b/cmd/chat/main.go index e09ce94..2b08914 100644 --- a/cmd/chat/main.go +++ b/cmd/chat/main.go @@ -3,19 +3,16 @@ package main import ( "context" clientv3 "go.etcd.io/etcd/client/v3" - "google.golang.org/grpc/grpclog" "google.golang.org/grpc/resolver" "sonet/api/gen/chat" "sonet/internal/chat/logic" "sonet/pkg/config" "sonet/pkg/grpc/discovery" "sonet/pkg/protocol/deliver" - "sonet/pkg/utils/logger" "sonet/pkg/utils/shutdown" ) func main() { - grpclog.SetLoggerV2(logger.Logger) conf := config.LoadConfig(nil, "cmd/chat") // registry and run... @@ -31,7 +28,7 @@ func main() { resolver.Register(etcdResolver) deli := deliver.NewDeliver(chat.Chat_ServiceDesc.ServiceName) - if err := deli.InitWithResolver(context.Background(), etcdResolver); err != nil { + if err := deli.InitWithResolver(context.Background()); err != nil { panic(err) } chatServer := logic.NewChatServer(deli) diff --git a/cmd/gateway_http/config.toml b/cmd/gateway_http/config.toml index f644a97..17217fb 100644 --- a/cmd/gateway_http/config.toml +++ b/cmd/gateway_http/config.toml @@ -1,5 +1,6 @@ [app] port = 7000 +aesTokenKey = "9Nz3Y6DES3msAFndz4QsJAECUIFkf+KKaRa+jRnNALk=" ignoreUrls = [ "/api/svc/auth/login", "/api/svc/auth/verify", diff --git a/cmd/gateway_http/main.go b/cmd/gateway_http/main.go index 99455e7..ab74610 100644 --- a/cmd/gateway_http/main.go +++ b/cmd/gateway_http/main.go @@ -7,17 +7,14 @@ import ( clientv3 "go.etcd.io/etcd/client/v3" "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" - "google.golang.org/grpc/grpclog" "google.golang.org/grpc/resolver" "net/http" 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" "sonet/pkg/utils/shutdown" - "sonet/pkg/utils/strs" ) type GatewayHttpConfig struct { @@ -27,8 +24,6 @@ type GatewayHttpConfig struct { } func main() { - grpclog.SetLoggerV2(logger.Logger) - appConf := &GatewayHttpConfig{} conf := config.LoadConfig(appConf, "cmd/gateway_http") @@ -41,9 +36,6 @@ func main() { shutdown.AddShutdownHook(etcdResolver.Close) resolver.Register(etcdResolver) - grpcFactory := generic.NewGpcGenericClientFactory(etcdResolver, grpc.WithTransportCredentials(insecure.NewCredentials())) - grpcFactory.Init() - // gin http server authFilter, err := config2.NewAuthFilter(appConf.AesTokenKey, appConf.IgnoreUrls) if err != nil { @@ -53,37 +45,17 @@ func main() { server := gin.Default() server.Use(authFilter.Filter) - group := server.Group("/api/svc") - group.POST("/:svc/:method", func(c *gin.Context) { - svc := strs.UpperInitialLetter(c.Param("svc")) - method := strs.UpperInitialLetter(c.Param("method")) - if svc == "" || method == "" { - c.JSON(http.StatusBadRequest, resp.Error("svc not found")) - return - } - - ctx := context.Background() - grpcClient, err := grpcFactory.GetClient(ctx, svc) - if err != nil { - c.JSON(http.StatusForbidden, resp.Error(err.Error())) - return - } - body := make(map[string]interface{}) - err = c.BindJSON(&body) - if err != nil { - c.JSON(http.StatusBadRequest, resp.Error("parse request body error: "+err.Error())) - return - } + // postal loadBalancer handler + postalBalancer := logic.NewPostalBalancer() + postalBalancer.Init(context.Background()) + server.GET("/api/lb/ws", postalBalancer.Endpoint) - res, err := grpcClient.InvokeUnaryJson(ctx, method, body) - if err != nil { - c.JSON(http.StatusInternalServerError, resp.Error(err.Error())) - return - } - // j, err := res.MarshalJSON() - // c.Render(http.StatusOK, RenderMarshaledJson{j}) - c.JSON(http.StatusOK, resp.Success(res)) - }) + // grpc services + grpcGroup := server.Group("/api/svc") + grpcFactory := generic.NewGpcGenericClientFactory(etcdResolver, grpc.WithTransportCredentials(insecure.NewCredentials())) + grpcFactory.Init() + grpcGenericHandler := logic.NewGrpcGenericHandler(grpcFactory) + grpcGenericHandler.Route(grpcGroup) go func() { err := server.Run(fmt.Sprintf(":%d", appConf.Port)) diff --git a/cmd/gateway_ws/config.toml b/cmd/gateway_ws/config.toml index 69ce25d..462af5f 100644 --- a/cmd/gateway_ws/config.toml +++ b/cmd/gateway_ws/config.toml @@ -1,18 +1,19 @@ [app] -httpPort = 7001 +httpPort = 7002 +endpointAddress = "127.0.0.1:7002" subjectCacheTopic = "wsgate:subject:" subjectLrcExpiration = "10m" subjectLrcCleanupInterval = "5m" [grpc] -address = ":7010" +address = ":7012" maxSendMsgSize = "8Mi" maxRecvMsgSize = "8Mi" readBufferSize = "8Ki" writeBufferSize = "8Ki" [grpc.register.attrs] -weight = 100 +weight = 10 [etcd] endpoints = ["124.222.131.236:3279"] diff --git a/cmd/gateway_ws/main.go b/cmd/gateway_ws/main.go index 3f296b4..8f60400 100644 --- a/cmd/gateway_ws/main.go +++ b/cmd/gateway_ws/main.go @@ -6,7 +6,6 @@ import ( clientv3 "go.etcd.io/etcd/client/v3" "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" - "google.golang.org/grpc/grpclog" "google.golang.org/grpc/resolver" "sonet/internal/gateway_ws/server" "sonet/internal/postal/logic" @@ -17,13 +16,13 @@ import ( "sonet/pkg/plugins/cache" "sonet/pkg/plugins/mq" "sonet/pkg/utils/conver" - "sonet/pkg/utils/logger" "sonet/pkg/utils/shutdown" "sync" ) type GatewayWsConfig struct { - HttpPort int + HttpPort int + EndpointAddress string SubjectCacheTopic string SubjectLrcExpiration string @@ -32,8 +31,6 @@ type GatewayWsConfig struct { // websocket server with postalService func main() { - grpclog.SetLoggerV2(logger.Logger) - appConf := &GatewayWsConfig{} conf := config.LoadConfig(appConf, "cmd/gateway_ws") @@ -68,7 +65,7 @@ func main() { // run postal server registry := discovery.NewRegister(etcdClient) shutdown.AddShutdownHook(registry.Stop) - postalServer := logic.NewPostalServer(sessionStore, subjectStore, clientFactory) + postalServer := logic.NewPostalServer(appConf.EndpointAddress, sessionStore, subjectStore, clientFactory) go func() { err = postalServer.Run(conf.Grpc, registry) if err != nil { diff --git a/cmd/mahjong/main.go b/cmd/mahjong/main.go index b79dd79..7905807 100644 --- a/cmd/mahjong/main.go +++ b/cmd/mahjong/main.go @@ -1,30 +1,5 @@ package main -import ( - "fmt" - "net" -) - func main() { - //listen, err := net.Listen("tcp", ":7878") - //addr := listen.(*net.TCPListener).Addr() - //port := addr.(*net.TCPAddr).Port - //fmt.Println(err, listen, addr, port) - // listen.(*net.TCPListener).Addr().(*net.TCPAddr).Port - - addr, err := net.ResolveTCPAddr("tcp", "192.168.1.110:7788") - fmt.Println(err) - fmt.Println(addr.Port) - if addr.IP == nil { - fmt.Println("ip is nil") - } - fmt.Println(addr.IP.String()) - - //addrPort, err := netip.ParseAddrPort(":7788") - //fmt.Println(err) - //fmt.Println(addrPort.Port()) - //fmt.Println(addrPort.Addr()) - //fmt.Println(addrPort.String()) - } diff --git a/internal/gateway_http/config/auth_filter.go b/internal/gateway_http/config/auth_filter.go index ca048fe..5a4ff9f 100644 --- a/internal/gateway_http/config/auth_filter.go +++ b/internal/gateway_http/config/auth_filter.go @@ -5,6 +5,7 @@ import ( "github.com/gin-gonic/gin" "net/http" "sonet/pkg/protocol/authorize" + "sonet/pkg/protocol/session" "sonet/pkg/utils/resp" "strings" ) @@ -53,3 +54,13 @@ func (f *AuthFilter) Filter(c *gin.Context) { c.Set(SubjectKey, subject) } + +func GetSubject(c *gin.Context) (subject *authorize.Subject, err error) { + val, ok := c.Get(SubjectKey) + if !ok { + err = session.UnauthorizedRequestError + return + } + subject = val.(*authorize.Subject) + return +} diff --git a/internal/gateway_http/logic/http_server.go b/internal/gateway_http/logic/http_server.go index 4c79103..035b294 100644 --- a/internal/gateway_http/logic/http_server.go +++ b/internal/gateway_http/logic/http_server.go @@ -1 +1,62 @@ package logic + +import ( + "context" + "github.com/gin-gonic/gin" + "net/http" + "sonet/internal/gateway_http/config" + "sonet/pkg/grpc/generic" + "sonet/pkg/protocol/session" + "sonet/pkg/utils/resp" + "sonet/pkg/utils/strs" +) + +type GrpcGenericHandler struct { + grpcFactory *generic.GrpcGenericClientFactory +} + +func NewGrpcGenericHandler(grpcFactory *generic.GrpcGenericClientFactory) *GrpcGenericHandler { + return &GrpcGenericHandler{ + grpcFactory: grpcFactory, + } +} + +func (h *GrpcGenericHandler) Route(route gin.IRoutes) { + route.POST("/:svc/:method", h.handler) +} + +func (h *GrpcGenericHandler) handler(c *gin.Context) { + ctx := context.Background() + subject, err := config.GetSubject(c) + if err == nil { + ctx = session.PutSubject(ctx, session.NewRpcSubject(subject.Uid)) + } + + svc := strs.UpperInitialLetter(c.Param("svc")) + method := strs.UpperInitialLetter(c.Param("method")) + if svc == "" || method == "" { + c.JSON(http.StatusBadRequest, resp.Error("svc not found")) + return + } + + grpcClient, err := h.grpcFactory.GetClient(ctx, svc) + if err != nil { + c.JSON(http.StatusForbidden, resp.Error(err.Error())) + return + } + body := make(map[string]interface{}) + err = c.BindJSON(&body) + if err != nil { + c.JSON(http.StatusBadRequest, resp.Error("parse request body error: "+err.Error())) + return + } + + res, err := grpcClient.InvokeUnaryJson(ctx, method, body) + if err != nil { + c.JSON(http.StatusInternalServerError, resp.Error(err.Error())) + return + } + // j, err := res.MarshalJSON() + // c.Render(http.StatusOK, RenderMarshaledJson{j}) + c.JSON(http.StatusOK, resp.Success(res)) +} diff --git a/internal/gateway_http/logic/postal_balancer.go b/internal/gateway_http/logic/postal_balancer.go new file mode 100644 index 0000000..b426443 --- /dev/null +++ b/internal/gateway_http/logic/postal_balancer.go @@ -0,0 +1,54 @@ +package logic + +import ( + "context" + "fmt" + "github.com/gin-gonic/gin" + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/protobuf/types/known/emptypb" + "net/http" + "sonet/api/gen/postal" + "sonet/internal/gateway_http/config" + "sonet/pkg/grpc/balancer" + "sonet/pkg/grpc/discovery" + "sonet/pkg/utils/resp" +) + +type PostalBalancer struct { + postalClient postal.PostalClient +} + +func NewPostalBalancer() *PostalBalancer { + return &PostalBalancer{} +} + +func (h *PostalBalancer) Init(ctx context.Context) { + balancer.InitConsistentHashBuilder() + + postalUrl := discovery.BuildResolverUrl(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)), + ) + if err != nil { + return + } + h.postalClient = postal.NewPostalClient(conn) +} + +func (h *PostalBalancer) Endpoint(c *gin.Context) { + subject, err := config.GetSubject(c) + if err != nil { + c.JSON(http.StatusBadRequest, resp.Fail(err.Error())) + return + } + ctx := context.WithValue(context.Background(), balancer.ConsistentHashKey, subject.Uid) + res, err := h.postalClient.Endpoint(ctx, &emptypb.Empty{}) + if err != nil { + c.JSON(http.StatusInternalServerError, resp.Error(err.Error())) + return + } + c.JSON(http.StatusOK, resp.Success(res.Endpoint)) +} diff --git a/internal/gateway_http/logic/postal_monitor.go b/internal/gateway_http/logic/postal_monitor.go new file mode 100644 index 0000000..c70323b --- /dev/null +++ b/internal/gateway_http/logic/postal_monitor.go @@ -0,0 +1,80 @@ +package logic + +import ( + "context" + "go.etcd.io/etcd/api/v3/mvccpb" + clientv3 "go.etcd.io/etcd/client/v3" + "sonet/api/gen/postal" + "sonet/pkg/grpc/discovery" + "sonet/pkg/utils/logger" +) + +type PostalMonitor struct { + client *clientv3.Client + keyPrefix string +} + +func NewPostalMonitor(client *clientv3.Client) *PostalMonitor { + return &PostalMonitor{ + client: client, + } +} + +func (m *PostalMonitor) Init(ctx context.Context) (err error) { + m.keyPrefix = discovery.BuildPrefix(discovery.Server{Name: postal.Postal_ServiceDesc.ServiceName}) + go m.watch(ctx) + m.build(ctx) + return +} + +func (m *PostalMonitor) Next() { + +} + +func (m *PostalMonitor) build(ctx context.Context) { + res, err := m.client.Get(ctx, m.keyPrefix, clientv3.WithPrefix()) + if err != nil { + return + } + //for _, kv := range res.Kvs { + // + //} + m.update(res.Kvs) + +} + +func (m *PostalMonitor) watch(ctx context.Context) { + w := m.client.Watch(ctx, m.keyPrefix, clientv3.WithPrefix()) + cancelCh := ctx.Done() + + for { + select { + case <-cancelCh: + return + case res := <-w: + if err := res.Err(); err != nil { + logger.Errorf("watch etcd instance error: %v\n", err) + continue + } + rebuild := false + eLoop: + for _, event := range res.Events { + switch event.Type { + case clientv3.EventTypePut: + fallthrough + case clientv3.EventTypeDelete: + rebuild = true + break eLoop + } + } + if rebuild { + go m.build(ctx) + } + } + } + +} + +func (m *PostalMonitor) update(value []*mvccpb.KeyValue) { + +} diff --git a/internal/postal/logic/postal_server.go b/internal/postal/logic/postal_server.go index 4074440..97214b0 100644 --- a/internal/postal/logic/postal_server.go +++ b/internal/postal/logic/postal_server.go @@ -19,19 +19,23 @@ import ( type PostalServer struct { postal.UnimplementedPostalServer + endpointAddress string // websocket 前端连接地址 broadcastAddress string // postal server 在注册中心注册的地址, 其他服务可直连访问 sessionStore *sync.Map // 在线用户conn存储 subjectStore cache.MultiLevelCache // online subject address cache, ws gateway集群所有在线用户存储 clientFactory *client.GrpcDirectClientFactory } -func NewPostalServer(sessionStore *sync.Map, +func NewPostalServer( + endpointAddress string, + sessionStore *sync.Map, subjectStore cache.MultiLevelCache, clientFactory *client.GrpcDirectClientFactory) *PostalServer { return &PostalServer{ - sessionStore: sessionStore, - subjectStore: subjectStore, - clientFactory: clientFactory, + endpointAddress: endpointAddress, + sessionStore: sessionStore, + subjectStore: subjectStore, + clientFactory: clientFactory, } } @@ -238,3 +242,10 @@ func (s *PostalServer) GroupLeave(ctx context.Context, leave *postal.ReqGroupLea func (s *PostalServer) GroupDissolve(ctx context.Context, dissolve *postal.ReqGroupDissolve) (*emptypb.Empty, error) { return nil, nil } + +func (s *PostalServer) Endpoint(context.Context, *emptypb.Empty) (*postal.ResEndpoint, error) { + res := &postal.ResEndpoint{ + Endpoint: s.endpointAddress, + } + return res, nil +} diff --git a/pkg/config/loader.go b/pkg/config/loader.go index 5bf8eba..d972334 100644 --- a/pkg/config/loader.go +++ b/pkg/config/loader.go @@ -23,6 +23,8 @@ func parseConfPathFlag(confPath string) (filePath, fileName, confName, confType // appConf service custom config // return common service configuration func LoadConfig(appConf any, confPathArg ...string) *Configuration { + initLogger() + confPath := "" confName := "config" confType := "toml" @@ -45,7 +47,7 @@ func LoadConfig(appConf any, confPathArg ...string) *Configuration { if confPath == "" { confPath = "./" } - logger.Infof("use config file: %s%s.%s, env prefix=%s\n", confPath, confName, confType, envPrefix) + logger.Infof("use config file: %s/%s.%s, env prefix=%s\n", confPath, confName, confType, envPrefix) viper.AddConfigPath(confPath) viper.SetConfigName(confName) diff --git a/pkg/config/logger.go b/pkg/config/logger.go new file mode 100644 index 0000000..a4d38e2 --- /dev/null +++ b/pkg/config/logger.go @@ -0,0 +1,16 @@ +package config + +import ( + "github.com/sirupsen/logrus" + "google.golang.org/grpc/grpclog" + "sonet/pkg/utils/logger" +) + +func initLogger() { + logrus.SetFormatter(&logrus.TextFormatter{ + ForceColors: true, + TimestampFormat: "2006-01-02 15:04:05", //时间格式 + FullTimestamp: true, + }) + grpclog.SetLoggerV2(logger.Logger) +} diff --git a/pkg/grpc/balancer/consistent_hash.go b/pkg/grpc/balancer/consistent_hash.go new file mode 100644 index 0000000..28a939d --- /dev/null +++ b/pkg/grpc/balancer/consistent_hash.go @@ -0,0 +1,98 @@ +package balancer + +import ( + "errors" + "fmt" + "google.golang.org/grpc/balancer" + "google.golang.org/grpc/balancer/base" + "google.golang.org/grpc/grpclog" + "google.golang.org/grpc/resolver" + "strconv" +) + +const ConsistentHash = "consistent_hash_x" + +var ConsistentHashKey = "consistent-hash" + +func InitConsistentHashBuilder() { + balancer.Register(newConsistentHashBuilder()) +} + +// newConsistentHashBuilder creates a new ConsistentHash balancer builder. +func newConsistentHashBuilder() balancer.Builder { + return base.NewBalancerBuilder( + ConsistentHash, + &consistentHashPickerBuilder{}, + base.Config{HealthCheck: true}, + ) +} + +type consistentHashPickerBuilder struct{} + +func (b *consistentHashPickerBuilder) Build(buildInfo base.PickerBuildInfo) balancer.Picker { + grpclog.Infof("consistentHashPicker: newPicker called with buildInfo: %v", buildInfo) + if len(buildInfo.ReadySCs) == 0 { + return base.NewErrPicker(balancer.ErrNoSubConnAvailable) + } + + picker := &consistentHashPicker{ + subConns: make(map[string]balancer.SubConn), + hash: NewKetama(DefaultReplicas, nil), + } + + for sc, conInfo := range buildInfo.ReadySCs { + weight := GetWeight(conInfo.Address) + for i := 0; i < weight; i++ { + node := wrapAddr(conInfo.Address.Addr, i) + picker.hash.Add(node) + picker.subConns[node] = sc + } + } + return picker +} + +type consistentHashPicker struct { + subConns map[string]balancer.SubConn + hash *Ketama +} + +func (p *consistentHashPicker) Pick(info balancer.PickInfo) (ret balancer.PickResult, err error) { + key, ok := info.Ctx.Value(ConsistentHashKey).(string) + if !ok || key == "" { + //key = strconv.Itoa(rand.Intn(65536)) + //grpclog.Warning("empty consistent hash key") + panic(errors.New("empty consistent hash key")) + } + targetAddr, ok := p.hash.Get(key) + if ok { + ret.SubConn = p.subConns[targetAddr] + } + return +} + +func wrapAddr(addr string, idx int) string { + return fmt.Sprintf("%s-%d", addr, idx) +} + +func GetWeight(addr resolver.Address) (weight int) { + weight = DefaultWeight + if addr.Attributes == nil { + return + } + + val := addr.Attributes.Value(WeightKey) + switch val.(type) { + case int: + weight = val.(int) + case string: + w, err := strconv.Atoi(val.(string)) + if err != nil { + grpclog.Errorf("instance weight format error: %v\n", val) + return + } + weight = w + default: + grpclog.Errorf("instance weight value type not string: %v\n", val) + } + return +} diff --git a/pkg/grpc/balancer/consistent_ketama.go b/pkg/grpc/balancer/consistent_ketama.go new file mode 100644 index 0000000..8d5d08a --- /dev/null +++ b/pkg/grpc/balancer/consistent_ketama.go @@ -0,0 +1,136 @@ +package balancer + +import ( + "hash/fnv" + "sort" + "strconv" + "sync" +) + +type HashFunc func(data []byte) uint32 + +var ( + DefaultReplicas = 10 + Salt = "this_is_salt" +) + +func DefaultHash(data []byte) uint32 { + f := fnv.New32() + _, err := f.Write(data) + if err != nil { + panic(err) + } + return f.Sum32() +} + +type Ketama struct { + sync.Mutex + hash HashFunc + replicas int + keys []int // Sorted keys + hashMap map[int]string +} + +func NewKetama(replicas int, fn HashFunc) *Ketama { + h := &Ketama{ + replicas: replicas, + hash: fn, + hashMap: make(map[int]string), + } + if h.replicas <= 0 { + h.replicas = DefaultReplicas + } + if h.hash == nil { + h.hash = DefaultHash + } + return h +} + +func (h *Ketama) IsEmpty() bool { + h.Lock() + defer h.Unlock() + + return len(h.keys) == 0 +} + +func (h *Ketama) Add(nodes ...string) { + h.Lock() + defer h.Unlock() + + for _, node := range nodes { + for i := 0; i < h.replicas; i++ { + key := int(h.hash([]byte(strconv.Itoa(i) + node + Salt))) + + if _, ok := h.hashMap[key]; !ok { + h.keys = append(h.keys, key) + } + h.hashMap[key] = node + } + } + sort.Ints(h.keys) +} + +func (h *Ketama) Remove(nodes ...string) { + h.Lock() + defer h.Unlock() + + deletedKey := make([]int, 0) + for _, node := range nodes { + for i := 0; i < h.replicas; i++ { + key := int(h.hash([]byte(strconv.Itoa(i) + node + Salt))) + + if _, ok := h.hashMap[key]; ok { + deletedKey = append(deletedKey, key) + delete(h.hashMap, key) + } + } + } + if len(deletedKey) > 0 { + h.deleteKeys(deletedKey) + } +} + +func (h *Ketama) deleteKeys(deletedKeys []int) { + sort.Ints(deletedKeys) + + index := 0 + count := 0 + for _, key := range deletedKeys { + for ; index < len(h.keys); index++ { + h.keys[index-count] = h.keys[index] + + if key == h.keys[index] { + count++ + index++ + break + } + } + } + + for ; index < len(h.keys); index++ { + h.keys[index-count] = h.keys[index] + } + + h.keys = h.keys[:len(h.keys)-count] +} + +func (h *Ketama) Get(key string) (string, bool) { + if h.IsEmpty() { + return "", false + } + + hash := int(h.hash([]byte(key + Salt))) + + h.Lock() + defer h.Unlock() + + idx := sort.Search(len(h.keys), func(i int) bool { + return h.keys[i] >= hash + }) + + if idx == len(h.keys) { + idx = 0 + } + str, ok := h.hashMap[h.keys[idx]] + return str, ok +} diff --git a/pkg/grpc/balancer/properties.go b/pkg/grpc/balancer/properties.go new file mode 100644 index 0000000..17b6979 --- /dev/null +++ b/pkg/grpc/balancer/properties.go @@ -0,0 +1,6 @@ +package balancer + +const ( + WeightKey = "weight" + DefaultWeight = 10 +) diff --git a/pkg/grpc/discovery/register.go b/pkg/grpc/discovery/register.go index e52a618..cc07778 100644 --- a/pkg/grpc/discovery/register.go +++ b/pkg/grpc/discovery/register.go @@ -14,7 +14,7 @@ import ( clientv3 "go.etcd.io/etcd/client/v3" ) -var DefaultRegisterTTL int64 = 10 +var DefaultRegisterTTL int64 = 30 func RegisterAddress(addr string) (fullAddr string, err error) { tcpAddr, err := net.ResolveTCPAddr("tcp", addr) diff --git a/pkg/grpc/discovery/resolver.go b/pkg/grpc/discovery/resolver.go index e39b3fa..1e68d16 100644 --- a/pkg/grpc/discovery/resolver.go +++ b/pkg/grpc/discovery/resolver.go @@ -21,7 +21,7 @@ type Resolver struct { closeCh chan struct{} watchCh clientv3.WatchChan cli *clientv3.Client - keyPrifix string + keyPrefix string srvAddrsList []resolver.Address cc resolver.ClientConn @@ -44,7 +44,7 @@ func (r *Resolver) Scheme() string { // Build creates a new resolver.Resolver for the given target func (r *Resolver) Build(target resolver.Target, cc resolver.ClientConn, opts resolver.BuildOptions) (rr resolver.Resolver, err error) { r.cc = cc - r.keyPrifix = BuildPrefix(Server{Name: target.Endpoint()}) + r.keyPrefix = BuildPrefix(Server{Name: target.Endpoint()}) if err = r.start(); err != nil { return nil, err } @@ -85,7 +85,7 @@ func (r *Resolver) start() error { // watch update events func (r *Resolver) watch() { ticker := time.NewTicker(time.Minute) - r.watchCh = r.cli.Watch(context.Background(), r.keyPrifix, clientv3.WithPrefix()) + r.watchCh = r.cli.Watch(context.Background(), r.keyPrefix, clientv3.WithPrefix()) for { select { @@ -148,7 +148,7 @@ func (r *Resolver) update(events []*clientv3.Event) { func (r *Resolver) sync() (err error) { ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) defer cancel() - res, err := r.cli.Get(ctx, r.keyPrifix, clientv3.WithPrefix()) + res, err := r.cli.Get(ctx, r.keyPrefix, clientv3.WithPrefix()) if err != nil { return } diff --git a/pkg/grpc/meta/meta.go b/pkg/grpc/meta/meta.go new file mode 100644 index 0000000..e754f86 --- /dev/null +++ b/pkg/grpc/meta/meta.go @@ -0,0 +1,18 @@ +package meta + +import ( + "context" + "google.golang.org/grpc/metadata" +) + +func GetUid(ctx context.Context, key string) (string, bool) { + md, ok := metadata.FromIncomingContext(ctx) + if !ok { + return "", false + } + vals := md.Get(key) + if len(vals) == 0 { + return "", false + } + return vals[len(vals)-1], true +} diff --git a/pkg/protocol/deliver/deliver.go b/pkg/protocol/deliver/deliver.go index 2acea4b..9b2c642 100644 --- a/pkg/protocol/deliver/deliver.go +++ b/pkg/protocol/deliver/deliver.go @@ -8,6 +8,7 @@ import ( "google.golang.org/grpc/credentials/insecure" "reflect" "sonet/api/gen/postal" + "sonet/pkg/grpc/balancer" "sonet/pkg/grpc/discovery" "sonet/pkg/utils/logger" "time" @@ -33,12 +34,14 @@ func NewDeliver(msgInServiceName string) *Deliver { } } -func (d *Deliver) InitWithResolver(ctx context.Context, resolver *discovery.Resolver) (err error) { - addr := fmt.Sprintf("%s:///%s", resolver.Scheme(), postal.Postal_ServiceDesc.ServiceName) - conn, err := grpc.DialContext(ctx, addr, +func (d *Deliver) InitWithResolver(ctx context.Context) (err error) { + balancer.InitConsistentHashBuilder() + + postalUrl := discovery.BuildResolverUrl(postal.Postal_ServiceDesc.ServiceName) + conn, err := grpc.DialContext(ctx, postalUrl, grpc.WithTransportCredentials(insecure.NewCredentials()), - // todo consistent hash lb - // grpc.WithDefaultServiceConfig(fmt.Sprintf(`{"loadBalancingPolicy":"%s"}`, balancer.ConsistentHash)), + // consistent hash lb + grpc.WithDefaultServiceConfig(fmt.Sprintf(`{"loadBalancingPolicy":"%s"}`, balancer.ConsistentHash)), ) if err != nil { return @@ -96,6 +99,8 @@ func (d *Deliver) deliver0(ctx context.Context, msg proto.Message, receivers []s Receiver: receivers[0], Msg: message, } + + ctx = context.WithValue(ctx, balancer.ConsistentHashKey, reqDeliver.Receiver) res, err := d.postal.Deliver(ctx, reqDeliver) if err != nil { return StatusError, err @@ -110,6 +115,7 @@ func (d *Deliver) deliver0(ctx context.Context, msg proto.Message, receivers []s } else { // deliver batch receiver + ctx = context.WithValue(ctx, balancer.ConsistentHashKey, receivers[0]) req := &postal.ReqDeliverBatch{Receivers: receivers, Msg: message} res, err := d.postal.DeliverBatch(ctx, req) if err != nil { diff --git a/pkg/protocol/protocol_test.go b/pkg/protocol/protocol_test.go index 8619b0f..6352476 100644 --- a/pkg/protocol/protocol_test.go +++ b/pkg/protocol/protocol_test.go @@ -6,7 +6,7 @@ import ( ) func TestProtocolCodec(t *testing.T) { - header := Header{ + header := &Header{ Magic: Magic, Type: 1, Status: 20, diff --git a/pkg/utils/logger/logger.go b/pkg/utils/logger/logger.go index 01ee9f1..c9660e4 100644 --- a/pkg/utils/logger/logger.go +++ b/pkg/utils/logger/logger.go @@ -7,7 +7,7 @@ import ( var Logger *SoLogger func init() { - Logger = &SoLogger{logger: logrus.New()} + Logger = &SoLogger{logger: logrus.StandardLogger()} } type SoLogger struct {