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.
 
 

188 lines
5.1 KiB

package server
import (
"context"
"github.com/gorilla/websocket"
"google.golang.org/protobuf/proto"
"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 // <string, *session.NetSubject> 当前连接用户,内存缓存
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
if client.Subject == nil && !(header.Svc == "Auth" && header.Target == "Verify") {
writeError(client, header.SeqId, session2.UnauthorizedRequestError.Error())
continue
}
// grpc generic call
ctx := context.Background()
grpcClient, err := c.grpcFactory.GetClient(ctx, header.Svc)
if err != nil {
logger.Error("get grpc generic client error: ", err)
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, err.Error())
logger.Error("grpc generic call error: ", err)
continue
}
// write response
payload.Body, err = resp.Marshal()
if err != nil {
logger.Error("generic call response marshal error: ", err)
writeError(client, header.SeqId, "server error")
continue
}
// auth verify success
if header.Svc == "Auth" && header.Target == "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
header.Target = resp.XXX_MessageName()
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
}
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
}