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.
 
 

140 lines
3.3 KiB

package deliver
import (
"context"
"errors"
"fmt"
"google.golang.org/protobuf/proto"
"sonet/api/gen/postal"
"sonet/pkg/plugins/mq"
"sonet/pkg/protocol/event"
"sonet/pkg/utils/collect"
"sonet/pkg/utils/logger"
)
type GroupLoader interface {
Load(uid string) (groupIds []string, err error)
}
type GroupDeliver struct {
svcName string
groupLoader GroupLoader
consumer mq.Consumer
postalPicker *PostalPicker
}
func NewGroupDeliver(
msgInServiceName string,
groupLoader GroupLoader,
consumer mq.Consumer,
) *GroupDeliver {
return &GroupDeliver{
svcName: msgInServiceName,
groupLoader: groupLoader,
consumer: consumer,
}
}
func (d *GroupDeliver) Init(ctx context.Context) (err error) {
err = d.subscribe()
return
}
// toPostalGid 加上 svc name 前缀避免和其他服务群组冲突
func (d *GroupDeliver) toPostalGid(gid string) string {
return d.svcName + "." + gid
}
// subscribe postal online events
func (d *GroupDeliver) subscribe() error {
consumerChannel := fmt.Sprintf("%s:%s", "deliver", d.svcName)
return d.consumer.Subscribe(event.TopicOnline, consumerChannel, func(message *mq.Message) (err error) {
online := &event.Online{}
if e := online.UnmarshalBinary(message.Body); e != nil {
logger.Error("deliver unmarshal online event payload error: ", e)
return
}
// load uid groups join to postal
groupIds, err := d.groupLoader.Load(online.Uid)
if err != nil {
logger.Error("deliver load groups error:", err)
return
}
if err = d.GroupJoin(context.Background(), online.Uid, groupIds); err != nil {
logger.Errorf("deliver uid %s group join error: %v", online.Uid, err)
return
}
return
})
}
// DeliverGroup 群组消息广播投递到所有 postal 节点
func (d *GroupDeliver) DeliverGroup(ctx context.Context, gid string, msg proto.Message, options ...Option) (err error) {
message, err := Proto2DeliverMessage(d.svcName, msg)
if err != nil {
return
}
clients := d.postalPicker.PickAll()
var errs []error
req := postal.ReqDeliverGroup{Gid: d.toPostalGid(gid), Msg: message}
for _, client := range clients {
_, err := client.DeliverGroup(ctx, &req)
if err != nil {
errs = append(errs, err)
}
}
if len(errs) > 0 {
err = errors.Join(errs...)
return
}
return
}
// GroupDissolve 群组消息广播投递到所有 postal 节点
func (d *GroupDeliver) GroupDissolve(ctx context.Context, gid string) (err error) {
clients := d.postalPicker.PickAll()
var errs []error
req := postal.ReqGroupDissolve{Gid: d.toPostalGid(gid)}
for _, client := range clients {
_, err := client.GroupDissolve(ctx, &req)
if err != nil {
errs = append(errs, err)
}
}
if len(errs) > 0 {
err = errors.Join(errs...)
return
}
return
}
func (d *GroupDeliver) GroupJoin(ctx context.Context, uid string, gids []string) (err error) {
if len(gids) == 0 {
return
}
gids = collect.Mapping(gids, d.toPostalGid)
client, err := d.postalPicker.Pick(uid)
if err != nil {
return
}
_, err = client.GroupJoin(ctx, &postal.ReqGroupJoin{Uid: uid, Gids: gids})
return
}
func (d *GroupDeliver) GroupLeave(ctx context.Context, uid string, gids []string) (err error) {
if len(gids) == 0 {
return
}
gids = collect.Mapping(gids, d.toPostalGid)
client, err := d.postalPicker.Pick(uid)
if err != nil {
return
}
req := &postal.ReqGroupLeave{Uid: uid, Gids: gids}
_, err = client.GroupLeave(ctx, req)
return
}