diff --git a/cmd/gateway_ws/config.toml b/cmd/gateway_ws/config.toml index 8443775..2d33c99 100644 --- a/cmd/gateway_ws/config.toml +++ b/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" diff --git a/cmd/gateway_ws/main.go b/cmd/gateway_ws/main.go index bd12434..4786c33 100644 --- a/cmd/gateway_ws/main.go +++ b/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) diff --git a/internal/gateway_ws/gnet_server/codec.go b/internal/gateway_ws/gnet_server/codec.go index 2379aaa..2d01e69 100644 --- a/internal/gateway_ws/gnet_server/codec.go +++ b/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) diff --git a/internal/gateway_ws/gnet_server/websocket.go b/internal/gateway_ws/gnet_server/websocket.go index 927c907..6b748b2 100644 --- a/internal/gateway_ws/gnet_server/websocket.go +++ b/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 // 当前连接用户,内存缓存 + 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 } diff --git a/pkg/config/logger.go b/pkg/config/logger.go index 9d9e437..019a935 100644 --- a/pkg/config/logger.go +++ b/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) } diff --git a/pkg/grpc/generic/generic_client_factory.go b/pkg/grpc/generic/generic_client_factory.go index d6974a2..db04a4e 100644 --- a/pkg/grpc/generic/generic_client_factory.go +++ b/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 } diff --git a/pkg/utils/collect/concurrent_map.go b/pkg/utils/collect/concurrent_map.go index 6a83f87..3cff73e 100644 --- a/pkg/utils/collect/concurrent_map.go +++ b/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