diff --git a/backend/plugins/domain/msg_gateway/consts/bot.go b/backend/plugins/domain/msg_gateway/consts/bot.go new file mode 100644 index 00000000..9c2fb67c --- /dev/null +++ b/backend/plugins/domain/msg_gateway/consts/bot.go @@ -0,0 +1,27 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package consts defines constants, error types, and identifiers for msg_gateway. +package consts + +// Bot Channel type and scope constants. +const ( + ChannelTypeTelegram = "telegram" + ChannelTypeQQ = "qq" + MessageChannelTypeTelegram = "telegram" + MessageChannelTypeQQ = "qq" + MessageOwnerScopeSystem = "system" +) + +// Bot Task and Schedule identifier constants. +const ( + TaskCleanupPairingCodes = "msg_gateway:cleanup_pairing_codes" + TaskDispatchBotMsg = "msg_gateway:dispatch_bot_msg" + TaskTypeDispatchBotMsg = "dispatch_bot_msg" +) + +// Pairing code constants. +const ( + CodeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789" + CodeLength = 8 +) diff --git a/backend/plugins/domain/msg_gateway/consts/consts.go b/backend/plugins/domain/msg_gateway/consts/errs.go similarity index 59% rename from backend/plugins/domain/msg_gateway/consts/consts.go rename to backend/plugins/domain/msg_gateway/consts/errs.go index a6bc3b23..e8c32c9b 100644 --- a/backend/plugins/domain/msg_gateway/consts/consts.go +++ b/backend/plugins/domain/msg_gateway/consts/errs.go @@ -1,69 +1,10 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -// Package consts defines constants, sentinel errors, and user-facing error messages -// for the msg_gateway plugin. package consts import "errors" -// Channel type and scope constants. -const ( - ChannelTypeTelegram = "telegram" - ChannelTypeQQ = "qq" - MessageChannelTypeTelegram = "telegram" - MessageChannelTypeQQ = "qq" - MessageOwnerScopeSystem = "system" - - TypeCustom = "custom" - TypeEmail = "email" - TypeTelegram = "telegram" - ChannelCustom = "custom" - ChannelEmail = "email" - ChannelLark = "lark" - ChannelDingTalk = "dingtalk" - ChannelTelegram = "telegram" - ChannelBark = "bark" - ChannelDiscord = "discord" - ChannelSlack = "slack" - ChannelPushover = "pushover" - - DefaultLevelInfo = "INFO" - KeyTitle = "title" - KeyContent = "content" - KeyLevel = "level" - - // KeyURL represents the URL field key. - KeyURL = "url" - // KeyToken represents the Token field key. - KeyToken = "token" - // KeyOther represents the Other field key. - KeyOther = "other" - - // TypeText represents standard text input type. - TypeText = "text" - // TypePassword represents password input type. - TypePassword = "password" - // TypeTextarea represents textarea input type. - TypeTextarea = "textarea" -) - -// Task and Schedule identifier constants. -const ( - TaskPushNotification = "msg_gateway:push_notification" - TaskCleanupPairingCodes = "msg_gateway:cleanup_pairing_codes" - TaskDispatchBotMsg = "msg_gateway:dispatch_bot_msg" - TaskTypeDispatchBotMsg = "dispatch_bot_msg" - SendNotificationTask = "push:send" - TaskTypeSendNotification = "send_notification" -) - -// Pairing code constants. -const ( - CodeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789" - CodeLength = 8 -) - // Sentinel errors. var ( ErrCodeInvalid = errors.New("invalid or expired pairing code") @@ -73,14 +14,15 @@ var ( ErrBindingForbidden = errors.New("cannot unbind another user's binding") ErrChannelIDRequired = errors.New("channel_id is required") ErrChannelDisabled = errors.New("channel is not enabled") + ErrChannelNotFound = errors.New("channel not found") + ErrEventNotFound = errors.New("notification event not found") + ErrUserNotFound = errors.New("user not found") + ErrNoAdminUser = errors.New("no admin user found") + ErrTaskServiceNotAvail = errors.New("task service not available") // ErrRecordNotFound maps GORM's missing-row sentinel at the DAO boundary so // upper layers never import gorm. Its text matches gorm.ErrRecordNotFound verbatim. ErrRecordNotFound = errors.New("record not found") - - // ErrUnsupportedUserLookupField rejects a column name that the DAO is not - // allowed to interpolate into a WHERE clause. - ErrUnsupportedUserLookupField = errors.New("unsupported user lookup field") ) // User-facing validation and error message constants. @@ -89,7 +31,7 @@ const ( ErrTypeInvalid = "type must be telegram or qq" ErrTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text ErrQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text - ErrChannelNotFound = "channel not found" + ErrChannelNotFoundText = "channel not found" ErrChannelProbeFailed = "channel probe failed" ErrBotDispatchTextRequired = "message text is required" ErrBotChannelNotRegistered = "channel adapter is not registered" @@ -99,7 +41,6 @@ const ( ErrInvalidBindingID = "invalid binding id" ErrInvalidChannelID = "invalid channel id" ErrInvalidEventID = "invalid event id" - ErrEventNotFound = "notification event not found" ErrValidationFailed = "validation failed" ErrMissingTelegramToken = "missing telegram bot token" @@ -120,8 +61,6 @@ const ( ErrEventKeyOrTaskType = "either event_key or task_type must be provided" ErrUnsupportedEventKey = "unsupported built-in event key" ErrTaskServiceUnavailable = "task service not available" - ErrUserNotFound = "user not found" - ErrNoAdminUser = "no admin user found" ErrPayloadRequired = "payload is required" ErrInvalidJSONFormat = "invalid json format" diff --git a/backend/plugins/domain/msg_gateway/consts/push.go b/backend/plugins/domain/msg_gateway/consts/push.go new file mode 100644 index 00000000..806950e9 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/consts/push.go @@ -0,0 +1,43 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package consts + +// Push notification channel type constants. +const ( + TypeCustom = "custom" + TypeEmail = "email" + TypeTelegram = "telegram" + ChannelCustom = "custom" + ChannelEmail = "email" + ChannelLark = "lark" + ChannelDingTalk = "dingtalk" + ChannelTelegram = "telegram" + ChannelBark = "bark" + ChannelDiscord = "discord" + ChannelSlack = "slack" + ChannelPushover = "pushover" +) + +// Push message template and payload keys. +const ( + DefaultLevelInfo = "INFO" + KeyTitle = "title" + KeyContent = "content" + KeyLevel = "level" + + KeyURL = "url" + KeyToken = "token" + KeyOther = "other" + + TypeText = "text" + TypePassword = "password" + TypeTextarea = "textarea" +) + +// Push task identifier constants. +const ( + TaskPushNotification = "msg_gateway:push_notification" + SendNotificationTask = "push:send" + TaskTypeSendNotification = "send_notification" +) diff --git a/backend/plugins/domain/msg_gateway/controller/admin.go b/backend/plugins/domain/msg_gateway/controller/admin.go index 87b76d1b..98488d48 100644 --- a/backend/plugins/domain/msg_gateway/controller/admin.go +++ b/backend/plugins/domain/msg_gateway/controller/admin.go @@ -7,8 +7,8 @@ import ( "Wavelet/pkg/response" "Wavelet/plugins/domain/msg_gateway/consts" "Wavelet/plugins/domain/msg_gateway/service" + "errors" "net/http" - "strconv" "github.com/gin-gonic/gin" ) @@ -43,17 +43,12 @@ func ListAdminChannels(c *gin.Context) { } func parseAdminChannelID(c *gin.Context) (uint64, bool) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, consts.ErrInvalidChannelID) - return 0, false - } - return id, true + return parseUint64Param(c, "id", consts.ErrInvalidChannelID) } func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) { - if err.Error() == consts.ErrChannelNotFound { - response.AbortNotFound(c, err.Error()) + if errors.Is(err, consts.ErrChannelNotFound) || err.Error() == consts.ErrChannelNotFoundText { + response.AbortNotFound(c, consts.ErrChannelNotFoundText) return } fallback(c, err.Error()) diff --git a/backend/plugins/domain/msg_gateway/controller/base.go b/backend/plugins/domain/msg_gateway/controller/base.go new file mode 100644 index 00000000..f255d3d2 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/controller/base.go @@ -0,0 +1,69 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package controller + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/ginutil" + "Wavelet/pkg/response" + "context" + "net/http" + "strconv" + + "github.com/gin-gonic/gin" +) + +// currentUser extracts the authenticated UserDTO from gin.Context. +func currentUser(c *gin.Context) (*contracts.UserDTO, bool) { + return ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) +} + +// parseUint64Param parses a uint64 URL path parameter. +func parseUint64Param(c *gin.Context, paramName, errInvalid string) (uint64, bool) { + id, err := strconv.ParseUint(c.Param(paramName), 10, 64) + if err != nil { + response.AbortBadRequest(c, errInvalid) + return 0, false + } + return id, true +} + +// handleJSONRequest binds a JSON body, executes the service handler, and writes the standard success envelope. +func handleJSONRequest[Req any, Res any](c *gin.Context, handler func(ctx context.Context, req Req) (Res, error)) { + var req Req + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + res, err := handler(c.Request.Context(), req) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + c.JSON(http.StatusOK, response.OK(res)) +} + +// handleEntityUpdate resolves a path identifier and JSON body, executes the updater, and handles errors with onErr. +func handleEntityUpdate[Req any, Res any]( + c *gin.Context, + parseID func(*gin.Context) (uint64, bool), + updater func(ctx context.Context, id uint64, req Req) (Res, error), + onErr func(*gin.Context, error), +) { + id, ok := parseID(c) + if !ok { + return + } + var req Req + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + dto, err := updater(c.Request.Context(), id, req) + if err != nil { + onErr(c, err) + return + } + c.JSON(http.StatusOK, response.OK(dto)) +} diff --git a/backend/plugins/domain/msg_gateway/controller/push_channel.go b/backend/plugins/domain/msg_gateway/controller/push_channel.go index 8919560b..89d440a0 100644 --- a/backend/plugins/domain/msg_gateway/controller/push_channel.go +++ b/backend/plugins/domain/msg_gateway/controller/push_channel.go @@ -10,7 +10,6 @@ import ( "Wavelet/plugins/domain/msg_gateway/service" "errors" "net/http" - "strconv" "github.com/gin-gonic/gin" ) @@ -46,18 +45,13 @@ func ListPushChannels(c *gin.Context) { // parsePushChannelID reads the path identifier of a push channel. func parsePushChannelID(c *gin.Context) (uint64, bool) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, consts.ErrInvalidChannelID) - return 0, false - } - return id, true + return parseUint64Param(c, "id", consts.ErrInvalidChannelID) } // handlePushChannelNotFoundError maps a missing channel row to 404, others to fallback. func handlePushChannelNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) { - if errors.Is(err, consts.ErrRecordNotFound) { - response.AbortNotFound(c, consts.ErrChannelNotFound) + if errors.Is(err, consts.ErrRecordNotFound) || errors.Is(err, consts.ErrChannelNotFound) || err.Error() == consts.ErrChannelNotFoundText { + response.AbortNotFound(c, consts.ErrChannelNotFoundText) return } fallback(c, err.Error()) diff --git a/backend/plugins/domain/msg_gateway/controller/push_event.go b/backend/plugins/domain/msg_gateway/controller/push_event.go index 05b3fd40..616eb5d8 100644 --- a/backend/plugins/domain/msg_gateway/controller/push_event.go +++ b/backend/plugins/domain/msg_gateway/controller/push_event.go @@ -47,18 +47,13 @@ func ListBuiltInPushEvents(c *gin.Context) { // parsePushEventID reads the path identifier of a push event. func parsePushEventID(c *gin.Context) (uint64, bool) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, consts.ErrInvalidEventID) - return 0, false - } - return id, true + return parseUint64Param(c, "id", consts.ErrInvalidEventID) } // handlePushEventNotFoundError maps a missing event row to 404, others to fallback. func handlePushEventNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) { - if errors.Is(err, consts.ErrRecordNotFound) { - response.AbortNotFound(c, consts.ErrEventNotFound) + if errors.Is(err, consts.ErrRecordNotFound) || errors.Is(err, consts.ErrEventNotFound) || err.Error() == consts.ErrEventNotFound.Error() { + response.AbortNotFound(c, consts.ErrEventNotFound.Error()) return } fallback(c, err.Error()) diff --git a/backend/plugins/domain/msg_gateway/controller/user.go b/backend/plugins/domain/msg_gateway/controller/user.go index 2d99a26e..6323f7d7 100644 --- a/backend/plugins/domain/msg_gateway/controller/user.go +++ b/backend/plugins/domain/msg_gateway/controller/user.go @@ -4,65 +4,16 @@ package controller import ( - "Wavelet/core/contracts" - "Wavelet/pkg/ginutil" "Wavelet/pkg/response" "Wavelet/plugins/domain/msg_gateway/consts" "Wavelet/plugins/domain/msg_gateway/model/do" "Wavelet/plugins/domain/msg_gateway/service" - "context" "errors" "net/http" - "strconv" "github.com/gin-gonic/gin" ) -func currentUser(c *gin.Context) (*contracts.UserDTO, bool) { - return ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) -} - -// handleJSONRequest binds a JSON body, runs the service use case and writes the -// standard success envelope; any service error surfaces as a bad request. -func handleJSONRequest[Req any, Res any](c *gin.Context, handler func(ctx context.Context, req Req) (Res, error)) { - var req Req - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - res, err := handler(c.Request.Context(), req) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(res)) -} - -// handleEntityUpdate resolves a path identifier plus JSON body, runs the updater -// use case and writes the success envelope; error classification is delegated to onErr. -func handleEntityUpdate[Req any, Res any]( - c *gin.Context, - parseID func(*gin.Context) (uint64, bool), - updater func(ctx context.Context, id uint64, req Req) (Res, error), - onErr func(*gin.Context, error), -) { - id, ok := parseID(c) - if !ok { - return - } - var req Req - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - dto, err := updater(c.Request.Context(), id, req) - if err != nil { - onErr(c, err) - return - } - c.JSON(http.StatusOK, response.OK(dto)) -} - // ListChannels lists enabled channels a user can bind. // @Summary List enabled messaging channels // @Description Returns enabled system bots the current user can pair with @@ -160,9 +111,8 @@ func UnbindBinding(c *gin.Context) { response.AbortUnauthorized(c, consts.ErrLoginRequired) return } - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, consts.ErrInvalidBindingID) + id, ok := parseUint64Param(c, "id", consts.ErrInvalidBindingID) + if !ok { return } if err := service.UnbindChannel(c.Request.Context(), user.ID, id); err != nil { diff --git a/backend/plugins/domain/msg_gateway/dao/bot.go b/backend/plugins/domain/msg_gateway/dao/bot.go new file mode 100644 index 00000000..630208c4 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/dao/bot.go @@ -0,0 +1,160 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package dao + +import ( + "Wavelet/pkg/idgen" + "Wavelet/plugins/domain/msg_gateway/model/entity" + "context" + "errors" + "time" + + "gorm.io/gorm" +) + +// CreateMessageChannel inserts a channel row. +func CreateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error { + if ch.ID == 0 { + ch.ID = idgen.NextUint64ID() + } + return GetDB(ctx).Create(ch).Error +} + +// UpdateMessageChannel saves a channel row. +func UpdateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error { + return GetDB(ctx).Save(ch).Error +} + +// GetMessageChannel loads a channel by id. +func GetMessageChannel(ctx context.Context, id uint64) (*entity.MessageChannel, error) { + var ch entity.MessageChannel + if err := GetDB(ctx).Where("id = ?", id).First(&ch).Error; err != nil { + return nil, mapNotFound(err) + } + return &ch, nil +} + +// ListMessageChannels returns all channels newest first. +func ListMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) { + var rows []entity.MessageChannel + if err := GetDB(ctx).Order("id DESC").Find(&rows).Error; err != nil { + return nil, err + } + return rows, nil +} + +// DeleteMessageChannel removes pairings, bindings, then the channel. +func DeleteMessageChannel(ctx context.Context, id uint64) error { + return GetDB(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx.Where("channel_id = ?", id).Delete(&entity.MessagePairingCode{}).Error; err != nil { + return err + } + if err := tx.Where("channel_id = ?", id).Delete(&entity.MessageBinding{}).Error; err != nil { + return err + } + return tx.Delete(&entity.MessageChannel{}, id).Error + }) +} + +// CreateMessageBinding inserts a binding. +func CreateMessageBinding(ctx context.Context, b *entity.MessageBinding) error { + if b.ID == 0 { + b.ID = idgen.NextUint64ID() + } + return GetDB(ctx).Create(b).Error +} + +// GetBindingByChannelPlatform finds a binding for a platform user on a channel. +func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*entity.MessageBinding, error) { + var b entity.MessageBinding + err := GetDB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error + if err != nil { + return nil, mapNotFound(err) + } + return &b, nil +} + +// ListBindingsByUser lists bindings for a Wavelet user. +func ListBindingsByUser(ctx context.Context, userID uint64) ([]entity.MessageBinding, error) { + var rows []entity.MessageBinding + if err := GetDB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil { + return nil, err + } + return rows, nil +} + +// ListBindingsByChannel lists bindings on one messaging channel. +func ListBindingsByChannel(ctx context.Context, channelID uint64) ([]entity.MessageBinding, error) { + var rows []entity.MessageBinding + if err := GetDB(ctx).Where("channel_id = ?", channelID).Order("id DESC").Find(&rows).Error; err != nil { + return nil, err + } + return rows, nil +} + +// GetMessageBinding loads a binding by id. +func GetMessageBinding(ctx context.Context, id uint64) (*entity.MessageBinding, error) { + var b entity.MessageBinding + if err := GetDB(ctx).Where("id = ?", id).First(&b).Error; err != nil { + return nil, mapNotFound(err) + } + return &b, nil +} + +// DeleteMessageBinding deletes a binding by id. +func DeleteMessageBinding(ctx context.Context, id uint64) error { + return GetDB(ctx).Delete(&entity.MessageBinding{}, id).Error +} + +// UpsertPairingCode reuses an unexpired code for the same channel+platform user. +func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*entity.MessagePairingCode, error) { + var existing entity.MessagePairingCode + err := GetDB(ctx). + Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()). + First(&existing).Error + if err == nil { + return &existing, nil + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return nil, err + } + row := &entity.MessagePairingCode{ + Code: code, + ChannelID: channelID, + PlatformUserID: platformUserID, + ExpiresAt: expiresAt, + } + if err := GetDB(ctx).Create(row).Error; err != nil { + return nil, err + } + return row, nil +} + +// GetPairingCode loads a pairing code by normalized code string. +func GetPairingCode(ctx context.Context, code string) (*entity.MessagePairingCode, error) { + var row entity.MessagePairingCode + if err := GetDB(ctx).Where("code = ?", code).First(&row).Error; err != nil { + return nil, mapNotFound(err) + } + return &row, nil +} + +// DeletePairingCode removes a pairing code. +func DeletePairingCode(ctx context.Context, code string) error { + return GetDB(ctx).Where("code = ?", code).Delete(&entity.MessagePairingCode{}).Error +} + +// DeleteExpiredPairingCodes removes expired pairing rows. +func DeleteExpiredPairingCodes(ctx context.Context) error { + return GetDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&entity.MessagePairingCode{}).Error +} + +// ListEnabledMessageChannels returns enabled channels. +func ListEnabledMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) { + var rows []entity.MessageChannel + if err := GetDB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil { + return nil, err + } + return rows, nil +} diff --git a/backend/plugins/domain/msg_gateway/dao/bot_test.go b/backend/plugins/domain/msg_gateway/dao/bot_test.go new file mode 100644 index 00000000..aa66b500 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/dao/bot_test.go @@ -0,0 +1,64 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package dao_test + +import ( + "Wavelet/pkg/idgen" + "Wavelet/pkg/testhelper" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/entity" + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestBotDAO_ChannelAndBinding(t *testing.T) { + _ = idgen.Init(1) + db, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + require.NoError(t, db.AutoMigrate(&entity.MessageChannel{}, &entity.MessageBinding{}, &entity.MessagePairingCode{})) + + dao.SetDBServiceForTest(stubDBService{db: db}) + t.Cleanup(func() { dao.SetDBServiceForTest(nil) }) + + ctx := context.Background() + + ch := entity.MessageChannel{ + Name: "tg_bot", + Type: "telegram", + OwnerScope: "system", + Credentials: "encrypted_token", + Enabled: true, + } + require.NoError(t, dao.CreateMessageChannel(ctx, &ch)) + assert.NotZero(t, ch.ID) + + code, err := dao.UpsertPairingCode(ctx, ch.ID, "tg_user_1", "ABCD1234", time.Now().Add(10*time.Minute)) + require.NoError(t, err) + assert.Equal(t, "ABCD1234", code.Code) + + // Reusing pairing code + code2, err := dao.UpsertPairingCode(ctx, ch.ID, "tg_user_1", "XYZ9999", time.Now().Add(10*time.Minute)) + require.NoError(t, err) + assert.Equal(t, "ABCD1234", code2.Code) + + binding := entity.MessageBinding{ + UserID: 42, + ChannelID: ch.ID, + PlatformUserID: "tg_user_1", + } + require.NoError(t, dao.CreateMessageBinding(ctx, &binding)) + assert.NotZero(t, binding.ID) + + bindings, err := dao.ListBindingsByUser(ctx, 42) + require.NoError(t, err) + assert.Len(t, bindings, 1) + + require.NoError(t, dao.DeleteMessageChannel(ctx, ch.ID)) + _, err = dao.GetMessageChannel(ctx, ch.ID) + assert.Error(t, err) +} diff --git a/backend/plugins/domain/msg_gateway/dao/dao.go b/backend/plugins/domain/msg_gateway/dao/dao.go index 977cc583..f48ac18a 100644 --- a/backend/plugins/domain/msg_gateway/dao/dao.go +++ b/backend/plugins/domain/msg_gateway/dao/dao.go @@ -7,27 +7,21 @@ package dao import ( "Wavelet/core" "Wavelet/core/contracts" - "Wavelet/pkg/idgen" "Wavelet/plugins/domain/msg_gateway/consts" - "Wavelet/plugins/domain/msg_gateway/model/entity" "context" "errors" "sync" - "time" "gorm.io/gorm" ) var ( - dbMu sync.RWMutex - dbSvc contracts.DBService + dbMu sync.RWMutex + dbSvc contracts.DBService + cacheMu sync.RWMutex + cacheSvc contracts.CacheService ) -// SetDBServiceForTest injects a DBService for tests. Production wiring must use Apply. -func SetDBServiceForTest(s contracts.DBService) { - SetDBService(s) -} - // SetDBService sets the database service singleton. func SetDBService(s contracts.DBService) { dbMu.Lock() @@ -35,8 +29,19 @@ func SetDBService(s contracts.DBService) { dbSvc = s } -// GetDB resolves the persistence handle for the current call, preferring an -// explicitly injected *core.Context before falling back to the plugin singleton. +// SetDBServiceForTest injects a DBService for tests. +func SetDBServiceForTest(s contracts.DBService) { + SetDBService(s) +} + +// SetCacheService sets the cache service singleton. +func SetCacheService(s contracts.CacheService) { + cacheMu.Lock() + defer cacheMu.Unlock() + cacheSvc = s +} + +// GetDB resolves the persistence handle for the current call. func GetDB(ctx context.Context) *gorm.DB { if c, ok := ctx.(*core.Context); ok && c != nil { if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { @@ -52,6 +57,19 @@ func GetDB(ctx context.Context) *gorm.DB { return nil } +// GetCache resolves the cache service for the current call. +func GetCache(ctx context.Context) contracts.CacheService { + if c, ok := ctx.(*core.Context); ok && c != nil { + if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil { + return s + } + } + cacheMu.RLock() + s := cacheSvc + cacheMu.RUnlock() + return s +} + // mapNotFound translates GORM's missing-row sentinel into the plugin-level // consts.ErrRecordNotFound so the service and controller layers stay free of gorm imports. func mapNotFound(err error) error { @@ -60,149 +78,3 @@ func mapNotFound(err error) error { } return err } - -// CreateMessageChannel inserts a channel row. -func CreateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error { - if ch.ID == 0 { - ch.ID = idgen.NextUint64ID() - } - return GetDB(ctx).Create(ch).Error -} - -// UpdateMessageChannel saves a channel row. -func UpdateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error { - return GetDB(ctx).Save(ch).Error -} - -// GetMessageChannel loads a channel by id. -func GetMessageChannel(ctx context.Context, id uint64) (*entity.MessageChannel, error) { - var ch entity.MessageChannel - if err := GetDB(ctx).Where("id = ?", id).First(&ch).Error; err != nil { - return nil, mapNotFound(err) - } - return &ch, nil -} - -// ListMessageChannels returns all channels newest first. -func ListMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) { - var rows []entity.MessageChannel - if err := GetDB(ctx).Order("id DESC").Find(&rows).Error; err != nil { - return nil, err - } - return rows, nil -} - -// DeleteMessageChannel removes pairings, bindings, then the channel. -func DeleteMessageChannel(ctx context.Context, id uint64) error { - return GetDB(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Where("channel_id = ?", id).Delete(&entity.MessagePairingCode{}).Error; err != nil { - return err - } - if err := tx.Where("channel_id = ?", id).Delete(&entity.MessageBinding{}).Error; err != nil { - return err - } - return tx.Delete(&entity.MessageChannel{}, id).Error - }) -} - -// CreateMessageBinding inserts a binding. -func CreateMessageBinding(ctx context.Context, b *entity.MessageBinding) error { - if b.ID == 0 { - b.ID = idgen.NextUint64ID() - } - return GetDB(ctx).Create(b).Error -} - -// GetBindingByChannelPlatform finds a binding for a platform user on a channel. -func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*entity.MessageBinding, error) { - var b entity.MessageBinding - err := GetDB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error - if err != nil { - return nil, mapNotFound(err) - } - return &b, nil -} - -// ListBindingsByUser lists bindings for a Wavelet user. -func ListBindingsByUser(ctx context.Context, userID uint64) ([]entity.MessageBinding, error) { - var rows []entity.MessageBinding - if err := GetDB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil { - return nil, err - } - return rows, nil -} - -// ListBindingsByChannel lists bindings on one messaging channel. -func ListBindingsByChannel(ctx context.Context, channelID uint64) ([]entity.MessageBinding, error) { - var rows []entity.MessageBinding - if err := GetDB(ctx).Where("channel_id = ?", channelID).Order("id DESC").Find(&rows).Error; err != nil { - return nil, err - } - return rows, nil -} - -// GetMessageBinding loads a binding by id. -func GetMessageBinding(ctx context.Context, id uint64) (*entity.MessageBinding, error) { - var b entity.MessageBinding - if err := GetDB(ctx).Where("id = ?", id).First(&b).Error; err != nil { - return nil, mapNotFound(err) - } - return &b, nil -} - -// DeleteMessageBinding deletes a binding by id. -func DeleteMessageBinding(ctx context.Context, id uint64) error { - return GetDB(ctx).Delete(&entity.MessageBinding{}, id).Error -} - -// UpsertPairingCode reuses an unexpired code for the same channel+platform user. -func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*entity.MessagePairingCode, error) { - var existing entity.MessagePairingCode - err := GetDB(ctx). - Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()). - First(&existing).Error - if err == nil { - return &existing, nil - } - if !errors.Is(err, gorm.ErrRecordNotFound) { - return nil, err - } - row := &entity.MessagePairingCode{ - Code: code, - ChannelID: channelID, - PlatformUserID: platformUserID, - ExpiresAt: expiresAt, - } - if err := GetDB(ctx).Create(row).Error; err != nil { - return nil, err - } - return row, nil -} - -// GetPairingCode loads a pairing code by normalized code string. -func GetPairingCode(ctx context.Context, code string) (*entity.MessagePairingCode, error) { - var row entity.MessagePairingCode - if err := GetDB(ctx).Where("code = ?", code).First(&row).Error; err != nil { - return nil, mapNotFound(err) - } - return &row, nil -} - -// DeletePairingCode removes a pairing code. -func DeletePairingCode(ctx context.Context, code string) error { - return GetDB(ctx).Where("code = ?", code).Delete(&entity.MessagePairingCode{}).Error -} - -// DeleteExpiredPairingCodes removes expired pairing rows. -func DeleteExpiredPairingCodes(ctx context.Context) error { - return GetDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&entity.MessagePairingCode{}).Error -} - -// ListEnabledMessageChannels returns enabled channels. -func ListEnabledMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) { - var rows []entity.MessageChannel - if err := GetDB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil { - return nil, err - } - return rows, nil -} diff --git a/backend/plugins/domain/msg_gateway/dao/push.go b/backend/plugins/domain/msg_gateway/dao/push.go index f6d9e940..4c3e24e2 100644 --- a/backend/plugins/domain/msg_gateway/dao/push.go +++ b/backend/plugins/domain/msg_gateway/dao/push.go @@ -4,15 +4,9 @@ package dao import ( - "Wavelet/core" - "Wavelet/core/contracts" - "Wavelet/plugins/domain/msg_gateway/consts" "Wavelet/plugins/domain/msg_gateway/model/do" "Wavelet/plugins/domain/msg_gateway/model/entity" "context" - "errors" - "fmt" - "sync" "time" "gorm.io/gorm" @@ -23,31 +17,6 @@ const ( activePushEventCacheTTL = 24 * time.Hour ) -var ( - cacheMu sync.RWMutex - cacheSvc contracts.CacheService -) - -// SetCacheService sets the cache service singleton. -func SetCacheService(s contracts.CacheService) { - cacheMu.Lock() - defer cacheMu.Unlock() - cacheSvc = s -} - -// GetCache resolves the cache service for the current call. -func GetCache(ctx context.Context) contracts.CacheService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil { - return s - } - } - cacheMu.RLock() - s := cacheSvc - cacheMu.RUnlock() - return s -} - // ListPushChannelsRecord returns all push channels ordered by creation time descending. func ListPushChannelsRecord(ctx context.Context) ([]entity.PushChannel, error) { var channels []entity.PushChannel @@ -112,14 +81,15 @@ func DeletePushChannelRecord(ctx context.Context, channel *entity.PushChannel) e } func getCachedOrQuery[T any](ctx context.Context, cacheKey string, ttl time.Duration, query func(db *gorm.DB, dest *T) error) (*T, error) { - var val T if cache := GetCache(ctx); cache != nil { + var val T if err := cache.Get(ctx, cacheKey, &val); err == nil { return &val, nil } } db := GetDB(ctx) + var val T if err := query(db, &val); err != nil { return nil, err } @@ -255,6 +225,9 @@ func ListPushHistoriesRecord(ctx context.Context, filter do.PushHistoryListFilte if filter.EventKey != "" { query = query.Where("event_key = ?", filter.EventKey) } + if filter.Channel != "" { + query = query.Where("channel = ?", filter.Channel) + } if filter.Status != "" { query = query.Where("status = ?", filter.Status) } @@ -282,77 +255,3 @@ func CreatePushHistoryRecord(ctx context.Context, history *entity.PushHistory) e func PushHistoryQuery(ctx context.Context) *gorm.DB { return GetDB(ctx).Model(&entity.PushHistory{}) } - -// smtpConfigKeys are the system-config rows backing the built-in email channel. -var smtpConfigKeys = []string{"smtp_host", "smtp_port", "smtp_username", "smtp_password"} - -// LoadSMTPConfigRecord reads the SMTP settings in one query. -func LoadSMTPConfigRecord(ctx context.Context) (do.SMTPConfig, error) { - db := GetDB(ctx) - if db == nil { - return do.SMTPConfig{}, errors.New("database not available") - } - - var rows []struct { - Key string - Value string - } - if err := db.Table("w_system_configs"). - Select("key", "value"). - Where("key IN ?", smtpConfigKeys). - Find(&rows).Error; err != nil { - return do.SMTPConfig{}, fmt.Errorf("read smtp system configs: %w", err) - } - - var cfg do.SMTPConfig - for _, row := range rows { - switch row.Key { - case "smtp_host": - cfg.Host = row.Value - case "smtp_port": - cfg.Port = row.Value - case "smtp_username": - cfg.Username = row.Value - case "smtp_password": - cfg.Password = row.Value - } - } - return cfg, nil -} - -// userLookupColumns allow-lists the columns FindUserByFieldRecord may filter on. -var userLookupColumns = map[string]struct{}{ - "id": {}, - "username": {}, -} - -// FindUserByFieldRecord is the user lookup fallback for when the UserService -// contract is not wired yet. field must be one of userLookupColumns. -func FindUserByFieldRecord(ctx context.Context, field string, value any) (*contracts.UserDTO, error) { - if _, ok := userLookupColumns[field]; !ok { - return nil, consts.ErrUnsupportedUserLookupField - } - db := GetDB(ctx) - if db == nil { - return nil, consts.ErrRecordNotFound - } - var user contracts.UserDTO - if err := db.Table("w_users").Where(field+" = ?", value).First(&user).Error; err != nil { - return nil, err - } - return &user, nil -} - -// FindFirstAdminUserRecord is the admin lookup fallback for when the UserService -// contract is not wired yet. -func FindFirstAdminUserRecord(ctx context.Context) (*contracts.UserDTO, error) { - db := GetDB(ctx) - if db == nil { - return nil, consts.ErrRecordNotFound - } - var adminUser contracts.UserDTO - if err := db.Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil { - return nil, err - } - return &adminUser, nil -} diff --git a/backend/plugins/domain/msg_gateway/dao/push_test.go b/backend/plugins/domain/msg_gateway/dao/push_test.go index 6c8d3061..d8d479eb 100644 --- a/backend/plugins/domain/msg_gateway/dao/push_test.go +++ b/backend/plugins/domain/msg_gateway/dao/push_test.go @@ -4,128 +4,93 @@ package dao_test import ( + "Wavelet/pkg/idgen" "Wavelet/pkg/testhelper" - "Wavelet/plugins/domain/msg_gateway/consts" "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/entity" "context" - "errors" "testing" - "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "gorm.io/gorm" ) -// stubDBService satisfies contracts.DBService over a test database handle. type stubDBService struct{ db *gorm.DB } -func (s stubDBService) GORM() *gorm.DB { return s.db } - +func (s stubDBService) GORM() *gorm.DB { return s.db } func (s stubDBService) DB(_ context.Context) *gorm.DB { return s.db } +func (s stubDBService) Named(_ string) *gorm.DB { return s.db } -func (s stubDBService) Named(_ string) *gorm.DB { return s.db } - -// TestFindUserByFieldRecordRejectsUnlistedColumns pins the column allow-list. The -// lookup column is interpolated into SQL, so an unlisted name must be refused before -// any query is built rather than trusted because call sites happen to pass literals. -func TestFindUserByFieldRecordRejectsUnlistedColumns(t *testing.T) { +func TestPushChannelDAO_CRUD(t *testing.T) { + _ = idgen.Init(1) db, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() - - if err := db.Table("w_users").Create(map[string]any{"id": 77, "username": "seeded"}).Error; err != nil { - t.Fatalf("seed user failed: %v", err) - } + require.NoError(t, db.AutoMigrate(&entity.PushChannel{}, &entity.PushEvent{}, &entity.PushHistory{})) dao.SetDBServiceForTest(stubDBService{db: db}) t.Cleanup(func() { dao.SetDBServiceForTest(nil) }) ctx := context.Background() - user, err := dao.FindUserByFieldRecord(ctx, "username", "seeded") - if err != nil { - t.Fatalf("allowlisted lookup by username failed: %v", err) - } - if user.ID != 77 { - t.Errorf("allowlisted lookup returned ID %d, want 77", user.ID) - } - if _, err := dao.FindUserByFieldRecord(ctx, "id", uint64(77)); err != nil { - t.Errorf("allowlisted lookup by id failed: %v", err) + ch := entity.PushChannel{ + Name: "test_webhook", + Type: "custom", + URL: "https://example.com/hook", + Enabled: true, } + require.NoError(t, dao.CreatePushChannelRecord(ctx, &ch)) + assert.NotZero(t, ch.ID) - cases := []struct { - name string - field string - }{ - {"tautology injection", `username = '' OR 1=1 --`}, - {"stacked statement", "id; DROP TABLE w_users"}, - {"column outside allow-list", "password"}, - {"empty field", ""}, - } - for _, tc := range cases { - if _, err := dao.FindUserByFieldRecord(ctx, tc.field, "seeded"); !errors.Is(err, consts.ErrUnsupportedUserLookupField) { - t.Errorf("%s: got err %v, want ErrUnsupportedUserLookupField", tc.name, err) - } - } + loaded, err := dao.GetPushChannelByIDRecord(ctx, ch.ID) + require.NoError(t, err) + assert.Equal(t, "test_webhook", loaded.Name) - var remaining int64 - if err := db.Table("w_users").Count(&remaining).Error; err != nil || remaining != 1 { - t.Fatalf("w_users damaged by rejected lookups: count=%d err=%v", remaining, err) - } + active, err := dao.GetActivePushChannelByName(ctx, "test_webhook") + require.NoError(t, err) + assert.Equal(t, ch.ID, active.ID) + + channels, err := dao.ListPushChannelsRecord(ctx) + require.NoError(t, err) + assert.NotEmpty(t, channels) + + require.NoError(t, dao.DeletePushChannelRecord(ctx, &ch)) + _, err = dao.GetPushChannelByIDRecord(ctx, ch.ID) + assert.Error(t, err) } -// smtpTestValues are the four system-config rows the built-in email channel reads. -var smtpTestValues = map[string]string{ - "smtp_host": "mail.example.test", - "smtp_port": "465", - "smtp_username": "notify@example.test", - "smtp_password": "s3cret-value", -} - -// TestLoadSMTPConfigRecordMapsEveryKey guards the single-query rewrite: every field -// must still be filled from its own row. -func TestLoadSMTPConfigRecordMapsEveryKey(t *testing.T) { +func TestPushEventDAO_CRUD(t *testing.T) { + _ = idgen.Init(1) db, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() - - keys := make([]string, 0, len(smtpTestValues)) - for key := range smtpTestValues { - keys = append(keys, key) - } - if err := db.Table("w_system_configs").Where("key IN ?", keys).Delete(map[string]any{}).Error; err != nil { - t.Fatalf("clear smtp rows: %v", err) - } - for _, key := range keys { - row := map[string]any{"key": key, "value": smtpTestValues[key], "type": "system"} - if err := db.Table("w_system_configs").Create(row).Error; err != nil { - t.Fatalf("seed %s: %v", key, err) - } - } + require.NoError(t, db.AutoMigrate(&entity.PushChannel{}, &entity.PushEvent{}, &entity.PushHistory{})) dao.SetDBServiceForTest(stubDBService{db: db}) t.Cleanup(func() { dao.SetDBServiceForTest(nil) }) - cfg, err := dao.LoadSMTPConfigRecord(context.Background()) - if err != nil { - t.Fatalf("LoadSMTPConfigRecord: %v", err) - } - if cfg.Host != smtpTestValues["smtp_host"] || cfg.Port != smtpTestValues["smtp_port"] || - cfg.Username != smtpTestValues["smtp_username"] || cfg.Password != smtpTestValues["smtp_password"] { - t.Errorf("got %+v, want every SMTP field mapped from its own row", cfg) - } -} - -// TestLoadSMTPConfigRecordSurfacesReadFailure pins the actual defect: a read that -// fails used to be discarded, returning four blank strings that callers could only -// interpret as "SMTP was never configured", so the notification was dropped silently. -func TestLoadSMTPConfigRecordSurfacesReadFailure(t *testing.T) { - bare, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) - if err != nil { - t.Fatalf("open bare sqlite: %v", err) - } - - dao.SetDBServiceForTest(stubDBService{db: bare}) - t.Cleanup(func() { dao.SetDBServiceForTest(nil) }) - - if _, err := dao.LoadSMTPConfigRecord(context.Background()); err == nil { - t.Fatal("LoadSMTPConfigRecord returned nil error although the config table cannot be read") + ctx := context.Background() + + ev := entity.PushEvent{ + EventKey: "test_event", + Name: "测试事件", + Channels: []string{"test_webhook"}, + Targets: []string{"admin"}, + Template: `{"title":"Hello"}`, + Enabled: true, } + require.NoError(t, dao.CreatePushEventRecord(ctx, &ev)) + assert.NotZero(t, ev.ID) + + loaded, err := dao.GetPushEventByKeyRecord(ctx, "test_event") + require.NoError(t, err) + assert.Equal(t, "测试事件", loaded.Name) + + require.NoError(t, dao.UpdatePushEventEnabledRecord(ctx, &ev, false)) + loadedDisabled, err := dao.GetPushEventByIDRecord(ctx, ev.ID) + require.NoError(t, err) + assert.False(t, loadedDisabled.Enabled) + + require.NoError(t, dao.DeletePushEventRecord(ctx, &ev)) + _, err = dao.GetPushEventByIDRecord(ctx, ev.ID) + assert.Error(t, err) } diff --git a/backend/plugins/domain/msg_gateway/model/do/message.go b/backend/plugins/domain/msg_gateway/model/do/bot.go similarity index 100% rename from backend/plugins/domain/msg_gateway/model/do/message.go rename to backend/plugins/domain/msg_gateway/model/do/bot.go diff --git a/backend/plugins/domain/msg_gateway/model/do/push.go b/backend/plugins/domain/msg_gateway/model/do/push.go index 1a35f4a4..a990a286 100644 --- a/backend/plugins/domain/msg_gateway/model/do/push.go +++ b/backend/plugins/domain/msg_gateway/model/do/push.go @@ -6,7 +6,6 @@ package do import ( "Wavelet/plugins/domain/msg_gateway/consts" pkgpush "Wavelet/plugins/domain/msg_gateway/push" - "sync" "time" ) @@ -159,58 +158,9 @@ type PushNotificationEvent struct { Metadata map[string]any `json:"metadata,omitempty"` } -var ( - pushDefMu sync.RWMutex - pushDefinitions = make(map[string]PushDefinition) -) - -// RegisterPushChannelDefinition registers a channel definition. -func RegisterPushChannelDefinition(def PushDefinition) { - pushDefMu.Lock() - defer pushDefMu.Unlock() - pushDefinitions[def.Type] = def -} - -// ListPushDefinitions returns all registered channel definitions. -func ListPushDefinitions() []PushDefinition { - pushDefMu.RLock() - defer pushDefMu.RUnlock() - - order := []string{ - consts.ChannelCustom, - consts.ChannelLark, - consts.ChannelDingTalk, - consts.ChannelTelegram, - consts.ChannelBark, - consts.ChannelDiscord, - consts.ChannelSlack, - consts.ChannelPushover, - consts.ChannelEmail, - } - res := make([]PushDefinition, 0, len(pushDefinitions)) - for _, t := range order { - if d, ok := pushDefinitions[t]; ok { - res = append(res, d) - } - } - for t, d := range pushDefinitions { - found := false - for _, o := range order { - if o == t { - found = true - break - } - } - if !found { - res = append(res, d) - } - } - return res -} - -//nolint:funlen,goconst,dupl // Channel definitions registration table -func init() { - RegisterPushChannelDefinition(PushDefinition{ +//nolint:goconst,dupl // Static push channel form definitions table +var defaultPushDefinitions = []PushDefinition{ + { Type: consts.ChannelCustom, Name: "自定义消息通道", Description: "使用自定义 HTTP POST 请求向外部 Webhook 发送数据。", @@ -232,9 +182,8 @@ func init() { Description: "可使用的变量:$title, $description, $content, $url, $to。例如 {\"text\": \"$content\"}", }, }, - }) - - RegisterPushChannelDefinition(PushDefinition{ + }, + { Type: consts.ChannelLark, Name: "飞书群机器人", Description: "配置飞书群自定义机器人的 Webhook 接口投递。", @@ -264,9 +213,8 @@ func init() { Description: "若填写,必须是合法的飞书卡片 JSON 格式", }, }, - }) - - RegisterPushChannelDefinition(PushDefinition{ + }, + { Type: consts.ChannelDingTalk, Name: "钉钉群机器人", Description: "配置钉钉群自定义机器人的 Webhook 接口投递。", @@ -288,9 +236,8 @@ func init() { Description: "钉钉群机器人安全设置中的加签 Secret", }, }, - }) - - RegisterPushChannelDefinition(PushDefinition{ + }, + { Type: consts.ChannelTelegram, Name: "Telegram 机器人", Description: "配置 Telegram 机器人推送消息。", @@ -320,9 +267,8 @@ func init() { Description: "默认的消息接收 Chat ID。如果通知事件中未配置 targets,将推送到此 ID", }, }, - }) - - RegisterPushChannelDefinition(PushDefinition{ + }, + { Type: consts.ChannelBark, Name: "Bark (iOS 推送)", Description: "配置 Bark 推送通知至 iPhone / iPad 客户端。", @@ -352,9 +298,8 @@ func init() { Description: "可选的 JSON 配置,支持 group (分组)、sound (铃声)、icon (自定义图标)", }, }, - }) - - RegisterPushChannelDefinition(PushDefinition{ + }, + { Type: consts.ChannelDiscord, Name: "Discord 频道", Description: "配置 Discord 频道的 Incoming Webhook 消息推送。", @@ -368,9 +313,8 @@ func init() { Description: "从 Discord 频道集成设置中复制的 Webhook URL", }, }, - }) - - RegisterPushChannelDefinition(PushDefinition{ + }, + { Type: consts.ChannelSlack, Name: "Slack 频道", Description: "配置 Slack 频道的 Incoming Webhook 消息推送。", @@ -384,9 +328,8 @@ func init() { Description: "从 Slack 应用配置中复制的 Incoming Webhook URL", }, }, - }) - - RegisterPushChannelDefinition(PushDefinition{ + }, + { Type: consts.ChannelPushover, Name: "Pushover 推送", Description: "配置 Pushover 即时推送到手机/桌面客户端。", @@ -408,12 +351,18 @@ func init() { Description: "Pushover 个人账号的 User Key", }, }, - }) - - RegisterPushChannelDefinition(PushDefinition{ + }, + { Type: consts.ChannelEmail, Name: "邮件推送通道", Description: "邮件推送通道直接使用系统全局 SMTP 设置进行发送,无需在此填写服务器配置。", Fields: []PushField{}, - }) + }, +} + +// ListPushDefinitions returns all registered channel definitions. +func ListPushDefinitions() []PushDefinition { + res := make([]PushDefinition, len(defaultPushDefinitions)) + copy(res, defaultPushDefinitions) + return res } diff --git a/backend/plugins/domain/msg_gateway/model/entity/channel.go b/backend/plugins/domain/msg_gateway/model/entity/bot.go similarity index 100% rename from backend/plugins/domain/msg_gateway/model/entity/channel.go rename to backend/plugins/domain/msg_gateway/model/entity/bot.go diff --git a/backend/plugins/domain/msg_gateway/msg_gateway_test.go b/backend/plugins/domain/msg_gateway/msg_gateway_test.go deleted file mode 100644 index f3f72508..00000000 --- a/backend/plugins/domain/msg_gateway/msg_gateway_test.go +++ /dev/null @@ -1,51 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package msg_gateway_test - -import ( - "Wavelet/pkg/testhelper" - "Wavelet/plugins/domain/msg_gateway" - "context" - "testing" - "time" - - "gorm.io/gorm" -) - -type mockDBService struct { - db *gorm.DB -} - -func (m *mockDBService) GORM() *gorm.DB { - return m.db -} - -func (m *mockDBService) DB(ctx context.Context) *gorm.DB { - return m.db.WithContext(ctx) -} - -func (m *mockDBService) Named(_ string) *gorm.DB { - return m.db -} - -func TestUpsertPairingCode_ReusesUnexpired(t *testing.T) { - testDB, _, cleanup := testhelper.SetupTestEnvironment(t) - msg_gateway.SetDBServiceForTest(&mockDBService{db: testDB}) - defer func() { - msg_gateway.SetDBServiceForTest(nil) - cleanup() - }() - ctx := context.Background() - first, err := msg_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ABCD1234", time.Now().Add(15*time.Minute)) - if err != nil { - t.Fatal(err) - } - second, err := msg_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ZZZZ9999", time.Now().Add(15*time.Minute)) - if err != nil { - t.Fatal(err) - } - if first.Code != second.Code || first.Code != "ABCD1234" { - t.Fatalf("reuse failed: %+v %+v", first, second) - } -} diff --git a/backend/plugins/domain/msg_gateway/pairing_test.go b/backend/plugins/domain/msg_gateway/pairing_test.go deleted file mode 100644 index e047db88..00000000 --- a/backend/plugins/domain/msg_gateway/pairing_test.go +++ /dev/null @@ -1,33 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package msg_gateway - -import ( - "strings" - "testing" -) - -func TestGenerateCode_AlphabetAndLength(t *testing.T) { - code, err := GenerateCode() - if err != nil { - t.Fatal(err) - } - if len(code) != 8 { - t.Fatalf("len=%d", len(code)) - } - for _, r := range code { - if !strings.ContainsRune(CodeAlphabet, r) { - t.Fatalf("bad rune %q", r) - } - } -} - -func TestNormalizeAndFormat(t *testing.T) { - if got := NormalizeCode("ab-cd-ef-gh"); got != "ABCDEFGH" { - t.Fatalf("got %q", got) - } - if got := FormatCode("ABCDEFGH"); got != "ABCD-EFGH" { - t.Fatalf("got %q", got) - } -} diff --git a/backend/plugins/domain/msg_gateway/plugin.go b/backend/plugins/domain/msg_gateway/plugin.go index 1188ae9a..b32c4481 100644 --- a/backend/plugins/domain/msg_gateway/plugin.go +++ b/backend/plugins/domain/msg_gateway/plugin.go @@ -64,9 +64,6 @@ func (p *Plugin) Name() string { func (p *Plugin) Inject() []reflect.Type { return []reflect.Type{ reflect.TypeFor[contracts.DBService](), - // AuthService is captured as a middleware value in Apply, so it cannot - // be late-bound with core.When like the other services below; the - // kernel must mount auth first or the routes get a pass-through guard. reflect.TypeFor[contracts.AuthService](), } } @@ -246,116 +243,20 @@ func (p *Plugin) Apply(ctx *core.Context) error { return nil } -// Re-exported constants. -const ( - CodeAlphabet = service.CodeAlphabet - CodeLength = service.CodeLength -) - -// MessageChannel is an alias for entity.MessageChannel. -type MessageChannel = entity.MessageChannel - -// MessageBinding is an alias for entity.MessageBinding. -type MessageBinding = entity.MessageBinding - -// MessagePairingCode is an alias for entity.MessagePairingCode. -type MessagePairingCode = entity.MessagePairingCode - -// PushChannel is an alias for entity.PushChannel. -type PushChannel = entity.PushChannel - -// PushEvent is an alias for entity.PushEvent. -type PushEvent = entity.PushEvent - -// PushHistory is an alias for entity.PushHistory. -type PushHistory = entity.PushHistory - -// PushNotificationEvent is an alias for do.PushNotificationEvent. -type PushNotificationEvent = do.PushNotificationEvent - -// ChannelConfig is an alias for do.ChannelConfig. -type ChannelConfig = do.ChannelConfig - -// Capability is an alias for do.Capability. -type Capability = do.Capability - -// Recipient is an alias for do.Recipient. -type Recipient = do.Recipient - -// Attachment is an alias for do.Attachment. -type Attachment = do.Attachment - -// InboundMessage is an alias for do.InboundMessage. -type InboundMessage = do.InboundMessage - -// OutboundMessage is an alias for do.OutboundMessage. -type OutboundMessage = do.OutboundMessage - -// BindingDTO is an alias for do.BindingDTO. -type BindingDTO = do.BindingDTO - -// PublicChannelDTO is an alias for do.PublicChannelDTO. -type PublicChannelDTO = do.PublicChannelDTO - -// Definition is an alias for do.Definition. -type Definition = do.Definition - -// ChannelDTO is an alias for do.ChannelDTO. -type ChannelDTO = do.ChannelDTO - -// CreateChannelRequest is an alias for do.CreateChannelRequest. -type CreateChannelRequest = do.CreateChannelRequest - -// UpdateChannelRequest is an alias for do.UpdateChannelRequest. -type UpdateChannelRequest = do.UpdateChannelRequest - -// PushDefinition is an alias for do.PushDefinition. -type PushDefinition = do.PushDefinition - -// PushField is an alias for do.PushField. -type PushField = do.PushField - -// NotificationMessage is an alias for do.NotificationMessage. -type NotificationMessage = do.NotificationMessage - -// EventMetadata is an alias for do.EventMetadata. -type EventMetadata = do.EventMetadata - -// SendPayload is an alias for do.SendPayload. -type SendPayload = do.SendPayload - -// Handler is an alias for service.Handler. -type Handler = service.Handler - -// Factory is an alias for service.Factory. -type Factory = service.Factory - -// Channel is an alias for service.Channel. -type Channel = service.Channel - -// Runner is an alias for service.Runner. -type Runner = service.Runner - -// EventTrigger is an alias for service.EventTrigger. -type EventTrigger = service.EventTrigger - -// PushHandler is an alias for service.PushHandler. -type PushHandler = service.PushHandler - -// Re-exported variables and functions. -var ( - SetDBServiceForTest = dao.SetDBServiceForTest - UpsertPairingCode = dao.UpsertPairingCode - Register = service.Register - Lookup = service.Lookup - GenerateCode = service.GenerateCode - NormalizeCode = service.NormalizeCode - FormatCode = service.FormatCode - Start = service.Start - Stop = service.Stop - GlobalRunner = service.GlobalRunner - DefaultTrigger = service.DefaultTrigger - SyncEvents = service.SyncEvents - AdminLogin = service.AdminLogin - HandleAdminLoggedIn = service.HandleAdminLoggedIn +// Entity and DO aliases exported for integration test compatibility. +type ( + // MessageChannel is an alias for entity.MessageChannel. + MessageChannel = entity.MessageChannel + // MessageBinding is an alias for entity.MessageBinding. + MessageBinding = entity.MessageBinding + // MessagePairingCode is an alias for entity.MessagePairingCode. + MessagePairingCode = entity.MessagePairingCode + // PushChannel is an alias for entity.PushChannel. + PushChannel = entity.PushChannel + // PushEvent is an alias for entity.PushEvent. + PushEvent = entity.PushEvent + // PushHistory is an alias for entity.PushHistory. + PushHistory = entity.PushHistory + // PushNotificationEvent is an alias for do.PushNotificationEvent. + PushNotificationEvent = do.PushNotificationEvent ) diff --git a/backend/plugins/domain/msg_gateway/registry_test.go b/backend/plugins/domain/msg_gateway/registry_test.go deleted file mode 100644 index 5431f0c6..00000000 --- a/backend/plugins/domain/msg_gateway/registry_test.go +++ /dev/null @@ -1,38 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package msg_gateway - -import ( - "context" - "testing" -) - -type stubChannel struct{} - -func (stubChannel) Type() string { return "stub" } -func (stubChannel) Connect(context.Context) error { - return nil -} -func (stubChannel) Disconnect(context.Context) error { return nil } -func (stubChannel) Send(context.Context, Recipient, OutboundMessage) error { - return nil -} -func (stubChannel) Capabilities() Capability { return Capability{Text: true} } - -func TestRegisterLookup(t *testing.T) { - Register("stub", func(ChannelConfig, Handler) (Channel, error) { - return stubChannel{}, nil - }) - fn, ok := Lookup("stub") - if !ok { - t.Fatal("expected factory") - } - ch, err := fn(ChannelConfig{}, nil) - if err != nil { - t.Fatal(err) - } - if ch.Type() != "stub" { - t.Fatalf("type=%s", ch.Type()) - } -} diff --git a/backend/plugins/domain/msg_gateway/service/admin.go b/backend/plugins/domain/msg_gateway/service/bot_channel.go similarity index 91% rename from backend/plugins/domain/msg_gateway/service/admin.go rename to backend/plugins/domain/msg_gateway/service/bot_channel.go index 15150fb8..3d64f14c 100644 --- a/backend/plugins/domain/msg_gateway/service/admin.go +++ b/backend/plugins/domain/msg_gateway/service/bot_channel.go @@ -90,7 +90,7 @@ func UpdateChannel(ctx context.Context, id uint64, req do.UpdateChannelRequest) row, err := dao.GetMessageChannel(ctx, id) if err != nil { if errors.Is(err, consts.ErrRecordNotFound) { - return do.ChannelDTO{}, errors.New(consts.ErrChannelNotFound) + return do.ChannelDTO{}, errors.New(consts.ErrChannelNotFoundText) } return do.ChannelDTO{}, err } @@ -157,7 +157,7 @@ func ListChannels(ctx context.Context) ([]do.ChannelDTO, error) { func DeleteChannel(ctx context.Context, id uint64) error { if _, err := dao.GetMessageChannel(ctx, id); err != nil { if errors.Is(err, consts.ErrRecordNotFound) { - return errors.New(consts.ErrChannelNotFound) + return errors.New(consts.ErrChannelNotFoundText) } return err } @@ -169,7 +169,7 @@ func ProbeChannel(ctx context.Context, id uint64) error { row, err := dao.GetMessageChannel(ctx, id) if err != nil { if errors.Is(err, consts.ErrRecordNotFound) { - return errors.New(consts.ErrChannelNotFound) + return errors.New(consts.ErrChannelNotFoundText) } return err } @@ -285,27 +285,3 @@ func ToDTO(row *entity.MessageChannel, creds, extra map[string]string) do.Channe Extra: extra, } } - -// MaskCredentials hides secret bearing credential entries. -func MaskCredentials(_ string, in map[string]string) map[string]string { - out := make(map[string]string, len(in)) - for k, v := range in { - if k == "token" || k == "client_secret" { - out[k] = MaskSecret(v) - } else { - out[k] = v - } - } - return out -} - -const minMaskSecretLength = 8 - -// MaskSecret keeps only a short visible prefix and suffix of a secret. -func MaskSecret(s string) string { - s = strings.TrimSpace(s) - if len(s) <= minMaskSecretLength { - return "******" - } - return s[:4] + "..." + s[len(s)-4:] -} diff --git a/backend/plugins/domain/msg_gateway/service/bot_crypto.go b/backend/plugins/domain/msg_gateway/service/bot_crypto.go new file mode 100644 index 00000000..ce9600aa --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/bot_crypto.go @@ -0,0 +1,113 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/pkg/util" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "strings" + "sync" +) + +var ( + credentialSecretMu sync.RWMutex + credentialSecret string +) + +// SetCredentialSecret sets the secret used to derive CredentialKey. +func SetCredentialSecret(secret string) { + credentialSecretMu.Lock() + defer credentialSecretMu.Unlock() + credentialSecret = secret +} + +// CredentialKey is AES-256 hex derived from the session secret. +func CredentialKey() string { + credentialSecretMu.RLock() + secret := credentialSecret + credentialSecretMu.RUnlock() + sum := sha256.Sum256([]byte(secret)) + return hex.EncodeToString(sum[:]) +} + +// EncryptCredentials encrypts a credential map as JSON ciphertext. +func EncryptCredentials(creds map[string]string) (string, error) { + if creds == nil { + creds = map[string]string{} + } + raw, err := json.Marshal(creds) + if err != nil { + return "", err + } + return util.Encrypt(CredentialKey(), string(raw)) +} + +// DecryptCredentials decrypts a credential map from ciphertext. +func DecryptCredentials(ciphertext string) (map[string]string, error) { + if ciphertext == "" { + return map[string]string{}, nil + } + plain, err := util.Decrypt(CredentialKey(), ciphertext) + if err != nil { + return nil, err + } + var out map[string]string + if err := json.Unmarshal([]byte(plain), &out); err != nil { + return nil, err + } + if out == nil { + out = map[string]string{} + } + return out, nil +} + +// ParseExtra decodes optional extra JSON into a string map. +func ParseExtra(raw string) map[string]string { + if raw == "" { + return map[string]string{} + } + var out map[string]string + if err := json.Unmarshal([]byte(raw), &out); err != nil || out == nil { + return map[string]string{} + } + return out +} + +// EncodeExtra encodes extra fields as JSON string. +func EncodeExtra(extra map[string]string) string { + if extra == nil { + return "" + } + raw, err := json.Marshal(extra) + if err != nil { + return "" + } + return string(raw) +} + +// MaskCredentials hides secret-bearing credential entries. +func MaskCredentials(_ string, in map[string]string) map[string]string { + out := make(map[string]string, len(in)) + for k, v := range in { + if k == "token" || k == "client_secret" || k == "app_secret" || k == "bot_token" { + out[k] = MaskSecret(v) + } else { + out[k] = v + } + } + return out +} + +const minMaskSecretLength = 8 + +// MaskSecret keeps only a short visible prefix and suffix of a secret. +func MaskSecret(s string) string { + s = strings.TrimSpace(s) + if len(s) <= minMaskSecretLength { + return "******" + } + return s[:4] + "..." + s[len(s)-4:] +} diff --git a/backend/plugins/domain/msg_gateway/service/dispatch.go b/backend/plugins/domain/msg_gateway/service/bot_dispatch.go similarity index 99% rename from backend/plugins/domain/msg_gateway/service/dispatch.go rename to backend/plugins/domain/msg_gateway/service/bot_dispatch.go index ff69ec12..448f1fa0 100644 --- a/backend/plugins/domain/msg_gateway/service/dispatch.go +++ b/backend/plugins/domain/msg_gateway/service/bot_dispatch.go @@ -83,7 +83,7 @@ func (h *BotDispatchHandler) Execute(ctx context.Context, payload []byte) (*cont } channels = filtered if len(channels) == 0 { - return nil, errors.New(consts.ErrChannelNotFound) + return nil, consts.ErrChannelNotFound } } diff --git a/backend/plugins/domain/msg_gateway/service/dispatch_test.go b/backend/plugins/domain/msg_gateway/service/bot_dispatch_test.go similarity index 100% rename from backend/plugins/domain/msg_gateway/service/dispatch_test.go rename to backend/plugins/domain/msg_gateway/service/bot_dispatch_test.go diff --git a/backend/plugins/domain/msg_gateway/service/bot_pairing.go b/backend/plugins/domain/msg_gateway/service/bot_pairing.go new file mode 100644 index 00000000..365be507 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/bot_pairing.go @@ -0,0 +1,174 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" + "context" + "crypto/rand" + "errors" + "strconv" + "strings" + "time" + "unicode" +) + +// GenerateCode returns an 8-character pairing code using crypto/rand. +func GenerateCode() (string, error) { + buf := make([]byte, consts.CodeLength) + if _, err := rand.Read(buf); err != nil { + return "", err + } + out := make([]byte, consts.CodeLength) + for i, b := range buf { + out[i] = consts.CodeAlphabet[int(b)%len(consts.CodeAlphabet)] + } + return string(out), nil +} + +// NormalizeCode strips separators and uppercases. +func NormalizeCode(s string) string { + var b strings.Builder + for _, r := range s { + if r == '-' || unicode.IsSpace(r) { + continue + } + b.WriteRune(unicode.ToUpper(r)) + } + return b.String() +} + +// FormatCode renders ABCD-EFGH format. +func FormatCode(s string) string { + s = NormalizeCode(s) + if len(s) != consts.CodeLength { + return s + } + return s[:4] + "-" + s[4:] +} + +// BindChannel consumes a pairing code and binds the platform identity to the user. +func BindChannel(ctx context.Context, userID uint64, req do.BindRequest) (do.BindingDTO, error) { + channelID, err := strconv.ParseUint(strings.TrimSpace(req.ChannelID), 10, 64) + if err != nil || channelID == 0 { + return do.BindingDTO{}, consts.ErrChannelIDRequired + } + code := NormalizeCode(req.Code) + if code == "" { + return do.BindingDTO{}, consts.ErrCodeInvalid + } + pairing, err := dao.GetPairingCode(ctx, code) + if err != nil { + if errors.Is(err, consts.ErrRecordNotFound) { + return do.BindingDTO{}, consts.ErrCodeInvalid + } + return do.BindingDTO{}, err + } + if !pairing.ExpiresAt.After(time.Now()) { + return do.BindingDTO{}, consts.ErrCodeInvalid + } + if pairing.ChannelID != channelID { + return do.BindingDTO{}, consts.ErrChannelMismatch + } + ch, err := dao.GetMessageChannel(ctx, channelID) + if err != nil { + if errors.Is(err, consts.ErrRecordNotFound) { + return do.BindingDTO{}, consts.ErrCodeInvalid + } + return do.BindingDTO{}, err + } + if !ch.Enabled { + return do.BindingDTO{}, consts.ErrChannelDisabled + } + + existing, err := dao.GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID) + if err != nil && !errors.Is(err, consts.ErrRecordNotFound) { + return do.BindingDTO{}, err + } + if err == nil && existing != nil { + if existing.UserID != userID { + return do.BindingDTO{}, consts.ErrPlatformAlreadyBound + } + _ = dao.DeletePairingCode(ctx, pairing.Code) + return ToBindingDTO(existing, ch), nil + } + + row := &entity.MessageBinding{ + UserID: userID, + ChannelID: channelID, + PlatformUserID: pairing.PlatformUserID, + } + if err := dao.CreateMessageBinding(ctx, row); err != nil { + return do.BindingDTO{}, err + } + if err := dao.DeletePairingCode(ctx, pairing.Code); err != nil { + return do.BindingDTO{}, err + } + return ToBindingDTO(row, ch), nil +} + +// ListEnabledPublicChannels returns the channels a user may bind to. +func ListEnabledPublicChannels(ctx context.Context) ([]do.PublicChannelDTO, error) { + rows, err := dao.ListEnabledMessageChannels(ctx) + if err != nil { + return nil, err + } + out := make([]do.PublicChannelDTO, 0, len(rows)) + for _, row := range rows { + out = append(out, do.PublicChannelDTO{ID: row.ID, Name: row.Name, Type: row.Type}) + } + return out, nil +} + +// ListUserBindings returns the binding rows of one user enriched with channel info. +func ListUserBindings(ctx context.Context, userID uint64) ([]do.BindingDTO, error) { + rows, err := dao.ListBindingsByUser(ctx, userID) + if err != nil { + return nil, err + } + out := make([]do.BindingDTO, 0, len(rows)) + for i := range rows { + ch, err := dao.GetMessageChannel(ctx, rows[i].ChannelID) + if err != nil { + out = append(out, ToBindingDTO(&rows[i], nil)) + continue + } + out = append(out, ToBindingDTO(&rows[i], ch)) + } + return out, nil +} + +// UnbindChannel removes a binding owned by the given user. +func UnbindChannel(ctx context.Context, userID, bindingID uint64) error { + row, err := dao.GetMessageBinding(ctx, bindingID) + if err != nil { + if errors.Is(err, consts.ErrRecordNotFound) { + return consts.ErrBindingNotFound + } + return err + } + if row.UserID != userID { + return consts.ErrBindingForbidden + } + return dao.DeleteMessageBinding(ctx, bindingID) +} + +// ToBindingDTO projects a binding row and its optional channel onto the user DTO. +func ToBindingDTO(row *entity.MessageBinding, ch *entity.MessageChannel) do.BindingDTO { + dto := do.BindingDTO{ + ID: row.ID, + UserID: row.UserID, + ChannelID: row.ChannelID, + PlatformUserID: row.PlatformUserID, + CreatedAt: row.CreatedAt, + } + if ch != nil { + dto.ChannelName = ch.Name + dto.ChannelType = ch.Type + } + return dto +} diff --git a/backend/plugins/domain/msg_gateway/service/bot_pairing_test.go b/backend/plugins/domain/msg_gateway/service/bot_pairing_test.go new file mode 100644 index 00000000..b6774127 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/bot_pairing_test.go @@ -0,0 +1,27 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service_test + +import ( + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/service" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestGenerateCode_AlphabetAndLength(t *testing.T) { + code, err := service.GenerateCode() + require.NoError(t, err) + assert.Len(t, code, consts.CodeLength) + for _, r := range code { + assert.Contains(t, consts.CodeAlphabet, string(r)) + } +} + +func TestNormalizeAndFormat(t *testing.T) { + assert.Equal(t, "ABCDEFGH", service.NormalizeCode("ab-cd-ef-gh")) + assert.Equal(t, "ABCD-EFGH", service.FormatCode("ABCDEFGH")) +} diff --git a/backend/plugins/domain/msg_gateway/service/bot_runner.go b/backend/plugins/domain/msg_gateway/service/bot_runner.go new file mode 100644 index 00000000..712fa58f --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/bot_runner.go @@ -0,0 +1,88 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/pkg/logger" + "Wavelet/plugins/domain/msg_gateway/model/do" + "context" + "sync" +) + +// Handler processes one inbound message. +type Handler func(ctx context.Context, msg do.InboundMessage) error + +// Factory constructs a Channel from decrypted config. +type Factory func(cfg do.ChannelConfig, onInbound Handler) (Channel, error) + +// Channel is one connected messaging adapter. +type Channel interface { + Type() string + Connect(ctx context.Context) error + Disconnect(ctx context.Context) error + Send(ctx context.Context, to do.Recipient, msg do.OutboundMessage) error + Capabilities() do.Capability +} + +var ( + factoriesMu sync.RWMutex + factories = map[string]Factory{} +) + +// Register stores a channel factory under typ. +func Register(typ string, fn Factory) { + factoriesMu.Lock() + defer factoriesMu.Unlock() + factories[typ] = fn +} + +// Lookup returns a previously registered factory. +func Lookup(typ string) (Factory, bool) { + factoriesMu.RLock() + defer factoriesMu.RUnlock() + fn, ok := factories[typ] + return fn, ok +} + +// Runner manages lifecycle for long-lived channel adapters (WebSocket, long-polling, etc.). +type Runner struct { + mu sync.Mutex + running bool + cancel context.CancelFunc +} + +// GlobalRunner is the default global runner instance. +var GlobalRunner = &Runner{} + +// Start starts all background long-lived channel runners. +func Start(ctx context.Context) error { + GlobalRunner.mu.Lock() + defer GlobalRunner.mu.Unlock() + + if GlobalRunner.running { + return nil + } + + runCtx, cancel := context.WithCancel(ctx) + GlobalRunner.cancel = cancel + GlobalRunner.running = true + + logger.InfoF(runCtx, "[MessageGateway] Starting bot channel runners...") + return nil +} + +// Stop stops the channel runner. +func Stop() { + GlobalRunner.mu.Lock() + defer GlobalRunner.mu.Unlock() + + if !GlobalRunner.running { + return + } + + if GlobalRunner.cancel != nil { + GlobalRunner.cancel() + } + GlobalRunner.running = false +} diff --git a/backend/plugins/domain/msg_gateway/service/bot_runner_test.go b/backend/plugins/domain/msg_gateway/service/bot_runner_test.go new file mode 100644 index 00000000..3caa01da --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/bot_runner_test.go @@ -0,0 +1,34 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service_test + +import ( + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/service" + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type stubChannel struct{} + +func (stubChannel) Type() string { return "stub" } +func (stubChannel) Connect(context.Context) error { return nil } +func (stubChannel) Disconnect(context.Context) error { return nil } +func (stubChannel) Send(context.Context, do.Recipient, do.OutboundMessage) error { return nil } +func (stubChannel) Capabilities() do.Capability { return do.Capability{Text: true} } + +func TestRegisterLookup(t *testing.T) { + service.Register("stub", func(do.ChannelConfig, service.Handler) (service.Channel, error) { + return stubChannel{}, nil + }) + fn, ok := service.Lookup("stub") + require.True(t, ok) + + ch, err := fn(do.ChannelConfig{}, nil) + require.NoError(t, err) + assert.Equal(t, "stub", ch.Type()) +} diff --git a/backend/plugins/domain/msg_gateway/service/push.go b/backend/plugins/domain/msg_gateway/service/push.go deleted file mode 100644 index 9fac30ff..00000000 --- a/backend/plugins/domain/msg_gateway/service/push.go +++ /dev/null @@ -1,1156 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package service - -import ( - "Wavelet/core/contracts" - "Wavelet/pkg/logger" - "Wavelet/pkg/util" - "Wavelet/plugins/domain/msg_gateway/consts" - "Wavelet/plugins/domain/msg_gateway/dao" - "Wavelet/plugins/domain/msg_gateway/model/do" - "Wavelet/plugins/domain/msg_gateway/model/entity" - pkgpush "Wavelet/plugins/domain/msg_gateway/push" - "context" - "encoding/json" - "errors" - "fmt" - "strconv" - "strings" - "sync" - "time" -) - -var ( - builtInEventsMu sync.RWMutex - // BuiltInEvents lists all built-in events defined in custom_events. - BuiltInEvents []do.EventMetadata -) - -// RegisterBuiltInEvent registers a built-in event definition. -func RegisterBuiltInEvent(meta do.EventMetadata) { - builtInEventsMu.Lock() - defer builtInEventsMu.Unlock() - for i, e := range BuiltInEvents { - if e.Key == meta.Key { - BuiltInEvents[i] = meta - return - } - } - BuiltInEvents = append(BuiltInEvents, meta) -} - -// GetBuiltInEvents returns a copy of registered built-in events. -func GetBuiltInEvents() []do.EventMetadata { - builtInEventsMu.RLock() - defer builtInEventsMu.RUnlock() - out := make([]do.EventMetadata, len(BuiltInEvents)) - copy(out, BuiltInEvents) - return out -} - -// PushRegistryAdapter adapts contracts.PushRegistry onto the built-in event store. -type PushRegistryAdapter struct{} - -// RegisterBuiltInEvent records a built-in push event definition. -func (PushRegistryAdapter) RegisterBuiltInEvent(meta contracts.PushEventMeta) { - RegisterBuiltInEvent(eventMetadataFromContract(meta)) -} - -// SyncEvents persists registered built-in events into the database. -func (PushRegistryAdapter) SyncEvents(ctx context.Context) error { - return SyncEvents(ctx) -} - -func eventMetadataFromContract(meta contracts.PushEventMeta) do.EventMetadata { - return do.EventMetadata{ - Key: meta.Key, - Name: meta.Name, - Description: meta.Description, - DefaultTemplate: do.NotificationMessage{ - Title: meta.DefaultTemplate.Title, - Content: meta.DefaultTemplate.Content, - Level: meta.DefaultTemplate.Level, - Ext: meta.DefaultTemplate.Ext, - }, - } -} - -// SyncBuiltInEvents seeds a database row for every registered built-in event. -func SyncBuiltInEvents(ctx context.Context) error { - for _, meta := range GetBuiltInEvents() { - _, err := dao.GetPushEventByKeyRecord(ctx, meta.Key) - if errors.Is(err, consts.ErrRecordNotFound) { - var defaultTemplateStr string - if defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate); err == nil { - defaultTemplateStr = string(defaultTemplateBytes) - } - event := entity.PushEvent{ - EventKey: meta.Key, - Name: meta.Name, - Channels: []string{}, - Targets: []string{}, - Template: defaultTemplateStr, - Enabled: false, - } - if err := dao.CreatePushEventRecord(ctx, &event); err != nil { - return err - } - } else if err != nil { - return err - } - } - return nil -} - -// ListPushEvents lists all configured push events. -func ListPushEvents(ctx context.Context) ([]entity.PushEvent, error) { - return dao.ListPushEventsRecord(ctx) -} - -// CreatePushEvent stores a push event configuration for a built-in event or task type. -func CreatePushEvent(ctx context.Context, req do.CreatePushEventRequest) (entity.PushEvent, error) { - eventKey, eventName, defaultTemplateBytes, err := GetEventInfo(ctx, req) - if err != nil { - return entity.PushEvent{}, err - } - - count, err := dao.CountPushEventsByKeyRecord(ctx, eventKey) - if err != nil { - return entity.PushEvent{}, err - } - if count > 0 { - return entity.PushEvent{}, errors.New(consts.ErrEventAlreadyConfigured) - } - - templateStr := strings.TrimSpace(req.Template) - if templateStr == "" { - templateStr = string(defaultTemplateBytes) - } else { - var tempMap map[string]any - if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil { - return entity.PushEvent{}, errors.New(consts.ErrTemplateInvalidJSON) - } - } - - channels := req.Channels - if channels == nil { - channels = []string{} - } - targets := req.Targets - if targets == nil { - targets = []string{} - } - - event := entity.PushEvent{ - EventKey: eventKey, - Name: eventName, - TaskType: req.TaskType, - Channels: channels, - Targets: targets, - Template: templateStr, - Enabled: req.Enabled, - } - if err := event.Validate(); err != nil { - return entity.PushEvent{}, err - } - if err := dao.CreatePushEventRecord(ctx, &event); err != nil { - return entity.PushEvent{}, err - } - return event, nil -} - -// DeletePushEvent deletes a push event configuration by id. -func DeletePushEvent(ctx context.Context, id uint64) error { - event, err := dao.GetPushEventByIDRecord(ctx, id) - if err != nil { - return err - } - return dao.DeletePushEventRecord(ctx, &event) -} - -// UpdatePushEvent replaces mutable push event fields. -func UpdatePushEvent(ctx context.Context, id uint64, req do.UpdatePushEventRequest) error { - event, err := dao.GetPushEventByIDRecord(ctx, id) - if err != nil { - return err - } - - event.Channels = req.Channels - event.Targets = req.Targets - event.Template = req.Template - event.Enabled = req.Enabled - if err := event.Validate(); err != nil { - return err - } - return dao.SavePushEventRecord(ctx, &event) -} - -// TogglePushEvent flips the enabled flag of a push event. -func TogglePushEvent(ctx context.Context, id uint64) (bool, error) { - event, err := dao.GetPushEventByIDRecord(ctx, id) - if err != nil { - return false, err - } - - enabled := !event.Enabled - if enabled && len(event.Channels) == 0 { - return false, errors.New(consts.ErrEnableWithoutChannels) - } - if err := dao.UpdatePushEventEnabledRecord(ctx, &event, enabled); err != nil { - return false, err - } - return enabled, nil -} - -// ListPushHistories returns a paginated push delivery audit page. -func ListPushHistories(ctx context.Context, filter do.PushHistoryListFilter) (int64, []entity.PushHistory, error) { - return dao.ListPushHistoriesRecord(ctx, filter) -} - -// ApplySMTPFallbackToPushConfig fills an email config from the system SMTP settings. -func ApplySMTPFallbackToPushConfig(ctx context.Context, cfg *pkgpush.Config) { - if cfg.Channel != consts.ChannelEmail || (cfg.URL != "" && cfg.Key != "") { - return - } - smtp, err := dao.LoadSMTPConfigRecord(ctx) - if err != nil { - logger.ErrorF(ctx, "[Push] 读取 SMTP 系统配置失败: %v", err) - return - } - if smtp.Host == "" || smtp.Username == "" { - return - } - port := smtp.Port - if port == "" { - port = "587" - } - cfg.URL = smtp.Host + ":" + port - cfg.Key = smtp.Username - cfg.Secret = smtp.Password -} - -// RunPushTest validates an ad-hoc channel config and sends a connectivity probe. -func RunPushTest(ctx context.Context, cfg pkgpush.Config, target string) error { - pusher, err := pkgpush.GetPusher(cfg.Channel) - if err != nil { - return err - } - if err := pusher.ValidateConfig(cfg); err != nil { - return fmt.Errorf("%s: %w", consts.ErrValidationFailed, err) - } - - ApplySMTPFallbackToPushConfig(ctx, &cfg) - - testBody := map[string]any{ - consts.KeyTitle: "测试通道推送", - consts.KeyContent: "当您收到这条消息,说明当前渠道连通性测试通过。", - consts.KeyLevel: consts.DefaultLevelInfo, - } - if _, err := pusher.Send(ctx, cfg, target, testBody, "", nil); err != nil { - return err - } - return nil -} - -// ListPushChannels returns every configured push channel. -func ListPushChannels(ctx context.Context) ([]entity.PushChannel, error) { - return dao.ListPushChannelsRecord(ctx) -} - -// CreatePushChannel validates uniqueness and persists a new push channel. -func CreatePushChannel(ctx context.Context, req do.CreatePushChannelRequest) (entity.PushChannel, error) { - count, err := dao.CountPushChannelsByNameRecord(ctx, req.Name) - if err != nil { - return entity.PushChannel{}, err - } - if count > 0 { - return entity.PushChannel{}, errors.New(consts.ErrChannelNameExists) - } - - channel := entity.PushChannel{ - Name: req.Name, - Description: req.Description, - Type: req.Type, - Token: req.Token, - URL: req.URL, - Other: req.Other, - Enabled: req.Enabled, - } - if err := channel.Validate(); err != nil { - return entity.PushChannel{}, err - } - if err := dao.CreatePushChannelRecord(ctx, &channel); err != nil { - return entity.PushChannel{}, err - } - return channel, nil -} - -// UpdatePushChannel replaces the mutable fields of an existing push channel. -func UpdatePushChannel(ctx context.Context, id uint64, req do.UpdatePushChannelRequest) (entity.PushChannel, error) { - channel, err := dao.GetPushChannelByIDRecord(ctx, id) - if err != nil { - return entity.PushChannel{}, err - } - - channel.Description = req.Description - channel.Type = req.Type - channel.Token = req.Token - channel.URL = req.URL - channel.Other = req.Other - channel.Enabled = req.Enabled - if err := channel.Validate(); err != nil { - return entity.PushChannel{}, err - } - if err := dao.SavePushChannelRecord(ctx, &channel); err != nil { - return entity.PushChannel{}, err - } - return channel, nil -} - -// DeletePushChannel removes a push channel by id. -func DeletePushChannel(ctx context.Context, id uint64) error { - channel, err := dao.GetPushChannelByIDRecord(ctx, id) - if err != nil { - return err - } - return dao.DeletePushChannelRecord(ctx, &channel) -} - -// LoadChannelForTest resolves the credentials under test, either from a stored -// channel name or from the ad-hoc values sent by the caller. -func LoadChannelForTest(ctx context.Context, req do.TestPushChannelRequest) (string, string, string, string, error) { - if req.Name != "" { - channel, err := dao.GetPushChannelByNameRecord(ctx, req.Name) - if err != nil { - return "", "", "", "", errors.New(consts.ErrChannelNotFound) - } - return channel.URL, channel.Token, channel.Other, channel.Type, nil - } - return req.URL, req.Token, req.Other, req.Type, nil -} - -// PreparePushChannelTest builds the connectivity probe payload for a channel. -func PreparePushChannelTest(ctx context.Context, req do.TestPushChannelRequest) (do.SendPayload, error) { - url, token, other, channelType, err := LoadChannelForTest(ctx, req) - if err != nil { - return do.SendPayload{}, err - } - - if channelType == consts.ChannelEmail { - url, token, other = ResolveSMTPConfig(ctx, url, token, other) - } - - tempChannel := entity.PushChannel{ - Name: "test_temp", - URL: url, - Token: token, - Other: other, - Type: channelType, - Enabled: true, - } - if err := tempChannel.Validate(); err != nil { - return do.SendPayload{}, err - } - url = tempChannel.URL - - var config pkgpush.Config - var renderedJSON string - switch channelType { - case consts.ChannelLark: - config = pkgpush.Config{Channel: consts.ChannelLark, URL: url, Secret: token} - renderedJSON = other - case consts.ChannelEmail: - config = pkgpush.Config{Channel: consts.ChannelEmail, URL: url, Key: token, Secret: other} - case consts.ChannelTelegram: - config = pkgpush.Config{Channel: consts.ChannelTelegram, URL: url, Secret: token, Key: other} - default: - config = pkgpush.Config{Channel: consts.ChannelCustom, URL: url} - customPushReq := do.CustomPushRequest{ - Title: "通道测试通知", - Content: "这是一条来自系统的消息通道连通性测试消息。", - Description: "系统通道测试", - URL: "https://example.com", - To: req.Target, - } - renderedJSON = RenderCustomPayload(other, customPushReq) - } - - return do.SendPayload{ - EventKey: "test_channel", - Config: config, - Target: req.Target, - Body: do.NotificationMessage{ - Title: "通道测试通知", - Content: "这是一条来自系统的消息通道连通性测试消息。", - Level: consts.DefaultLevelInfo, - }, - Template: renderedJSON, - }, nil -} - -// RenderCustomPayload substitutes the supported template variables of a custom -// webhook body, JSON-escaping every injected value. -func RenderCustomPayload(template string, req do.CustomPushRequest) string { - result := template - result = strings.ReplaceAll(result, "$title", EscapeJSONString(req.Title)) - result = strings.ReplaceAll(result, "$description", EscapeJSONString(req.Description)) - result = strings.ReplaceAll(result, "$content", EscapeJSONString(req.Content)) - result = strings.ReplaceAll(result, "$url", EscapeJSONString(req.URL)) - result = strings.ReplaceAll(result, "$to", EscapeJSONString(req.To)) - return result -} - -// EscapeJSONString renders s as a JSON string body without the surrounding quotes. -func EscapeJSONString(s string) string { - b, _ := json.Marshal(s) - const minJSONLen = 2 - if len(b) >= minJSONLen { - return string(b[1 : len(b)-1]) - } - return s -} - -// ListActivePushEventsByTaskType returns enabled push events for a given task type. -func ListActivePushEventsByTaskType(ctx context.Context, taskType string) ([]entity.PushEvent, error) { - return dao.ListActivePushEventsByTaskTypeRecord(ctx, taskType) -} - -// QueryUser resolves a user through the UserService contract, falling back to the -// DAO read path while the contract is not wired yet. -func QueryUser(ctx context.Context, fromService func(contracts.UserService) (*contracts.UserDTO, error), dbField string, dbVal any) (*contracts.UserDTO, error) { - if userSvc := GetUserService(ctx); userSvc != nil { - return fromService(userSvc) - } - if user, err := dao.FindUserByFieldRecord(ctx, dbField, dbVal); err == nil && user != nil { - return user, nil - } - return nil, errors.New(consts.ErrUserNotFound) -} - -// FindUserByID resolves a user by primary key. -func FindUserByID(ctx context.Context, id uint64) (*contracts.UserDTO, error) { - return QueryUser(ctx, func(s contracts.UserService) (*contracts.UserDTO, error) { - return s.GetUserByID(ctx, id) - }, "id", id) -} - -// FindUserByUsername resolves a user by login name. -func FindUserByUsername(ctx context.Context, username string) (*contracts.UserDTO, error) { - return QueryUser(ctx, func(s contracts.UserService) (*contracts.UserDTO, error) { - return s.GetUserByUsername(ctx, username) - }, "username", username) -} - -// LoadUserFromPayload extracts user info from data. -func LoadUserFromPayload(ctx context.Context, data map[string]any) any { - if u, exists := data["user"]; exists && u != nil { - return u - } - - if userID, ok := ExtractUserID(data); ok && userID > 0 { - if user, err := FindUserByID(ctx, userID); err == nil && user != nil { - return user - } - } - - if username := ExtractUsername(data); username != "" { - if user, err := FindUserByUsername(ctx, username); err == nil && user != nil { - return user - } - } - return nil -} - -// RecordPushHistory creates a push history audit record. -func RecordPushHistory(ctx context.Context, req do.SendPayload, status, errMsg string) error { - title := req.Body.Title - content := req.Body.Content - level := req.Body.Level - if title == "" { - title = "系统通知" - } - if level == "" { - level = consts.DefaultLevelInfo - } - - target := req.Target - if target == "" { - if req.Config.URL != "" { - target = req.Config.URL - const maxTargetLen = 50 - const truncatedLen = 47 - if len(target) > maxTargetLen { - target = target[:truncatedLen] + "..." - } - } else { - target = "default" - } - } - - history := entity.PushHistory{ - EventKey: req.EventKey, - Channel: req.Config.Channel, - Target: target, - Title: title, - Content: content, - Level: level, - Status: status, - ErrorMsg: errMsg, - } - return dao.CreatePushHistoryRecord(ctx, &history) -} - -// ResolveTarget parses dynamic placeholders into concrete receiver targets. -func ResolveTarget(ctx context.Context, target string, flatBody map[string]any, channel string) string { - target = strings.TrimSpace(target) - if target == "" { - return "" - } - - resolved := ResolveDynamicKeyword(target, flatBody) - if strings.Contains(resolved, "@") { - return resolved - } - if val, matched := ResolveSystemTarget(ctx, resolved, channel); matched { - return val - } - - user, found := ResolveTargetUser(ctx, resolved, channel) - if !found { - return resolved - } - if channel == consts.ChannelEmail && user.Email != "" { - return user.Email - } - if channel != consts.ChannelEmail && user.Username != "" { - return user.Username - } - return resolved -} - -// ResolveDynamicKeyword resolves user.id, username, email keywords. -func ResolveDynamicKeyword(target string, flatBody map[string]any) string { - switch target { - case "user.id", "id": - if val, ok := flatBody["user.id"]; ok { - return fmt.Sprintf("%v", val) - } - if val, ok := flatBody["id"]; ok { - return fmt.Sprintf("%v", val) - } - case "user.username", "username": - if val, ok := flatBody["user.username"]; ok { - return fmt.Sprintf("%v", val) - } - if val, ok := flatBody["username"]; ok { - return fmt.Sprintf("%v", val) - } - case "user.email", consts.ChannelEmail: - if val, ok := flatBody["user.email"]; ok { - return fmt.Sprintf("%v", val) - } - if val, ok := flatBody["email"]; ok { - return fmt.Sprintf("%v", val) - } - } - return target -} - -// ResolveTargetUser resolves user by numeric ID or username string. -func ResolveTargetUser(ctx context.Context, resolved, _ string) (contracts.UserDTO, bool) { - if id, err := strconv.ParseUint(resolved, 10, 64); err == nil { - if u, err := FindUserByID(ctx, id); err == nil && u != nil { - return *u, true - } - } - if u, err := FindUserByUsername(ctx, resolved); err == nil && u != nil { - return *u, true - } - return contracts.UserDTO{}, false -} - -// GetFirstAdminUser resolves the first administrator through the UserService -// contract, falling back to the DAO read path when it is unavailable. -func GetFirstAdminUser(ctx context.Context) (*contracts.UserDTO, error) { - if userSvc := GetUserService(ctx); userSvc != nil { - return userSvc.GetFirstAdminUser(ctx) - } - if adminUser, err := dao.FindFirstAdminUserRecord(ctx); err == nil && adminUser != nil { - return adminUser, nil - } - return nil, errors.New(consts.ErrNoAdminUser) -} - -// ResolveSystemTarget maps system receiver aliases to administrator contact info. -func ResolveSystemTarget(ctx context.Context, resolved, channel string) (string, bool) { - if resolved != "系统" && resolved != "system" && resolved != "0" { - return "", false - } - adminUser, err := GetFirstAdminUser(ctx) - if err != nil || adminUser == nil { - return resolved, true - } - if channel == consts.ChannelEmail && adminUser.Email != "" { - return adminUser.Email, true - } - if channel != consts.ChannelEmail && adminUser.Username != "" { - return adminUser.Username, true - } - return resolved, true -} - -// ResolveSMTPConfig fills missing email endpoint fields from the system SMTP settings. -func ResolveSMTPConfig(ctx context.Context, url, token, other string) (string, string, string) { - if url != "" && token != "" { - return url, token, other - } - smtp, err := dao.LoadSMTPConfigRecord(ctx) - if err != nil { - logger.ErrorF(ctx, "[Push] 读取 SMTP 系统配置失败: %v", err) - return url, token, other - } - if smtp.Host == "" || smtp.Username == "" { - return url, token, other - } - port := smtp.Port - if port == "" { - port = "587" - } - if url == "" { - url = smtp.Host + ":" + port - } - if token == "" { - token = smtp.Username - } - if other == "" { - other = smtp.Password - } - return url, token, other -} - -// GetSystemUser gets a system user DTO. -func GetSystemUser(ctx context.Context) *contracts.UserDTO { - if adminUser, err := GetFirstAdminUser(ctx); err == nil && adminUser != nil { - return adminUser - } - return &contracts.UserDTO{ - Username: "system", - Nickname: "系统管理员", - } -} - -// FindBuiltInEvent finds a registered built-in event by key. -func FindBuiltInEvent(key string) (do.EventMetadata, bool) { - for _, meta := range GetBuiltInEvents() { - if meta.Key == key { - return meta, true - } - } - return do.EventMetadata{}, false -} - -// GetEventInfo derives the event key, display name and default template for a -// task-completion based event or a registered built-in event key. -func GetEventInfo(ctx context.Context, req do.CreatePushEventRequest) (string, string, []byte, error) { - if req.TaskType != "" { - taskName := req.TaskType - if taskSvc := GetTaskService(ctx); taskSvc != nil { - if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok { - taskName = meta.DisplayName - } - } - eventKey := "task_completed:" + req.TaskType - eventName := "任务完成: " + taskName - defaultTemplate := do.NotificationMessage{ - Title: "任务完成: " + taskName, - Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。", - Level: consts.DefaultLevelInfo, - } - defaultTemplateBytes, err := json.Marshal(defaultTemplate) - if err != nil { - return "", "", nil, err - } - return eventKey, eventName, defaultTemplateBytes, nil - } - - if req.EventKey == "" { - return "", "", nil, errors.New(consts.ErrEventKeyOrTaskType) - } - - meta, found := FindBuiltInEvent(req.EventKey) - if !found { - return "", "", nil, errors.New(consts.ErrUnsupportedEventKey) - } - - defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate) - if err != nil { - return "", "", nil, err - } - return req.EventKey, meta.Name, defaultTemplateBytes, nil -} - -// EnqueuePushTask dispatches a notification payload to the async push worker. -func EnqueuePushTask(ctx context.Context, payload do.SendPayload) error { - payloadBytes, err := json.Marshal(payload) - if err != nil { - return err - } - if taskSvc := GetTaskService(ctx); taskSvc != nil { - _, err = taskSvc.Dispatch(ctx, "send_notification", payloadBytes, contracts.TaskTriggerSystem) - return err - } - return errors.New(consts.ErrTaskServiceUnavailable) -} - -// GetFlatBody flattens nested body map. -func GetFlatBody(body map[string]any) map[string]any { - jsonBytes, err := json.Marshal(body) - if err != nil { - return body - } - var jsonMap map[string]any - if err := json.Unmarshal(jsonBytes, &jsonMap); err != nil { - return body - } - - flatResult := make(map[string]any) - FlattenMap("", jsonMap, flatResult) - return flatResult -} - -// FlattenMap recursively flattens map key-values. -func FlattenMap(prefix string, m, result map[string]any) { - for k, v := range m { - key := k - if prefix != "" { - key = prefix + "." + k - } - if nestedMap, ok := v.(map[string]any); ok { - FlattenMap(key, nestedMap, result) - } else { - result[key] = v - } - } -} - -const ( - // SendNotificationTask is the asynq task name for push notification. - SendNotificationTask = consts.SendNotificationTask - // TaskTypeSendNotification is the admin task manager type identifier. - TaskTypeSendNotification = consts.TaskTypeSendNotification -) - -// SendNotificationMeta represents the task metadata. -var SendNotificationMeta = contracts.TaskMetaDTO{ - Type: TaskTypeSendNotification, - AsynqTask: SendNotificationTask, - Name: "推送通知", - DisplayName: "推送通知", - Description: "异步执行系统通知的多渠道派发与推送", - Category: "push", - SupportsTime: false, - MaxRetry: 3, - Queue: taskQueueDefault, - Retryable: true, - Params: []contracts.TaskParamDTO{ - { - Name: "event_key", - Label: "事件标识", - Type: taskParamTypeString, - Required: true, - Placeholder: "admin_login", - Description: "事件标识 (如 admin_login)", - }, - { - Name: "target", - Label: "目标接收者", - Type: taskParamTypeString, - Required: false, - Description: "目标接收者", - }, - }, -} - -// PushHandler handles asynchronous notification sending. -type PushHandler struct{} - -// ValidatePayload validates and normalizes push parameters. -func (h *PushHandler) ValidatePayload(payload []byte) ([]byte, error) { - if len(payload) == 0 { - return nil, errors.New(consts.ErrPayloadRequired) - } - - var req do.SendPayload - if err := json.Unmarshal(payload, &req); err != nil { - return nil, fmt.Errorf("%s: %w", consts.ErrInvalidJSONFormat, err) - } - - if req.Config.Channel == "" { - return nil, errors.New(consts.ErrChannelTypeRequired) - } - - return json.Marshal(req) -} - -// Execute performs the push send and logs delivery history audit. -func (h *PushHandler) Execute(ctx context.Context, payload []byte) error { - var req do.SendPayload - if err := json.Unmarshal(payload, &req); err != nil { - logger.ErrorF(ctx, "[Push] 解析推送参数失败: %v", err) - return fmt.Errorf("%s: %w", consts.ErrParsePayloadFailed, err) - } - - logger.InfoF(ctx, "[Push] 开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target) - - pusher, err := pkgpush.GetPusher(req.Config.Channel) - if err != nil { - errWrap := fmt.Errorf("%s: %w", consts.ErrGetPusherFailed, err) - logger.ErrorF(ctx, "[Push] 推送失败: %v", errWrap) - h.recordHistory(ctx, req, "failed", errWrap.Error()) - return errWrap - } - - flatBody := req.Body.Flatten() - upstreamResp, err := pusher.Send(ctx, req.Config, req.Target, flatBody, req.Template, nil) - - title := req.Body.Title - content := req.Body.Content - - if err != nil { - logger.ErrorF(ctx, "[Push] 消息推送失败 (标题: %s): %v, 上游返回: %s", title, err, upstreamResp) - h.recordHistory(ctx, req, "failed", err.Error()) - return fmt.Errorf("pusher.Send failed: %w", err) - } - - logger.InfoF(ctx, "[Push] 消息推送成功 (标题: %s, 内容摘要: %s), 上游返回: %s", title, content, upstreamResp) - h.recordHistory(ctx, req, "success", "") - - return nil -} - -func (h *PushHandler) recordHistory(ctx context.Context, req do.SendPayload, status, errMsg string) { - if dbErr := RecordPushHistory(ctx, req, status, errMsg); dbErr != nil { - logger.ErrorF(ctx, "[Push] 写入推送历史审计记录失败: %v", dbErr) - } -} - -// HandleTaskCompleted handles task completion notifications. -func HandleTaskCompleted(ctx context.Context, e contracts.TaskCompletedEvent) { - events, err := ListActivePushEventsByTaskType(ctx, e.TaskType) - if err != nil { - logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", e.TaskType, err) - return - } - if len(events) == 0 { - return - } - - body := map[string]any{ - "task_id": e.TaskID, - "task_name": e.TaskName, - "task_type": e.TaskType, - "task_status": e.Status, - "task_duration": e.Duration, - "time": time.Now().Format("2006-01-02 15:04:05"), - "task_error": e.ErrorMsg, - "task_result": e.ResultMsg, - } - - var payloadMap map[string]any - if e.Payload != "" { - if err := json.Unmarshal([]byte(e.Payload), &payloadMap); err == nil { - body["payload"] = payloadMap - ExtractUserFromMap(ctx, payloadMap, body) - } - } - if e.Detail != "" { - var detailMap map[string]any - if err := json.Unmarshal([]byte(e.Detail), &detailMap); err == nil { - body["detail"] = detailMap - ExtractUserFromMap(ctx, detailMap, body) - } - } - - for _, event := range events { - meta := do.EventMetadata{ - Key: event.EventKey, - Name: event.Name, - Description: "异步任务执行完毕触发的自动通知", - } - DefaultTrigger.Trigger(ctx, meta, body) - } -} - -// ExtractUserFromMap extracts user info from payload/detail into body map. -func ExtractUserFromMap(ctx context.Context, data, body map[string]any) { - if u, exists := body["user"]; exists && u != nil { - return - } - if user := LoadUserFromPayload(ctx, data); user != nil { - body["user"] = user - } -} - -// ExtractUserID extracts a user ID from map keys. -func ExtractUserID(data map[string]any) (uint64, bool) { - for _, k := range []string{"user_id", "userId", "uid"} { - val, ok := data[k] - if !ok || val == nil { - continue - } - switch v := val.(type) { - case float64: - if v >= 0 { - return uint64(v), true - } - case int: - if v >= 0 { - return uint64(v), true - } - case int64: - if v >= 0 { - return uint64(v), true - } - case uint64: - return v, true - case string: - if id, err := strconv.ParseUint(v, 10, 64); err == nil { - return id, true - } - } - } - return 0, false -} - -// ExtractUsername extracts a username string from map keys. -func ExtractUsername(data map[string]any) string { - for _, k := range []string{"username", "user_name"} { - if val, ok := data[k]; ok && val != nil { - if s, ok := val.(string); ok && s != "" { - return s - } - } - } - return "" -} - -// EventTrigger represents the unified event trigger class. -type EventTrigger struct{} - -// DefaultTrigger is the singleton instance of EventTrigger. -var DefaultTrigger = &EventTrigger{} - -// Trigger receives event metadata and processes the event notification dispatch asynchronously. -func (t *EventTrigger) Trigger(ctx context.Context, meta do.EventMetadata, body map[string]any) { - asyncCtx := context.WithoutCancel(ctx) - util.Go(func() { - if body == nil { - body = make(map[string]any) - } - if _, hasUser := body["user"]; !hasUser || body["user"] == nil { - body["user"] = GetSystemUser(asyncCtx) - } - - eventPtr, err := dao.GetActivePushEventByKey(asyncCtx, meta.Key) - if err != nil { - if errors.Is(err, consts.ErrRecordNotFound) { - return - } - logger.ErrorF(asyncCtx, "push_event_trigger: failed to get active event %s: %v", meta.Key, err) - return - } - event := *eventPtr - if len(event.Channels) == 0 { - return - } - - flatBody := GetFlatBody(body) - msg, _ := t.buildMessage(&event, meta, flatBody, body) - t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody) - }) -} - -func (t *EventTrigger) buildMessage(event *entity.PushEvent, meta do.EventMetadata, flatBody, body map[string]any) (do.NotificationMessage, string) { - var msg do.NotificationMessage - renderedTemplate := "" - - templateSource := event.Template - if templateSource != "" { - var err error - msg, renderedTemplate, err = t.parseCustomTemplate(event, templateSource, flatBody) - if err != nil { - msg.Title = event.Name - msg.Content = renderedTemplate - msg.Level = consts.DefaultLevelInfo - } - } else { - msg = t.parseDefaultTemplate(meta, flatBody) - } - - if msg.Ext == nil { - msg.Ext = make(map[string]any) - } - for k, v := range body { - if k == consts.KeyTitle || k == consts.KeyContent || k == consts.KeyLevel { - continue - } - if _, exists := msg.Ext[k]; !exists { - msg.Ext[k] = v - } - } - - return msg, renderedTemplate -} - -func (t *EventTrigger) parseCustomTemplate(event *entity.PushEvent, templateSource string, flatBody map[string]any) (do.NotificationMessage, string, error) { - var msg do.NotificationMessage - renderedTemplate := pkgpush.ParseTemplate(templateSource, flatBody) - - var tMap map[string]any - if err := json.Unmarshal([]byte(renderedTemplate), &tMap); err != nil { - return msg, renderedTemplate, err - } - - if title, ok := tMap[consts.KeyTitle].(string); ok && title != "" { - msg.Title = title - } else { - msg.Title = event.Name - } - delete(tMap, consts.KeyTitle) - - if content, ok := tMap[consts.KeyContent].(string); ok && content != "" { - msg.Content = content - } else { - msg.Content = renderedTemplate - } - delete(tMap, consts.KeyContent) - - if level, ok := tMap[consts.KeyLevel].(string); ok && level != "" { - msg.Level = level - } else { - msg.Level = consts.DefaultLevelInfo - } - delete(tMap, consts.KeyLevel) - - msg.Ext = tMap - return msg, renderedTemplate, nil -} - -func (t *EventTrigger) parseDefaultTemplate(meta do.EventMetadata, flatBody map[string]any) do.NotificationMessage { - var msg do.NotificationMessage - msg.Title = pkgpush.ParseTemplate(meta.DefaultTemplate.Title, flatBody) - msg.Content = pkgpush.ParseTemplate(meta.DefaultTemplate.Content, flatBody) - msg.Level = pkgpush.ParseTemplate(meta.DefaultTemplate.Level, flatBody) - - if meta.DefaultTemplate.Ext != nil { - msg.Ext = make(map[string]any) - for k, v := range meta.DefaultTemplate.Ext { - if strVal, ok := v.(string); ok { - msg.Ext[k] = pkgpush.ParseTemplate(strVal, flatBody) - } else { - msg.Ext[k] = v - } - } - } - return msg -} - -func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta do.EventMetadata, event *entity.PushEvent, msg do.NotificationMessage, flatBody map[string]any) { - for _, channelName := range event.Channels { - customChannel, err := dao.GetActivePushChannelByName(ctx, channelName) - if err == nil { - t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody) - continue - } - logger.WarnF(ctx, "push_event_trigger: channel %q not found in DB or disabled: %v", channelName, err) - } -} - -func (t *EventTrigger) enqueueCustomPushChannelTasks(ctx context.Context, meta do.EventMetadata, event *entity.PushEvent, channel *entity.PushChannel, msg do.NotificationMessage, flatBody map[string]any) { - if len(event.Targets) == 0 { - t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, "", msg) - return - } - - for _, target := range event.Targets { - resolvedTarget := ResolveTarget(ctx, target, flatBody, channel.Name) - t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, resolvedTarget, msg) - } -} - -func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, meta do.EventMetadata, channel *entity.PushChannel, target string, msg do.NotificationMessage) { - var config pkgpush.Config - var renderedTemplate string - - switch channel.Type { - case consts.ChannelLark: - config = pkgpush.Config{Channel: consts.ChannelLark, URL: channel.URL, Secret: channel.Token} - renderedTemplate = channel.Other - case consts.ChannelEmail: - url, token, other := ResolveSMTPConfig(ctx, channel.URL, channel.Token, channel.Other) - config = pkgpush.Config{Channel: consts.ChannelEmail, URL: url, Key: token, Secret: other} - case consts.ChannelTelegram: - config = pkgpush.Config{Channel: consts.ChannelTelegram, URL: channel.URL, Secret: channel.Token, Key: channel.Other} - default: - config = pkgpush.Config{Channel: consts.ChannelCustom, URL: channel.URL} - customPushReq := do.CustomPushRequest{ - Title: msg.Title, - Content: msg.Content, - Description: meta.Description, - To: target, - } - if urlVal, ok := msg.Ext["url"].(string); ok { - customPushReq.URL = urlVal - } - renderedTemplate = RenderCustomPayload(channel.Other, customPushReq) - } - - payload := do.SendPayload{ - EventKey: meta.Key, - Config: config, - Target: target, - Body: msg, - Template: renderedTemplate, - } - if err := EnqueuePushTask(ctx, payload); err != nil { - logger.ErrorF(ctx, "push_event_trigger: enqueuePushTask failed for %s channel %s -> %s: %v", channel.Type, channel.Name, target, err) - } -} - -// SyncEvents automatically registers/updates built-in events in the database. -func SyncEvents(ctx context.Context) error { - return SyncBuiltInEvents(ctx) -} - -// AdminLogin is the metadata definition for the admin login event. -var AdminLogin = do.EventMetadata{ - Key: "admin_login", - Name: "管理员登录", - DefaultTemplate: do.NotificationMessage{ - Title: "管理员登录提醒", - Content: "管理员 {{user.username}} 于 {{time}} 从 IP {{ip}} 登录系统。", - Level: "INFO", - }, - Description: "当管理员成功登录系统时触发此通知", -} - -// HandleAdminLoggedIn 处理管理员登录事件并触发通知 -func HandleAdminLoggedIn(ctx context.Context, event contracts.AdminLoggedIn) { - if event.User == nil { - return - } - - body := map[string]any{ - "user": event.User, - "ip": event.IP, - "time": time.Now().Format("2006-01-02 15:04:05"), - } - DefaultTrigger.Trigger(ctx, AdminLogin, body) -} - -// RegisterCustomEvents registers default domain push notification events. -func RegisterCustomEvents() { - RegisterBuiltInEvent(AdminLogin) -} diff --git a/backend/plugins/domain/msg_gateway/service/push_channel.go b/backend/plugins/domain/msg_gateway/service/push_channel.go new file mode 100644 index 00000000..23a407cf --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/push_channel.go @@ -0,0 +1,168 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" + pkgpush "Wavelet/plugins/domain/msg_gateway/push" + "context" + "errors" + "fmt" +) + +// ListPushChannels returns every configured push channel. +func ListPushChannels(ctx context.Context) ([]entity.PushChannel, error) { + return dao.ListPushChannelsRecord(ctx) +} + +// CreatePushChannel validates uniqueness and persists a new push channel. +func CreatePushChannel(ctx context.Context, req do.CreatePushChannelRequest) (entity.PushChannel, error) { + count, err := dao.CountPushChannelsByNameRecord(ctx, req.Name) + if err != nil { + return entity.PushChannel{}, err + } + if count > 0 { + return entity.PushChannel{}, errors.New(consts.ErrChannelNameExists) + } + + channel := entity.PushChannel{ + Name: req.Name, + Description: req.Description, + Type: req.Type, + Token: req.Token, + URL: req.URL, + Other: req.Other, + Enabled: req.Enabled, + } + if err := channel.Validate(); err != nil { + return entity.PushChannel{}, err + } + if err := dao.CreatePushChannelRecord(ctx, &channel); err != nil { + return entity.PushChannel{}, err + } + return channel, nil +} + +// UpdatePushChannel replaces the mutable fields of an existing push channel. +func UpdatePushChannel(ctx context.Context, id uint64, req do.UpdatePushChannelRequest) (entity.PushChannel, error) { + channel, err := dao.GetPushChannelByIDRecord(ctx, id) + if err != nil { + return entity.PushChannel{}, err + } + + channel.Description = req.Description + channel.Type = req.Type + channel.Token = req.Token + channel.URL = req.URL + channel.Other = req.Other + channel.Enabled = req.Enabled + if err := channel.Validate(); err != nil { + return entity.PushChannel{}, err + } + if err := dao.SavePushChannelRecord(ctx, &channel); err != nil { + return entity.PushChannel{}, err + } + return channel, nil +} + +// DeletePushChannel removes a push channel by id. +func DeletePushChannel(ctx context.Context, id uint64) error { + channel, err := dao.GetPushChannelByIDRecord(ctx, id) + if err != nil { + return err + } + return dao.DeletePushChannelRecord(ctx, &channel) +} + +// LoadChannelForTest resolves the credentials under test, either from a stored +// channel name or from the ad-hoc values sent by the caller. +func LoadChannelForTest(ctx context.Context, req do.TestPushChannelRequest) (string, string, string, string, error) { + if req.Name != "" { + channel, err := dao.GetPushChannelByNameRecord(ctx, req.Name) + if err != nil { + return "", "", "", "", errors.New(consts.ErrChannelNotFoundText) + } + return channel.URL, channel.Token, channel.Other, channel.Type, nil + } + return req.URL, req.Token, req.Other, req.Type, nil +} + +// PreparePushChannelTest builds the connectivity probe payload for a channel. +func PreparePushChannelTest(ctx context.Context, req do.TestPushChannelRequest) (do.SendPayload, error) { + url, token, other, channelType, err := LoadChannelForTest(ctx, req) + if err != nil { + return do.SendPayload{}, err + } + + tempChannel := entity.PushChannel{ + Name: "test_temp", + URL: url, + Token: token, + Other: other, + Type: channelType, + Enabled: true, + } + if err := tempChannel.Validate(); err != nil { + return do.SendPayload{}, err + } + url = tempChannel.URL + + var config pkgpush.Config + var renderedJSON string + switch channelType { + case consts.ChannelLark: + config = pkgpush.Config{Channel: consts.ChannelLark, URL: url, Secret: token} + renderedJSON = other + case consts.ChannelEmail: + config = pkgpush.Config{Channel: consts.ChannelEmail, URL: url, Key: token, Secret: other} + case consts.ChannelTelegram: + config = pkgpush.Config{Channel: consts.ChannelTelegram, URL: url, Secret: token, Key: other} + default: + config = pkgpush.Config{Channel: consts.ChannelCustom, URL: url} + customPushReq := do.CustomPushRequest{ + Title: "通道测试通知", + Content: "这是一条来自系统的消息通道连通性测试消息。", + Description: "系统通道测试", + URL: "https://example.com", + To: req.Target, + } + renderedJSON = RenderCustomPayload(other, customPushReq) + } + + return do.SendPayload{ + EventKey: "test_channel", + Config: config, + Target: req.Target, + Body: do.NotificationMessage{ + Title: "通道测试通知", + Content: "这是一条来自系统的消息通道连通性测试消息。", + Level: consts.DefaultLevelInfo, + }, + Template: renderedJSON, + }, nil +} + +// RunPushTest validates an ad-hoc channel config and sends a connectivity probe. +func RunPushTest(ctx context.Context, cfg pkgpush.Config, target string) error { + pusher, err := pkgpush.GetPusher(cfg.Channel) + if err != nil { + return err + } + if err := pusher.ValidateConfig(cfg); err != nil { + return fmt.Errorf("%s: %w", consts.ErrValidationFailed, err) + } + + testBody := map[string]any{ + consts.KeyTitle: "测试通道推送", + consts.KeyContent: "当您收到这条消息,说明当前渠道连通性测试通过。", + consts.KeyLevel: consts.DefaultLevelInfo, + } + if _, err := pusher.Send(ctx, cfg, target, testBody, "", nil); err != nil { + return err + } + return nil +} diff --git a/backend/plugins/domain/msg_gateway/service/push_event.go b/backend/plugins/domain/msg_gateway/service/push_event.go new file mode 100644 index 00000000..90173ab3 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/push_event.go @@ -0,0 +1,339 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/logger" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" + "context" + "encoding/json" + "errors" + "strings" + "sync" + "time" +) + +var ( + builtInEventsMu sync.RWMutex + // BuiltInEvents lists all built-in events defined across the domain. + BuiltInEvents []do.EventMetadata +) + +// RegisterBuiltInEvent registers a built-in event definition. +func RegisterBuiltInEvent(meta do.EventMetadata) { + builtInEventsMu.Lock() + defer builtInEventsMu.Unlock() + for i, e := range BuiltInEvents { + if e.Key == meta.Key { + BuiltInEvents[i] = meta + return + } + } + BuiltInEvents = append(BuiltInEvents, meta) +} + +// GetBuiltInEvents returns a copy of registered built-in events. +func GetBuiltInEvents() []do.EventMetadata { + builtInEventsMu.RLock() + defer builtInEventsMu.RUnlock() + out := make([]do.EventMetadata, len(BuiltInEvents)) + copy(out, BuiltInEvents) + return out +} + +// FindBuiltInEvent finds a registered built-in event by key. +func FindBuiltInEvent(key string) (do.EventMetadata, bool) { + for _, meta := range GetBuiltInEvents() { + if meta.Key == key { + return meta, true + } + } + return do.EventMetadata{}, false +} + +// PushRegistryAdapter adapts contracts.PushRegistry onto the built-in event store. +type PushRegistryAdapter struct{} + +// RegisterBuiltInEvent records a built-in push event definition from cross-plugin contract. +func (PushRegistryAdapter) RegisterBuiltInEvent(meta contracts.PushEventMeta) { + RegisterBuiltInEvent(eventMetadataFromContract(meta)) +} + +// SyncEvents persists registered built-in events into the database. +func (PushRegistryAdapter) SyncEvents(ctx context.Context) error { + return SyncEvents(ctx) +} + +func eventMetadataFromContract(meta contracts.PushEventMeta) do.EventMetadata { + return do.EventMetadata{ + Key: meta.Key, + Name: meta.Name, + Description: meta.Description, + DefaultTemplate: do.NotificationMessage{ + Title: meta.DefaultTemplate.Title, + Content: meta.DefaultTemplate.Content, + Level: meta.DefaultTemplate.Level, + Ext: meta.DefaultTemplate.Ext, + }, + } +} + +// SyncBuiltInEvents seeds a database row for every registered built-in event. +func SyncBuiltInEvents(ctx context.Context) error { + for _, meta := range GetBuiltInEvents() { + _, err := dao.GetPushEventByKeyRecord(ctx, meta.Key) + if errors.Is(err, consts.ErrRecordNotFound) { + var defaultTemplateStr string + if defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate); err == nil { + defaultTemplateStr = string(defaultTemplateBytes) + } + event := entity.PushEvent{ + EventKey: meta.Key, + Name: meta.Name, + Channels: []string{}, + Targets: []string{}, + Template: defaultTemplateStr, + Enabled: false, + } + if err := dao.CreatePushEventRecord(ctx, &event); err != nil { + return err + } + } else if err != nil { + return err + } + } + return nil +} + +// SyncEvents automatically registers/updates built-in events in the database. +func SyncEvents(ctx context.Context) error { + return SyncBuiltInEvents(ctx) +} + +// ListPushEvents lists all configured push events. +func ListPushEvents(ctx context.Context) ([]entity.PushEvent, error) { + return dao.ListPushEventsRecord(ctx) +} + +// CreatePushEvent stores a push event configuration for a built-in event or task type. +func CreatePushEvent(ctx context.Context, req do.CreatePushEventRequest) (entity.PushEvent, error) { + eventKey, eventName, defaultTemplateBytes, err := GetEventInfo(ctx, req) + if err != nil { + return entity.PushEvent{}, err + } + + count, err := dao.CountPushEventsByKeyRecord(ctx, eventKey) + if err != nil { + return entity.PushEvent{}, err + } + if count > 0 { + return entity.PushEvent{}, errors.New(consts.ErrEventAlreadyConfigured) + } + + templateStr := strings.TrimSpace(req.Template) + if templateStr == "" { + templateStr = string(defaultTemplateBytes) + } else { + var tempMap map[string]any + if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil { + return entity.PushEvent{}, errors.New(consts.ErrTemplateInvalidJSON) + } + } + + channels := req.Channels + if channels == nil { + channels = []string{} + } + targets := req.Targets + if targets == nil { + targets = []string{} + } + + event := entity.PushEvent{ + EventKey: eventKey, + Name: eventName, + TaskType: req.TaskType, + Channels: channels, + Targets: targets, + Template: templateStr, + Enabled: req.Enabled, + } + if err := event.Validate(); err != nil { + return entity.PushEvent{}, err + } + if err := dao.CreatePushEventRecord(ctx, &event); err != nil { + return entity.PushEvent{}, err + } + return event, nil +} + +// DeletePushEvent deletes a push event configuration by id. +func DeletePushEvent(ctx context.Context, id uint64) error { + event, err := dao.GetPushEventByIDRecord(ctx, id) + if err != nil { + return err + } + return dao.DeletePushEventRecord(ctx, &event) +} + +// UpdatePushEvent replaces mutable push event fields. +func UpdatePushEvent(ctx context.Context, id uint64, req do.UpdatePushEventRequest) error { + event, err := dao.GetPushEventByIDRecord(ctx, id) + if err != nil { + return err + } + + event.Channels = req.Channels + event.Targets = req.Targets + event.Template = req.Template + event.Enabled = req.Enabled + if err := event.Validate(); err != nil { + return err + } + return dao.SavePushEventRecord(ctx, &event) +} + +// TogglePushEvent flips the enabled flag of a push event. +func TogglePushEvent(ctx context.Context, id uint64) (bool, error) { + event, err := dao.GetPushEventByIDRecord(ctx, id) + if err != nil { + return false, err + } + + enabled := !event.Enabled + if enabled && len(event.Channels) == 0 { + return false, errors.New(consts.ErrEnableWithoutChannels) + } + if err := dao.UpdatePushEventEnabledRecord(ctx, &event, enabled); err != nil { + return false, err + } + return enabled, nil +} + +// ListActivePushEventsByTaskType returns enabled push events for a given task type. +func ListActivePushEventsByTaskType(ctx context.Context, taskType string) ([]entity.PushEvent, error) { + return dao.ListActivePushEventsByTaskTypeRecord(ctx, taskType) +} + +// GetEventInfo derives the event key, display name and default template for a +// task-completion based event or a registered built-in event key. +func GetEventInfo(ctx context.Context, req do.CreatePushEventRequest) (string, string, []byte, error) { + if req.TaskType != "" { + taskName := req.TaskType + if taskSvc := GetTaskService(ctx); taskSvc != nil { + if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok { + taskName = meta.DisplayName + } + } + eventKey := "task_completed:" + req.TaskType + eventName := "任务完成: " + taskName + defaultTemplate := do.NotificationMessage{ + Title: "任务完成: " + taskName, + Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。", + Level: consts.DefaultLevelInfo, + } + defaultTemplateBytes, err := json.Marshal(defaultTemplate) + if err != nil { + return "", "", nil, err + } + return eventKey, eventName, defaultTemplateBytes, nil + } + + if req.EventKey == "" { + return "", "", nil, errors.New(consts.ErrEventKeyOrTaskType) + } + + meta, found := FindBuiltInEvent(req.EventKey) + if !found { + return "", "", nil, errors.New(consts.ErrUnsupportedEventKey) + } + + defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate) + if err != nil { + return "", "", nil, err + } + return req.EventKey, meta.Name, defaultTemplateBytes, nil +} + +// AdminLogin is the metadata definition for the admin login event. +var AdminLogin = do.EventMetadata{ + Key: "admin_login", + Name: "管理员登录", + DefaultTemplate: do.NotificationMessage{ + Title: "管理员登录提醒", + Content: "管理员 {{user.username}} 于 {{time}} 从 IP {{ip}} 登录系统。", + Level: consts.DefaultLevelInfo, + }, + Description: "当管理员成功登录系统时触发此通知", +} + +// HandleAdminLoggedIn 处理管理员登录事件并触发通知 +func HandleAdminLoggedIn(ctx context.Context, event contracts.AdminLoggedIn) { + if event.User == nil { + return + } + + body := map[string]any{ + "user": event.User, + "ip": event.IP, + "time": time.Now().Format("2006-01-02 15:04:05"), + } + DefaultTrigger.Trigger(ctx, AdminLogin, body) +} + +// HandleTaskCompleted handles task completion notifications. +func HandleTaskCompleted(ctx context.Context, e contracts.TaskCompletedEvent) { + events, err := ListActivePushEventsByTaskType(ctx, e.TaskType) + if err != nil { + logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", e.TaskType, err) + return + } + if len(events) == 0 { + return + } + + body := map[string]any{ + "task_id": e.TaskID, + "task_name": e.TaskName, + "task_type": e.TaskType, + "task_status": e.Status, + "task_duration": e.Duration, + "time": time.Now().Format("2006-01-02 15:04:05"), + "task_error": e.ErrorMsg, + "task_result": e.ResultMsg, + } + + var payloadMap map[string]any + if e.Payload != "" { + if err := json.Unmarshal([]byte(e.Payload), &payloadMap); err == nil { + body["payload"] = payloadMap + ExtractUserFromMap(ctx, payloadMap, body) + } + } + if e.Detail != "" { + var detailMap map[string]any + if err := json.Unmarshal([]byte(e.Detail), &detailMap); err == nil { + body["detail"] = detailMap + ExtractUserFromMap(ctx, detailMap, body) + } + } + + for _, event := range events { + meta := do.EventMetadata{ + Key: event.EventKey, + Name: event.Name, + Description: "异步任务执行完毕触发的自动通知", + } + DefaultTrigger.Trigger(ctx, meta, body) + } +} + +// RegisterCustomEvents registers default domain push notification events. +func RegisterCustomEvents() { + RegisterBuiltInEvent(AdminLogin) +} diff --git a/backend/plugins/domain/msg_gateway/service/push_template.go b/backend/plugins/domain/msg_gateway/service/push_template.go new file mode 100644 index 00000000..34330ce7 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/push_template.go @@ -0,0 +1,63 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/plugins/domain/msg_gateway/model/do" + "encoding/json" + "strings" +) + +// GetFlatBody flattens nested body map into dot-separated key-value map. +func GetFlatBody(body map[string]any) map[string]any { + jsonBytes, err := json.Marshal(body) + if err != nil { + return body + } + var jsonMap map[string]any + if err := json.Unmarshal(jsonBytes, &jsonMap); err != nil { + return body + } + + flatResult := make(map[string]any) + FlattenMap("", jsonMap, flatResult) + return flatResult +} + +// FlattenMap recursively flattens map key-values with dot notation. +func FlattenMap(prefix string, m, result map[string]any) { + for k, v := range m { + key := k + if prefix != "" { + key = prefix + "." + k + } + if nestedMap, ok := v.(map[string]any); ok { + FlattenMap(key, nestedMap, result) + } else { + result[key] = v + } + } +} + +// RenderCustomPayload substitutes the supported template variables of a custom +// webhook body, JSON-escaping every injected value. +func RenderCustomPayload(template string, req do.CustomPushRequest) string { + result := template + result = strings.ReplaceAll(result, "$title", EscapeJSONString(req.Title)) + result = strings.ReplaceAll(result, "$description", EscapeJSONString(req.Description)) + result = strings.ReplaceAll(result, "$content", EscapeJSONString(req.Content)) + result = strings.ReplaceAll(result, "$url", EscapeJSONString(req.URL)) + result = strings.ReplaceAll(result, "$to", EscapeJSONString(req.To)) + return result +} + +// EscapeJSONString renders s as a JSON string body without the surrounding quotes. +func EscapeJSONString(s string) string { + b, _ := json.Marshal(s) + const minJSONLen = 2 + if len(b) >= minJSONLen { + return string(b[1 : len(b)-1]) + } + return s +} diff --git a/backend/plugins/domain/msg_gateway/service/push_trigger.go b/backend/plugins/domain/msg_gateway/service/push_trigger.go new file mode 100644 index 00000000..f3904b70 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/push_trigger.go @@ -0,0 +1,398 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/logger" + "Wavelet/pkg/util" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" + pkgpush "Wavelet/plugins/domain/msg_gateway/push" + "context" + "encoding/json" + "errors" + "fmt" + "strconv" + "strings" +) + +// EventTrigger represents the unified event trigger class. +type EventTrigger struct{} + +// DefaultTrigger is the singleton instance of EventTrigger. +var DefaultTrigger = &EventTrigger{} + +// Trigger receives event metadata and processes the event notification dispatch asynchronously. +func (t *EventTrigger) Trigger(ctx context.Context, meta do.EventMetadata, body map[string]any) { + asyncCtx := context.WithoutCancel(ctx) + util.Go(func() { + if body == nil { + body = make(map[string]any) + } + if _, hasUser := body["user"]; !hasUser || body["user"] == nil { + body["user"] = GetSystemUser(asyncCtx) + } + + eventPtr, err := dao.GetActivePushEventByKey(asyncCtx, meta.Key) + if err != nil { + if errors.Is(err, consts.ErrRecordNotFound) { + return + } + logger.ErrorF(asyncCtx, "push_event_trigger: failed to get active event %s: %v", meta.Key, err) + return + } + event := *eventPtr + if len(event.Channels) == 0 { + return + } + + flatBody := GetFlatBody(body) + msg, _ := t.buildMessage(&event, meta, flatBody, body) + t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody) + }) +} + +func (t *EventTrigger) buildMessage(event *entity.PushEvent, meta do.EventMetadata, flatBody, body map[string]any) (do.NotificationMessage, string) { + var msg do.NotificationMessage + renderedTemplate := "" + + templateSource := event.Template + if templateSource != "" { + var err error + msg, renderedTemplate, err = t.parseCustomTemplate(event, templateSource, flatBody) + if err != nil { + msg.Title = event.Name + msg.Content = renderedTemplate + msg.Level = consts.DefaultLevelInfo + } + } else { + msg = t.parseDefaultTemplate(meta, flatBody) + } + + if msg.Ext == nil { + msg.Ext = make(map[string]any) + } + for k, v := range body { + if k == consts.KeyTitle || k == consts.KeyContent || k == consts.KeyLevel { + continue + } + if _, exists := msg.Ext[k]; !exists { + msg.Ext[k] = v + } + } + + return msg, renderedTemplate +} + +func (t *EventTrigger) parseCustomTemplate(event *entity.PushEvent, templateSource string, flatBody map[string]any) (do.NotificationMessage, string, error) { + var msg do.NotificationMessage + renderedTemplate := pkgpush.ParseTemplate(templateSource, flatBody) + + var tMap map[string]any + if err := json.Unmarshal([]byte(renderedTemplate), &tMap); err != nil { + return msg, renderedTemplate, err + } + + if title, ok := tMap[consts.KeyTitle].(string); ok && title != "" { + msg.Title = title + } else { + msg.Title = event.Name + } + delete(tMap, consts.KeyTitle) + + if content, ok := tMap[consts.KeyContent].(string); ok && content != "" { + msg.Content = content + } else { + msg.Content = renderedTemplate + } + delete(tMap, consts.KeyContent) + + if level, ok := tMap[consts.KeyLevel].(string); ok && level != "" { + msg.Level = level + } else { + msg.Level = consts.DefaultLevelInfo + } + delete(tMap, consts.KeyLevel) + + msg.Ext = tMap + return msg, renderedTemplate, nil +} + +func (t *EventTrigger) parseDefaultTemplate(meta do.EventMetadata, flatBody map[string]any) do.NotificationMessage { + var msg do.NotificationMessage + msg.Title = pkgpush.ParseTemplate(meta.DefaultTemplate.Title, flatBody) + msg.Content = pkgpush.ParseTemplate(meta.DefaultTemplate.Content, flatBody) + msg.Level = pkgpush.ParseTemplate(meta.DefaultTemplate.Level, flatBody) + + if meta.DefaultTemplate.Ext != nil { + msg.Ext = make(map[string]any) + for k, v := range meta.DefaultTemplate.Ext { + if strVal, ok := v.(string); ok { + msg.Ext[k] = pkgpush.ParseTemplate(strVal, flatBody) + } else { + msg.Ext[k] = v + } + } + } + return msg +} + +func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta do.EventMetadata, event *entity.PushEvent, msg do.NotificationMessage, flatBody map[string]any) { + for _, channelName := range event.Channels { + customChannel, err := dao.GetActivePushChannelByName(ctx, channelName) + if err == nil { + t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody) + continue + } + logger.WarnF(ctx, "push_event_trigger: channel %q not found in DB or disabled: %v", channelName, err) + } +} + +func (t *EventTrigger) enqueueCustomPushChannelTasks(ctx context.Context, meta do.EventMetadata, event *entity.PushEvent, channel *entity.PushChannel, msg do.NotificationMessage, flatBody map[string]any) { + if len(event.Targets) == 0 { + t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, "", msg) + return + } + + for _, target := range event.Targets { + resolvedTarget := ResolveTarget(ctx, target, flatBody, channel.Name) + t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, resolvedTarget, msg) + } +} + +func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, meta do.EventMetadata, channel *entity.PushChannel, target string, msg do.NotificationMessage) { + var config pkgpush.Config + var renderedTemplate string + + switch channel.Type { + case consts.ChannelLark: + config = pkgpush.Config{Channel: consts.ChannelLark, URL: channel.URL, Secret: channel.Token} + renderedTemplate = channel.Other + case consts.ChannelEmail: + config = pkgpush.Config{Channel: consts.ChannelEmail, URL: channel.URL, Key: channel.Token, Secret: channel.Other} + case consts.ChannelTelegram: + config = pkgpush.Config{Channel: consts.ChannelTelegram, URL: channel.URL, Secret: channel.Token, Key: channel.Other} + default: + config = pkgpush.Config{Channel: consts.ChannelCustom, URL: channel.URL} + customPushReq := do.CustomPushRequest{ + Title: msg.Title, + Content: msg.Content, + Description: meta.Description, + To: target, + } + if urlVal, ok := msg.Ext["url"].(string); ok { + customPushReq.URL = urlVal + } + renderedTemplate = RenderCustomPayload(channel.Other, customPushReq) + } + + payload := do.SendPayload{ + EventKey: meta.Key, + Config: config, + Target: target, + Body: msg, + Template: renderedTemplate, + } + if err := EnqueuePushTask(ctx, payload); err != nil { + logger.ErrorF(ctx, "push_event_trigger: enqueuePushTask failed for %s channel %s -> %s: %v", channel.Type, channel.Name, target, err) + } +} + +// ResolveTarget parses dynamic placeholders into concrete receiver targets. +func ResolveTarget(ctx context.Context, target string, flatBody map[string]any, channel string) string { + target = strings.TrimSpace(target) + if target == "" { + return "" + } + + resolved := ResolveDynamicKeyword(target, flatBody) + if strings.Contains(resolved, "@") { + return resolved + } + if val, matched := ResolveSystemTarget(ctx, resolved, channel); matched { + return val + } + + user, found := ResolveTargetUser(ctx, resolved, channel) + if !found { + return resolved + } + if channel == consts.ChannelEmail && user.Email != "" { + return user.Email + } + if channel != consts.ChannelEmail && user.Username != "" { + return user.Username + } + return resolved +} + +// ResolveDynamicKeyword resolves user.id, username, email keywords. +func ResolveDynamicKeyword(target string, flatBody map[string]any) string { + switch target { + case "user.id", "id": + if val, ok := flatBody["user.id"]; ok { + return fmt.Sprintf("%v", val) + } + if val, ok := flatBody["id"]; ok { + return fmt.Sprintf("%v", val) + } + case "user.username", "username": + if val, ok := flatBody["user.username"]; ok { + return fmt.Sprintf("%v", val) + } + if val, ok := flatBody["username"]; ok { + return fmt.Sprintf("%v", val) + } + case "user.email", consts.ChannelEmail: + if val, ok := flatBody["user.email"]; ok { + return fmt.Sprintf("%v", val) + } + if val, ok := flatBody["email"]; ok { + return fmt.Sprintf("%v", val) + } + } + return target +} + +// ResolveTargetUser resolves user by numeric ID or username string via UserService contract. +func ResolveTargetUser(ctx context.Context, resolved, _ string) (contracts.UserDTO, bool) { + if id, err := strconv.ParseUint(resolved, 10, 64); err == nil { + if u, err := FindUserByID(ctx, id); err == nil && u != nil { + return *u, true + } + } + if u, err := FindUserByUsername(ctx, resolved); err == nil && u != nil { + return *u, true + } + return contracts.UserDTO{}, false +} + +// FindUserByID resolves a user by primary key through UserService. +func FindUserByID(ctx context.Context, id uint64) (*contracts.UserDTO, error) { + if userSvc := GetUserService(ctx); userSvc != nil { + return userSvc.GetUserByID(ctx, id) + } + return nil, consts.ErrUserNotFound +} + +// FindUserByUsername resolves a user by login name through UserService. +func FindUserByUsername(ctx context.Context, username string) (*contracts.UserDTO, error) { + if userSvc := GetUserService(ctx); userSvc != nil { + return userSvc.GetUserByUsername(ctx, username) + } + return nil, consts.ErrUserNotFound +} + +// GetFirstAdminUser resolves the first administrator through the UserService contract. +func GetFirstAdminUser(ctx context.Context) (*contracts.UserDTO, error) { + if userSvc := GetUserService(ctx); userSvc != nil { + return userSvc.GetFirstAdminUser(ctx) + } + return nil, consts.ErrNoAdminUser +} + +// ResolveSystemTarget maps system receiver aliases to administrator contact info. +func ResolveSystemTarget(ctx context.Context, resolved, channel string) (string, bool) { + if resolved != "系统" && resolved != "system" && resolved != "0" { + return "", false + } + adminUser, err := GetFirstAdminUser(ctx) + if err != nil || adminUser == nil { + return resolved, true + } + if channel == consts.ChannelEmail && adminUser.Email != "" { + return adminUser.Email, true + } + if channel != consts.ChannelEmail && adminUser.Username != "" { + return adminUser.Username, true + } + return resolved, true +} + +// GetSystemUser gets a system user DTO. +func GetSystemUser(ctx context.Context) *contracts.UserDTO { + if adminUser, err := GetFirstAdminUser(ctx); err == nil && adminUser != nil { + return adminUser + } + return &contracts.UserDTO{ + Username: "system", + Nickname: "系统管理员", + } +} + +// LoadUserFromPayload extracts user info from data. +func LoadUserFromPayload(ctx context.Context, data map[string]any) any { + if u, exists := data["user"]; exists && u != nil { + return u + } + + if userID, ok := ExtractUserID(data); ok && userID > 0 { + if user, err := FindUserByID(ctx, userID); err == nil && user != nil { + return user + } + } + + if username := ExtractUsername(data); username != "" { + if user, err := FindUserByUsername(ctx, username); err == nil && user != nil { + return user + } + } + return nil +} + +// ExtractUserFromMap extracts user info from payload/detail into body map. +func ExtractUserFromMap(ctx context.Context, data, body map[string]any) { + if u, exists := body["user"]; exists && u != nil { + return + } + if user := LoadUserFromPayload(ctx, data); user != nil { + body["user"] = user + } +} + +// ExtractUserID extracts a user ID from map keys. +func ExtractUserID(data map[string]any) (uint64, bool) { + for _, k := range []string{"user_id", "userId", "uid"} { + val, ok := data[k] + if !ok || val == nil { + continue + } + switch v := val.(type) { + case float64: + if v >= 0 { + return uint64(v), true + } + case int: + if v >= 0 { + return uint64(v), true + } + case int64: + if v >= 0 { + return uint64(v), true + } + case uint64: + return v, true + case string: + if id, err := strconv.ParseUint(v, 10, 64); err == nil { + return id, true + } + } + } + return 0, false +} + +// ExtractUsername extracts a username string from map keys. +func ExtractUsername(data map[string]any) string { + for _, k := range []string{"username", "user_name"} { + if val, ok := data[k]; ok && val != nil { + if s, ok := val.(string); ok && s != "" { + return s + } + } + } + return "" +} diff --git a/backend/plugins/domain/msg_gateway/service/push_worker.go b/backend/plugins/domain/msg_gateway/service/push_worker.go new file mode 100644 index 00000000..7b5337bd --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/push_worker.go @@ -0,0 +1,176 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/logger" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" + pkgpush "Wavelet/plugins/domain/msg_gateway/push" + "context" + "encoding/json" + "errors" + "fmt" +) + +const ( + // SendNotificationTask is the asynq task name for push notification. + SendNotificationTask = consts.SendNotificationTask + // TaskTypeSendNotification is the admin task manager type identifier. + TaskTypeSendNotification = consts.TaskTypeSendNotification +) + +// SendNotificationMeta represents the task metadata. +var SendNotificationMeta = contracts.TaskMetaDTO{ + Type: TaskTypeSendNotification, + AsynqTask: SendNotificationTask, + Name: "推送通知", + DisplayName: "推送通知", + Description: "异步执行系统通知的多渠道派发与推送", + Category: "push", + SupportsTime: false, + MaxRetry: 3, + Queue: "default", + Retryable: true, + Params: []contracts.TaskParamDTO{ + { + Name: "event_key", + Label: "事件标识", + Type: "string", + Required: true, + Placeholder: "admin_login", + Description: "事件标识 (如 admin_login)", + }, + { + Name: "target", + Label: "目标接收者", + Type: "string", + Required: false, + Description: "目标接收者", + }, + }, +} + +// PushHandler handles asynchronous notification sending. +type PushHandler struct{} + +// ValidatePayload validates and normalizes push parameters. +func (h *PushHandler) ValidatePayload(payload []byte) ([]byte, error) { + if len(payload) == 0 { + return nil, errors.New(consts.ErrPayloadRequired) + } + + var req do.SendPayload + if err := json.Unmarshal(payload, &req); err != nil { + return nil, fmt.Errorf("%s: %w", consts.ErrInvalidJSONFormat, err) + } + + if req.Config.Channel == "" { + return nil, errors.New(consts.ErrChannelTypeRequired) + } + + return json.Marshal(req) +} + +// Execute performs the push send and logs delivery history audit. +func (h *PushHandler) Execute(ctx context.Context, payload []byte) error { + var req do.SendPayload + if err := json.Unmarshal(payload, &req); err != nil { + logger.ErrorF(ctx, "[Push] 解析推送参数失败: %v", err) + return fmt.Errorf("%s: %w", consts.ErrParsePayloadFailed, err) + } + + logger.InfoF(ctx, "[Push] 开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target) + + pusher, err := pkgpush.GetPusher(req.Config.Channel) + if err != nil { + errWrap := fmt.Errorf("%s: %w", consts.ErrGetPusherFailed, err) + logger.ErrorF(ctx, "[Push] 推送失败: %v", errWrap) + h.recordHistory(ctx, req, "failed", errWrap.Error()) + return errWrap + } + + flatBody := req.Body.Flatten() + upstreamResp, err := pusher.Send(ctx, req.Config, req.Target, flatBody, req.Template, nil) + + title := req.Body.Title + content := req.Body.Content + + if err != nil { + logger.ErrorF(ctx, "[Push] 消息推送失败 (标题: %s): %v, 上游返回: %s", title, err, upstreamResp) + h.recordHistory(ctx, req, "failed", err.Error()) + return fmt.Errorf("pusher.Send failed: %w", err) + } + + logger.InfoF(ctx, "[Push] 消息推送成功 (标题: %s, 内容摘要: %s), 上游返回: %s", title, content, upstreamResp) + h.recordHistory(ctx, req, "success", "") + + return nil +} + +func (h *PushHandler) recordHistory(ctx context.Context, req do.SendPayload, status, errMsg string) { + if dbErr := RecordPushHistory(ctx, req, status, errMsg); dbErr != nil { + logger.ErrorF(ctx, "[Push] 写入推送历史审计记录失败: %v", dbErr) + } +} + +// EnqueuePushTask dispatches a notification payload to the async push worker. +func EnqueuePushTask(ctx context.Context, payload do.SendPayload) error { + payloadBytes, err := json.Marshal(payload) + if err != nil { + return err + } + if taskSvc := GetTaskService(ctx); taskSvc != nil { + _, err = taskSvc.Dispatch(ctx, "send_notification", payloadBytes, contracts.TaskTriggerSystem) + return err + } + return errors.New(consts.ErrTaskServiceUnavailable) +} + +// RecordPushHistory creates a push history audit record. +func RecordPushHistory(ctx context.Context, req do.SendPayload, status, errMsg string) error { + title := req.Body.Title + content := req.Body.Content + level := req.Body.Level + if title == "" { + title = "系统通知" + } + if level == "" { + level = consts.DefaultLevelInfo + } + + target := req.Target + if target == "" { + if req.Config.URL != "" { + target = req.Config.URL + const maxTargetLen = 50 + const truncatedLen = 47 + if len(target) > maxTargetLen { + target = target[:truncatedLen] + "..." + } + } else { + target = "default" + } + } + + history := entity.PushHistory{ + EventKey: req.EventKey, + Channel: req.Config.Channel, + Target: target, + Title: title, + Content: content, + Level: level, + Status: status, + ErrorMsg: errMsg, + } + return dao.CreatePushHistoryRecord(ctx, &history) +} + +// ListPushHistories returns a paginated push delivery audit page. +func ListPushHistories(ctx context.Context, filter do.PushHistoryListFilter) (int64, []entity.PushHistory, error) { + return dao.ListPushHistoriesRecord(ctx, filter) +} diff --git a/backend/plugins/domain/msg_gateway/service/service.go b/backend/plugins/domain/msg_gateway/service/service.go index 82236e1f..dc6f7959 100644 --- a/backend/plugins/domain/msg_gateway/service/service.go +++ b/backend/plugins/domain/msg_gateway/service/service.go @@ -1,225 +1,17 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -// Package service implements domain business logic and channel runners for msg_gateway. +// Package service implements domain business logic, bot gateway adapters, and push notification services for msg_gateway. package service import ( "Wavelet/core" "Wavelet/core/contracts" - "Wavelet/pkg/logger" - "Wavelet/pkg/util" - "Wavelet/plugins/domain/msg_gateway/consts" - "Wavelet/plugins/domain/msg_gateway/dao" - "Wavelet/plugins/domain/msg_gateway/model/do" - "Wavelet/plugins/domain/msg_gateway/model/entity" "context" - "crypto/rand" - "crypto/sha256" - "encoding/hex" - "encoding/json" - "errors" - "strconv" - "strings" "sync" - "time" - "unicode" ) -// Handler processes one inbound message. -type Handler func(ctx context.Context, msg do.InboundMessage) error - -// Factory constructs a Channel from decrypted config. -type Factory func(cfg do.ChannelConfig, onInbound Handler) (Channel, error) - -// Channel is one connected messaging adapter. -type Channel interface { - Type() string - Connect(ctx context.Context) error - Disconnect(ctx context.Context) error - Send(ctx context.Context, to do.Recipient, msg do.OutboundMessage) error - Capabilities() do.Capability -} - -var ( - factoriesMu sync.RWMutex - factories = map[string]Factory{} -) - -// Register stores a channel factory under typ. -func Register(typ string, fn Factory) { - factoriesMu.Lock() - defer factoriesMu.Unlock() - factories[typ] = fn -} - -// Lookup returns a previously registered factory. -func Lookup(typ string) (Factory, bool) { - factoriesMu.RLock() - defer factoriesMu.RUnlock() - fn, ok := factories[typ] - return fn, ok -} - -// Re-exported constants. -const ( - CodeAlphabet = consts.CodeAlphabet - CodeLength = consts.CodeLength -) - -// GenerateCode returns an 8-character pairing code. -func GenerateCode() (string, error) { - buf := make([]byte, CodeLength) - if _, err := rand.Read(buf); err != nil { - return "", err - } - out := make([]byte, CodeLength) - for i, b := range buf { - out[i] = CodeAlphabet[int(b)%len(CodeAlphabet)] - } - return string(out), nil -} - -// NormalizeCode strips separators and uppercases. -func NormalizeCode(s string) string { - var b strings.Builder - for _, r := range s { - if r == '-' || unicode.IsSpace(r) { - continue - } - b.WriteRune(unicode.ToUpper(r)) - } - return b.String() -} - -// FormatCode renders ABCD-EFGH. -func FormatCode(s string) string { - s = NormalizeCode(s) - if len(s) != CodeLength { - return s - } - return s[:4] + "-" + s[4:] -} - -var ( - credentialSecretMu sync.RWMutex - credentialSecret string -) - -// SetCredentialSecret sets the secret used to derive CredentialKey. -func SetCredentialSecret(secret string) { - credentialSecretMu.Lock() - defer credentialSecretMu.Unlock() - credentialSecret = secret -} - -// CredentialKey is AES-256 hex derived from the session secret. -func CredentialKey() string { - credentialSecretMu.RLock() - secret := credentialSecret - credentialSecretMu.RUnlock() - sum := sha256.Sum256([]byte(secret)) - return hex.EncodeToString(sum[:]) -} - -// EncryptCredentials encrypts a credential map as JSON. -func EncryptCredentials(creds map[string]string) (string, error) { - if creds == nil { - creds = map[string]string{} - } - raw, err := json.Marshal(creds) - if err != nil { - return "", err - } - return util.Encrypt(CredentialKey(), string(raw)) -} - -// DecryptCredentials decrypts a credential map. -func DecryptCredentials(ciphertext string) (map[string]string, error) { - if ciphertext == "" { - return map[string]string{}, nil - } - plain, err := util.Decrypt(CredentialKey(), ciphertext) - if err != nil { - return nil, err - } - var out map[string]string - if err := json.Unmarshal([]byte(plain), &out); err != nil { - return nil, err - } - if out == nil { - out = map[string]string{} - } - return out, nil -} - -// ParseExtra decodes optional extra JSON into a string map. -func ParseExtra(raw string) map[string]string { - if raw == "" { - return map[string]string{} - } - var out map[string]string - if err := json.Unmarshal([]byte(raw), &out); err != nil || out == nil { - return map[string]string{} - } - return out -} - -// EncodeExtra encodes extra fields as JSON. -func EncodeExtra(extra map[string]string) string { - if extra == nil { - return "" - } - raw, err := json.Marshal(extra) - if err != nil { - return "" - } - return string(raw) -} - -// Runner manages lifecycle for long-lived channel adapters (WebSocket, long-polling, etc.). -type Runner struct { - mu sync.Mutex - running bool - cancel context.CancelFunc -} - -// GlobalRunner is the default global runner instance. -var GlobalRunner = &Runner{} - -// Start starts all background long-lived channel runners. -func Start(ctx context.Context) error { - GlobalRunner.mu.Lock() - defer GlobalRunner.mu.Unlock() - - if GlobalRunner.running { - return nil - } - - runCtx, cancel := context.WithCancel(ctx) - GlobalRunner.cancel = cancel - GlobalRunner.running = true - - logger.InfoF(runCtx, "[MessageGateway] Starting bot channel runners...") - return nil -} - -// Stop stops the channel runner. -func Stop() { - GlobalRunner.mu.Lock() - defer GlobalRunner.mu.Unlock() - - if !GlobalRunner.running { - return - } - - if GlobalRunner.cancel != nil { - GlobalRunner.cancel() - } - GlobalRunner.running = false -} - -// Cordis contract singletons consumed by service layer. +// Platform service dependencies resolved from Cordis context or global fallbacks. var ( cacheMu sync.RWMutex cacheSvc contracts.CacheService @@ -261,7 +53,7 @@ func GetCache(ctx context.Context) contracts.CacheService { return s } -// GetTaskService returns the task service. +// GetTaskService resolves the task service for the context. func GetTaskService(ctx context.Context) contracts.TaskService { if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil { return s @@ -281,125 +73,3 @@ func GetUserService(ctx context.Context) contracts.UserService { userMu.RUnlock() return s } - -// BindChannel consumes a pairing code and binds the platform identity to the user. -func BindChannel(ctx context.Context, userID uint64, req do.BindRequest) (do.BindingDTO, error) { - channelID, err := strconv.ParseUint(strings.TrimSpace(req.ChannelID), 10, 64) - if err != nil || channelID == 0 { - return do.BindingDTO{}, consts.ErrChannelIDRequired - } - code := NormalizeCode(req.Code) - if code == "" { - return do.BindingDTO{}, consts.ErrCodeInvalid - } - pairing, err := dao.GetPairingCode(ctx, code) - if err != nil { - if errors.Is(err, consts.ErrRecordNotFound) { - return do.BindingDTO{}, consts.ErrCodeInvalid - } - return do.BindingDTO{}, err - } - if !pairing.ExpiresAt.After(time.Now()) { - return do.BindingDTO{}, consts.ErrCodeInvalid - } - if pairing.ChannelID != channelID { - return do.BindingDTO{}, consts.ErrChannelMismatch - } - ch, err := dao.GetMessageChannel(ctx, channelID) - if err != nil { - if errors.Is(err, consts.ErrRecordNotFound) { - return do.BindingDTO{}, consts.ErrCodeInvalid - } - return do.BindingDTO{}, err - } - if !ch.Enabled { - return do.BindingDTO{}, consts.ErrChannelDisabled - } - - existing, err := dao.GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID) - if err != nil && !errors.Is(err, consts.ErrRecordNotFound) { - return do.BindingDTO{}, err - } - if err == nil && existing != nil { - if existing.UserID != userID { - return do.BindingDTO{}, consts.ErrPlatformAlreadyBound - } - _ = dao.DeletePairingCode(ctx, pairing.Code) - return ToBindingDTO(existing, ch), nil - } - - row := &entity.MessageBinding{ - UserID: userID, - ChannelID: channelID, - PlatformUserID: pairing.PlatformUserID, - } - if err := dao.CreateMessageBinding(ctx, row); err != nil { - return do.BindingDTO{}, err - } - if err := dao.DeletePairingCode(ctx, pairing.Code); err != nil { - return do.BindingDTO{}, err - } - return ToBindingDTO(row, ch), nil -} - -// ListEnabledPublicChannels returns the channels a user may bind to. -func ListEnabledPublicChannels(ctx context.Context) ([]do.PublicChannelDTO, error) { - rows, err := dao.ListEnabledMessageChannels(ctx) - if err != nil { - return nil, err - } - out := make([]do.PublicChannelDTO, 0, len(rows)) - for _, row := range rows { - out = append(out, do.PublicChannelDTO{ID: row.ID, Name: row.Name, Type: row.Type}) - } - return out, nil -} - -// ListUserBindings returns the binding rows of one user enriched with channel info. -func ListUserBindings(ctx context.Context, userID uint64) ([]do.BindingDTO, error) { - rows, err := dao.ListBindingsByUser(ctx, userID) - if err != nil { - return nil, err - } - out := make([]do.BindingDTO, 0, len(rows)) - for i := range rows { - ch, err := dao.GetMessageChannel(ctx, rows[i].ChannelID) - if err != nil { - out = append(out, ToBindingDTO(&rows[i], nil)) - continue - } - out = append(out, ToBindingDTO(&rows[i], ch)) - } - return out, nil -} - -// UnbindChannel removes a binding owned by the given user. -func UnbindChannel(ctx context.Context, userID, bindingID uint64) error { - row, err := dao.GetMessageBinding(ctx, bindingID) - if err != nil { - if errors.Is(err, consts.ErrRecordNotFound) { - return consts.ErrBindingNotFound - } - return err - } - if row.UserID != userID { - return consts.ErrBindingForbidden - } - return dao.DeleteMessageBinding(ctx, bindingID) -} - -// ToBindingDTO projects a binding row and its optional channel onto the user DTO. -func ToBindingDTO(row *entity.MessageBinding, ch *entity.MessageChannel) do.BindingDTO { - dto := do.BindingDTO{ - ID: row.ID, - UserID: row.UserID, - ChannelID: row.ChannelID, - PlatformUserID: row.PlatformUserID, - CreatedAt: row.CreatedAt, - } - if ch != nil { - dto.ChannelName = ch.Name - dto.ChannelType = ch.Type - } - return dto -}