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.
261 lines
6.8 KiB
261 lines
6.8 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/discovery" |
|
"sonet/pkg/grpc/interceptor" |
|
"sonet/pkg/plugins/mq" |
|
"sonet/pkg/protocol/event" |
|
"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 存储 |
|
postalClusterServer *PostalClusterServer |
|
producer mq.Producer |
|
} |
|
|
|
func NewPostalServer( |
|
endpointAddress string, |
|
sessionStore session.Store, |
|
groupStore *collect.ConcurrentMap[string, *group.Group], |
|
postalClusterServer *PostalClusterServer, |
|
producer mq.Producer) *PostalServer { |
|
return &PostalServer{ |
|
endpointAddress: endpointAddress, |
|
sessionStore: sessionStore, |
|
groupStore: groupStore, |
|
postalClusterServer: postalClusterServer, |
|
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.Infof("uid %s online", uid) |
|
// publish mq online event todo delay 5s |
|
online := &event.Online{Uid: uid} |
|
data, err := online.MarshalBinary() |
|
if err != nil { |
|
logger.Errorf("marshal %s online event error: %v", uid, err) |
|
continue |
|
} |
|
if err = s.producer.Publish(online.Topic(), data); err != nil { |
|
logger.Errorf("publish %s online event error: %v", uid, err) |
|
} |
|
|
|
case channel := <-s.sessionStore.OnDelete(): |
|
//logger.Infof("uid %s offline", channel.Uid) |
|
for _, gid := range channel.Groups() { |
|
if g, ok := s.groupStore.Load(gid); ok { |
|
g.Leave(channel.Uid) |
|
} |
|
} |
|
// publish mq offline event todo delay 5s |
|
offline := event.Offline{Uid: channel.Uid} |
|
data, err := offline.MarshalBinary() |
|
if err != nil { |
|
logger.Errorf("marshal %s offline event error: %v", channel.Uid, err) |
|
continue |
|
} |
|
if err = s.producer.Publish(offline.Topic(), data); err != nil { |
|
logger.Errorf("publish %s offline event error: %v", channel.Uid, err) |
|
} |
|
} |
|
} |
|
} |
|
|
|
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.Push(req.Msg) |
|
if err != nil { |
|
return |
|
} |
|
res = &postal.ResDeliver{Ok: true} |
|
return |
|
} |
|
// 向一致性 hash 下一个节点传递 |
|
// picker reqCluster offset +3: |
|
for _, offset := range []int32{3, -3} { |
|
res, err = s.postalClusterServer.Redeliver(ctx, &postal.ReqRedeliver{ |
|
Receivers: []string{req.Receiver}, |
|
Msg: req.Msg, |
|
Offset: offset, |
|
}) |
|
if err == nil { |
|
res = &postal.ResDeliver{Ok: true} |
|
return |
|
} |
|
} |
|
return |
|
} |
|
|
|
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 redeliverReceivers []string |
|
for _, receiverId := range req.Receivers { |
|
receiver, ok := s.sessionStore.Load(receiverId) |
|
if !ok { |
|
redeliverReceivers = append(redeliverReceivers, receiverId) |
|
continue |
|
} |
|
err = receiver.Push(req.Msg) |
|
if err != nil { |
|
logger.Errorf("deliver to %s error: ", receiver, err) |
|
} |
|
} |
|
|
|
res = &postal.ResDeliver{Ok: true} |
|
if len(redeliverReceivers) == 0 { |
|
return |
|
} |
|
|
|
// 向一致性 hash 下一个节点传递 |
|
for _, offset := range []int32{3, -3} { |
|
res, err = s.postalClusterServer.Redeliver(ctx, &postal.ReqRedeliver{ |
|
Receivers: redeliverReceivers, |
|
Msg: req.Msg, |
|
Offset: offset, |
|
}) |
|
if err == nil { |
|
return |
|
} |
|
} |
|
return |
|
} |
|
|
|
// 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(req.Msg) |
|
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) |
|
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 |
|
}
|
|
|