Browse Source

postal consistent hash picker

master
tangmingyou 3 years ago
parent
commit
e111b2e686
  1. 20
      api/chat.proto
  2. 19
      api/postal.proto
  3. 4
      cmd/chat/config.toml
  4. 17
      cmd/chat/main.go
  5. 2
      cmd/gateway_ws/main.go
  6. 11
      cmd/mahjong/main.go
  7. 52
      internal/chat/logic/chat_server.go
  8. 16
      internal/chat/logic/room_loader.go
  9. 11
      internal/mahjong/store/store.go
  10. 61
      internal/postal/logic/postal_cluster_server.go
  11. 3
      internal/postal/logic/postal_server.go
  12. 45
      pkg/grpc/balancer/consistent_hash.go
  13. 17
      pkg/grpc/discovery/discovery.go
  14. 2
      pkg/grpc/discovery/etcd_naming.go
  15. 2
      pkg/grpc/generic/generic_client.go
  16. 110
      pkg/protocol/deliver/consistent_hash_picker.go
  17. 51
      pkg/protocol/deliver/consistent_hash_picker_test.go
  18. 151
      pkg/protocol/deliver/deliver.go
  19. 218
      pkg/protocol/deliver/deliver_group.go
  20. 136
      pkg/protocol/deliver/deprecated/d_deliver.go
  21. 140
      pkg/protocol/deliver/group_deliver.go
  22. 45
      pkg/protocol/deliver/postal.go
  23. 40
      pkg/protocol/deliver/postal_node.go
  24. 172
      pkg/protocol/deliver/postal_picker.go
  25. 43
      pkg/utils/nets/ip.go
  26. 28
      pkg/utils/nets/ip_test.go

20
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;
}

19
api/postal.proto

@ -80,10 +80,24 @@ message ReqGroupDissolve {
string gid = 1;
}
message ResEndpoint {
string endpoint = 1;
map<string, string> 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<string, string> extra = 2;
}

4
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

17
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)

2
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 {

11
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

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

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

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

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

3
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())}

45
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

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

2
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:
}

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

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

51
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")
}
}

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

218
pkg/protocol/deliver/deliver_group.go

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

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

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

45
pkg/protocol/deliver/postal.go

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

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

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

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

28
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")
}
}
Loading…
Cancel
Save