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
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 |
|
}
|
|
|