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.
256 lines
6.5 KiB
256 lines
6.5 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 |
|
SessionUidKey = "uid" |
|
SessionGwsConnKey = "gws" |
|
) |
|
|
|
type GwsHandler struct { |
|
postalServerAddress string |
|
grpcFactory *generic.GrpcGenericClientFactory |
|
sessionStore session.Store // <string, *session.NetSubject> 当前连接用户,内存缓存 |
|
} |
|
|
|
func NewGwsHandler(postalServerAddress string, |
|
grpcFactory *generic.GrpcGenericClientFactory, |
|
sessionStore session.Store, |
|
) *GwsHandler { |
|
return &GwsHandler{ |
|
grpcFactory: grpcFactory, |
|
postalServerAddress: postalServerAddress, |
|
sessionStore: sessionStore, |
|
} |
|
} |
|
|
|
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(SessionUidKey) |
|
if !ok { |
|
return |
|
} |
|
uid := val.(string) |
|
c.sessionStore.Delete(uid) |
|
logger.Info("subject offline: ", 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 { |
|
return |
|
} |
|
|
|
val, ok := socket.Session().Load(SessionGwsConnKey) |
|
if !ok { |
|
logger.Error("gws socket session error: conn not exists") |
|
return |
|
} |
|
conn := val.(*gwsConn) |
|
|
|
val, authorized := socket.Session().Load(SessionUidKey) |
|
var sessionUid string |
|
if authorized { |
|
sessionUid = val.(string) |
|
} |
|
|
|
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 |
|
} |
|
|
|
// 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 session |
|
if authorized { |
|
ctx = session.PutSubject(ctx, session.NewRpcSubject(sessionUid)) |
|
} |
|
|
|
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 |
|
sessionUid, channel, err = c.extractAuthVerify(conn, payload.Body) |
|
if err != nil { |
|
logger.Error("extractAuthVerify error: ", err) |
|
writeError(conn, header.SeqId, "server error") |
|
return |
|
} |
|
socket.Session().Store(SessionUidKey, sessionUid) |
|
// channel |
|
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 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 |
|
}
|
|
|