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] [app]
httpPort = 7001 httpPort = 7001
endpointAddress = "192.168.110.41:7001" endpointAddress = "192.168.1.3:7001"
subjectCacheTopic = "wsgate:subject:" subjectCacheTopic = "wsgate:subject:"
subjectLrcExpiration = "10m" subjectLrcExpiration = "10m"
subjectLrcCleanupInterval = "5m" subjectLrcCleanupInterval = "5m"

33
cmd/gateway_ws/main.go

@ -2,17 +2,18 @@ package main
import ( import (
"context" "context"
"fmt"
"github.com/panjf2000/gnet/v2"
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" "runtime"
"sonet/internal/gateway_ws/gws_server" "sonet/internal/gateway_ws/gnet_server"
"sonet/internal/postal/group" "sonet/internal/postal/group"
"sonet/internal/postal/logic" "sonet/internal/postal/logic"
"sonet/pkg/config" "sonet/pkg/config"
"sonet/pkg/grpc/client" "sonet/pkg/grpc/client"
"sonet/pkg/grpc/discovery" "sonet/pkg/grpc/discovery"
"sonet/pkg/grpc/discovery/etcd"
"sonet/pkg/grpc/generic" "sonet/pkg/grpc/generic"
"sonet/pkg/plugins/mq" "sonet/pkg/plugins/mq"
"sonet/pkg/protocol/deliver" "sonet/pkg/protocol/deliver"
@ -59,23 +60,37 @@ func main() {
groupStore := collect.NewConcurrentMap[string, *group.Group](128, func(k string) string { return k }) groupStore := collect.NewConcurrentMap[string, *group.Group](128, func(k string) string { return k })
// run websocket server // run websocket server
postalAddr := etcd.MustRegisterAddress(conf.Grpc.Address) //postalAddr := etcd.MustRegisterAddress(conf.Grpc.Address)
//connHandler := server.NewConnHandler(postalAddr, grpcFactory, sessionStore, subjectStore) //connHandler := server.NewConnHandler(postalAddr, grpcFactory, sessionStore, subjectStore)
//httpServer := server.NewHttpServer(connHandler) //httpServer := server.NewHttpServer(connHandler)
gwsHandler := gws_server.NewGwsHandler(postalAddr, grpcFactory, sessionStore) //gwsHandler := gws_server.NewGwsHandler(grpcFactory, sessionStore)
httpServer := gws_server.NewGwsServer(gwsHandler) //httpServer := gws_server.NewGwsServer(gwsHandler)
httpServer.Init() //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() { 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 { if err != nil {
panic(err) panic(err)
} }
}() }()
// run postal cluster server // run postal cluster server
ctx, cancel := context.WithCancel(context.Background())
shutdown.AddHook(cancel)
postalPicker := deliver.NewPostalPicker(dis, grpc.WithTransportCredentials(insecure.NewCredentials())) postalPicker := deliver.NewPostalPicker(dis, grpc.WithTransportCredentials(insecure.NewCredentials()))
if err = postalPicker.Init(ctx); err != nil { if err = postalPicker.Init(ctx); err != nil {
panic(err) 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 buf := &w.buf
tmpReader := bytes.NewReader(buf.Bytes()) tmpReader := bytes.NewReader(buf.Bytes())
oldLen := tmpReader.Len() oldLen := tmpReader.Len()
logging.Infof("do Upgrade") // logging.Infof("do Upgrade")
hs, err := ws.Upgrade(readWrite{tmpReader, c}) hs, err := ws.Upgrade(readWrite{tmpReader, c})
skipN := oldLen - tmpReader.Len() skipN := oldLen - tmpReader.Len()
@ -75,7 +75,7 @@ func (w *wsCodec) readBufferBytes(c gnet.Conn) gnet.Action {
return gnet.None return gnet.None
} }
func (w *wsCodec) Decode(c gnet.Conn) (outs []wsutil.Message, err error) { func (w *wsCodec) Decode(c gnet.Conn) (outs []wsutil.Message, err error) {
fmt.Println("do Decode") // fmt.Println("do Decode")
messages, err := w.readWsMessages() messages, err := w.readWsMessages()
if err != nil { if err != nil {
logging.Infof("Error reading message! %v", err) logging.Infof("Error reading message! %v", err)

182
internal/gateway_ws/gnet_server/websocket.go

@ -2,10 +2,13 @@ package gnet_server
import ( import (
"context" "context"
"flag"
"fmt" "fmt"
"github.com/gobwas/ws" "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"
"sonet/pkg/protocol/session" "sonet/pkg/protocol/session"
"sonet/pkg/utils/logger" "sonet/pkg/utils/logger"
@ -16,6 +19,11 @@ import (
"github.com/panjf2000/gnet/v2/pkg/logging" "github.com/panjf2000/gnet/v2/pkg/logging"
) )
type gnetRequest struct {
conn *gnetConn
payload *protocol.Payload
}
type GnetWsServer struct { type GnetWsServer struct {
gnet.BuiltinEventEngine gnet.BuiltinEventEngine
@ -23,17 +31,28 @@ type GnetWsServer struct {
multicore bool multicore bool
eng gnet.Engine eng gnet.Engine
connected int64 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{ 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 // dispatch 多个 goroutine 进行请求 dispatch
@ -42,8 +61,107 @@ func (wss *GnetWsServer) dispatch(ctx context.Context) {
select { select {
case <-ctx.Done(): case <-ctx.Done():
return return
case payload := <-wss.requests: case request := <-wss.requests:
fmt.Println(payload.Header) 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 // payload -> processing channel
request := &gnetRequest{conn: conn, payload: payload}
select { select {
case wss.requests <- payload: case wss.requests <- request:
default: default:
writeError(conn, header.SeqId, "server busy") 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 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) { func writeError(conn session.NetConn, seqId int32, errMsg string) {
header := &protocol.Header{} header := &protocol.Header{}
payload := &protocol.Payload{Header: header} payload := &protocol.Payload{Header: header}
@ -168,17 +312,13 @@ func writeError(conn session.NetConn, seqId int32, errMsg string) {
return return
} }
func main() { var reg = regexp.MustCompile("^rpc error:.*?desc ?= ?(.*)$")
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()
wss := &GnetWsServer{addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), multicore: multicore} // unwrapRpcError 提取 grpc err: fmt.Sprintf("rpc error: code = %s desc = %s", s.Code(), s.Message())
func unwrapRpcError(err string) string {
// Start serving! finds := reg.FindStringSubmatch(err)
log.Println("server exits:", gnet.Run(wss, wss.addr, gnet.WithMulticore(multicore), gnet.WithReusePort(true), gnet.WithTicker(true))) if len(finds) > 1 {
return finds[1]
}
return err
} }

4
pkg/config/logger.go

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

20
pkg/grpc/generic/generic_client_factory.go

@ -4,13 +4,13 @@ import (
"context" "context"
"fmt" "fmt"
"google.golang.org/grpc" "google.golang.org/grpc"
"sync" "sonet/pkg/utils/collect"
) )
type GrpcGenericClientFactory struct { type GrpcGenericClientFactory struct {
scheme string scheme string
defaultOpts []grpc.DialOption defaultOpts []grpc.DialOption
clientCache *sync.Map clientCache *collect.ConcurrentMap[string, *GrpcGenericClient]
} }
func NewGpcGenericClientFactory(scheme string, defaultOpts ...grpc.DialOption) *GrpcGenericClientFactory { func NewGpcGenericClientFactory(scheme string, defaultOpts ...grpc.DialOption) *GrpcGenericClientFactory {
@ -21,7 +21,8 @@ func NewGpcGenericClientFactory(scheme string, defaultOpts ...grpc.DialOption) *
} }
func (f *GrpcGenericClientFactory) Init() { 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) { 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) { func (f *GrpcGenericClientFactory) GetClient(ctx context.Context, serviceName string, opts ...grpc.DialOption) (client *GrpcGenericClient, err error) {
val, ok := f.clientCache.Load(serviceName) client, err, _ = f.clientCache.ComputeIfAbsentE(serviceName, func(serviceName string) (*GrpcGenericClient, error) {
if ok { return f.NewClient(ctx, serviceName, opts...)
client = val.(*GrpcGenericClient) })
return
}
client, err = f.NewClient(ctx, serviceName, opts...)
if err != nil {
return
}
f.clientCache.Store(serviceName, client)
return 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) 填充并返回 // ComputeIfAbsent 加载, 如果值不存在使用 mapping(k) 填充并返回
// mapped 值是否是 mapping(k) 填充的 // mapped 值是否是 mapping(k) 填充的
func (m *ConcurrentMap[K, V]) ComputeIfAbsent(k K, mapping func(k K) V) (res V, mapped bool) { 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) segment := m.segment(k)
lock := m.segmentsLock[segment] lock := m.segmentsLock[segment]
lock.RLock() lock.RLock()
@ -131,7 +140,10 @@ func (m *ConcurrentMap[K, V]) ComputeIfAbsent(k K, mapping func(k K) V) (res V,
} }
// write mapping value // write mapping value
res = mapping(k) res, err = mapping(k)
if err != nil {
return
}
m.segmentsMap[segment][k] = res m.segmentsMap[segment][k] = res
mapped = true mapped = true
return return

Loading…
Cancel
Save