Browse Source

postal cluster dispatch

master
tangmingyou 3 years ago
parent
commit
6f40a94ec8
  1. 148
      benchmark/main.go
  2. 3
      build-all.bat
  3. 2
      cmd/chat/main.go
  4. 2
      cmd/gateway_ws/main.go
  5. 45
      internal/gateway_ws/gws_server/conn_handler.go
  6. 15
      internal/postal/group/group.go
  7. 10
      internal/postal/logic/postal_cluster_server.go
  8. 34
      internal/postal/logic/postal_server.go
  9. 8
      pkg/grpc/discovery/discovery.go
  10. 59
      pkg/grpc/discovery/etcd_naming.go
  11. 4
      pkg/protocol/deliver/group_deliver.go
  12. 23
      pkg/protocol/session/channel.go

148
benchmark/main.go

@ -3,33 +3,46 @@ package main
import ( import (
"context" "context"
"encoding/base64" "encoding/base64"
"encoding/json"
"fmt" "fmt"
"github.com/bytedance/sonic" "github.com/bytedance/sonic"
"github.com/gorilla/websocket" "github.com/gorilla/websocket"
clientv3 "go.etcd.io/etcd/client/v3"
"google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure"
"google.golang.org/protobuf/proto" "google.golang.org/protobuf/proto"
"io"
"math/rand" "math/rand"
"net" "net"
"net/http"
"runtime" "runtime"
"sonet/api/gen/auth" "sonet/api/gen/auth"
"sonet/api/gen/chat" "sonet/api/gen/chat"
"sonet/pkg/config" "sonet/pkg/config"
"sonet/pkg/grpc/discovery"
"sonet/pkg/protocol" "sonet/pkg/protocol"
"sonet/pkg/protocol/deliver"
"sonet/pkg/utils/logger" "sonet/pkg/utils/logger"
"sonet/pkg/utils/security" "sonet/pkg/utils/security"
"sonet/pkg/utils/shutdown" "sonet/pkg/utils/shutdown"
"strconv" "strconv"
"strings"
"sync" "sync"
"sync/atomic" "sync/atomic"
"time" "time"
) )
var ( var (
mockUsers = 100 benchmarkMode = "deliver" // deliver / groupDeliver
eachUserSend = 100000 mockUsers = 2000
eachUserSend = 100
mockNetUsers []*NetUser mockNetUsers []*NetUser
sendCounter int64 = 0 sendCounter int64 = 0
receiverCounter int64 = 0 receiverCounter int64 = 0
wsUrls = []string{"ws://192.168.110.36:7001", "ws://192.168.110.36:7003"} //wsUrls = []string{"ws://192.168.110.36:7001", "ws://192.168.110.36:7003"}
gatewayHttp = "http://192.168.110.36:7000"
httpClient *http.Client
etcdClient *clientv3.Client
seqId int32 seqId int32
callbacks map[int32]func(res any) callbacks map[int32]func(res any)
callbackMutex *sync.Mutex callbackMutex *sync.Mutex
@ -42,10 +55,21 @@ func init() {
callbacks = make(map[int32]func(res any), 128) callbacks = make(map[int32]func(res any), 128)
callbackMutex = &sync.Mutex{} callbackMutex = &sync.Mutex{}
httpClient = &http.Client{Timeout: 10 * time.Second}
var err error
etcdClient, err = clientv3.New(clientv3.Config{
Endpoints: []string{"124.222.131.236:3279"},
Username: "root",
Password: "sopod@etcd",
})
if err != nil {
panic(err)
}
} }
func main() { func main() {
benchmark() // benchmark()
benchmarkGroup()
} }
func records(ctx context.Context) { func records(ctx context.Context) {
@ -83,11 +107,15 @@ func benchmark() {
// initial uids // initial uids
for i := 0; i < mockUsers; i++ { for i := 0; i < mockUsers; i++ {
uid := strconv.Itoa(110000 + i) uid := strconv.Itoa(110000 + i)
conn, err := getConn() token, err := getToken(uid)
if err != nil { if err != nil {
panic(err) panic(err)
} }
err = handleConn(ctx, uid, conn) conn, err := getConn(token)
if err != nil {
panic(err)
}
err = handleConn(ctx, uid, token, conn)
if err != nil { if err != nil {
panic(err) panic(err)
} }
@ -108,6 +136,75 @@ func benchmark() {
logger.Infof("total receiver:%d, send:%d\n", receiverCounter, sendCounter) logger.Infof("total receiver:%d, send:%d\n", receiverCounter, sendCounter)
} }
func benchmarkGroup() {
benchmarkMode = "groupDeliver"
runtime.GOMAXPROCS(runtime.NumCPU())
ctx, cancel := context.WithCancel(context.Background())
shutdown.AddHook(cancel)
// postal group deliver
dis := discovery.NewEtcdDiscovery(etcdClient)
picker := deliver.NewPostalPicker(dis, grpc.WithTransportCredentials(insecure.NewCredentials()))
if err := picker.Init(ctx); err != nil {
panic(err)
}
groupDeliver := deliver.NewGroupDeliver(chat.Chat_ServiceDesc.ServiceName, picker, nil, nil)
if err := groupDeliver.Init(ctx); err != nil {
panic(err)
}
groupId := "9527"
// initial mock users, join to postal group
mockNetUsers = make([]*NetUser, mockUsers)
for i := 0; i < mockUsers; i++ {
uid := strconv.Itoa(110000 + i)
token, err := getToken(uid)
if err != nil {
panic(err)
}
conn, err := getConn(token)
if err != nil {
panic(err)
}
err = handleConn(ctx, uid, token, conn)
if err != nil {
panic(err)
}
mockNetUsers[i] = &NetUser{
Uid: uid,
Conn: conn,
}
// join group
err = groupDeliver.GroupJoin(context.Background(), uid, []string{groupId})
if err != nil {
panic(err)
}
}
shutdown.AddHook(func() {
if err := groupDeliver.GroupDissolve(context.Background(), groupId); err != nil {
logger.Error("dissolve group error: ", err)
return
}
logger.Info("test group dissolved")
})
go records(ctx)
// send group message
message := &chat.ChatMessage{Sender: "100001", Content: "hello"}
for i := 0; i < eachUserSend; i++ {
err := groupDeliver.DeliverGroup(context.Background(), groupId, message)
if err != nil {
logger.Error("deliver group error:", err)
}
atomic.AddInt64(&sendCounter, 1)
}
shutdown.Await()
}
func sendBatchChatMessage(ctx context.Context, uid string, conn *websocket.Conn, count int) { func sendBatchChatMessage(ctx context.Context, uid string, conn *websocket.Conn, count int) {
r := rand.New(rand.NewSource(time.Now().UnixMilli())) r := rand.New(rand.NewSource(time.Now().UnixMilli()))
for i := 0; i < count; i++ { for i := 0; i < count; i++ {
@ -132,9 +229,30 @@ func sendBatchChatMessage(ctx context.Context, uid string, conn *websocket.Conn,
} }
} }
func getConn() (conn *websocket.Conn, err error) { func getConn(token string) (conn *websocket.Conn, err error) {
wsUrl := wsUrls[rand.Intn(len(wsUrls))] req, err := http.NewRequest("GET", gatewayHttp+"/api/lb/ws", io.LimitReader(nil, 0))
conn, _, err = websocket.DefaultDialer.Dial(wsUrl+"/ws", nil) if err != nil {
return
}
req.Header.Set("Authorization", token)
resp, err := httpClient.Do(req)
if err != nil {
logger.Error("get ws endpoint error: ", err)
return
}
defer resp.Body.Close()
bytes, err := io.ReadAll(resp.Body)
if err != nil {
return
}
body := make(map[string]any, 2)
err = json.Unmarshal(bytes, &body)
if err != nil {
return
}
wsUrl := body["data"].(map[string]any)["ws"].(string)
// wsUrl := wsUrls[rand.Intn(len(wsUrls))]
conn, _, err = websocket.DefaultDialer.Dial("ws://"+wsUrl+"/ws", nil)
if err != nil { if err != nil {
return return
} }
@ -161,14 +279,10 @@ func getToken(uid string) (token string, err error) {
} }
// handleConn listen and auth verify // handleConn listen and auth verify
func handleConn(ctx context.Context, uid string, conn *websocket.Conn) (err error) { func handleConn(ctx context.Context, uid string, token string, conn *websocket.Conn) (err error) {
go listen(ctx, conn) go listen(ctx, conn)
// handshake // handshake
token, err := getToken(uid)
if err != nil {
return
}
args := &auth.ReqVerify{Token: token} args := &auth.ReqVerify{Token: token}
channel, err := send(conn, auth.Auth_ServiceDesc.ServiceName, "Verify", args) channel, err := send(conn, auth.Auth_ServiceDesc.ServiceName, "Verify", args)
@ -186,7 +300,7 @@ func handleConn(ctx context.Context, uid string, conn *websocket.Conn) (err erro
func listen(ctx context.Context, conn *websocket.Conn) { func listen(ctx context.Context, conn *websocket.Conn) {
var e error var e error
defer func() { defer func() {
if e != nil { if e != nil && !strings.Contains(e.Error(), "close") {
fmt.Println("conn error: ", e) fmt.Println("conn error: ", e)
} }
}() }()
@ -221,6 +335,10 @@ func listen(ctx context.Context, conn *websocket.Conn) {
} }
header := payload.Header header := payload.Header
if benchmarkMode == "groupDeliver" {
atomic.AddInt64(&receiverCounter, 1)
}
if header.Type == 4 { if header.Type == 4 {
callbackMutex.Lock() callbackMutex.Lock()
callback, ok := callbacks[header.SeqId] callback, ok := callbacks[header.SeqId]

3
build-all.bat

@ -0,0 +1,3 @@
@title build all sonet
build.bat gateway_ws && build.bat gateway_http && build.bat auth && build.bat chat && build.bat mahjong

2
cmd/chat/main.go

@ -42,7 +42,7 @@ func main() {
if err != nil { if err != nil {
panic(err) panic(err)
} }
groupDeli := deliver.NewGroupDeliver(chat.Chat_ServiceDesc.ServiceName, logic.NewRoomLoader(), consumer) groupDeli := deliver.NewGroupDeliver(chat.Chat_ServiceDesc.ServiceName, picker, logic.NewRoomLoader(), consumer)
chatServer := logic.NewChatServer(deli, groupDeli) chatServer := logic.NewChatServer(deli, groupDeli)
go func() { go func() {

2
cmd/gateway_ws/main.go

@ -5,6 +5,7 @@ import (
clientv3 "go.etcd.io/etcd/client/v3" clientv3 "go.etcd.io/etcd/client/v3"
"google.golang.org/grpc" "google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure" "google.golang.org/grpc/credentials/insecure"
"runtime"
"sonet/internal/gateway_ws/gws_server" "sonet/internal/gateway_ws/gws_server"
"sonet/internal/postal/group" "sonet/internal/postal/group"
"sonet/internal/postal/logic" "sonet/internal/postal/logic"
@ -31,6 +32,7 @@ type GatewayWsConfig struct {
// websocket server with postalService // websocket server with postalService
func main() { func main() {
runtime.GOMAXPROCS(runtime.NumCPU())
appConf := &GatewayWsConfig{} appConf := &GatewayWsConfig{}
conf := config.LoadConfig(appConf, "cmd/gateway_ws") conf := config.LoadConfig(appConf, "cmd/gateway_ws")

45
internal/gateway_ws/gws_server/conn_handler.go

@ -7,6 +7,7 @@ import (
"google.golang.org/protobuf/proto" "google.golang.org/protobuf/proto"
"regexp" "regexp"
"sonet/api/gen/auth" "sonet/api/gen/auth"
"sonet/api/gen/postal"
"sonet/pkg/grpc/generic" "sonet/pkg/grpc/generic"
"sonet/pkg/protocol" "sonet/pkg/protocol"
"sonet/pkg/protocol/session" "sonet/pkg/protocol/session"
@ -140,14 +141,16 @@ func (c *GwsHandler) OnMessage(socket *gws.Conn, message *gws.Message) {
// auth verify success // auth verify success
if service == "Auth" && method == "Verify" { if service == "Auth" && method == "Verify" {
sessionUid, err = c.extractAuthVerify(conn, payload.Body) var channel *session.Channel
sessionUid, channel, err = c.extractAuthVerify(conn, payload.Body)
if err != nil { if err != nil {
logger.Error("extractAuthVerify error: ", err) logger.Error("extractAuthVerify error: ", err)
writeError(conn, header.SeqId, "server error") writeError(conn, header.SeqId, "server error")
return return
} }
socket.Session().Store(SessionUidKey, sessionUid) socket.Session().Store(SessionUidKey, sessionUid)
// todo online event // channel
go c.dispatch(channel)
} }
header.Type = protocol.TypeResponse header.Type = protocol.TypeResponse
@ -163,7 +166,40 @@ func (c *GwsHandler) OnMessage(socket *gws.Conn, message *gws.Message) {
} }
} }
func (c *GwsHandler) extractAuthVerify(conn *gwsConn, resp []byte) (sessionUid string, err error) { 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
}
// todo close
func (c *GwsHandler) dispatch(channel *session.Channel) {
for {
msg := channel.Ready()
bytes, err := encodeDeliverMessage(msg)
if err != nil {
logger.Error("encode dispatch message error:", err)
continue
}
err = channel.Conn.Write(bytes)
if err != nil {
logger.Error("write dispatch message error:", err)
}
}
}
func (c *GwsHandler) extractAuthVerify(conn *gwsConn, resp []byte) (sessionUid string, channel *session.Channel, err error) {
authSubject := &auth.Subject{} authSubject := &auth.Subject{}
err = proto.Unmarshal(resp, authSubject) err = proto.Unmarshal(resp, authSubject)
if err != nil { if err != nil {
@ -180,7 +216,8 @@ func (c *GwsHandler) extractAuthVerify(conn *gwsConn, resp []byte) (sessionUid s
logger.Infof("close subject old conn: %s\n", sessionUid) logger.Infof("close subject old conn: %s\n", sessionUid)
} }
// 存储session到内存 // 存储session到内存
c.sessionStore.Store(sessionUid, session.NewChannel(sessionUid, conn)) channel = session.NewChannel(sessionUid, conn)
c.sessionStore.Store(sessionUid, channel)
logger.Infof("subject online: %s\n", sessionUid) logger.Infof("subject online: %s\n", sessionUid)
return return
} }

15
internal/postal/group/group.go

@ -3,6 +3,7 @@ package group
import ( import (
"errors" "errors"
"github.com/redis/go-redis/v9" "github.com/redis/go-redis/v9"
"sonet/api/gen/postal"
"sonet/pkg/protocol/session" "sonet/pkg/protocol/session"
"sonet/pkg/utils/logger" "sonet/pkg/utils/logger"
"sync" "sync"
@ -36,27 +37,27 @@ func (d *RedisPostalDao) LoadGroupIdsByUid(uid string) (gids []string, err error
type Group struct { type Group struct {
sync.RWMutex sync.RWMutex
Gid string Gid string
uids map[string]session.NetConn uids map[string]*session.Channel
} }
func NewGroup(gid string) *Group { func NewGroup(gid string) *Group {
return &Group{ return &Group{
Gid: gid, Gid: gid,
uids: make(map[string]session.NetConn), uids: make(map[string]*session.Channel, 6),
} }
} }
func (g *Group) Write(data []byte) { func (g *Group) Write(msg *postal.Message) {
g.RLock() g.RLock()
defer g.RUnlock() defer g.RUnlock()
for uid, conn := range g.uids { for uid, channel := range g.uids {
if err := conn.Write(data); err != nil { if err := channel.Push(msg); err != nil {
logger.Errorf("group send %s.%s error: ", g.Gid, uid, err) logger.Errorf("group send %s.%s error: ", g.Gid, uid, err)
} }
} }
} }
func (g *Group) Join(uid string, conn session.NetConn) { func (g *Group) Join(uid string, conn *session.Channel) {
g.Lock() g.Lock()
defer g.Unlock() defer g.Unlock()
g.uids[uid] = conn g.uids[uid] = conn
@ -74,7 +75,7 @@ func (g *Group) Dismiss() {
} }
func (g *Group) Load(uid string) (conn session.NetConn, ok bool) { func (g *Group) Load(uid string) (conn *session.Channel, ok bool) {
g.RLock() g.RLock()
defer g.RUnlock() defer g.RUnlock()
conn, ok = g.uids[uid] conn, ok = g.uids[uid]

10
internal/postal/logic/postal_cluster_server.go

@ -69,10 +69,10 @@ func (s *PostalClusterServer) Run(postalAddr string, opts ...grpc.ServerOption)
} }
func (s *PostalClusterServer) Redeliver(ctx context.Context, req *postal.ReqRedeliver) (res *postal.ResDeliver, err error) { func (s *PostalClusterServer) Redeliver(ctx context.Context, req *postal.ReqRedeliver) (res *postal.ResDeliver, err error) {
bytes, err := encodeDeliverMessage(req.Msg) //bytes, err := encodeDeliverMessage(req.Msg)
if err != nil { //if err != nil {
return // return
} //}
var redelivers []string var redelivers []string
for _, receiver := range req.Receivers { for _, receiver := range req.Receivers {
channel, ok := s.sessionStore.Load(receiver) channel, ok := s.sessionStore.Load(receiver)
@ -80,7 +80,7 @@ func (s *PostalClusterServer) Redeliver(ctx context.Context, req *postal.ReqRede
redelivers = append(redelivers, receiver) redelivers = append(redelivers, receiver)
continue continue
} }
err = channel.Conn.Write(bytes) err = channel.Push(req.Msg)
if err != nil { if err != nil {
logger.Errorf("redeliver channel %s write error: %v", receiver, err) logger.Errorf("redeliver channel %s write error: %v", receiver, err)
continue continue

34
internal/postal/logic/postal_server.go

@ -131,13 +131,13 @@ func (s *PostalServer) processSessionStoreEvent(ctx context.Context) {
func (s *PostalServer) Deliver(ctx context.Context, req *postal.ReqDeliver) (res *postal.ResDeliver, err error) { func (s *PostalServer) Deliver(ctx context.Context, req *postal.ReqDeliver) (res *postal.ResDeliver, err error) {
receiver, ok := s.sessionStore.Load(req.Receiver) receiver, ok := s.sessionStore.Load(req.Receiver)
if ok { if ok {
var bytes []byte //var bytes []byte
bytes, err = encodeDeliverMessage(req.Msg) //bytes, err = encodeDeliverMessage(req.Msg)
if err != nil { //if err != nil {
return // return
} //}
// write msg // write msg
err = receiver.Conn.Write(bytes) err = receiver.Push(req.Msg)
if err != nil { if err != nil {
return return
} }
@ -161,10 +161,10 @@ func (s *PostalServer) Deliver(ctx context.Context, req *postal.ReqDeliver) (res
} }
func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverBatch) (res *postal.ResDeliver, err error) { func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverBatch) (res *postal.ResDeliver, err error) {
bytes, err := encodeDeliverMessage(req.Msg) //bytes, err := encodeDeliverMessage(req.Msg)
if err != nil { //if err != nil {
return // return
} //}
var redeliverReceivers []string var redeliverReceivers []string
for _, receiverId := range req.Receivers { for _, receiverId := range req.Receivers {
@ -173,7 +173,7 @@ func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverB
redeliverReceivers = append(redeliverReceivers, receiverId) redeliverReceivers = append(redeliverReceivers, receiverId)
continue continue
} }
err = receiver.Conn.Write(bytes) err = receiver.Push(req.Msg)
if err != nil { if err != nil {
logger.Errorf("deliver to %s error: ", receiver, err) logger.Errorf("deliver to %s error: ", receiver, err)
} }
@ -200,17 +200,17 @@ func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverB
// DeliverGroup postal broadcast // DeliverGroup postal broadcast
func (s *PostalServer) DeliverGroup(ctx context.Context, req *postal.ReqDeliverGroup) (res *postal.ResDeliver, err error) { func (s *PostalServer) DeliverGroup(ctx context.Context, req *postal.ReqDeliverGroup) (res *postal.ResDeliver, err error) {
bytes, err := encodeDeliverMessage(req.Msg) //bytes, err := encodeDeliverMessage(req.Msg)
if err != nil { //if err != nil {
return // return
} //}
res = &postal.ResDeliver{Ok: true} res = &postal.ResDeliver{Ok: true}
g, ok := s.groupStore.Load(req.Gid) g, ok := s.groupStore.Load(req.Gid)
if !ok { if !ok {
return return
} }
g.Write(bytes) g.Write(req.Msg)
return return
} }
@ -223,7 +223,7 @@ func (s *PostalServer) GroupJoin(ctx context.Context, req *postal.ReqGroupJoin)
} }
for _, gid := range req.Gids { for _, gid := range req.Gids {
g, _ := s.groupStore.ComputeIfAbsent(gid, func(gid string) *group.Group { return group.NewGroup(gid) }) g, _ := s.groupStore.ComputeIfAbsent(gid, func(gid string) *group.Group { return group.NewGroup(gid) })
g.Join(req.Uid, channel.Conn) g.Join(req.Uid, channel)
channel.GroupJoin(gid) channel.GroupJoin(gid)
} }
return return

8
pkg/grpc/discovery/discovery.go

@ -3,6 +3,7 @@ package discovery
import ( import (
"context" "context"
"google.golang.org/grpc/resolver" "google.golang.org/grpc/resolver"
"sonet/pkg/utils/logger"
"strconv" "strconv"
) )
@ -47,8 +48,11 @@ func (s Server) GetWeight() (weight int) {
if !ok { if !ok {
return return
} }
if w, err := strconv.Atoi(v); err != nil { w, err := strconv.Atoi(v)
weight = w if err != nil {
logger.Warning("failed parse discovery server attr weight: ", v)
return
} }
weight = w
return return
} }

59
pkg/grpc/discovery/etcd_naming.go

@ -75,17 +75,70 @@ func (r *EtcdDiscovery) Registry(ctx context.Context, server Server) (err error)
) )
// keepalive lease // keepalive lease
keepAliveCh, err := r.client.KeepAlive(context.Background(), lease.ID) keepCtx, keepCancel := context.WithCancel(context.Background())
keepAliveCh, err := r.client.KeepAlive(keepCtx, lease.ID)
if err != nil {
keepCancel()
logger.Error("registry keepalive error: ", err)
return
}
// ticker := time.NewTicker(time.Second * time.Duration(DefaultRegisterTTL))
w := r.client.Watch(ctx, endpointKey)
go func() { go func() {
for { for {
select { select {
case <-ctx.Done(): case <-ctx.Done():
logger.Info("registry keepalive done") logger.Info("registry keepalive done")
ctx, c := context.WithTimeout(context.Background(), time.Second*2) ctx, c := context.WithTimeout(context.Background(), time.Second*3)
_, _ = r.client.Revoke(ctx, lease.ID) _, _ = r.client.Revoke(ctx, lease.ID)
c() c()
keepCancel()
//ticker.Stop()
return return
case _ = <-keepAliveCh: case <-keepAliveCh:
//logger.Infof("keepalive: lease id %d, %+v", lease.ID, res)
case res := <-w:
if err := res.Err(); err != nil {
logger.Errorf("registry watch endpoint %s error: %v", endpointKey, err)
continue
}
deleted := false
for _, event := range res.Events {
if event.Type == clientv3.EventTypeDelete {
deleted = true
}
}
// endpoint key被删除, 重放
if deleted {
logger.Infof("registry endpoint %s deleted ", endpointKey)
// 删除旧的 lease
keepCancel()
_, _ = r.client.Revoke(context.Background(), lease.ID)
lease, err = r.client.Grant(ctx, DefaultRegisterTTL)
if err != nil {
logger.Error("registry grant lease again error: ", err)
}
err = em.AddEndpoint(ctx,
endpointKey,
endpoints.Endpoint{
Addr: addr,
Metadata: meta,
},
clientv3.WithLease(lease.ID),
)
if err != nil {
logger.Error("refresh endpoint error: ", err)
continue
}
// keepalive lease
keepCtx, keepCancel = context.WithCancel(context.Background())
keepAliveCh, err = r.client.KeepAlive(keepCtx, lease.ID)
if err != nil {
logger.Error("registry keepalive error: ", err)
}
}
//case <-ticker.C: // 定时重放防止etcd中key被删除
} }
} }
}() }()

4
pkg/protocol/deliver/group_deliver.go

@ -25,18 +25,22 @@ type GroupDeliver struct {
func NewGroupDeliver( func NewGroupDeliver(
msgInServiceName string, msgInServiceName string,
postalPicker *PostalPicker,
groupLoader GroupLoader, groupLoader GroupLoader,
consumer mq.Consumer, consumer mq.Consumer,
) *GroupDeliver { ) *GroupDeliver {
return &GroupDeliver{ return &GroupDeliver{
svcName: msgInServiceName, svcName: msgInServiceName,
postalPicker: postalPicker,
groupLoader: groupLoader, groupLoader: groupLoader,
consumer: consumer, consumer: consumer,
} }
} }
func (d *GroupDeliver) Init(ctx context.Context) (err error) { func (d *GroupDeliver) Init(ctx context.Context) (err error) {
if d.consumer != nil {
err = d.subscribe() err = d.subscribe()
}
return return
} }

23
pkg/protocol/session/session.go → pkg/protocol/session/channel.go

@ -1,6 +1,12 @@
package session package session
import "sync" import (
"errors"
"sonet/api/gen/postal"
"sync"
)
var ErrChannelFullMsgDropped = errors.New("channel full, msg dropped")
// NetConn 各类型连接的 write 接口 // NetConn 各类型连接的 write 接口
type NetConn interface { type NetConn interface {
@ -10,6 +16,7 @@ type NetConn interface {
type Channel struct { type Channel struct {
Uid string Uid string
Conn NetConn Conn NetConn
ch chan *postal.Message
groups []string // 记录 uid 对应的群组列表 groups []string // 记录 uid 对应的群组列表
lock *sync.Mutex lock *sync.Mutex
} }
@ -18,10 +25,24 @@ func NewChannel(uid string, conn NetConn) *Channel {
return &Channel{ return &Channel{
Uid: uid, Uid: uid,
Conn: conn, Conn: conn,
ch: make(chan *postal.Message, 16),
lock: &sync.Mutex{}, lock: &sync.Mutex{},
} }
} }
func (c *Channel) Push(msg *postal.Message) (err error) {
select {
case c.ch <- msg:
default:
err = ErrChannelFullMsgDropped
}
return
}
func (c *Channel) Ready() *postal.Message {
return <-c.ch
}
func (c *Channel) GroupJoin(gid string) { func (c *Channel) GroupJoin(gid string) {
c.lock.Lock() c.lock.Lock()
defer c.lock.Unlock() defer c.lock.Unlock()
Loading…
Cancel
Save