You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 

220 lines
5.8 KiB

package server
import (
"context"
"fmt"
"github.com/gorilla/websocket"
"google.golang.org/protobuf/proto"
"regexp"
"runtime/debug"
"sonet/api/gen/auth"
"sonet/pkg/grpc/generic"
"sonet/pkg/plugins/cache"
"sonet/pkg/protocol"
"sonet/pkg/protocol/session"
"sonet/pkg/utils/logger"
"time"
)
type ConnHandler struct {
postalServerAddress string
grpcFactory *generic.GrpcGenericClientFactory
sessionStore session.Store // <string, *session.NetSubject> 当前连接用户,内存缓存
subjectStore cache.MultiLevelCache
}
func NewConnHandler(postalServerAddress string,
grpcFactory *generic.GrpcGenericClientFactory,
sessionStore session.Store,
subjectStore cache.MultiLevelCache) *ConnHandler {
return &ConnHandler{
grpcFactory: grpcFactory,
postalServerAddress: postalServerAddress,
sessionStore: sessionStore,
subjectStore: subjectStore,
}
}
func (c *ConnHandler) handleConn(conn *websocket.Conn) {
wsConn := newWsConn(conn)
var sessionUid string
defer func() {
// 捕获其他错误
if r := recover(); r != nil {
logger.Error("NetClient recover error: ", r)
// 输出堆栈信息
logger.Error("NetClient recover error stack: ", string(debug.Stack()))
}
}()
defer func() {
if err := conn.Close(); err != nil {
logger.Error("conn close: ", err)
}
// TODO 连接关闭,mq发送关闭事件
// 删除连接
if sessionUid != "" {
c.sessionStore.Delete(sessionUid)
err := c.subjectStore.Del(context.Background(), sessionUid)
if err != nil {
logger.Error("del cluster subject store uid error: ", sessionUid)
}
logger.Info("subject offline: ", sessionUid)
}
}()
for {
if err := conn.SetReadDeadline(time.Time{}); err != nil {
return
}
t, message, err := conn.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 t != websocket.BinaryMessage {
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 sessionUid == "" && !(service == "Auth" && method == "Verify") {
writeError(wsConn, header.SeqId, session.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(wsConn, header.SeqId, fmt.Sprintf("service %s not avaliable: %s", header.Svc, err.Error()))
continue
}
// put session
if sessionUid != "" {
ctx = session.PutSubject(ctx, session.NewRpcSubject(sessionUid))
}
resp, err := grpcClient.InvokeUnary(ctx, header.Target, payload.Body)
if err != nil {
writeError(wsConn, 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(wsConn, header.SeqId, "server error")
continue
}
header.Target = resp.XXX_MessageName()
}
// auth verify success
if service == "Auth" && method == "Verify" {
sessionUid, err = c.extractAuthVerify(wsConn, payload.Body)
if err != nil {
logger.Error("extractAuthVerify error: ", err)
writeError(wsConn, 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
}
err = wsConn.Write(resMessage)
if err != nil {
logger.Errorf("write 2 %s response error: \n", sessionUid, err)
}
}
}
func (c *ConnHandler) extractAuthVerify(conn *wsConn, resp []byte) (sessionUid string, err error) {
authSubject := &auth.Subject{}
err = proto.Unmarshal(resp, authSubject)
if err != nil {
logger.Error("grpc generic call error: ", err)
return
}
// 设置连接身份信息
sessionUid = authSubject.Uid
gateSubject := &session.GateSubject{
Uid: authSubject.Uid,
Online: 1,
Time: time.Now().UnixMilli(),
Gate: c.postalServerAddress,
}
// 存储session
err = c.subjectStore.Set(context.Background(), sessionUid, gateSubject) // store to redis cluster cache
if err != nil {
return
}
// 关闭旧的链接
oldConn, ok := c.sessionStore.Load(sessionUid)
if ok {
// TODO nats offline / force load subject target cluster call offline
if err = oldConn.(*wsConn).conn.Close(); err != nil {
logger.Error("old conn close error: ", err)
}
logger.Infof("close subject old conn: %s\n", sessionUid)
}
// 存储session到内存
c.sessionStore.Store(sessionUid, conn)
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
}