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.
 
 

289 lines
7.4 KiB

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 // <string, *session.NetSubject> 当前连接用户,内存缓存
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
}