Compare commits

...

5 Commits

Author SHA1 Message Date
tangmingyou aa943bd761 benchmark group bugfix 3 years ago
tangmingyou 0319c33631 gateway dispatch 3 years ago
tangmingyou ea9b49c82f gnet gateway 3 years ago
tangmingyou ab143b0c3f gnet gateway 3 years ago
tangmingyou eaad4b2da9 benchmark 3 years ago
  1. 8
      api/chat.proto
  2. 249
      benchmark/main.go
  3. 3
      cmd/chat/main.go
  4. 1
      cmd/gateway_ws/config.toml
  5. 26
      cmd/gateway_ws/main.go
  6. 16
      deploy_k8s/docker_compose.yml
  7. 75
      deploy_k8s/gateway_ws.yaml
  8. 6
      go.mod
  9. 13
      go.sum
  10. 34
      internal/chat/logic/chat_server.go
  11. 161
      internal/gateway_ws/gnet_server/codec.go
  12. 23
      internal/gateway_ws/gnet_server/gnet_conn.go
  13. 324
      internal/gateway_ws/gnet_server/websocket.go
  14. 167
      internal/gateway_ws/gws_server/conn_handler.go
  15. 1
      internal/gateway_ws/gws_server/gws_conn.go
  16. 8
      internal/postal/logic/postal_server.go
  17. 1
      pkg/config/config.go
  18. 4
      pkg/config/loader.go
  19. 6
      pkg/config/logger.go
  20. 20
      pkg/grpc/generic/generic_client_factory.go
  21. 5
      pkg/grpc/interceptor/recover_interceptor.go
  22. 8
      pkg/protocol/protocol.go
  23. 3
      pkg/protocol/session/channel.go
  24. 2
      pkg/protocol/session/store.go
  25. 20
      pkg/utils/collect/concurrent_map.go
  26. 62
      pkg/utils/logger/logger.go

8
api/chat.proto

@ -19,6 +19,8 @@ service Chat {
rpc RoomInfo(ReqRoomInfo) returns (ResRoomInfo);
rpc RoomDissolve(ReqRoomDissolve) returns (google.protobuf.Empty);
rpc RoomList(ReqRoomList) returns (ResRoomList);
}
@ -35,7 +37,7 @@ message ReqSend {
}
message ReqRoomSend {
string gid = 1;
string Rid = 1;
string message = 2;
}
@ -83,6 +85,10 @@ message ResRoomInfo {
Room room = 1;
}
message ReqRoomDissolve {
string rid = 1;
}
message ReqRoomList {
}

249
benchmark/main.go

