From a664e186b04a710c7a81de72141f80844b6aa21e Mon Sep 17 00:00:00 2001 From: tangmingyou Date: Fri, 2 Feb 2024 20:42:35 +0800 Subject: [PATCH] deliver group --- benchmark/main.go | 44 +----- cmd/gateway_ws/main.go | 17 ++- .../gateway_ws/gws_server/conn_handler.go | 8 +- internal/postal/group/group.go | 52 +++++++ internal/postal/logic/postal.go | 24 ++++ internal/postal/logic/postal_server.go | 130 +++++++++++++++--- pkg/protocol/deliver/deliver_group.go | 27 ++-- pkg/protocol/deliver/postal.go | 10 ++ pkg/protocol/event/event.go | 22 +++ pkg/protocol/session/net_conn.go | 5 - pkg/protocol/session/session.go | 51 +++++++ pkg/protocol/session/store.go | 79 ++++++----- pkg/utils/collect/concurrent_map.go | 35 ++++- pkg/utils/collect/concurrent_map_test.go | 43 ++++++ 14 files changed, 427 insertions(+), 120 deletions(-) create mode 100644 internal/postal/logic/postal.go create mode 100644 pkg/protocol/event/event.go delete mode 100644 pkg/protocol/session/net_conn.go create mode 100644 pkg/protocol/session/session.go diff --git a/benchmark/main.go b/benchmark/main.go index d37bdcd..3df4abd 100644 --- a/benchmark/main.go +++ b/benchmark/main.go @@ -6,16 +6,13 @@ import ( "fmt" "github.com/bytedance/sonic" "github.com/gorilla/websocket" - clientv3 "go.etcd.io/etcd/client/v3" "google.golang.org/protobuf/proto" "math/rand" "net" "runtime" "sonet/api/gen/auth" "sonet/api/gen/chat" - "sonet/api/gen/postal" "sonet/pkg/config" - "sonet/pkg/grpc/discovery" "sonet/pkg/protocol" "sonet/pkg/utils/logger" "sonet/pkg/utils/security" @@ -47,6 +44,10 @@ func init() { callbackMutex = &sync.Mutex{} } +func main() { + benchmark() +} + func records(ctx context.Context) { ticker := time.NewTicker(time.Second) defer ticker.Stop() @@ -71,42 +72,7 @@ type NetUser struct { Conn *websocket.Conn } -func main() { - client, err := clientv3.New(clientv3.Config{ - Endpoints: []string{"124.222.131.236:3279"}, - Username: "root", - Password: "sopod@etcd", - }) - if err != nil { - panic(err) - } - - servers, err := discovery.ResolveAll(context.Background(), client, postal.Postal_ServiceDesc.ServiceName) - if err != nil { - panic(err) - } - fmt.Printf("%+v\n", servers) - - ctx, cancel := context.WithCancel(context.Background()) - ch := discovery.Watch(ctx, client, postal.Postal_ServiceDesc.ServiceName) - - go func() { - for { - select { - case <-ctx.Done(): - fmt.Println("done2") - return - case servers := <-ch: - fmt.Printf("watch services: %+v\n", servers) - } - } - }() - - shutdown.AddHook(cancel) - shutdown.Await() -} - -func main2() { +func benchmark() { runtime.GOMAXPROCS(runtime.NumCPU()) // go prof.StartPprof(":8888") diff --git a/cmd/gateway_ws/main.go b/cmd/gateway_ws/main.go index 36e7471..b110658 100644 --- a/cmd/gateway_ws/main.go +++ b/cmd/gateway_ws/main.go @@ -7,6 +7,7 @@ import ( "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" "sonet/internal/gateway_ws/gws_server" + "sonet/internal/postal/group" "sonet/internal/postal/logic" "sonet/pkg/config" "sonet/pkg/grpc/client" @@ -16,6 +17,7 @@ import ( "sonet/pkg/plugins/cache" "sonet/pkg/plugins/mq" "sonet/pkg/protocol/session" + "sonet/pkg/utils/collect" "sonet/pkg/utils/conver" "sonet/pkg/utils/shutdown" ) @@ -41,7 +43,7 @@ func main() { } shutdown.AddHook(func() { _ = etcdClient.Close() }, shutdown.WithOrderBack()) - subjectStore, err := getSubjectMultiCache(appConf, conf.Redis, conf.Nats) + subjectStore, producer, err := getSubjectMultiCache(appConf, conf.Redis, conf.Nats) if err != nil { panic(err) } @@ -58,7 +60,8 @@ func main() { ) grpcFactory.Init() - sessionStore := session.NewMapStore() + sessionStore := session.NewMapStore(128) + groupStore := collect.NewConcurrentMap[string, *group.Group](128, func(k string) string { return k }) // run websocket server postalAddr := etcd.MustRegisterAddress(conf.Grpc.Address) @@ -77,7 +80,7 @@ func main() { // run postal server clientFactory := client.NewGrpcDirectClientFactory(grpc.WithTransportCredentials(insecure.NewCredentials())) - postalServer := logic.NewPostalServer(appConf.EndpointAddress, sessionStore, subjectStore, clientFactory) + postalServer := logic.NewPostalServer(appConf.EndpointAddress, sessionStore, groupStore, subjectStore, clientFactory, producer) go func() { err = postalServer.Run(conf.Grpc, dis) if err != nil { @@ -97,7 +100,7 @@ func main() { shutdown.Await() } -func getSubjectMultiCache(appConf *GatewayWsConfig, redisOptions redis.Options, natsOptions nats.Options) (*cache.LocalRemoteCache, error) { +func getSubjectMultiCache(appConf *GatewayWsConfig, redisOptions redis.Options, natsOptions nats.Options) (*cache.LocalRemoteCache, mq.Producer, error) { // initial cache... rdb := redis.NewClient(&redisOptions) subjectRedisCache := cache.NewRedisCache(appConf.SubjectCacheTopic, rdb) @@ -106,12 +109,12 @@ func getSubjectMultiCache(appConf *GatewayWsConfig, redisOptions redis.Options, // nats mq producer, err := mq.NewNatsProducer(natsOptions) if err != nil { - return nil, err + return nil, nil, err } shutdown.AddHook(func() { producer.Stop() }) consumer, err := mq.NewNatsConsumer(natsOptions) if err != nil { - return nil, err + return nil, nil, err } shutdown.AddHook(func() { consumer.Stop() }) @@ -125,5 +128,5 @@ func getSubjectMultiCache(appConf *GatewayWsConfig, redisOptions redis.Options, Consumer: consumer, } subjectLrc, err := cache.NewLocalRemoteCache(subjectLrcOpts) - return subjectLrc, err + return subjectLrc, producer, err } diff --git a/internal/gateway_ws/gws_server/conn_handler.go b/internal/gateway_ws/gws_server/conn_handler.go index bcf305e..fa61961 100644 --- a/internal/gateway_ws/gws_server/conn_handler.go +++ b/internal/gateway_ws/gws_server/conn_handler.go @@ -56,6 +56,7 @@ func (c *GwsHandler) OnClose(socket *gws.Conn, err error) { } uid := val.(string) c.sessionStore.Delete(uid) + // todo offline event if err = c.subjectStore.Del(context.Background(), uid); err != nil { logger.Error("del offline subject store error: ", err) } @@ -148,6 +149,7 @@ func (c *GwsHandler) OnMessage(socket *gws.Conn, message *gws.Message) { return } socket.Session().Store(SessionUidKey, sessionUid) + // todo online event } header.Type = protocol.TypeResponse @@ -185,14 +187,14 @@ func (c *GwsHandler) extractAuthVerify(conn *gwsConn, resp []byte) (sessionUid s } // 关闭旧的链接 - oldConn, ok := c.sessionStore.Load(sessionUid) + oldChannel, ok := c.sessionStore.Load(sessionUid) if ok { // TODO nats offline / force load subject target cluster call offline - oldConn.(*gwsConn).conn.WriteClose(1000, nil) + oldChannel.Conn.(*gwsConn).conn.WriteClose(1000, nil) logger.Infof("close subject old conn: %s\n", sessionUid) } // 存储session到内存 - c.sessionStore.Store(sessionUid, conn) + c.sessionStore.Store(sessionUid, session.NewChannel(sessionUid, conn)) logger.Infof("subject online: %s\n", sessionUid) return } diff --git a/internal/postal/group/group.go b/internal/postal/group/group.go index e84a706..3e880d0 100644 --- a/internal/postal/group/group.go +++ b/internal/postal/group/group.go @@ -3,6 +3,9 @@ package group import ( "errors" "github.com/redis/go-redis/v9" + "sonet/pkg/protocol/session" + "sonet/pkg/utils/logger" + "sync" ) var ErrNotExists = errors.New("not exists") @@ -28,3 +31,52 @@ type RedisPostalDao struct { func (d *RedisPostalDao) LoadGroupIdsByUid(uid string) (gids []string, err error) { return } + +// Group 群组 +type Group struct { + sync.RWMutex + Gid string + uids map[string]session.NetConn +} + +func NewGroup(gid string) *Group { + return &Group{ + Gid: gid, + uids: make(map[string]session.NetConn), + } +} + +func (g *Group) Write(data []byte) { + g.RLock() + defer g.RUnlock() + for uid, conn := range g.uids { + if err := conn.Write(data); err != nil { + logger.Errorf("group send %s.%s error: ", g.Gid, uid, err) + } + } +} + +func (g *Group) Join(uid string, conn session.NetConn) { + g.Lock() + defer g.Unlock() + g.uids[uid] = conn +} + +// Leave delete uid +func (g *Group) Leave(uid string) { + g.Lock() + defer g.Unlock() + delete(g.uids, uid) +} + +// Dismiss 解散 +func (g *Group) Dismiss() { + +} + +func (g *Group) Load(uid string) (conn session.NetConn, ok bool) { + g.RLock() + defer g.RUnlock() + conn, ok = g.uids[uid] + return +} diff --git a/internal/postal/logic/postal.go b/internal/postal/logic/postal.go new file mode 100644 index 0000000..958cbf9 --- /dev/null +++ b/internal/postal/logic/postal.go @@ -0,0 +1,24 @@ +package logic + +import ( + "sonet/api/gen/postal" + "sonet/pkg/protocol" + "sonet/pkg/utils/logger" +) + +func encodeDeliverMessage(msg *postal.Message) (bytes []byte, err error) { + header := &protocol.Header{ + Magic: protocol.Magic, + Type: protocol.TypeNotice, + UrlType: 1, + SerializeType: 1, + Svc: msg.Svc, + Target: msg.Msg, + } + payload := &protocol.Payload{Header: header, Body: msg.Body} + bytes, err = protocol.EncodeSo(payload) + if err != nil { + logger.Error("encode deliver message error: ", err) + } + return +} diff --git a/internal/postal/logic/postal_server.go b/internal/postal/logic/postal_server.go index b221ecb..c295fd2 100644 --- a/internal/postal/logic/postal_server.go +++ b/internal/postal/logic/postal_server.go @@ -2,41 +2,51 @@ package logic import ( "context" + "errors" "google.golang.org/grpc" "google.golang.org/grpc/reflection" "google.golang.org/protobuf/types/known/emptypb" "net" "sonet/api/gen/postal" + "sonet/internal/postal/group" "sonet/pkg/config" "sonet/pkg/grpc/client" "sonet/pkg/grpc/discovery" "sonet/pkg/grpc/interceptor" "sonet/pkg/plugins/cache" + "sonet/pkg/plugins/mq" "sonet/pkg/protocol" "sonet/pkg/protocol/session" + "sonet/pkg/utils/collect" "sonet/pkg/utils/logger" "sonet/pkg/utils/shutdown" ) type PostalServer struct { postal.UnimplementedPostalServer - endpointAddress string // websocket 前端连接地址 - broadcastAddress string // postal server 在注册中心注册的地址, 其他服务可直连访问 - sessionStore session.Store // 在线用户conn存储 - subjectStore cache.MultiLevelCache // online subject address cache, ws gateway集群所有在线用户存储 + endpointAddress string // websocket 前端连接地址 + broadcastAddress string // postal server 在注册中心注册的地址, 其他服务可直连访问 + sessionStore session.Store // k:uid 在线用户conn存储 + groupStore *collect.ConcurrentMap[string, *group.Group] // k:gid 在线用户 groups 存储 + subjectStore cache.MultiLevelCache // online subject address cache, ws gateway集群所有在线用户存储 clientFactory *client.GrpcDirectClientFactory + producer mq.Producer } func NewPostalServer( endpointAddress string, sessionStore session.Store, + groupStore *collect.ConcurrentMap[string, *group.Group], subjectStore cache.MultiLevelCache, - clientFactory *client.GrpcDirectClientFactory) *PostalServer { + clientFactory *client.GrpcDirectClientFactory, + producer mq.Producer) *PostalServer { return &PostalServer{ endpointAddress: endpointAddress, sessionStore: sessionStore, + groupStore: groupStore, subjectStore: subjectStore, clientFactory: clientFactory, + producer: producer, } } @@ -72,6 +82,8 @@ func (s *PostalServer) Run(conf config.GrpcConfig, registry discovery.Registry) } shutdown.AddHook(cancel) + go s.processSessionStoreEvent(ctx) + // 其他服务直连地址 s.broadcastAddress = register.Addr @@ -81,6 +93,29 @@ func (s *PostalServer) Run(conf config.GrpcConfig, registry discovery.Registry) return } +func (s *PostalServer) processSessionStoreEvent(ctx context.Context) { + for { + select { + case <-ctx.Done(): + return + + case uid := <-s.sessionStore.OnStore(): + logger.Info("uid %s online", uid) + // todo mq event..., delay 5s + // s.producer.Publish() + + case channel := <-s.sessionStore.OnDelete(): + logger.Info("uid %s offline", channel.Uid) + for _, gid := range channel.Groups() { + if g, ok := s.groupStore.Load(gid); ok { + g.Leave(channel.Uid) + } + } + // todo mq event... + } + } +} + func (s *PostalServer) deliverMessage(msg *postal.Message, conn session.NetConn) error { header := &protocol.Header{ Magic: protocol.Magic, @@ -129,20 +164,30 @@ func (s *PostalServer) clusterGateReceivers(ctx context.Context, receivers []str } return } -func (s *PostalServer) Deliver(ctx context.Context, req *postal.ReqDeliver) (*postal.ResDeliver, error) { + +func (s *PostalServer) Deliver(ctx context.Context, req *postal.ReqDeliver) (res *postal.ResDeliver, err error) { receiver, ok := s.sessionStore.Load(req.Receiver) if ok { - err := s.deliverMessage(req.Msg, receiver) + var bytes []byte + bytes, err = encodeDeliverMessage(req.Msg) if err != nil { - return nil, err + return } - return &postal.ResDeliver{Ok: true}, nil + // write msg + err = receiver.Conn.Write(bytes) + if err != nil { + return + } + res = &postal.ResDeliver{Ok: true} + return } // 用户连接不在当前gateway + // todo 向一致性 hash 下一个节点传递 gateReceivers, offline := s.clusterGateReceivers(ctx, []string{req.Receiver}) if offline != nil && len(offline) > 0 { - return &postal.ResDeliver{Code: int32(postal.DeliverResult_ReceiverOffline.Number())}, nil + res = &postal.ResDeliver{Code: int32(postal.DeliverResult_ReceiverOffline.Number())} + return } for postalAddr := range gateReceivers { @@ -171,7 +216,12 @@ func (s *PostalServer) Deliver(ctx context.Context, req *postal.ReqDeliver) (*po return nil, nil } -func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverBatch) (*postal.ResDeliver, error) { +func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverBatch) (res *postal.ResDeliver, err error) { + bytes, err := encodeDeliverMessage(req.Msg) + if err != nil { + return + } + var redirectReceivers []string for _, receiverId := range req.Receivers { receiver, ok := s.sessionStore.Load(receiverId) @@ -179,7 +229,7 @@ func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverB redirectReceivers = append(redirectReceivers, receiverId) continue } - err := s.deliverMessage(req.Msg, receiver) + err = receiver.Conn.Write(bytes) if err != nil { logger.Errorf("deliver to %s error: ", receiver, err) } @@ -189,6 +239,7 @@ func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverB return &postal.ResDeliver{Ok: true}, nil } + // todo 向一致性 hash 下个节点传递 gateReceivers, offline := s.clusterGateReceivers(ctx, redirectReceivers) if len(offline) > 0 { logger.Warning("offline redirect receivers: ", offline) @@ -224,22 +275,59 @@ func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverB return &postal.ResDeliver{Ok: true}, nil } -// DeliverGroup group: -func (s *PostalServer) DeliverGroup(ctx context.Context, group *postal.ReqDeliverGroup) (*postal.ResDeliver, error) { +// DeliverGroup postal broadcast +func (s *PostalServer) DeliverGroup(ctx context.Context, req *postal.ReqDeliverGroup) (res *postal.ResDeliver, err error) { + bytes, err := encodeDeliverMessage(req.Msg) + if err != nil { + return + } - return nil, nil + res = &postal.ResDeliver{Ok: true} + g, ok := s.groupStore.Load(req.Gid) + if !ok { + return + } + g.Write(bytes) + return } -func (s *PostalServer) GroupJoin(ctx context.Context, join *postal.ReqGroupJoin) (*emptypb.Empty, error) { - return nil, nil +// GroupJoin uid consistent hash +func (s *PostalServer) GroupJoin(ctx context.Context, req *postal.ReqGroupJoin) (emp *emptypb.Empty, err error) { + channel, ok := s.sessionStore.Load(req.Uid) + if !ok { + err = errors.New("uid not online") + return + } + for _, gid := range req.Gids { + g, _ := s.groupStore.ComputeIfAbsent(gid, func(gid string) *group.Group { return group.NewGroup(gid) }) + g.Join(req.Uid, channel.Conn) + channel.GroupJoin(gid) + } + return } -func (s *PostalServer) GroupLeave(ctx context.Context, leave *postal.ReqGroupLeave) (*emptypb.Empty, error) { - return nil, nil +// GroupLeave uid consistent hash +func (s *PostalServer) GroupLeave(ctx context.Context, req *postal.ReqGroupLeave) (emp *emptypb.Empty, err error) { + channel, ok := s.sessionStore.Load(req.Uid) + if !ok { + err = errors.New("uid not online") + return + } + for _, gid := range req.Gids { + g, ok := s.groupStore.Load(gid) + if !ok { + continue + } + g.Leave(req.Uid) + channel.GroupLeave(gid) + } + return } -func (s *PostalServer) GroupDissolve(ctx context.Context, dissolve *postal.ReqGroupDissolve) (*emptypb.Empty, error) { - return nil, nil +// GroupDissolve postal broadcast +func (s *PostalServer) GroupDissolve(ctx context.Context, req *postal.ReqGroupDissolve) (emp *emptypb.Empty, err error) { + s.groupStore.Delete(req.Gid) + return } func (s *PostalServer) Endpoint(context.Context, *emptypb.Empty) (*postal.ResEndpoint, error) { diff --git a/pkg/protocol/deliver/deliver_group.go b/pkg/protocol/deliver/deliver_group.go index b8d5399..e8d447c 100644 --- a/pkg/protocol/deliver/deliver_group.go +++ b/pkg/protocol/deliver/deliver_group.go @@ -135,21 +135,32 @@ func (d *GroupDeliver) DeliverGroup(ctx context.Context, gid string, msg proto.M if err != nil { return } - req := postal.ReqDeliverGroup{ - Gid: gid, - Msg: message, - } + req := postal.ReqDeliverGroup{Gid: gid, Msg: message} for _, p := range d.postals { p.DeliverGroup(&req) } return } -func (d *GroupDeliver) GroupJoin(ctx context.Context, uid string, gids []string) (err error) { - req := &postal.ReqGroupJoin{ - Uid: uid, - Gids: gids, +func (d *GroupDeliver) GroupDissolve(gid string) { + req := postal.ReqGroupDissolve{Gid: gid} + for _, p := range d.postals { + p.GroupDissolve(&req) } +} + +func (d *GroupDeliver) GroupJoin(ctx context.Context, uid string, gids []string) (err error) { + req := &postal.ReqGroupJoin{Uid: uid, Gids: gids} + + ctx = context.WithValue(ctx, balancer.ConsistentHashKey, uid) _, 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/postal.go b/pkg/protocol/deliver/postal.go index 05f4993..bbd60b9 100644 --- a/pkg/protocol/deliver/postal.go +++ b/pkg/protocol/deliver/postal.go @@ -26,6 +26,7 @@ func (p *Postal) Close() { } } +// DeliverGroup 群消息发送 func (p *Postal) DeliverGroup(req *postal.ReqDeliverGroup) { // todo put in chan res, err := p.client.DeliverGroup(context.Background(), req) @@ -33,3 +34,12 @@ func (p *Postal) DeliverGroup(req *postal.ReqDeliverGroup) { 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/event/event.go b/pkg/protocol/event/event.go new file mode 100644 index 0000000..3fc7258 --- /dev/null +++ b/pkg/protocol/event/event.go @@ -0,0 +1,22 @@ +package event + +type Event[T any] interface { + Topic() string + Payload2Bytes(T) []byte + Bytes2Payload([]byte) T +} + +type online struct { +} + +func (o *online) Topic() string { + return "postal_event_online" +} + +func (o *online) Payload2Bytes(t string) []byte { + return []byte(t) +} + +func (o *online) Bytes2Payload(bytes []byte) string { + return string(bytes) +} diff --git a/pkg/protocol/session/net_conn.go b/pkg/protocol/session/net_conn.go deleted file mode 100644 index f6a3aeb..0000000 --- a/pkg/protocol/session/net_conn.go +++ /dev/null @@ -1,5 +0,0 @@ -package session - -type NetConn interface { - Write([]byte) error -} diff --git a/pkg/protocol/session/session.go b/pkg/protocol/session/session.go new file mode 100644 index 0000000..db7aecb --- /dev/null +++ b/pkg/protocol/session/session.go @@ -0,0 +1,51 @@ +package session + +import "sync" + +// NetConn 各类型连接的 write 接口 +type NetConn interface { + Write([]byte) error +} + +type Channel struct { + Uid string + Conn NetConn + groups []string // 记录 uid 对应的群组列表 + lock *sync.Mutex +} + +func NewChannel(uid string, conn NetConn) *Channel { + return &Channel{ + Uid: uid, + Conn: conn, + lock: &sync.Mutex{}, + } +} + +func (c *Channel) GroupJoin(gid string) { + c.lock.Lock() + defer c.lock.Unlock() + for _, g := range c.groups { + if g == gid { // 已加入 + return + } + } + c.groups = append(c.groups, gid) +} + +func (c *Channel) GroupLeave(gid string) { + c.lock.Lock() + defer c.lock.Unlock() + + for i := 0; i < len(c.groups); i++ { + if c.groups[i] == gid { // 从i前移一位 + copy(c.groups[i:], c.groups[i+1:]) + c.groups = c.groups[0 : len(c.groups)-1] + return + } + } +} + +func (c *Channel) Groups() []string { + return c.groups +} diff --git a/pkg/protocol/session/store.go b/pkg/protocol/session/store.go index 7e3ed26..32d1643 100644 --- a/pkg/protocol/session/store.go +++ b/pkg/protocol/session/store.go @@ -1,57 +1,68 @@ package session -import "sync" +import ( + "sonet/pkg/utils/collect" + "sonet/pkg/utils/logger" +) type Store interface { - Load(key string) (value NetConn, exist bool) - Delete(key string) - Store(key string, value NetConn) - Range(f func(key string, value NetConn) bool) + Load(uid string) (channel *Channel, exist bool) + Delete(uid string) + Store(uid string, channel *Channel) + Range(f func(key string, channel *Channel) bool) + OnStore() <-chan string // Store 函数调用 + OnDelete() <-chan *Channel // Delete 函数调用 } -func NewMapStore() Store { +func NewMapStore(concurrentLevel int) Store { return &mapStore{ - data: make(map[string]NetConn), + m: collect.NewConcurrentMap[string, *Channel](concurrentLevel, func(k string) string { return k }), + storeCh: make(chan string, 128), + deleteCh: make(chan *Channel, 64), } } type mapStore struct { - sync.RWMutex - data map[string]NetConn + m *collect.ConcurrentMap[string, *Channel] + storeCh chan string + deleteCh chan *Channel } -func (c *mapStore) Len() int { - c.RLock() - defer c.RUnlock() - return len(c.data) +func (c *mapStore) Load(uid string) (channel *Channel, exist bool) { + channel, exist = c.m.Load(uid) + return } -func (c *mapStore) Load(key string) (value NetConn, exist bool) { - c.RLock() - defer c.RUnlock() - value, exist = c.data[key] - return +func (c *mapStore) Delete(uid string) { + ch, ok := c.m.LoadAndDelete(uid) + if !ok { + return + } + + select { + case c.deleteCh <- ch: + default: + logger.Warning("session store delete channel fulled, %s", uid) + } } -func (c *mapStore) Delete(key string) { - c.Lock() - defer c.Unlock() - delete(c.data, key) +func (c *mapStore) Store(uid string, channel *Channel) { + c.m.Store(uid, channel) + + c.storeCh <- uid + // logger.Warningf("session store store channel fulled, %s", uid) } -func (c *mapStore) Store(key string, value NetConn) { - c.Lock() - defer c.Unlock() - c.data[key] = value +func (c *mapStore) Range(f func(key string, channel *Channel) bool) { + c.m.Range(f) } -func (c *mapStore) Range(f func(key string, value NetConn) bool) { - c.RLock() - defer c.RUnlock() +// OnStore 用户上线 +func (c *mapStore) OnStore() <-chan string { + return c.storeCh +} - for k, v := range c.data { - if !f(k, v) { - return - } - } +// OnDelete 用户离线 +func (c *mapStore) OnDelete() <-chan *Channel { + return c.deleteCh } diff --git a/pkg/utils/collect/concurrent_map.go b/pkg/utils/collect/concurrent_map.go index 54518dd..6a83f87 100644 --- a/pkg/utils/collect/concurrent_map.go +++ b/pkg/utils/collect/concurrent_map.go @@ -7,9 +7,9 @@ import ( // ConcurrentMap 分段锁 map, 提升并发性 type ConcurrentMap[K comparable, V any] struct { - hashKeyFunc func(K) string - equalsFunc func(v1, v2 V) bool - counter int64 + hashKeyFunc func(K) string + equalsFunc func(v1, v2 V) bool + // counter int64 segments int segmentsMap []map[K]V segmentsLock []*sync.RWMutex @@ -108,6 +108,35 @@ func (m *ConcurrentMap[K, V]) Swap(k K, v V) (prev V, loaded bool) { return } +// ComputeIfAbsent 加载, 如果值不存在使用 mapping(k) 填充并返回 +// mapped 值是否是 mapping(k) 填充的 +func (m *ConcurrentMap[K, V]) ComputeIfAbsent(k K, mapping func(k K) V) (res V, mapped bool) { + segment := m.segment(k) + lock := m.segmentsLock[segment] + lock.RLock() + v, ok := m.segmentsMap[segment][k] + lock.RUnlock() + if ok { + res = v + return + } + + lock.Lock() + defer lock.Unlock() + + // double check + if v, ok = m.segmentsMap[segment][k]; ok { + res = v + return + } + + // write mapping value + res = mapping(k) + m.segmentsMap[segment][k] = res + mapped = true + return +} + func fnv32Hash(k string) uint32 { f := fnv.New32() _, err := f.Write([]byte(k)) diff --git a/pkg/utils/collect/concurrent_map_test.go b/pkg/utils/collect/concurrent_map_test.go index c5fe944..f9261f9 100644 --- a/pkg/utils/collect/concurrent_map_test.go +++ b/pkg/utils/collect/concurrent_map_test.go @@ -3,6 +3,7 @@ package collect import ( "fmt" "math/rand" + "strconv" "sync" "testing" "time" @@ -54,3 +55,45 @@ func TestConcurrentMap(t *testing.T) { } wg.Wait() } + +type counter struct { + c int +} + +func (c *counter) increment() { + c.c += 1 +} + +func TestComputeIfAbsent(t *testing.T) { + cm := NewConcurrentMap[int, *counter](16, func(k int) string { return strconv.Itoa(k) }) + concurrent := 1000 + add := 10 + wg := &sync.WaitGroup{} + wg.Add(concurrent) + + for i := 0; i < concurrent; i++ { + go func() { + defer wg.Done() + + r := rand.New(rand.NewSource(time.Now().UnixMilli())) + v, mapped := cm.ComputeIfAbsent(r.Intn(10), func(k int) *counter { + return &counter{} + }) + if !mapped { + return + } + for i := 0; i < add; i++ { + v.increment() + } + }() + } + + wg.Wait() + + cm.Range(func(k int, v *counter) bool { + if v.c != add { + t.Error("error value...") + } + return true + }) +}