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
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 |
|
}
|
|
|