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, postalPicker *PostalPicker, groupLoader GroupLoader, consumer mq.Consumer, ) *GroupDeliver { return &GroupDeliver{ svcName: msgInServiceName, postalPicker: postalPicker, groupLoader: groupLoader, consumer: consumer, } } func (d *GroupDeliver) Init(ctx context.Context) (err error) { if d.consumer != nil { 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 }