mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-11 01:36:37 +08:00
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:
@@ -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
|
||||||
|
)
|
||||||
+6
-67
@@ -1,69 +1,10 @@
|
|||||||
// Copyright 2026 Arctel.net
|
// Copyright 2026 Arctel.net
|
||||||
// SPDX-License-Identifier: Apache-2.0
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
// Package consts defines constants, sentinel errors, and user-facing error messages
|
|
||||||
// for the msg_gateway plugin.
|
|
||||||
package consts
|
package consts
|
||||||
|
|
||||||
import "errors"
|
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.
|
// Sentinel errors.
|
||||||
var (
|
var (
|
||||||
ErrCodeInvalid = errors.New("invalid or expired pairing code")
|
ErrCodeInvalid = errors.New("invalid or expired pairing code")
|
||||||
@@ -73,14 +14,15 @@ var (
|
|||||||
ErrBindingForbidden = errors.New("cannot unbind another user's binding")
|
ErrBindingForbidden = errors.New("cannot unbind another user's binding")
|
||||||
ErrChannelIDRequired = errors.New("channel_id is required")
|
ErrChannelIDRequired = errors.New("channel_id is required")
|
||||||
ErrChannelDisabled = errors.New("channel is not enabled")
|
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
|
// ErrRecordNotFound maps GORM's missing-row sentinel at the DAO boundary so
|
||||||
// upper layers never import gorm. Its text matches gorm.ErrRecordNotFound verbatim.
|
// upper layers never import gorm. Its text matches gorm.ErrRecordNotFound verbatim.
|
||||||
ErrRecordNotFound = errors.New("record not found")
|
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.
|
// User-facing validation and error message constants.
|
||||||
@@ -89,7 +31,7 @@ const (
|
|||||||
ErrTypeInvalid = "type must be telegram or qq"
|
ErrTypeInvalid = "type must be telegram or qq"
|
||||||
ErrTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text
|
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
|
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"
|
ErrChannelProbeFailed = "channel probe failed"
|
||||||
ErrBotDispatchTextRequired = "message text is required"
|
ErrBotDispatchTextRequired = "message text is required"
|
||||||
ErrBotChannelNotRegistered = "channel adapter is not registered"
|
ErrBotChannelNotRegistered = "channel adapter is not registered"
|
||||||
@@ -99,7 +41,6 @@ const (
|
|||||||
ErrInvalidBindingID = "invalid binding id"
|
ErrInvalidBindingID = "invalid binding id"
|
||||||
ErrInvalidChannelID = "invalid channel id"
|
ErrInvalidChannelID = "invalid channel id"
|
||||||
ErrInvalidEventID = "invalid event id"
|
ErrInvalidEventID = "invalid event id"
|
||||||
ErrEventNotFound = "notification event not found"
|
|
||||||
ErrValidationFailed = "validation failed"
|
ErrValidationFailed = "validation failed"
|
||||||
|
|
||||||
ErrMissingTelegramToken = "missing telegram bot token"
|
ErrMissingTelegramToken = "missing telegram bot token"
|
||||||
@@ -120,8 +61,6 @@ const (
|
|||||||
ErrEventKeyOrTaskType = "either event_key or task_type must be provided"
|
ErrEventKeyOrTaskType = "either event_key or task_type must be provided"
|
||||||
ErrUnsupportedEventKey = "unsupported built-in event key"
|
ErrUnsupportedEventKey = "unsupported built-in event key"
|
||||||
ErrTaskServiceUnavailable = "task service not available"
|
ErrTaskServiceUnavailable = "task service not available"
|
||||||
ErrUserNotFound = "user not found"
|
|
||||||
ErrNoAdminUser = "no admin user found"
|
|
||||||
|
|
||||||
ErrPayloadRequired = "payload is required"
|
ErrPayloadRequired = "payload is required"
|
||||||
ErrInvalidJSONFormat = "invalid json format"
|
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/pkg/response"
|
||||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||||
"Wavelet/plugins/domain/msg_gateway/service"
|
"Wavelet/plugins/domain/msg_gateway/service"
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
)
|
)
|
||||||
@@ -43,17 +43,12 @@ func ListAdminChannels(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func parseAdminChannelID(c *gin.Context) (uint64, bool) {
|
func parseAdminChannelID(c *gin.Context) (uint64, bool) {
|
||||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
return parseUint64Param(c, "id", consts.ErrInvalidChannelID)
|
||||||
if err != nil {
|
|
||||||
response.AbortBadRequest(c, consts.ErrInvalidChannelID)
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
return id, true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
|
func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
|
||||||
if err.Error() == consts.ErrChannelNotFound {
|
if errors.Is(err, consts.ErrChannelNotFound) || err.Error() == consts.ErrChannelNotFoundText {
|
||||||
response.AbortNotFound(c, err.Error())
|
response.AbortNotFound(c, consts.ErrChannelNotFoundText)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
fallback(c, err.Error())
|
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"
|
"Wavelet/plugins/domain/msg_gateway/service"
|
||||||
"errors"
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
)
|
)
|
||||||
@@ -46,18 +45,13 @@ func ListPushChannels(c *gin.Context) {
|
|||||||
|
|
||||||
// parsePushChannelID reads the path identifier of a push channel.
|
// parsePushChannelID reads the path identifier of a push channel.
|
||||||
func parsePushChannelID(c *gin.Context) (uint64, bool) {
|
func parsePushChannelID(c *gin.Context) (uint64, bool) {
|
||||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
return parseUint64Param(c, "id", consts.ErrInvalidChannelID)
|
||||||
if err != nil {
|
|
||||||
response.AbortBadRequest(c, consts.ErrInvalidChannelID)
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
return id, true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// handlePushChannelNotFoundError maps a missing channel row to 404, others to fallback.
|
// 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)) {
|
func handlePushChannelNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
|
||||||
if errors.Is(err, consts.ErrRecordNotFound) {
|
if errors.Is(err, consts.ErrRecordNotFound) || errors.Is(err, consts.ErrChannelNotFound) || err.Error() == consts.ErrChannelNotFoundText {
|
||||||
response.AbortNotFound(c, consts.ErrChannelNotFound)
|
response.AbortNotFound(c, consts.ErrChannelNotFoundText)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
fallback(c, err.Error())
|
fallback(c, err.Error())
|
||||||
|
|||||||
@@ -47,18 +47,13 @@ func ListBuiltInPushEvents(c *gin.Context) {
|
|||||||
|
|
||||||
// parsePushEventID reads the path identifier of a push event.
|
// parsePushEventID reads the path identifier of a push event.
|
||||||
func parsePushEventID(c *gin.Context) (uint64, bool) {
|
func parsePushEventID(c *gin.Context) (uint64, bool) {
|
||||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
return parseUint64Param(c, "id", consts.ErrInvalidEventID)
|
||||||
if err != nil {
|
|
||||||
response.AbortBadRequest(c, consts.ErrInvalidEventID)
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
return id, true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// handlePushEventNotFoundError maps a missing event row to 404, others to fallback.
|
// 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)) {
|
func handlePushEventNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
|
||||||
if errors.Is(err, consts.ErrRecordNotFound) {
|
if errors.Is(err, consts.ErrRecordNotFound) || errors.Is(err, consts.ErrEventNotFound) || err.Error() == consts.ErrEventNotFound.Error() {
|
||||||
response.AbortNotFound(c, consts.ErrEventNotFound)
|
response.AbortNotFound(c, consts.ErrEventNotFound.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
fallback(c, err.Error())
|
fallback(c, err.Error())
|
||||||
|
|||||||
@@ -4,65 +4,16 @@
|
|||||||
package controller
|
package controller
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"Wavelet/core/contracts"
|
|
||||||
"Wavelet/pkg/ginutil"
|
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||||
"Wavelet/plugins/domain/msg_gateway/service"
|
"Wavelet/plugins/domain/msg_gateway/service"
|
||||||
"context"
|
|
||||||
"errors"
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"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.
|
// ListChannels lists enabled channels a user can bind.
|
||||||
// @Summary List enabled messaging channels
|
// @Summary List enabled messaging channels
|
||||||
// @Description Returns enabled system bots the current user can pair with
|
// @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)
|
response.AbortUnauthorized(c, consts.ErrLoginRequired)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
id, ok := parseUint64Param(c, "id", consts.ErrInvalidBindingID)
|
||||||
if err != nil {
|
if !ok {
|
||||||
response.AbortBadRequest(c, consts.ErrInvalidBindingID)
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if err := service.UnbindChannel(c.Request.Context(), user.ID, id); err != nil {
|
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)
|
||||||
|
}
|
||||||
@@ -7,27 +7,21 @@ package dao
|
|||||||
import (
|
import (
|
||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/idgen"
|
|
||||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
|
||||||
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
dbMu sync.RWMutex
|
dbMu sync.RWMutex
|
||||||
dbSvc contracts.DBService
|
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.
|
// SetDBService sets the database service singleton.
|
||||||
func SetDBService(s contracts.DBService) {
|
func SetDBService(s contracts.DBService) {
|
||||||
dbMu.Lock()
|
dbMu.Lock()
|
||||||
@@ -35,8 +29,19 @@ func SetDBService(s contracts.DBService) {
|
|||||||
dbSvc = s
|
dbSvc = s
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetDB resolves the persistence handle for the current call, preferring an
|
// SetDBServiceForTest injects a DBService for tests.
|
||||||
// explicitly injected *core.Context before falling back to the plugin singleton.
|
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 {
|
func GetDB(ctx context.Context) *gorm.DB {
|
||||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != 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
|
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
|
// mapNotFound translates GORM's missing-row sentinel into the plugin-level
|
||||||
// consts.ErrRecordNotFound so the service and controller layers stay free of gorm imports.
|
// consts.ErrRecordNotFound so the service and controller layers stay free of gorm imports.
|
||||||
func mapNotFound(err error) error {
|
func mapNotFound(err error) error {
|
||||||
@@ -60,149 +78,3 @@ func mapNotFound(err error) error {
|
|||||||
}
|
}
|
||||||
return err
|
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
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -4,15 +4,9 @@
|
|||||||
package dao
|
package dao
|
||||||
|
|
||||||
import (
|
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/do"
|
||||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"sync"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
@@ -23,31 +17,6 @@ const (
|
|||||||
activePushEventCacheTTL = 24 * time.Hour
|
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.
|
// ListPushChannelsRecord returns all push channels ordered by creation time descending.
|
||||||
func ListPushChannelsRecord(ctx context.Context) ([]entity.PushChannel, error) {
|
func ListPushChannelsRecord(ctx context.Context) ([]entity.PushChannel, error) {
|
||||||
var channels []entity.PushChannel
|
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) {
|
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 {
|
if cache := GetCache(ctx); cache != nil {
|
||||||
|
var val T
|
||||||
if err := cache.Get(ctx, cacheKey, &val); err == nil {
|
if err := cache.Get(ctx, cacheKey, &val); err == nil {
|
||||||
return &val, nil
|
return &val, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
db := GetDB(ctx)
|
db := GetDB(ctx)
|
||||||
|
var val T
|
||||||
if err := query(db, &val); err != nil {
|
if err := query(db, &val); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -255,6 +225,9 @@ func ListPushHistoriesRecord(ctx context.Context, filter do.PushHistoryListFilte
|
|||||||
if filter.EventKey != "" {
|
if filter.EventKey != "" {
|
||||||
query = query.Where("event_key = ?", filter.EventKey)
|
query = query.Where("event_key = ?", filter.EventKey)
|
||||||
}
|
}
|
||||||
|
if filter.Channel != "" {
|
||||||
|
query = query.Where("channel = ?", filter.Channel)
|
||||||
|
}
|
||||||
if filter.Status != "" {
|
if filter.Status != "" {
|
||||||
query = query.Where("status = ?", 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 {
|
func PushHistoryQuery(ctx context.Context) *gorm.DB {
|
||||||
return GetDB(ctx).Model(&entity.PushHistory{})
|
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
|
package dao_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"Wavelet/pkg/idgen"
|
||||||
"Wavelet/pkg/testhelper"
|
"Wavelet/pkg/testhelper"
|
||||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
|
||||||
"Wavelet/plugins/domain/msg_gateway/dao"
|
"Wavelet/plugins/domain/msg_gateway/dao"
|
||||||
|
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
// stubDBService satisfies contracts.DBService over a test database handle.
|
|
||||||
type stubDBService struct{ db *gorm.DB }
|
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) 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 }
|
func TestPushChannelDAO_CRUD(t *testing.T) {
|
||||||
|
_ = idgen.Init(1)
|
||||||
// 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) {
|
|
||||||
db, _, cleanup := testhelper.SetupTestEnvironment(t)
|
db, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
require.NoError(t, db.AutoMigrate(&entity.PushChannel{}, &entity.PushEvent{}, &entity.PushHistory{}))
|
||||||
if err := db.Table("w_users").Create(map[string]any{"id": 77, "username": "seeded"}).Error; err != nil {
|
|
||||||
t.Fatalf("seed user failed: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
dao.SetDBServiceForTest(stubDBService{db: db})
|
dao.SetDBServiceForTest(stubDBService{db: db})
|
||||||
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
||||||
|
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
user, err := dao.FindUserByFieldRecord(ctx, "username", "seeded")
|
ch := entity.PushChannel{
|
||||||
if err != nil {
|
Name: "test_webhook",
|
||||||
t.Fatalf("allowlisted lookup by username failed: %v", err)
|
Type: "custom",
|
||||||
}
|
URL: "https://example.com/hook",
|
||||||
if user.ID != 77 {
|
Enabled: true,
|
||||||
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)
|
|
||||||
}
|
}
|
||||||
|
require.NoError(t, dao.CreatePushChannelRecord(ctx, &ch))
|
||||||
|
assert.NotZero(t, ch.ID)
|
||||||
|
|
||||||
cases := []struct {
|
loaded, err := dao.GetPushChannelByIDRecord(ctx, ch.ID)
|
||||||
name string
|
require.NoError(t, err)
|
||||||
field string
|
assert.Equal(t, "test_webhook", loaded.Name)
|
||||||
}{
|
|
||||||
{"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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
var remaining int64
|
active, err := dao.GetActivePushChannelByName(ctx, "test_webhook")
|
||||||
if err := db.Table("w_users").Count(&remaining).Error; err != nil || remaining != 1 {
|
require.NoError(t, err)
|
||||||
t.Fatalf("w_users damaged by rejected lookups: count=%d err=%v", remaining, 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.
|
func TestPushEventDAO_CRUD(t *testing.T) {
|
||||||
var smtpTestValues = map[string]string{
|
_ = idgen.Init(1)
|
||||||
"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) {
|
|
||||||
db, _, cleanup := testhelper.SetupTestEnvironment(t)
|
db, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
require.NoError(t, db.AutoMigrate(&entity.PushChannel{}, &entity.PushEvent{}, &entity.PushHistory{}))
|
||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
dao.SetDBServiceForTest(stubDBService{db: db})
|
dao.SetDBServiceForTest(stubDBService{db: db})
|
||||||
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
||||||
|
|
||||||
cfg, err := dao.LoadSMTPConfigRecord(context.Background())
|
ctx := context.Background()
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("LoadSMTPConfigRecord: %v", err)
|
ev := entity.PushEvent{
|
||||||
}
|
EventKey: "test_event",
|
||||||
if cfg.Host != smtpTestValues["smtp_host"] || cfg.Port != smtpTestValues["smtp_port"] ||
|
Name: "测试事件",
|
||||||
cfg.Username != smtpTestValues["smtp_username"] || cfg.Password != smtpTestValues["smtp_password"] {
|
Channels: []string{"test_webhook"},
|
||||||
t.Errorf("got %+v, want every SMTP field mapped from its own row", cfg)
|
Targets: []string{"admin"},
|
||||||
}
|
Template: `{"title":"Hello"}`,
|
||||||
}
|
Enabled: true,
|
||||||
|
|
||||||
// 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")
|
|
||||||
}
|
}
|
||||||
|
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 (
|
import (
|
||||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||||
pkgpush "Wavelet/plugins/domain/msg_gateway/push"
|
pkgpush "Wavelet/plugins/domain/msg_gateway/push"
|
||||||
"sync"
|
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -159,58 +158,9 @@ type PushNotificationEvent struct {
|
|||||||
Metadata map[string]any `json:"metadata,omitempty"`
|
Metadata map[string]any `json:"metadata,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
var (
|
//nolint:goconst,dupl // Static push channel form definitions table
|
||||||
pushDefMu sync.RWMutex
|
var defaultPushDefinitions = []PushDefinition{
|
||||||
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{
|
|
||||||
Type: consts.ChannelCustom,
|
Type: consts.ChannelCustom,
|
||||||
Name: "自定义消息通道",
|
Name: "自定义消息通道",
|
||||||
Description: "使用自定义 HTTP POST 请求向外部 Webhook 发送数据。",
|
Description: "使用自定义 HTTP POST 请求向外部 Webhook 发送数据。",
|
||||||
@@ -232,9 +182,8 @@ func init() {
|
|||||||
Description: "可使用的变量:$title, $description, $content, $url, $to。例如 {\"text\": \"$content\"}",
|
Description: "可使用的变量:$title, $description, $content, $url, $to。例如 {\"text\": \"$content\"}",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
})
|
},
|
||||||
|
{
|
||||||
RegisterPushChannelDefinition(PushDefinition{
|
|
||||||
Type: consts.ChannelLark,
|
Type: consts.ChannelLark,
|
||||||
Name: "飞书群机器人",
|
Name: "飞书群机器人",
|
||||||
Description: "配置飞书群自定义机器人的 Webhook 接口投递。",
|
Description: "配置飞书群自定义机器人的 Webhook 接口投递。",
|
||||||
@@ -264,9 +213,8 @@ func init() {
|
|||||||
Description: "若填写,必须是合法的飞书卡片 JSON 格式",
|
Description: "若填写,必须是合法的飞书卡片 JSON 格式",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
})
|
},
|
||||||
|
{
|
||||||
RegisterPushChannelDefinition(PushDefinition{
|
|
||||||
Type: consts.ChannelDingTalk,
|
Type: consts.ChannelDingTalk,
|
||||||
Name: "钉钉群机器人",
|
Name: "钉钉群机器人",
|
||||||
Description: "配置钉钉群自定义机器人的 Webhook 接口投递。",
|
Description: "配置钉钉群自定义机器人的 Webhook 接口投递。",
|
||||||
@@ -288,9 +236,8 @@ func init() {
|
|||||||
Description: "钉钉群机器人安全设置中的加签 Secret",
|
Description: "钉钉群机器人安全设置中的加签 Secret",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
})
|
},
|
||||||
|
{
|
||||||
RegisterPushChannelDefinition(PushDefinition{
|
|
||||||
Type: consts.ChannelTelegram,
|
Type: consts.ChannelTelegram,
|
||||||
Name: "Telegram 机器人",
|
Name: "Telegram 机器人",
|
||||||
Description: "配置 Telegram 机器人推送消息。",
|
Description: "配置 Telegram 机器人推送消息。",
|
||||||
@@ -320,9 +267,8 @@ func init() {
|
|||||||
Description: "默认的消息接收 Chat ID。如果通知事件中未配置 targets,将推送到此 ID",
|
Description: "默认的消息接收 Chat ID。如果通知事件中未配置 targets,将推送到此 ID",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
})
|
},
|
||||||
|
{
|
||||||
RegisterPushChannelDefinition(PushDefinition{
|
|
||||||
Type: consts.ChannelBark,
|
Type: consts.ChannelBark,
|
||||||
Name: "Bark (iOS 推送)",
|
Name: "Bark (iOS 推送)",
|
||||||
Description: "配置 Bark 推送通知至 iPhone / iPad 客户端。",
|
Description: "配置 Bark 推送通知至 iPhone / iPad 客户端。",
|
||||||
@@ -352,9 +298,8 @@ func init() {
|
|||||||
Description: "可选的 JSON 配置,支持 group (分组)、sound (铃声)、icon (自定义图标)",
|
Description: "可选的 JSON 配置,支持 group (分组)、sound (铃声)、icon (自定义图标)",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
})
|
},
|
||||||
|
{
|
||||||
RegisterPushChannelDefinition(PushDefinition{
|
|
||||||
Type: consts.ChannelDiscord,
|
Type: consts.ChannelDiscord,
|
||||||
Name: "Discord 频道",
|
Name: "Discord 频道",
|
||||||
Description: "配置 Discord 频道的 Incoming Webhook 消息推送。",
|
Description: "配置 Discord 频道的 Incoming Webhook 消息推送。",
|
||||||
@@ -368,9 +313,8 @@ func init() {
|
|||||||
Description: "从 Discord 频道集成设置中复制的 Webhook URL",
|
Description: "从 Discord 频道集成设置中复制的 Webhook URL",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
})
|
},
|
||||||
|
{
|
||||||
RegisterPushChannelDefinition(PushDefinition{
|
|
||||||
Type: consts.ChannelSlack,
|
Type: consts.ChannelSlack,
|
||||||
Name: "Slack 频道",
|
Name: "Slack 频道",
|
||||||
Description: "配置 Slack 频道的 Incoming Webhook 消息推送。",
|
Description: "配置 Slack 频道的 Incoming Webhook 消息推送。",
|
||||||
@@ -384,9 +328,8 @@ func init() {
|
|||||||
Description: "从 Slack 应用配置中复制的 Incoming Webhook URL",
|
Description: "从 Slack 应用配置中复制的 Incoming Webhook URL",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
})
|
},
|
||||||
|
{
|
||||||
RegisterPushChannelDefinition(PushDefinition{
|
|
||||||
Type: consts.ChannelPushover,
|
Type: consts.ChannelPushover,
|
||||||
Name: "Pushover 推送",
|
Name: "Pushover 推送",
|
||||||
Description: "配置 Pushover 即时推送到手机/桌面客户端。",
|
Description: "配置 Pushover 即时推送到手机/桌面客户端。",
|
||||||
@@ -408,12 +351,18 @@ func init() {
|
|||||||
Description: "Pushover 个人账号的 User Key",
|
Description: "Pushover 个人账号的 User Key",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
})
|
},
|
||||||
|
{
|
||||||
RegisterPushChannelDefinition(PushDefinition{
|
|
||||||
Type: consts.ChannelEmail,
|
Type: consts.ChannelEmail,
|
||||||
Name: "邮件推送通道",
|
Name: "邮件推送通道",
|
||||||
Description: "邮件推送通道直接使用系统全局 SMTP 设置进行发送,无需在此填写服务器配置。",
|
Description: "邮件推送通道直接使用系统全局 SMTP 设置进行发送,无需在此填写服务器配置。",
|
||||||
Fields: []PushField{},
|
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -64,9 +64,6 @@ func (p *Plugin) Name() string {
|
|||||||
func (p *Plugin) Inject() []reflect.Type {
|
func (p *Plugin) Inject() []reflect.Type {
|
||||||
return []reflect.Type{
|
return []reflect.Type{
|
||||||
reflect.TypeFor[contracts.DBService](),
|
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](),
|
reflect.TypeFor[contracts.AuthService](),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -246,116 +243,20 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Re-exported constants.
|
// Entity and DO aliases exported for integration test compatibility.
|
||||||
const (
|
type (
|
||||||
CodeAlphabet = service.CodeAlphabet
|
// MessageChannel is an alias for entity.MessageChannel.
|
||||||
CodeLength = service.CodeLength
|
MessageChannel = entity.MessageChannel
|
||||||
)
|
// MessageBinding is an alias for entity.MessageBinding.
|
||||||
|
MessageBinding = entity.MessageBinding
|
||||||
// MessageChannel is an alias for entity.MessageChannel.
|
// MessagePairingCode is an alias for entity.MessagePairingCode.
|
||||||
type MessageChannel = entity.MessageChannel
|
MessagePairingCode = entity.MessagePairingCode
|
||||||
|
// PushChannel is an alias for entity.PushChannel.
|
||||||
// MessageBinding is an alias for entity.MessageBinding.
|
PushChannel = entity.PushChannel
|
||||||
type MessageBinding = entity.MessageBinding
|
// PushEvent is an alias for entity.PushEvent.
|
||||||
|
PushEvent = entity.PushEvent
|
||||||
// MessagePairingCode is an alias for entity.MessagePairingCode.
|
// PushHistory is an alias for entity.PushHistory.
|
||||||
type MessagePairingCode = entity.MessagePairingCode
|
PushHistory = entity.PushHistory
|
||||||
|
// PushNotificationEvent is an alias for do.PushNotificationEvent.
|
||||||
// PushChannel is an alias for entity.PushChannel.
|
PushNotificationEvent = do.PushNotificationEvent
|
||||||
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
|
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+3
-27
@@ -90,7 +90,7 @@ func UpdateChannel(ctx context.Context, id uint64, req do.UpdateChannelRequest)
|
|||||||
row, err := dao.GetMessageChannel(ctx, id)
|
row, err := dao.GetMessageChannel(ctx, id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, consts.ErrRecordNotFound) {
|
if errors.Is(err, consts.ErrRecordNotFound) {
|
||||||
return do.ChannelDTO{}, errors.New(consts.ErrChannelNotFound)
|
return do.ChannelDTO{}, errors.New(consts.ErrChannelNotFoundText)
|
||||||
}
|
}
|
||||||
return do.ChannelDTO{}, err
|
return do.ChannelDTO{}, err
|
||||||
}
|
}
|
||||||
@@ -157,7 +157,7 @@ func ListChannels(ctx context.Context) ([]do.ChannelDTO, error) {
|
|||||||
func DeleteChannel(ctx context.Context, id uint64) error {
|
func DeleteChannel(ctx context.Context, id uint64) error {
|
||||||
if _, err := dao.GetMessageChannel(ctx, id); err != nil {
|
if _, err := dao.GetMessageChannel(ctx, id); err != nil {
|
||||||
if errors.Is(err, consts.ErrRecordNotFound) {
|
if errors.Is(err, consts.ErrRecordNotFound) {
|
||||||
return errors.New(consts.ErrChannelNotFound)
|
return errors.New(consts.ErrChannelNotFoundText)
|
||||||
}
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -169,7 +169,7 @@ func ProbeChannel(ctx context.Context, id uint64) error {
|
|||||||
row, err := dao.GetMessageChannel(ctx, id)
|
row, err := dao.GetMessageChannel(ctx, id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, consts.ErrRecordNotFound) {
|
if errors.Is(err, consts.ErrRecordNotFound) {
|
||||||
return errors.New(consts.ErrChannelNotFound)
|
return errors.New(consts.ErrChannelNotFoundText)
|
||||||
}
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -285,27 +285,3 @@ func ToDTO(row *entity.MessageChannel, creds, extra map[string]string) do.Channe
|
|||||||
Extra: extra,
|
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:]
|
||||||
|
}
|
||||||
+1
-1
@@ -83,7 +83,7 @@ func (h *BotDispatchHandler) Execute(ctx context.Context, payload []byte) (*cont
|
|||||||
}
|
}
|
||||||
channels = filtered
|
channels = filtered
|
||||||
if len(channels) == 0 {
|
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
|
// Copyright 2026 Arctel.net
|
||||||
// SPDX-License-Identifier: Apache-2.0
|
// 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
|
package service
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"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"
|
"context"
|
||||||
"crypto/rand"
|
|
||||||
"crypto/sha256"
|
|
||||||
"encoding/hex"
|
|
||||||
"encoding/json"
|
|
||||||
"errors"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
|
||||||
"unicode"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// Handler processes one inbound message.
|
// Platform service dependencies resolved from Cordis context or global fallbacks.
|
||||||
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.
|
|
||||||
var (
|
var (
|
||||||
cacheMu sync.RWMutex
|
cacheMu sync.RWMutex
|
||||||
cacheSvc contracts.CacheService
|
cacheSvc contracts.CacheService
|
||||||
@@ -261,7 +53,7 @@ func GetCache(ctx context.Context) contracts.CacheService {
|
|||||||
return s
|
return s
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetTaskService returns the task service.
|
// GetTaskService resolves the task service for the context.
|
||||||
func GetTaskService(ctx context.Context) contracts.TaskService {
|
func GetTaskService(ctx context.Context) contracts.TaskService {
|
||||||
if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil {
|
if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil {
|
||||||
return s
|
return s
|
||||||
@@ -281,125 +73,3 @@ func GetUserService(ctx context.Context) contracts.UserService {
|
|||||||
userMu.RUnlock()
|
userMu.RUnlock()
|
||||||
return s
|
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
|
|
||||||
}
|
|
||||||
|
|||||||
Reference in New Issue
Block a user