package gws_server import ( "context" "fmt" "github.com/lxzan/gws" "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" "time" ) const ( PingInterval = 5 * time.Second PingWait = 10 * time.Second SessionGwsConnKey = "gws" ) type soRequest struct { conn *gwsConn payload *protocol.Payload } type GwsHandler struct { grpcFactory *generic.GrpcGenericClientFactory sessionStore session.Store // 当前连接用户,内存缓存 requests chan *soRequest } func NewGwsHandler( grpcFactory *generic.GrpcGenericClientFactory, sessionStore session.Store, ) *GwsHandler { return &GwsHandler{ grpcFactory: grpcFactory, sessionStore: sessionStore, } } func (c *GwsHandler) Init(ctx context.Context) { c.requests = make(chan *soRequest, 1024*16) // todo 针对不同服务的协程池 for i := 0; i < 32; i++ { go c.dispatchRequest(ctx) } } // 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 // grpc generic call ctx := context.Background() ctx2, cancel2 := context.WithTimeout(ctx, time.Second*3) defer cancel2() grpcClient, err := c.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 = c.extractAuthVerify(conn, payload.Body) if err != nil { logger.Error("extractAuthVerify error: ", err) writeError(conn, header.SeqId, "server error") return } // channel todo dispatch on conn open go c.dispatch(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 (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) { header := &protocol.Header{ Magic: protocol.Magic, Type: protocol.TypeNotice, UrlType: 1, SerializeType: 1, Svc: msg.Svc, Target: msg.Msg, } payload := &protocol.Payload{Header: header, Body: msg.Body} bytes, err = protocol.EncodeSo(payload) if err != nil { logger.Error("encode deliver message error: ", err) } return } // todo close func (c *GwsHandler) dispatch(channel *session.Channel) { for { msg := channel.Ready() bytes, err := encodeDeliverMessage(msg) if err != nil { logger.Error("encode dispatch message error:", err) continue } err = channel.Conn.Write(bytes) if err != nil { logger.Error("write dispatch message error:", err) } } } func (c *GwsHandler) extractAuthVerify(conn *gwsConn, resp []byte) (sessionUid string, channel *session.Channel, err error) { authSubject := &auth.Subject{} 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 { oldChannel.Conn.(*gwsConn).conn.WriteClose(1000, nil) 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 } 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 } 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 }