refactor(msg_gateway): restructure and rename message_gateway aligned with custom_example

This commit is contained in:
ryan
2026-09-02 22:50:04 +08:00
parent 87e3bfd0e6
commit 8395dd5019
57 changed files with 2150 additions and 2120 deletions
@@ -0,0 +1,164 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package qq implements the official QQ Bot C2C adapter.
package qq
import (
"Wavelet/pkg/logger"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/model/do"
"Wavelet/plugins/domain/msg_gateway/service"
"context"
"fmt"
"strings"
"sync"
"time"
"github.com/tencent-connect/botgo"
"github.com/tencent-connect/botgo/dto"
"github.com/tencent-connect/botgo/event"
"github.com/tencent-connect/botgo/openapi"
"github.com/tencent-connect/botgo/token"
"golang.org/x/oauth2"
)
// qqEvent is a testable inbound envelope.
type qqEvent struct {
Kind string
UserID string
Text string
MessageID string
}
// Adapter is an official QQ Bot C2C channel.
type Adapter struct {
cfg do.ChannelConfig
onInbound service.Handler
api openapi.OpenAPI
tokenSrc oauth2.TokenSource
cancel context.CancelFunc
mu sync.Mutex
disconnected bool
}
// New constructs a QQ adapter.
func New(cfg do.ChannelConfig, onInbound service.Handler) (service.Channel, error) {
if strings.TrimSpace(cfg.Credentials["app_id"]) == "" || strings.TrimSpace(cfg.Credentials["app_secret"]) == "" {
return nil, fmt.Errorf("qq: app_id and app_secret are required")
}
return &Adapter{cfg: cfg, onInbound: onInbound}, nil
}
// Type returns qq.
func (a *Adapter) Type() string { return consts.ChannelTypeQQ }
// Capabilities reports C2C text/media support.
func (a *Adapter) Capabilities() do.Capability {
return do.Capability{Text: true, Image: true, File: true, Reply: true}
}
// Connect starts the official WebSocket session (C2C intent).
func (a *Adapter) Connect(ctx context.Context) error {
credentials := &token.QQBotCredentials{
AppID: a.cfg.Credentials["app_id"],
AppSecret: a.cfg.Credentials["app_secret"],
}
tokSrc := token.NewQQBotTokenSource(credentials)
runCtx, cancel := context.WithCancel(ctx)
if err := token.StartRefreshAccessToken(runCtx, tokSrc); err != nil {
cancel()
return fmt.Errorf("qq: refresh token: %w", err)
}
var api openapi.OpenAPI
const apiTimeout = 5 * time.Second
if strings.EqualFold(strings.TrimSpace(a.cfg.Extra["sandbox"]), "true") {
api = botgo.NewSandboxOpenAPI(credentials.AppID, tokSrc).WithTimeout(apiTimeout)
} else {
api = botgo.NewOpenAPI(credentials.AppID, tokSrc).WithTimeout(apiTimeout)
}
wsAP, err := api.WS(ctx, nil, "")
if err != nil {
cancel()
return fmt.Errorf("qq: websocket ap: %w", err)
}
intent := event.RegisterHandlers(event.C2CMessageEventHandler(func(_ *dto.WSPayload, data *dto.WSC2CMessageData) error {
authorID := ""
if data != nil && data.Author != nil {
authorID = data.Author.ID
}
text := ""
id := ""
if data != nil {
text = data.Content
id = data.ID
}
a.handleEvent(runCtx, qqEvent{Kind: "c2c", UserID: authorID, Text: text, MessageID: id})
return nil
}))
a.mu.Lock()
a.api = api
a.tokenSrc = tokSrc
a.cancel = cancel
a.disconnected = false
a.mu.Unlock()
util.Go(func() {
if err := botgo.NewSessionManager().Start(wsAP, tokSrc, &intent); err != nil {
logger.ErrorF(runCtx, "qq session stopped: %v", err)
}
})
return nil
}
// Disconnect stops token refresh and drops further inbound events.
func (a *Adapter) Disconnect(_ context.Context) error {
a.mu.Lock()
defer a.mu.Unlock()
a.disconnected = true
if a.cancel != nil {
a.cancel()
a.cancel = nil
}
return nil
}
// Send posts a C2C text reply.
func (a *Adapter) Send(ctx context.Context, to do.Recipient, msg do.OutboundMessage) error {
a.mu.Lock()
api := a.api
a.mu.Unlock()
if api == nil {
return fmt.Errorf("qq: not connected")
}
_, err := api.PostC2CMessage(ctx, to.PlatformUserID, &dto.MessageToCreate{
Content: msg.Text,
MsgID: msg.ReplyToID,
})
return err
}
func (a *Adapter) handleEvent(ctx context.Context, ev qqEvent) {
if ev.Kind != "c2c" {
return
}
a.mu.Lock()
disconnected := a.disconnected
a.mu.Unlock()
if disconnected || a.onInbound == nil {
return
}
_ = a.onInbound(ctx, do.InboundMessage{
ChannelID: a.cfg.ID,
ChannelType: consts.ChannelTypeQQ,
PlatformUserID: ev.UserID,
ChatID: ev.UserID,
MessageID: ev.MessageID,
Text: ev.Text,
})
}
@@ -0,0 +1,41 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package qq
import (
"Wavelet/plugins/domain/msg_gateway/model/do"
"context"
"testing"
)
func TestHandleEvent_DropsNonC2C(t *testing.T) {
var got int
a := &Adapter{onInbound: func(_ context.Context, _ do.InboundMessage) error {
got++
return nil
}}
a.handleEvent(context.Background(), qqEvent{Kind: "group", UserID: "u1", Text: "hi"})
if got != 0 {
t.Fatal("non-C2C must be ignored")
}
}
func TestHandleEvent_C2CText(t *testing.T) {
var got do.InboundMessage
a := &Adapter{cfg: do.ChannelConfig{ID: 3}, onInbound: func(_ context.Context, msg do.InboundMessage) error {
got = msg
return nil
}}
a.handleEvent(context.Background(), qqEvent{Kind: "c2c", UserID: "openid-1", Text: "hello", MessageID: "m1"})
if got.Text != "hello" || got.PlatformUserID != "openid-1" || got.ChannelID != 3 {
t.Fatalf("%+v", got)
}
}
func TestNew_RequiresCreds(t *testing.T) {
_, err := New(do.ChannelConfig{}, nil)
if err == nil {
t.Fatal("expected error")
}
}
@@ -0,0 +1,182 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package telegram implements the Telegram private-chat adapter.
package telegram
import (
"Wavelet/pkg/logger"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/model/do"
"Wavelet/plugins/domain/msg_gateway/service"
"context"
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
"time"
tele "gopkg.in/telebot.v4"
)
// Adapter is a Telegram private-chat channel.
type Adapter struct {
cfg do.ChannelConfig
onInbound service.Handler
bot *tele.Bot
}
// New constructs a Telegram adapter. Call service.Register from the runner.
func New(cfg do.ChannelConfig, onInbound service.Handler) (service.Channel, error) {
if strings.TrimSpace(cfg.Credentials["bot_token"]) == "" {
return nil, fmt.Errorf("telegram: bot_token is required")
}
return &Adapter{cfg: cfg, onInbound: onInbound}, nil
}
// Type returns telegram.
func (a *Adapter) Type() string { return consts.ChannelTypeTelegram }
// Capabilities reports private-chat media support.
func (a *Adapter) Capabilities() do.Capability {
return do.Capability{Text: true, Image: true, File: true, Reply: true}
}
// longPollWindow is how long Telegram may hold a getUpdates call open before
// returning empty. telebot converts it with int(timeout / time.Second), so a
// bare integer here would mean nanoseconds, send timeout=0 and turn the poller
// into a tight loop against the Bot API.
const longPollWindow = 10 * time.Second
// buildTeleSettings assembles the telebot settings.
func buildTeleSettings(cfg do.ChannelConfig) tele.Settings {
pref := tele.Settings{
Token: cfg.Credentials["bot_token"],
Poller: &tele.LongPoller{Timeout: longPollWindow},
}
if base := strings.TrimSpace(cfg.Extra["base_url"]); base != "" {
pref.URL = strings.TrimSuffix(base, "/")
}
return pref
}
// Connect starts long polling.
func (a *Adapter) Connect(ctx context.Context) error {
bot, err := tele.NewBot(buildTeleSettings(a.cfg))
if err != nil {
return fmt.Errorf("telegram: new bot: %w", err)
}
a.bot = bot
bot.Handle(tele.OnText, func(c tele.Context) error {
a.handleTeleMessage(ctx, c.Message())
return nil
})
bot.Handle(tele.OnPhoto, func(c tele.Context) error {
a.handleTeleMessage(ctx, c.Message())
return nil
})
bot.Handle(tele.OnDocument, func(c tele.Context) error {
a.handleTeleMessage(ctx, c.Message())
return nil
})
util.Go(func() {
bot.Start()
})
util.Go(func() {
<-ctx.Done()
bot.Stop()
})
return nil
}
// Disconnect stops the bot.
func (a *Adapter) Disconnect(_ context.Context) error {
if a.bot != nil {
a.bot.Stop()
}
return nil
}
// Send replies to a private chat.
func (a *Adapter) Send(_ context.Context, to do.Recipient, msg do.OutboundMessage) error {
if a.bot == nil {
return fmt.Errorf("telegram: not connected")
}
chatID, err := strconv.ParseInt(to.ChatID, 10, 64)
if err != nil {
return fmt.Errorf("telegram: chat id: %w", err)
}
_, err = a.bot.Send(tele.ChatID(chatID), msg.Text)
return err
}
func (a *Adapter) handleTeleMessage(ctx context.Context, m *tele.Message) {
if m == nil || m.Chat == nil || m.Chat.Type != tele.ChatPrivate {
return
}
if a.onInbound == nil {
return
}
msg := do.InboundMessage{
ChannelID: a.cfg.ID,
ChannelType: consts.ChannelTypeTelegram,
PlatformUserID: strconv.FormatInt(m.Sender.ID, 10),
ChatID: strconv.FormatInt(m.Chat.ID, 10),
MessageID: strconv.Itoa(m.ID),
Text: m.Text,
}
if m.Caption != "" && msg.Text == "" {
msg.Text = m.Caption
}
if a.bot != nil {
dir, attachments := a.downloadMedia(m)
if dir != "" {
defer func() {
if err := os.RemoveAll(dir); err != nil {
logger.WarnF(ctx, "telegram: 清理临时媒体目录 %q 失败: %v", dir, err)
}
}()
}
msg.Attachments = attachments
}
_ = a.onInbound(ctx, msg)
}
// downloadMedia fetches message media into a scratch directory, returned so the
// caller can remove it once the inbound handler no longer needs the paths.
// An empty dir means nothing was downloaded.
func (a *Adapter) downloadMedia(m *tele.Message) (string, []do.Attachment) {
var files []*tele.File
var names []string
if m.Photo != nil {
files = append(files, m.Photo.MediaFile())
names = append(names, "photo.jpg")
}
if m.Document != nil {
files = append(files, &m.Document.File)
name := m.Document.FileName
if name == "" {
name = "file"
}
names = append(names, name)
}
if len(files) == 0 {
return "", nil
}
dir, err := os.MkdirTemp("", "wg-tg-*")
if err != nil {
return "", []do.Attachment{{Error: err.Error()}}
}
out := make([]do.Attachment, 0, len(files))
for i, f := range files {
path := filepath.Join(dir, names[i])
if err := a.bot.Download(f, path); err != nil {
out = append(out, do.Attachment{FileName: names[i], Error: err.Error()})
continue
}
out = append(out, do.Attachment{Path: path, FileName: names[i]})
}
return dir, out
}
@@ -0,0 +1,78 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package telegram
import (
"Wavelet/plugins/domain/msg_gateway/model/do"
"context"
"testing"
"time"
tele "gopkg.in/telebot.v4"
)
// TestBuildTeleSettingsLongPollWindow 回归:LongPoller.Timeout 是 time.Duration,
// telebot 以 int(timeout/time.Second) 下发给 getUpdates。写成裸整数会被解释为
// 纳秒,令 timeout=0,长轮询退化为对 Bot API 的空转轮询。
func TestBuildTeleSettingsLongPollWindow(t *testing.T) {
pref := buildTeleSettings(do.ChannelConfig{
Credentials: map[string]string{"bot_token": "token"},
Extra: map[string]string{"base_url": "https://tg.example.com/api/"},
})
poller, ok := pref.Poller.(*tele.LongPoller)
if !ok {
t.Fatalf("expected *tele.LongPoller, got %T", pref.Poller)
}
if got := int(poller.Timeout / time.Second); got != 10 {
t.Errorf("getUpdates would receive timeout=%d seconds, want 10", got)
}
if pref.URL != "https://tg.example.com/api" {
t.Errorf("base_url trailing slash should be trimmed, got %q", pref.URL)
}
}
func TestHandleUpdate_DropsGroups(t *testing.T) {
var got int
a := &Adapter{onInbound: func(_ context.Context, _ do.InboundMessage) error {
got++
return nil
}}
a.handleTeleMessage(context.Background(), &tele.Message{
ID: 1,
Text: "hi",
Chat: &tele.Chat{ID: -100, Type: tele.ChatGroup},
Sender: &tele.User{ID: 1},
})
if got != 0 {
t.Fatalf("group must be ignored")
}
}
func TestHandleUpdate_PrivateText(t *testing.T) {
var got do.InboundMessage
a := &Adapter{
cfg: do.ChannelConfig{ID: 7, Type: "telegram"},
onInbound: func(_ context.Context, msg do.InboundMessage) error {
got = msg
return nil
},
}
a.handleTeleMessage(context.Background(), &tele.Message{
ID: 9,
Text: "hi",
Chat: &tele.Chat{ID: 42, Type: tele.ChatPrivate},
Sender: &tele.User{ID: 42},
})
if got.Text != "hi" || got.PlatformUserID != "42" || got.ChannelID != 7 {
t.Fatalf("%+v", got)
}
}
func TestNew_RequiresToken(t *testing.T) {
_, err := New(do.ChannelConfig{}, nil)
if err == nil {
t.Fatal("expected error")
}
}
@@ -0,0 +1,131 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package consts defines constants, sentinel errors, and user-facing error messages
// for the msg_gateway plugin.
package consts
import "errors"
// Channel type and scope constants.
const (
ChannelTypeTelegram = "telegram"
ChannelTypeQQ = "qq"
MessageChannelTypeTelegram = "telegram"
MessageChannelTypeQQ = "qq"
MessageOwnerScopeSystem = "system"
TypeCustom = "custom"
TypeEmail = "email"
TypeTelegram = "telegram"
ChannelCustom = "custom"
ChannelEmail = "email"
ChannelLark = "lark"
ChannelDingTalk = "dingtalk"
ChannelTelegram = "telegram"
ChannelBark = "bark"
ChannelDiscord = "discord"
ChannelSlack = "slack"
ChannelPushover = "pushover"
DefaultLevelInfo = "INFO"
KeyTitle = "title"
KeyContent = "content"
KeyLevel = "level"
// KeyURL represents the URL field key.
KeyURL = "url"
// KeyToken represents the Token field key.
KeyToken = "token"
// KeyOther represents the Other field key.
KeyOther = "other"
// TypeText represents standard text input type.
TypeText = "text"
// TypePassword represents password input type.
TypePassword = "password"
// TypeTextarea represents textarea input type.
TypeTextarea = "textarea"
)
// Task and Schedule identifier constants.
const (
TaskPushNotification = "msg_gateway:push_notification"
TaskCleanupPairingCodes = "msg_gateway:cleanup_pairing_codes"
TaskDispatchBotMsg = "msg_gateway:dispatch_bot_msg"
TaskTypeDispatchBotMsg = "dispatch_bot_msg"
SendNotificationTask = "push:send"
TaskTypeSendNotification = "send_notification"
)
// Pairing code constants.
const (
CodeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
CodeLength = 8
)
// Sentinel errors.
var (
ErrCodeInvalid = errors.New("invalid or expired pairing code")
ErrChannelMismatch = errors.New("pairing code does not match channel")
ErrPlatformAlreadyBound = errors.New("this platform account is already bound")
ErrBindingNotFound = errors.New("binding not found")
ErrBindingForbidden = errors.New("cannot unbind another user's binding")
ErrChannelIDRequired = errors.New("channel_id is required")
ErrChannelDisabled = errors.New("channel is not enabled")
// ErrRecordNotFound maps GORM's missing-row sentinel at the DAO boundary so
// upper layers never import gorm. Its text matches gorm.ErrRecordNotFound verbatim.
ErrRecordNotFound = errors.New("record not found")
// ErrUnsupportedUserLookupField rejects a column name that the DAO is not
// allowed to interpolate into a WHERE clause.
ErrUnsupportedUserLookupField = errors.New("unsupported user lookup field")
)
// User-facing validation and error message constants.
const (
ErrNameRequired = "name is required"
ErrTypeInvalid = "type must be telegram or qq"
ErrTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text
ErrQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text
ErrChannelNotFound = "channel not found"
ErrChannelProbeFailed = "channel probe failed"
ErrBotDispatchTextRequired = "message text is required"
ErrBotChannelNotRegistered = "channel adapter is not registered"
MaskedSecret = "********"
ErrLoginRequired = "login required"
ErrInvalidBindingID = "invalid binding id"
ErrInvalidChannelID = "invalid channel id"
ErrInvalidEventID = "invalid event id"
ErrEventNotFound = "notification event not found"
ErrValidationFailed = "validation failed"
ErrMissingTelegramToken = "missing telegram bot token"
ErrMissingQQCredentials = "missing qq app_id or app_secret" //nolint:gosec // user-facing validation text
ErrQQTokenFetchFailed = "qq token fetch failed" //nolint:gosec // user-facing validation text
ErrQQEmptyToken = "qq returned empty access token" //nolint:gosec // user-facing validation text
ErrTelegramGetMeFailed = "telegram getMe failed"
ErrTelegramNotOK = "telegram returned ok=false"
ErrAPIBaseInvalid = "api_base must start with http:// or https://"
ErrChannelNameExists = "channel name already exists"
ErrChannelNameRequired = "channel name is required"
ErrChannelTypeRequired = "channel type is required"
ErrEventKeyRequired = "event_key is required"
ErrEventAlreadyConfigured = "this notification event is already configured"
ErrTemplateInvalidJSON = "custom template is not a valid JSON format"
ErrEnableWithoutChannels = "cannot enable event without any push channels configured"
ErrEventKeyOrTaskType = "either event_key or task_type must be provided"
ErrUnsupportedEventKey = "unsupported built-in event key"
ErrTaskServiceUnavailable = "task service not available"
ErrUserNotFound = "user not found"
ErrNoAdminUser = "no admin user found"
ErrPayloadRequired = "payload is required"
ErrInvalidJSONFormat = "invalid json format"
ErrParsePayloadFailed = "parse payload failed"
ErrGetPusherFailed = "get pusher failed"
ErrPusherSendFailed = "pusher.Send failed"
)
@@ -0,0 +1,139 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package controller
import (
"Wavelet/pkg/response"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/service"
"net/http"
"strconv"
"github.com/gin-gonic/gin"
)
// ListAdminChannelDefinitions returns form schemas for supported channel types.
// @Summary List message gateway channel definitions
// @Description Returns form field definitions for Telegram and QQ channels
// @Tags admin-message-gateway
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]do.Definition}
// @Router /api/v1/admin/message-gateway/channels/definitions [get]
func ListAdminChannelDefinitions(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(service.ListDefinitions()))
}
// ListAdminChannels lists configured messaging channels with secrets masked.
// @Summary List message gateway channels
// @Description Returns all messaging channels; secrets are masked
// @Tags admin-message-gateway
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]do.ChannelDTO}
// @Router /api/v1/admin/message-gateway/channels [get]
func ListAdminChannels(c *gin.Context) {
rows, err := service.ListChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(rows))
}
func parseAdminChannelID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, consts.ErrInvalidChannelID)
return 0, false
}
return id, true
}
func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
if err.Error() == consts.ErrChannelNotFound {
response.AbortNotFound(c, err.Error())
return
}
fallback(c, err.Error())
}
// CreateAdminChannel creates a messaging channel.
// @Summary Create message gateway channel
// @Description Creates a Telegram or QQ channel with encrypted credentials
// @Tags admin-message-gateway
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body do.CreateChannelRequest true "create body"
// @Success 200 {object} response.Any{data=do.ChannelDTO}
// @Failure 400 {object} response.Any
// @Router /api/v1/admin/message-gateway/channels [post]
func CreateAdminChannel(c *gin.Context) {
handleJSONRequest(c, service.CreateChannel)
}
// UpdateAdminChannel patches a messaging channel. Empty secrets keep the previous values.
// @Summary Update message gateway channel
// @Description Updates a channel; empty secrets keep the current ciphertext
// @Tags admin-message-gateway
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "channel id"
// @Param request body do.UpdateChannelRequest true "update body"
// @Success 200 {object} response.Any{data=do.ChannelDTO}
// @Failure 400 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/admin/message-gateway/channels/{id} [patch]
func UpdateAdminChannel(c *gin.Context) {
handleEntityUpdate(c, parseAdminChannelID, service.UpdateChannel, func(c *gin.Context, err error) {
handleAdminChannelError(c, err, response.AbortBadRequest)
})
}
// DeleteAdminChannel removes a channel and its bindings/pairing codes.
// @Summary Delete message gateway channel
// @Description Deletes a channel and cascaded bindings and pairing codes
// @Tags admin-message-gateway
// @Produce json
// @Security SessionCookie
// @Param id path int true "channel id"
// @Success 200 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/admin/message-gateway/channels/{id} [delete]
func DeleteAdminChannel(c *gin.Context) {
id, ok := parseAdminChannelID(c)
if !ok {
return
}
if err := service.DeleteChannel(c.Request.Context(), id); err != nil {
handleAdminChannelError(c, err, response.AbortInternal)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// TestAdminChannel probes stored credentials (Telegram getMe or QQ token).
// @Summary Test message gateway channel
// @Description Probes stored credentials without returning secrets
// @Tags admin-message-gateway
// @Produce json
// @Security SessionCookie
// @Param id path int true "channel id"
// @Success 200 {object} response.Any
// @Failure 400 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/admin/message-gateway/channels/{id}/test [post]
func TestAdminChannel(c *gin.Context) {
id, ok := parseAdminChannelID(c)
if !ok {
return
}
if err := service.ProbeChannel(c.Request.Context(), id); err != nil {
handleAdminChannelError(c, err, response.AbortBadRequest)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,146 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package controller
import (
"Wavelet/pkg/response"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/model/do"
"Wavelet/plugins/domain/msg_gateway/service"
"errors"
"net/http"
"strconv"
"github.com/gin-gonic/gin"
)
// ListPushChannelDefinitions 获取各种消息通道的表单配置定义列表
// @Summary 获取所有消息通道配置字段定义
// @Description 返回系统支持的所有消息通道类型的动态表单定义,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any "通道配置定义列表"
// @Router /api/v1/admin/push/channels/definitions [get]
func ListPushChannelDefinitions(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(do.ListPushDefinitions()))
}
// ListPushChannels 获取消息通道列表
// @Summary 获取所有消息通道
// @Description 返回系统配置的所有消息通道列表,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]entity.PushChannel} "消息通道列表"
// @Router /api/v1/admin/push/channels [get]
func ListPushChannels(c *gin.Context) {
channels, err := service.ListPushChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(channels))
}
// parsePushChannelID reads the path identifier of a push channel.
func parsePushChannelID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, consts.ErrInvalidChannelID)
return 0, false
}
return id, true
}
// handlePushChannelNotFoundError maps a missing channel row to 404, others to fallback.
func handlePushChannelNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
if errors.Is(err, consts.ErrRecordNotFound) {
response.AbortNotFound(c, consts.ErrChannelNotFound)
return
}
fallback(c, err.Error())
}
// CreatePushChannel 创建消息通道
// @Summary 创建消息通道
// @Description 新建一个消息通道配置,需要管理员权限
// @Tags admin-push
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body do.CreatePushChannelRequest true "创建参数"
// @Success 200 {object} response.Any{data=entity.PushChannel} "创建成功"
// @Router /api/v1/admin/push/channels [post]
func CreatePushChannel(c *gin.Context) {
handleJSONRequest(c, service.CreatePushChannel)
}
// UpdatePushChannel 更新消息通道
// @Summary 更新消息通道
// @Description 修改消息通道配置,需要管理员权限
// @Tags admin-push
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path uint64 true "通道ID"
// @Param request body do.UpdatePushChannelRequest true "更新参数"
// @Success 200 {object} response.Any{data=entity.PushChannel} "更新成功"
// @Router /api/v1/admin/push/channels/{id} [put]
func UpdatePushChannel(c *gin.Context) {
handleEntityUpdate(c, parsePushChannelID, service.UpdatePushChannel, func(c *gin.Context, err error) {
handlePushChannelNotFoundError(c, err, response.AbortInternal)
})
}
// DeletePushChannel 删除消息通道
// @Summary 删除消息通道
// @Description 根据ID删除消息通道,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Param id path uint64 true "通道ID"
// @Success 200 {object} response.Any "删除成功"
// @Router /api/v1/admin/push/channels/{id} [delete]
func DeletePushChannel(c *gin.Context) {
id, ok := parsePushChannelID(c)
if !ok {
return
}
if err := service.DeletePushChannel(c.Request.Context(), id); err != nil {
handlePushChannelNotFoundError(c, err, response.AbortInternal)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// TestPushChannel 测试通道连通性
// @Summary 测试通道连通性
// @Description 触发一次临时的或现有的通道连通性推送测试,需要管理员权限
// @Tags admin-push
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body do.TestPushChannelRequest true "测试参数"
// @Success 200 {object} response.Any "测试触发成功"
// @Router /api/v1/admin/push/channels/test [post]
func TestPushChannel(c *gin.Context) {
var req do.TestPushChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
payload, err := service.PreparePushChannelTest(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := service.EnqueuePushTask(c.Request.Context(), payload); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,218 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package controller
import (
"Wavelet/pkg/response"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/model/do"
"Wavelet/plugins/domain/msg_gateway/service"
"errors"
"net/http"
"strconv"
"github.com/gin-gonic/gin"
)
// ListPushEvents 获取通知事件列表
// @Summary 获取所有通知事件
// @Description 返回系统配置的通知事件列表,包括预置和自定义事件,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]entity.PushEvent} "通知事件列表"
// @Router /api/v1/admin/push/events [get]
func ListPushEvents(c *gin.Context) {
ctx := c.Request.Context()
events, err := service.ListPushEvents(ctx)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(events))
}
// ListBuiltInPushEvents 获取内置通知事件列表
// @Summary 获取所有内置通知事件
// @Description 返回系统定义的所有内置通知事件元数据,供前端下拉框选择,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any "内置通知事件列表"
// @Router /api/v1/admin/push/events/builtin [get]
func ListBuiltInPushEvents(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(service.GetBuiltInEvents()))
}
// parsePushEventID reads the path identifier of a push event.
func parsePushEventID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, consts.ErrInvalidEventID)
return 0, false
}
return id, true
}
// handlePushEventNotFoundError maps a missing event row to 404, others to fallback.
func handlePushEventNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
if errors.Is(err, consts.ErrRecordNotFound) {
response.AbortNotFound(c, consts.ErrEventNotFound)
return
}
fallback(c, err.Error())
}
// CreatePushEvent 创建通知事件
// @Summary 创建通知事件
// @Description 绑定系统内置事件或异步任务、推送渠道、接收目标并创建通知事件配置,需要管理员权限
// @Tags admin-push
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body do.CreatePushEventRequest true "创建参数"
// @Success 200 {object} response.Any{data=entity.PushEvent} "创建成功"
// @Router /api/v1/admin/push/events [post]
func CreatePushEvent(c *gin.Context) {
handleJSONRequest(c, service.CreatePushEvent)
}
// DeletePushEvent 删除通知事件配置
// @Summary 删除通知事件配置
// @Description 删除数据库中的特定通知事件配置,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Param id path int true "事件 ID"
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Router /api/v1/admin/push/events/{id} [delete]
func DeletePushEvent(c *gin.Context) {
id, ok := parsePushEventID(c)
if !ok {
return
}
if err := service.DeletePushEvent(c.Request.Context(), id); err != nil {
handlePushEventNotFoundError(c, err, response.AbortInternal)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// UpdatePushEvent 更新通知事件
// @Summary 更新通知事件
// @Description 更新已有通知事件的推送渠道、接收目标和内容模板,需要管理员权限
// @Tags admin-push
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "事件 ID"
// @Param request body do.UpdatePushEventRequest true "更新参数"
// @Success 200 {object} response.Any{data=string} "修改成功"
// @Router /api/v1/admin/push/events/{id} [put]
func UpdatePushEvent(c *gin.Context) {
id, ok := parsePushEventID(c)
if !ok {
return
}
var req do.UpdatePushEventRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := service.UpdatePushEvent(c.Request.Context(), id, req); err != nil {
handlePushEventNotFoundError(c, err, response.AbortBadRequest)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// TogglePushEvent 快捷切换通知事件启用状态
// @Summary 快捷切换通知事件启用状态
// @Description 启用或禁用指定的通知事件
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Param id path int true "事件 ID"
// @Success 200 {object} response.Any{data=string} "切换成功"
// @Router /api/v1/admin/push/events/{id}/toggle [post]
func TogglePushEvent(c *gin.Context) {
id, ok := parsePushEventID(c)
if !ok {
return
}
enabled, err := service.TogglePushEvent(c.Request.Context(), id)
if err != nil {
handlePushEventNotFoundError(c, err, response.AbortBadRequest)
return
}
c.JSON(http.StatusOK, response.OK(enabled))
}
// ListPushHistories 分页获取通知推送历史
// @Summary 分页获取通知推送历史
// @Description 返回分页的通知历史日志数据,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Param page query int false "当前页码"
// @Param page_size query int false "分页大小"
// @Param event_key query string false "过滤事件名称"
// @Param status query string false "过滤发送状态"
// @Success 200 {object} response.Any "推送历史列表"
// @Router /api/v1/admin/push/histories [get]
func ListPushHistories(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
if page < 1 {
page = 1
}
if pageSize < 1 {
pageSize = 20
}
total, results, err := service.ListPushHistories(c.Request.Context(), do.PushHistoryListFilter{
EventKey: c.Query("event_key"),
Channel: c.Query("channel"),
Status: c.Query("status"),
Page: page,
PageSize: pageSize,
})
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(map[string]any{
"total": total,
"results": results,
}))
}
// TestPush 测试推送通道发送
// @Summary 测试推送通道发送
// @Description 接收临时通知渠道配置并在本地同步调用 Pusher.Send 发送测试消息
// @Tags admin-push
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body do.TestPushRequest true "测试请求体"
// @Success 200 {object} response.Any{data=string} "测试成功"
// @Router /api/v1/admin/push/test [post]
func TestPush(c *gin.Context) {
var req do.TestPushRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := service.RunPushTest(c.Request.Context(), req.Config, req.Target); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,64 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package controller provides HTTP endpoints for msg_gateway.
package controller
import (
"Wavelet/core/extpoints"
"github.com/gin-gonic/gin"
)
// RegisterUserRoutes mounts user-facing message gateway endpoints.
func RegisterUserRoutes(r extpoints.RouterExtension, loginMW gin.HandlerFunc) {
mg := r.Group("/message-gateway", loginMW)
{
mg.GET("/channels", ListChannels)
mg.GET("/bindings", ListBindings)
mg.POST("/bindings", BindBinding)
mg.DELETE("/bindings/:id", UnbindBinding)
}
}
// RegisterAdminRoutes mounts admin message-gateway APIs under /admin.
func RegisterAdminRoutes(adminRouter extpoints.RouterExtension, loginMW, adminMW gin.HandlerFunc) {
g := adminRouter.Group("/message-gateway", loginMW, adminMW)
{
g.GET("/channels/definitions", ListAdminChannelDefinitions)
g.GET("/channels", ListAdminChannels)
g.POST("/channels", CreateAdminChannel)
g.PATCH("/channels/:id", UpdateAdminChannel)
g.DELETE("/channels/:id", DeleteAdminChannel)
g.POST("/channels/:id/test", TestAdminChannel)
}
}
// RegisterAdminPushRoutes mounts admin push notification APIs under /admin.
func RegisterAdminPushRoutes(adminRouter extpoints.RouterExtension, loginMW, adminMW gin.HandlerFunc) {
adminPushGroup := adminRouter.Group("/push", loginMW, adminMW)
{
events := adminPushGroup.Group("/events")
{
events.GET("", ListPushEvents)
events.GET("/builtin", ListBuiltInPushEvents)
events.POST("", CreatePushEvent)
events.PUT("/:id", UpdatePushEvent)
events.DELETE("/:id", DeletePushEvent)
events.POST("/:id/toggle", TogglePushEvent)
}
adminPushGroup.GET("/histories", ListPushHistories)
adminPushGroup.POST("/test", TestPush)
channels := adminPushGroup.Group("/channels")
{
channels.GET("/definitions", ListPushChannelDefinitions)
channels.GET("", ListPushChannels)
channels.POST("", CreatePushChannel)
channels.PUT("/:id", UpdatePushChannel)
channels.DELETE("/:id", DeletePushChannel)
channels.POST("/test", TestPushChannel)
}
}
}
@@ -0,0 +1,181 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package controller
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/model/do"
"Wavelet/plugins/domain/msg_gateway/service"
"context"
"errors"
"net/http"
"strconv"
"github.com/gin-gonic/gin"
)
func currentUser(c *gin.Context) (*contracts.UserDTO, bool) {
return ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
}
// handleJSONRequest binds a JSON body, runs the service use case and writes the
// standard success envelope; any service error surfaces as a bad request.
func handleJSONRequest[Req any, Res any](c *gin.Context, handler func(ctx context.Context, req Req) (Res, error)) {
var req Req
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
res, err := handler(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(res))
}
// handleEntityUpdate resolves a path identifier plus JSON body, runs the updater
// use case and writes the success envelope; error classification is delegated to onErr.
func handleEntityUpdate[Req any, Res any](
c *gin.Context,
parseID func(*gin.Context) (uint64, bool),
updater func(ctx context.Context, id uint64, req Req) (Res, error),
onErr func(*gin.Context, error),
) {
id, ok := parseID(c)
if !ok {
return
}
var req Req
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := updater(c.Request.Context(), id, req)
if err != nil {
onErr(c, err)
return
}
c.JSON(http.StatusOK, response.OK(dto))
}
// ListChannels lists enabled channels a user can bind.
// @Summary List enabled messaging channels
// @Description Returns enabled system bots the current user can pair with
// @Tags message-gateway
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]do.PublicChannelDTO}
// @Failure 401 {object} response.Any
// @Router /api/v1/message-gateway/channels [get]
func ListChannels(c *gin.Context) {
if user, ok := currentUser(c); !ok || user == nil {
response.AbortUnauthorized(c, consts.ErrLoginRequired)
return
}
rows, err := service.ListEnabledPublicChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(rows))
}
// ListBindings lists the current user's bot bindings.
// @Summary List message gateway bindings
// @Description Returns the current user's bound messaging channels
// @Tags message-gateway
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]do.BindingDTO}
// @Failure 401 {object} response.Any
// @Router /api/v1/message-gateway/bindings [get]
func ListBindings(c *gin.Context) {
user, ok := currentUser(c)
if !ok || user == nil {
response.AbortUnauthorized(c, consts.ErrLoginRequired)
return
}
rows, err := service.ListUserBindings(c.Request.Context(), user.ID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(rows))
}
// BindBinding consumes a pairing code and binds the platform identity.
// @Summary Bind a messaging channel
// @Description Binds the current user to a platform identity using a one-time pairing code
// @Tags message-gateway
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body do.BindRequest true "bind body"
// @Success 200 {object} response.Any{data=do.BindingDTO}
// @Failure 400 {object} response.Any
// @Failure 409 {object} response.Any
// @Router /api/v1/message-gateway/bindings [post]
func BindBinding(c *gin.Context) {
user, ok := currentUser(c)
if !ok || user == nil {
response.AbortUnauthorized(c, consts.ErrLoginRequired)
return
}
var req do.BindRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := service.BindChannel(c.Request.Context(), user.ID, req)
if err != nil {
if errors.Is(err, consts.ErrPlatformAlreadyBound) {
response.AbortConflict(c, err.Error())
return
}
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(dto))
}
// UnbindBinding removes the current user's binding.
// @Summary Unbind a messaging channel
// @Description Removes a binding owned by the current user
// @Tags message-gateway
// @Produce json
// @Security SessionCookie
// @Param id path int true "binding id"
// @Success 200 {object} response.Any
// @Failure 403 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/message-gateway/bindings/{id} [delete]
func UnbindBinding(c *gin.Context) {
user, ok := currentUser(c)
if !ok || user == nil {
response.AbortUnauthorized(c, consts.ErrLoginRequired)
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, consts.ErrInvalidBindingID)
return
}
if err := service.UnbindChannel(c.Request.Context(), user.ID, id); err != nil {
if errors.Is(err, consts.ErrBindingNotFound) {
response.AbortNotFound(c, err.Error())
return
}
if errors.Is(err, consts.ErrBindingForbidden) {
response.AbortForbidden(c, err.Error())
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,208 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dao provides database persistence and caching for the msg_gateway plugin.
package dao
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/idgen"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/model/entity"
"context"
"errors"
"sync"
"time"
"gorm.io/gorm"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
)
// SetDBServiceForTest injects a DBService for tests. Production wiring must use Apply.
func SetDBServiceForTest(s contracts.DBService) {
SetDBService(s)
}
// SetDBService sets the database service singleton.
func SetDBService(s contracts.DBService) {
dbMu.Lock()
defer dbMu.Unlock()
dbSvc = s
}
// GetDB resolves the persistence handle for the current call, preferring an
// explicitly injected *core.Context before falling back to the plugin singleton.
func GetDB(ctx context.Context) *gorm.DB {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
return s.DB(ctx)
}
}
dbMu.RLock()
s := dbSvc
dbMu.RUnlock()
if s != nil {
return s.DB(ctx)
}
return nil
}
// mapNotFound translates GORM's missing-row sentinel into the plugin-level
// consts.ErrRecordNotFound so the service and controller layers stay free of gorm imports.
func mapNotFound(err error) error {
if errors.Is(err, gorm.ErrRecordNotFound) {
return consts.ErrRecordNotFound
}
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
}
@@ -0,0 +1,358 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package dao
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/model/do"
"Wavelet/plugins/domain/msg_gateway/model/entity"
"context"
"errors"
"fmt"
"sync"
"time"
"gorm.io/gorm"
)
const (
activePushChannelCacheTTL = 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.
func ListPushChannelsRecord(ctx context.Context) ([]entity.PushChannel, error) {
var channels []entity.PushChannel
if err := GetDB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
return nil, err
}
return channels, nil
}
// GetPushChannelByIDRecord loads a push channel by primary key.
func GetPushChannelByIDRecord(ctx context.Context, id uint64) (entity.PushChannel, error) {
var channel entity.PushChannel
if err := GetDB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
return entity.PushChannel{}, mapNotFound(err)
}
return channel, nil
}
// GetPushChannelByNameRecord loads a push channel by its unique name.
func GetPushChannelByNameRecord(ctx context.Context, name string) (*entity.PushChannel, error) {
var channel entity.PushChannel
if err := GetDB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
return nil, mapNotFound(err)
}
return &channel, nil
}
// CountPushChannelsByNameRecord returns how many channels share the given name.
func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, error) {
var count int64
if err := GetDB(ctx).Model(&entity.PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CreatePushChannelRecord persists a new channel and invalidates cache.
func CreatePushChannelRecord(ctx context.Context, channel *entity.PushChannel) error {
if err := GetDB(ctx).Create(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
// SavePushChannelRecord updates a channel and invalidates cache.
func SavePushChannelRecord(ctx context.Context, channel *entity.PushChannel) error {
if err := GetDB(ctx).Save(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
// DeletePushChannelRecord removes a channel and invalidates cache.
func DeletePushChannelRecord(ctx context.Context, channel *entity.PushChannel) error {
if err := GetDB(ctx).Delete(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
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 err := cache.Get(ctx, cacheKey, &val); err == nil {
return &val, nil
}
}
db := GetDB(ctx)
if err := query(db, &val); err != nil {
return nil, err
}
if cache := GetCache(ctx); cache != nil {
_ = cache.Set(ctx, cacheKey, val, ttl)
}
return &val, nil
}
// GetActivePushChannelByName loads an enabled push channel, preferring the cache layer.
func GetActivePushChannelByName(ctx context.Context, name string) (*entity.PushChannel, error) {
channel, err := getCachedOrQuery(ctx, "push:channel:active:"+name, activePushChannelCacheTTL, func(db *gorm.DB, dest *entity.PushChannel) error {
return db.Where("name = ? AND enabled = ?", name, true).First(dest).Error
})
if err != nil {
return nil, mapNotFound(err)
}
return channel, nil
}
// DeleteActivePushChannelCache drops the cached enabled-channel entry.
func DeleteActivePushChannelCache(ctx context.Context, name string) {
if cache := GetCache(ctx); cache != nil {
_ = cache.Delete(ctx, "push:channel:active:"+name)
}
}
// ListPushEventsRecord returns all push events ordered by creation time descending.
func ListPushEventsRecord(ctx context.Context) ([]entity.PushEvent, error) {
var events []entity.PushEvent
if err := GetDB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
return nil, err
}
return events, nil
}
// GetPushEventByIDRecord loads a push event by primary key.
func GetPushEventByIDRecord(ctx context.Context, id uint64) (entity.PushEvent, error) {
var event entity.PushEvent
if err := GetDB(ctx).First(&event, id).Error; err != nil {
return entity.PushEvent{}, mapNotFound(err)
}
return event, nil
}
// GetPushEventByKeyRecord loads a push event by event key.
func GetPushEventByKeyRecord(ctx context.Context, key string) (entity.PushEvent, error) {
var event entity.PushEvent
if err := GetDB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
return entity.PushEvent{}, mapNotFound(err)
}
return event, nil
}
// CountPushEventsByKeyRecord returns how many events use the given event key.
func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error) {
var count int64
if err := GetDB(ctx).Model(&entity.PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CreatePushEventRecord persists a new push event and invalidates cache.
func CreatePushEventRecord(ctx context.Context, event *entity.PushEvent) error {
if err := GetDB(ctx).Create(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// SavePushEventRecord updates a push event and invalidates cache.
func SavePushEventRecord(ctx context.Context, event *entity.PushEvent) error {
if err := GetDB(ctx).Save(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// UpdatePushEventEnabledRecord toggles the enabled flag for a push event.
func UpdatePushEventEnabledRecord(ctx context.Context, event *entity.PushEvent, enabled bool) error {
event.Enabled = enabled
if err := GetDB(ctx).Model(event).Update("enabled", enabled).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// DeletePushEventRecord removes a push event and invalidates cache.
func DeletePushEventRecord(ctx context.Context, event *entity.PushEvent) error {
if err := GetDB(ctx).Delete(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// ListActivePushEventsByTaskTypeRecord returns enabled events bound to a task type.
func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]entity.PushEvent, error) {
var events []entity.PushEvent
if err := GetDB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil {
return nil, err
}
return events, nil
}
// GetActivePushEventByKey loads an enabled push event, preferring the cache layer.
func GetActivePushEventByKey(ctx context.Context, key string) (*entity.PushEvent, error) {
event, err := getCachedOrQuery(ctx, "push:event:active:"+key, activePushEventCacheTTL, func(db *gorm.DB, dest *entity.PushEvent) error {
return db.Where("event_key = ? AND enabled = ?", key, true).First(dest).Error
})
if err != nil {
return nil, mapNotFound(err)
}
return event, nil
}
// DeleteActivePushEventCache drops the cached enabled-event entry.
func DeleteActivePushEventCache(ctx context.Context, key string) {
if cache := GetCache(ctx); cache != nil {
_ = cache.Delete(ctx, "push:event:active:"+key)
}
}
// ListPushHistoriesRecord returns paginated push history records.
func ListPushHistoriesRecord(ctx context.Context, filter do.PushHistoryListFilter) (int64, []entity.PushHistory, error) {
query := GetDB(ctx).Model(&entity.PushHistory{}).Order("created_at DESC")
if filter.EventKey != "" {
query = query.Where("event_key = ?", filter.EventKey)
}
if filter.Status != "" {
query = query.Where("status = ?", filter.Status)
}
var total int64
if err := query.Count(&total).Error; err != nil {
return 0, nil, err
}
var results []entity.PushHistory
offset := (filter.Page - 1) * filter.PageSize
if err := query.Offset(offset).Limit(filter.PageSize).Find(&results).Error; err != nil {
return 0, nil, err
}
return total, results, nil
}
// CreatePushHistoryRecord persists a push history audit record.
func CreatePushHistoryRecord(ctx context.Context, history *entity.PushHistory) error {
return GetDB(ctx).Create(history).Error
}
// PushHistoryQuery returns a scoped query builder for push histories.
func PushHistoryQuery(ctx context.Context) *gorm.DB {
return GetDB(ctx).Model(&entity.PushHistory{})
}
// smtpConfigKeys are the system-config rows backing the built-in email channel.
var smtpConfigKeys = []string{"smtp_host", "smtp_port", "smtp_username", "smtp_password"}
// LoadSMTPConfigRecord reads the SMTP settings in one query.
func LoadSMTPConfigRecord(ctx context.Context) (do.SMTPConfig, error) {
db := GetDB(ctx)
if db == nil {
return do.SMTPConfig{}, errors.New("database not available")
}
var rows []struct {
Key string
Value string
}
if err := db.Table("w_system_configs").
Select("key", "value").
Where("key IN ?", smtpConfigKeys).
Find(&rows).Error; err != nil {
return do.SMTPConfig{}, fmt.Errorf("read smtp system configs: %w", err)
}
var cfg do.SMTPConfig
for _, row := range rows {
switch row.Key {
case "smtp_host":
cfg.Host = row.Value
case "smtp_port":
cfg.Port = row.Value
case "smtp_username":
cfg.Username = row.Value
case "smtp_password":
cfg.Password = row.Value
}
}
return cfg, nil
}
// userLookupColumns allow-lists the columns FindUserByFieldRecord may filter on.
var userLookupColumns = map[string]struct{}{
"id": {},
"username": {},
}
// FindUserByFieldRecord is the user lookup fallback for when the UserService
// contract is not wired yet. field must be one of userLookupColumns.
func FindUserByFieldRecord(ctx context.Context, field string, value any) (*contracts.UserDTO, error) {
if _, ok := userLookupColumns[field]; !ok {
return nil, consts.ErrUnsupportedUserLookupField
}
db := GetDB(ctx)
if db == nil {
return nil, consts.ErrRecordNotFound
}
var user contracts.UserDTO
if err := db.Table("w_users").Where(field+" = ?", value).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
// FindFirstAdminUserRecord is the admin lookup fallback for when the UserService
// contract is not wired yet.
func FindFirstAdminUserRecord(ctx context.Context) (*contracts.UserDTO, error) {
db := GetDB(ctx)
if db == nil {
return nil, consts.ErrRecordNotFound
}
var adminUser contracts.UserDTO
if err := db.Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
return nil, err
}
return &adminUser, nil
}
@@ -0,0 +1,131 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package dao_test
import (
"Wavelet/pkg/testhelper"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/dao"
"context"
"errors"
"testing"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
// stubDBService satisfies contracts.DBService over a test database handle.
type stubDBService struct{ db *gorm.DB }
func (s stubDBService) GORM() *gorm.DB { return s.db }
func (s stubDBService) DB(_ context.Context) *gorm.DB { return s.db }
func (s stubDBService) Named(_ string) *gorm.DB { return s.db }
// TestFindUserByFieldRecordRejectsUnlistedColumns pins the column allow-list. The
// lookup column is interpolated into SQL, so an unlisted name must be refused before
// any query is built rather than trusted because call sites happen to pass literals.
func TestFindUserByFieldRecordRejectsUnlistedColumns(t *testing.T) {
db, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
if err := db.Table("w_users").Create(map[string]any{"id": 77, "username": "seeded"}).Error; err != nil {
t.Fatalf("seed user failed: %v", err)
}
dao.SetDBServiceForTest(stubDBService{db: db})
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
ctx := context.Background()
user, err := dao.FindUserByFieldRecord(ctx, "username", "seeded")
if err != nil {
t.Fatalf("allowlisted lookup by username failed: %v", err)
}
if user.ID != 77 {
t.Errorf("allowlisted lookup returned ID %d, want 77", user.ID)
}
if _, err := dao.FindUserByFieldRecord(ctx, "id", uint64(77)); err != nil {
t.Errorf("allowlisted lookup by id failed: %v", err)
}
cases := []struct {
name string
field string
}{
{"tautology injection", `username = '' OR 1=1 --`},
{"stacked statement", "id; DROP TABLE w_users"},
{"column outside allow-list", "password"},
{"empty field", ""},
}
for _, tc := range cases {
if _, err := dao.FindUserByFieldRecord(ctx, tc.field, "seeded"); !errors.Is(err, consts.ErrUnsupportedUserLookupField) {
t.Errorf("%s: got err %v, want ErrUnsupportedUserLookupField", tc.name, err)
}
}
var remaining int64
if err := db.Table("w_users").Count(&remaining).Error; err != nil || remaining != 1 {
t.Fatalf("w_users damaged by rejected lookups: count=%d err=%v", remaining, err)
}
}
// smtpTestValues are the four system-config rows the built-in email channel reads.
var smtpTestValues = map[string]string{
"smtp_host": "mail.example.test",
"smtp_port": "465",
"smtp_username": "notify@example.test",
"smtp_password": "s3cret-value",
}
// TestLoadSMTPConfigRecordMapsEveryKey guards the single-query rewrite: every field
// must still be filled from its own row.
func TestLoadSMTPConfigRecordMapsEveryKey(t *testing.T) {
db, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
keys := make([]string, 0, len(smtpTestValues))
for key := range smtpTestValues {
keys = append(keys, key)
}
if err := db.Table("w_system_configs").Where("key IN ?", keys).Delete(map[string]any{}).Error; err != nil {
t.Fatalf("clear smtp rows: %v", err)
}
for _, key := range keys {
row := map[string]any{"key": key, "value": smtpTestValues[key], "type": "system"}
if err := db.Table("w_system_configs").Create(row).Error; err != nil {
t.Fatalf("seed %s: %v", key, err)
}
}
dao.SetDBServiceForTest(stubDBService{db: db})
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
cfg, err := dao.LoadSMTPConfigRecord(context.Background())
if err != nil {
t.Fatalf("LoadSMTPConfigRecord: %v", err)
}
if cfg.Host != smtpTestValues["smtp_host"] || cfg.Port != smtpTestValues["smtp_port"] ||
cfg.Username != smtpTestValues["smtp_username"] || cfg.Password != smtpTestValues["smtp_password"] {
t.Errorf("got %+v, want every SMTP field mapped from its own row", cfg)
}
}
// TestLoadSMTPConfigRecordSurfacesReadFailure pins the actual defect: a read that
// fails used to be discarded, returning four blank strings that callers could only
// interpret as "SMTP was never configured", so the notification was dropped silently.
func TestLoadSMTPConfigRecordSurfacesReadFailure(t *testing.T) {
bare, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open bare sqlite: %v", err)
}
dao.SetDBServiceForTest(stubDBService{db: bare})
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
if _, err := dao.LoadSMTPConfigRecord(context.Background()); err == nil {
t.Fatal("LoadSMTPConfigRecord returned nil error although the config table cannot be read")
}
}
@@ -0,0 +1,93 @@
-- +goose Up
-- +goose StatementBegin
CREATE TABLE IF NOT EXISTS w_message_channels (
id BIGINT PRIMARY KEY,
name VARCHAR(128) NOT NULL,
type VARCHAR(32) NOT NULL,
owner_scope VARCHAR(16) NOT NULL DEFAULT 'system',
owner_id BIGINT NULL,
enabled BOOLEAN NOT NULL DEFAULT TRUE,
credentials TEXT NOT NULL DEFAULT '',
extra TEXT NOT NULL DEFAULT '',
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_message_channels_type ON w_message_channels (type);
CREATE TABLE IF NOT EXISTS w_message_bindings (
id BIGINT PRIMARY KEY,
user_id BIGINT NOT NULL,
channel_id BIGINT NOT NULL,
platform_user_id VARCHAR(128) NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_message_bindings_channel_platform
ON w_message_bindings (channel_id, platform_user_id);
CREATE INDEX IF NOT EXISTS idx_w_message_bindings_user ON w_message_bindings (user_id);
CREATE TABLE IF NOT EXISTS w_message_pairing_codes (
code VARCHAR(16) PRIMARY KEY,
channel_id BIGINT NOT NULL,
platform_user_id VARCHAR(128) NOT NULL,
expires_at TIMESTAMPTZ NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_message_pairing_lookup
ON w_message_pairing_codes (channel_id, platform_user_id);
CREATE TABLE IF NOT EXISTS w_push_events (
id BIGINT PRIMARY KEY,
event_key VARCHAR(80) NOT NULL,
name VARCHAR(100) NOT NULL,
task_type VARCHAR(100) NOT NULL DEFAULT '',
channels TEXT NOT NULL DEFAULT '',
targets TEXT NOT NULL DEFAULT '',
template TEXT NOT NULL DEFAULT '',
enabled BOOLEAN NOT NULL DEFAULT FALSE,
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_push_events_key ON w_push_events(event_key);
CREATE INDEX IF NOT EXISTS idx_w_push_events_enabled ON w_push_events(enabled);
CREATE INDEX IF NOT EXISTS idx_w_push_events_task_type ON w_push_events(task_type);
CREATE TABLE IF NOT EXISTS w_push_channels (
id BIGINT PRIMARY KEY,
name VARCHAR(80) NOT NULL,
description VARCHAR(255) NOT NULL DEFAULT '',
type VARCHAR(50) NOT NULL DEFAULT 'custom',
token VARCHAR(100) NOT NULL DEFAULT '',
url TEXT NOT NULL DEFAULT '',
other TEXT NOT NULL DEFAULT '',
enabled BOOLEAN NOT NULL DEFAULT TRUE,
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_push_channels_name ON w_push_channels(name);
CREATE INDEX IF NOT EXISTS idx_w_push_channels_enabled ON w_push_channels(enabled);
CREATE TABLE IF NOT EXISTS w_push_histories (
id BIGINT PRIMARY KEY,
event_key VARCHAR(80) NOT NULL,
channel VARCHAR(50) NOT NULL,
target VARCHAR(255) NOT NULL,
title VARCHAR(255) NOT NULL,
content TEXT NOT NULL,
level VARCHAR(20) NOT NULL,
status VARCHAR(20) NOT NULL,
error_msg TEXT NOT NULL DEFAULT '',
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_push_histories_event ON w_push_histories(event_key);
CREATE INDEX IF NOT EXISTS idx_w_push_histories_created ON w_push_histories(created_at);
-- +goose StatementEnd
-- +goose Down
-- +goose StatementBegin
DROP TABLE IF EXISTS w_push_histories;
DROP TABLE IF EXISTS w_push_channels;
DROP TABLE IF EXISTS w_push_events;
DROP TABLE IF EXISTS w_message_pairing_codes;
DROP TABLE IF EXISTS w_message_bindings;
DROP TABLE IF EXISTS w_message_channels;
-- +goose StatementEnd
@@ -0,0 +1,93 @@
-- +goose Up
-- +goose StatementBegin
CREATE TABLE IF NOT EXISTS w_message_channels (
id BIGINT PRIMARY KEY,
name VARCHAR(128) NOT NULL,
type VARCHAR(32) NOT NULL,
owner_scope VARCHAR(16) NOT NULL DEFAULT 'system',
owner_id BIGINT NULL,
enabled BOOLEAN NOT NULL DEFAULT 1,
credentials TEXT NOT NULL DEFAULT '',
extra TEXT NOT NULL DEFAULT '',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_message_channels_type ON w_message_channels (type);
CREATE TABLE IF NOT EXISTS w_message_bindings (
id BIGINT PRIMARY KEY,
user_id BIGINT NOT NULL,
channel_id BIGINT NOT NULL,
platform_user_id VARCHAR(128) NOT NULL,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_message_bindings_channel_platform
ON w_message_bindings (channel_id, platform_user_id);
CREATE INDEX IF NOT EXISTS idx_w_message_bindings_user ON w_message_bindings (user_id);
CREATE TABLE IF NOT EXISTS w_message_pairing_codes (
code VARCHAR(16) PRIMARY KEY,
channel_id BIGINT NOT NULL,
platform_user_id VARCHAR(128) NOT NULL,
expires_at DATETIME NOT NULL,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_message_pairing_lookup
ON w_message_pairing_codes (channel_id, platform_user_id);
CREATE TABLE IF NOT EXISTS w_push_events (
id BIGINT PRIMARY KEY,
event_key VARCHAR(80) NOT NULL,
name VARCHAR(100) NOT NULL,
task_type VARCHAR(100) NOT NULL DEFAULT '',
channels TEXT NOT NULL DEFAULT '',
targets TEXT NOT NULL DEFAULT '',
template TEXT NOT NULL DEFAULT '',
enabled BOOLEAN NOT NULL DEFAULT 0,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_push_events_key ON w_push_events(event_key);
CREATE INDEX IF NOT EXISTS idx_w_push_events_enabled ON w_push_events(enabled);
CREATE INDEX IF NOT EXISTS idx_w_push_events_task_type ON w_push_events(task_type);
CREATE TABLE IF NOT EXISTS w_push_channels (
id BIGINT PRIMARY KEY,
name VARCHAR(80) NOT NULL,
description VARCHAR(255) NOT NULL DEFAULT '',
type VARCHAR(50) NOT NULL DEFAULT 'custom',
token VARCHAR(100) NOT NULL DEFAULT '',
url TEXT NOT NULL DEFAULT '',
other TEXT NOT NULL DEFAULT '',
enabled BOOLEAN NOT NULL DEFAULT 1,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_push_channels_name ON w_push_channels(name);
CREATE INDEX IF NOT EXISTS idx_w_push_channels_enabled ON w_push_channels(enabled);
CREATE TABLE IF NOT EXISTS w_push_histories (
id BIGINT PRIMARY KEY,
event_key VARCHAR(80) NOT NULL,
channel VARCHAR(50) NOT NULL,
target VARCHAR(255) NOT NULL,
title VARCHAR(255) NOT NULL,
content TEXT NOT NULL,
level VARCHAR(20) NOT NULL,
status VARCHAR(20) NOT NULL,
error_msg TEXT NOT NULL DEFAULT '',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_push_histories_event ON w_push_histories(event_key);
CREATE INDEX IF NOT EXISTS idx_w_push_histories_created ON w_push_histories(created_at);
-- +goose StatementEnd
-- +goose Down
-- +goose StatementBegin
DROP TABLE IF EXISTS w_push_histories;
DROP TABLE IF EXISTS w_push_channels;
DROP TABLE IF EXISTS w_push_events;
DROP TABLE IF EXISTS w_message_pairing_codes;
DROP TABLE IF EXISTS w_message_bindings;
DROP TABLE IF EXISTS w_message_channels;
-- +goose StatementEnd
@@ -0,0 +1,124 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package do defines domain objects, DTOs, and request/response payloads for msg_gateway.
package do
import "time"
// Capability describes what an adapter can send and receive.
type Capability struct {
Text bool
Image bool
File bool
Reply bool
Group bool
}
// ChannelConfig is the decrypted runtime config passed to a factory.
type ChannelConfig struct {
ID uint64
Type string
Name string
Credentials map[string]string
Extra map[string]string
}
// Recipient is the outbound destination on a platform.
type Recipient struct {
ChatID string
PlatformUserID string
}
// Attachment is a downloaded inbound file sitting on local disk.
type Attachment struct {
Path string
FileName string
MIME string
Error string
}
// InboundMessage is a normalized private-chat message.
type InboundMessage struct {
ChannelID uint64
ChannelType string
PlatformUserID string
ChatID string
MessageID string
Text string
Attachments []Attachment
BindingUserID *uint64
}
// OutboundMessage is a reply or probe send.
type OutboundMessage struct {
Text string
ReplyToID string
Attachments []Attachment
}
// BindRequest is the user bind body.
type BindRequest struct {
ChannelID string `json:"channel_id"`
Code string `json:"code"`
}
// BindingDTO is a user-facing binding row.
type BindingDTO struct {
ID uint64 `json:"id,string"`
UserID uint64 `json:"user_id,string"`
ChannelID uint64 `json:"channel_id,string"`
ChannelName string `json:"channel_name"`
ChannelType string `json:"channel_type"`
PlatformUserID string `json:"platform_user_id"`
CreatedAt time.Time `json:"created_at"`
}
// PublicChannelDTO is an enabled channel a user can bind to.
type PublicChannelDTO struct {
ID uint64 `json:"id,string"`
Name string `json:"name"`
Type string `json:"type"`
}
// Field is one admin form field.
type Field struct {
Key string `json:"key"`
Type string `json:"type"`
Required bool `json:"required"`
}
// Definition describes a channel type form.
type Definition struct {
Type string `json:"type"`
Fields []Field `json:"fields"`
}
// ChannelDTO represents a channel for admin consumption.
type ChannelDTO struct {
ID uint64 `json:"id,string"`
Name string `json:"name"`
Type string `json:"type"`
OwnerScope string `json:"owner_scope"`
OwnerID *uint64 `json:"owner_id,string,omitempty"`
Enabled bool `json:"enabled"`
Credentials map[string]string `json:"credentials"`
Extra map[string]string `json:"extra"`
}
// CreateChannelRequest is admin create payload.
type CreateChannelRequest struct {
Name string `json:"name"`
Type string `json:"type"`
Enabled *bool `json:"enabled"`
Credentials map[string]string `json:"credentials"`
Extra map[string]string `json:"extra"`
}
// UpdateChannelRequest is admin update payload.
type UpdateChannelRequest struct {
Name string `json:"name"`
Enabled *bool `json:"enabled"`
Credentials map[string]string `json:"credentials"`
Extra map[string]string `json:"extra"`
}
@@ -0,0 +1,419 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package do
import (
"Wavelet/plugins/domain/msg_gateway/consts"
pkgpush "Wavelet/plugins/domain/msg_gateway/push"
"sync"
"time"
)
// SMTPConfig mirrors the system SMTP settings consumed by the push service.
type SMTPConfig struct {
Host string
Port string
Username string
Password string
}
// PushField represents a form field configuration for a channel.
type PushField struct {
Key string `json:"key"`
Label string `json:"label"`
Type string `json:"type"`
Required bool `json:"required"`
Placeholder string `json:"placeholder"`
Description string `json:"description"`
}
// PushDefinition represents the metadata and form schema for a notification channel.
type PushDefinition struct {
Type string `json:"type"`
Name string `json:"name"`
Description string `json:"description"`
Fields []PushField `json:"fields"`
}
// CreatePushChannelRequest is the create channel request payload.
type CreatePushChannelRequest struct {
Name string `json:"name" binding:"required"`
Description string `json:"description"`
Type string `json:"type" binding:"required"`
Token string `json:"token"`
URL string `json:"url"`
Other string `json:"other"`
Enabled bool `json:"enabled"`
}
// UpdatePushChannelRequest is the update channel request payload.
type UpdatePushChannelRequest struct {
Description string `json:"description"`
Type string `json:"type" binding:"required"`
Token string `json:"token"`
URL string `json:"url"`
Other string `json:"other"`
Enabled bool `json:"enabled"`
}
// TestPushChannelRequest is the test channel request payload.
type TestPushChannelRequest struct {
Name string `json:"name"`
Type string `json:"type"`
Token string `json:"token"`
URL string `json:"url"`
Other string `json:"other"`
Target string `json:"target"`
}
// CustomPushRequest contains custom webhook parameters.
type CustomPushRequest struct {
Title string `json:"title" form:"title"`
Description string `json:"description" form:"description"`
Content string `json:"content" form:"content"`
URL string `json:"url" form:"url"`
To string `json:"to" form:"to"`
Token string `json:"token" form:"token"`
}
// CreatePushEventRequest is the request body for creating a push event.
type CreatePushEventRequest struct {
EventKey string `json:"event_key"`
TaskType string `json:"task_type"`
Channels []string `json:"channels"`
Targets []string `json:"targets"`
Template string `json:"template"`
Enabled bool `json:"enabled"`
}
// UpdatePushEventRequest is the request body for updating a push event.
type UpdatePushEventRequest struct {
Channels []string `json:"channels"`
Targets []string `json:"targets"`
Template string `json:"template" binding:"required"`
Enabled bool `json:"enabled"`
}
// TestPushRequest is the request body for testing push config.
type TestPushRequest struct {
Config pkgpush.Config `json:"config" binding:"required"`
Target string `json:"target"`
}
// NotificationMessage represents the structured notification message payload.
type NotificationMessage struct {
Title string `json:"title"`
Content string `json:"content"`
Level string `json:"level"`
Ext map[string]any `json:"ext,omitempty"`
}
// Flatten converts the structured NotificationMessage back to a flat map.
func (m NotificationMessage) Flatten() map[string]any {
res := map[string]any{
consts.KeyTitle: m.Title,
consts.KeyContent: m.Content,
consts.KeyLevel: m.Level,
}
for k, v := range m.Ext {
res[k] = v
}
return res
}
// EventMetadata represents the metadata of a push notification event.
type EventMetadata struct {
Key string `json:"key"`
Name string `json:"name"`
DefaultTemplate NotificationMessage `json:"default_template"`
Description string `json:"description"`
}
// SendPayload is the async push dispatch payload consumed by the notification worker.
type SendPayload struct {
EventKey string `json:"event_key"`
Config pkgpush.Config `json:"config"`
Target string `json:"target"`
Body NotificationMessage `json:"body"`
Template string `json:"template"`
}
// PushHistoryListFilter filters push history pagination queries.
type PushHistoryListFilter struct {
EventKey string
Channel string
Status string
StartTime *time.Time
EndTime *time.Time
Page int
PageSize int
}
// PushNotificationEvent defines the payload for eventbus notification trigger.
type PushNotificationEvent struct {
UserID uint64 `json:"user_id"`
Channel string `json:"channel"`
Title string `json:"title"`
Content string `json:"content"`
Metadata map[string]any `json:"metadata,omitempty"`
}
var (
pushDefMu sync.RWMutex
pushDefinitions = make(map[string]PushDefinition)
)
// RegisterPushChannelDefinition registers a channel definition.
func RegisterPushChannelDefinition(def PushDefinition) {
pushDefMu.Lock()
defer pushDefMu.Unlock()
pushDefinitions[def.Type] = def
}
// ListPushDefinitions returns all registered channel definitions.
func ListPushDefinitions() []PushDefinition {
pushDefMu.RLock()
defer pushDefMu.RUnlock()
order := []string{
consts.ChannelCustom,
consts.ChannelLark,
consts.ChannelDingTalk,
consts.ChannelTelegram,
consts.ChannelBark,
consts.ChannelDiscord,
consts.ChannelSlack,
consts.ChannelPushover,
consts.ChannelEmail,
}
res := make([]PushDefinition, 0, len(pushDefinitions))
for _, t := range order {
if d, ok := pushDefinitions[t]; ok {
res = append(res, d)
}
}
for t, d := range pushDefinitions {
found := false
for _, o := range order {
if o == t {
found = true
break
}
}
if !found {
res = append(res, d)
}
}
return res
}
//nolint:funlen,goconst,dupl // Channel definitions registration table
func init() {
RegisterPushChannelDefinition(PushDefinition{
Type: consts.ChannelCustom,
Name: "自定义消息通道",
Description: "使用自定义 HTTP POST 请求向外部 Webhook 发送数据。",
Fields: []PushField{
{
Key: consts.KeyURL,
Label: "请求地址",
Type: consts.TypeText,
Required: true,
Placeholder: "在此填写完整的请求地址,必须使用 HTTPS 协议",
Description: "接口请求的完整 HTTPS URL,例如 https://api.example.com/webhook",
},
{
Key: consts.KeyOther,
Label: "请求体 (JSON)",
Type: consts.TypeTextarea,
Required: true,
Placeholder: "在此输入请求体,支持模板变量,必须为合法的 JSON 格式",
Description: "可使用的变量:$title, $description, $content, $url, $to。例如 {\"text\": \"$content\"}",
},
},
})
RegisterPushChannelDefinition(PushDefinition{
Type: consts.ChannelLark,
Name: "飞书群机器人",
Description: "配置飞书群自定义机器人的 Webhook 接口投递。",
Fields: []PushField{
{
Key: consts.KeyURL,
Label: "Webhook 地址",
Type: consts.TypeText,
Required: true,
Placeholder: "https://open.feishu.cn/open-apis/bot/v2/hook/YOUR_TOKEN",
Description: "从飞书群机器人设置中复制的 Webhook URL",
},
{
Key: consts.KeyToken,
Label: "签名校验密钥 (Secret) (可选)",
Type: consts.TypeText,
Required: false,
Placeholder: "可选,若机器人启用了安全设置中的签名校验,请在此输入",
Description: "飞书群机器人安全设置中的签名校验 Key",
},
{
Key: consts.KeyOther,
Label: "自定义卡片 JSON 模版 (可选)",
Type: consts.TypeTextarea,
Required: false,
Placeholder: "可选,留空则默认使用系统内置的精美互动卡片",
Description: "若填写,必须是合法的飞书卡片 JSON 格式",
},
},
})
RegisterPushChannelDefinition(PushDefinition{
Type: consts.ChannelDingTalk,
Name: "钉钉群机器人",
Description: "配置钉钉群自定义机器人的 Webhook 接口投递。",
Fields: []PushField{
{
Key: consts.KeyURL,
Label: "Webhook 地址",
Type: consts.TypeText,
Required: true,
Placeholder: "https://oapi.dingtalk.com/robot/send?access_token=YOUR_TOKEN",
Description: "从钉钉群机器人设置中获取的完整 Webhook URL",
},
{
Key: consts.KeyToken,
Label: "加签密钥 (Secret) (可选)",
Type: consts.TypeText,
Required: false,
Placeholder: "可选,若机器人启用了安全设置中的加签校验,请在此输入 SEC 开头的密钥",
Description: "钉钉群机器人安全设置中的加签 Secret",
},
},
})
RegisterPushChannelDefinition(PushDefinition{
Type: consts.ChannelTelegram,
Name: "Telegram 机器人",
Description: "配置 Telegram 机器人推送消息。",
Fields: []PushField{
{
Key: consts.KeyURL,
Label: "API 基础地址 (可选)",
Type: consts.TypeText,
Required: false,
Placeholder: "https://api.telegram.org",
Description: "接口请求的 HTTPS 基础地址,留空默认为 https://api.telegram.org",
},
{
Key: consts.KeyToken,
Label: "机器人 Token (Bot Token)",
Type: consts.TypePassword,
Required: true,
Placeholder: "在此输入 Telegram 机器人的 Bot Token",
Description: "通过 BotFather 申请到的机器人 Access Token",
},
{
Key: consts.KeyOther,
Label: "默认会话 ID (Chat ID) (可选)",
Type: consts.TypeText,
Required: false,
Placeholder: "例如 -100123456789 或 @channel_name",
Description: "默认的消息接收 Chat ID。如果通知事件中未配置 targets,将推送到此 ID",
},
},
})
RegisterPushChannelDefinition(PushDefinition{
Type: consts.ChannelBark,
Name: "Bark (iOS 推送)",
Description: "配置 Bark 推送通知至 iPhone / iPad 客户端。",
Fields: []PushField{
{
Key: consts.KeyToken,
Label: "设备 Key (Device Key)",
Type: consts.TypeText,
Required: true,
Placeholder: "Bark App 首页显示的 Device Key",
Description: "从 Bark App 复制的设备专属 Key",
},
{
Key: consts.KeyURL,
Label: "Bark 服务器地址 (可选)",
Type: consts.TypeText,
Required: false,
Placeholder: "https://api.day.app",
Description: "Bark 服务器地址,留空默认使用官方公共服务器 https://api.day.app",
},
{
Key: consts.KeyOther,
Label: "额外配置 JSON (可选)",
Type: consts.TypeTextarea,
Required: false,
Placeholder: "{\"group\": \"Wavelet\", \"sound\": \"minuet\", \"icon\": \"https://...\"}",
Description: "可选的 JSON 配置,支持 group (分组)、sound (铃声)、icon (自定义图标)",
},
},
})
RegisterPushChannelDefinition(PushDefinition{
Type: consts.ChannelDiscord,
Name: "Discord 频道",
Description: "配置 Discord 频道的 Incoming Webhook 消息推送。",
Fields: []PushField{
{
Key: consts.KeyURL,
Label: "Webhook 地址",
Type: consts.TypeText,
Required: true,
Placeholder: "https://discord.com/api/webhooks/...",
Description: "从 Discord 频道集成设置中复制的 Webhook URL",
},
},
})
RegisterPushChannelDefinition(PushDefinition{
Type: consts.ChannelSlack,
Name: "Slack 频道",
Description: "配置 Slack 频道的 Incoming Webhook 消息推送。",
Fields: []PushField{
{
Key: consts.KeyURL,
Label: "Webhook 地址",
Type: consts.TypeText,
Required: true,
Placeholder: "https://hooks.slack.com/services/...",
Description: "从 Slack 应用配置中复制的 Incoming Webhook URL",
},
},
})
RegisterPushChannelDefinition(PushDefinition{
Type: consts.ChannelPushover,
Name: "Pushover 推送",
Description: "配置 Pushover 即时推送到手机/桌面客户端。",
Fields: []PushField{
{
Key: consts.KeyToken,
Label: "应用 Token (App Token)",
Type: consts.TypePassword,
Required: true,
Placeholder: "Pushover 创建应用生成的 API Token / Key",
Description: "从 Pushover 控制台创建的 Application API Token",
},
{
Key: consts.KeyURL,
Label: "用户 Key (User Key)",
Type: consts.TypeText,
Required: true,
Placeholder: "Pushover 账号主页的 User Key",
Description: "Pushover 个人账号的 User Key",
},
},
})
RegisterPushChannelDefinition(PushDefinition{
Type: consts.ChannelEmail,
Name: "邮件推送通道",
Description: "邮件推送通道直接使用系统全局 SMTP 设置进行发送,无需在此填写服务器配置。",
Fields: []PushField{},
})
}
@@ -0,0 +1,56 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package entity defines GORM table mapping entities for msg_gateway.
package entity
import "time"
// MessageChannel is an admin-configured messaging adapter entity.
type MessageChannel struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
Type string `json:"type" gorm:"size:32;not null"`
Name string `json:"name" gorm:"size:64;not null"`
OwnerScope string `json:"owner_scope" gorm:"size:32;not null;default:'system'"`
OwnerID *uint64 `json:"owner_id,omitempty"`
Credentials string `json:"credentials" gorm:"type:text;not null"`
Extra string `json:"extra" gorm:"type:text"`
Enabled bool `json:"enabled" gorm:"default:false;not null"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (MessageChannel) TableName() string {
return "w_message_channels"
}
// MessageBinding maps a platform user to a Wavelet user on one channel.
type MessageBinding struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
ChannelID uint64 `json:"channel_id" gorm:"not null;index"`
PlatformUserID string `json:"platform_user_id" gorm:"size:128;not null;index"`
UserID uint64 `json:"user_id" gorm:"not null;index"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName 表名
func (MessageBinding) TableName() string {
return "w_message_bindings"
}
// MessagePairingCode is a one-time bind code.
type MessagePairingCode struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
Code string `json:"code" gorm:"size:32;uniqueIndex;not null"`
ChannelID uint64 `json:"channel_id" gorm:"not null;index"`
PlatformUserID string `json:"platform_user_id" gorm:"size:128;not null;index"`
UserID uint64 `json:"user_id" gorm:"not null;index"`
ExpiresAt time.Time `json:"expires_at" gorm:"not null;index"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName 表名
func (MessagePairingCode) TableName() string {
return "w_message_pairing_codes"
}
@@ -0,0 +1,95 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package entity
import (
"Wavelet/plugins/domain/msg_gateway/consts"
"errors"
"strings"
"time"
)
// PushChannel 消息通道实体
type PushChannel struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"size:100;not null"`
Description string `json:"description" gorm:"size:255"`
Type string `json:"type" gorm:"size:50;not null;index"`
URL string `json:"url" gorm:"type:text"`
Token string `json:"token" gorm:"type:text"`
Other string `json:"other" gorm:"type:text"`
Enabled bool `json:"enabled" gorm:"index;not null;default:true"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
}
// TableName 指定 GORM 表名
func (PushChannel) TableName() string {
return "w_push_channels"
}
// Validate 验证与标准化字段
func (c *PushChannel) Validate() error {
c.Name = strings.TrimSpace(c.Name)
if c.Name == "" {
return errors.New(consts.ErrChannelNameRequired)
}
c.Type = strings.TrimSpace(c.Type)
if c.Type == "" {
return errors.New(consts.ErrChannelTypeRequired)
}
return nil
}
// PushEvent 系统通知事件实体
type PushEvent struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
EventKey string `json:"event_key" gorm:"uniqueIndex;size:80;not null"`
Name string `json:"name" gorm:"size:100;not null"`
TaskType string `json:"task_type" gorm:"size:100;index;not null;default:''"`
Channels []string `json:"channels" gorm:"type:text;serializer:json;not null"`
Targets []string `json:"targets" gorm:"type:text;serializer:json;not null"`
Template string `json:"template" gorm:"type:text;not null"`
Enabled bool `json:"enabled" gorm:"index;not null;default:false"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
}
// TableName 指定 GORM 表名
func (PushEvent) TableName() string {
return "w_push_events"
}
// Validate 验证 PushEvent 实体字段
func (e *PushEvent) Validate() error {
e.EventKey = strings.TrimSpace(e.EventKey)
if e.EventKey == "" {
return errors.New(consts.ErrEventKeyRequired)
}
e.Name = strings.TrimSpace(e.Name)
if e.Name == "" {
return errors.New(consts.ErrNameRequired)
}
return nil
}
// PushHistory 推送日志/历史实体
type PushHistory struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
EventKey string `json:"event_key" gorm:"size:80;not null;index"`
Channel string `json:"channel" gorm:"size:50;not null;index"`
Target string `json:"target" gorm:"size:255;not null"`
Title string `json:"title" gorm:"size:255;not null"`
Content string `json:"content" gorm:"type:text;not null"`
Level string `json:"level" gorm:"size:20;not null;default:'INFO'"`
Status string `json:"status" gorm:"size:20;not null;index"`
ErrorMsg string `json:"error_msg" gorm:"type:text"`
Payload string `json:"payload" gorm:"type:text"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
}
// TableName 指定 GORM 表名
func (PushHistory) TableName() string {
return "w_push_histories"
}
@@ -0,0 +1,51 @@
// 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)
}
}
@@ -0,0 +1,33 @@
// 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)
}
}
@@ -0,0 +1,361 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package msg_gateway provides the Bot gateway, multi-channel notification dispatching, and asynchronous push worker domain plugin for Cordis.
package msg_gateway
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/msg_gateway/channels/qq"
"Wavelet/plugins/domain/msg_gateway/channels/telegram"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/controller"
"Wavelet/plugins/domain/msg_gateway/dao"
"Wavelet/plugins/domain/msg_gateway/model/do"
"Wavelet/plugins/domain/msg_gateway/model/entity"
"Wavelet/plugins/domain/msg_gateway/service"
"context"
"embed"
"reflect"
"github.com/gin-gonic/gin"
)
//go:embed migrations/*/*.sql
var mgMigrations embed.FS
// Option configures the msg_gateway plugin.
type Option func(*Plugin)
// WithAutoStartRunner enables automatic bot runner startup in the background.
func WithAutoStartRunner(enable bool) Option {
return func(p *Plugin) {
p.autoStartRunner = enable
}
}
// Plugin implements core.Plugin to provide Bot gateway and notification dispatch domain services.
type Plugin struct {
autoStartRunner bool
cancelRunner context.CancelFunc
}
// New creates a new msg_gateway domain plugin.
func New(opts ...Option) *Plugin {
p := &Plugin{}
for _, opt := range opts {
if opt != nil {
opt(p)
}
}
return p
}
// Name returns the unique identifier for the msg_gateway domain plugin.
func (p *Plugin) Name() string {
return "msg_gateway"
}
// Inject declares required dependencies for the msg_gateway domain plugin.
func (p *Plugin) Inject() []reflect.Type {
return []reflect.Type{
reflect.TypeFor[contracts.DBService](),
// AuthService is captured as a middleware value in Apply, so it cannot
// be late-bound with core.When like the other services below; the
// kernel must mount auth first or the routes get a pass-through guard.
reflect.TypeFor[contracts.AuthService](),
}
}
// Manifest returns the plugin metadata.
func (p *Plugin) Manifest() core.Manifest {
return core.Manifest{
Name: "msg_gateway",
Version: "1.0.0",
Description: "Bot gateway, multi-channel notification push, and async worker dispatch plugin",
Author: "Wavelet Team",
}
}
type mgAppConfig struct {
SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"`
}
// DeclareConfig declares configuration bindings for the msg_gateway plugin.
func (p *Plugin) DeclareConfig() []core.ConfigBinding {
return []core.ConfigBinding{
{Prefix: "app", Target: &mgAppConfig{}},
}
}
// Apply registers msg_gateway migrations, routes, tasks, schedules, events, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
var cfg mgAppConfig
if err := ctx.Config().Bind("app", &cfg); err == nil && cfg.SessionSecret != "" {
service.SetCredentialSecret(cfg.SessionSecret)
}
core.Bind[contracts.DBService](ctx, dao.SetDBService)
core.Bind[contracts.CacheService](ctx, func(cache contracts.CacheService) {
dao.SetCacheService(cache)
service.SetCacheService(cache)
})
core.Bind[contracts.TaskService](ctx, service.SetTaskService)
core.Bind[contracts.UserService](ctx, service.SetUserService)
ctx.OnDispose(func() error {
dao.SetDBService(nil)
dao.SetCacheService(nil)
service.SetCacheService(nil)
service.SetTaskService(nil)
service.SetUserService(nil)
return nil
})
// 0. Resolve auth service for middleware (via IoC, not direct import)
denyAuth := ginutil.AuthUnavailable()
loginMW := denyAuth
adminMW := denyAuth
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
loginMW = mw
}
if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok {
adminMW = mw
}
}
// 1. Register migrations
ctx.Migrations().Register("msg_gateway", mgMigrations)
// 2. Register User HTTP Routes
controller.RegisterUserRoutes(ctx.Router().Group("/api/v1"), loginMW)
// 3. Register Admin Message Gateway HTTP Routes
controller.RegisterAdminRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW)
// 4. Register Admin Push HTTP Routes
controller.RegisterAdminPushRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW)
service.Register(consts.MessageChannelTypeTelegram, telegram.New)
service.Register(consts.MessageChannelTypeQQ, qq.New)
const defaultTaskRetry = 3
pushHandler := &service.PushHandler{}
// 5. Register background tasks
ctx.Task().Register(consts.TaskPushNotification, func(c context.Context, payload []byte) error {
return pushHandler.Execute(c, payload)
},
extpoints.WithTaskType("push_notification"),
extpoints.WithTaskName("消息网关推送通知"),
extpoints.WithTaskDescription("异步执行系统通知的多渠道派发与推送"),
extpoints.WithTaskCategory("push"),
extpoints.WithTaskRetry(defaultTaskRetry),
extpoints.WithTaskQueue("default"),
extpoints.WithTaskRetryable(true),
)
ctx.Task().Register(service.SendNotificationTask, func(c context.Context, payload []byte) error {
return pushHandler.Execute(c, payload)
}, extpoints.WithTaskMeta(service.SendNotificationMeta), extpoints.WithTaskRetry(defaultTaskRetry))
ctx.Task().Register(service.TaskDispatchBotMsg, &service.BotDispatchHandler{},
extpoints.WithTaskMeta(service.BotDispatchMeta))
ctx.Task().Register(consts.TaskCleanupPairingCodes, func(c context.Context, _ []byte) error {
return dao.DeleteExpiredPairingCodes(c)
},
extpoints.WithTaskType("cleanup_pairing_codes"),
extpoints.WithTaskName("清理过期配对码"),
extpoints.WithTaskDescription("定时清理已过期的平台 Bot 配对码"),
extpoints.WithTaskCategory("messaging"),
extpoints.WithTaskRetry(defaultTaskRetry),
extpoints.WithTaskQueue("default"),
extpoints.WithTaskRetryable(true),
)
// 6. Register Cron Schedules
ctx.Schedule().RegisterCron("*/10 * * * *", consts.TaskCleanupPairingCodes, map[string]any{"action": "cleanup"})
// 7. Register EventBus listeners for decoupled push triggers
ctx.Events().On("notification:push", func(c context.Context, e do.PushNotificationEvent) error {
meta := do.EventMetadata{
Key: "eventbus:" + e.Channel,
Name: e.Title,
DefaultTemplate: do.NotificationMessage{
Title: e.Title,
Content: e.Content,
Level: consts.DefaultLevelInfo,
Ext: e.Metadata,
},
Description: "EventBus triggered notification",
}
service.DefaultTrigger.Trigger(c, meta, map[string]any{
"user.id": e.UserID,
"title": e.Title,
"content": e.Content,
})
return nil
})
// 8. Register task completed event listener
ctx.Events().On(contracts.EventTopicTaskCompleted, func(c context.Context, e contracts.TaskCompletedEvent) error {
service.HandleTaskCompleted(c, e)
return nil
})
// 9. Register built-in domain events and provide PushRegistry
service.RegisterCustomEvents()
core.Provide[contracts.PushRegistry](ctx, service.PushRegistryAdapter{})
// 10. Register Settings Schemas
ctx.Settings().Register(extpoints.SettingSchema{
Key: "msg_gateway.pairing_code_expiry_minutes",
Default: 15,
Description: "Expiry duration for bot pairing codes in minutes",
Type: "integer",
Category: "messaging",
})
ctx.Settings().Register(extpoints.SettingSchema{
Key: "msg_gateway.max_bindings_per_user",
Default: 5,
Description: "Maximum platform bot bindings per user",
Type: "integer",
Category: "messaging",
})
// 11. Optional runner start & lifecycle
if p.autoStartRunner {
runnerCtx, cancel := context.WithCancel(ctx.GoContext())
p.cancelRunner = cancel
util.Go(func() {
_ = service.Start(runnerCtx)
})
}
ctx.OnDispose(func() error {
if p.cancelRunner != nil {
p.cancelRunner()
}
return nil
})
return nil
}
// Re-exported constants.
const (
CodeAlphabet = service.CodeAlphabet
CodeLength = service.CodeLength
)
// MessageChannel is an alias for entity.MessageChannel.
type MessageChannel = entity.MessageChannel
// MessageBinding is an alias for entity.MessageBinding.
type MessageBinding = entity.MessageBinding
// MessagePairingCode is an alias for entity.MessagePairingCode.
type MessagePairingCode = entity.MessagePairingCode
// PushChannel is an alias for entity.PushChannel.
type PushChannel = entity.PushChannel
// PushEvent is an alias for entity.PushEvent.
type PushEvent = entity.PushEvent
// PushHistory is an alias for entity.PushHistory.
type PushHistory = entity.PushHistory
// PushNotificationEvent is an alias for do.PushNotificationEvent.
type PushNotificationEvent = do.PushNotificationEvent
// ChannelConfig is an alias for do.ChannelConfig.
type ChannelConfig = do.ChannelConfig
// Capability is an alias for do.Capability.
type Capability = do.Capability
// Recipient is an alias for do.Recipient.
type Recipient = do.Recipient
// Attachment is an alias for do.Attachment.
type Attachment = do.Attachment
// InboundMessage is an alias for do.InboundMessage.
type InboundMessage = do.InboundMessage
// OutboundMessage is an alias for do.OutboundMessage.
type OutboundMessage = do.OutboundMessage
// BindingDTO is an alias for do.BindingDTO.
type BindingDTO = do.BindingDTO
// PublicChannelDTO is an alias for do.PublicChannelDTO.
type PublicChannelDTO = do.PublicChannelDTO
// Definition is an alias for do.Definition.
type Definition = do.Definition
// ChannelDTO is an alias for do.ChannelDTO.
type ChannelDTO = do.ChannelDTO
// CreateChannelRequest is an alias for do.CreateChannelRequest.
type CreateChannelRequest = do.CreateChannelRequest
// UpdateChannelRequest is an alias for do.UpdateChannelRequest.
type UpdateChannelRequest = do.UpdateChannelRequest
// PushDefinition is an alias for do.PushDefinition.
type PushDefinition = do.PushDefinition
// PushField is an alias for do.PushField.
type PushField = do.PushField
// NotificationMessage is an alias for do.NotificationMessage.
type NotificationMessage = do.NotificationMessage
// EventMetadata is an alias for do.EventMetadata.
type EventMetadata = do.EventMetadata
// SendPayload is an alias for do.SendPayload.
type SendPayload = do.SendPayload
// Handler is an alias for service.Handler.
type Handler = service.Handler
// Factory is an alias for service.Factory.
type Factory = service.Factory
// Channel is an alias for service.Channel.
type Channel = service.Channel
// Runner is an alias for service.Runner.
type Runner = service.Runner
// EventTrigger is an alias for service.EventTrigger.
type EventTrigger = service.EventTrigger
// PushHandler is an alias for service.PushHandler.
type PushHandler = service.PushHandler
// Re-exported variables and functions.
var (
SetDBServiceForTest = dao.SetDBServiceForTest
UpsertPairingCode = dao.UpsertPairingCode
Register = service.Register
Lookup = service.Lookup
GenerateCode = service.GenerateCode
NormalizeCode = service.NormalizeCode
FormatCode = service.FormatCode
Start = service.Start
Stop = service.Stop
GlobalRunner = service.GlobalRunner
DefaultTrigger = service.DefaultTrigger
SyncEvents = service.SyncEvents
AdminLogin = service.AdminLogin
HandleAdminLoggedIn = service.HandleAdminLoggedIn
)
@@ -0,0 +1,54 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package msg_gateway_test
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/msg_gateway"
"Wavelet/plugins/domain/msg_gateway/service"
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestPushRegistry(t *testing.T) {
ctx := core.NewContext(context.Background())
require.NoError(t, msg_gateway.New().Apply(ctx))
registry, err := core.Inject[contracts.PushRegistry](ctx)
require.NoError(t, err)
require.NotNil(t, registry)
const key = "test.push_registry.probe"
registry.RegisterBuiltInEvent(contracts.PushEventMeta{
Key: key,
Name: "Push Registry Probe",
Description: "observability probe for contracts.PushRegistry",
DefaultTemplate: contracts.PushNotificationTemplate{
Title: "Probe Title",
Content: "Probe Content",
Level: "INFO",
Ext: map[string]any{"source": "test"},
},
})
found := false
for _, ev := range service.GetBuiltInEvents() {
if ev.Key != key {
continue
}
found = true
assert.Equal(t, "Push Registry Probe", ev.Name)
assert.Equal(t, "observability probe for contracts.PushRegistry", ev.Description)
assert.Equal(t, "Probe Title", ev.DefaultTemplate.Title)
assert.Equal(t, "Probe Content", ev.DefaultTemplate.Content)
assert.Equal(t, "INFO", ev.DefaultTemplate.Level)
assert.Equal(t, map[string]any{"source": "test"}, ev.DefaultTemplate.Ext)
break
}
require.True(t, found, "registered key %q should be visible via GetBuiltInEvents", key)
}
@@ -0,0 +1,61 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package msg_gateway_test
import (
"Wavelet/core"
"Wavelet/plugins/domain/msg_gateway"
"context"
"io/fs"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestMsgGatewayPluginUnit(t *testing.T) {
ctx := core.NewContext(context.Background())
p := msg_gateway.New()
assert.Equal(t, "msg_gateway", p.Name())
assert.Equal(t, "1.0.0", p.Manifest().Version)
require.NoError(t, p.Apply(ctx))
// Verify migrations
entry, ok := ctx.Migrations().Get("msg_gateway")
require.True(t, ok)
entries, err := fs.ReadDir(entry.FS, entry.Dir)
require.NoError(t, err)
assert.NotEmpty(t, entries)
// Verify tasks
task, ok := ctx.Tasks().Get("msg_gateway:push_notification")
require.True(t, ok)
assert.Equal(t, 3, task.Retry)
// Verify schedules
sched, ok := ctx.Schedules().Get("msg_gateway:cleanup_pairing_codes")
require.True(t, ok)
assert.Equal(t, "*/10 * * * *", sched.Spec)
// Verify settings
setting, ok := ctx.Settings().Get("msg_gateway.max_bindings_per_user")
require.True(t, ok)
assert.Equal(t, 5, setting.Default)
}
// TestEveryScheduleHasTaskHandler 回归:RegisterCron 仅登记调度;若同名任务从未
// Register,则每次触发都投递到无人处理的任务类型,清理逻辑静默失效。
func TestEveryScheduleHasTaskHandler(t *testing.T) {
ctx := core.NewContext(context.Background())
require.NoError(t, msg_gateway.New().Apply(ctx))
schedules := ctx.Schedules().Schedules()
require.NotEmpty(t, schedules)
for _, sched := range schedules {
_, ok := ctx.Tasks().Get(sched.TaskType)
assert.Truef(t, ok, "schedule %q dispatches to task %q, which is never registered",
sched.Spec, sched.TaskType)
}
}
@@ -0,0 +1,68 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"errors"
"fmt"
"strings"
"github.com/nikoksr/notify"
"github.com/nikoksr/notify/service/bark"
)
func init() {
Register("bark", &BarkPusher{})
}
// BarkPusher 基于 nikoksr/notify 的 Bark iOS 客户端通知推送实现
type BarkPusher struct{}
// Send 发送 Bark 通知
func (p *BarkPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, _ map[string]any) (string, error) {
deviceKey := cfg.Key
if deviceKey == "" {
deviceKey = cfg.Secret
}
if target != "" {
deviceKey = target
}
if deviceKey == "" {
return "", errors.New("bark: device key is required")
}
serverURL := strings.TrimRight(cfg.URL, "/")
if serverURL == "" {
serverURL = bark.DefaultServerURL
}
title := bodyTitle(body)
content := bodyContent(body, "%s: %v", "\n")
barkService := bark.NewWithServers(deviceKey, serverURL)
notifier := notify.New()
notifier.UseServices(barkService)
if err := notifier.Send(ctx, title, content); err != nil {
return "", fmt.Errorf("bark: notify send failed: %w", err)
}
return "ok", nil
}
// ValidateConfig 校验 Bark 配置
func (p *BarkPusher) ValidateConfig(cfg Config) error {
deviceKey := cfg.Key
if deviceKey == "" {
deviceKey = cfg.Secret
}
if deviceKey == "" {
return errors.New("device key is required")
}
if cfg.URL != "" && !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
return errors.New("server URL must start with http:// or https://")
}
return nil
}
@@ -0,0 +1,125 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"net/http"
"net/http/httptest"
"testing"
)
func TestDingTalkPusher(t *testing.T) {
pusher, err := GetPusher("dingtalk")
if err != nil {
t.Fatalf("failed to get dingtalk pusher: %v", err)
}
err = pusher.ValidateConfig(Config{URL: "https://oapi.dingtalk.com/robot/send?access_token=test"})
if err != nil {
t.Errorf("ValidateConfig failed: %v", err)
}
err = pusher.ValidateConfig(Config{})
if err == nil {
t.Errorf("expected error for empty config, got nil")
}
}
func TestBarkPusher(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{"code":200,"message":"success"}`))
}))
defer server.Close()
pusher, err := GetPusher("bark")
if err != nil {
t.Fatalf("failed to get bark pusher: %v", err)
}
err = pusher.ValidateConfig(Config{Key: "device_key_123"})
if err != nil {
t.Errorf("ValidateConfig failed: %v", err)
}
_, err = pusher.Send(context.Background(), Config{
URL: server.URL,
Key: "device_key_123",
}, "", map[string]any{
"title": "Alert",
"content": "Bark notification",
}, "", nil)
if err != nil {
t.Errorf("Send failed: %v", err)
}
}
func TestDiscordPusher(t *testing.T) {
pusher, err := GetPusher("discord")
if err != nil {
t.Fatalf("failed to get discord pusher: %v", err)
}
err = pusher.ValidateConfig(Config{Key: "bot_token_123"})
if err != nil {
t.Errorf("ValidateConfig failed: %v", err)
}
err = pusher.ValidateConfig(Config{})
if err == nil {
t.Errorf("expected error for empty config, got nil")
}
}
func TestSlackPusher(t *testing.T) {
pusher, err := GetPusher("slack")
if err != nil {
t.Fatalf("failed to get slack pusher: %v", err)
}
err = pusher.ValidateConfig(Config{Key: "xoxb-123456"})
if err != nil {
t.Errorf("ValidateConfig failed: %v", err)
}
err = pusher.ValidateConfig(Config{})
if err == nil {
t.Errorf("expected error for empty config, got nil")
}
}
func TestPushoverPusher(t *testing.T) {
pusher, err := GetPusher("pushover")
if err != nil {
t.Fatalf("failed to get pushover pusher: %v", err)
}
err = pusher.ValidateConfig(Config{Key: "app_token_123"})
if err != nil {
t.Errorf("ValidateConfig failed: %v", err)
}
err = pusher.ValidateConfig(Config{})
if err == nil {
t.Errorf("expected error for empty config, got nil")
}
}
func TestLarkPusher(t *testing.T) {
pusher, err := GetPusher("lark")
if err != nil {
t.Fatalf("failed to get lark pusher: %v", err)
}
err = pusher.ValidateConfig(Config{URL: "https://open.feishu.cn/open-apis/bot/v2/hook/xxx"})
if err != nil {
t.Errorf("ValidateConfig failed: %v", err)
}
err = pusher.ValidateConfig(Config{})
if err == nil {
t.Errorf("expected error for empty config, got nil")
}
}
@@ -0,0 +1,97 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"Wavelet/pkg/httppool"
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strings"
)
func init() {
Register("custom", &CustomPusher{})
}
// maxCustomResponseBytes 限制读取 Webhook 响应体的最大字节数,防止无界读取。
const maxCustomResponseBytes = 4096
// CustomPusher 自定义 Webhook 发送实现
type CustomPusher struct{}
// Send 发送自定义 webhook
func (p *CustomPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, template string, _ map[string]any) (string, error) {
if cfg.URL == "" {
return "", errors.New("custom: URL is required")
}
var reqBody []byte
if template != "" {
// 替换模板中的 {{key}} 占位符
rendered := ParseTemplate(template, body)
reqBody = []byte(rendered)
} else {
// 兜底:直接把 body 转为 JSON 字符串发送
var err error
reqBody, err = json.Marshal(body)
if err != nil {
return "", fmt.Errorf("custom: marshal body failed: %w", err)
}
}
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.URL, bytes.NewReader(reqBody))
if err != nil {
return "", fmt.Errorf("custom: create http request failed: %w", err)
}
httpReq.Header.Set("Content-Type", "application/json")
// 如果配置了 Key 且格式为 "HeaderName:HeaderValue",我们可以附加测试用 Header
if cfg.Key != "" && strings.Contains(cfg.Key, ":") {
parts := strings.SplitN(cfg.Key, ":", 2) //nolint:mnd
httpReq.Header.Set(strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1]))
}
client := httppool.NewClient(defaultHTTPClientTimeout)
resp, err := client.Do(httpReq)
if err != nil {
return "", fmt.Errorf("custom: http request failed: %w", err)
}
defer func() { _ = resp.Body.Close() }()
bodyBytes, _ := io.ReadAll(io.LimitReader(resp.Body, maxCustomResponseBytes))
upstreamResp := strings.TrimSpace(string(bodyBytes))
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return upstreamResp, fmt.Errorf("custom: http status %s", resp.Status)
}
// 部分 Webhook(如企业微信、钉钉)即使业务失败也返回 HTTP 200,
// 仅当响应体包含非零 errcode 时才判定为发送失败,避免审计记录误报成功。
var apiResp struct {
ErrCode int `json:"errcode"`
ErrMsg string `json:"errmsg"`
}
if err := json.Unmarshal(bodyBytes, &apiResp); err == nil && apiResp.ErrCode != 0 {
return upstreamResp, fmt.Errorf("custom: webhook rejected: errcode=%d errmsg=%q", apiResp.ErrCode, apiResp.ErrMsg)
}
return upstreamResp, nil
}
// ValidateConfig 校验自定义配置
func (p *CustomPusher) ValidateConfig(cfg Config) error {
if cfg.URL == "" {
return errors.New("webhook URL is required")
}
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
return errors.New("webhook URL must start with http:// or https://")
}
return nil
}
@@ -0,0 +1,91 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestCustomPusherSend_ResponseBodyErrcode(t *testing.T) {
tests := []struct {
name string
statusCode int
body string
wantErr bool
wantErrMsg string
}{
{
name: "wechat business error returns HTTP 200 with non-zero errcode",
statusCode: http.StatusOK,
body: `{"errcode":93000,"errmsg":"invalid request data"}`,
wantErr: true,
wantErrMsg: "errcode=93000",
},
{
name: "wechat success returns errcode 0",
statusCode: http.StatusOK,
body: `{"errcode":0,"errmsg":"ok"}`,
wantErr: false,
},
{
name: "json response without errcode is tolerated",
statusCode: http.StatusOK,
body: `{"success":true}`,
wantErr: false,
},
{
name: "non-json response body is tolerated",
statusCode: http.StatusOK,
body: "ok",
wantErr: false,
},
{
name: "empty response body is tolerated",
statusCode: http.StatusNoContent,
body: "",
wantErr: false,
},
{
name: "http error status still fails",
statusCode: http.StatusInternalServerError,
body: `{"errcode":0,"errmsg":"ok"}`,
wantErr: true,
wantErrMsg: "http status",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(tt.statusCode)
_, _ = w.Write([]byte(tt.body))
}))
defer srv.Close()
pusher := &CustomPusher{}
upstreamResp, err := pusher.Send(context.Background(),
Config{Channel: "custom", URL: srv.URL},
"",
map[string]any{"title": "t", "content": "c"},
`{"title":"$title","content":"$content"}`,
nil,
)
if tt.wantErr {
require.Error(t, err)
assert.Contains(t, err.Error(), tt.wantErrMsg)
return
}
assert.NoError(t, err)
if tt.body != "" {
assert.Contains(t, upstreamResp, tt.body)
}
})
}
}
@@ -0,0 +1,66 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"errors"
"fmt"
"net/url"
"strings"
"github.com/nikoksr/notify"
"github.com/nikoksr/notify/service/dingding"
)
func init() {
Register("dingtalk", &DingTalkPusher{})
}
// DingTalkPusher 基于 nikoksr/notify 的钉钉机器人推送实现
type DingTalkPusher struct{}
// Send 发送钉钉通知
func (p *DingTalkPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, _ string, _ map[string]any) (string, error) {
token := cfg.Key
if token == "" {
if u, err := url.Parse(cfg.URL); err == nil {
token = u.Query().Get("access_token")
}
}
if token == "" {
token = cfg.URL
}
if token == "" {
return "", errors.New("dingtalk: access token or webhook URL is required")
}
title := bodyTitle(body)
content := bodyContent(body, "**%s**: %v", "\n\n")
dingService := dingding.New(&dingding.Config{
Token: token,
Secret: cfg.Secret,
})
notifier := notify.New()
notifier.UseServices(dingService)
if err := notifier.Send(ctx, title, content); err != nil {
return "", fmt.Errorf("dingtalk: notify send failed: %w", err)
}
return "ok", nil
}
// ValidateConfig 校验钉钉配置
func (p *DingTalkPusher) ValidateConfig(cfg Config) error {
if cfg.URL == "" && cfg.Key == "" {
return errors.New("webhook URL or access token is required")
}
if cfg.URL != "" && !strings.HasPrefix(cfg.URL, "https://") {
return errors.New("webhook URL must use https:// protocol")
}
return nil
}
@@ -0,0 +1,72 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"errors"
"fmt"
"github.com/nikoksr/notify"
"github.com/nikoksr/notify/service/discord"
)
func init() {
Register("discord", &DiscordPusher{})
}
// DiscordPusher 基于 nikoksr/notify 的 Discord 推送实现
type DiscordPusher struct{}
// Send 发送 Discord 通知
func (p *DiscordPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, _ map[string]any) (string, error) {
botToken := cfg.Key
if botToken == "" {
botToken = cfg.Secret
}
channelID := cfg.URL
if target != "" {
channelID = target
}
if channelID == "" {
channelID = cfg.Other
}
if botToken == "" {
return "", errors.New("discord: bot token is required")
}
if channelID == "" {
return "", errors.New("discord: channel ID is required")
}
title := bodyTitle(body)
content := bodyContent(body, "**%s**: %v", "\n")
discordService := discord.New()
if err := discordService.AuthenticateWithBotToken(botToken); err != nil {
return "", fmt.Errorf("discord: auth failed: %w", err)
}
discordService.AddReceivers(channelID)
notifier := notify.New()
notifier.UseServices(discordService)
if err := notifier.Send(ctx, title, content); err != nil {
return "", fmt.Errorf("discord: notify send failed: %w", err)
}
return "ok", nil
}
// ValidateConfig 校验 Discord 配置
func (p *DiscordPusher) ValidateConfig(cfg Config) error {
botToken := cfg.Key
if botToken == "" {
botToken = cfg.Secret
}
if botToken == "" {
return errors.New("bot token is required")
}
return nil
}
@@ -0,0 +1,78 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
pkgmail "Wavelet/pkg/mail"
"context"
"errors"
"fmt"
"net"
"strconv"
)
func init() {
Register("email", &EmailPusher{})
}
// EmailPusher 基于 pkg/mail 的 SMTP 邮件推送实现
type EmailPusher struct{}
// Send 发送邮件
func (p *EmailPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, ext map[string]any) (string, error) {
if cfg.URL == "" || cfg.Key == "" || cfg.Secret == "" {
return "", errors.New("email: SMTP configuration (url, key, secret) is incomplete")
}
if target == "" {
return "", errors.New("email: target email address is required")
}
title := bodyTitle(body)
content := bodyContent(body, "<p><b>%s</b>: %v</p>", "")
fromName := "System Notification"
if ext != nil {
if fn, ok := ext["from_name"].(string); ok && fn != "" {
fromName = fn
}
}
htmlBody := fmt.Sprintf(`<html><body><h2>%s</h2><div>%s</div></body></html>`, title, content)
host, portStr, err := net.SplitHostPort(cfg.URL)
port := 25
if err != nil {
host = cfg.URL
} else if p, err := strconv.Atoi(portStr); err == nil && p > 0 {
port = p
}
mailCfg := pkgmail.Config{
Host: host,
Port: port,
Username: cfg.Key,
Password: cfg.Secret,
FromName: fromName,
}
if err := pkgmail.SendMail(ctx, mailCfg, target, title, htmlBody); err != nil {
return "", fmt.Errorf("email: send smtp mail failed: %w", err)
}
return "", nil
}
// ValidateConfig 校验邮件 SMTP 配置
func (p *EmailPusher) ValidateConfig(cfg Config) error {
if cfg.URL == "" {
return errors.New("SMTP host:port is required")
}
if cfg.Key == "" {
return errors.New("SMTP username is required")
}
if cfg.Secret == "" {
return errors.New("SMTP password is required")
}
return nil
}
@@ -0,0 +1,65 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"testing"
)
func TestEmailPusherValidateConfig(t *testing.T) {
pusher := &EmailPusher{}
tests := []struct {
name string
cfg Config
wantErr bool
}{
{
name: "empty url",
cfg: Config{URL: "", Key: "user", Secret: "pass"},
wantErr: true,
},
{
name: "empty key",
cfg: Config{URL: "smtp.example.com:587", Key: "", Secret: "pass"},
wantErr: true,
},
{
name: "empty secret",
cfg: Config{URL: "smtp.example.com:587", Key: "user", Secret: ""},
wantErr: true,
},
{
name: "valid config",
cfg: Config{URL: "smtp.example.com:587", Key: "user", Secret: "pass"},
wantErr: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := pusher.ValidateConfig(tt.cfg)
if (err != nil) != tt.wantErr {
t.Errorf("ValidateConfig() error = %v, wantErr %v", err, tt.wantErr)
}
})
}
}
func TestEmailPusherSendValidation(t *testing.T) {
pusher := &EmailPusher{}
// Missing target
_, err := pusher.Send(context.Background(), Config{URL: "127.0.0.1:25", Key: "u", Secret: "p"}, "", map[string]any{"title": "hi"}, "", nil)
if err == nil {
t.Errorf("expected error for empty target, got nil")
}
// Missing config
_, err = pusher.Send(context.Background(), Config{}, "test@example.com", map[string]any{"title": "hi"}, "", nil)
if err == nil {
t.Errorf("expected error for empty config, got nil")
}
}
@@ -0,0 +1,52 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"errors"
"fmt"
"strings"
"github.com/nikoksr/notify"
"github.com/nikoksr/notify/service/lark"
)
func init() {
Register("lark", &LarkPusher{})
}
// LarkPusher 基于 nikoksr/notify 的飞书 Webhook 机器人推送实现
type LarkPusher struct{}
// Send 发送飞书通知
func (p *LarkPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, _ string, _ map[string]any) (string, error) {
if cfg.URL == "" {
return "", errors.New("lark: webhook URL is required")
}
title := bodyTitle(body)
content := bodyContent(body, "**%s**: %v", "\n")
larkService := lark.NewWebhookService(cfg.URL)
notifier := notify.New()
notifier.UseServices(larkService)
if err := notifier.Send(ctx, title, content); err != nil {
return "", fmt.Errorf("lark: notify send failed: %w", err)
}
return "ok", nil
}
// ValidateConfig 校验飞书机器人配置
func (p *LarkPusher) ValidateConfig(cfg Config) error {
if cfg.URL == "" {
return errors.New("webhook URL is required")
}
if !strings.HasPrefix(cfg.URL, "https://") {
return errors.New("webhook URL must use https:// protocol")
}
return nil
}
@@ -0,0 +1,69 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package push 提供解耦的、无外部业务依赖 of 通知推送底层实现
package push
import (
"context"
"fmt"
"sync"
"time"
)
const (
defaultTitle = "系统通知"
levelInfo = "INFO"
defaultHTTPClientTimeout = 10 * time.Second
maxResponseBodyBytes = 4096
)
// Config 基础通知渠道配置
type Config struct {
Channel string `json:"channel"` // 渠道名称,例如 "lark", "custom", "email" 等,唯一标识
URL string `json:"url,omitempty"` // Webhook 地址或 SMTP 地址
Secret string `json:"secret,omitempty"` // 签名密钥或 SMTP 密码/Token
Key string `json:"key,omitempty"` // AppID 或 SMTP 用户名
Other string `json:"other,omitempty"` // 附加配置 (如 ChatID / UserKey / 扩展 JSON)
Ext map[string]any `json:"ext,omitempty"` // 预留拓展 JSON 配置
}
// Pusher 通知推送渠道接口
type Pusher interface {
// Send 发送通知消息
// target: 发送目标 (如邮箱地址或特定用户标识;若为 bot 机器人此项为空)
// body: 消息体数据 (含默认字段如 title, content, level)
// template: 消息卡片/模板 JSON (可选)
// ext: 预留的单次发送拓展数据
// 返回 upstreamResp: 上游服务返回的响应内容(如 Webhook 响应体),用于任务日志审计;无响应时为空字符串
Send(ctx context.Context, cfg Config, target string, body map[string]any, template string, ext map[string]any) (upstreamResp string, err error)
// ValidateConfig 校验渠道配置合法性
ValidateConfig(cfg Config) error
}
var (
pushersMu sync.RWMutex
pushers = make(map[string]Pusher)
)
// Register 注册一个推送渠道实现
func Register(channelType string, pusher Pusher) {
pushersMu.Lock()
defer pushersMu.Unlock()
if pusher == nil {
panic("push: Register pusher is nil")
}
pushers[channelType] = pusher
}
// GetPusher 获取指定类型的推送渠道实现
func GetPusher(channelType string) (Pusher, error) {
pushersMu.RLock()
defer pushersMu.RUnlock()
pusher, ok := pushers[channelType]
if !ok {
return nil, fmt.Errorf("push: unknown channel type %q", channelType)
}
return pusher, nil
}
@@ -0,0 +1,69 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"errors"
"fmt"
"github.com/nikoksr/notify"
"github.com/nikoksr/notify/service/pushover"
)
func init() {
Register("pushover", &PushoverPusher{})
}
// PushoverPusher 基于 nikoksr/notify 的 Pushover 移动端推送实现
type PushoverPusher struct{}
// Send 发送 Pushover 通知
func (p *PushoverPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, _ map[string]any) (string, error) {
appToken := cfg.Key
if appToken == "" {
appToken = cfg.Secret
}
if appToken == "" {
return "", errors.New("pushover: app token is required")
}
userKey := cfg.URL
if userKey == "" {
userKey = cfg.Other
}
if target != "" {
userKey = target
}
if userKey == "" {
return "", errors.New("pushover: user key is required")
}
title := bodyTitle(body)
content := bodyContent(body, "%s: %v", "\n")
poService := pushover.New(appToken)
poService.AddReceivers(userKey)
notifier := notify.New()
notifier.UseServices(poService)
if err := notifier.Send(ctx, title, content); err != nil {
return "", fmt.Errorf("pushover: notify send failed: %w", err)
}
return "ok", nil
}
// ValidateConfig 校验 Pushover 配置
func (p *PushoverPusher) ValidateConfig(cfg Config) error {
appToken := cfg.Key
if appToken == "" {
appToken = cfg.Secret
}
if appToken == "" {
return errors.New("app token is required")
}
return nil
}
@@ -0,0 +1,69 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"errors"
"fmt"
"github.com/nikoksr/notify"
"github.com/nikoksr/notify/service/slack"
)
func init() {
Register("slack", &SlackPusher{})
}
// SlackPusher 基于 nikoksr/notify 的 Slack 推送实现
type SlackPusher struct{}
// Send 发送 Slack 通知
func (p *SlackPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, _ map[string]any) (string, error) {
token := cfg.Key
if token == "" {
token = cfg.Secret
}
channelID := cfg.URL
if target != "" {
channelID = target
}
if channelID == "" {
channelID = cfg.Other
}
if token == "" {
return "", errors.New("slack: bot/api token is required")
}
if channelID == "" {
return "", errors.New("slack: channel ID is required")
}
title := bodyTitle(body)
content := bodyContent(body, "*%s*: %v", "\n")
slackService := slack.New(token)
slackService.AddReceivers(channelID)
notifier := notify.New()
notifier.UseServices(slackService)
if err := notifier.Send(ctx, title, content); err != nil {
return "", fmt.Errorf("slack: notify send failed: %w", err)
}
return "ok", nil
}
// ValidateConfig 校验 Slack 配置
func (p *SlackPusher) ValidateConfig(cfg Config) error {
token := cfg.Key
if token == "" {
token = cfg.Secret
}
if token == "" {
return errors.New("slack token is required")
}
return nil
}
@@ -0,0 +1,75 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"errors"
"fmt"
"strconv"
"github.com/nikoksr/notify"
"github.com/nikoksr/notify/service/telegram"
)
func init() {
Register("telegram", &TelegramPusher{})
}
// TelegramPusher 基于 nikoksr/notify 的 Telegram 机器人推送实现
type TelegramPusher struct{}
// Send 执行 Telegram 消息发送
func (p *TelegramPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, _ map[string]any) (string, error) {
botToken := cfg.Secret
if botToken == "" {
botToken = cfg.Key
}
if botToken == "" {
return "", errors.New("telegram: bot token is required")
}
chatIDStr := target
if chatIDStr == "" {
chatIDStr = cfg.Other
}
if chatIDStr == "" {
return "", errors.New("telegram: chat_id is required")
}
chatID, err := strconv.ParseInt(chatIDStr, 10, 64)
if err != nil {
return "", fmt.Errorf("telegram: invalid chat_id %q: %w", chatIDStr, err)
}
title := bodyTitle(body)
content := bodyContent(body, "%s: %v", "\n")
tgService, err := telegram.New(botToken)
if err != nil {
return "", fmt.Errorf("telegram: init service failed: %w", err)
}
tgService.AddReceivers(chatID)
notifier := notify.New()
notifier.UseServices(tgService)
if err := notifier.Send(ctx, title, content); err != nil {
return "", fmt.Errorf("telegram: notify send failed: %w", err)
}
return "ok", nil
}
// ValidateConfig 校验 Telegram 机器人配置
func (p *TelegramPusher) ValidateConfig(cfg Config) error {
token := cfg.Secret
if token == "" {
token = cfg.Key
}
if token == "" {
return errors.New("bot token is required")
}
return nil
}
@@ -0,0 +1,28 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"testing"
)
func TestTelegramPusherValidation(t *testing.T) {
pusher := &TelegramPusher{}
err := pusher.ValidateConfig(Config{Secret: "123456:ABC-DEF1234ghIkl-zyx57W2v1u123ew11"})
if err != nil {
t.Errorf("ValidateConfig failed: %v", err)
}
err = pusher.ValidateConfig(Config{})
if err == nil {
t.Errorf("expected error for empty config, got nil")
}
_, err = pusher.Send(context.Background(), Config{Secret: "123:token"}, "not-a-number", map[string]any{"title": "test"}, "", nil)
if err == nil {
t.Errorf("expected error for invalid chat_id, got nil")
}
}
@@ -0,0 +1,346 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"bytes"
"encoding/json"
"fmt"
"maps"
"regexp"
"slices"
"strconv"
"strings"
"text/template"
"time"
)
var (
// placeholderRegex matches {{ ... }} tags
placeholderRegex = regexp.MustCompile(`\{\{\s*([^}]+?)\s*\}\}`)
// identifierRegex matches simple identifiers like name or user.username
identifierRegex = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_.]*$`)
)
// jsonMap is a map that serializes to JSON when printed as a string in templates.
type jsonMap map[string]any
func (m jsonMap) String() string {
b, err := json.Marshal(map[string]any(m))
if err != nil {
return fmt.Sprintf("%v", map[string]any(m))
}
return string(b)
}
func (m jsonMap) MarshalJSON() ([]byte, error) {
return json.Marshal(map[string]any(m))
}
// jsonSlice is a slice that serializes to JSON when printed as a string in templates.
type jsonSlice []any
func (s jsonSlice) String() string {
b, err := json.Marshal([]any(s))
if err != nil {
return fmt.Sprintf("%v", []any(s))
}
return string(b)
}
func (s jsonSlice) MarshalJSON() ([]byte, error) {
return json.Marshal([]any(s))
}
var defaultFuncMap = template.FuncMap{
"default": func(fallback any, val any) any {
if val == nil {
return fallback
}
switch v := val.(type) {
case string:
if v == "" {
return fallback
}
case bool:
if !v {
return fallback
}
case int:
if v == 0 {
return fallback
}
case int32:
if v == 0 {
return fallback
}
case int64:
if v == 0 {
return fallback
}
case float64:
if v == 0 {
return fallback
}
}
return val
},
"toJson": func(v any) string {
b, err := json.Marshal(v)
if err != nil {
return fmt.Sprint(v)
}
return string(b)
},
"upper": strings.ToUpper,
"lower": strings.ToLower,
"trim": strings.TrimSpace,
"dateFormat": func(format string, t any) string {
switch v := t.(type) {
case time.Time:
return v.Format(format)
case *time.Time:
if v != nil {
return v.Format(format)
}
}
return fmt.Sprint(t)
},
}
// hasKey checks if a dot-delimited or plain key exists in body
func hasKey(body map[string]any, key string) bool {
if _, ok := body[key]; ok {
return true
}
parts := strings.Split(key, ".")
var cur any = body
for _, part := range parts {
m, ok := cur.(map[string]any)
if !ok {
return false
}
val, exists := m[part]
if !exists {
return false
}
cur = val
}
return true
}
// normalizeTemplate converts legacy {{key}} / {{user.name}} into Go template {{.user.name}}
// while preserving Go template keywords, dot expressions, pipelines, and missing placeholders.
func normalizeTemplate(tmpl string, body map[string]any) string {
return placeholderRegex.ReplaceAllStringFunc(tmpl, func(match string) string {
sub := strings.TrimSpace(match[2 : len(match)-2])
if sub == "" {
return match
}
// If it's already a dot expression or special variable ($...)
if strings.HasPrefix(sub, ".") || strings.HasPrefix(sub, "$") {
return match
}
// If it's a known Go template keyword or block
firstWord := strings.Fields(sub)[0]
switch firstWord {
case "if", "else", "end", "range", "with", "template", "define", "block", "nil", "true", "false":
return match
}
// Check if it's a pipeline like `key | default "val"`
if strings.Contains(sub, "|") {
const pipelineSplitCount = 2
parts := strings.SplitN(sub, "|", pipelineSplitCount)
left := strings.TrimSpace(parts[0])
right := strings.TrimSpace(parts[1])
if identifierRegex.MatchString(left) && !strings.HasPrefix(left, ".") && !strings.HasPrefix(left, "$") {
return fmt.Sprintf("{{ .%s | %s }}", left, right)
}
return match
}
// Simple identifier: if present in body, convert to dot expression; otherwise preserve as is for fallback
if identifierRegex.MatchString(sub) {
if hasKey(body, sub) {
return fmt.Sprintf("{{ .%s }}", sub)
}
// Missing key: keep original text so fallback or literal is preserved
return match
}
return match
})
}
// prepareContext pre-processes the body map so that:
// 1. Dotted keys like "user.name" are expanded to nested map structure.
// 2. Complex structs, slices, and maps have JSON-friendly string representations when directly interpolated.
func prepareContext(body map[string]any) jsonMap {
if body == nil {
return make(jsonMap)
}
ctx := make(jsonMap, len(body))
for k, v := range body {
formatted := formatContextValue(v)
ctx[k] = formatted
// If key contains '.', expand into nested hierarchy
if strings.Contains(k, ".") {
parts := strings.Split(k, ".")
cur := ctx
for i := 0; i < len(parts)-1; i++ {
sub, ok := cur[parts[i]].(jsonMap)
if !ok {
sub = make(jsonMap)
cur[parts[i]] = sub
}
cur = sub
}
cur[parts[len(parts)-1]] = formatted
}
}
return ctx
}
// formatContextValue formats slices and maps to JSON representation for direct string printing,
// while preserving basic scalar types for template functions.
func formatContextValue(v any) any {
if v == nil {
return ""
}
switch val := v.(type) {
case string, int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64, float32, float64, bool, time.Time:
return val
case []byte:
return string(val)
case map[string]any:
jm := make(jsonMap, len(val))
for k, subVal := range val {
jm[k] = formatContextValue(subVal)
}
return jm
case []any:
js := make(jsonSlice, len(val))
for i, subVal := range val {
js[i] = formatContextValue(subVal)
}
return js
case []string:
js := make(jsonSlice, len(val))
for i, subVal := range val {
js[i] = subVal
}
return js
default:
return val
}
}
// ParseTemplate parses template strings by replacing {{placeholder}} structures with values from body.
// It supports Go text/template expressions (e.g. if/else, pipelines, default, toJson) as well as legacy {{key}} placeholders.
func ParseTemplate(templateStr string, body map[string]any) string {
if templateStr == "" {
return ""
}
normalized := normalizeTemplate(templateStr, body)
ctx := prepareContext(body)
tmpl, err := template.New("push_tmpl").
Funcs(defaultFuncMap).
Option("missingkey=zero").
Parse(normalized)
if err != nil {
return fallbackReplace(templateStr, body)
}
var buf bytes.Buffer
if err = tmpl.Execute(&buf, ctx); err != nil {
return fallbackReplace(templateStr, body)
}
return buf.String()
}
func fallbackReplace(template string, body map[string]any) string {
var buf strings.Builder
buf.Grow(len(template))
i := 0
for {
pos := strings.Index(template[i:], "{{")
if pos == -1 {
buf.WriteString(template[i:])
break
}
buf.WriteString(template[i : i+pos])
i += pos + 2
endPos := strings.Index(template[i:], "}}")
if endPos == -1 {
buf.WriteString("{{")
buf.WriteString(template[i:])
break
}
key := strings.TrimSpace(template[i : i+endPos])
key = strings.TrimPrefix(key, ".")
if val, ok := body[key]; ok {
buf.WriteString(formatValue(val))
} else {
buf.WriteString("{{")
buf.WriteString(template[i : i+endPos])
buf.WriteString("}}")
}
i += endPos + 2
}
return buf.String()
}
func formatValue(v any) string {
if v == nil {
return ""
}
switch val := v.(type) {
case string:
return val
case []byte:
return string(val)
case int:
return strconv.Itoa(val)
case int32:
return strconv.FormatInt(int64(val), 10)
case int64:
return strconv.FormatInt(val, 10)
case float64:
return strconv.FormatFloat(val, 'f', -1, 64)
case bool:
return strconv.FormatBool(val)
default:
b, err := json.Marshal(v)
if err == nil {
return string(b)
}
return fmt.Sprintf("%v", v)
}
}
// bodyTitle returns the notification title, falling back to the default.
func bodyTitle(body map[string]any) string {
if t, ok := body["title"].(string); ok && t != "" {
return t
}
return defaultTitle
}
// bodyContent returns the notification body, rendering every entry with format
// (a "%s … %v" pair) and joining them with sep when no content field is given.
// Entries render in sorted key order so identical bodies always produce
// identical text.
func bodyContent(body map[string]any, format, sep string) string {
if c, ok := body["content"].(string); ok && c != "" {
return c
}
parts := make([]string, 0, len(body))
for _, k := range slices.Sorted(maps.Keys(body)) {
parts = append(parts, fmt.Sprintf(format, k, body[k]))
}
return strings.Join(parts, sep)
}
@@ -0,0 +1,150 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestParseTemplate(t *testing.T) {
tests := []struct {
name string
template string
body map[string]any
expected string
}{
{
name: "simple replacement",
template: "hello {{name}}",
body: map[string]any{"name": "world"},
expected: "hello world",
},
{
name: "multiple replacements",
template: "{{greeting}} {{name}}!",
body: map[string]any{"greeting": "Hello", "name": "Alice"},
expected: "Hello Alice!",
},
{
name: "missing key preserves placeholder",
template: "hello {{name}} and {{other}}",
body: map[string]any{"name": "world"},
expected: "hello world and {{other}}",
},
{
name: "unbalanced placeholders",
template: "hello {{name",
body: map[string]any{"name": "world"},
expected: "hello {{name",
},
{
name: "nil value",
template: "val: {{val}}",
body: map[string]any{"val": nil},
expected: "val: ",
},
{
name: "basic types",
template: "int: {{i}}, float: {{f}}, bool: {{b}}",
body: map[string]any{"i": 123, "f": 45.67, "b": true},
expected: "int: 123, float: 45.67, bool: true",
},
{
name: "complex type slice",
template: "items: {{items}}",
body: map[string]any{"items": []string{"a", "b"}},
expected: `items: ["a","b"]`,
},
{
name: "complex type map",
template: "obj: {{obj}}",
body: map[string]any{"obj": map[string]any{"key": "value"}},
expected: `obj: {"key":"value"}`,
},
{
name: "nested property from flat map",
template: "hello {{user.username}}",
body: map[string]any{"user.username": "Alice"},
expected: "hello Alice",
},
{
name: "nested property from nested map",
template: "hello {{user.username}}",
body: map[string]any{"user": map[string]any{"username": "Bob"}},
expected: "hello Bob",
},
{
name: "go template dot syntax",
template: "hello {{.user.username}}",
body: map[string]any{"user": map[string]any{"username": "Charlie"}},
expected: "hello Charlie",
},
{
name: "default value helper fallback",
template: "hello {{.nickname | default \"Guest\"}}",
body: map[string]any{"nickname": ""},
expected: "hello Guest",
},
{
name: "default value helper provided",
template: "hello {{.nickname | default \"Guest\"}}",
body: map[string]any{"nickname": "David"},
expected: "hello David",
},
{
name: "conditional if else true",
template: "{{if .is_admin}}Admin: {{.name}}{{else}}User: {{.name}}{{end}}",
body: map[string]any{"is_admin": true, "name": "Eve"},
expected: "Admin: Eve",
},
{
name: "conditional if else false",
template: "{{if .is_admin}}Admin: {{.name}}{{else}}User: {{.name}}{{end}}",
body: map[string]any{"is_admin": false, "name": "Frank"},
expected: "User: Frank",
},
{
name: "upper and lower helper",
template: "{{.title | upper}} - {{.level | lower}}",
body: map[string]any{"title": "Warning", "level": "INFO"},
expected: "WARNING - info",
},
{
name: "toJson helper",
template: "payload: {{toJson .data}}",
body: map[string]any{"data": map[string]any{"status": "ok"}},
expected: `payload: {"status":"ok"}`,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := ParseTemplate(tt.template, tt.body)
assert.Equal(t, tt.expected, result)
})
}
}
// Synthesized content must render in a stable order, otherwise two identical
// notifications produce different text on every send.
func TestBodyContentFallbackIsDeterministic(t *testing.T) {
body := map[string]any{
"zebra": 1,
"alpha": 2,
"mike": 3,
"charlie": 4,
"yankee": 5,
}
first := bodyContent(body, "%s=%v", ",")
for i := 1; i <= 50; i++ {
if got := bodyContent(body, "%s=%v", ","); got != first {
t.Fatalf("bodyContent order changed on call %d: %q != %q", i, got, first)
}
}
assert.Equal(t, "alpha=2,charlie=4,mike=3,yankee=5,zebra=1", first)
}
@@ -0,0 +1,38 @@
// 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())
}
}
@@ -0,0 +1,311 @@
// 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"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strings"
"time"
"github.com/tencent-connect/botgo/token"
)
const defaultTelegramAPI = "https://api.telegram.org"
// ListDefinitions returns the admin form schema of every supported channel type.
func ListDefinitions() []do.Definition {
return []do.Definition{
{
Type: consts.MessageChannelTypeTelegram,
Fields: []do.Field{
{Key: "token", Type: consts.TypePassword, Required: true},
{Key: "api_base", Type: consts.TypeText, Required: false},
},
},
{
Type: consts.MessageChannelTypeQQ,
Fields: []do.Field{
{Key: "app_id", Type: consts.TypeText, Required: true},
{Key: "client_secret", Type: "password", Required: true},
},
},
}
}
// CreateChannel validates the admin payload and persists an encrypted channel.
func CreateChannel(ctx context.Context, req do.CreateChannelRequest) (do.ChannelDTO, error) {
name := strings.TrimSpace(req.Name)
if name == "" {
return do.ChannelDTO{}, errors.New(consts.ErrNameRequired)
}
channelType := strings.TrimSpace(req.Type)
if channelType != consts.MessageChannelTypeTelegram && channelType != consts.MessageChannelTypeQQ {
return do.ChannelDTO{}, errors.New(consts.ErrTypeInvalid)
}
creds := req.Credentials
if creds == nil {
creds = map[string]string{}
}
if err := ValidateCredentials(channelType, creds, false); err != nil {
return do.ChannelDTO{}, err
}
cipher, err := EncryptCredentials(creds)
if err != nil {
return do.ChannelDTO{}, err
}
extra := req.Extra
if extra == nil {
extra = map[string]string{}
}
enabled := true
if req.Enabled != nil {
enabled = *req.Enabled
}
row := &entity.MessageChannel{
Name: name,
Type: channelType,
OwnerScope: consts.MessageOwnerScopeSystem,
Enabled: enabled,
Credentials: cipher,
Extra: EncodeExtra(extra),
}
if err := dao.CreateMessageChannel(ctx, row); err != nil {
return do.ChannelDTO{}, err
}
return ToDTO(row, creds, extra), nil
}
// UpdateChannel patches a channel; empty secrets keep the stored ciphertext.
func UpdateChannel(ctx context.Context, id uint64, req do.UpdateChannelRequest) (do.ChannelDTO, error) {
row, err := dao.GetMessageChannel(ctx, id)
if err != nil {
if errors.Is(err, consts.ErrRecordNotFound) {
return do.ChannelDTO{}, errors.New(consts.ErrChannelNotFound)
}
return do.ChannelDTO{}, err
}
creds, err := DecryptCredentials(row.Credentials)
if err != nil {
return do.ChannelDTO{}, err
}
extra := ParseExtra(row.Extra)
if name := strings.TrimSpace(req.Name); name != "" {
row.Name = name
}
if req.Enabled != nil {
row.Enabled = *req.Enabled
}
if req.Extra != nil {
extra = req.Extra
}
if len(req.Credentials) > 0 {
merged := make(map[string]string, len(creds))
for k, v := range creds {
merged[k] = v
}
for k, v := range req.Credentials {
if strings.TrimSpace(v) == "" {
continue
}
merged[k] = v
}
if err := ValidateCredentials(row.Type, merged, true); err != nil {
return do.ChannelDTO{}, err
}
creds = merged
}
cipher, err := EncryptCredentials(creds)
if err != nil {
return do.ChannelDTO{}, err
}
row.Credentials = cipher
row.Extra = EncodeExtra(extra)
if err := dao.UpdateMessageChannel(ctx, row); err != nil {
return do.ChannelDTO{}, err
}
return ToDTO(row, creds, extra), nil
}
// ListChannels returns every channel with secrets masked.
func ListChannels(ctx context.Context) ([]do.ChannelDTO, error) {
rows, err := dao.ListMessageChannels(ctx)
if err != nil {
return nil, err
}
out := make([]do.ChannelDTO, 0, len(rows))
for i := range rows {
creds, _ := DecryptCredentials(rows[i].Credentials)
extra := ParseExtra(rows[i].Extra)
out = append(out, ToDTO(&rows[i], creds, extra))
}
return out, nil
}
// DeleteChannel removes a channel together with its bindings and pairing codes.
func DeleteChannel(ctx context.Context, id uint64) error {
if _, err := dao.GetMessageChannel(ctx, id); err != nil {
if errors.Is(err, consts.ErrRecordNotFound) {
return errors.New(consts.ErrChannelNotFound)
}
return err
}
return dao.DeleteMessageChannel(ctx, id)
}
// ProbeChannel verifies the stored credentials against the upstream platform.
func ProbeChannel(ctx context.Context, id uint64) error {
row, err := dao.GetMessageChannel(ctx, id)
if err != nil {
if errors.Is(err, consts.ErrRecordNotFound) {
return errors.New(consts.ErrChannelNotFound)
}
return err
}
creds, err := DecryptCredentials(row.Credentials)
if err != nil {
return err
}
switch row.Type {
case consts.MessageChannelTypeTelegram:
return ProbeTelegram(ctx, creds)
case consts.MessageChannelTypeQQ:
return ProbeQQ(ctx, creds)
default:
return errors.New(consts.ErrTypeInvalid)
}
}
// ProbeTelegram calls getMe to confirm the bot token is usable.
func ProbeTelegram(ctx context.Context, creds map[string]string) error {
tok := creds["token"]
if strings.TrimSpace(tok) == "" {
return errors.New(consts.ErrMissingTelegramToken)
}
base := creds["api_base"]
base = strings.TrimRight(strings.TrimSpace(base), "/")
if base == "" {
base = defaultTelegramAPI
}
url := fmt.Sprintf("%s/bot%s/getMe", base, tok)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return err
}
client := &http.Client{Timeout: 10 * time.Second}
resp, err := client.Do(req)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
body, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("%s (%d): %s", consts.ErrTelegramGetMeFailed, resp.StatusCode, string(body))
}
var res struct {
OK bool `json:"ok"`
}
if err := json.Unmarshal(body, &res); err != nil {
return err
}
if !res.OK {
return fmt.Errorf("%s: %s", consts.ErrTelegramNotOK, string(body))
}
return nil
}
// ProbeQQ exchanges the app credentials for an access token.
func ProbeQQ(_ context.Context, creds map[string]string) error {
appID := strings.TrimSpace(creds["app_id"])
secret := strings.TrimSpace(creds["client_secret"])
if appID == "" || secret == "" {
return errors.New(consts.ErrMissingQQCredentials)
}
credentials := &token.QQBotCredentials{
AppID: appID,
AppSecret: secret,
}
tokSrc := token.NewQQBotTokenSource(credentials)
tok, err := tokSrc.Token()
if err != nil {
return fmt.Errorf("%s: %w", consts.ErrQQTokenFetchFailed, err)
}
if tok == nil || tok.AccessToken == "" {
return errors.New(consts.ErrQQEmptyToken)
}
return nil
}
// ValidateCredentials checks the admin submitted credentials for a channel type.
func ValidateCredentials(t string, creds map[string]string, isUpdate bool) error {
switch t {
case consts.MessageChannelTypeTelegram:
tok := creds["token"]
if strings.TrimSpace(tok) == "" && !isUpdate {
return errors.New(consts.ErrTelegramTokenRequired)
}
if base, ok := creds["api_base"]; ok && strings.TrimSpace(base) != "" {
if !strings.HasPrefix(base, "http://") && !strings.HasPrefix(base, "https://") {
return errors.New(consts.ErrAPIBaseInvalid)
}
}
case consts.MessageChannelTypeQQ:
appID := creds["app_id"]
secret := creds["client_secret"]
if (strings.TrimSpace(appID) == "" || strings.TrimSpace(secret) == "") && !isUpdate {
return errors.New(consts.ErrQQCredentialsRequired)
}
default:
return errors.New(consts.ErrTypeInvalid)
}
return nil
}
// ToDTO projects a channel row onto the admin DTO with credentials masked.
func ToDTO(row *entity.MessageChannel, creds, extra map[string]string) do.ChannelDTO {
return do.ChannelDTO{
ID: row.ID,
Name: row.Name,
Type: row.Type,
OwnerScope: row.OwnerScope,
OwnerID: row.OwnerID,
Enabled: row.Enabled,
Credentials: MaskCredentials(row.Type, creds),
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,191 @@
// 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"
"fmt"
"strings"
)
const (
// TaskDispatchBotMsg is the queue pattern for bot downlink dispatch.
TaskDispatchBotMsg = consts.TaskDispatchBotMsg
// TaskTypeDispatchBotMsg is the admin type identifier for bot downlink dispatch.
TaskTypeDispatchBotMsg = consts.TaskTypeDispatchBotMsg
taskQueueDefault = "default"
taskParamTypeString = "string"
paramNameText = "text"
)
// BotDispatchMeta describes the bot downlink dispatch task.
var BotDispatchMeta = contracts.TaskMetaDTO{
Type: TaskTypeDispatchBotMsg,
AsynqTask: TaskDispatchBotMsg,
Name: "分发 Bot 消息",
DisplayName: "分发 Bot 消息",
Description: "向已绑定的平台账号异步下发 Bot 文本消息",
Category: "messaging",
Queue: taskQueueDefault,
Retryable: true,
Params: []contracts.TaskParamDTO{
{Name: paramNameText, Label: "消息内容", Type: consts.TypeText, Required: true, Placeholder: "要发送的文本", Description: "下发给绑定用户的文本"},
{Name: "channel_id", Label: "频道 ID", Type: taskParamTypeString, Required: false, Placeholder: "留空表示全部启用频道", Description: "仅向指定频道的绑定发送"},
{Name: "user_id", Label: "用户 ID", Type: taskParamTypeString, Required: false, Placeholder: "留空表示频道下全部绑定", Description: "仅向指定 Wavelet 用户的绑定发送"},
},
}
type botDispatchPayload struct {
Text string `json:"text"`
ChannelID uint64 `json:"channel_id,string"`
UserID uint64 `json:"user_id,string"`
}
// BotDispatchHandler sends a text message through enabled bot channels.
type BotDispatchHandler struct{}
// ValidatePayload requires a non-empty message body.
func (h *BotDispatchHandler) ValidatePayload(payload []byte) ([]byte, error) {
p, err := parseBotDispatchPayload(payload)
if err != nil {
return nil, err
}
return json.Marshal(p)
}
// Execute delivers the text to matching channel bindings.
func (h *BotDispatchHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) {
p, err := parseBotDispatchPayload(payload)
if err != nil {
return nil, err
}
channels, err := dao.ListEnabledMessageChannels(ctx)
if err != nil {
return nil, err
}
if p.ChannelID != 0 {
filtered := channels[:0]
for i := range channels {
if channels[i].ID == p.ChannelID {
filtered = append(filtered, channels[i])
}
}
channels = filtered
if len(channels) == 0 {
return nil, errors.New(consts.ErrChannelNotFound)
}
}
sent := 0
failed := 0
for i := range channels {
n, ferr := dispatchOnChannel(ctx, &channels[i], p.UserID, p.Text)
sent += n
failed += ferr
}
msg := fmt.Sprintf("Bot 消息已尝试发送,成功 %d,失败 %d", sent, failed)
if svc := GetTaskService(ctx); svc != nil {
svc.AppendLog(ctx, "%s", msg)
}
if sent == 0 && failed > 0 {
return nil, errors.New(msg)
}
return &contracts.TaskResultDTO{Message: msg}, nil
}
func parseBotDispatchPayload(payload []byte) (botDispatchPayload, error) {
var p botDispatchPayload
if len(payload) > 0 {
if err := json.Unmarshal(payload, &p); err != nil {
return p, fmt.Errorf("%s: %w", consts.ErrInvalidJSONFormat, err)
}
}
p.Text = strings.TrimSpace(p.Text)
if p.Text == "" {
return p, errors.New(consts.ErrBotDispatchTextRequired)
}
return p, nil
}
func dispatchOnChannel(ctx context.Context, row *entity.MessageChannel, userID uint64, text string) (sent, failed int) {
factory, ok := Lookup(row.Type)
if !ok {
logger.ErrorF(ctx, "bot dispatch: %s type=%s", consts.ErrBotChannelNotRegistered, row.Type)
return 0, 1
}
cfg, err := channelConfigFromRow(row)
if err != nil {
logger.ErrorF(ctx, "bot dispatch: decode channel %d: %v", row.ID, err)
return 0, 1
}
ch, err := factory(cfg, nil)
if err != nil {
logger.ErrorF(ctx, "bot dispatch: create adapter %d: %v", row.ID, err)
return 0, 1
}
if err := ch.Connect(ctx); err != nil {
logger.ErrorF(ctx, "bot dispatch: connect channel %d: %v", row.ID, err)
return 0, 1
}
defer func() { _ = ch.Disconnect(ctx) }()
bindings, err := dao.ListBindingsByChannel(ctx, row.ID)
if err != nil {
logger.ErrorF(ctx, "bot dispatch: list bindings %d: %v", row.ID, err)
return 0, 1
}
for i := range bindings {
if userID != 0 && bindings[i].UserID != userID {
continue
}
to := do.Recipient{
ChatID: bindings[i].PlatformUserID,
PlatformUserID: bindings[i].PlatformUserID,
}
if err := ch.Send(ctx, to, do.OutboundMessage{Text: text}); err != nil {
logger.ErrorF(ctx, "bot dispatch: send channel=%d user=%d: %v", row.ID, bindings[i].UserID, err)
failed++
continue
}
sent++
}
return sent, failed
}
func channelConfigFromRow(row *entity.MessageChannel) (do.ChannelConfig, error) {
creds, err := DecryptCredentials(row.Credentials)
if err != nil {
return do.ChannelConfig{}, err
}
if creds == nil {
creds = map[string]string{}
}
if creds["bot_token"] == "" && creds["token"] != "" {
creds["bot_token"] = creds["token"]
}
if creds["app_secret"] == "" && creds["client_secret"] != "" {
creds["app_secret"] = creds["client_secret"]
}
extra := ParseExtra(row.Extra)
if extra["base_url"] == "" && creds["api_base"] != "" {
extra["base_url"] = creds["api_base"]
}
return do.ChannelConfig{
ID: row.ID,
Type: row.Type,
Name: row.Name,
Credentials: creds,
Extra: extra,
}, nil
}
@@ -0,0 +1,46 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service_test
import (
"Wavelet/plugins/domain/msg_gateway/dao"
"Wavelet/plugins/domain/msg_gateway/model/entity"
"Wavelet/plugins/domain/msg_gateway/service"
"context"
"path/filepath"
"testing"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
type dispatchTestDB struct{ db *gorm.DB }
func (m *dispatchTestDB) GORM() *gorm.DB { return m.db }
func (m *dispatchTestDB) DB(ctx context.Context) *gorm.DB { return m.db.WithContext(ctx) }
func (m *dispatchTestDB) Named(_ string) *gorm.DB { return m.db }
func TestBotDispatchValidatePayload(t *testing.T) {
h := &service.BotDispatchHandler{}
_, err := h.ValidatePayload([]byte(`{}`))
require.Error(t, err)
_, err = h.ValidatePayload([]byte(`{"text":"hello"}`))
require.NoError(t, err)
}
func TestBotDispatchNoChannels(t *testing.T) {
testDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "dispatch.db")), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, testDB.AutoMigrate(&entity.MessageChannel{}, &entity.MessageBinding{}))
dao.SetDBServiceForTest(&dispatchTestDB{db: testDB})
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
h := &service.BotDispatchHandler{}
res, err := h.Execute(context.Background(), []byte(`{"text":"hello"}`))
require.NoError(t, err)
require.NotNil(t, res)
assert.Contains(t, res.Message, "成功 0")
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,405 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package service implements domain business logic and channel runners for msg_gateway.
package service
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/dao"
"Wavelet/plugins/domain/msg_gateway/model/do"
"Wavelet/plugins/domain/msg_gateway/model/entity"
"context"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"strconv"
"strings"
"sync"
"time"
"unicode"
)
// Handler processes one inbound message.
type Handler func(ctx context.Context, msg do.InboundMessage) error
// Factory constructs a Channel from decrypted config.
type Factory func(cfg do.ChannelConfig, onInbound Handler) (Channel, error)
// Channel is one connected messaging adapter.
type Channel interface {
Type() string
Connect(ctx context.Context) error
Disconnect(ctx context.Context) error
Send(ctx context.Context, to do.Recipient, msg do.OutboundMessage) error
Capabilities() do.Capability
}
var (
factoriesMu sync.RWMutex
factories = map[string]Factory{}
)
// Register stores a channel factory under typ.
func Register(typ string, fn Factory) {
factoriesMu.Lock()
defer factoriesMu.Unlock()
factories[typ] = fn
}
// Lookup returns a previously registered factory.
func Lookup(typ string) (Factory, bool) {
factoriesMu.RLock()
defer factoriesMu.RUnlock()
fn, ok := factories[typ]
return fn, ok
}
// Re-exported constants.
const (
CodeAlphabet = consts.CodeAlphabet
CodeLength = consts.CodeLength
)
// GenerateCode returns an 8-character pairing code.
func GenerateCode() (string, error) {
buf := make([]byte, CodeLength)
if _, err := rand.Read(buf); err != nil {
return "", err
}
out := make([]byte, CodeLength)
for i, b := range buf {
out[i] = CodeAlphabet[int(b)%len(CodeAlphabet)]
}
return string(out), nil
}
// NormalizeCode strips separators and uppercases.
func NormalizeCode(s string) string {
var b strings.Builder
for _, r := range s {
if r == '-' || unicode.IsSpace(r) {
continue
}
b.WriteRune(unicode.ToUpper(r))
}
return b.String()
}
// FormatCode renders ABCD-EFGH.
func FormatCode(s string) string {
s = NormalizeCode(s)
if len(s) != CodeLength {
return s
}
return s[:4] + "-" + s[4:]
}
var (
credentialSecretMu sync.RWMutex
credentialSecret string
)
// SetCredentialSecret sets the secret used to derive CredentialKey.
func SetCredentialSecret(secret string) {
credentialSecretMu.Lock()
defer credentialSecretMu.Unlock()
credentialSecret = secret
}
// CredentialKey is AES-256 hex derived from the session secret.
func CredentialKey() string {
credentialSecretMu.RLock()
secret := credentialSecret
credentialSecretMu.RUnlock()
sum := sha256.Sum256([]byte(secret))
return hex.EncodeToString(sum[:])
}
// EncryptCredentials encrypts a credential map as JSON.
func EncryptCredentials(creds map[string]string) (string, error) {
if creds == nil {
creds = map[string]string{}
}
raw, err := json.Marshal(creds)
if err != nil {
return "", err
}
return util.Encrypt(CredentialKey(), string(raw))
}
// DecryptCredentials decrypts a credential map.
func DecryptCredentials(ciphertext string) (map[string]string, error) {
if ciphertext == "" {
return map[string]string{}, nil
}
plain, err := util.Decrypt(CredentialKey(), ciphertext)
if err != nil {
return nil, err
}
var out map[string]string
if err := json.Unmarshal([]byte(plain), &out); err != nil {
return nil, err
}
if out == nil {
out = map[string]string{}
}
return out, nil
}
// ParseExtra decodes optional extra JSON into a string map.
func ParseExtra(raw string) map[string]string {
if raw == "" {
return map[string]string{}
}
var out map[string]string
if err := json.Unmarshal([]byte(raw), &out); err != nil || out == nil {
return map[string]string{}
}
return out
}
// EncodeExtra encodes extra fields as JSON.
func EncodeExtra(extra map[string]string) string {
if extra == nil {
return ""
}
raw, err := json.Marshal(extra)
if err != nil {
return ""
}
return string(raw)
}
// Runner manages lifecycle for long-lived channel adapters (WebSocket, long-polling, etc.).
type Runner struct {
mu sync.Mutex
running bool
cancel context.CancelFunc
}
// GlobalRunner is the default global runner instance.
var GlobalRunner = &Runner{}
// Start starts all background long-lived channel runners.
func Start(ctx context.Context) error {
GlobalRunner.mu.Lock()
defer GlobalRunner.mu.Unlock()
if GlobalRunner.running {
return nil
}
runCtx, cancel := context.WithCancel(ctx)
GlobalRunner.cancel = cancel
GlobalRunner.running = true
logger.InfoF(runCtx, "[MessageGateway] Starting bot channel runners...")
return nil
}
// Stop stops the channel runner.
func Stop() {
GlobalRunner.mu.Lock()
defer GlobalRunner.mu.Unlock()
if !GlobalRunner.running {
return
}
if GlobalRunner.cancel != nil {
GlobalRunner.cancel()
}
GlobalRunner.running = false
}
// Cordis contract singletons consumed by service layer.
var (
cacheMu sync.RWMutex
cacheSvc contracts.CacheService
taskMu sync.RWMutex
taskSvc contracts.TaskService
userMu sync.RWMutex
userSvc contracts.UserService
)
// SetCacheService sets the cache service.
func SetCacheService(s contracts.CacheService) {
cacheMu.Lock()
defer cacheMu.Unlock()
cacheSvc = s
}
// SetTaskService sets the task service.
func SetTaskService(s contracts.TaskService) {
taskMu.Lock()
defer taskMu.Unlock()
taskSvc = s
}
// SetUserService sets the user service.
func SetUserService(s contracts.UserService) {
userMu.Lock()
defer userMu.Unlock()
userSvc = s
}
// GetCache resolves the cache service for the context.
func GetCache(ctx context.Context) contracts.CacheService {
if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil {
return s
}
cacheMu.RLock()
s := cacheSvc
cacheMu.RUnlock()
return s
}
// GetTaskService returns the task service.
func GetTaskService(ctx context.Context) contracts.TaskService {
if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil {
return s
}
taskMu.RLock()
defer taskMu.RUnlock()
return taskSvc
}
// GetUserService resolves the user service for the context.
func GetUserService(ctx context.Context) contracts.UserService {
if s, err := core.InjectFrom[contracts.UserService](ctx); err == nil && s != nil {
return s
}
userMu.RLock()
s := userSvc
userMu.RUnlock()
return s
}
// BindChannel consumes a pairing code and binds the platform identity to the user.
func BindChannel(ctx context.Context, userID uint64, req do.BindRequest) (do.BindingDTO, error) {
channelID, err := strconv.ParseUint(strings.TrimSpace(req.ChannelID), 10, 64)
if err != nil || channelID == 0 {
return do.BindingDTO{}, consts.ErrChannelIDRequired
}
code := NormalizeCode(req.Code)
if code == "" {
return do.BindingDTO{}, consts.ErrCodeInvalid
}
pairing, err := dao.GetPairingCode(ctx, code)
if err != nil {
if errors.Is(err, consts.ErrRecordNotFound) {
return do.BindingDTO{}, consts.ErrCodeInvalid
}
return do.BindingDTO{}, err
}
if !pairing.ExpiresAt.After(time.Now()) {
return do.BindingDTO{}, consts.ErrCodeInvalid
}
if pairing.ChannelID != channelID {
return do.BindingDTO{}, consts.ErrChannelMismatch
}
ch, err := dao.GetMessageChannel(ctx, channelID)
if err != nil {
if errors.Is(err, consts.ErrRecordNotFound) {
return do.BindingDTO{}, consts.ErrCodeInvalid
}
return do.BindingDTO{}, err
}
if !ch.Enabled {
return do.BindingDTO{}, consts.ErrChannelDisabled
}
existing, err := dao.GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID)
if err != nil && !errors.Is(err, consts.ErrRecordNotFound) {
return do.BindingDTO{}, err
}
if err == nil && existing != nil {
if existing.UserID != userID {
return do.BindingDTO{}, consts.ErrPlatformAlreadyBound
}
_ = dao.DeletePairingCode(ctx, pairing.Code)
return ToBindingDTO(existing, ch), nil
}
row := &entity.MessageBinding{
UserID: userID,
ChannelID: channelID,
PlatformUserID: pairing.PlatformUserID,
}
if err := dao.CreateMessageBinding(ctx, row); err != nil {
return do.BindingDTO{}, err
}
if err := dao.DeletePairingCode(ctx, pairing.Code); err != nil {
return do.BindingDTO{}, err
}
return ToBindingDTO(row, ch), nil
}
// ListEnabledPublicChannels returns the channels a user may bind to.
func ListEnabledPublicChannels(ctx context.Context) ([]do.PublicChannelDTO, error) {
rows, err := dao.ListEnabledMessageChannels(ctx)
if err != nil {
return nil, err
}
out := make([]do.PublicChannelDTO, 0, len(rows))
for _, row := range rows {
out = append(out, do.PublicChannelDTO{ID: row.ID, Name: row.Name, Type: row.Type})
}
return out, nil
}
// ListUserBindings returns the binding rows of one user enriched with channel info.
func ListUserBindings(ctx context.Context, userID uint64) ([]do.BindingDTO, error) {
rows, err := dao.ListBindingsByUser(ctx, userID)
if err != nil {
return nil, err
}
out := make([]do.BindingDTO, 0, len(rows))
for i := range rows {
ch, err := dao.GetMessageChannel(ctx, rows[i].ChannelID)
if err != nil {
out = append(out, ToBindingDTO(&rows[i], nil))
continue
}
out = append(out, ToBindingDTO(&rows[i], ch))
}
return out, nil
}
// UnbindChannel removes a binding owned by the given user.
func UnbindChannel(ctx context.Context, userID, bindingID uint64) error {
row, err := dao.GetMessageBinding(ctx, bindingID)
if err != nil {
if errors.Is(err, consts.ErrRecordNotFound) {
return consts.ErrBindingNotFound
}
return err
}
if row.UserID != userID {
return consts.ErrBindingForbidden
}
return dao.DeleteMessageBinding(ctx, bindingID)
}
// ToBindingDTO projects a binding row and its optional channel onto the user DTO.
func ToBindingDTO(row *entity.MessageBinding, ch *entity.MessageChannel) do.BindingDTO {
dto := do.BindingDTO{
ID: row.ID,
UserID: row.UserID,
ChannelID: row.ChannelID,
PlatformUserID: row.PlatformUserID,
CreatedAt: row.CreatedAt,
}
if ch != nil {
dto.ChannelName = ch.Name
dto.ChannelType = ch.Type
}
return dto
}