package server import ( "context" "fmt" "github.com/gorilla/websocket" "google.golang.org/protobuf/proto" "regexp" "runtime/debug" "sonet/api/gen/auth" "sonet/internal/gateway_ws/session" "sonet/pkg/grpc/generic" "sonet/pkg/plugins/cache" "sonet/pkg/protocol" session2 "sonet/pkg/protocol/session" "sonet/pkg/utils/logger" "sync" "time" ) type ConnHandler struct { postalServerAddress string grpcFactory *generic.GrpcGenericClientFactory sessionStore *sync.Map // 当前连接用户,内存缓存 subjectStore cache.MultiLevelCache } func NewConnHandler(postalServerAddress string, grpcFactory *generic.GrpcGenericClientFactory, sessionStore *sync.Map, subjectStore cache.MultiLevelCache) *ConnHandler { return &ConnHandler{ grpcFactory: grpcFactory, postalServerAddress: postalServerAddress, sessionStore: sessionStore, subjectStore: subjectStore, } } func (c *ConnHandler) handleConn(conn *websocket.Conn) { client := session.NewNetClient(conn, ReadDeadline, WriteDeadline) defer func() { // 捕获其他错误 if r := recover(); r != nil { logger.Error("NetClient recover error: ", r) // 输出堆栈信息 logger.Error("NetClient recover error stack: ", string(debug.Stack())) } }() defer func() { client.Close() // TODO 连接关闭,mq发送关闭事件 // 删除连接 if client.Subject != nil { c.sessionStore.Delete(client.Subject.Uid) err := c.subjectStore.Del(context.Background(), client.Subject.Uid) if err != nil { logger.Error("del cluster subject store uid error: ", client.Subject.Uid) } logger.Info("subject offline: ", client.Subject.Uid) client.Subject = nil } }() for { ignore, message, err := client.ReadMessage() if err != nil { if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway) { logger.Error("unexpected close error: ", err) } else { // 读失败 logger.Errorf("NetClient ReadMessage error: %T, %v", err, err) } return } if ignore { continue } payload, err := protocol.DecodeSo(message) if err != nil { logger.Errorf("decode message error: len=%d", len(message), err) return } header := payload.Header service, method := header.Svc, header.Target if client.Subject == nil && !(service == "Auth" && method == "Verify") { writeError(client, header.SeqId, session2.UnauthorizedRequestError.Error()) continue } // grpc generic call ctx := context.Background() grpcClient, err := c.grpcFactory.GetClient(ctx, service) if err != nil { logger.Error("get grpc generic client error: ", err) writeError(client, header.SeqId, fmt.Sprintf("service %s not avaliable: %s", header.Svc, err.Error())) continue } // put session if client.Subject != nil { ctx = session2.PutSubject(ctx, session2.NewRpcSubject(client.Subject.Uid)) } resp, err := grpcClient.InvokeUnary(ctx, header.Target, payload.Body) if err != nil { // todo 提取 err: fmt.Sprintf("rpc error: code = %s desc = %s", s.Code(), s.Message()) writeError(client, header.SeqId, unwrapRpcError(err.Error())) logger.Error("grpc generic call error: ", err) continue } // 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(client, header.SeqId, "server error") continue } header.Target = resp.XXX_MessageName() } // auth verify success if service == "Auth" && method == "Verify" { err = c.extractAuthVerify(client, payload.Body) if err != nil { logger.Error("extractAuthVerify error: ", err) writeError(client, header.SeqId, "server error") continue } } header.Type = protocol.TypeResponse resMessage, err := protocol.EncodeSo(payload) if err != nil { logger.Error("grpc generic call error: ", err) continue } client.MustWrite(resMessage) } } func (c *ConnHandler) extractAuthVerify(netClient *session.NetClient, resp []byte) (err error) { authSubject := &auth.Subject{} err = proto.Unmarshal(resp, authSubject) if err != nil { logger.Error("grpc generic call error: ", err) return } // 设置连接身份信息 subject := &session.Subject{ Uid: authSubject.Uid, Online: 1, Time: time.Now().UnixMilli(), Gate: c.postalServerAddress, } // 存储session err = c.subjectStore.Set(context.Background(), subject.Uid, subject) // store to redis cluster cache if err != nil { return } // 关闭旧的链接 oldOnline, ok := c.sessionStore.Load(subject.Uid) if ok { oldOnline.(*session.NetSubject).Client.Close() // TODO nats offline / force load subject target cluster call offline logger.Infof("close subject old conn: %s\n", subject.Uid) } netClient.Subject = session.NewNetSubject(subject.Uid, netClient) c.sessionStore.Store(netClient.Subject.Uid, netClient.Subject) logger.Infof("subject online: %s\n", subject.Uid) return } var reg = regexp.MustCompile("^rpc error:.*?desc ?= ?(.*)$") func unwrapRpcError(err string) string { finds := reg.FindStringSubmatch(err) if len(finds) > 1 { return finds[1] } return err } func writeError(client *session.NetClient, 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 } client.MustWrite(message) return }