Browse Source

postal group api

master
tangmingyou 3 years ago
parent
commit
a618a498b0
  1. 6
      api/postal.proto
  2. 8
      internal/gateway_ws/gws_server/conn_handler.go
  3. 99
      internal/postal/group/concurrent_map.go
  4. 46
      internal/postal/group/concurrent_map_test.go
  5. 7
      internal/postal/group/group.go
  6. 4
      internal/postal/logic/postal_server.go
  7. 4
      pkg/grpc/generic/desc_source/desc_source.go

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

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

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

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

7
internal/gateway_ws/dao/group.go → 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)

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

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

Loading…
Cancel
Save