diff --git a/api/postal.proto b/api/postal.proto index 548c066..460a1f1 100644 --- a/api/postal.proto +++ b/api/postal.proto @@ -13,7 +13,7 @@ service Postal { rpc DeliverBatch(ReqDeliverBatch) returns(ResDeliver); rpc DeliverGroup(ReqDeliverGroup) returns(ResDeliver); - rpc GroupCreate(ReqGroupCreate) returns(ResGroupCreate); + // rpc GroupCreate(ReqGroupCreate) returns(ResGroupCreate); rpc GroupJoin(ReqGroupJoin) returns(google.protobuf.Empty); rpc GroupLeave(ReqGroupLeave) returns(google.protobuf.Empty); rpc GroupDissolve(ReqGroupDissolve) returns(google.protobuf.Empty); @@ -67,12 +67,12 @@ message ResGroupCreate { message ReqGroupJoin { string gid = 1; - string uid = 2; + repeated string uid = 2; } message ReqGroupLeave { string gid = 1; - string uid = 2; + repeated string uid = 2; } message ReqGroupDissolve { @@ -82,7 +82,7 @@ message ReqGroupDissolve { // 集群之间接口调用,重定向消息... service PostalCluster { // socket不在当前节点,重新投递消息 - rpc Redirect(ReqRedirect) returns(ResDeliver); + rpc Redirect(ReqRedirect) returns(ResDeliver); } message ReqRedirect { diff --git a/internal/gateway_ws/gws_server/conn_handler.go b/internal/gateway_ws/gws_server/conn_handler.go index 2c9d891..bcf305e 100644 --- a/internal/gateway_ws/gws_server/conn_handler.go +++ b/internal/gateway_ws/gws_server/conn_handler.go @@ -50,11 +50,15 @@ func (c *GwsHandler) OnClose(socket *gws.Conn, err error) { if err != nil { logger.Error("gws conn close error: ", err) } - uid, ok := socket.Session().Load(SessionUidKey) + val, ok := socket.Session().Load(SessionUidKey) if !ok { return } - c.sessionStore.Delete(uid.(string)) + uid := val.(string) + c.sessionStore.Delete(uid) + if err = c.subjectStore.Del(context.Background(), uid); err != nil { + logger.Error("del offline subject store error: ", err) + } logger.Info("subject offline: ", uid) } diff --git a/internal/postal/group/concurrent_map.go b/internal/postal/group/concurrent_map.go new file mode 100644 index 0000000..e5f3213 --- /dev/null +++ b/internal/postal/group/concurrent_map.go @@ -0,0 +1,99 @@ +package group + +import ( + "hash/fnv" + "sync" +) + +// ConcurrentMap 分段锁 map, 提升并发性 +type ConcurrentMap[K comparable, V any] struct { + hashKeyFunc func(K) string + equalsFunc func(v1, v2 V) bool + counter int64 + segments int + segmentsMap []map[K]V + segmentsLock []*sync.RWMutex +} + +// NewConcurrentMap 分段锁并发 map +// segments: 分段数 +// hashKeyFunc: key转string函数 +// equalsFunc: value比较函数, Put时新旧值相同则不返回旧值 +func NewConcurrentMap[K comparable, V any](segments int, hashKeyFunc func(K) string) *ConcurrentMap[K, V] { + m := &ConcurrentMap[K, V]{ + hashKeyFunc: hashKeyFunc, + segments: segments, + segmentsMap: make([]map[K]V, segments), + segmentsLock: make([]*sync.RWMutex, segments), + } + for i := 0; i < segments; i++ { + m.segmentsMap[i] = make(map[K]V, 16) + m.segmentsLock[i] = &sync.RWMutex{} + } + return m +} + +// segment 根据 key 确定分段 +func (m *ConcurrentMap[K, V]) segment(k K) int { + hashK := m.hashKeyFunc(k) + hash := fnv32Hash(hashK) + return int(hash) % m.segments +} + +// Store 放置新值 +func (m *ConcurrentMap[K, V]) Store(k K, v V) { // (old V, hasOld bool) // 返回旧值 + segment := m.segment(k) + lock := m.segmentsLock[segment] + lock.Lock() + defer lock.Unlock() + //if prev, ok := m.segmentsMap[segment][k]; ok { + // if !m.equalsFunc(v, prev) { // 两值不同返回旧值 + // old, hasOld = prev, true + // } + //} + m.segmentsMap[segment][k] = v + return +} + +func (m *ConcurrentMap[K, V]) Load(k K) (v V, ok bool) { + segment := m.segment(k) + lock := m.segmentsLock[segment] + lock.RLock() + defer lock.RUnlock() + v, ok = m.segmentsMap[segment][k] + return +} + +func (m *ConcurrentMap[K, V]) Delete(k K) { + segment := m.segment(k) + lock := m.segmentsLock[segment] + lock.Lock() + defer lock.Unlock() + delete(m.segmentsMap[segment], k) +} + +func (m *ConcurrentMap[K, V]) Range(f func(key, value any) bool) { + for i := 0; i < m.segments; i++ { + lock := m.segmentsLock[i] + func() { + lock.RLock() + defer lock.RUnlock() + + for k, v := range m.segmentsMap[i] { + if !f(k, v) { + i = m.segments // stop range + return + } + } + }() + } +} + +func fnv32Hash(k string) uint32 { + f := fnv.New32() + _, err := f.Write([]byte(k)) + if err != nil { + panic(err) + } + return f.Sum32() +} diff --git a/internal/postal/group/concurrent_map_test.go b/internal/postal/group/concurrent_map_test.go new file mode 100644 index 0000000..841adb2 --- /dev/null +++ b/internal/postal/group/concurrent_map_test.go @@ -0,0 +1,46 @@ +package group + +import ( + "fmt" + "math/rand" + "sync" + "testing" + "time" +) + +func TestConcurrentMap(t *testing.T) { + cm := NewConcurrentMap[string, string](16, func(k string) string { return k }) + concurrent := 1000 + + wg := sync.WaitGroup{} + wg.Add(concurrent) + + for i := 0; i < concurrent; i++ { + go func(loop int) { + for j := 0; j < concurrent; j++ { + k := fmt.Sprintf("%d-%d", loop, j) + cm.Store(k, k) + } + wg.Done() + }(i) + } + wg.Wait() + + wg.Add(concurrent) + + for i := 0; i < concurrent; i++ { + go func() { + r := rand.New(rand.NewSource(time.Now().UnixMilli())) + for i := 0; i < concurrent; i++ { + k := fmt.Sprintf("%d-%d", r.Intn(concurrent), r.Intn(concurrent)) + v, ok := cm.Load(k) + if !ok || v != k { + t.Errorf("value load error: %s", k) + } + } + wg.Done() + }() + } + + wg.Wait() +} diff --git a/internal/gateway_ws/dao/group.go b/internal/postal/group/group.go similarity index 62% rename from internal/gateway_ws/dao/group.go rename to internal/postal/group/group.go index bf8b86d..e84a706 100644 --- a/internal/gateway_ws/dao/group.go +++ b/internal/postal/group/group.go @@ -1,4 +1,4 @@ -package dao +package group import ( "errors" @@ -8,6 +8,11 @@ import ( var ErrNotExists = errors.New("not exists") // PostalDao persistent postal group api +// 1.uid online, send online event +// 2.onOnlineEvent: chat -> chat groups to postal, game -> game groups to postal +// 2.postal create uid create or join memory group linkedList +// 3.postalA: gid-1(uid1, uid2, uid3); postalB: gid1(uid4, uid5, uid6) +// extra: room bit位判断是否存在 type PostalDao interface { LoadGroupIdsByUid(uid string) ([]string, error) GroupCreate(gid string) (string, error) diff --git a/internal/postal/logic/postal_server.go b/internal/postal/logic/postal_server.go index eeb1ad9..4830010 100644 --- a/internal/postal/logic/postal_server.go +++ b/internal/postal/logic/postal_server.go @@ -222,11 +222,9 @@ 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) { - return nil, nil -} -func (s *PostalServer) GroupCreate(ctx context.Context, create *postal.ReqGroupCreate) (*postal.ResGroupCreate, error) { return nil, nil } diff --git a/pkg/grpc/generic/desc_source/desc_source.go b/pkg/grpc/generic/desc_source/desc_source.go index 1ea0895..2707472 100644 --- a/pkg/grpc/generic/desc_source/desc_source.go +++ b/pkg/grpc/generic/desc_source/desc_source.go @@ -4,7 +4,7 @@ import ( "context" "errors" "fmt" - "io/ioutil" + "os" "sync" "github.com/golang/protobuf/proto" @@ -40,7 +40,7 @@ type DescriptorSource interface { func DescriptorSourceFromProtoSets(fileNames ...string) (DescriptorSource, error) { files := &descpb.FileDescriptorSet{} for _, fileName := range fileNames { - b, err := ioutil.ReadFile(fileName) + b, err := os.ReadFile(fileName) if err != nil { return nil, fmt.Errorf("could not load protoset file %q: %v", fileName, err) }