Browse Source

deliver group

master
tangmingyou 3 years ago
parent
commit
a664e186b0
  1. 44
      benchmark/main.go
  2. 17
      cmd/gateway_ws/main.go
  3. 8
      internal/gateway_ws/gws_server/conn_handler.go
  4. 52
      internal/postal/group/group.go
  5. 24
      internal/postal/logic/postal.go
  6. 124
      internal/postal/logic/postal_server.go
  7. 27
      pkg/protocol/deliver/deliver_group.go
  8. 10
      pkg/protocol/deliver/postal.go
  9. 22
      pkg/protocol/event/event.go
  10. 5
      pkg/protocol/session/net_conn.go
  11. 51
      pkg/protocol/session/session.go
  12. 77
      pkg/protocol/session/store.go
  13. 31
      pkg/utils/collect/concurrent_map.go
  14. 43
      pkg/utils/collect/concurrent_map_test.go

44
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")

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

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

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

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

124
internal/postal/logic/postal_server.go

@ -2,18 +2,22 @@ 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"
)
@ -22,21 +26,27 @@ type PostalServer struct {
postal.UnimplementedPostalServer
endpointAddress string // websocket 前端连接地址
broadcastAddress string // postal server 在注册中心注册的地址, 其他服务可直连访问
sessionStore session.Store // 在线用户conn存储
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) {

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

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

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

5
pkg/protocol/session/net_conn.go

@ -1,5 +0,0 @@
package session
type NetConn interface {
Write([]byte) error
}

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

77
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]
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
}

31
pkg/utils/collect/concurrent_map.go

@ -9,7 +9,7 @@ import (
type ConcurrentMap[K comparable, V any] struct {
hashKeyFunc func(K) string
equalsFunc func(v1, v2 V) bool
counter int64
// 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))

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

Loading…
Cancel
Save