diff --git a/cmd/wire_gen.go b/cmd/wire_gen.go index 5fe906a35..6eb17eaa8 100644 --- a/cmd/wire_gen.go +++ b/cmd/wire_gen.go @@ -77,21 +77,21 @@ func newServer(hub2 *hub.Hub, db *gorm.DB, repo repository.Repository, fs storag viewerManager := viewer.NewManager(hub2) wsStreamer := ws2.NewStreamer(hub2, viewerManager, webrtcv3Manager, logger) serverOriginString := provideServerOriginString(c2) - notificationService := notification.NewService(repo, manager, messageManager, fileManager, hub2, logger, client, wsStreamer, viewerManager, serverOriginString) - ogpService, err := ogp.NewServiceImpl(repo, logger) + esEngineConfig := provideESEngineConfig(c2) + engine, err := initSearchServiceIfAvailable(messageManager, manager, repo, logger, esEngineConfig) if err != nil { return nil, err } - rbacRBAC, err := rbac.New(repo) + notificationService := notification.NewService(repo, manager, messageManager, fileManager, hub2, logger, client, wsStreamer, viewerManager, serverOriginString, engine) + ogpService, err := ogp.NewServiceImpl(repo, logger) if err != nil { return nil, err } - oidcService := provideOIDCService(c2, repo, rbacRBAC) - esEngineConfig := provideESEngineConfig(c2) - engine, err := initSearchServiceIfAvailable(messageManager, manager, repo, logger, esEngineConfig) + rbacRBAC, err := rbac.New(repo) if err != nil { return nil, err } + oidcService := provideOIDCService(c2, repo, rbacRBAC) roomStateManager := provideQallRoomStateManager(c2, hub2) soundboard, err := provideQallSoundboard(repo, fs, logger, hub2) if err != nil { diff --git a/service/notification/handlers.go b/service/notification/handlers.go index 136ad4c3d..4eda3e8bc 100644 --- a/service/notification/handlers.go +++ b/service/notification/handlers.go @@ -17,6 +17,7 @@ import ( "github.com/traPtitech/traQ/repository" "github.com/traPtitech/traQ/service/fcm" "github.com/traPtitech/traQ/service/qall" + "github.com/traPtitech/traQ/service/search" "github.com/traPtitech/traQ/service/viewer" "github.com/traPtitech/traQ/service/ws" "github.com/traPtitech/traQ/utils/message" @@ -284,9 +285,14 @@ func messageCreatedHandler(ns *Service, ev hub.Message) { func messageUpdatedHandler(ns *Service, ev hub.Message) { cid := ev.Fields["message"].(*model.Message).ChannelID + mid := ev.Fields["message_id"].(uuid.UUID) + citedChannels, err := ns.cache.Get(context.Background(), mid) + if err != nil { + return + } wsEventType := "MESSAGE_UPDATED" wsPayload := map[string]interface{}{ - "id": ev.Fields["message_id"].(uuid.UUID), + "id": mid, } var targetFunc ws.TargetFunc @@ -294,6 +300,7 @@ func messageUpdatedHandler(ns *Service, ev hub.Message) { // 公開チャンネル targetFunc = ws.Or( ws.TargetChannelViewers(cid), + ws.TargetChannelsViewers(citedChannels), ws.TargetTimelineStreamingEnabled(), ) } else { @@ -306,9 +313,12 @@ func messageUpdatedHandler(ns *Service, ev hub.Message) { func messageDeletedHandler(ns *Service, ev hub.Message) { cid := ev.Fields["message"].(*model.Message).ChannelID + mid := ev.Fields["message_id"].(uuid.UUID) + citedChannels := ns.getCitedChannelIDs(context.Background(), mid) + wsEventType := "MESSAGE_DELETED" wsPayload := map[string]interface{}{ - "id": ev.Fields["message_id"].(uuid.UUID), + "id": mid, } var targetFunc ws.TargetFunc @@ -316,6 +326,7 @@ func messageDeletedHandler(ns *Service, ev hub.Message) { // 公開チャンネル targetFunc = ws.Or( ws.TargetChannelViewers(cid), + ws.TargetChannelsViewers(citedChannels), ws.TargetTimelineStreamingEnabled(), ) } else { @@ -743,3 +754,18 @@ func broadcast(ns *Service, wsEventType string, wsPayload interface{}) { func userMulticast(ns *Service, userID uuid.UUID, wsEventType string, wsPayload interface{}) { go ns.ws.WriteMessage(wsEventType, wsPayload, ws.TargetUsers(userID)) } + +func (ns *Service) getCitedChannelIDs(_ context.Context, messageID uuid.UUID) []uuid.UUID { + query := search.Query{} + query.Citation = optional.From(messageID) + + res, err := ns.search.Do(&query) + if err != nil { + return nil + } + result := []uuid.UUID{} + for _, hit := range res.Hits() { + result = append(result, hit.GetChannelID()) + } + return result +} diff --git a/service/notification/service.go b/service/notification/service.go index a23cdaab0..e207e5231 100644 --- a/service/notification/service.go +++ b/service/notification/service.go @@ -1,7 +1,12 @@ package notification import ( + "context" + "time" + + "github.com/gofrs/uuid" "github.com/leandro-lugaresi/hub" + "github.com/motoki317/sc" "go.uber.org/zap" "github.com/traPtitech/traQ/repository" @@ -9,6 +14,7 @@ import ( "github.com/traPtitech/traQ/service/fcm" "github.com/traPtitech/traQ/service/file" "github.com/traPtitech/traQ/service/message" + "github.com/traPtitech/traQ/service/search" "github.com/traPtitech/traQ/service/variable" "github.com/traPtitech/traQ/service/viewer" "github.com/traPtitech/traQ/service/ws" @@ -26,10 +32,12 @@ type Service struct { ws *ws.Streamer vm *viewer.Manager origin string + search search.Engine + cache *sc.Cache[uuid.UUID, []uuid.UUID] } // NewService 通知サービスを作成して起動します -func NewService(repo repository.Repository, cm channel.Manager, mm message.Manager, fm file.Manager, hub *hub.Hub, logger *zap.Logger, fcm fcm.Client, ws *ws.Streamer, vm *viewer.Manager, origin variable.ServerOriginString) *Service { +func NewService(repo repository.Repository, cm channel.Manager, mm message.Manager, fm file.Manager, hub *hub.Hub, logger *zap.Logger, fcm fcm.Client, ws *ws.Streamer, vm *viewer.Manager, origin variable.ServerOriginString, search search.Engine) *Service { service := &Service{ repo: repo, cm: cm, @@ -41,7 +49,16 @@ func NewService(repo repository.Repository, cm channel.Manager, mm message.Manag ws: ws, vm: vm, origin: string(origin), + search: search, } + service.cache = sc.NewMust( + func(_ context.Context, messageID uuid.UUID) ([]uuid.UUID, error) { + return service.getCitedChannelIDs(context.Background(), messageID), nil + }, + time.Minute, + time.Minute, + sc.WithLRUBackend(1000), + ) go func() { topics := make([]string, 0, len(handlerMap)) for k := range handlerMap { diff --git a/service/ws/target_func.go b/service/ws/target_func.go index 07e257e9f..fb0903712 100644 --- a/service/ws/target_func.go +++ b/service/ws/target_func.go @@ -48,6 +48,19 @@ func TargetChannelViewers(channelID uuid.UUID) TargetFunc { } } +// TargetChannelViewersの複数形 +func TargetChannelsViewers(channelIDs []uuid.UUID) TargetFunc { + return func(s Session) bool { + c, _ := s.ViewState() + for _, id := range channelIDs { + if c == id { + return true + } + } + return false + } +} + // TargetTimelineStreamingEnabled タイムラインストリーミングが有効なコネクションを対象に送信します func TargetTimelineStreamingEnabled() TargetFunc { return func(s Session) bool {