diff --git a/api/chat.proto b/api/chat.proto index bb26a03..8f9766c 100644 --- a/api/chat.proto +++ b/api/chat.proto @@ -44,8 +44,8 @@ message ResSend { string result = 2; } -message Group { - string gid = 1; +message Room { + string rid = 1; string masterUid = 2; int32 userLimit = 3; int32 userNum = 4; @@ -55,37 +55,37 @@ message Group { } message ReqRoomCreate { - string gname = 2; + string rname = 2; } message ResRoomCreate { - Group group = 1; + Room room = 1; } message ReqRoomJoin { - string gid = 1; + string rid = 1; } message ReqRoomLeave { - string gid = 1; + string rid = 1; } message ReqRoomKickOut { - string gid = 1; + string rid = 1; string uid = 2; } message ReqRoomInfo { - string gid = 1; + string rid = 1; } message ResRoomInfo { - Group group = 1; + Room room = 1; } message ReqRoomList { } message ResRoomList { - repeated Group groups = 1; + repeated Room rooms = 1; } diff --git a/api/postal.proto b/api/postal.proto index 25aed7e..a239b4e 100644 --- a/api/postal.proto +++ b/api/postal.proto @@ -80,10 +80,24 @@ message ReqGroupDissolve { string gid = 1; } +message ResEndpoint { + string endpoint = 1; + map extra = 2; +} + // 集群之间接口调用,重定向消息... service PostalCluster { // socket不在当前节点,重新投递消息 rpc Redirect(ReqRedirect) returns(ResDeliver); + + rpc ReDeliver(ReqReDeliver) returns(ResDeliver); +} + +message ReqReDeliver { + string receiver = 1; + Message msg = 2; + int32 offset = 7; // consistent hash node offset, <0逆时针, >0顺时针, =0丢弃包 + repeated int64 nodes = 8; // [ip:port,] 出现重复节点丢弃包 } message ReqRedirect { @@ -94,8 +108,3 @@ message ReqRedirect { ReqDeliverBatch deliverBatch = 11; ReqDeliverGroup deliverGroup = 12; } - -message ResEndpoint { - string endpoint = 1; - map extra = 2; -} diff --git a/cmd/chat/config.toml b/cmd/chat/config.toml index fe07e5d..68e9b1e 100644 --- a/cmd/chat/config.toml +++ b/cmd/chat/config.toml @@ -38,3 +38,7 @@ DSN = "root:sopod_mysql2347-@tcp(124.222.131.236:3666)/groups?charset=utf8&parse [prometheus] enable = false port = 7039 + +[nats] +Url = "nats://nats.sopod@124.222.131.236:3222" +RetryOnFailedConnect = true diff --git a/cmd/chat/main.go b/cmd/chat/main.go index 51d91ee..87b6be7 100644 --- a/cmd/chat/main.go +++ b/cmd/chat/main.go @@ -3,10 +3,13 @@ package main import ( "context" clientv3 "go.etcd.io/etcd/client/v3" + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" "sonet/api/gen/chat" "sonet/internal/chat/logic" "sonet/pkg/config" "sonet/pkg/grpc/discovery" + "sonet/pkg/plugins/mq" "sonet/pkg/protocol/deliver" "sonet/pkg/utils/shutdown" ) @@ -28,11 +31,19 @@ func main() { // etcd discovery dis := discovery.NewEtcdDiscovery(etcdClient) - deli := deliver.NewDeliver(chat.Chat_ServiceDesc.ServiceName) - if err := deli.InitWithResolver(context.Background(), dis); err != nil { + ctx, cancel := context.WithCancel(context.Background()) + shutdown.AddHook(cancel) + picker := deliver.NewPostalPicker(dis, grpc.WithTransportCredentials(insecure.NewCredentials())) + if err = picker.Init(ctx); err != nil { panic(err) } - chatServer := logic.NewChatServer(deli, nil) + deli := deliver.NewDeliver(chat.Chat_ServiceDesc.ServiceName, picker) + consumer, err := mq.NewNatsConsumer(conf.Nats) + if err != nil { + panic(err) + } + groupDeli := deliver.NewGroupDeliver(chat.Chat_ServiceDesc.ServiceName, logic.NewRoomLoader(), consumer) + chatServer := logic.NewChatServer(deli, groupDeli) go func() { err = chatServer.Run(conf.Grpc, dis) diff --git a/cmd/gateway_ws/main.go b/cmd/gateway_ws/main.go index b110658..e9889d6 100644 --- a/cmd/gateway_ws/main.go +++ b/cmd/gateway_ws/main.go @@ -89,7 +89,7 @@ func main() { }() // run postal cluster server - postalClusterServer := logic.NewPostalClusterServer(postalServer) + postalClusterServer := logic.NewPostalClusterServer(postalServer, nil) go func() { err := postalClusterServer.Run(conf.Grpc.Address, config.GetGrpcOptions(conf.Grpc)...) if err != nil { diff --git a/cmd/mahjong/main.go b/cmd/mahjong/main.go index 29e48ca..31f2328 100644 --- a/cmd/mahjong/main.go +++ b/cmd/mahjong/main.go @@ -11,7 +11,6 @@ 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,9 +24,6 @@ func main() { panic(err) } - registry := etcd.NewRegister(etcdClient) - shutdown.AddHook(registry.Stop) - // init deliver dis := discovery.NewEtcdDiscovery(etcdClient) resolver, err := dis.Resolver() @@ -35,10 +31,13 @@ func main() { panic(err) } - deli := deliver.NewDeliver(mahjong.Mahjong_ServiceDesc.ServiceName) - if err := deli.InitWithResolver(context.Background(), dis); err != nil { + ctx, cancel := context.WithCancel(context.Background()) + shutdown.AddHook(cancel) + picker := deliver.NewPostalPicker(dis, grpc.WithTransportCredentials(insecure.NewCredentials())) + if err = picker.Init(ctx); err != nil { panic(err) } + deli := deliver.NewDeliver(mahjong.Mahjong_ServiceDesc.ServiceName, picker) mjStore := store.NewStore(deli) // auth conn diff --git a/internal/chat/logic/chat_server.go b/internal/chat/logic/chat_server.go index dd684c4..b273ebb 100644 --- a/internal/chat/logic/chat_server.go +++ b/internal/chat/logic/chat_server.go @@ -24,16 +24,16 @@ import ( type ChatServer struct { chat.UnimplementedChatServer deliver *deliver.Deliver - deliverGroup *deliver.GroupDeliver + groupDeliver *deliver.GroupDeliver roomId int64 - rooms *collect.ConcurrentMap[string, *chat.Group] + rooms *collect.ConcurrentMap[string, *chat.Room] } -func NewChatServer(deliver *deliver.Deliver, deliverGroup *deliver.GroupDeliver) *ChatServer { - rooms := collect.NewConcurrentMap[string, *chat.Group](32, func(k string) string { return k }) +func NewChatServer(deliver *deliver.Deliver, groupDeliver *deliver.GroupDeliver) *ChatServer { + rooms := collect.NewConcurrentMap[string, *chat.Room](32, func(k string) string { return k }) return &ChatServer{ deliver: deliver, - deliverGroup: deliverGroup, + groupDeliver: groupDeliver, roomId: 90000, rooms: rooms, } @@ -106,7 +106,7 @@ func (s *ChatServer) RoomSend(ctx context.Context, req *chat.ReqRoomSend) (res * Gid: req.Gid, Content: req.Message, } - err = s.deliverGroup.DeliverGroup(ctx, req.Gid, gMessage) + err = s.groupDeliver.DeliverGroup(ctx, req.Gid, gMessage) return } @@ -115,26 +115,26 @@ func (s *ChatServer) RoomCreate(ctx context.Context, req *chat.ReqRoomCreate) (r if err != nil { return nil, err } - logger.Infof("create room: %s", req.Gname) + logger.Infof("create room: %s", req.Rname) - group := &chat.Group{ - Gid: strconv.Itoa(int(atomic.AddInt64(&s.roomId, 1))), + room := &chat.Room{ + Rid: strconv.Itoa(int(atomic.AddInt64(&s.roomId, 1))), MasterUid: subject.Uid, UserLimit: 10000, UserNum: 1, - Gname: req.Gname, + Gname: req.Rname, CreateUid: subject.Uid, CreateAt: time.Now().UnixMilli(), } - // join postal group - err = s.deliverGroup.GroupJoin(ctx, subject.Uid, []string{group.Gid}) + // join postal room + err = s.groupDeliver.GroupJoin(ctx, subject.Uid, []string{room.Rid}) if err != nil { return } - s.rooms.Store(group.Gid, group) - res = &chat.ResRoomCreate{Group: group} + s.rooms.Store(room.Rid, room) + res = &chat.ResRoomCreate{Room: room} return } @@ -144,14 +144,14 @@ func (s *ChatServer) RoomJoin(ctx context.Context, req *chat.ReqRoomJoin) (_ *em return nil, err } - room, ok := s.rooms.Load(req.Gid) + room, ok := s.rooms.Load(req.Rid) if !ok { - err = fmt.Errorf("room %s not exists", req.Gid) + err = fmt.Errorf("room %s not exists", req.Rid) return } atomic.AddInt32(&room.UserNum, 1) - err = s.deliverGroup.GroupJoin(ctx, subject.Uid, []string{req.Gid}) + err = s.groupDeliver.GroupJoin(ctx, subject.Uid, []string{req.Rid}) return } @@ -161,14 +161,14 @@ func (s *ChatServer) RoomLeave(ctx context.Context, req *chat.ReqRoomLeave) (_ * return nil, err } - room, ok := s.rooms.Load(req.Gid) + room, ok := s.rooms.Load(req.Rid) if !ok { - err = fmt.Errorf("room %s not exists", req.Gid) + err = fmt.Errorf("room %s not exists", req.Rid) return } atomic.AddInt32(&room.UserNum, -1) - err = s.deliverGroup.GroupLeave(ctx, subject.Uid, []string{req.Gid}) + err = s.groupDeliver.GroupLeave(ctx, subject.Uid, []string{req.Rid}) return } @@ -177,22 +177,22 @@ func (s *ChatServer) RoomKickOut(ctx context.Context, req *chat.ReqRoomKickOut) } func (s *ChatServer) RoomInfo(ctx context.Context, req *chat.ReqRoomInfo) (res *chat.ResRoomInfo, err error) { - room, ok := s.rooms.Load(req.Gid) + room, ok := s.rooms.Load(req.Rid) if !ok { - err = fmt.Errorf("room %s not exists", req.Gid) + err = fmt.Errorf("room %s not exists", req.Rid) return } - res = &chat.ResRoomInfo{Group: room} + res = &chat.ResRoomInfo{Room: room} return } func (s *ChatServer) RoomList(ctx context.Context, req *chat.ReqRoomList) (res *chat.ResRoomList, err error) { - rooms := make([]*chat.Group, 16) - s.rooms.Range(func(_ string, room *chat.Group) bool { + rooms := make([]*chat.Room, 16) + s.rooms.Range(func(_ string, room *chat.Room) bool { rooms = append(rooms, room) return true }) - res = &chat.ResRoomList{Groups: rooms} + res = &chat.ResRoomList{Rooms: rooms} return } diff --git a/internal/chat/logic/room_loader.go b/internal/chat/logic/room_loader.go new file mode 100644 index 0000000..d7c8010 --- /dev/null +++ b/internal/chat/logic/room_loader.go @@ -0,0 +1,16 @@ +package logic + +import "sonet/pkg/utils/logger" + +// RoomLoader todo 从持久存储加载用户群关系 +type RoomLoader struct { +} + +func NewRoomLoader() *RoomLoader { + return &RoomLoader{} +} + +func (l *RoomLoader) Load(uid string) (groupIds []string, err error) { + logger.Infof("load %s join rooms", uid) + return +} diff --git a/internal/mahjong/store/store.go b/internal/mahjong/store/store.go index 29376be..838f3ab 100644 --- a/internal/mahjong/store/store.go +++ b/internal/mahjong/store/store.go @@ -66,9 +66,14 @@ func (s *Store) StorePlayer(player *game.MjPlayer) { } else { receivers = setState.IncludeReceivers } - if len(receivers) > 0 { - // deliver player state change - _, _ = s.deliver.DeliverBatch(context.Background(), msg, receivers) + if len(receivers) == 0 { + continue + } + + _, err = s.deliver.DeliverBatch(context.Background(), msg, receivers) + if err != nil { + logger.Error("deliver batch store state error: ", err) + continue } } } diff --git a/internal/postal/logic/postal_cluster_server.go b/internal/postal/logic/postal_cluster_server.go index 437c429..7a95611 100644 --- a/internal/postal/logic/postal_cluster_server.go +++ b/internal/postal/logic/postal_cluster_server.go @@ -7,7 +7,10 @@ import ( "google.golang.org/grpc" "net" "sonet/api/gen/postal" + "sonet/pkg/protocol/deliver" + "sonet/pkg/protocol/session" "sonet/pkg/utils/logger" + "sonet/pkg/utils/nets" ) // PostalClusterPortOffset 相对与 postal grpc server 接口偏移量 @@ -35,11 +38,14 @@ var ( type PostalClusterServer struct { postal.UnimplementedPostalClusterServer postalServer postal.PostalServer // 当前节点 postal server + postalPicker *deliver.PostalPicker + sessionStore session.Store } -func NewPostalClusterServer(postalServer postal.PostalServer) *PostalClusterServer { +func NewPostalClusterServer(postalServer postal.PostalServer, postalPicker *deliver.PostalPicker) *PostalClusterServer { return &PostalClusterServer{ postalServer: postalServer, + postalPicker: postalPicker, } } @@ -75,3 +81,56 @@ func (s *PostalClusterServer) Redirect(ctx context.Context, req *postal.ReqRedir } return nil, errors.New(fmt.Sprintf("not found redirect type %d", req.RedirectMethod)) } + +func (s *PostalClusterServer) ReDeliver(ctx context.Context, req *postal.ReqReDeliver) (res *postal.ResDeliver, err error) { + channel, ok := s.sessionStore.Load(req.Receiver) + if ok { + bytes, err := encodeDeliverMessage(req.Msg) + if err != nil { + return + } + err = channel.Conn.Write(bytes) + if err != nil { + return + } + res = &postal.ResDeliver{Ok: true} + return + } + + // 重新投递 + var next int + next, req.Offset = convergeOffset(req.Offset) + if next == 0 || req.Offset == 0 { + // 重投次数结束 + return + } + node, _, err := s.postalPicker.PickOffsetNode(req.Receiver, next) + if err != nil { + return + } + addr, err := nets.Address2i64(node.GetAddr()) + if err != nil { + return + } + for _, nodeI64 := range req.Nodes { + if addr == nodeI64 { + // 重复投递节点, 结束投递 + return + } + } + + req.Nodes = append(req.Nodes, addr) + + return +} + +// convergeOffset 向0收敛offset +func convergeOffset(offset int32) (int, int32) { + if offset < 0 { + return -1, offset + 1 + } + if offset > 0 { + return 1, offset - 1 + } + return 0, 0 +} diff --git a/internal/postal/logic/postal_server.go b/internal/postal/logic/postal_server.go index 151999a..c206aa4 100644 --- a/internal/postal/logic/postal_server.go +++ b/internal/postal/logic/postal_server.go @@ -199,9 +199,10 @@ func (s *PostalServer) Deliver(ctx context.Context, req *postal.ReqDeliver) (res res = &postal.ResDeliver{Ok: true} return } + // todo 向一致性 hash 下一个节点传递 + // todo picker reqCluster offset +-3 // 用户连接不在当前gateway - // todo 向一致性 hash 下一个节点传递 gateReceivers, offline := s.clusterGateReceivers(ctx, []string{req.Receiver}) if offline != nil && len(offline) > 0 { res = &postal.ResDeliver{Code: int32(postal.DeliverResult_ReceiverOffline.Number())} diff --git a/pkg/grpc/balancer/consistent_hash.go b/pkg/grpc/balancer/consistent_hash.go index 93eaf02..a054676 100644 --- a/pkg/grpc/balancer/consistent_hash.go +++ b/pkg/grpc/balancer/consistent_hash.go @@ -8,6 +8,7 @@ import ( "google.golang.org/grpc/balancer/base" "google.golang.org/grpc/grpclog" "google.golang.org/grpc/resolver" + "sonet/pkg/protocol/deliver" "sonet/pkg/utils/logger" "strconv" ) @@ -32,27 +33,51 @@ func newConsistentHashBuilder() balancer.Builder { type consistentHashPickerBuilder struct{} func (b *consistentHashPickerBuilder) Build(buildInfo base.PickerBuildInfo) balancer.Picker { - // grpclog.Infof("consistentHashPicker: newPicker called with buildInfo: %v", buildInfo) + // logger.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), - } + subConns := make(map[string]balancer.SubConn) + var nodes []*deliver.PickNode 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 + node := &deliver.PickNode{ + Key: conInfo.Address.Addr, + Weight: weight, } + nodes = append(nodes, node) + subConns[node.Key] = sc + } + + picker := deliver.NewConsistentHashPicker(nodes, deliver.DefaultReplicas, deliver.DefaultSalt) + picker.Init() + return &postalConsistentHashPicker{ + subConns: subConns, + picker: picker, } - return picker } +type postalConsistentHashPicker struct { + subConns map[string]balancer.SubConn + picker *deliver.ConsistentHashPicker +} + +func (p *postalConsistentHashPicker) Pick(info balancer.PickInfo) (ret balancer.PickResult, err error) { + key, ok := info.Ctx.Value(ConsistentHashKey).(string) + if !ok || key == "" { + panic(errors.New("empty consistent hash key")) + } + node, ok := p.picker.Pick(key) + if ok { + ret.SubConn = p.subConns[node.Key] + } + return +} + +// consistentHashPicker +// Deprecated type consistentHashPicker struct { subConns map[string]balancer.SubConn hash *Ketama diff --git a/pkg/grpc/discovery/discovery.go b/pkg/grpc/discovery/discovery.go index 8ea4794..2e8f9a3 100644 --- a/pkg/grpc/discovery/discovery.go +++ b/pkg/grpc/discovery/discovery.go @@ -3,6 +3,7 @@ package discovery import ( "context" "google.golang.org/grpc/resolver" + "strconv" ) type Registry interface { @@ -35,3 +36,19 @@ type Server struct { Addr string `json:"addr"` // 地址 Attrs map[string]string `json:"attrs"` // attributes } + +func (s Server) GetWeight() (weight int) { + weight = 1 + if len(s.Attrs) == 0 { + return + } + + v, ok := s.Attrs["weight"] + if !ok { + return + } + if w, err := strconv.Atoi(v); err != nil { + weight = w + } + return +} diff --git a/pkg/grpc/discovery/etcd_naming.go b/pkg/grpc/discovery/etcd_naming.go index 8d3f99c..ad402e2 100644 --- a/pkg/grpc/discovery/etcd_naming.go +++ b/pkg/grpc/discovery/etcd_naming.go @@ -82,8 +82,8 @@ func (r *EtcdDiscovery) Registry(ctx context.Context, server Server) (err error) 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) + c() return case _ = <-keepAliveCh: } diff --git a/pkg/grpc/generic/generic_client.go b/pkg/grpc/generic/generic_client.go index a8640d9..41cd30c 100644 --- a/pkg/grpc/generic/generic_client.go +++ b/pkg/grpc/generic/generic_client.go @@ -86,7 +86,7 @@ func (c *GrpcGenericClient) InvokeUnary(ctx context.Context, method string, reqB ms := time.Now().UnixMilli() resp, err = c.invokeUnary0(ctx, method, reqMessage, opts...) if delay := time.Now().UnixMilli() - ms; delay > 1000 { - logger.Warningf("api %s process use %sms\n", method, delay) + logger.Warningf("api %s process use %dms\n", method, delay) } return } diff --git a/pkg/protocol/deliver/consistent_hash_picker.go b/pkg/protocol/deliver/consistent_hash_picker.go new file mode 100644 index 0000000..ac0943d --- /dev/null +++ b/pkg/protocol/deliver/consistent_hash_picker.go @@ -0,0 +1,110 @@ +package deliver + +import ( + "hash/fnv" + "sort" + "strconv" +) + +type PickNode struct { + Key string + Weight int +} + +type ConsistentHashPicker struct { + nodes []*PickNode + replicas int + salt string + hashKeys []uint32 // sorted hashKeys + hashNodes map[uint32]*PickNode + length int +} + +func NewConsistentHashPicker(nodes []*PickNode, replicas int, salt string) *ConsistentHashPicker { + return &ConsistentHashPicker{ + nodes: nodes, + replicas: replicas, + salt: salt, + hashKeys: make([]uint32, 0, len(nodes)*replicas), + hashNodes: make(map[uint32]*PickNode, len(nodes)*replicas), + } +} + +var ( + DefaultReplicas = 10 + DefaultSalt = "this_is_salt" +) + +func (c *ConsistentHashPicker) hashFnv32(data []byte) uint32 { + f := fnv.New32() + _, err := f.Write(data) + if err != nil { + panic(err) + } + return f.Sum32() +} + +// Init 构建hash环 +func (c *ConsistentHashPicker) Init() { + for _, node := range c.nodes { + weight := node.Weight + if weight < 1 { + weight = 1 + } + for i := 0; i < weight; i++ { + for j := 0; j < c.replicas; j++ { + key := c.hashFnv32([]byte(strconv.Itoa(i) + node.Key + strconv.Itoa(j) + c.salt)) + + if _, ok := c.hashNodes[key]; !ok { + c.hashKeys = append(c.hashKeys, key) + } + c.hashNodes[key] = node + } + } + } + sort.Slice(c.hashKeys, func(i, j int) bool { + return c.hashKeys[i] < c.hashKeys[j] + }) + c.length = len(c.hashKeys) +} + +func (c *ConsistentHashPicker) Pick(key string) (node *PickNode, ok bool) { + if c.length == 0 { + return + } + hash := c.hashFnv32([]byte(key + c.salt)) + + idx := sort.Search(c.length, func(i int) bool { // 二分查找最小为true的index + return c.hashKeys[i] >= hash + }) + if idx == c.length { + idx = 0 + } + node, ok = c.hashNodes[c.hashKeys[idx]] + return +} + +// PickOffset +// key: pick key +// offset: 顺时针第n个节点 +// return node: pick node +// return same: node是否是offset=0时的相同节点 +func (c *ConsistentHashPicker) PickOffset(key string, offset int) (node *PickNode, same bool, ok bool) { + if c.length == 0 { + return + } + hash := c.hashFnv32([]byte(key + c.salt)) + idx := sort.Search(c.length, func(i int) bool { + return c.hashKeys[i] >= hash + }) + if idx == c.length { + idx = 0 + } + hit := (idx + offset) % c.length + if hit < 0 { + hit = c.length + hit + } + node, ok = c.hashNodes[c.hashKeys[hit]] + same = c.hashNodes[c.hashKeys[idx]] == node + return +} diff --git a/pkg/protocol/deliver/consistent_hash_picker_test.go b/pkg/protocol/deliver/consistent_hash_picker_test.go new file mode 100644 index 0000000..eb257f2 --- /dev/null +++ b/pkg/protocol/deliver/consistent_hash_picker_test.go @@ -0,0 +1,51 @@ +package deliver + +import ( + "strconv" + "testing" +) + +func TestConsistentHashPicker(t *testing.T) { + var nodes []*PickNode + var hits = make(map[string]int, 10) + for i := 0; i < 10; i++ { + node := &PickNode{ + Key: "node-" + strconv.Itoa(i), + Weight: 10, + } + if i < 5 { + node.Weight = 5 // 遵循 weight 进行负载均衡 + } + nodes = append(nodes, node) + hits[node.Key] = 0 + } + picker := NewConsistentHashPicker(nodes, DefaultReplicas, DefaultSalt) + picker.Init() + + for i := 0; i < 100000; i++ { + node1, ok1 := picker.Pick(strconv.Itoa(i)) + node2, ok2 := picker.Pick(strconv.Itoa(i)) + if !ok1 || !ok2 { + t.Error("pick non node") + } + if node1 != node2 { + t.Error("pick same key not same node") + } + hits[node1.Key]++ + } + + for nodeKey, hit := range hits { + if hit == 0 { + t.Errorf("node %s not picked", nodeKey) + } + } + + node1, same1, ok1 := picker.PickOffset("100", 1) + node2, same2, ok2 := picker.PickOffset("100", 1) + if !ok1 || !ok2 { + t.Error("pick non node") + } + if node1 != node2 || same1 != same2 { + t.Error("pick same key not same node") + } +} diff --git a/pkg/protocol/deliver/deliver.go b/pkg/protocol/deliver/deliver.go index b72597c..3d8378b 100644 --- a/pkg/protocol/deliver/deliver.go +++ b/pkg/protocol/deliver/deliver.go @@ -2,14 +2,10 @@ package deliver import ( "context" - "fmt" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" + "errors" "google.golang.org/protobuf/proto" "reflect" "sonet/api/gen/postal" - "sonet/pkg/grpc/balancer" - "sonet/pkg/grpc/discovery" "sonet/pkg/utils/logger" "time" ) @@ -22,62 +18,7 @@ const ( StatusReceiverOffline Status = 10 ) -// Deliver n包,通知消息投递 -type Deliver struct { - svcName string - postal postal.PostalClient -} - -func NewDeliver(msgInServiceName string) *Deliver { - return &Deliver{ - svcName: msgInServiceName, - } -} - -func (d *Deliver) InitWithResolver(ctx context.Context, resolver discovery.GrpcResolver, opts ...grpc.DialOption) (err error) { - balancer.InitConsistentHashBuilder() - rb, err := resolver.Resolver() - if err != nil { - return - } - - postalUrl := resolver.DialUrl(postal.Postal_ServiceDesc.ServiceName) - var options []grpc.DialOption - // consistent hash lb - options = append(options, grpc.WithDefaultServiceConfig(fmt.Sprintf(`{"loadBalancingPolicy":"%s"}`, balancer.ConsistentHash))) - options = append(options, grpc.WithResolvers(rb)) - options = append(options, grpc.WithTransportCredentials(insecure.NewCredentials())) - options = append(options, opts...) - conn, err := grpc.DialContext(ctx, postalUrl, options...) - if err != nil { - return - } - d.postal = postal.NewPostalClient(conn) - return -} - -func (d *Deliver) InitWithAddr(postalAddr string, opts ...grpc.DialOption) (err error) { - // Conn *grpc.ClientConn - var options []grpc.DialOption - options = append(options, grpc.WithTransportCredentials(insecure.NewCredentials())) - options = append(options, opts...) - conn, err := grpc.Dial(postalAddr, options...) - if err != nil { - return - } - d.postal = postal.NewPostalClient(conn) - return -} - -func (d *Deliver) Deliver(ctx context.Context, msg proto.Message, receiver string, options ...Option) (Status, error) { - return d.deliver0(ctx, msg, []string{receiver}, options...) -} - -func (d *Deliver) DeliverBatch(ctx context.Context, msg proto.Message, receivers []string, options ...Option) (Status, error) { - return d.deliver0(ctx, msg, receivers, options...) -} - -func protoMessage2Deliver(svcName string, msg proto.Message) (message *postal.Message, err error) { +func Proto2DeliverMessage(svcName string, msg proto.Message) (message *postal.Message, err error) { body, err := proto.Marshal(msg) if err != nil { return @@ -92,60 +33,66 @@ func protoMessage2Deliver(svcName string, msg proto.Message) (message *postal.Me return } -func (d *Deliver) deliver0(ctx context.Context, msg proto.Message, receivers []string, options ...Option) (status Status, err error) { - opts := defaultOptions - if options != nil { - for _, opt := range options { - opt.f(&opts) - } +// Deliver n包,通知消息投递 +type Deliver struct { + serviceName string + postalPicker *PostalPicker +} + +func NewDeliver(svcName string, postalPicker *PostalPicker) *Deliver { + return &Deliver{ + serviceName: svcName, + postalPicker: postalPicker, } - if receivers == nil || len(receivers) == 0 { - // return StatusError, errors.New("receivers is empty") - return StatusSuccess, nil +} + +func (d *Deliver) Deliver(ctx context.Context, msg proto.Message, receiver string, options ...Option) (Status, error) { + client, err := d.postalPicker.Pick(receiver) + if err != nil { + return StatusError, err } - // encode msg - message, err := protoMessage2Deliver(d.svcName, msg) + message, err := Proto2DeliverMessage(d.serviceName, msg) + if err != nil { + logger.Error("encode proto message error: ", err) + return StatusError, err + } + req := &postal.ReqDeliver{Receiver: receiver, Msg: message} + _, err = client.Deliver(ctx, req) if err != nil { return StatusError, err } + return StatusSuccess, nil +} - // deliver to gateway - if len(receivers) == 1 { - // deliver one receiver - reqDeliver := &postal.ReqDeliver{ - Receiver: receivers[0], - Msg: message, - } +func (d *Deliver) DeliverBatch(ctx context.Context, msg proto.Message, receivers []string, options ...Option) (status Status, err error) { + // encode msg + message, err := Proto2DeliverMessage(d.serviceName, msg) + if err != nil { + logger.Error("encode proto message error: ", err) + return StatusError, err + } - ctx = context.WithValue(ctx, balancer.ConsistentHashKey, reqDeliver.Receiver) - res, err := d.postal.Deliver(ctx, reqDeliver) + // 将消息负载均衡pick好后分批发送 + nodeReceivers := make(map[postal.PostalClient][]string, 5) + for _, receiver := range receivers { + client, err := d.postalPicker.Pick(receiver) if err != nil { return StatusError, err } - // TODO res code - if res.Ok { - status = StatusSuccess - } else { - status = StatusError - } - logger.Info("deliver result: ", err, res) - } else { + nodeReceivers[client] = append(nodeReceivers[client], receiver) + } - // 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) + var errs []error + for node, receivers := range nodeReceivers { + _, err := node.DeliverBatch(ctx, &postal.ReqDeliverBatch{Msg: message, Receivers: receivers}) if err != nil { - logger.Error("deliver error: ", err) - return StatusError, err - } - if res.Ok { - status = StatusSuccess - } else { - status = StatusError + errs = append(errs, err) } - logger.Info("deliver batch result: ", res) } - return + if len(errs) > 0 { + err = errors.Join(errs...) + return StatusError, err + } + return StatusSuccess, nil } diff --git a/pkg/protocol/deliver/deliver_group.go b/pkg/protocol/deliver/deliver_group.go deleted file mode 100644 index dd0b4f5..0000000 --- a/pkg/protocol/deliver/deliver_group.go +++ /dev/null @@ -1,218 +0,0 @@ -package deliver - -import ( - "context" - "fmt" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" - "google.golang.org/protobuf/proto" - "sonet/api/gen/postal" - "sonet/pkg/grpc/balancer" - "sonet/pkg/grpc/discovery" - "sonet/pkg/plugins/mq" - "sonet/pkg/protocol/event" - "sonet/pkg/utils/logger" - "sync" - "sync/atomic" -) - -type GroupLoader interface { - Load(uid string) (groupIds []string, err error) -} - -type GroupDeliver struct { - svcName string - groupLoader GroupLoader - consumer mq.Consumer - resolver discovery.Resolver - postal postal.PostalClient // postal consistent hash client - directPostalDialOptions []grpc.DialOption - directPostals *atomic.Value //map[string]*Postal , postal server 直连客户端 - lock *sync.RWMutex -} - -func NewGroupDeliver( - msgInServiceName string, - groupLoader GroupLoader, - consumer mq.Consumer, - resolver discovery.Resolver, - directPostalDialOptions []grpc.DialOption, -) *GroupDeliver { - directPostals := &atomic.Value{} - directPostals.Store(make(map[string]*Postal, 3)) - - return &GroupDeliver{ - svcName: msgInServiceName, - groupLoader: groupLoader, - consumer: consumer, - resolver: resolver, - directPostalDialOptions: directPostalDialOptions, - directPostals: directPostals, - lock: &sync.RWMutex{}, - } -} - -func (d *GroupDeliver) Init(ctx context.Context, grpcResolver discovery.GrpcResolver, opts ...grpc.DialOption) (err error) { - err = d.initPostal(ctx, grpcResolver, opts...) - if err != nil { - return - } - - // initial all postal direct 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 - } - - err = d.subscribe() - return -} - -// toPostalGid 加上 svc name 前缀避免和其他服务群组冲突 -func (d *GroupDeliver) toPostalGid(gid string) string { - return d.svcName + "." + gid -} - -// subscribe postal online events -func (d *GroupDeliver) subscribe() error { - consumerChannel := fmt.Sprintf("%s:%s", "deliver", d.svcName) - return d.consumer.Subscribe(event.TopicOnline, consumerChannel, func(message *mq.Message) (err error) { - online := &event.Online{} - if e := online.UnmarshalBinary(message.Body); e != nil { - logger.Error("deliver unmarshal online event payload error: ", e) - return - } - // load uid groups join to postal - groupIds, err := d.groupLoader.Load(online.Uid) - if err != nil { - logger.Error("deliver load groups error:", err) - return - } - if err = d.GroupJoin(context.Background(), online.Uid, groupIds); err != nil { - logger.Errorf("deliver uid %s group join error: %v", online.Uid, err) - return - } - return - }) -} - -func (d *GroupDeliver) initPostal(ctx context.Context, resolver discovery.GrpcResolver, opts ...grpc.DialOption) (err error) { - balancer.InitConsistentHashBuilder() - rb, err := resolver.Resolver() - if err != nil { - return - } - - postalUrl := resolver.DialUrl(postal.Postal_ServiceDesc.ServiceName) - var options []grpc.DialOption - // consistent hash lb - options = append(options, grpc.WithDefaultServiceConfig(fmt.Sprintf(`{"loadBalancingPolicy":"%s"}`, balancer.ConsistentHash))) - options = append(options, grpc.WithResolvers(rb)) - options = append(options, grpc.WithTransportCredentials(insecure.NewCredentials())) - options = append(options, opts...) - conn, err := grpc.DialContext(ctx, postalUrl, options...) - if err != nil { - return - } - d.postal = postal.NewPostalClient(conn) - 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) - directPostals := d.directPostals.Load().(map[string]*Postal) - for _, server := range servers { - if p, ok := directPostals[server.Addr]; ok { - postals[server.Addr] = p - continue - } - // new connection - conn, err := grpc.DialContext(context.Background(), server.Addr, d.directPostalDialOptions...) - if err != nil { - logger.Errorf("dial postal server %+v error: %v", server.Addr, err) - continue - } - postals[server.Addr] = NewPostal(conn) - } - - // close old connection - oldPostals := directPostals - d.directPostals.Store(postals) - for addr, p := range oldPostals { - if _, ok := postals[addr]; !ok { - p.Close() - } - } -} - -func (d *GroupDeliver) DeliverGroup(ctx context.Context, gid string, msg proto.Message, options ...Option) (err error) { - gid = d.toPostalGid(gid) - - // 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.directPostals.Load().(map[string]*Postal) { - p.DeliverGroup(&req) - } - return -} - -func (d *GroupDeliver) GroupDissolve(gid string) { - gid = d.toPostalGid(gid) - - req := postal.ReqGroupDissolve{Gid: gid} - for _, p := range d.directPostals.Load().(map[string]*Postal) { - p.GroupDissolve(&req) - } -} - -func (d *GroupDeliver) GroupJoin(ctx context.Context, uid string, gids []string) (err error) { - if len(gids) == 0 { - return - } - for i := 0; i < len(gids); i++ { - gids[i] = d.toPostalGid(gids[i]) - } - - ctx = context.WithValue(ctx, balancer.ConsistentHashKey, uid) - req := &postal.ReqGroupJoin{Uid: uid, Gids: gids} - _, err = d.postal.GroupJoin(ctx, req) - return -} - -func (d *GroupDeliver) GroupLeave(ctx context.Context, uid string, gids []string) (err error) { - req := &postal.ReqGroupLeave{Uid: uid, Gids: gids} - - ctx = context.WithValue(ctx, balancer.ConsistentHashKey, uid) - _, err = d.postal.GroupLeave(ctx, req) - return -} diff --git a/pkg/protocol/deliver/deprecated/d_deliver.go b/pkg/protocol/deliver/deprecated/d_deliver.go new file mode 100644 index 0000000..854fcdc --- /dev/null +++ b/pkg/protocol/deliver/deprecated/d_deliver.go @@ -0,0 +1,136 @@ +package deprecated + +import ( + "context" + "fmt" + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/protobuf/proto" + "sonet/api/gen/postal" + "sonet/pkg/grpc/balancer" + "sonet/pkg/grpc/discovery" + "sonet/pkg/protocol/deliver" + "sonet/pkg/utils/logger" +) + +type Status int16 + +const ( + StatusSuccess Status = 1 + StatusError Status = 2 + StatusReceiverOffline Status = 10 +) + +// DDeliver n包,通知消息投递 +// Deprecated +type DDeliver struct { + svcName string + postal postal.PostalClient +} + +func NewDDeliver(msgInServiceName string) *DDeliver { + return &DDeliver{ + svcName: msgInServiceName, + } +} + +func (d *DDeliver) InitWithResolver(ctx context.Context, resolver discovery.GrpcResolver, opts ...grpc.DialOption) (err error) { + balancer.InitConsistentHashBuilder() + rb, err := resolver.Resolver() + if err != nil { + return + } + + postalUrl := resolver.DialUrl(postal.Postal_ServiceDesc.ServiceName) + var options []grpc.DialOption + // consistent hash lb + options = append(options, grpc.WithDefaultServiceConfig(fmt.Sprintf(`{"loadBalancingPolicy":"%s"}`, balancer.ConsistentHash))) + options = append(options, grpc.WithResolvers(rb)) + options = append(options, grpc.WithTransportCredentials(insecure.NewCredentials())) + options = append(options, opts...) + conn, err := grpc.DialContext(ctx, postalUrl, options...) + if err != nil { + return + } + d.postal = postal.NewPostalClient(conn) + return +} + +func (d *DDeliver) InitWithAddr(postalAddr string, opts ...grpc.DialOption) (err error) { + // Conn *grpc.ClientConn + var options []grpc.DialOption + options = append(options, grpc.WithTransportCredentials(insecure.NewCredentials())) + options = append(options, opts...) + conn, err := grpc.Dial(postalAddr, options...) + if err != nil { + return + } + d.postal = postal.NewPostalClient(conn) + return +} + +func (d *DDeliver) Deliver(ctx context.Context, msg proto.Message, receiver string, options ...deliver.Option) (Status, error) { + return d.deliver0(ctx, msg, []string{receiver}, options...) +} + +func (d *DDeliver) DeliverBatch(ctx context.Context, msg proto.Message, receivers []string, options ...deliver.Option) (Status, error) { + return d.deliver0(ctx, msg, receivers, options...) +} + +func (d *DDeliver) deliver0(ctx context.Context, msg proto.Message, receivers []string, options ...deliver.Option) (status Status, err error) { + //opts := deliver.defaultOptions + //if options != nil { + // for _, opt := range options { + // opt.f(&opts) + // } + //} + if receivers == nil || len(receivers) == 0 { + // return StatusError, errors.New("receivers is empty") + return StatusSuccess, nil + } + + // encode msg + message, err := deliver.Proto2DeliverMessage(d.svcName, msg) + if err != nil { + return StatusError, err + } + + // deliver to gateway + if len(receivers) == 1 { + // deliver one receiver + reqDeliver := &postal.ReqDeliver{ + 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 + } + // TODO res code + if res.Ok { + status = StatusSuccess + } else { + status = StatusError + } + logger.Info("deliver result: ", err, res) + } 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 { + logger.Error("deliver error: ", err) + return StatusError, err + } + if res.Ok { + status = StatusSuccess + } else { + status = StatusError + } + logger.Info("deliver batch result: ", res) + } + return +} diff --git a/pkg/protocol/deliver/group_deliver.go b/pkg/protocol/deliver/group_deliver.go new file mode 100644 index 0000000..4c7fa47 --- /dev/null +++ b/pkg/protocol/deliver/group_deliver.go @@ -0,0 +1,140 @@ +package deliver + +import ( + "context" + "errors" + "fmt" + "google.golang.org/protobuf/proto" + "sonet/api/gen/postal" + "sonet/pkg/plugins/mq" + "sonet/pkg/protocol/event" + "sonet/pkg/utils/collect" + "sonet/pkg/utils/logger" +) + +type GroupLoader interface { + Load(uid string) (groupIds []string, err error) +} + +type GroupDeliver struct { + svcName string + groupLoader GroupLoader + consumer mq.Consumer + postalPicker *PostalPicker +} + +func NewGroupDeliver( + msgInServiceName string, + groupLoader GroupLoader, + consumer mq.Consumer, +) *GroupDeliver { + return &GroupDeliver{ + svcName: msgInServiceName, + groupLoader: groupLoader, + consumer: consumer, + } +} + +func (d *GroupDeliver) Init(ctx context.Context) (err error) { + err = d.subscribe() + return +} + +// toPostalGid 加上 svc name 前缀避免和其他服务群组冲突 +func (d *GroupDeliver) toPostalGid(gid string) string { + return d.svcName + "." + gid +} + +// subscribe postal online events +func (d *GroupDeliver) subscribe() error { + consumerChannel := fmt.Sprintf("%s:%s", "deliver", d.svcName) + return d.consumer.Subscribe(event.TopicOnline, consumerChannel, func(message *mq.Message) (err error) { + online := &event.Online{} + if e := online.UnmarshalBinary(message.Body); e != nil { + logger.Error("deliver unmarshal online event payload error: ", e) + return + } + // load uid groups join to postal + groupIds, err := d.groupLoader.Load(online.Uid) + if err != nil { + logger.Error("deliver load groups error:", err) + return + } + if err = d.GroupJoin(context.Background(), online.Uid, groupIds); err != nil { + logger.Errorf("deliver uid %s group join error: %v", online.Uid, err) + return + } + return + }) +} + +// DeliverGroup 群组消息广播投递到所有 postal 节点 +func (d *GroupDeliver) DeliverGroup(ctx context.Context, gid string, msg proto.Message, options ...Option) (err error) { + message, err := Proto2DeliverMessage(d.svcName, msg) + if err != nil { + return + } + clients := d.postalPicker.PickAll() + + var errs []error + req := postal.ReqDeliverGroup{Gid: d.toPostalGid(gid), Msg: message} + for _, client := range clients { + _, err := client.DeliverGroup(ctx, &req) + if err != nil { + errs = append(errs, err) + } + } + if len(errs) > 0 { + err = errors.Join(errs...) + return + } + return +} + +// GroupDissolve 群组消息广播投递到所有 postal 节点 +func (d *GroupDeliver) GroupDissolve(ctx context.Context, gid string) (err error) { + clients := d.postalPicker.PickAll() + + var errs []error + req := postal.ReqGroupDissolve{Gid: d.toPostalGid(gid)} + for _, client := range clients { + _, err := client.GroupDissolve(ctx, &req) + if err != nil { + errs = append(errs, err) + } + } + if len(errs) > 0 { + err = errors.Join(errs...) + return + } + return +} + +func (d *GroupDeliver) GroupJoin(ctx context.Context, uid string, gids []string) (err error) { + if len(gids) == 0 { + return + } + gids = collect.Mapping(gids, d.toPostalGid) + + client, err := d.postalPicker.Pick(uid) + if err != nil { + return + } + _, err = client.GroupJoin(ctx, &postal.ReqGroupJoin{Uid: uid, Gids: gids}) + return +} + +func (d *GroupDeliver) GroupLeave(ctx context.Context, uid string, gids []string) (err error) { + if len(gids) == 0 { + return + } + gids = collect.Mapping(gids, d.toPostalGid) + + client, err := d.postalPicker.Pick(uid) + if err != nil { + return + } + req := &postal.ReqGroupLeave{Uid: uid, Gids: gids} + _, err = client.GroupLeave(ctx, req) + return +} diff --git a/pkg/protocol/deliver/postal.go b/pkg/protocol/deliver/postal.go deleted file mode 100644 index bbd60b9..0000000 --- a/pkg/protocol/deliver/postal.go +++ /dev/null @@ -1,45 +0,0 @@ -package deliver - -import ( - "context" - "google.golang.org/grpc" - "sonet/api/gen/postal" - "sonet/pkg/utils/logger" -) - -type Postal struct { - conn *grpc.ClientConn - client postal.PostalClient -} - -func NewPostal(conn *grpc.ClientConn) *Postal { - return &Postal{ - conn: conn, - client: postal.NewPostalClient(conn), - } -} - -func (p *Postal) Close() { - err := p.conn.Close() - if err != nil { - logger.Error("postal close error: ", err) - } -} - -// DeliverGroup 群消息发送 -func (p *Postal) DeliverGroup(req *postal.ReqDeliverGroup) { - // todo put in chan - res, err := p.client.DeliverGroup(context.Background(), req) - if err != nil || !res.Ok { - logger.Errorf("deliver to postal failed: %+v, %v", res, err) - } -} - -// GroupDissolve 解散群组 -func (p *Postal) GroupDissolve(req *postal.ReqGroupDissolve) { - // todo put in chan - _, err := p.client.GroupDissolve(context.Background(), req) - if err != nil { - logger.Errorf("group dissolve failed: %v", err) - } -} diff --git a/pkg/protocol/deliver/postal_node.go b/pkg/protocol/deliver/postal_node.go new file mode 100644 index 0000000..531f76a --- /dev/null +++ b/pkg/protocol/deliver/postal_node.go @@ -0,0 +1,40 @@ +package deliver + +import ( + "google.golang.org/grpc" + "sonet/api/gen/postal" + "sonet/pkg/utils/logger" +) + +type PostalNode struct { + addr string + conn *grpc.ClientConn + client postal.PostalClient +} + +func NewPostalNode(addr string, conn *grpc.ClientConn) *PostalNode { + return &PostalNode{ + addr: addr, + conn: conn, + client: postal.NewPostalClient(conn), + } +} + +func (p *PostalNode) Close() { + err := p.conn.Close() + if err != nil { + logger.Error("postal close error: ", err) + } +} + +func (p *PostalNode) GetAddr() string { + return p.addr +} + +func (p *PostalNode) GetConn() *grpc.ClientConn { + return p.conn +} + +func (p *PostalNode) GetClient() postal.PostalClient { + return p.client +} diff --git a/pkg/protocol/deliver/postal_picker.go b/pkg/protocol/deliver/postal_picker.go new file mode 100644 index 0000000..688d747 --- /dev/null +++ b/pkg/protocol/deliver/postal_picker.go @@ -0,0 +1,172 @@ +package deliver + +import ( + "context" + "errors" + "fmt" + "google.golang.org/grpc" + "sonet/api/gen/postal" + "sonet/pkg/grpc/discovery" + "sonet/pkg/utils/collect" + "sonet/pkg/utils/logger" + "sync/atomic" +) + +type PostalPicker struct { + resolver discovery.Resolver + directPostalDialOptions []grpc.DialOption + directPostalNodes atomic.Value // map[string]*PostalNode , postal server 直连客户端 + picker atomic.Value // *ConsistentHashPicker +} + +func NewPostalPicker(resolver discovery.Resolver, directPostalDialOptions ...grpc.DialOption) *PostalPicker { + return &PostalPicker{ + resolver: resolver, + directPostalDialOptions: directPostalDialOptions, + } +} + +func (p *PostalPicker) Init(ctx context.Context) (err error) { + // initial all postal direct clients + servers, err := p.resolver.ResolveAll(ctx, postal.Postal_ServiceDesc.ServiceName) + if err != nil { + return + } + p.newPostalServers(servers) + + // watch postal server instance + err = p.watchPostal(ctx) + if err != nil { + return + } + return +} + +// watchPostal watch postal service list +func (p *PostalPicker) watchPostal(ctx context.Context) (err error) { + ch, err := p.resolver.Watch(ctx, postal.Postal_ServiceDesc.ServiceName) + if err != nil { + return + } + go func() { + for { + select { + case <-ctx.Done(): + return + case servers := <-ch: + p.newPostalServers(servers) + } + } + }() + return +} + +// newPostalServers postal集群节点增减 +func (p *PostalPicker) newPostalServers(servers []discovery.Server) { + newPostal := make(map[string]*PostalNode) + var pickServers []discovery.Server + oldPostal, ok := p.directPostalNodes.Load().(map[string]*PostalNode) + if !ok { + oldPostal = make(map[string]*PostalNode) + } + + for _, server := range servers { + if p, ok := oldPostal[server.Addr]; ok { + newPostal[server.Addr] = p + continue + } + // new connection + conn, err := grpc.DialContext(context.Background(), server.Addr, p.directPostalDialOptions...) + if err != nil { + logger.Errorf("dial postal server %+v error: %v", server, err) + continue + } + newPostal[server.Addr] = NewPostalNode(server.Addr, conn) + pickServers = append(pickServers, server) + } + // 重新构建 consistent hash picker + p.newConsistentHash(pickServers) + + // close old connection + p.directPostalNodes.Store(newPostal) + for addr, p := range oldPostal { + if _, ok := newPostal[addr]; !ok { + p.Close() + } + } +} + +func (p *PostalPicker) newConsistentHash(servers []discovery.Server) { + nodes := collect.Mapping(servers, func(server discovery.Server) *PickNode { + return &PickNode{ + Key: server.Addr, + Weight: server.GetWeight(), + } + }) + + picker := NewConsistentHashPicker(nodes, DefaultReplicas, DefaultSalt) + picker.Init() + p.picker.Store(picker) +} + +func (p *PostalPicker) Pick(key string) (client postal.PostalClient, err error) { + postalNode, err := p.PickNode(key) + if err != nil { + return + } + client = postalNode.client + return +} + +func (p *PostalPicker) PickNode(key string) (postalNode *PostalNode, err error) { + picker := p.picker.Load().(*ConsistentHashPicker) + node, ok := picker.Pick(key) + if !ok { + err = errors.New("no available postal server") + return + } + postalNode, ok = p.directPostalNodes.Load().(map[string]*PostalNode)[node.Key] + if !ok { + err = fmt.Errorf("no conn for postal server %s", node.Key) + return + } + return +} + +// PickOffset +// return same: client 是否与 offset=0 时相同 +func (p *PostalPicker) PickOffset(key string, offset int) (client postal.PostalClient, same bool, err error) { + postalNode, same, err := p.PickOffsetNode(key, offset) + if err != nil { + return + } + client = postalNode.client + return +} + +func (p *PostalPicker) PickOffsetNode(key string, offset int) (postalNode *PostalNode, same bool, err error) { + picker := p.picker.Load().(*ConsistentHashPicker) + node, same, ok := picker.PickOffset(key, offset) + if !ok { + err = errors.New("no available postal server") + return + } + postalNode, ok = p.directPostalNodes.Load().(map[string]*PostalNode)[node.Key] + if !ok { + err = fmt.Errorf("no conn for postal server %s", node.Key) + return + } + return +} + +func (p *PostalPicker) PickAll() (clients []postal.PostalClient) { + nodes, ok := p.directPostalNodes.Load().(map[string]*PostalNode) + if !ok { + return + } + clients = make([]postal.PostalClient, 0, len(nodes)) + for _, node := range nodes { + clients = append(clients, node.client) + } + return +} diff --git a/pkg/utils/nets/ip.go b/pkg/utils/nets/ip.go index 65a2f83..183f0d6 100644 --- a/pkg/utils/nets/ip.go +++ b/pkg/utils/nets/ip.go @@ -2,7 +2,12 @@ package nets import ( "errors" + "fmt" + "math" + "math/big" "net" + "strconv" + "strings" ) // GetHostIpv4 获取本地内网IP @@ -40,3 +45,41 @@ func getAllIPV4(filter func(net.IP) bool) (ips []string, err error) { } return } + +func Address2i64(addr string) (int64, error) { + idx := strings.Index(addr, ":") + if idx < 0 { + return 0, fmt.Errorf("invalid addr %s", addr) + } + ip, port := addr[0:idx], addr[idx+1:] + portI, err := strconv.Atoi(port) + if err != nil { + return 0, err + } + ipI, err := Ip2i64(ip) + if err != nil { + return 0, err + } + return (int64(portI) << 32) | ipI, nil +} + +func I642Address(addr int64) string { + port := addr >> 32 + ip := addr & math.MaxUint32 + return fmt.Sprintf("%s:%d", I642Ip(ip), port) +} + +func Ip2i64(ip string) (int64, error) { + ip4 := net.ParseIP(ip).To4() + if ip4 == nil { + return 0, fmt.Errorf("invalid ip %s", ip) + } + ret := big.NewInt(0) + ret.SetBytes(ip4) + return ret.Int64(), nil +} + +func I642Ip(ip int64) string { + return fmt.Sprintf("%d.%d.%d.%d", + byte(ip>>24), byte(ip>>16), byte(ip>>8), byte(ip)) +} diff --git a/pkg/utils/nets/ip_test.go b/pkg/utils/nets/ip_test.go new file mode 100644 index 0000000..ae9c999 --- /dev/null +++ b/pkg/utils/nets/ip_test.go @@ -0,0 +1,28 @@ +package nets + +import "testing" + +func TestIpConvert(t *testing.T) { + ip := "192.168.1.110" + // ip := "255.255.255.255" + i64, err := Ip2i64(ip) + if err != nil { + t.Error(err) + } + ip2 := I642Ip(int64(int32(i64))) + if ip2 != ip { + t.Error("ip parse error") + } +} + +func TestAddressConvert(t *testing.T) { + addr := "192.168.1.110:8080" + i64, err := Address2i64(addr) + if err != nil { + t.Error(err) + } + addr2 := I642Address(i64) + if addr2 != addr { + t.Error("address parse error") + } +}