Browse Source

postal consistent hash balancer

master
tangmingyou 3 years ago
parent
commit
f8b57c3e4b
  1. 8
      api/postal.proto
  2. 3
      cmd/auth/main.go
  3. 5
      cmd/chat/main.go
  4. 1
      cmd/gateway_http/config.toml
  5. 50
      cmd/gateway_http/main.go
  6. 7
      cmd/gateway_ws/config.toml
  7. 9
      cmd/gateway_ws/main.go
  8. 25
      cmd/mahjong/main.go
  9. 11
      internal/gateway_http/config/auth_filter.go
  10. 61
      internal/gateway_http/logic/http_server.go
  11. 54
      internal/gateway_http/logic/postal_balancer.go
  12. 80
      internal/gateway_http/logic/postal_monitor.go
  13. 19
      internal/postal/logic/postal_server.go
  14. 4
      pkg/config/loader.go
  15. 16
      pkg/config/logger.go
  16. 98
      pkg/grpc/balancer/consistent_hash.go
  17. 136
      pkg/grpc/balancer/consistent_ketama.go
  18. 6
      pkg/grpc/balancer/properties.go
  19. 2
      pkg/grpc/discovery/register.go
  20. 8
      pkg/grpc/discovery/resolver.go
  21. 18
      pkg/grpc/meta/meta.go
  22. 16
      pkg/protocol/deliver/deliver.go
  23. 2
      pkg/protocol/protocol_test.go
  24. 2
      pkg/utils/logger/logger.go

8
api/postal.proto