@ -4,6 +4,7 @@ import (
"context"
"encoding/base64"
"encoding/json"
"errors"
"flag"
"fmt"
"github.com/bytedance/sonic"
@ -16,6 +17,7 @@ import (
"math/rand"
"net"
"net/http"
"reflect"
"runtime"
"sonet/api/gen/auth"
"sonet/api/gen/chat"
@ -23,6 +25,7 @@ import (
"sonet/pkg/grpc/discovery"
"sonet/pkg/protocol"
"sonet/pkg/protocol/deliver"
"sonet/pkg/utils/collect"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/security"
"sonet/pkg/utils/shutdown"
@ -34,28 +37,29 @@ import (
)
var (
gatewayHttp = "http://192.168.110.41:7000"
benchmarkMode = "deliver" // deliver/group
mockUsers = 2000
eachUserSend = 100
mockNetUsers []*NetUser
sendCounter int64 = 0
receiverCounter int64 = 0
//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
callbacks map[int32]func(res any)
callbackMutex *sync.Mutex
callbacks *collect.ConcurrentMap[int32, func(res any)]
useMsSum int64
useMsAvg int64
)
func init() {
config.InitLogger()
config.InitLogger(false)
callbacks = make(map[int32]func(res any), 128)
callbackMutex = &sync.Mutex{}
//callbacks = make(map[int32]func(res any), 128)
//callbackMutex = &sync.Mutex{}
callbacks = collect.NewConcurrentMap[int32, func(res any)](16, func(k int32) string {
return strconv.Itoa(int(k))
})
httpClient = &http.Client{Timeout: 10 * time.Second}
var err error
etcdClient, err = clientv3.New(clientv3.Config{
@ -68,10 +72,15 @@ func init() {
}
}
// main -mode=deliver -users=10 -send=10
func main() {
flag.StringVar(&benchmarkMode, "mode", "deliver", "benchmark mode: deliver/group")
flag.IntVar(&mockUsers, "users", 10, "mock users")
runtime.GOMAXPROCS(runtime.NumCPU())
flag.StringVar(&gatewayHttp, "gateway", "http://192.168.110.41:7000", "http gateway address") // 124.222.131.236:30830
flag.StringVar(&benchmarkMode, "mode", "group", "benchmark mode: deliver/group")
flag.IntVar(&mockUsers, "users", 100, "mock users")
flag.IntVar(&eachUserSend, "send", 10, "each user send msg count")
flag.Parse()
if benchmarkMode == "group" {
benchmarkGroup()
@ -109,29 +118,10 @@ func benchmark() {
// go prof.StartPprof(":8888")
ctx, cancel := context.WithCancel(context.Background())
go records(ctx)
mockNetUsers = make([]*NetUser, mockUsers)
// initial uids
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,
}
}
mockUsersConnect(ctx)
for i := 0; i < mockUsers; i++ {
netUser := mockNetUsers[i]
@ -145,9 +135,6 @@ func benchmark() {
}
func benchmarkGroup() {
benchmarkMode = "groupDeliver"
runtime.GOMAXPROCS(runtime.NumCPU())
ctx, cancel := context.WithCancel(context.Background())
shutdown.AddHook(cancel)
@ -163,76 +150,121 @@ func benchmarkGroup() {
}
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)
mockUsersConnect(ctx)
// create room
channel, err := send(mockNetUsers[0].Conn, chat.Chat_ServiceDesc.ServiceName, "RoomCreate", &chat.ReqRoomCreate{Rname: "benchmark"})
if err != nil {
panic(err)
return
}
conn, err := getConn(token)
if err != nil {
// wait result
res := <-channel
if err, failed := res.(error); failed {
panic(err)
} else {
groupId = res.(*chat.ResRoomCreate).Room.Rid
}
err = handleConn(ctx, uid, token, conn)
for _, user := range mockNetUsers {
channel, err := send(user.Conn, chat.Chat_ServiceDesc.ServiceName, "RoomJoin", &chat.ReqRoomJoin{Rid: groupId})
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 {
// wait result
if err, failed := (<-channel).(error); failed {
panic(err)
}
}
shutdown.AddHook(func() {
if err := groupDeliver.GroupDissolve(context.Background(), groupId); err != nil {
logger.Error("dissolve group error: ", err)
// delete room
channel, err := send(mockNetUsers[0].Conn, chat.Chat_ServiceDesc.ServiceName, "RoomDissolve", &chat.ReqRoomDissolve{Rid: groupId})
if err != nil {
panic(err)
return
}
// wait result
if err, failed := (<-channel).(error); failed {
panic(err)
}
logger.Info("test group dissolved")
})
go records(ctx)
// send group message
message := &chat.ChatMessage{Sender: "100001", Content: "hello"}
args := &chat.ReqRoomSend{Rid: groupId, Message: "hi"}
for i := 0; i < eachUserSend; i++ {
err := groupDeliver.DeliverGroup(context.Background(), groupId, message)
channel, err := send(mockNetUsers[0].Conn, chat.Chat_ServiceDesc.ServiceName, "RoomSend", args)
if err != nil {
logger.Error("deliver group error:", err)
logger.Error("room send error: ", err)
continue
}
// wait result
if err, failed := (<-channel).(error); failed {
logger.Error("room send failed: ", err)
}
atomic.AddInt64(&sendCounter, 1)
}
shutdown.Await()
}
// mockUsersConnect 并行快速创建链接
func mockUsersConnect(ctx context.Context) {
// initial mock users, join to postal group
mockNetUsers = make([]*NetUser, mockUsers)
concurrent := 100
wg := &sync.WaitGroup{}
wg.Add(concurrent)
each := (mockUsers / concurrent) + 1
for i := 0; i < concurrent; i++ {
//begin, end := i*each, (i+1)*each
//fmt.Printf("%d: %d~%d \n", i, begin, end)
go func(segment int) {
defer wg.Done()
begin, end := segment*each, (segment+1)*each
for i := begin; i < end && 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}
}
}(i)
}
wg.Wait()
}
func sendBatchChatMessage(ctx context.Context, uid string, conn *websocket.Conn, count int) {
r := rand.New(rand.NewSource(time.Now().UnixMilli()))
for i := 0; i < count; i++ {
receiverUid := mockNetUsers[r.Intn(mockUsers)].Uid
args := &chat.ReqSend{
Receiver: receiverUid,
Content: "hello",
Content: "hi",
}
channel, err := send(conn, chat.Chat_ServiceDesc.ServiceName, "send", args)
channel, err := send(conn, chat.Chat_ServiceDesc.ServiceName, "Send", args)
if err != nil {
fmt.Println("send failed:", err)
return
}
res := <-channel // wait result
err, failed := res.(error)
if failed {
// wait result
if err, failed := (<-channel).(error); failed {
fmt.Println("send failed:", err)
} else {
// fmt.Println("send ok:", res)
}
}
}
@ -297,11 +329,16 @@ func handleConn(ctx context.Context, uid string, token string, conn *websocket.C
if err != nil {
return
}
res := <-channel
// fmt.Printf("%v\n", res)
ch := time.After(time.Second * 5)
select {
case <-ch:
err = errors.New("req verify timeout")
case res := <-channel:
if e, failed := res.(error); failed {
err = e
}
}
return
}
@ -347,21 +384,6 @@ func listen(ctx context.Context, conn *websocket.Conn) {
atomic.AddInt64(&receiverCounter, 1)
}
if header.Type == 4 {
callbackMutex.Lock()
callback, ok := callbacks[header.SeqId]
if ok {
delete(callbacks, header.SeqId)
}
callbackMutex.Unlock()
if !ok {
e = fmt.Errorf("callback %d not found", header.SeqId)
return
}
callback(err)
continue
}
// notice message
if header.Type == protocol.TypeNotice {
switch header.Target {
@ -377,46 +399,55 @@ func listen(ctx context.Context, conn *websocket.Conn) {
continue
}
callback, ok := callbacks.LoadAndDelete(header.SeqId)
if !ok {
logger.Errorf("callback %d not found", header.SeqId)
continue
}
if header.Type == protocol.TypeError {
err = errors.New(string(payload.Body))
callback(err)
continue
}
// rpc response
var msg proto.Message
switch header.Svc {
case auth.Auth_ServiceDesc.ServiceName:
switch header.Target {
case "Subject":
// fmt.Println("res verify...")
msg = &auth.Subject{}
}
case chat.Chat_ServiceDesc.ServiceName:
switch header.Target {
case "ResSend":
// fmt.Println("res send...")
msg = &chat.ResSend{}
if svc, ok := protoStructs[header.Svc]; ok {
if s, ok := svc[header.Target]; ok {
val := reflect.New(reflect.TypeOf(s).Elem())
msg = val.Interface().(proto.Message)
}
}
if msg == nil {
logger.Warning("unknown rpc response: ", len(payload.Body), header)
if msg == nil { // protobuf.Empty
// logger.Warning("unknown rpc response: ", len(payload.Body), header)
callback(nil)
continue
}
err = proto.Unmarshal(payload.Body, msg)
if e != nil {
if err != nil {
logger.Error("unmarshal rcp res body error: ", payload, err)
return
}
callbackMutex.Lock()
callback, ok := callbacks[header.SeqId]
if ok {
delete(callbacks, header.SeqId)
}
callbackMutex.Unlock()
if !ok {
e = fmt.Errorf("callback %d not found", header.SeqId)
callback(fmt.Errorf("unmarshal rcp res body error: %v", err))
return
}
callback(msg)
}
}
var protoStructs = map[string]map[string]proto.Message{
auth.Auth_ServiceDesc.ServiceName: {
"Subject": &auth.Subject{},
},
chat.Chat_ServiceDesc.ServiceName: {
"ResSend": &chat.ResSend{},
"ResRoomCreate": &chat.ResRoomCreate{},
"ResRoomInfo": &chat.ResRoomInfo{},
"ResRoomList": &chat.ResRoomList{},
},
}
func send(conn *websocket.Conn, svc, method string, reqArgs proto.Message) (channel chan any, err error) {
seq := atomic.AddInt32(&seqId, 1)
header := &protocol.Header{
@ -443,16 +474,14 @@ func send(conn *websocket.Conn, svc, method string, reqArgs proto.Message) (chan
channel = make(chan any)
// ready callback
callbackMutex.Lock()
callbacks[seq] = func(res any) {
callbacks.Store(seq, func(res any) {
ms := time.Now().UnixMilli() - begin
msSum := atomic.AddInt64(&useMsSum, ms)
counter := atomic.AddInt64(&receiverCounter, 1)
atomic.StoreInt64(&useMsAvg, msSum/counter)
channel <- res
}
callbackMutex.Unlock()
})
// send message
err = conn.WriteMessage(websocket.BinaryMessage, bytes)

3
cmd/chat/main.go

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

1
cmd/gateway_ws/config.toml

@ -6,6 +6,7 @@ subjectLrcExpiration = "10m"
subjectLrcCleanupInterval = "5m"
[grpc]
log = false
address = ":7011" # postal cluster offset port +1000=8011
maxSendMsgSize = "8Mi"
maxRecvMsgSize = "8Mi"

26
cmd/gateway_ws/main.go

@ -12,7 +12,6 @@ import (
"sonet/pkg/config"
"sonet/pkg/grpc/client"
"sonet/pkg/grpc/discovery"
"sonet/pkg/grpc/discovery/etcd"
"sonet/pkg/grpc/generic"
"sonet/pkg/plugins/mq"
"sonet/pkg/protocol/deliver"
@ -58,14 +57,17 @@ func main() {
sessionStore := session.NewMapStore(128)
groupStore := collect.NewConcurrentMap[string, *group.Group](128, func(k string) string { return k })
ctx, cancel := context.WithCancel(context.Background())
shutdown.AddHook(cancel)
// run websocket server
postalAddr := etcd.MustRegisterAddress(conf.Grpc.Address)
//postalAddr := etcd.MustRegisterAddress(conf.Grpc.Address)
//connHandler := server.NewConnHandler(postalAddr, grpcFactory, sessionStore, subjectStore)
//httpServer := server.NewHttpServer(connHandler)
gwsHandler := gws_server.NewGwsHandler(postalAddr, grpcFactory, sessionStore)
gwsHandler := gws_server.NewGwsHandler(grpcFactory, sessionStore)
gwsHandler.Init(ctx)
httpServer := gws_server.NewGwsServer(gwsHandler)
httpServer.Init()
go func() {
err = httpServer.Run(appConf.HttpPort)
if err != nil {
@ -73,9 +75,21 @@ func main() {
}
}()
//go func() {
// gnetWsServer := gnet_server.NewGnetWsServer(grpcFactory, sessionStore)
// gnetWsServer.Init(ctx)
// err := gnet.Run(gnetWsServer,
// fmt.Sprintf("tcp://0.0.0.0:%d", appConf.HttpPort),
// gnet.WithMulticore(true),
// gnet.WithReusePort(true),
// gnet.WithTicker(true),
// )
// if err != nil {
// panic(err)
// }
//}()
// run postal cluster server
ctx, cancel := context.WithCancel(context.Background())
shutdown.AddHook(cancel)
postalPicker := deliver.NewPostalPicker(dis, grpc.WithTransportCredentials(insecure.NewCredentials()))
if err = postalPicker.Init(ctx); err != nil {
panic(err)

16
deploy_k8s/docker_compose.yml

@ -8,12 +8,26 @@ services:
- SO_APP.ENDPOINTADDRESS=192.168.110.36:7001
- SO_GRPC.ADDRESS=:7011
svr-gateway-ws-2:
image: so_gateway_ws:1.0.0
network_mode: host
environment:
- SO_APP.HTTPPORT=7002
- SO_APP.ENDPOINTADDRESS=192.168.110.36:7002
- SO_GRPC.ADDRESS=:7012
svr-gateway-ws-3:
image: so_gateway_ws:1.0.0
network_mode: host
environment:
- SO_APP.HTTPPORT=7003
- SO_APP.ENDPOINTADDRESS=192.168.110.36:7003
- SO_GRPC.ADDRESS=:7012
- SO_GRPC.ADDRESS=:7013
svr-gateway-ws-4:
image: so_gateway_ws:1.0.0
network_mode: host
environment:
- SO_APP.HTTPPORT=7004
- SO_APP.ENDPOINTADDRESS=192.168.110.36:7004
- SO_GRPC.ADDRESS=:7014
svr-chat:
image: so_chat:1.0.0

75
deploy_k8s/gateway_ws.yaml

@ -144,3 +144,78 @@ spec:
targetPort: 9100
selector:
app: gateway-ws2
---
apiVersion: v1
kind: Pod
metadata:
name: gateway-ws3
namespace: sopod
labels:
app: gateway-ws3
spec:
containers:
- name: gateway-ws3
image: 10.0.16.5:3050/so_gateway_ws:1.0.0
env:
- name: SO_ETCD.ENDPOINTS
value: etcd-svc:2379
- name: SO_ETCD.USERNAME
value: ""
- name: SO_ETCD.PASSWORD
value: ""
- name: SO_PROMETHEUS.ENABLE
value: "false"
- name: SO_PROMETHEUS.PORT
value: "9100"
- name: SO_APP.HTTPPORT
value: "7003"
- name: SO_APP.ENDPOINTADDRESS
value: 124.222.131.236:30803
- name: SO_REDIS.ADDR
value: "redis-svc:6379"
- name: SO_REDIS.PASSWORD
value: ""
- name: SO_NATS.URL
value: nats://nats-svc:4222
imagePullPolicy: Always
ports:
- containerPort: 7003
- containerPort: 30803
- containerPort: 9100
---
apiVersion: v1
kind: Service
metadata:
name: gateway-ws3-svc
namespace: sopod
labels:
app: gateway-ws3-svc
spec:
type: NodePort
ports:
- name: gateway-ws3
port: 7003
targetPort: 7003
nodePort: 30803
selector:
app: gateway-ws3
---
apiVersion: v1
kind: Service
metadata:
name: gateway-ws3-metrics-svc
namespace: sopod
labels:
app: gateway-ws3-metrics-svc
spec:
type: ClusterIP
ports:
- name: metrics
port: 9100
targetPort: 9100
selector:
app: gateway-ws3

6
go.mod

@ -6,12 +6,14 @@ require (
github.com/bytedance/sonic v1.10.2
github.com/dsnet/golib/unitconv v1.0.2
github.com/gin-gonic/gin v1.9.1
github.com/gobwas/ws v1.3.2
github.com/golang/protobuf v1.5.3
github.com/google/uuid v1.4.0
github.com/gorilla/websocket v1.5.1
github.com/jhump/protoreflect v1.15.4
github.com/lxzan/gws v1.7.0
github.com/nats-io/nats.go v1.31.0
github.com/panjf2000/gnet/v2 v2.3.4
github.com/patrickmn/go-cache v2.1.0+incompatible
github.com/redis/go-redis/v9 v9.4.0
github.com/sirupsen/logrus v1.9.3
@ -39,6 +41,8 @@ require (
github.com/go-playground/universal-translator v0.18.1 // indirect
github.com/go-playground/validator/v10 v10.14.0 // indirect
github.com/go-sql-driver/mysql v1.7.0 // indirect
github.com/gobwas/httphead v0.1.0 // indirect
github.com/gobwas/pool v0.2.1 // indirect
github.com/goccy/go-json v0.10.2 // indirect
github.com/gogo/protobuf v1.3.2 // indirect
github.com/hashicorp/hcl v1.0.0 // indirect
@ -65,6 +69,7 @@ require (
github.com/subosito/gotenv v1.6.0 // indirect
github.com/twitchyliquid64/golang-asm v0.15.1 // indirect
github.com/ugorji/go/codec v1.2.11 // indirect
github.com/valyala/bytebufferpool v1.0.0 // indirect
go.etcd.io/etcd/client/pkg/v3 v3.5.11 // indirect
go.uber.org/atomic v1.9.0 // indirect
go.uber.org/multierr v1.9.0 // indirect
@ -80,5 +85,6 @@ require (
google.golang.org/genproto/googleapis/api v0.0.0-20231106174013-bbf56f31fb17 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20231120223509-83a465c0220f // indirect
gopkg.in/ini.v1 v1.67.0 // indirect
gopkg.in/natefinch/lumberjack.v2 v2.2.1 // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect
)

13
go.sum

@ -46,6 +46,12 @@ github.com/go-playground/validator/v10 v10.14.0 h1:vgvQWe3XCz3gIeFDm/HnTIbj6UGmg
github.com/go-playground/validator/v10 v10.14.0/go.mod h1:9iXMNT7sEkjXb0I+enO7QXmzG6QCsPWY4zveKFVRSyU=
github.com/go-sql-driver/mysql v1.7.0 h1:ueSltNNllEqE3qcWBTD0iQd3IpL/6U+mJxLkazJ7YPc=
github.com/go-sql-driver/mysql v1.7.0/go.mod h1:OXbVy3sEdcQ2Doequ6Z5BW6fXNQTmx+9S1MCJN5yJMI=
github.com/gobwas/httphead v0.1.0 h1:exrUm0f4YX0L7EBwZHuCF4GDp8aJfVeBrlLQrs6NqWU=
github.com/gobwas/httphead v0.1.0/go.mod h1:O/RXo79gxV8G+RqlR/otEwx4Q36zl9rqC5u12GKvMCM=
github.com/gobwas/pool v0.2.1 h1:xfeeEhW7pwmX8nuLVlqbzVc7udMDrwetjEv+TZIz1og=
github.com/gobwas/pool v0.2.1/go.mod h1:q8bcK0KcYlCgd9e7WYLm9LpyS+YeLd8JVDW6WezmKEw=
github.com/gobwas/ws v1.3.2 h1:zlnbNHxumkRvfPWgfXu8RBwyNR1x8wh9cf5PTOCqs9Q=
github.com/gobwas/ws v1.3.2/go.mod h1:hRKAFb8wOxFROYNsT1bqfWnhX+b5MFeJM9r2ZSwg/KY=
github.com/goccy/go-json v0.10.2 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU=
github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MGFi0w8I=
github.com/godbus/dbus/v5 v5.0.4/go.mod h1:xhWf0FNVPg57R7Z0UbKHbJfkEywrmjJnf7w5xrFpKfA=
@ -106,6 +112,9 @@ github.com/nats-io/nkeys v0.4.7 h1:RwNJbbIdYCoClSDNY7QVKZlyb/wfT6ugvFCiKy6vDvI=
github.com/nats-io/nkeys v0.4.7/go.mod h1:kqXRgRDPlGy7nGaEDMuYzmiJCIAAWDK0IMBtDmGD0nc=
github.com/nats-io/nuid v1.0.1 h1:5iA8DT8V7q8WK2EScv2padNa/rTESc1KdnPw4TC2paw=
github.com/nats-io/nuid v1.0.1/go.mod h1:19wcPz3Ph3q0Jbyiqsd0kePYG7A95tJPxeL+1OSON2c=
github.com/panjf2000/ants/v2 v2.8.2 h1:D1wfANttg8uXhC9149gRt1PDQ+dLVFjNXkCEycMcvQQ=
github.com/panjf2000/gnet/v2 v2.3.4 h1:+ASHt+Wxr0KIzlk5FsLBbegCc4US7iVCdZ1QbUyw17g=
github.com/panjf2000/gnet/v2 v2.3.4/go.mod h1:0mTLWq4zMEXyQ35BY094dNWYnXfIdDg0mOlmZJflaXE=
github.com/patrickmn/go-cache v2.1.0+incompatible h1:HRMgzkcYKYpi3C8ajMPV8OFXaaRUnok+kx1WdO15EQc=
github.com/patrickmn/go-cache v2.1.0+incompatible/go.mod h1:3Qf8kWWT7OJRJbdiICTKqZju1ZixQ/KpMGzzAfe6+WQ=
github.com/pelletier/go-toml/v2 v2.1.0 h1:FnwAJ4oYMvbT/34k9zzHuZNrhlz48GB3/s6at6/MHO4=
@ -150,6 +159,8 @@ github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS
github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08=
github.com/ugorji/go/codec v1.2.11 h1:BMaWp1Bb6fHwEtbplGBGJ498wD+LKlNSl25MjdZY4dU=
github.com/ugorji/go/codec v1.2.11/go.mod h1:UNopzCgEMSXjBc6AOMqYvWC1ktqTAfzJZUZgYf6w6lg=
github.com/valyala/bytebufferpool v1.0.0 h1:GqA5TC/0021Y/b9FG4Oi9Mr3q7XYx6KllzawFIhcdPw=
github.com/valyala/bytebufferpool v1.0.0/go.mod h1:6bBcMArwyJ5K/AmCkWv1jt77kVWyCJ6HpOuEn7z0Csc=
github.com/yuin/goldmark v1.1.27/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74=
github.com/yuin/goldmark v1.2.1/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74=
github.com/yuin/goldmark v1.3.5/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
@ -240,6 +251,8 @@ gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8
gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15 h1:YR8cESwS4TdDjEe65xsg0ogRM/Nc3DYOhEAlW+xobZo=
gopkg.in/ini.v1 v1.67.0 h1:Dgnx+6+nfE+IfzjUEISNeydPJh9AXNNsWbGP9KzCsOA=
gopkg.in/ini.v1 v1.67.0/go.mod h1:pNLf8WUiyNEtQjuu5G5vTm06TEv9tsIgeAvK8hOrP4k=
gopkg.in/natefinch/lumberjack.v2 v2.2.1 h1:bBRl1b0OH9s/DuPhuXpNl+VtCaJXFZ5/uEFST95x9zc=
gopkg.in/natefinch/lumberjack.v2 v2.2.1/go.mod h1:YD8tP3GAjkrDg1eZH7EGmyESg/lsYskCTPBJVb9jqSc=
gopkg.in/yaml.v2 v2.2.8/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
gopkg.in/yaml.v2 v2.4.0 h1:D8xgwECY7CYvx+Y2n4sBz93Jn9JRvxdiyyo8CTfuKaY=
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=

34
internal/chat/logic/chat_server.go

@ -98,22 +98,23 @@ func (s *ChatServer) Send(ctx context.Context, req *chat.ReqSend) (*chat.ResSend
func (s *ChatServer) RoomSend(ctx context.Context, req *chat.ReqRoomSend) (res *chat.ResSend, err error) {
subject, err := session.GetSubject(ctx)
if err != nil {
return nil, err
return
}
gMessage := &chat.ChatMessage{
Type: 1,
Sender: subject.Uid,
Gid: req.Gid,
Gid: req.Rid,
Content: req.Message,
}
err = s.groupDeliver.DeliverGroup(ctx, req.Gid, gMessage)
err = s.groupDeliver.DeliverGroup(ctx, req.Rid, gMessage)
res = &chat.ResSend{}
return
}
func (s *ChatServer) RoomCreate(ctx context.Context, req *chat.ReqRoomCreate) (res *chat.ResRoomCreate, err error) {
subject, err := session.GetSubject(ctx)
if err != nil {
return nil, err
return
}
logger.Infof("create room: %s", req.Rname)
@ -138,10 +139,11 @@ func (s *ChatServer) RoomCreate(ctx context.Context, req *chat.ReqRoomCreate) (r
return
}
func (s *ChatServer) RoomJoin(ctx context.Context, req *chat.ReqRoomJoin) (_ *emptypb.Empty, err error) {
func (s *ChatServer) RoomJoin(ctx context.Context, req *chat.ReqRoomJoin) (emp *emptypb.Empty, err error) {
emp = &emptypb.Empty{}
subject, err := session.GetSubject(ctx)
if err != nil {
return nil, err
return
}
room, ok := s.rooms.Load(req.Rid)
@ -155,10 +157,11 @@ func (s *ChatServer) RoomJoin(ctx context.Context, req *chat.ReqRoomJoin) (_ *em
return
}
func (s *ChatServer) RoomLeave(ctx context.Context, req *chat.ReqRoomLeave) (_ *emptypb.Empty, err error) {
func (s *ChatServer) RoomLeave(ctx context.Context, req *chat.ReqRoomLeave) (emp *emptypb.Empty, err error) {
emp = &emptypb.Empty{}
subject, err := session.GetSubject(ctx)
if err != nil {
return nil, err
return
}
room, ok := s.rooms.Load(req.Rid)
@ -172,7 +175,8 @@ func (s *ChatServer) RoomLeave(ctx context.Context, req *chat.ReqRoomLeave) (_ *
return
}
func (s *ChatServer) RoomKickOut(ctx context.Context, req *chat.ReqRoomKickOut) (_ *emptypb.Empty, err error) {
func (s *ChatServer) RoomKickOut(ctx context.Context, req *chat.ReqRoomKickOut) (emp *emptypb.Empty, err error) {
emp = &emptypb.Empty{}
return
}
@ -187,6 +191,18 @@ func (s *ChatServer) RoomInfo(ctx context.Context, req *chat.ReqRoomInfo) (res *
return
}
// RoomDissolve 解散群组
func (s *ChatServer) RoomDissolve(ctx context.Context, req *chat.ReqRoomDissolve) (emp *emptypb.Empty, err error) {
emp = &emptypb.Empty{}
err = s.groupDeliver.GroupDissolve(ctx, req.Rid)
if err != nil {
return
}
s.rooms.Delete(req.Rid)
return
}
func (s *ChatServer) RoomList(ctx context.Context, req *chat.ReqRoomList) (res *chat.ResRoomList, err error) {
rooms := make([]*chat.Room, 16)
s.rooms.Range(func(_ string, room *chat.Room) bool {

161
internal/gateway_ws/gnet_server/codec.go

@ -0,0 +1,161 @@
package gnet_server
import (
"bytes"
"fmt"
"github.com/gobwas/ws"
"github.com/gobwas/ws/wsutil"
"github.com/panjf2000/gnet/v2"
"github.com/panjf2000/gnet/v2/pkg/logging"
"io"
)
type wsCodec struct {
upgraded bool // 链接是否升级
buf bytes.Buffer // 从实际socket中读取到的数据缓存
wsMsgBuf wsMessageBuf // ws 消息缓存
}
type wsMessageBuf struct {
firstHeader *ws.Header
curHeader *ws.Header
cachedBuf bytes.Buffer
}
type readWrite struct {
io.Reader
io.Writer
}
func (w *wsCodec) upgrade(c gnet.Conn) (ok bool, action gnet.Action) {
if w.upgraded {
ok = true
return
}
buf := &w.buf
tmpReader := bytes.NewReader(buf.Bytes())
oldLen := tmpReader.Len()
// logging.Infof("do Upgrade")
hs, err := ws.Upgrade(readWrite{tmpReader, c})
skipN := oldLen - tmpReader.Len()
if err != nil {
if err == io.EOF || err == io.ErrUnexpectedEOF { //数据不完整
return
}
buf.Next(skipN)
logging.Infof("conn[%v] [err=%v]", c.RemoteAddr().String(), err.Error())
action = gnet.Close
return
}
buf.Next(skipN)
logging.Infof("conn[%v] upgrade websocket protocol! Handshake: %v", c.RemoteAddr().String(), hs)
if err != nil {
logging.Infof("conn[%v] [err=%v]", c.RemoteAddr().String(), err.Error())
action = gnet.Close
return
}
ok = true
w.upgraded = true
return
}
func (w *wsCodec) readBufferBytes(c gnet.Conn) gnet.Action {
size := c.InboundBuffered()
buf := make([]byte, size, size)
read, err := c.Read(buf)
if err != nil {
logging.Infof("read err! %w", err)
return gnet.Close
}
if read < size {
logging.Infof("read bytes len err! size: %d read: %d", size, read)
return gnet.Close
}
w.buf.Write(buf)
return gnet.None
}
func (w *wsCodec) Decode(c gnet.Conn) (outs []wsutil.Message, err error) {
// fmt.Println("do Decode")
messages, err := w.readWsMessages()
if err != nil {
logging.Infof("Error reading message! %v", err)
return nil, err
}
if messages == nil || len(messages) <= 0 { //没有读到完整数据 不处理
return
}
for _, message := range messages {
if message.OpCode.IsControl() {
err = wsutil.HandleClientControlMessage(c, message)
if err != nil {
return
}
continue
}
if message.OpCode == ws.OpText || message.OpCode == ws.OpBinary {
outs = append(outs, message)
}
}
return
}
func (w *wsCodec) readWsMessages() (messages []wsutil.Message, err error) {
msgBuf := &w.wsMsgBuf
in := &w.buf
for {
if msgBuf.curHeader == nil {
if in.Len() < ws.MinHeaderSize { //头长度至少是2
return
}
var head ws.Header
if in.Len() >= ws.MaxHeaderSize {
head, err = ws.ReadHeader(in)
if err != nil {
return messages, err
}
} else { //有可能不完整,构建新的 reader 读取 head 读取成功才实际对 in 进行读操作
tmpReader := bytes.NewReader(in.Bytes())
oldLen := tmpReader.Len()
head, err = ws.ReadHeader(tmpReader)
skipN := oldLen - tmpReader.Len()
if err != nil {
if err == io.EOF || err == io.ErrUnexpectedEOF { //数据不完整
return messages, nil
}
in.Next(skipN)
return nil, err
}
in.Next(skipN)
}
msgBuf.curHeader = &head
err = ws.WriteHeader(&msgBuf.cachedBuf, head)
if err != nil {
return nil, err
}
}
dataLen := (int)(msgBuf.curHeader.Length)
if dataLen > 0 {
if in.Len() >= dataLen {
_, err = io.CopyN(&msgBuf.cachedBuf, in, int64(dataLen))
if err != nil {
return
}
} else { //数据不完整
fmt.Println(in.Len(), dataLen)
logging.Infof("incomplete data")
return
}
}
if msgBuf.curHeader.Fin { //当前 header 已经是一个完整消息
messages, err = wsutil.ReadClientMessage(&msgBuf.cachedBuf, messages)
if err != nil {
return nil, err
}
msgBuf.cachedBuf.Reset()
} else {
logging.Infof("The data is split into multiple frames")
}
msgBuf.curHeader = nil
}
}

23
internal/gateway_ws/gnet_server/gnet_conn.go

@ -0,0 +1,23 @@
package gnet_server
import (
"github.com/gobwas/ws"
"github.com/gobwas/ws/wsutil"
"github.com/panjf2000/gnet/v2"
)
type gnetConn struct {
conn gnet.Conn
uid string
wsCodec *wsCodec
}
func newGnetConn(conn gnet.Conn) *gnetConn {
return &gnetConn{
conn: conn,
}
}
func (c *gnetConn) Write(msg []byte) error {
return wsutil.WriteServerMessage(c.conn, ws.OpBinary, msg)
}

324
internal/gateway_ws/gnet_server/websocket.go

@ -0,0 +1,324 @@
package gnet_server
import (
"context"
"fmt"
"github.com/gobwas/ws"
"google.golang.org/protobuf/proto"
"regexp"
"sonet/api/gen/auth"
"sonet/api/gen/postal"
"sonet/pkg/grpc/generic"
"sonet/pkg/protocol"
"sonet/pkg/protocol/session"
"sonet/pkg/utils/logger"
"sync/atomic"
"time"
"github.com/panjf2000/gnet/v2"
"github.com/panjf2000/gnet/v2/pkg/logging"
)
type gnetRequest struct {
conn *gnetConn
payload *protocol.Payload
}
type GnetWsServer struct {
gnet.BuiltinEventEngine
addr string
multicore bool
eng gnet.Engine
connected int64
grpcFactory *generic.GrpcGenericClientFactory
sessionStore session.Store // <string, *session.NetSubject> 当前连接用户,内存缓存
requests chan *gnetRequest
}
func NewGnetWsServer(
grpcFactory *generic.GrpcGenericClientFactory,
sessionStore session.Store) *GnetWsServer {
return &GnetWsServer{
grpcFactory: grpcFactory,
sessionStore: sessionStore,
}
}
func (wss *GnetWsServer) Init(ctx context.Context) {
wss.requests = make(chan *gnetRequest, 1024*16)
// todo 针对不同服务的协程池
for i := 0; i < 32; i++ {
go wss.dispatch(ctx)
}
}
// dispatch 多个 goroutine 进行请求 dispatch
func (wss *GnetWsServer) dispatch(ctx context.Context) {
for {
select {
case <-ctx.Done():
return
case request := <-wss.requests:
func() {
conn, payload := request.conn, request.payload
header := payload.Header
service, method := header.Svc, header.Target
// grpc generic call
ctx := context.Background()
ctx2, cancel2 := context.WithTimeout(ctx, time.Second*3)
defer cancel2()
grpcClient, err := wss.grpcFactory.GetClient(ctx2, service)
if err != nil {
logger.Error("get grpc generic client error: ", err)
writeError(conn, header.SeqId, fmt.Sprintf("service %s not avaliable: %s", header.Svc, err.Error()))
return
}
// put grpc request session
if conn.uid != "" {
ctx = session.PutSubject(ctx, session.NewRpcSubject(conn.uid))
}
ctx3, cancel3 := context.WithTimeout(ctx, time.Second*10)
defer cancel3()
resp, err := grpcClient.InvokeUnary(ctx3, header.Target, payload.Body)
if err != nil {
writeError(conn, header.SeqId, unwrapRpcError(err.Error()))
logger.Error("grpc generic call error: ", err)
return
}
// write response
if resp == nil { // proto.Empty
header.Target = ""
} else {
payload.Body, err = resp.Marshal()
if err != nil {
logger.Error("generic call response marshal error: ", err)
writeError(conn, header.SeqId, "server error")
return
}
header.Target = resp.XXX_MessageName()
}
// auth verify success
if service == "Auth" && method == "Verify" {
var channel *session.Channel
conn.uid, channel, err = wss.extractAuthVerify(conn, payload.Body)
if err != nil {
logger.Error("extractAuthVerify error: ", err)
writeError(conn, header.SeqId, "server error")
return
}
// channel
go wss.processWriting(channel)
}
header.Type = protocol.TypeResponse
resMessage, err := protocol.EncodeSo(payload)
if err != nil {
logger.Error("grpc generic call error: ", err)
return
}
err = conn.Write(resMessage)
if err != nil {
logger.Error("gws write response error: ", err)
return
}
}()
}
}
}
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 *GnetWsServer) processWriting(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 (wss *GnetWsServer) OnBoot(eng gnet.Engine) gnet.Action {
wss.eng = eng
logging.Infof("echo server with multi-core=%t is listening on %s", wss.multicore, wss.addr)
return gnet.None
}
func (wss *GnetWsServer) OnOpen(c gnet.Conn) ([]byte, gnet.Action) {
// wsCodec, uid, gnetConn
conn := newGnetConn(c)
conn.wsCodec = new(wsCodec)
c.SetContext(conn)
// c.SetContext(new(wsCodec))
atomic.AddInt64(&wss.connected, 1)
return nil, gnet.None
}
func (wss *GnetWsServer) OnClose(c gnet.Conn, err error) (action gnet.Action) {
if err != nil {
logging.Warnf("error occurred on connection=%s, %v\n", c.RemoteAddr().String(), err)
}
atomic.AddInt64(&wss.connected, -1)
logging.Infof("conn[%v] disconnected", c.RemoteAddr().String())
return gnet.None
}
func (wss *GnetWsServer) OnTraffic(c gnet.Conn) (action gnet.Action) {
conn := c.Context().(*gnetConn)
if conn.wsCodec.readBufferBytes(c) == gnet.Close {
return gnet.Close
}
ok, action := conn.wsCodec.upgrade(c)
if !ok {
return
}
if conn.wsCodec.buf.Len() <= 0 {
return gnet.None
}
messages, err := conn.wsCodec.Decode(c)
if err != nil {
return gnet.Close
}
if messages == nil {
return
}
authorized := conn.uid != ""
for _, message := range messages {
if message.OpCode != ws.OpBinary {
logger.Info("receive message type: ", message.OpCode)
continue
}
// decode payload
payload, err := protocol.DecodeSo(message.Payload)
if err != nil {
logger.Errorf("decode message error: len=%d", len(message.Payload), err)
return gnet.Close
}
header := payload.Header
service, method := header.Svc, header.Target
// authorization
if !authorized {
if !(service == "Auth" && method == "Verify") {
writeError(conn, header.SeqId, session.UnauthorizedRequestError.Error())
return gnet.Close
}
}
// payload -> processing channel
request := &gnetRequest{conn: conn, payload: payload}
select {
case wss.requests <- request:
default:
writeError(conn, header.SeqId, "server busy")
}
//msgLen := len(message.Payload)
//if msgLen > 128 {
// logging.Infof("conn[%v] receive [op=%v] [msg=%v..., len=%d]", c.RemoteAddr().String(), message.OpCode, string(message.Payload[:128]), len(message.Payload))
//} else {
// logging.Infof("conn[%v] receive [op=%v] [msg=%v, len=%d]", c.RemoteAddr().String(), message.OpCode, string(message.Payload), len(message.Payload))
//}
//// This is the echo server
//err = wsutil.WriteServerMessage(c, message.OpCode, message.Payload)
//if err != nil {
// logging.Infof("conn[%v] [err=%v]", c.RemoteAddr().String(), err.Error())
// return gnet.Close
//}
}
return gnet.None
}
func (wss *GnetWsServer) OnTick() (delay time.Duration, action gnet.Action) {
logging.Infof("[connected-count=%v]", atomic.LoadInt64(&wss.connected))
return 10 * time.Second, gnet.None
}
func (c *GnetWsServer) extractAuthVerify(conn *gnetConn, resp []byte) (sessionUid string, channel *session.Channel, err error) {
authSubject := &auth.Subject{}
err = proto.Unmarshal(resp, authSubject)
if err != nil {
logger.Error("grpc generic call error: ", err)
return
}
// 设置连接身份信息
sessionUid = authSubject.Uid
// 关闭旧的链接
oldChannel, ok := c.sessionStore.Load(sessionUid)
if ok {
if err = oldChannel.Conn.(*gnetConn).conn.Close(); err != nil {
logger.Error("close gnet conn error: ", err)
}
logger.Infof("close subject old conn: %s\n", sessionUid)
}
// 存储session到内存
channel = session.NewChannel(sessionUid, conn)
c.sessionStore.Store(sessionUid, channel)
logger.Infof("subject online: %s\n", sessionUid)
return
}
func writeError(conn session.NetConn, seqId int32, errMsg string) {
header := &protocol.Header{}
payload := &protocol.Payload{Header: header}
header.Magic = protocol.Magic
header.Type = protocol.TypeError
header.Status = 50
header.SeqId = seqId
payload.Body = []byte(errMsg)
message, err := protocol.EncodeSo(payload)
if err != nil {
logger.Error("NetClient writeError encode error: ", err)
return
}
err = conn.Write(message)
if err != nil {
logger.Error("net conn write error msg fail: ", err)
return
}
return
}
var reg = regexp.MustCompile("^rpc error:.*?desc ?= ?(.*)$")
// unwrapRpcError 提取 grpc err: fmt.Sprintf("rpc error: code = %s desc = %s", s.Code(), s.Message())
func unwrapRpcError(err string) string {
finds := reg.FindStringSubmatch(err)
if len(finds) > 1 {
return finds[1]
}
return err
}

167
internal/gateway_ws/gws_server/conn_handler.go

@ -18,89 +18,50 @@ import (
const (
PingInterval = 5 * time.Second
PingWait = 10 * time.Second
SessionUidKey = "uid"
SessionGwsConnKey = "gws"
)
type soRequest struct {
conn *gwsConn
payload *protocol.Payload
}
type GwsHandler struct {
postalServerAddress string
grpcFactory *generic.GrpcGenericClientFactory
sessionStore session.Store // <string, *session.NetSubject> 当前连接用户,内存缓存
requests chan *soRequest
}
func NewGwsHandler(postalServerAddress string,
func NewGwsHandler(
grpcFactory *generic.GrpcGenericClientFactory,
sessionStore session.Store,
) *GwsHandler {
return &GwsHandler{
grpcFactory: grpcFactory,
postalServerAddress: postalServerAddress,
sessionStore: sessionStore,
}
}
func (c *GwsHandler) OnOpen(socket *gws.Conn) {
_ = socket.SetDeadline(time.Time{}) // time.Now().Add(PingInterval + PingWait)
socket.Session().Store(SessionGwsConnKey, newGwsConn(socket))
}
func (c *GwsHandler) Init(ctx context.Context) {
c.requests = make(chan *soRequest, 1024*16)
func (c *GwsHandler) OnClose(socket *gws.Conn, err error) {
if err != nil {
logger.Error("gws conn close error: ", err)
}
val, ok := socket.Session().Load(SessionUidKey)
if !ok {
return
// todo 针对不同服务的协程池
for i := 0; i < 32; i++ {
go c.dispatchRequest(ctx)
}
uid := val.(string)
c.sessionStore.Delete(uid)
logger.Info("subject offline: ", uid)
}
func (c *GwsHandler) OnPing(socket *gws.Conn, payload []byte) {
_ = socket.SetDeadline(time.Now().Add(PingInterval + PingWait))
_ = socket.WritePong(nil)
}
func (c *GwsHandler) OnPong(socket *gws.Conn, payload []byte) {}
func (c *GwsHandler) OnMessage(socket *gws.Conn, message *gws.Message) {
defer func() {
err := message.Close()
if err != nil {
logger.Error("gws message close error: ", err)
}
}()
if message.Opcode != gws.OpcodeBinary {
return
}
val, ok := socket.Session().Load(SessionGwsConnKey)
if !ok {
logger.Error("gws socket session error: conn not exists")
return
}
conn := val.(*gwsConn)
val, authorized := socket.Session().Load(SessionUidKey)
var sessionUid string
if authorized {
sessionUid = val.(string)
}
messageBytes := message.Bytes()
payload, err := protocol.DecodeSo(messageBytes)
if err != nil {
logger.Errorf("decode message error: len=%d", len(messageBytes), err)
// dispatchRequest 多个 goroutine 进行请求 dispatch
func (c *GwsHandler) dispatchRequest(ctx context.Context) {
for {
select {
case <-ctx.Done():
return
}
case request := <-c.requests:
func() {
conn, payload := request.conn, request.payload
header := payload.Header
service, method := header.Svc, header.Target
if !authorized && !(service == "Auth" && method == "Verify") {
writeError(conn, header.SeqId, session.UnauthorizedRequestError.Error())
return
}
// grpc generic call
ctx := context.Background()
@ -112,11 +73,11 @@ func (c *GwsHandler) OnMessage(socket *gws.Conn, message *gws.Message) {
writeError(conn, header.SeqId, fmt.Sprintf("service %s not avaliable: %s", header.Svc, err.Error()))
return
}
// put session
if authorized {
ctx = session.PutSubject(ctx, session.NewRpcSubject(sessionUid))
}
// put grpc request session
if conn.uid != "" {
ctx = session.PutSubject(ctx, session.NewRpcSubject(conn.uid))
}
ctx3, cancel3 := context.WithTimeout(ctx, time.Second*10)
defer cancel3()
resp, err := grpcClient.InvokeUnary(ctx3, header.Target, payload.Body)
@ -142,14 +103,13 @@ func (c *GwsHandler) OnMessage(socket *gws.Conn, message *gws.Message) {
// auth verify success
if service == "Auth" && method == "Verify" {
var channel *session.Channel
sessionUid, channel, err = c.extractAuthVerify(conn, payload.Body)
conn.uid, channel, err = c.extractAuthVerify(conn, payload.Body)
if err != nil {
logger.Error("extractAuthVerify error: ", err)
writeError(conn, header.SeqId, "server error")
return
}
socket.Session().Store(SessionUidKey, sessionUid)
// channel
// channel todo dispatch on conn open
go c.dispatch(channel)
}
@ -164,6 +124,79 @@ func (c *GwsHandler) OnMessage(socket *gws.Conn, message *gws.Message) {
logger.Error("gws write response error: ", err)
return
}
}()
}
}
}
func (c *GwsHandler) OnOpen(socket *gws.Conn) {
_ = socket.SetDeadline(time.Time{}) // time.Now().Add(PingInterval + PingWait)
socket.Session().Store(SessionGwsConnKey, newGwsConn(socket))
}
func (c *GwsHandler) OnClose(socket *gws.Conn, err error) {
if err != nil {
logger.Error("gws conn close error: ", err)
}
val, ok := socket.Session().Load(SessionGwsConnKey)
if !ok {
return
}
conn := val.(*gwsConn)
if conn.uid != "" {
c.sessionStore.Delete(conn.uid)
//logger.Info("subject offline: ", conn.uid)
}
}
func (c *GwsHandler) OnPing(socket *gws.Conn, payload []byte) {
_ = socket.SetDeadline(time.Now().Add(PingInterval + PingWait))
_ = socket.WritePong(nil)
}
func (c *GwsHandler) OnPong(socket *gws.Conn, payload []byte) {}
func (c *GwsHandler) OnMessage(socket *gws.Conn, message *gws.Message) {
defer func() {
err := message.Close()
if err != nil {
logger.Error("gws message close error: ", err)
}
}()
if message.Opcode != gws.OpcodeBinary {
logger.Info("receive websocket message type: ", message.Opcode)
return
}
val, ok := socket.Session().Load(SessionGwsConnKey)
if !ok {
logger.Error("gws socket session error: conn not exists")
return
}
conn := val.(*gwsConn)
authorized := conn.uid != ""
messageBytes := message.Bytes()
payload, err := protocol.DecodeSo(messageBytes)
if err != nil {
logger.Errorf("decode message error: len=%d", len(messageBytes), err)
return
}
header := payload.Header
service, method := header.Svc, header.Target
if !authorized && !(service == "Auth" && method == "Verify") {
writeError(conn, header.SeqId, session.UnauthorizedRequestError.Error())
return
}
// payload -> processing channel
request := &soRequest{conn: conn, payload: payload}
select {
case c.requests <- request:
default:
writeError(conn, header.SeqId, "server busy")
}
}
func encodeDeliverMessage(msg *postal.Message) (bytes []byte, err error) {
@ -218,7 +251,7 @@ func (c *GwsHandler) extractAuthVerify(conn *gwsConn, resp []byte) (sessionUid s
// 存储session到内存
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
}

1
internal/gateway_ws/gws_server/gws_conn.go

@ -5,6 +5,7 @@ import (
)
type gwsConn struct {
uid string
conn *gws.Conn
}

8
internal/postal/logic/postal_server.go

@ -95,7 +95,7 @@ func (s *PostalServer) processSessionStoreEvent(ctx context.Context) {
return
case uid := <-s.sessionStore.OnStore():
logger.Info("uid %s online", uid)
//logger.Infof("uid %s online", uid)
// publish mq online event todo delay 5s
online := &event.Online{Uid: uid}
data, err := online.MarshalBinary()
@ -108,7 +108,7 @@ func (s *PostalServer) processSessionStoreEvent(ctx context.Context) {
}
case channel := <-s.sessionStore.OnDelete():
logger.Info("uid %s offline", channel.Uid)
//logger.Infof("uid %s offline", channel.Uid)
for _, gid := range channel.Groups() {
if g, ok := s.groupStore.Load(gid); ok {
g.Leave(channel.Uid)
@ -179,8 +179,9 @@ func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverB
}
}
res = &postal.ResDeliver{Ok: true}
if len(redeliverReceivers) == 0 {
return &postal.ResDeliver{Ok: true}, nil
return
}
// 向一致性 hash 下一个节点传递
@ -191,7 +192,6 @@ func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverB
Offset: offset,
})
if err == nil {
res = &postal.ResDeliver{Ok: true}
return
}
}

1
pkg/config/config.go

@ -22,6 +22,7 @@ type Configuration struct {
}
type GrpcConfig struct {
Log bool
Address string
NoReflection bool
MaxSendMsgSize string

4
pkg/config/loader.go

@ -25,8 +25,6 @@ func parseConfPathFlag(confPath string) (filePath, fileName, confName, confType
// appConf service custom config
// return common service configuration
func LoadConfig(appConf any, confPathArg ...string) *Configuration {
InitLogger()
confPath := ""
confName := "config"
confType := "toml"
@ -83,6 +81,8 @@ func LoadConfig(appConf any, confPathArg ...string) *Configuration {
panic(err)
}
}
InitLogger(conf.Grpc.Log)
return conf
}

6
pkg/config/logger.go

@ -6,11 +6,15 @@ import (
"sonet/pkg/utils/logger"
)
func InitLogger() {
func InitLogger(grpcLog bool) {
logrus.SetFormatter(&logrus.TextFormatter{
ForceColors: true,
TimestampFormat: "2006-01-02 15:04:05", //时间格式
FullTimestamp: true,
})
// logging grpc
if grpcLog {
grpclog.SetLoggerV2(logger.Logger)
}
}

20
pkg/grpc/generic/generic_client_factory.go

@ -4,13 +4,13 @@ import (
"context"
"fmt"
"google.golang.org/grpc"
"sync"
"sonet/pkg/utils/collect"
)
type GrpcGenericClientFactory struct {
scheme string
defaultOpts []grpc.DialOption
clientCache *sync.Map
clientCache *collect.ConcurrentMap[string, *GrpcGenericClient]
}
func NewGpcGenericClientFactory(scheme string, defaultOpts ...grpc.DialOption) *GrpcGenericClientFactory {
@ -21,7 +21,8 @@ func NewGpcGenericClientFactory(scheme string, defaultOpts ...grpc.DialOption) *
}
func (f *GrpcGenericClientFactory) Init() {
f.clientCache = &sync.Map{}
// f.clientCache = &sync.Map{}
f.clientCache = collect.NewConcurrentMap[string, *GrpcGenericClient](8, func(serviceName string) string { return serviceName })
}
func (f *GrpcGenericClientFactory) NewClient(ctx context.Context, serviceName string, opts ...grpc.DialOption) (client *GrpcGenericClient, err error) {
@ -39,15 +40,8 @@ func (f *GrpcGenericClientFactory) NewClient(ctx context.Context, serviceName st
}
func (f *GrpcGenericClientFactory) GetClient(ctx context.Context, serviceName string, opts ...grpc.DialOption) (client *GrpcGenericClient, err error) {
val, ok := f.clientCache.Load(serviceName)
if ok {
client = val.(*GrpcGenericClient)
return
}
client, err = f.NewClient(ctx, serviceName, opts...)
if err != nil {
return
}
f.clientCache.Store(serviceName, client)
client, err, _ = f.clientCache.ComputeIfAbsentE(serviceName, func(serviceName string) (*GrpcGenericClient, error) {
return f.NewClient(ctx, serviceName, opts...)
})
return
}

5
pkg/grpc/interceptor/recover_interceptor.go

@ -7,9 +7,10 @@ import (
"google.golang.org/grpc/grpclog"
"google.golang.org/protobuf/types/known/emptypb"
"runtime/debug"
"sonet/pkg/utils/logger"
)
func RecoverInterceptor(ctx context.Context, req any, _ *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (resp any, err error) {
func RecoverInterceptor(ctx context.Context, req any, server *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (resp any, err error) {
defer func() {
if r := recover(); r != nil {
switch r.(type) {
@ -29,6 +30,8 @@ func RecoverInterceptor(ctx context.Context, req any, _ *grpc.UnaryServerInfo, h
if err == nil {
if empty, ok := resp.(*emptypb.Empty); ok && empty == nil {
resp = &emptypb.Empty{} // grpc: error while marshaling: proto: Marshal called with nil
} else if resp == nil {
logger.Warningf("grpc request: has no response: method=%s, %v", server.FullMethod, req)
}
}
return

8
pkg/protocol/protocol.go

@ -61,7 +61,9 @@ func DecodeSo(bytes []byte) (payload *Payload, err error) {
header.SeqId = int32(binary.BigEndian.Uint32(bytes[4:8]))
// error message
if header.Type == TypeError {
payload.Body = bytes[8:]
body := bytes[8:]
payload.Body = make([]byte, len(body))
copy(payload.Body, body)
return
}
@ -84,7 +86,9 @@ func DecodeSo(bytes []byte) (payload *Payload, err error) {
return
}
payload.Body = bytes[cursor:]
body := bytes[cursor:]
payload.Body = make([]byte, len(body))
copy(payload.Body, body)
return
}

3
pkg/protocol/session/channel.go

@ -13,6 +13,7 @@ type NetConn interface {
Write([]byte) error
}
// Channel todo close channel
type Channel struct {
Uid string
Conn NetConn
@ -25,7 +26,7 @@ func NewChannel(uid string, conn NetConn) *Channel {
return &Channel{
Uid: uid,
Conn: conn,
ch: make(chan *postal.Message, 16),
ch: make(chan *postal.Message, 32),
lock: &sync.Mutex{},
}
}

2
pkg/protocol/session/store.go

@ -42,7 +42,7 @@ func (c *mapStore) Delete(uid string) {
select {
case c.deleteCh <- ch:
default:
logger.Warning("session store delete channel fulled, %s", uid)
logger.Warningf("session store delete channel fulled, %s", uid)
}
}

20
pkg/utils/collect/concurrent_map.go

@ -92,10 +92,12 @@ func (m *ConcurrentMap[K, V]) Range(f func(key K, value V) bool) {
}
}
func (m *ConcurrentMap[K, V]) LoadAndDelete(k K) (prev V, loaded bool) {
func (m *ConcurrentMap[K, V]) LoadAndDelete(k K) (value V, loaded bool) {
m.update(k, func(segment map[K]V) {
prev, loaded = segment[k]
value, loaded = segment[k]
if loaded {
delete(segment, k)
}
})
return
}
@ -111,6 +113,15 @@ func (m *ConcurrentMap[K, V]) Swap(k K, v V) (prev V, loaded bool) {
// ComputeIfAbsent 加载, 如果值不存在使用 mapping(k) 填充并返回
// mapped 值是否是 mapping(k) 填充的
func (m *ConcurrentMap[K, V]) ComputeIfAbsent(k K, mapping func(k K) V) (res V, mapped bool) {
res, _, mapped = m.ComputeIfAbsentE(k, func(k K) (V, error) {
return mapping(k), nil
})
return
}
// ComputeIfAbsentE 加载, 如果值不存在使用 mapping(k) 填充并返回
// mapped 值是否是 mapping(k) 填充的
func (m *ConcurrentMap[K, V]) ComputeIfAbsentE(k K, mapping func(k K) (V, error)) (res V, err error, mapped bool) {
segment := m.segment(k)
lock := m.segmentsLock[segment]
lock.RLock()
@ -131,7 +142,10 @@ func (m *ConcurrentMap[K, V]) ComputeIfAbsent(k K, mapping func(k K) V) (res V,
}
// write mapping value
res = mapping(k)
res, err = mapping(k)
if err != nil {
return
}
m.segmentsMap[segment][k] = res
mapped = true
return

62
pkg/utils/logger/logger.go

@ -2,6 +2,9 @@ package logger
import (
"github.com/sirupsen/logrus"
"runtime"
"strconv"
"strings"
)
var Logger *SoLogger
@ -14,51 +17,94 @@ type SoLogger struct {
logger *logrus.Logger
}
func (l *SoLogger) enhance(args []any) (ret []any) {
//logrus.SetReportCaller(true)
ret = args
for i := 2; i < 5; i++ {
_, file, line, ok := runtime.Caller(i)
if !ok {
return
}
if strings.HasSuffix(file, "/pkg/utils/logger/exported.go") ||
strings.HasSuffix(file, "/pkg/utils/logger/logger.go") {
continue
}
//caller = fmt.Sprintf("%s:%d ", file, line)
caller := file + ":" + strconv.Itoa(line) + " "
if arg0, ok := ret[0].(string); ok {
ret[0] = caller + arg0
return
}
ret = make([]any, len(args)+1)
ret[0] = caller
copy(ret[1:], args)
return
}
return
}
func (l *SoLogger) enhancef(format string, args []any) (retf string, ret []any) {
args2 := make([]any, len(args)+1)
args2[0] = format
copy(args2[1:], args)
ret = l.enhance(args2)
retf = ret[0].(string)
ret = ret[1:]
return
}
func (l *SoLogger) Info(args ...any) {
l.logger.Info(args...)
l.logger.Info(l.enhance(args)...)
}
func (l *SoLogger) Infoln(args ...any) {
l.logger.Infoln(args...)
l.logger.Infoln(l.enhance(args)...)
}
func (l *SoLogger) Infof(format string, args ...any) {
format, args = l.enhancef(format, args)
l.logger.Infof(format, args...)
}
func (l *SoLogger) Warning(args ...any) {
l.logger.Warning(args...)
l.logger.Warning(l.enhance(args)...)
}
func (l *SoLogger) Warningln(args ...any) {
l.logger.Warningln(args...)
l.logger.Warningln(l.enhance(args)...)
}
func (l *SoLogger) Warningf(format string, args ...any) {
format, args = l.enhancef(format, args)
l.logger.Warningf(format, args...)
}
func (l *SoLogger) Error(args ...any) {
l.logger.Error(args...)
l.logger.Error(l.enhance(args)...)
}
func (l *SoLogger) Errorln(args ...any) {
l.logger.Errorln(args...)
l.logger.Errorln(l.enhance(args)...)
}
func (l *SoLogger) Errorf(format string, args ...any) {
format, args = l.enhancef(format, args)
l.logger.Errorf(format, args...)
}
func (l *SoLogger) Fatal(args ...any) {
l.logger.Fatal(args...)
l.logger.Fatal(l.enhance(args)...)
}
func (l *SoLogger) Fatalln(args ...any) {
l.logger.Fatalln(args...)
l.logger.Fatalln(l.enhance(args)...)
}
func (l *SoLogger) Fatalf(format string, args ...any) {
format, args = l.enhancef(format, args)
l.logger.Fatalf(format, args...)
}

Loading…
Cancel
Save