mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 23:56:37 +08:00
refactor(msg_gateway): restructure and rename message_gateway aligned with custom_example
This commit is contained in:
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user