refactor(msg_gateway): decouple bot gateway and push notification architecture

- Split shared monolithic consts into bot, push, and errs with typed sentinel errors
- Restructure model layer into distinct bot and push subdomains
- Refactor DAO layer to enforce single-owner principle and remove cross-table raw SQL queries
- Decompose 1150+ line service/push.go into push_channel, push_event, push_trigger, push_worker, and push_template
- Clean up controller layer with generic request handlers and parameter validation in controller/base.go
- Streamline plugin.go to core Cordis lifecycle orchestration and remove re-export bloat
- Verify all unit tests, race tests, Cordis architecture rules, and Swagger generation pass cleanly
This commit is contained in:
ryan
2026-09-02 23:21:36 +08:00
parent 8395dd5019
commit 1e19d8114a
35 changed files with 2103 additions and 2333 deletions
@@ -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
)
@@ -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"
@@ -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"
)
@@ -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())
@@ -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))
}
@@ -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())
@@ -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())
@@ -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 {
@@ -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
}
@@ -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)
}
+30 -158
View File
@@ -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
}
+5 -106
View File
@@ -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
}
@@ -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)
}
@@ -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
}
@@ -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)
}
}
@@ -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)
}
}
+16 -115
View File
@@ -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
)
@@ -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())
}
}
@@ -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:]
}
@@ -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:]
}
@@ -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
}
}
@@ -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
}
@@ -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"))
}
@@ -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
}
@@ -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())
}
File diff suppressed because it is too large Load Diff
@@ -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
}
@@ -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)
}
@@ -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
}
@@ -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 ""
}
@@ -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)
}
@@ -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
}