Browse Source

gnet gateway

master
tangmingyou 3 years ago
parent
commit
ea9b49c82f
  1. 2
      cmd/gateway_ws/config.toml
  2. 33
      cmd/gateway_ws/main.go
  3. 4
      internal/gateway_ws/gnet_server/codec.go
  4. 182
      internal/gateway_ws/gnet_server/websocket.go
  5. 4
      pkg/config/logger.go
  6. 20
      pkg/grpc/generic/generic_client_factory.go
  7. 14
      pkg/utils/collect/concurrent_map.go

2
cmd/gateway_ws/config.toml

@ -1,6 +1,6 @@
[app]
httpPort = 7001
endpointAddress = "192.168.110.41:7001"
endpointAddress = "192.168.1.3:7001"
subjectCacheTopic = "wsgate:subject:"
subjectLrcExpiration = "10m"
subjectLrcCleanupInterval = "5m"

33
cmd/gateway_ws/main.go

@ -2,17 +2,18 @@ package main
import (
"context"
"fmt"
"github.com/panjf2000/gnet/v2"
clientv3 "go.etcd.io/etcd/client/v3"
"google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure"
"runtime"
"sonet/internal/gateway_ws/gws_server"
"sonet/internal/gateway_ws/gnet_server"
"sonet/internal/postal/group"
"sonet/internal/postal/logic"
"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"
@ -59,23 +60,37 @@ func main() {
groupStore := collect.NewConcurrentMap[string, *group.Group](128, func(k string) string { return k })
// 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)
httpServer := gws_server.NewGwsServer(gwsHandler)
httpServer.Init()
//gwsHandler := gws_server.NewGwsHandler(grpcFactory, sessionStore)
//httpServer := gws_server.NewGwsServer(gwsHandler)
//httpServer.Init()
//go func() {
// err = httpServer.Run(appConf.HttpPort)
// if err != nil {
// panic(err)
// }
//}()
ctx, cancel := context.WithCancel(context.Background())
shutdown.AddHook(cancel)
go func() {
err = httpServer.Run(appConf.HttpPort)
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)

4
internal/gateway_ws/gnet_server/codec.go

@ -35,7 +35,7 @@ func (w *wsCodec) upgrade(c gnet.Conn) (ok bool, action gnet.Action) {
buf := &w.buf
tmpReader := bytes.NewReader(buf.Bytes())
oldLen := tmpReader.Len()
logging.Infof("do Upgrade")
// logging.Infof("do Upgrade")
hs, err := ws.Upgrade(readWrite{tmpReader, c})
skipN := oldLen - tmpReader.Len()
@ -75,7 +75,7 @@ func (w *wsCodec) readBufferBytes(c gnet.Conn) gnet.Action {
return gnet.None
}
func (w *wsCodec) Decode(c gnet.Conn) (outs []wsutil.Message, err error) {
fmt.Println("do Decode")
// fmt.Println("do Decode")
messages, err := w.readWsMessages()
if err != nil {
logging.Infof("Error reading message! %v", err)

182
internal/gateway_ws/gnet_server/websocket.go

@ -2,10 +2,13 @@ package gnet_server
import (
"context"
"flag"
"fmt"
"github.com/gobwas/ws"
"log"
"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"
@ -16,6 +19,11 @@ import (
"github.com/panjf2000/gnet/v2/pkg/logging"
)
type gnetRequest struct {
conn *gnetConn
payload *protocol.Payload
}
type GnetWsServer struct {
gnet.BuiltinEventEngine
@ -23,17 +31,28 @@ type GnetWsServer struct {
multicore bool
eng gnet.Engine
connected int64
requests chan *protocol.Payload
grpcFactory *generic.GrpcGenericClientFactory
sessionStore session.Store // <string, *session.NetSubject> 当前连接用户,内存缓存
requests chan *gnetRequest
}
func NewGnetWsServer() *GnetWsServer {
func NewGnetWsServer(
grpcFactory *generic.GrpcGenericClientFactory,
sessionStore session.Store) *GnetWsServer {
return &GnetWsServer{
requests: make(chan *protocol.Payload, 1024*16),
grpcFactory: grpcFactory,
sessionStore: sessionStore,
}
}
func (wss *GnetWsServer) Init() {
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
@ -42,8 +61,107 @@ func (wss *GnetWsServer) dispatch(ctx context.Context) {
select {
case <-ctx.Done():
return
case payload := <-wss.requests:
fmt.Println(payload.Header)
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)
}
}
}
@ -120,8 +238,9 @@ func (wss *GnetWsServer) OnTraffic(c gnet.Conn) (action gnet.Action) {
}
// payload -> processing channel
request := &gnetRequest{conn: conn, payload: payload}
select {
case wss.requests <- payload:
case wss.requests <- request:
default:
writeError(conn, header.SeqId, "server busy")
}
@ -146,6 +265,31 @@ func (wss *GnetWsServer) OnTick() (delay time.Duration, action gnet.Action) {
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}
@ -168,17 +312,13 @@ func writeError(conn session.NetConn, seqId int32, errMsg string) {
return
}
func main() {
var port int
var multicore bool
// Example command: go run main.go --port 8080 --multicore=true
flag.IntVar(&port, "port", 9080, "server port")
flag.BoolVar(&multicore, "multicore", true, "multicore")
flag.Parse()
var reg = regexp.MustCompile("^rpc error:.*?desc ?= ?(.*)$")
wss := &GnetWsServer{addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), multicore: multicore}
// Start serving!
log.Println("server exits:", gnet.Run(wss, wss.addr, gnet.WithMulticore(multicore), gnet.WithReusePort(true), gnet.WithTicker(true)))
// 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
}

4
pkg/config/logger.go

@ -2,8 +2,6 @@ package config
import (
"github.com/sirupsen/logrus"
"google.golang.org/grpc/grpclog"
"sonet/pkg/utils/logger"
)
func InitLogger() {
@ -12,5 +10,5 @@ func InitLogger() {
TimestampFormat: "2006-01-02 15:04:05", //时间格式
FullTimestamp: true,
})
grpclog.SetLoggerV2(logger.Logger)
// 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
}

14
pkg/utils/collect/concurrent_map.go

@ -111,6 +111,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 +140,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

Loading…
Cancel
Save