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.
 
 

338 lines
9.2 KiB

package logic
import (
"context"
"errors"
"google.golang.org/grpc"
"google.golang.org/grpc/reflection"
"google.golang.org/protobuf/types/known/emptypb"
"net"
"sonet/api/gen/postal"
"sonet/internal/postal/group"
"sonet/pkg/config"
"sonet/pkg/grpc/client"
"sonet/pkg/grpc/discovery"
"sonet/pkg/grpc/interceptor"
"sonet/pkg/plugins/cache"
"sonet/pkg/plugins/mq"
"sonet/pkg/protocol"
"sonet/pkg/protocol/session"
"sonet/pkg/utils/collect"
"sonet/pkg/utils/logger"
"sonet/pkg/utils/shutdown"
)
type PostalServer struct {
postal.UnimplementedPostalServer
endpointAddress string // websocket 前端连接地址
broadcastAddress string // postal server 在注册中心注册的地址, 其他服务可直连访问
sessionStore session.Store // k:uid 在线用户conn存储
groupStore *collect.ConcurrentMap[string, *group.Group] // k:gid 在线用户 groups 存储
subjectStore cache.MultiLevelCache // online subject address cache, ws gateway集群所有在线用户存储
clientFactory *client.GrpcDirectClientFactory
producer mq.Producer
}
func NewPostalServer(
endpointAddress string,
sessionStore session.Store,
groupStore *collect.ConcurrentMap[string, *group.Group],
subjectStore cache.MultiLevelCache,
clientFactory *client.GrpcDirectClientFactory,
producer mq.Producer) *PostalServer {
return &PostalServer{
endpointAddress: endpointAddress,
sessionStore: sessionStore,
groupStore: groupStore,
subjectStore: subjectStore,
clientFactory: clientFactory,
producer: producer,
}
}
func (s *PostalServer) Run(conf config.GrpcConfig, registry discovery.Registry) (err error) {
server := grpc.NewServer(
config.GetGrpcOptions(
conf,
grpc.UnaryInterceptor(interceptor.RecoverInterceptor),
)...,
)
if !conf.NoReflection {
// 注册反射服务
reflection.Register(server)
}
postal.RegisterPostalServer(server, s)
listen, err := net.Listen("tcp", conf.Address)
if err != nil {
return
}
// registry discovery
register := conf.Register
if register.Name == "" {
register.Name = postal.Postal_ServiceDesc.ServiceName
}
if register.Addr == "" {
register.Addr = conf.Address
}
ctx, cancel := context.WithCancel(context.Background())
err = registry.Registry(ctx, register)
if err != nil {
panic(err)
}
shutdown.AddHook(cancel)
go s.processSessionStoreEvent(ctx)
// 其他服务直连地址
s.broadcastAddress = register.Addr
// run serve
logger.Infof("%s grpc server running %s\n", register.Name, listen.Addr().String())
err = server.Serve(listen)
return
}
func (s *PostalServer) processSessionStoreEvent(ctx context.Context) {
for {
select {
case <-ctx.Done():
return
case uid := <-s.sessionStore.OnStore():
logger.Info("uid %s online", uid)
// todo mq event..., delay 5s
// s.producer.Publish()
case channel := <-s.sessionStore.OnDelete():
logger.Info("uid %s offline", channel.Uid)
for _, gid := range channel.Groups() {
if g, ok := s.groupStore.Load(gid); ok {
g.Leave(channel.Uid)
}
}
// todo mq event...
}
}
}
func (s *PostalServer) deliverMessage(msg *postal.Message, conn session.NetConn) 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("req deliver encode notice error: ", err)
return err
}
return conn.Write(bytes)
}
// receiverGates 找receiver在集群内哪些其他节点
func (s *PostalServer) clusterGateReceivers(ctx context.Context, receivers []string) (gateReceivers map[string][]string, offline []string) {
gateReceivers = make(map[string][]string)
for _, receiver := range receivers {
subject := &session.GateSubject{}
err := s.subjectStore.Load(ctx, receiver, subject)
if err != nil {
offline = append(offline, receiver)
if err != cache.NotExists {
logger.Error("load from subject store error: ", err)
offline = append(offline, receiver)
}
continue
}
if subject.Gate == s.broadcastAddress {
offline = append(offline, receiver)
// 清除失效缓存
err := s.subjectStore.Del(ctx, receiver)
if err != nil {
logger.Error("del subject store error: ", receiver, err)
}
continue
}
// put receiver gate addr
gateReceivers[subject.Gate] = append(gateReceivers[subject.Gate], receiver)
}
return
}
func (s *PostalServer) Deliver(ctx context.Context, req *postal.ReqDeliver) (res *postal.ResDeliver, err error) {
receiver, ok := s.sessionStore.Load(req.Receiver)
if ok {
var bytes []byte
bytes, err = encodeDeliverMessage(req.Msg)
if err != nil {
return
}
// write msg
err = receiver.Conn.Write(bytes)
if err != nil {
return
}
res = &postal.ResDeliver{Ok: true}
return
}
// 用户连接不在当前gateway
// todo 向一致性 hash 下一个节点传递
gateReceivers, offline := s.clusterGateReceivers(ctx, []string{req.Receiver})
if offline != nil && len(offline) > 0 {
res = &postal.ResDeliver{Code: int32(postal.DeliverResult_ReceiverOffline.Number())}
return
}
for postalAddr := range gateReceivers {
// gateway集群中转发消息
clusterAddr, err := PostalAddr2Cluster(postalAddr)
if err != nil {
logger.Errorf("parse postal server addr error: %s", postalAddr, err)
continue
}
conn, err := s.clientFactory.GetConn(context.Background(), clusterAddr)
if err != nil {
logger.Error("get postal cluster conn error: ", err)
return nil, err
}
clusterClient := postal.NewPostalClusterClient(conn)
reqRedirect := &postal.ReqRedirect{
Ttl: 3, // TODO 转发n次就丢弃
RedirectMethod: methodRedirectDeliver,
Deliver: req,
}
resp, err := clusterClient.Redirect(ctx, reqRedirect)
return resp, err
}
return nil, nil
}
func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverBatch) (res *postal.ResDeliver, err error) {
bytes, err := encodeDeliverMessage(req.Msg)
if err != nil {
return
}
var redirectReceivers []string
for _, receiverId := range req.Receivers {
receiver, ok := s.sessionStore.Load(receiverId)
if !ok {
redirectReceivers = append(redirectReceivers, receiverId)
continue
}
err = receiver.Conn.Write(bytes)
if err != nil {
logger.Errorf("deliver to %s error: ", receiver, err)
}
}
if len(redirectReceivers) == 0 {
return &postal.ResDeliver{Ok: true}, nil
}
// todo 向一致性 hash 下个节点传递
gateReceivers, offline := s.clusterGateReceivers(ctx, redirectReceivers)
if len(offline) > 0 {
logger.Warning("offline redirect receivers: ", offline)
}
if len(gateReceivers) > 0 {
for postalAddr, receivers := range gateReceivers {
// gateway集群中转发消息
clusterAddr, err := PostalAddr2Cluster(postalAddr)
if err != nil {
logger.Errorf("parse postal server addr error: %s", postalAddr, err)
continue
}
conn, err := s.clientFactory.GetConn(ctx, clusterAddr)
if err != nil {
logger.Error("get postal cluster conn error: ", err)
continue
}
clusterClient := postal.NewPostalClusterClient(conn)
req.Receivers = receivers
reqRedirect := &postal.ReqRedirect{
Ttl: 3, // TODO 转发n次就丢弃
RedirectMethod: methodRedirectDeliverBatch,
DeliverBatch: req,
}
_, err = clusterClient.Redirect(ctx, reqRedirect)
if err != nil {
logger.Error("redirect batch error: ", err)
}
}
}
return &postal.ResDeliver{Ok: true}, nil
}
// DeliverGroup postal broadcast
func (s *PostalServer) DeliverGroup(ctx context.Context, req *postal.ReqDeliverGroup) (res *postal.ResDeliver, err error) {
bytes, err := encodeDeliverMessage(req.Msg)
if err != nil {
return
}
res = &postal.ResDeliver{Ok: true}
g, ok := s.groupStore.Load(req.Gid)
if !ok {
return
}
g.Write(bytes)
return
}
// GroupJoin uid consistent hash
func (s *PostalServer) GroupJoin(ctx context.Context, req *postal.ReqGroupJoin) (emp *emptypb.Empty, err error) {
channel, ok := s.sessionStore.Load(req.Uid)
if !ok {
err = errors.New("uid not online")
return
}
for _, gid := range req.Gids {
g, _ := s.groupStore.ComputeIfAbsent(gid, func(gid string) *group.Group { return group.NewGroup(gid) })
g.Join(req.Uid, channel.Conn)
channel.GroupJoin(gid)
}
return
}
// GroupLeave uid consistent hash
func (s *PostalServer) GroupLeave(ctx context.Context, req *postal.ReqGroupLeave) (emp *emptypb.Empty, err error) {
channel, ok := s.sessionStore.Load(req.Uid)
if !ok {
err = errors.New("uid not online")
return
}
for _, gid := range req.Gids {
g, ok := s.groupStore.Load(gid)
if !ok {
continue
}
g.Leave(req.Uid)
channel.GroupLeave(gid)
}
return
}
// GroupDissolve postal broadcast
func (s *PostalServer) GroupDissolve(ctx context.Context, req *postal.ReqGroupDissolve) (emp *emptypb.Empty, err error) {
s.groupStore.Delete(req.Gid)
return
}
func (s *PostalServer) Endpoint(context.Context, *emptypb.Empty) (*postal.ResEndpoint, error) {
res := &postal.ResEndpoint{
Endpoint: s.endpointAddress,
}
return res, nil
}