@ -17,6 +17,9 @@ service Postal {
rpc GroupJoin(ReqGroupJoin) returns(google.protobuf.Empty); rpc GroupJoin(ReqGroupJoin) returns(google.protobuf.Empty);
rpc GroupLeave(ReqGroupLeave) returns(google.protobuf.Empty); rpc GroupLeave(ReqGroupLeave) returns(google.protobuf.Empty);
rpc GroupDissolve(ReqGroupDissolve) returns(google.protobuf.Empty); rpc GroupDissolve(ReqGroupDissolve) returns(google.protobuf.Empty);
// , websocket addr
rpc Endpoint(google.protobuf.Empty) returns(ResEndpoint);
} }
message Message { message Message {
@ -90,3 +93,8 @@ message ReqRedirect {
ReqDeliverBatch deliverBatch = 11; ReqDeliverBatch deliverBatch = 11;
ReqDeliverGroup deliverGroup = 12; ReqDeliverGroup deliverGroup = 12;
} }
message ResEndpoint {
string endpoint = 1;
map<string, string> extra = 2;
}

3
cmd/auth/main.go

@ -2,12 +2,10 @@ package main
import ( import (
clientv3 "go.etcd.io/etcd/client/v3" clientv3 "go.etcd.io/etcd/client/v3"
"google.golang.org/grpc/grpclog"
"sonet/internal/auth/data" "sonet/internal/auth/data"
"sonet/internal/auth/logic" "sonet/internal/auth/logic"
"sonet/pkg/config" "sonet/pkg/config"
"sonet/pkg/grpc/discovery" "sonet/pkg/grpc/discovery"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/shutdown" "sonet/pkg/utils/shutdown"
) )
@ -16,7 +14,6 @@ type AuthConfig struct {
} }
func main() { func main() {
grpclog.SetLoggerV2(logger.Logger)
appConf := &AuthConfig{} appConf := &AuthConfig{}
conf := config.LoadConfig(appConf, "cmd/auth") conf := config.LoadConfig(appConf, "cmd/auth")

5
cmd/chat/main.go

@ -3,19 +3,16 @@ package main
import ( import (
"context" "context"
clientv3 "go.etcd.io/etcd/client/v3" clientv3 "go.etcd.io/etcd/client/v3"
"google.golang.org/grpc/grpclog"
"google.golang.org/grpc/resolver" "google.golang.org/grpc/resolver"
"sonet/api/gen/chat" "sonet/api/gen/chat"
"sonet/internal/chat/logic" "sonet/internal/chat/logic"
"sonet/pkg/config" "sonet/pkg/config"
"sonet/pkg/grpc/discovery" "sonet/pkg/grpc/discovery"
"sonet/pkg/protocol/deliver" "sonet/pkg/protocol/deliver"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/shutdown" "sonet/pkg/utils/shutdown"
) )
func main() { func main() {
grpclog.SetLoggerV2(logger.Logger)
conf := config.LoadConfig(nil, "cmd/chat") conf := config.LoadConfig(nil, "cmd/chat")
// registry and run... // registry and run...
@ -31,7 +28,7 @@ func main() {
resolver.Register(etcdResolver) resolver.Register(etcdResolver)
deli := deliver.NewDeliver(chat.Chat_ServiceDesc.ServiceName) 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) panic(err)
} }
chatServer := logic.NewChatServer(deli) chatServer := logic.NewChatServer(deli)

1
cmd/gateway_http/config.toml

@ -1,5 +1,6 @@
[app] [app]
port = 7000 port = 7000
aesTokenKey = "9Nz3Y6DES3msAFndz4QsJAECUIFkf+KKaRa+jRnNALk="
ignoreUrls = [ ignoreUrls = [
"/api/svc/auth/login", "/api/svc/auth/login",
"/api/svc/auth/verify", "/api/svc/auth/verify",

50
cmd/gateway_http/main.go

@ -7,17 +7,14 @@ import (
clientv3 "go.etcd.io/etcd/client/v3" clientv3 "go.etcd.io/etcd/client/v3"
"google.golang.org/grpc" "google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure" "google.golang.org/grpc/credentials/insecure"
"google.golang.org/grpc/grpclog"
"google.golang.org/grpc/resolver" "google.golang.org/grpc/resolver"
"net/http" "net/http"
config2 "sonet/internal/gateway_http/config" config2 "sonet/internal/gateway_http/config"
"sonet/internal/gateway_http/logic"
"sonet/pkg/config" "sonet/pkg/config"
"sonet/pkg/grpc/discovery" "sonet/pkg/grpc/discovery"
"sonet/pkg/grpc/generic" "sonet/pkg/grpc/generic"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/resp"
"sonet/pkg/utils/shutdown" "sonet/pkg/utils/shutdown"
"sonet/pkg/utils/strs"
) )
type GatewayHttpConfig struct { type GatewayHttpConfig struct {
@ -27,8 +24,6 @@ type GatewayHttpConfig struct {
} }
func main() { func main() {
grpclog.SetLoggerV2(logger.Logger)
appConf := &GatewayHttpConfig{} appConf := &GatewayHttpConfig{}
conf := config.LoadConfig(appConf, "cmd/gateway_http") conf := config.LoadConfig(appConf, "cmd/gateway_http")
@ -41,9 +36,6 @@ func main() {
shutdown.AddShutdownHook(etcdResolver.Close) shutdown.AddShutdownHook(etcdResolver.Close)
resolver.Register(etcdResolver) resolver.Register(etcdResolver)
grpcFactory := generic.NewGpcGenericClientFactory(etcdResolver, grpc.WithTransportCredentials(insecure.NewCredentials()))
grpcFactory.Init()
// gin http server // gin http server
authFilter, err := config2.NewAuthFilter(appConf.AesTokenKey, appConf.IgnoreUrls) authFilter, err := config2.NewAuthFilter(appConf.AesTokenKey, appConf.IgnoreUrls)
if err != nil { if err != nil {
@ -53,37 +45,17 @@ func main() {
server := gin.Default() server := gin.Default()
server.Use(authFilter.Filter) server.Use(authFilter.Filter)
group := server.Group("/api/svc") // postal loadBalancer handler
group.POST("/:svc/:method", func(c *gin.Context) { postalBalancer := logic.NewPostalBalancer()
svc := strs.UpperInitialLetter(c.Param("svc")) postalBalancer.Init(context.Background())
method := strs.UpperInitialLetter(c.Param("method")) server.GET("/api/lb/ws", postalBalancer.Endpoint)
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
}
res, err := grpcClient.InvokeUnaryJson(ctx, method, body) // grpc services
if err != nil { grpcGroup := server.Group("/api/svc")
c.JSON(http.StatusInternalServerError, resp.Error(err.Error())) grpcFactory := generic.NewGpcGenericClientFactory(etcdResolver, grpc.WithTransportCredentials(insecure.NewCredentials()))
return grpcFactory.Init()
} grpcGenericHandler := logic.NewGrpcGenericHandler(grpcFactory)
// j, err := res.MarshalJSON() grpcGenericHandler.Route(grpcGroup)
// c.Render(http.StatusOK, RenderMarshaledJson{j})
c.JSON(http.StatusOK, resp.Success(res))
})
go func() { go func() {
err := server.Run(fmt.Sprintf(":%d", appConf.Port)) err := server.Run(fmt.Sprintf(":%d", appConf.Port))

7
cmd/gateway_ws/config.toml

@ -1,18 +1,19 @@
[app] [app]
httpPort = 7001 httpPort = 7002
endpointAddress = "127.0.0.1:7002"
subjectCacheTopic = "wsgate:subject:" subjectCacheTopic = "wsgate:subject:"
subjectLrcExpiration = "10m" subjectLrcExpiration = "10m"
subjectLrcCleanupInterval = "5m" subjectLrcCleanupInterval = "5m"
[grpc] [grpc]
address = ":7010" address = ":7012"
maxSendMsgSize = "8Mi" maxSendMsgSize = "8Mi"
maxRecvMsgSize = "8Mi" maxRecvMsgSize = "8Mi"
readBufferSize = "8Ki" readBufferSize = "8Ki"
writeBufferSize = "8Ki" writeBufferSize = "8Ki"
[grpc.register.attrs] [grpc.register.attrs]
weight = 100 weight = 10
[etcd] [etcd]
endpoints = ["124.222.131.236:3279"] endpoints = ["124.222.131.236:3279"]

9
cmd/gateway_ws/main.go

@ -6,7 +6,6 @@ import (
clientv3 "go.etcd.io/etcd/client/v3" clientv3 "go.etcd.io/etcd/client/v3"
"google.golang.org/grpc" "google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure" "google.golang.org/grpc/credentials/insecure"
"google.golang.org/grpc/grpclog"
"google.golang.org/grpc/resolver" "google.golang.org/grpc/resolver"
"sonet/internal/gateway_ws/server" "sonet/internal/gateway_ws/server"
"sonet/internal/postal/logic" "sonet/internal/postal/logic"
@ -17,13 +16,13 @@ import (
"sonet/pkg/plugins/cache" "sonet/pkg/plugins/cache"
"sonet/pkg/plugins/mq" "sonet/pkg/plugins/mq"
"sonet/pkg/utils/conver" "sonet/pkg/utils/conver"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/shutdown" "sonet/pkg/utils/shutdown"
"sync" "sync"
) )
type GatewayWsConfig struct { type GatewayWsConfig struct {
HttpPort int HttpPort int
EndpointAddress string
SubjectCacheTopic string SubjectCacheTopic string
SubjectLrcExpiration string SubjectLrcExpiration string
@ -32,8 +31,6 @@ type GatewayWsConfig struct {
// websocket server with postalService // websocket server with postalService
func main() { func main() {
grpclog.SetLoggerV2(logger.Logger)
appConf := &GatewayWsConfig{} appConf := &GatewayWsConfig{}
conf := config.LoadConfig(appConf, "cmd/gateway_ws") conf := config.LoadConfig(appConf, "cmd/gateway_ws")
@ -68,7 +65,7 @@ func main() {
// run postal server // run postal server
registry := discovery.NewRegister(etcdClient) registry := discovery.NewRegister(etcdClient)
shutdown.AddShutdownHook(registry.Stop) shutdown.AddShutdownHook(registry.Stop)
postalServer := logic.NewPostalServer(sessionStore, subjectStore, clientFactory) postalServer := logic.NewPostalServer(appConf.EndpointAddress, sessionStore, subjectStore, clientFactory)
go func() { go func() {
err = postalServer.Run(conf.Grpc, registry) err = postalServer.Run(conf.Grpc, registry)
if err != nil { if err != nil {

25
cmd/mahjong/main.go

@ -1,30 +1,5 @@
package main package main
import (
"fmt"
"net"
)
func main() { 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())
} }

11
internal/gateway_http/config/auth_filter.go

@ -5,6 +5,7 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"net/http" "net/http"
"sonet/pkg/protocol/authorize" "sonet/pkg/protocol/authorize"
"sonet/pkg/protocol/session"
"sonet/pkg/utils/resp" "sonet/pkg/utils/resp"
"strings" "strings"
) )
@ -53,3 +54,13 @@ func (f *AuthFilter) Filter(c *gin.Context) {
c.Set(SubjectKey, subject) 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
}

61
internal/gateway_http/logic/http_server.go

@ -1 +1,62 @@
package logic 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))
}

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

80
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) {
}

19
internal/postal/logic/postal_server.go

@ -19,19 +19,23 @@ import (
type PostalServer struct { type PostalServer struct {
postal.UnimplementedPostalServer postal.UnimplementedPostalServer
endpointAddress string // websocket 前端连接地址
broadcastAddress string // postal server 在注册中心注册的地址, 其他服务可直连访问 broadcastAddress string // postal server 在注册中心注册的地址, 其他服务可直连访问
sessionStore *sync.Map // 在线用户conn存储 sessionStore *sync.Map // 在线用户conn存储
subjectStore cache.MultiLevelCache // online subject address cache, ws gateway集群所有在线用户存储 subjectStore cache.MultiLevelCache // online subject address cache, ws gateway集群所有在线用户存储
clientFactory *client.GrpcDirectClientFactory clientFactory *client.GrpcDirectClientFactory
} }
func NewPostalServer(sessionStore *sync.Map, func NewPostalServer(
endpointAddress string,
sessionStore *sync.Map,
subjectStore cache.MultiLevelCache, subjectStore cache.MultiLevelCache,
clientFactory *client.GrpcDirectClientFactory) *PostalServer { clientFactory *client.GrpcDirectClientFactory) *PostalServer {
return &PostalServer{ return &PostalServer{
sessionStore: sessionStore, endpointAddress: endpointAddress,
subjectStore: subjectStore, sessionStore: sessionStore,
clientFactory: clientFactory, 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) { func (s *PostalServer) GroupDissolve(ctx context.Context, dissolve *postal.ReqGroupDissolve) (*emptypb.Empty, error) {
return nil, nil return nil, nil
} }
func (s *PostalServer) Endpoint(context.Context, *emptypb.Empty) (*postal.ResEndpoint, error) {
res := &postal.ResEndpoint{
Endpoint: s.endpointAddress,
}
return res, nil
}

4
pkg/config/loader.go

@ -23,6 +23,8 @@ func parseConfPathFlag(confPath string) (filePath, fileName, confName, confType
// appConf service custom config // appConf service custom config
// return common service configuration // return common service configuration
func LoadConfig(appConf any, confPathArg ...string) *Configuration { func LoadConfig(appConf any, confPathArg ...string) *Configuration {
initLogger()
confPath := "" confPath := ""
confName := "config" confName := "config"
confType := "toml" confType := "toml"
@ -45,7 +47,7 @@ func LoadConfig(appConf any, confPathArg ...string) *Configuration {
if confPath == "" { if confPath == "" {
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.AddConfigPath(confPath)
viper.SetConfigName(confName) viper.SetConfigName(confName)

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

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

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

6
pkg/grpc/balancer/properties.go

@ -0,0 +1,6 @@
package balancer
const (
WeightKey = "weight"
DefaultWeight = 10
)

2
pkg/grpc/discovery/register.go

@ -14,7 +14,7 @@ import (
clientv3 "go.etcd.io/etcd/client/v3" clientv3 "go.etcd.io/etcd/client/v3"
) )
var DefaultRegisterTTL int64 = 10 var DefaultRegisterTTL int64 = 30
func RegisterAddress(addr string) (fullAddr string, err error) { func RegisterAddress(addr string) (fullAddr string, err error) {
tcpAddr, err := net.ResolveTCPAddr("tcp", addr) tcpAddr, err := net.ResolveTCPAddr("tcp", addr)

8
pkg/grpc/discovery/resolver.go

@ -21,7 +21,7 @@ type Resolver struct {
closeCh chan struct{} closeCh chan struct{}
watchCh clientv3.WatchChan watchCh clientv3.WatchChan
cli *clientv3.Client cli *clientv3.Client
keyPrifix string keyPrefix string
srvAddrsList []resolver.Address srvAddrsList []resolver.Address
cc resolver.ClientConn cc resolver.ClientConn
@ -44,7 +44,7 @@ func (r *Resolver) Scheme() string {
// Build creates a new resolver.Resolver for the given target // 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) { func (r *Resolver) Build(target resolver.Target, cc resolver.ClientConn, opts resolver.BuildOptions) (rr resolver.Resolver, err error) {
r.cc = cc r.cc = cc
r.keyPrifix = BuildPrefix(Server{Name: target.Endpoint()}) r.keyPrefix = BuildPrefix(Server{Name: target.Endpoint()})
if err = r.start(); err != nil { if err = r.start(); err != nil {
return nil, err return nil, err
} }
@ -85,7 +85,7 @@ func (r *Resolver) start() error {
// watch update events // watch update events
func (r *Resolver) watch() { func (r *Resolver) watch() {
ticker := time.NewTicker(time.Minute) 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 { for {
select { select {
@ -148,7 +148,7 @@ func (r *Resolver) update(events []*clientv3.Event) {
func (r *Resolver) sync() (err error) { func (r *Resolver) sync() (err error) {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel() 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 { if err != nil {
return return
} }

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

16
pkg/protocol/deliver/deliver.go

@ -8,6 +8,7 @@ import (
"google.golang.org/grpc/credentials/insecure" "google.golang.org/grpc/credentials/insecure"
"reflect" "reflect"
"sonet/api/gen/postal" "sonet/api/gen/postal"
"sonet/pkg/grpc/balancer"
"sonet/pkg/grpc/discovery" "sonet/pkg/grpc/discovery"
"sonet/pkg/utils/logger" "sonet/pkg/utils/logger"
"time" "time"
@ -33,12 +34,14 @@ func NewDeliver(msgInServiceName string) *Deliver {
} }
} }
func (d *Deliver) InitWithResolver(ctx context.Context, resolver *discovery.Resolver) (err error) { func (d *Deliver) InitWithResolver(ctx context.Context) (err error) {
addr := fmt.Sprintf("%s:///%s", resolver.Scheme(), postal.Postal_ServiceDesc.ServiceName) balancer.InitConsistentHashBuilder()
conn, err := grpc.DialContext(ctx, addr,
postalUrl := discovery.BuildResolverUrl(postal.Postal_ServiceDesc.ServiceName)
conn, err := grpc.DialContext(ctx, postalUrl,
grpc.WithTransportCredentials(insecure.NewCredentials()), grpc.WithTransportCredentials(insecure.NewCredentials()),
// todo consistent hash lb // consistent hash lb
// grpc.WithDefaultServiceConfig(fmt.Sprintf(`{"loadBalancingPolicy":"%s"}`, balancer.ConsistentHash)), grpc.WithDefaultServiceConfig(fmt.Sprintf(`{"loadBalancingPolicy":"%s"}`, balancer.ConsistentHash)),
) )
if err != nil { if err != nil {
return return
@ -96,6 +99,8 @@ func (d *Deliver) deliver0(ctx context.Context, msg proto.Message, receivers []s
Receiver: receivers[0], Receiver: receivers[0],
Msg: message, Msg: message,
} }
ctx = context.WithValue(ctx, balancer.ConsistentHashKey, reqDeliver.Receiver)
res, err := d.postal.Deliver(ctx, reqDeliver) res, err := d.postal.Deliver(ctx, reqDeliver)
if err != nil { if err != nil {
return StatusError, err return StatusError, err
@ -110,6 +115,7 @@ func (d *Deliver) deliver0(ctx context.Context, msg proto.Message, receivers []s
} else { } else {
// deliver batch receiver // deliver batch receiver
ctx = context.WithValue(ctx, balancer.ConsistentHashKey, receivers[0])
req := &postal.ReqDeliverBatch{Receivers: receivers, Msg: message} req := &postal.ReqDeliverBatch{Receivers: receivers, Msg: message}
res, err := d.postal.DeliverBatch(ctx, req) res, err := d.postal.DeliverBatch(ctx, req)
if err != nil { if err != nil {

2
pkg/protocol/protocol_test.go

@ -6,7 +6,7 @@ import (
) )
func TestProtocolCodec(t *testing.T) { func TestProtocolCodec(t *testing.T) {
header := Header{ header := &Header{
Magic: Magic, Magic: Magic,
Type: 1, Type: 1,
Status: 20, Status: 20,

2
pkg/utils/logger/logger.go

@ -7,7 +7,7 @@ import (
var Logger *SoLogger var Logger *SoLogger
func init() { func init() {
Logger = &SoLogger{logger: logrus.New()} Logger = &SoLogger{logger: logrus.StandardLogger()}
} }
type SoLogger struct { type SoLogger struct {

Loading…
Cancel
Save