mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-07 16:16:37 +08:00
feat(cordis): add OpenFlare Cordis 架构改造设计
docs(changelog): 修正表述笔误
refactor(cordis): 磁盘缓存改用上上游能力并清理本地副本
按上游/下游归属规约:类型断言守卫已回流 Wavelet(f3d85d5,附回归用例),
本仓库删除 OpenFlare/plugins/server/pkg/cache 整包并改 import 到
Wavelet/pkg/cache/disk,同步后与上游零漂移。
验证:go build 通过;go test ./... exit 0(137 包 ok);256 条路由对拍与
232 条 swagger 操作均零差异;make build-all 四进制;前端零改动。
docs(cordis): 记录 T1 清理结果与五个复用阻塞点
refactor(cordis): server 复用上游 pkg 能力并删除等价本地副本
按上游/下游归属规约清理重复实现,删除 7 个与上游等价的本地包并改 import:
shared/response→pkg/response、pkg/{logger,mail,trace,httppool,cache/ram}→
上游同名包、infra/persistence/batchwriter→pkg/batchwriter。逐项核过差异:
httppool 逐字节相同;logger 的 Config 字段完全一致;response 的 7 个 Abort*
一致;cache/ram 换过去顺带把裸 go 变回带 panic 恢复的 util.Go。
两处非等价差异按语义处理:
- batchwriter.Stats 与 status DTO 原为类型别名,改为消费侧逐字段转换,
避免 model 反向依赖基础设施类型;
- 上游 pkg/idgen 要求显式 Init(本地副本为懒加载自动初始化),本次保留本地
副本,待与 infra 初始化一并迁移(已登记在清理计划)。
验证:go build 通过;go test ./... exit 0(138 包 ok);256 条路由对拍零差异;
make swagger 232 条操作零增减,且归一化后与旧文档深度相等——差异仅为
response.Any / logger.LogEntry 两个定义名随包路径改名,接口形状未变。
chore(cordis): 回流内核与 pkg/util 通用能力并清理 vendoring 污染
按新增的上游/下游归属规约:HandleRaw/BasePath 与版本比较、网络、格式化助手
属通用能力,已提交到 Wavelet 分支 feat/cordis-router-raw-routes,本仓库改为
纯同步获取(pkg/util 已零漂移),补丁登记保留至上游合并。
同时修掉我此前 git add -A 造成的污染:首次 vendoring 把上游工作区里被
gitignore 的运行期产物一起提交进来(upload 的 diskcache 缓存块 650 个与
driver_http/dist 前端构建物 380 个,共 12872 行/1030 文件)。sync-upstream.sh
现显式排除 uploads/dist/data/*.db,.gitignore 补上对应兜底规则。
AGENTS.md 增加上游/下游改动归属规约,并把仍指向前 Cordis 布局的硬性约束
(internal/router + Serve、internal/repository/logstore、internal/platform/bootstrap、
internal/cmd)改到当前插件路径。
验证:go build 通过;go test ./... exit 0(144 包 ok);make swagger 232 条
操作与基线逐条一致;make build-all 四进制;gofmt 干净。
feat(cordis): server 插件化并改由内核挂载控制面路由
新增 plugins/server/plugin.go:Apply 以 ctx.Router().Group(app.api_prefix)
声明根级与 /v1 全部路由;33 个注册函数由 *gin.RouterGroup 改为
core.RouterExtension,RegisterCollection 改用内核新增的 HandleRaw 保留
尾部斜杠变体,AdminMiddlewares 返回 []any(Go 不允许把 []T 展开为 ...any)。
删除 router.Serve 与 registerRoutes,装配根改为 core.App +
driver_http.New(WithEngine(router.BuildEngine())),监听、信号与优雅退出归内核;
前端 SPA 的 NoRoute 兜底因内核暂无贡献点而保留在引擎层。
路由保真证据:plugin_parity_test 对拍 baseline/routes-engine.txt 的 256 条
(方法 路径) 零差异;go test ./... exit 0(144 包 ok,含真实 handler 的
openflare/integration 用例走同一条挂载路径);make swagger 232 条操作与基线
逐条一致;golangci-lint 0 issues;make build-all 四进制;embed_frontend
标签编译通过;前端零改动。
已知待补:带 Redis 的实机 HTTP 冒烟(本机 6379 未启动,session store 与
改造前一样在建店阶段即 fatal),以及 bootstrap 的任务/设置/迁移注册迁入 Apply。
feat(core): RouterExtension 增加 HandleRaw 与 BasePath 以保真尾部斜杠路由
server 插件化的前置:Handle 经 cleanPath 会剥掉尾部斜杠,无法表达
/resource 与 /resource/ 两条不同路由,而 OpenFlare 有 20 个历史 list
端点两者都注册且部署关闭了 RedirectTrailingSlash,缺失即 404。新增
HandleRaw 与 BasePath(作用域包装器同样登记反注册),补 extpoints 用例;
并把 router.Serve 拆出 BuildEngine 以便交给 driver_http.WithEngine 复用,
新增路由表导出 harness,固化 256 条 (方法 路径) 基线供插件化对拍。
上游补丁登记于 backend/OpenFlare/upstream-patches.md,同步脚本改为按目录
前缀输出差异并在同步后提醒确认补丁是否仍在。
验证:go build 通过;go test ./... exit 0(143 包 ok);gofmt 干净。
docs(cordis): 记录 server 插件接入内核的可行路径与内核能力缺口
feat(cordis): agent/relay/flared 落地为内核驱动插件
三个边缘守护进程各新增 plugin.go,实现 core.Plugin + core.Driver
(自定义 DriverType 与同名 profile),装配与生命周期从 main 迁入
Apply/Start/Stop:Apply 负责 JSON 配置加载、运行环境与用户确保、
openresty/frps/frpc 管理器与各服务装配;Start 以 util.Go 拉起阻塞式
runner 与 GeoIP 周期更新;Stop 收敛主循环结果并在超时时报错而非静默。
入口改为 core.NewApp(core.WithProfile(...)) + Prepare/Run,保持
-config 旗标、默认路径、退出码与启动/停止日志不变。
验证:go build 通过;go test ./... exit 0(143 包 ok,含 3 个插件身份
与配置失败路径测试);make build-all 四进制产出;三进制实跑缺失配置
均 exit 1 且错误链保留 load {agent,relay,flared} config 原因;gofmt 干净。
refactor(cordis): 按功能职责拆分为 4 个插件与 share 共享层
backend/OpenFlare 不再平铺遗留分层,改为 plugins/{server,agent,relay,flared}
加 share/:控制面业务(openflare/admin/oauth/user/upload/cap/config/health 与
repository/model/infra/router 等支撑层)归 server;三个边缘守护进程各自成插件;
被两个以上插件消费的 protocol/geoip/wsclient/render/pagesarchive/edge 归 share。
同时把 pkg/util 与 buildinfo 合并回上游 pkg(上游已覆盖全部符号,仅 8 个函数与
2 个类型为 OpenFlare 独有,已一并迁入),装配根统一到 backend/cmd(含三个 daemon
入口),Dockerfile 与 release 工作流的构建路径和 -X 注入路径同步更新。
验证:go build 通过;go test ./... exit 0(141 包 ok);make swagger exit 0 且
232 条 API 操作与基线逐条一致;make build-all 产出 4 进制;-X 注入经二进制
strings 实测生效;日志后端直连门禁改写为按 server 插件业务域扫描并在扫描数为 0
时报错(防门禁静默失效);前端零改动。
feat(cordis): 落地 backend/share 共享层与上游同步脚本
跨插件共享资源(控制消息协议、GeoIP+iputil、边缘守护进程日志)从下游包
移入 backend/share,并声明其只能依赖 core/pkg 与标准/第三方库,禁止反向
引用下游业务与具体插件实现;新增 scripts/sync-upstream.sh 只覆盖
backend/{core,pkg,plugins},同步后 --check 报告零差异,证明与上游逐字一致。
go build 通过,go test ./... exit 0(142 包 ok),前端零改动。
refactor(cordis): 采用与 Wavelet 同构的单模块布局并引入上游内核
按上游结构落位:backend/{core,pkg,plugins} 为 Wavelet 上游拷贝,OpenFlare
全部业务收拢到上游 downstream 所对应的位置 backend/OpenFlare/,模块名保持
Wavelet 以保证上游 import 路径逐字一致、同步零改写;三个 daemon 入口移至
backend/OpenFlare/cmd,backend/cmd 与 main.go 作为控制面装配根。
行为不变:go build 通过,142 个测试包全绿(含上游插件测试),232 条 API
操作与改造前逐条一致,四进制产物正常,前端零改动。swagger 暂只扫描下游代码,
待 P4 挂载上游路由后再纳入 plugins/。
style: 修正模块路径改写导致的 import 分组排序漂移
refactor(layout): Go 代码迁入 backend/ 并将模块名简化为 OpenFlare
对齐上游 Wavelet 的仓库布局,为以第二 module 形态 vendoring Cordis 内核与
平台插件做准备:模块路径整体改写为 OpenFlare,Go 目标加 cd backend,
swaggo 产物移至 backend/docs 并把 json/yaml 复制回 docs/ 供站点消费,
Dockerfile 与 release 工作流的构建目录、ldflags 模块路径同步更新。
行为保持不变:232 条路由与改造前逐条一致,95 个测试包全绿,
四进制产物正常,前端零改动。
chore(cordis): 落地改造计划与 schema/路由基线
新增 legacy_dump_test 迁移快照 harness:在临时 sqlite 库上按生产顺序
(goose.UpTo → zone 导入 → goose.Up)跑完 76 个历史迁移并导出 schema 与
版本序列,作为改造前后一致性门禁的唯一事实来源。同时记录 232 条路由清单
与 foundation 实施计划。
docs(cordis): add OpenFlare Cordis 架构改造设计
明确上游以第二 module 形态 vendoring 进 backend/Wavelet、4 个插件
(server/agent/relay/flared) 全部装载内核,并规定保留 76 个历史 goose
迁移 + 一次性版本 stamp 桥接的迁移方案,配套三方 schema 一致性门禁,
确保已部署库不重跑历史、不丢数据。
This commit is contained in:
@@ -0,0 +1,163 @@
|
||||
// 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/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_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 model.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 model.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 model.ChannelTypeQQ }
|
||||
|
||||
// Capabilities reports C2C text/media support.
|
||||
func (a *Adapter) Capabilities() model.Capability {
|
||||
return model.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 model.Recipient, msg model.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, model.InboundMessage{
|
||||
ChannelID: a.cfg.ID,
|
||||
ChannelType: model.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/message_gateway/model"
|
||||
"context"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestHandleEvent_DropsNonC2C(t *testing.T) {
|
||||
var got int
|
||||
a := &Adapter{onInbound: func(ctx context.Context, msg model.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 model.InboundMessage
|
||||
a := &Adapter{cfg: model.ChannelConfig{ID: 3}, onInbound: func(ctx context.Context, msg model.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(model.ChannelConfig{}, nil)
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
// 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/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_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 model.ChannelConfig
|
||||
onInbound service.Handler
|
||||
bot *tele.Bot
|
||||
}
|
||||
|
||||
// New constructs a Telegram adapter. Call service.Register from the runner.
|
||||
func New(cfg model.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 model.ChannelTypeTelegram }
|
||||
|
||||
// Capabilities reports private-chat media support.
|
||||
func (a *Adapter) Capabilities() model.Capability {
|
||||
return model.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 model.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 model.Recipient, msg model.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 := model.InboundMessage{
|
||||
ChannelID: a.cfg.ID,
|
||||
ChannelType: model.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, []model.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 "", []model.Attachment{{Error: err.Error()}}
|
||||
}
|
||||
out := make([]model.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, model.Attachment{FileName: names[i], Error: err.Error()})
|
||||
continue
|
||||
}
|
||||
out = append(out, model.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/message_gateway/model"
|
||||
"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(model.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(ctx context.Context, msg model.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 model.InboundMessage
|
||||
a := &Adapter{
|
||||
cfg: model.ChannelConfig{ID: 7, Type: "telegram"},
|
||||
onInbound: func(ctx context.Context, msg model.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(model.ChannelConfig{}, nil)
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package errs defines error sentinels and user-facing error message constants
|
||||
// for the message_gateway plugin.
|
||||
package errs
|
||||
|
||||
import "errors"
|
||||
|
||||
// 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 repository 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 repository 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"
|
||||
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 handler
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_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=[]model.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=[]model.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, errs.ErrInvalidChannelID)
|
||||
return 0, false
|
||||
}
|
||||
return id, true
|
||||
}
|
||||
|
||||
func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
|
||||
if err.Error() == errs.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 model.CreateChannelRequest true "create body"
|
||||
// @Success 200 {object} response.Any{data=model.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 model.UpdateChannelRequest true "update body"
|
||||
// @Success 200 {object} response.Any{data=model.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,182 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package handler provides HTTP endpoints for message_gateway.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_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=[]model.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, errs.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=[]model.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, errs.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 model.BindRequest true "bind body"
|
||||
// @Success 200 {object} response.Any{data=model.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, errs.ErrLoginRequired)
|
||||
return
|
||||
}
|
||||
var req model.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, errs.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, errs.ErrLoginRequired)
|
||||
return
|
||||
}
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errs.ErrInvalidBindingID)
|
||||
return
|
||||
}
|
||||
if err := service.UnbindChannel(c.Request.Context(), user.ID, id); err != nil {
|
||||
if errors.Is(err, errs.ErrBindingNotFound) {
|
||||
response.AbortNotFound(c, err.Error())
|
||||
return
|
||||
}
|
||||
if errors.Is(err, errs.ErrBindingForbidden) {
|
||||
response.AbortForbidden(c, err.Error())
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
@@ -0,0 +1,96 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package handler
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_gateway/service"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ListPushChannelDefinitions returns channel definitions.
|
||||
func ListPushChannelDefinitions(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(model.ListPushDefinitions()))
|
||||
}
|
||||
|
||||
// ListPushChannels lists configured push channels.
|
||||
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, errs.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, errs.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, errs.ErrChannelNotFound)
|
||||
return
|
||||
}
|
||||
fallback(c, err.Error())
|
||||
}
|
||||
|
||||
// CreatePushChannel creates a push channel.
|
||||
func CreatePushChannel(c *gin.Context) {
|
||||
handleJSONRequest(c, service.CreatePushChannel)
|
||||
}
|
||||
|
||||
// UpdatePushChannel updates a push channel.
|
||||
func UpdatePushChannel(c *gin.Context) {
|
||||
handleEntityUpdate(c, parsePushChannelID, service.UpdatePushChannel, func(c *gin.Context, err error) {
|
||||
handlePushChannelNotFoundError(c, err, response.AbortInternal)
|
||||
})
|
||||
}
|
||||
|
||||
// DeletePushChannel deletes a push channel.
|
||||
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 tests connectivity of a push channel.
|
||||
func TestPushChannel(c *gin.Context) {
|
||||
var req model.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,148 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package handler
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_gateway/service"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ListPushEvents lists configured push events.
|
||||
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 lists system built-in push event definitions.
|
||||
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, errs.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, errs.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, errs.ErrEventNotFound)
|
||||
return
|
||||
}
|
||||
fallback(c, err.Error())
|
||||
}
|
||||
|
||||
// CreatePushEvent creates a new push event configuration.
|
||||
func CreatePushEvent(c *gin.Context) {
|
||||
handleJSONRequest(c, service.CreatePushEvent)
|
||||
}
|
||||
|
||||
// DeletePushEvent deletes a push event configuration by ID.
|
||||
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 updates an existing push event.
|
||||
func UpdatePushEvent(c *gin.Context) {
|
||||
id, ok := parsePushEventID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
var req model.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 toggles the enabled state of a push event.
|
||||
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 returns paginated push notification delivery histories.
|
||||
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(), model.PushHistoryListFilter{
|
||||
EventKey: c.Query("event_key"),
|
||||
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 executes a synchronous push test using the specified config.
|
||||
func TestPush(c *gin.Context) {
|
||||
var req model.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,63 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package handler
|
||||
|
||||
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,51 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_gateway_test
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/message_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)
|
||||
message_gateway.SetDBServiceForTest(&mockDBService{db: testDB})
|
||||
defer func() {
|
||||
message_gateway.SetDBServiceForTest(nil)
|
||||
cleanup()
|
||||
}()
|
||||
ctx := context.Background()
|
||||
first, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ABCD1234", time.Now().Add(15*time.Minute))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
second, err := message_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,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,46 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
// 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,242 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package model defines the domain entities, DTOs, and schemas for message_gateway.
|
||||
package model
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Channel type and scope constants.
|
||||
const (
|
||||
ChannelTypeTelegram = "telegram"
|
||||
ChannelTypeQQ = "qq"
|
||||
MessageChannelTypeTelegram = "telegram"
|
||||
MessageChannelTypeQQ = "qq"
|
||||
MessageOwnerScopeSystem = "system"
|
||||
|
||||
TypeCustom = "custom"
|
||||
TypeEmail = "email"
|
||||
TypeTelegram = "telegram"
|
||||
)
|
||||
|
||||
// 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
|
||||
}
|
||||
|
||||
// MessageChannel is an admin-configured messaging adapter.
|
||||
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"
|
||||
}
|
||||
|
||||
// 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(errs.ErrChannelNameRequired)
|
||||
}
|
||||
c.Type = strings.TrimSpace(c.Type)
|
||||
if c.Type == "" {
|
||||
return errors.New(errs.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(errs.ErrEventKeyRequired)
|
||||
}
|
||||
e.Name = strings.TrimSpace(e.Name)
|
||||
if e.Name == "" {
|
||||
return errors.New(errs.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"
|
||||
}
|
||||
|
||||
// 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"`
|
||||
}
|
||||
|
||||
// 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"`
|
||||
}
|
||||
@@ -0,0 +1,313 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
||||
)
|
||||
|
||||
// Push channel and payload constants.
|
||||
const (
|
||||
ChannelCustom = "custom"
|
||||
ChannelEmail = "email"
|
||||
ChannelLark = "lark"
|
||||
ChannelTelegram = "telegram"
|
||||
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"
|
||||
)
|
||||
|
||||
// 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 (original json structure).
|
||||
func (m NotificationMessage) Flatten() map[string]any {
|
||||
res := map[string]any{
|
||||
KeyTitle: m.Title,
|
||||
KeyContent: m.Content,
|
||||
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载荷 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
|
||||
}
|
||||
|
||||
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{ChannelCustom, ChannelLark, ChannelTelegram, 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
|
||||
}
|
||||
|
||||
func init() {
|
||||
RegisterPushChannelDefinition(PushDefinition{
|
||||
Type: ChannelCustom,
|
||||
Name: "自定义消息通道",
|
||||
Description: "使用自定义 HTTP POST 请求向外部 Webhook 发送数据。",
|
||||
Fields: []PushField{
|
||||
{
|
||||
Key: KeyURL,
|
||||
Label: "请求地址",
|
||||
Type: TypeText,
|
||||
Required: true,
|
||||
Placeholder: "在此填写完整的请求地址,必须使用 HTTPS 协议",
|
||||
Description: "接口请求的完整 HTTPS URL,例如 https://api.example.com/webhook",
|
||||
},
|
||||
{
|
||||
Key: KeyOther,
|
||||
Label: "请求体 (JSON)",
|
||||
Type: TypeTextarea,
|
||||
Required: true,
|
||||
Placeholder: "在此输入请求体,支持模板变量,必须为合法的 JSON 格式",
|
||||
Description: "可使用的变量:$title, $description, $content, $url, $to。例如 {\"text\": \"$content\"}",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
RegisterPushChannelDefinition(PushDefinition{
|
||||
Type: ChannelLark,
|
||||
Name: "飞书群机器人",
|
||||
Description: "配置飞书群自定义机器人的 Webhook 接口投递。",
|
||||
Fields: []PushField{
|
||||
{
|
||||
Key: KeyURL,
|
||||
Label: "Webhook 地址",
|
||||
Type: TypeText,
|
||||
Required: true,
|
||||
Placeholder: "https://open.feishu.cn/open-apis/bot/v2/hook/YOUR_TOKEN",
|
||||
Description: "从飞书群机器人设置中复制的 Webhook URL",
|
||||
},
|
||||
{
|
||||
Key: KeyToken,
|
||||
Label: "签名校验密钥 (Secret) (可选)",
|
||||
Type: TypeText,
|
||||
Required: false,
|
||||
Placeholder: "可选,若机器人启用了安全设置中的签名校验,请在此输入",
|
||||
Description: "飞书群机器人安全设置中的签名校验 Key",
|
||||
},
|
||||
{
|
||||
Key: KeyOther,
|
||||
Label: "自定义卡片 JSON 模版 (可选)",
|
||||
Type: TypeTextarea,
|
||||
Required: false,
|
||||
Placeholder: "可选,留空则默认使用系统内置的精美互动卡片",
|
||||
Description: "若填写,必须是合法的飞书卡片 JSON 格式",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
RegisterPushChannelDefinition(PushDefinition{
|
||||
Type: ChannelTelegram,
|
||||
Name: "Telegram 机器人",
|
||||
Description: "配置 Telegram 机器人推送消息。",
|
||||
Fields: []PushField{
|
||||
{
|
||||
Key: KeyURL,
|
||||
Label: "API 基础地址 (可选)",
|
||||
Type: TypeText,
|
||||
Required: false,
|
||||
Placeholder: "https://api.telegram.org",
|
||||
Description: "接口请求的 HTTPS 基础地址,留空默认为 https://api.telegram.org",
|
||||
},
|
||||
{
|
||||
Key: KeyToken,
|
||||
Label: "机器人 Token (Bot Token)",
|
||||
Type: TypePassword,
|
||||
Required: true,
|
||||
Placeholder: "在此输入 Telegram 机器人的 Bot Token",
|
||||
Description: "通过 BotFather 申请到的机器人 Access Token",
|
||||
},
|
||||
{
|
||||
Key: KeyOther,
|
||||
Label: "默认会话 ID (Chat ID) (可选)",
|
||||
Type: TypeText,
|
||||
Required: false,
|
||||
Placeholder: "例如 -100123456789 或 @channel_name",
|
||||
Description: "默认的消息接收 Chat ID。如果通知事件中未配置 targets,将推送到此 ID",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
RegisterPushChannelDefinition(PushDefinition{
|
||||
Type: ChannelEmail,
|
||||
Name: "邮件推送通道",
|
||||
Description: "邮件推送通道直接使用系统全局 SMTP 设置进行发送,无需在此填写服务器配置。",
|
||||
Fields: []PushField{},
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_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,384 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package message_gateway provides the Bot gateway, multi-channel notification dispatching, and asynchronous push worker domain plugin for Cordis.
|
||||
package message_gateway
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/core/extpoints"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/message_gateway/handler"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_gateway/repository"
|
||||
"Wavelet/plugins/domain/message_gateway/service"
|
||||
"context"
|
||||
"embed"
|
||||
"reflect"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
//go:embed migrations/*/*.sql
|
||||
var mgMigrations embed.FS
|
||||
|
||||
// Option configures the message_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 message_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 message_gateway domain plugin.
|
||||
func (p *Plugin) Name() string {
|
||||
return "message_gateway"
|
||||
}
|
||||
|
||||
// Inject declares required dependencies for the message_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: "message_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 message_gateway plugin.
|
||||
func (p *Plugin) DeclareConfig() []core.ConfigBinding {
|
||||
return []core.ConfigBinding{
|
||||
{Prefix: "app", Target: &mgAppConfig{}},
|
||||
}
|
||||
}
|
||||
|
||||
// Apply registers message_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)
|
||||
}
|
||||
// 0. Bind DBService, CacheService, TaskService, UserService
|
||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||
repository.SetDBService(db)
|
||||
} else {
|
||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||
repository.SetDBService(db)
|
||||
})
|
||||
}
|
||||
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
|
||||
repository.SetCacheService(cache)
|
||||
service.SetCacheService(cache)
|
||||
} else {
|
||||
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
|
||||
repository.SetCacheService(cache)
|
||||
service.SetCacheService(cache)
|
||||
})
|
||||
}
|
||||
if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil {
|
||||
service.SetTaskService(taskSvc)
|
||||
} else {
|
||||
core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) {
|
||||
service.SetTaskService(taskSvc)
|
||||
})
|
||||
}
|
||||
if uSvc, err := core.Inject[contracts.UserService](ctx); err == nil && uSvc != nil {
|
||||
service.SetUserService(uSvc)
|
||||
} else {
|
||||
core.When[contracts.UserService](ctx, func(uSvc contracts.UserService) {
|
||||
service.SetUserService(uSvc)
|
||||
})
|
||||
}
|
||||
ctx.OnDispose(func() error {
|
||||
repository.SetDBService(nil)
|
||||
repository.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("message_gateway", mgMigrations)
|
||||
|
||||
// 2. Register User HTTP Routes
|
||||
handler.RegisterUserRoutes(ctx.Router().Group("/api/v1"), loginMW)
|
||||
|
||||
// 3. Register Admin Message Gateway HTTP Routes
|
||||
handler.RegisterAdminRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW)
|
||||
|
||||
// 4. Register Admin Push HTTP Routes
|
||||
handler.RegisterAdminPushRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW)
|
||||
|
||||
const defaultTaskRetry = 3
|
||||
pushHandler := &service.PushHandler{}
|
||||
|
||||
// 5. Register background tasks
|
||||
ctx.Task().Register("message_gateway:push_notification", 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("message_gateway:dispatch_bot_msg", func(_ context.Context, _ []byte) error {
|
||||
return nil
|
||||
},
|
||||
extpoints.WithTaskType("dispatch_bot_msg"),
|
||||
extpoints.WithTaskName("分发 Bot 消息"),
|
||||
extpoints.WithTaskDescription("异步处理与分发 Bot 下行消息"),
|
||||
extpoints.WithTaskCategory("messaging"),
|
||||
extpoints.WithTaskQueue("default"),
|
||||
)
|
||||
|
||||
ctx.Task().Register("message_gateway:cleanup_pairing_codes", func(c context.Context, _ []byte) error {
|
||||
return repository.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 * * * *", "message_gateway:cleanup_pairing_codes", map[string]any{"action": "cleanup"})
|
||||
|
||||
// 7. Register EventBus listeners for decoupled push triggers
|
||||
ctx.Events().On("notification:push", func(c context.Context, e model.PushNotificationEvent) error {
|
||||
meta := model.EventMetadata{
|
||||
Key: "eventbus:" + e.Channel,
|
||||
Name: e.Title,
|
||||
DefaultTemplate: model.NotificationMessage{
|
||||
Title: e.Title,
|
||||
Content: e.Content,
|
||||
Level: model.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
|
||||
service.RegisterCustomEvents()
|
||||
|
||||
// 10. Register Settings Schemas
|
||||
ctx.Settings().Register(extpoints.SettingSchema{
|
||||
Key: "message_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: "message_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 model.MessageChannel.
|
||||
type MessageChannel = model.MessageChannel
|
||||
|
||||
// MessageBinding is an alias for model.MessageBinding.
|
||||
type MessageBinding = model.MessageBinding
|
||||
|
||||
// MessagePairingCode is an alias for model.MessagePairingCode.
|
||||
type MessagePairingCode = model.MessagePairingCode
|
||||
|
||||
// PushChannel is an alias for model.PushChannel.
|
||||
type PushChannel = model.PushChannel
|
||||
|
||||
// PushEvent is an alias for model.PushEvent.
|
||||
type PushEvent = model.PushEvent
|
||||
|
||||
// PushHistory is an alias for model.PushHistory.
|
||||
type PushHistory = model.PushHistory
|
||||
|
||||
// PushNotificationEvent is an alias for model.PushNotificationEvent.
|
||||
type PushNotificationEvent = model.PushNotificationEvent
|
||||
|
||||
// ChannelConfig is an alias for model.ChannelConfig.
|
||||
type ChannelConfig = model.ChannelConfig
|
||||
|
||||
// Capability is an alias for model.Capability.
|
||||
type Capability = model.Capability
|
||||
|
||||
// Recipient is an alias for model.Recipient.
|
||||
type Recipient = model.Recipient
|
||||
|
||||
// Attachment is an alias for model.Attachment.
|
||||
type Attachment = model.Attachment
|
||||
|
||||
// InboundMessage is an alias for model.InboundMessage.
|
||||
type InboundMessage = model.InboundMessage
|
||||
|
||||
// OutboundMessage is an alias for model.OutboundMessage.
|
||||
type OutboundMessage = model.OutboundMessage
|
||||
|
||||
// BindingDTO is an alias for model.BindingDTO.
|
||||
type BindingDTO = model.BindingDTO
|
||||
|
||||
// PublicChannelDTO is an alias for model.PublicChannelDTO.
|
||||
type PublicChannelDTO = model.PublicChannelDTO
|
||||
|
||||
// Definition is an alias for model.Definition.
|
||||
type Definition = model.Definition
|
||||
|
||||
// ChannelDTO is an alias for model.ChannelDTO.
|
||||
type ChannelDTO = model.ChannelDTO
|
||||
|
||||
// CreateChannelRequest is an alias for model.CreateChannelRequest.
|
||||
type CreateChannelRequest = model.CreateChannelRequest
|
||||
|
||||
// UpdateChannelRequest is an alias for model.UpdateChannelRequest.
|
||||
type UpdateChannelRequest = model.UpdateChannelRequest
|
||||
|
||||
// PushDefinition is an alias for model.PushDefinition.
|
||||
type PushDefinition = model.PushDefinition
|
||||
|
||||
// PushField is an alias for model.PushField.
|
||||
type PushField = model.PushField
|
||||
|
||||
// NotificationMessage is an alias for model.NotificationMessage.
|
||||
type NotificationMessage = model.NotificationMessage
|
||||
|
||||
// EventMetadata is an alias for model.EventMetadata.
|
||||
type EventMetadata = model.EventMetadata
|
||||
|
||||
// SendPayload is an alias for model.SendPayload.
|
||||
type SendPayload = model.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 = repository.SetDBServiceForTest
|
||||
UpsertPairingCode = repository.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,61 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_gateway_test
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/plugins/domain/message_gateway"
|
||||
"context"
|
||||
"io/fs"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestMessageGatewayPluginUnit(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
p := message_gateway.New()
|
||||
assert.Equal(t, "message_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("message_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("message_gateway:push_notification")
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, 3, task.Retry)
|
||||
|
||||
// Verify schedules
|
||||
sched, ok := ctx.Schedules().Get("message_gateway:cleanup_pairing_codes")
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "*/10 * * * *", sched.Spec)
|
||||
|
||||
// Verify settings
|
||||
setting, ok := ctx.Settings().Get("message_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, message_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,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,104 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/util"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/smtp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("email", &EmailPusher{})
|
||||
}
|
||||
|
||||
// EmailPusher 极简 SMTP 邮件推送实现 (静态、解耦)
|
||||
type EmailPusher struct{}
|
||||
|
||||
// sanitizeEmailHeader removes CR/LF bytes so untrusted values cannot inject
|
||||
// additional email headers (email header injection).
|
||||
func sanitizeEmailHeader(v string) string {
|
||||
v = strings.ReplaceAll(v, "\r", "")
|
||||
v = strings.ReplaceAll(v, "\n", "")
|
||||
return v
|
||||
}
|
||||
|
||||
// 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>", "")
|
||||
|
||||
// 邮件头和体
|
||||
from := cfg.Key
|
||||
to := target
|
||||
|
||||
// 如果 ext 中指定了 from_name,我们在 From 头部包含它
|
||||
fromName := "System Notification"
|
||||
if ext != nil {
|
||||
if fn, ok := ext["from_name"].(string); ok && fn != "" {
|
||||
fromName = fn
|
||||
}
|
||||
}
|
||||
|
||||
subjectHeader := fmt.Sprintf("Subject: %s\r\n", sanitizeEmailHeader(title))
|
||||
fromHeader := fmt.Sprintf("From: %s <%s>\r\n", sanitizeEmailHeader(fromName), sanitizeEmailHeader(from))
|
||||
toHeader := fmt.Sprintf("To: %s\r\n", sanitizeEmailHeader(to))
|
||||
mimeHeader := "MIME-version: 1.0;\r\nContent-Type: text/html; charset=\"UTF-8\";\r\n\r\n"
|
||||
|
||||
// 拼装完整的邮件报文
|
||||
// 简单的 HTML 正文渲染
|
||||
htmlBody := fmt.Sprintf(`<html><body><h2>%s</h2><div>%s</div></body></html>`, title, content)
|
||||
msg := []byte(fromHeader + toHeader + subjectHeader + mimeHeader + htmlBody + "\r\n")
|
||||
|
||||
// 解析 Host 和 Port
|
||||
host, port, err := net.SplitHostPort(cfg.URL)
|
||||
if err != nil {
|
||||
host = cfg.URL
|
||||
port = "25" // 默认 SMTP 端口
|
||||
}
|
||||
|
||||
auth := smtp.PlainAuth("", cfg.Key, cfg.Secret, host)
|
||||
|
||||
// 异步超时处理
|
||||
errChan := make(chan error, 1)
|
||||
util.Go(func() {
|
||||
errChan <- smtp.SendMail(host+":"+port, auth, from, []string{to}, msg)
|
||||
})
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return "", ctx.Err()
|
||||
case err := <-errChan:
|
||||
if 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,26 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestSanitizeEmailHeader(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{"plain", "System Notification", "System Notification"},
|
||||
{"crlf stripped", "alert\r\nBcc: attacker@example.com", "alertBcc: attacker@example.com"},
|
||||
{"cr stripped", "a\rb", "ab"},
|
||||
{"lf stripped", "a\nb", "ab"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := sanitizeEmailHeader(tt.input); got != tt.want {
|
||||
t.Errorf("sanitizeEmailHeader(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,256 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/httppool"
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("lark", &LarkPusher{})
|
||||
}
|
||||
|
||||
const (
|
||||
msgTypeInteractive = "interactive"
|
||||
)
|
||||
|
||||
// LarkPusher 飞书 Webhook 机器人推送实现
|
||||
type LarkPusher struct{}
|
||||
|
||||
type larkTextContent struct {
|
||||
Text string `json:"text"`
|
||||
}
|
||||
|
||||
type larkCardHeaderTitle struct {
|
||||
Content string `json:"content"`
|
||||
Tag string `json:"tag"`
|
||||
}
|
||||
|
||||
type larkCardHeader struct {
|
||||
Template string `json:"template"` // "blue", "orange", "red" etc.
|
||||
Title larkCardHeaderTitle `json:"title"`
|
||||
}
|
||||
|
||||
type larkCardElementText struct {
|
||||
Content string `json:"content"`
|
||||
Tag string `json:"tag"` // "lark_md"
|
||||
}
|
||||
|
||||
type larkCardElement struct {
|
||||
Tag string `json:"tag"` // "div"
|
||||
Text larkCardElementText `json:"text"`
|
||||
}
|
||||
|
||||
type larkCardContent struct {
|
||||
Header larkCardHeader `json:"header"`
|
||||
Elements []larkCardElement `json:"elements"`
|
||||
}
|
||||
|
||||
type larkMessageRequest struct {
|
||||
MessageType string `json:"msg_type"`
|
||||
Timestamp string `json:"timestamp,omitempty"`
|
||||
Sign string `json:"sign,omitempty"`
|
||||
Content larkTextContent `json:"content,omitempty"`
|
||||
Card *larkCardContent `json:"card,omitempty"`
|
||||
}
|
||||
|
||||
type larkMessageResponse struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
|
||||
// Send 执行飞书消息发送
|
||||
//
|
||||
//nolint:nestif,cyclop
|
||||
func (p *LarkPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, template string, _ map[string]any) (string, error) {
|
||||
if cfg.URL == "" {
|
||||
return "", errors.New("lark: URL is required")
|
||||
}
|
||||
|
||||
var req larkMessageRequest
|
||||
|
||||
// 1. 如果有自定义模板,我们尝试进行解析
|
||||
if template != "" {
|
||||
rendered := ParseTemplate(template, body)
|
||||
|
||||
// 尝试解析原生的 Lark Card
|
||||
var customCard larkCardContent
|
||||
var rawMap map[string]any
|
||||
_ = json.Unmarshal([]byte(rendered), &rawMap)
|
||||
|
||||
if rawMap != nil && rawMap["elements"] != nil {
|
||||
// 如果包含 elements 字段,说明是用户定制的原生飞书卡片 JSON
|
||||
if err := json.Unmarshal([]byte(rendered), &customCard); err == nil {
|
||||
req.MessageType = msgTypeInteractive
|
||||
req.Card = &customCard
|
||||
} else {
|
||||
req.MessageType = "text"
|
||||
req.Content.Text = rendered
|
||||
}
|
||||
} else {
|
||||
// 说明配置的是系统统一通知消息 of JSON 模板:{"title": "...", "content": "...", "level": "..."}
|
||||
type larkNotificationMessage struct {
|
||||
Title string `json:"title"`
|
||||
Content string `json:"content"`
|
||||
Level string `json:"level"`
|
||||
}
|
||||
var msg larkNotificationMessage
|
||||
if err := json.Unmarshal([]byte(rendered), &msg); err == nil && (msg.Title != "" || msg.Content != "") {
|
||||
title := msg.Title
|
||||
if title == "" {
|
||||
title = defaultTitle
|
||||
}
|
||||
content := msg.Content
|
||||
level := strings.ToUpper(msg.Level)
|
||||
if level == "" {
|
||||
level = levelInfo
|
||||
}
|
||||
|
||||
headerColor := "blue"
|
||||
switch level {
|
||||
case "IMPORTANT":
|
||||
headerColor = "orange"
|
||||
case "CRITICAL":
|
||||
headerColor = "red"
|
||||
}
|
||||
|
||||
req.MessageType = msgTypeInteractive
|
||||
req.Card = &larkCardContent{
|
||||
Header: larkCardHeader{
|
||||
Template: headerColor,
|
||||
Title: larkCardHeaderTitle{
|
||||
Content: title,
|
||||
Tag: "plain_text",
|
||||
},
|
||||
},
|
||||
Elements: []larkCardElement{
|
||||
{
|
||||
Tag: "div",
|
||||
Text: larkCardElementText{
|
||||
Content: content,
|
||||
Tag: "lark_md",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
} else {
|
||||
// 兜底:如果无法按 JSON 解析出结构化字段,当做普通文本发送
|
||||
req.MessageType = "text"
|
||||
req.Content.Text = rendered
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// 2. 如果无模板,默认生成一个精美的飞书互动卡片
|
||||
title := bodyTitle(body)
|
||||
content := bodyContent(body, "**%s**: %v", "\n")
|
||||
level := bodyLevel(body)
|
||||
|
||||
// 根据级别确定飞书卡片头部的背景色模板
|
||||
headerColor := "blue"
|
||||
switch level {
|
||||
case "IMPORTANT":
|
||||
headerColor = "orange"
|
||||
case "CRITICAL":
|
||||
headerColor = "red"
|
||||
}
|
||||
|
||||
req.MessageType = msgTypeInteractive
|
||||
req.Card = &larkCardContent{
|
||||
Header: larkCardHeader{
|
||||
Template: headerColor,
|
||||
Title: larkCardHeaderTitle{
|
||||
Content: title,
|
||||
Tag: "plain_text",
|
||||
},
|
||||
},
|
||||
Elements: []larkCardElement{
|
||||
{
|
||||
Tag: "div",
|
||||
Text: larkCardElementText{
|
||||
Content: content,
|
||||
Tag: "lark_md",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 计算签名 (如果配置了 secret)
|
||||
if cfg.Secret != "" {
|
||||
timestamp := time.Now().Unix()
|
||||
sign, err := larkSign(cfg.Secret, timestamp)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lark: sign failed: %w", err)
|
||||
}
|
||||
req.Timestamp = strconv.FormatInt(timestamp, 10)
|
||||
req.Sign = sign
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lark: marshal request failed: %w", err)
|
||||
}
|
||||
|
||||
// 4. 发送 POST 请求
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.URL, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lark: create http request failed: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
client := httppool.NewClient(defaultHTTPClientTimeout)
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lark: http request failed: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("lark: http status %s", resp.Status)
|
||||
}
|
||||
|
||||
var res larkMessageResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&res); err != nil {
|
||||
return "", fmt.Errorf("lark: decode response failed: %w", err)
|
||||
}
|
||||
|
||||
if res.Code != 0 {
|
||||
return "", fmt.Errorf("lark: send message failed, code %d: %s", res.Code, res.Msg)
|
||||
}
|
||||
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验飞书配置
|
||||
func (p *LarkPusher) 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
|
||||
}
|
||||
|
||||
func larkSign(secret string, timestamp int64) (string, error) {
|
||||
stringToSign := fmt.Sprintf("%v", timestamp) + "\n" + secret
|
||||
h := hmac.New(sha256.New, []byte(stringToSign))
|
||||
_, err := h.Write(nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.StdEncoding.EncodeToString(h.Sum(nil)), nil
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
// 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
|
||||
)
|
||||
|
||||
// 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 用户名
|
||||
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,140 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/httppool"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("telegram", &TelegramPusher{})
|
||||
}
|
||||
|
||||
// TelegramPusher Telegram 机器人推送实现
|
||||
type TelegramPusher struct{}
|
||||
|
||||
type telegramMessageRequest struct {
|
||||
ChatID string `json:"chat_id"`
|
||||
Text string `json:"text"`
|
||||
ParseMode string `json:"parse_mode,omitempty"`
|
||||
}
|
||||
|
||||
type telegramErrorResponse struct {
|
||||
Ok bool `json:"ok"`
|
||||
ErrorCode int `json:"error_code"`
|
||||
Description string `json:"description"`
|
||||
}
|
||||
|
||||
// Send 执行 Telegram 消息发送
|
||||
func (p *TelegramPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, template string, _ map[string]any) (string, error) {
|
||||
if cfg.Secret == "" {
|
||||
return "", errors.New("telegram: Bot Token (Secret) is required")
|
||||
}
|
||||
|
||||
chatID := target
|
||||
if chatID == "" {
|
||||
chatID = cfg.Key // Use default chat ID (Key) if target is blank
|
||||
}
|
||||
if chatID == "" {
|
||||
return "", errors.New("telegram: chat_id (target or default Key) is required")
|
||||
}
|
||||
|
||||
baseURL := cfg.URL
|
||||
if baseURL == "" {
|
||||
baseURL = "https://api.telegram.org"
|
||||
}
|
||||
baseURL = strings.TrimSuffix(baseURL, "/")
|
||||
|
||||
title := bodyTitle(body)
|
||||
content := bodyContent(body, "<b>%s</b>: %v", "\n")
|
||||
level := bodyLevel(body)
|
||||
|
||||
var text string
|
||||
if template != "" {
|
||||
text = ParseTemplate(template, body)
|
||||
} else {
|
||||
text = fmt.Sprintf("<b>[%s] %s</b>\n\n%s", escapeHTML(level), escapeHTML(title), escapeHTML(content))
|
||||
}
|
||||
|
||||
// Try sending with HTML parse mode
|
||||
err := p.sendMessage(ctx, baseURL, cfg.Secret, chatID, text, "HTML")
|
||||
if err != nil {
|
||||
// Fallback: send as plain text without parse mode
|
||||
plainText := text
|
||||
if template == "" {
|
||||
plainText = fmt.Sprintf("[%s] %s\n\n%s", level, title, content)
|
||||
}
|
||||
fallbackErr := p.sendMessage(ctx, baseURL, cfg.Secret, chatID, plainText, "")
|
||||
if fallbackErr != nil {
|
||||
return "", fmt.Errorf("telegram: send message failed (fallback also failed): %w (original HTML error: %w)", fallbackErr, err)
|
||||
}
|
||||
}
|
||||
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验 Telegram 配置
|
||||
func (p *TelegramPusher) ValidateConfig(cfg Config) error {
|
||||
if cfg.Secret == "" {
|
||||
return errors.New("bot Token (Secret) is required")
|
||||
}
|
||||
if cfg.URL != "" {
|
||||
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
|
||||
return errors.New("API base URL must start with http:// or https://")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *TelegramPusher) sendMessage(ctx context.Context, baseURL, token, chatID, text, parseMode string) error {
|
||||
apiURL := fmt.Sprintf("%s/bot%s/sendMessage", baseURL, token)
|
||||
|
||||
reqPayload := telegramMessageRequest{
|
||||
ChatID: chatID,
|
||||
Text: text,
|
||||
ParseMode: parseMode,
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(reqPayload)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal request failed: %w", err)
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return fmt.Errorf("create http request failed: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
client := httppool.NewClient(defaultHTTPClientTimeout)
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil {
|
||||
return fmt.Errorf("http request failed: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
var errRes telegramErrorResponse
|
||||
if decodeErr := json.NewDecoder(resp.Body).Decode(&errRes); decodeErr == nil {
|
||||
return fmt.Errorf("http status %d: %s", resp.StatusCode, errRes.Description)
|
||||
}
|
||||
return fmt.Errorf("http status %s", resp.Status)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func escapeHTML(s string) string {
|
||||
s = strings.ReplaceAll(s, "&", "&")
|
||||
s = strings.ReplaceAll(s, "<", "<")
|
||||
s = strings.ReplaceAll(s, ">", ">")
|
||||
return s
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestTelegramPusher_Send(t *testing.T) {
|
||||
t.Run("successful send with HTML parse mode", func(t *testing.T) {
|
||||
var receivedReq telegramMessageRequest
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
assert.Equal(t, "/botmy-token/sendMessage", r.URL.Path)
|
||||
assert.Equal(t, http.MethodPost, r.Method)
|
||||
assert.Equal(t, "application/json", r.Header.Get("Content-Type"))
|
||||
|
||||
err := json.NewDecoder(r.Body).Decode(&receivedReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(`{"ok": true}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
pusher := &TelegramPusher{}
|
||||
cfg := Config{
|
||||
Channel: "telegram",
|
||||
URL: server.URL,
|
||||
Secret: "my-token",
|
||||
}
|
||||
body := map[string]any{
|
||||
"title": "Alert",
|
||||
"content": "Host down",
|
||||
"level": "CRITICAL",
|
||||
}
|
||||
_, err := pusher.Send(context.Background(), cfg, "123456", body, "", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "123456", receivedReq.ChatID)
|
||||
assert.Contains(t, receivedReq.Text, "[CRITICAL] Alert")
|
||||
assert.Contains(t, receivedReq.Text, "Host down")
|
||||
assert.Equal(t, "HTML", receivedReq.ParseMode)
|
||||
})
|
||||
|
||||
t.Run("fallback to plain text on HTML error", func(t *testing.T) {
|
||||
var requests []*telegramMessageRequest
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var req telegramMessageRequest
|
||||
err := json.NewDecoder(r.Body).Decode(&req)
|
||||
require.NoError(t, err)
|
||||
requests = append(requests, &req)
|
||||
|
||||
if len(requests) == 1 {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
_, _ = w.Write([]byte(`{"ok": false, "error_code": 400, "description": "Bad Request: can't parse entities"}`))
|
||||
} else {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(`{"ok": true}`))
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
pusher := &TelegramPusher{}
|
||||
cfg := Config{
|
||||
Channel: "telegram",
|
||||
URL: server.URL,
|
||||
Secret: "my-token",
|
||||
}
|
||||
body := map[string]any{
|
||||
"title": "Alert & Info",
|
||||
"content": "A < B comparison",
|
||||
"level": "INFO",
|
||||
}
|
||||
_, err := pusher.Send(context.Background(), cfg, "123456", body, "", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Len(t, requests, 2)
|
||||
assert.Equal(t, "HTML", requests[0].ParseMode)
|
||||
assert.Equal(t, "", requests[1].ParseMode)
|
||||
assert.Contains(t, requests[1].Text, "[INFO] Alert & Info")
|
||||
assert.Contains(t, requests[1].Text, "A < B comparison")
|
||||
})
|
||||
|
||||
t.Run("validation error", func(t *testing.T) {
|
||||
pusher := &TelegramPusher{}
|
||||
cfg := Config{
|
||||
Channel: "telegram",
|
||||
URL: "https://api.telegram.org",
|
||||
}
|
||||
err := pusher.ValidateConfig(cfg)
|
||||
assert.Error(t, err)
|
||||
|
||||
cfg = Config{
|
||||
Channel: "telegram",
|
||||
URL: "ftp://api.telegram.org",
|
||||
Secret: "token",
|
||||
}
|
||||
err = pusher.ValidateConfig(cfg)
|
||||
assert.Error(t, err)
|
||||
|
||||
cfg = Config{
|
||||
Channel: "telegram",
|
||||
Secret: "token",
|
||||
}
|
||||
err = pusher.ValidateConfig(cfg)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,111 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"maps"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ParseTemplate parses template strings by replacing {{placeholder}} structures with values from body.
|
||||
// It is a single-pass parser designed for high performance and low allocations.
|
||||
func ParseTemplate(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
|
||||
}
|
||||
// Write prefix
|
||||
buf.WriteString(template[i : i+pos])
|
||||
i += pos + 2 // skip "{{"
|
||||
|
||||
endPos := strings.Index(template[i:], "}}")
|
||||
if endPos == -1 {
|
||||
// Unbalanced "{{"
|
||||
buf.WriteString("{{")
|
||||
buf.WriteString(template[i:])
|
||||
break
|
||||
}
|
||||
key := template[i : i+endPos]
|
||||
if val, ok := body[key]; ok {
|
||||
buf.WriteString(formatValue(val))
|
||||
} else {
|
||||
// Keep the placeholder if key not found
|
||||
buf.WriteString("{{")
|
||||
buf.WriteString(key)
|
||||
buf.WriteString("}}")
|
||||
}
|
||||
i += endPos + 2 // skip "}}"
|
||||
}
|
||||
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:
|
||||
// If it's a map, slice, or struct, marshal it to JSON.
|
||||
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)
|
||||
}
|
||||
|
||||
// bodyLevel returns the upper-cased notification level, falling back to INFO.
|
||||
func bodyLevel(body map[string]any) string {
|
||||
if l, ok := body["level"].(string); ok && l != "" {
|
||||
return strings.ToUpper(l)
|
||||
}
|
||||
return levelInfo
|
||||
}
|
||||
@@ -0,0 +1,96 @@
|
||||
// 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"}`,
|
||||
},
|
||||
}
|
||||
|
||||
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 message_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,364 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"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) ([]model.PushChannel, error) {
|
||||
var channels []model.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) (model.PushChannel, error) {
|
||||
var channel model.PushChannel
|
||||
if err := GetDB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
|
||||
return model.PushChannel{}, mapNotFound(err)
|
||||
}
|
||||
return channel, nil
|
||||
}
|
||||
|
||||
// GetPushChannelByNameRecord loads a push channel by its unique name.
|
||||
func GetPushChannelByNameRecord(ctx context.Context, name string) (*model.PushChannel, error) {
|
||||
var channel model.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(&model.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 *model.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 *model.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 *model.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) (*model.PushChannel, error) {
|
||||
channel, err := getCachedOrQuery(ctx, "push:channel:active:"+name, activePushChannelCacheTTL, func(db *gorm.DB, dest *model.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) ([]model.PushEvent, error) {
|
||||
var events []model.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) (model.PushEvent, error) {
|
||||
var event model.PushEvent
|
||||
if err := GetDB(ctx).First(&event, id).Error; err != nil {
|
||||
return model.PushEvent{}, mapNotFound(err)
|
||||
}
|
||||
return event, nil
|
||||
}
|
||||
|
||||
// GetPushEventByKeyRecord loads a push event by event key.
|
||||
func GetPushEventByKeyRecord(ctx context.Context, key string) (model.PushEvent, error) {
|
||||
var event model.PushEvent
|
||||
if err := GetDB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
|
||||
return model.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(&model.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 *model.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 *model.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 *model.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 *model.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) ([]model.PushEvent, error) {
|
||||
var events []model.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) (*model.PushEvent, error) {
|
||||
event, err := getCachedOrQuery(ctx, "push:event:active:"+key, activePushEventCacheTTL, func(db *gorm.DB, dest *model.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 model.PushHistoryListFilter) (int64, []model.PushHistory, error) {
|
||||
query := GetDB(ctx).Model(&model.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 []model.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 *model.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(&model.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.
|
||||
//
|
||||
// A key that is simply absent leaves its field empty, which is how an unconfigured
|
||||
// mailer is represented. A read that fails is returned as an error, so callers
|
||||
// cannot mistake an unhealthy database for "no SMTP configured" and silently drop
|
||||
// the notification.
|
||||
func LoadSMTPConfigRecord(ctx context.Context) (model.SMTPConfig, error) {
|
||||
db := GetDB(ctx)
|
||||
if db == nil {
|
||||
return model.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 model.SMTPConfig{}, fmt.Errorf("read smtp system configs: %w", err)
|
||||
}
|
||||
|
||||
var cfg model.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.
|
||||
// The column name is concatenated into the WHERE clause, so anything not listed
|
||||
// here must never reach the database.
|
||||
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, errs.ErrUnsupportedUserLookupField
|
||||
}
|
||||
db := GetDB(ctx)
|
||||
if db == nil {
|
||||
return nil, errs.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, errs.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 repository_test
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/repository"
|
||||
"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)
|
||||
}
|
||||
|
||||
repository.SetDBServiceForTest(stubDBService{db: db})
|
||||
t.Cleanup(func() { repository.SetDBServiceForTest(nil) })
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
user, err := repository.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 := repository.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 := repository.FindUserByFieldRecord(ctx, tc.field, "seeded"); !errors.Is(err, errs.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)
|
||||
}
|
||||
}
|
||||
|
||||
repository.SetDBServiceForTest(stubDBService{db: db})
|
||||
t.Cleanup(func() { repository.SetDBServiceForTest(nil) })
|
||||
|
||||
cfg, err := repository.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)
|
||||
}
|
||||
|
||||
repository.SetDBServiceForTest(stubDBService{db: bare})
|
||||
t.Cleanup(func() { repository.SetDBServiceForTest(nil) })
|
||||
|
||||
if _, err := repository.LoadSMTPConfigRecord(context.Background()); err == nil {
|
||||
t.Fatal("LoadSMTPConfigRecord returned nil error although the config table cannot be read")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package repository provides data persistence for the message_gateway plugin.
|
||||
package repository
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"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
|
||||
// errs.ErrRecordNotFound so the service and handler layers stay free of gorm imports.
|
||||
func mapNotFound(err error) error {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return errs.ErrRecordNotFound
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// CreateMessageChannel inserts a channel row.
|
||||
func CreateMessageChannel(ctx context.Context, ch *model.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 *model.MessageChannel) error {
|
||||
return GetDB(ctx).Save(ch).Error
|
||||
}
|
||||
|
||||
// GetMessageChannel loads a channel by id.
|
||||
func GetMessageChannel(ctx context.Context, id uint64) (*model.MessageChannel, error) {
|
||||
var ch model.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) ([]model.MessageChannel, error) {
|
||||
var rows []model.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(&model.MessagePairingCode{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("channel_id = ?", id).Delete(&model.MessageBinding{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&model.MessageChannel{}, id).Error
|
||||
})
|
||||
}
|
||||
|
||||
// CreateMessageBinding inserts a binding.
|
||||
func CreateMessageBinding(ctx context.Context, b *model.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) (*model.MessageBinding, error) {
|
||||
var b model.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) ([]model.MessageBinding, error) {
|
||||
var rows []model.MessageBinding
|
||||
if err := GetDB(ctx).Where("user_id = ?", userID).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) (*model.MessageBinding, error) {
|
||||
var b model.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(&model.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) (*model.MessagePairingCode, error) {
|
||||
var existing model.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 := &model.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) (*model.MessagePairingCode, error) {
|
||||
var row model.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(&model.MessagePairingCode{}).Error
|
||||
}
|
||||
|
||||
// DeleteExpiredPairingCodes removes expired pairing rows.
|
||||
func DeleteExpiredPairingCodes(ctx context.Context) error {
|
||||
return GetDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&model.MessagePairingCode{}).Error
|
||||
}
|
||||
|
||||
// ListEnabledMessageChannels returns enabled channels.
|
||||
func ListEnabledMessageChannels(ctx context.Context) ([]model.MessageChannel, error) {
|
||||
var rows []model.MessageChannel
|
||||
if err := GetDB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
@@ -0,0 +1,310 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_gateway/repository"
|
||||
"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() []model.Definition {
|
||||
return []model.Definition{
|
||||
{
|
||||
Type: model.MessageChannelTypeTelegram,
|
||||
Fields: []model.Field{
|
||||
{Key: "token", Type: "password", Required: true},
|
||||
{Key: "api_base", Type: "text", Required: false},
|
||||
},
|
||||
},
|
||||
{
|
||||
Type: model.MessageChannelTypeQQ,
|
||||
Fields: []model.Field{
|
||||
{Key: "app_id", Type: "text", 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 model.CreateChannelRequest) (model.ChannelDTO, error) {
|
||||
name := strings.TrimSpace(req.Name)
|
||||
if name == "" {
|
||||
return model.ChannelDTO{}, errors.New(errs.ErrNameRequired)
|
||||
}
|
||||
channelType := strings.TrimSpace(req.Type)
|
||||
if channelType != model.MessageChannelTypeTelegram && channelType != model.MessageChannelTypeQQ {
|
||||
return model.ChannelDTO{}, errors.New(errs.ErrTypeInvalid)
|
||||
}
|
||||
creds := req.Credentials
|
||||
if creds == nil {
|
||||
creds = map[string]string{}
|
||||
}
|
||||
if err := ValidateCredentials(channelType, creds, false); err != nil {
|
||||
return model.ChannelDTO{}, err
|
||||
}
|
||||
cipher, err := EncryptCredentials(creds)
|
||||
if err != nil {
|
||||
return model.ChannelDTO{}, err
|
||||
}
|
||||
extra := req.Extra
|
||||
if extra == nil {
|
||||
extra = map[string]string{}
|
||||
}
|
||||
enabled := true
|
||||
if req.Enabled != nil {
|
||||
enabled = *req.Enabled
|
||||
}
|
||||
row := &model.MessageChannel{
|
||||
Name: name,
|
||||
Type: channelType,
|
||||
OwnerScope: model.MessageOwnerScopeSystem,
|
||||
Enabled: enabled,
|
||||
Credentials: cipher,
|
||||
Extra: EncodeExtra(extra),
|
||||
}
|
||||
if err := repository.CreateMessageChannel(ctx, row); err != nil {
|
||||
return model.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 model.UpdateChannelRequest) (model.ChannelDTO, error) {
|
||||
row, err := repository.GetMessageChannel(ctx, id)
|
||||
if err != nil {
|
||||
if errors.Is(err, errs.ErrRecordNotFound) {
|
||||
return model.ChannelDTO{}, errors.New(errs.ErrChannelNotFound)
|
||||
}
|
||||
return model.ChannelDTO{}, err
|
||||
}
|
||||
creds, err := DecryptCredentials(row.Credentials)
|
||||
if err != nil {
|
||||
return model.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 model.ChannelDTO{}, err
|
||||
}
|
||||
creds = merged
|
||||
}
|
||||
|
||||
cipher, err := EncryptCredentials(creds)
|
||||
if err != nil {
|
||||
return model.ChannelDTO{}, err
|
||||
}
|
||||
row.Credentials = cipher
|
||||
row.Extra = EncodeExtra(extra)
|
||||
if err := repository.UpdateMessageChannel(ctx, row); err != nil {
|
||||
return model.ChannelDTO{}, err
|
||||
}
|
||||
return ToDTO(row, creds, extra), nil
|
||||
}
|
||||
|
||||
// ListChannels returns every channel with secrets masked.
|
||||
func ListChannels(ctx context.Context) ([]model.ChannelDTO, error) {
|
||||
rows, err := repository.ListMessageChannels(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]model.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 := repository.GetMessageChannel(ctx, id); err != nil {
|
||||
if errors.Is(err, errs.ErrRecordNotFound) {
|
||||
return errors.New(errs.ErrChannelNotFound)
|
||||
}
|
||||
return err
|
||||
}
|
||||
return repository.DeleteMessageChannel(ctx, id)
|
||||
}
|
||||
|
||||
// ProbeChannel verifies the stored credentials against the upstream platform.
|
||||
func ProbeChannel(ctx context.Context, id uint64) error {
|
||||
row, err := repository.GetMessageChannel(ctx, id)
|
||||
if err != nil {
|
||||
if errors.Is(err, errs.ErrRecordNotFound) {
|
||||
return errors.New(errs.ErrChannelNotFound)
|
||||
}
|
||||
return err
|
||||
}
|
||||
creds, err := DecryptCredentials(row.Credentials)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
switch row.Type {
|
||||
case model.MessageChannelTypeTelegram:
|
||||
return ProbeTelegram(ctx, creds)
|
||||
case model.MessageChannelTypeQQ:
|
||||
return ProbeQQ(ctx, creds)
|
||||
default:
|
||||
return errors.New(errs.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(errs.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", errs.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", errs.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(errs.ErrMissingQQCredentials)
|
||||
}
|
||||
credentials := &token.QQBotCredentials{
|
||||
AppID: appID,
|
||||
AppSecret: secret,
|
||||
}
|
||||
tokSrc := token.NewQQBotTokenSource(credentials)
|
||||
tok, err := tokSrc.Token()
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s: %w", errs.ErrQQTokenFetchFailed, err)
|
||||
}
|
||||
if tok == nil || tok.AccessToken == "" {
|
||||
return errors.New(errs.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 model.MessageChannelTypeTelegram:
|
||||
tok := creds["token"]
|
||||
if strings.TrimSpace(tok) == "" && !isUpdate {
|
||||
return errors.New(errs.ErrTelegramTokenRequired)
|
||||
}
|
||||
if base, ok := creds["api_base"]; ok && strings.TrimSpace(base) != "" {
|
||||
if !strings.HasPrefix(base, "http://") && !strings.HasPrefix(base, "https://") {
|
||||
return errors.New(errs.ErrAPIBaseInvalid)
|
||||
}
|
||||
}
|
||||
case model.MessageChannelTypeQQ:
|
||||
appID := creds["app_id"]
|
||||
secret := creds["client_secret"]
|
||||
if (strings.TrimSpace(appID) == "" || strings.TrimSpace(secret) == "") && !isUpdate {
|
||||
return errors.New(errs.ErrQQCredentialsRequired)
|
||||
}
|
||||
default:
|
||||
return errors.New(errs.ErrTypeInvalid)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ToDTO projects a channel row onto the admin DTO with credentials masked.
|
||||
func ToDTO(row *model.MessageChannel, creds, extra map[string]string) model.ChannelDTO {
|
||||
return model.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:]
|
||||
}
|
||||
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 message_gateway.
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_gateway/repository"
|
||||
"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 model.InboundMessage) error
|
||||
|
||||
// Factory constructs a Channel from decrypted config.
|
||||
type Factory func(cfg model.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 model.Recipient, msg model.OutboundMessage) error
|
||||
Capabilities() model.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
|
||||
}
|
||||
|
||||
// CodeAlphabet excludes easily confused runes 0/O/1/I.
|
||||
const CodeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
|
||||
|
||||
// CodeLength is the raw pairing code size.
|
||||
const CodeLength = 8
|
||||
|
||||
// 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 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
|
||||
}
|
||||
|
||||
// GetTaskService returns the task service.
|
||||
func GetTaskService() contracts.TaskService {
|
||||
taskMu.RLock()
|
||||
defer taskMu.RUnlock()
|
||||
return taskSvc
|
||||
}
|
||||
|
||||
// GetUserService resolves the user service for the context.
|
||||
func GetUserService(ctx context.Context) contracts.UserService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.UserService](c); 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 model.BindRequest) (model.BindingDTO, error) {
|
||||
channelID, err := strconv.ParseUint(strings.TrimSpace(req.ChannelID), 10, 64)
|
||||
if err != nil || channelID == 0 {
|
||||
return model.BindingDTO{}, errs.ErrChannelIDRequired
|
||||
}
|
||||
code := NormalizeCode(req.Code)
|
||||
if code == "" {
|
||||
return model.BindingDTO{}, errs.ErrCodeInvalid
|
||||
}
|
||||
pairing, err := repository.GetPairingCode(ctx, code)
|
||||
if err != nil {
|
||||
if errors.Is(err, errs.ErrRecordNotFound) {
|
||||
return model.BindingDTO{}, errs.ErrCodeInvalid
|
||||
}
|
||||
return model.BindingDTO{}, err
|
||||
}
|
||||
if !pairing.ExpiresAt.After(time.Now()) {
|
||||
return model.BindingDTO{}, errs.ErrCodeInvalid
|
||||
}
|
||||
if pairing.ChannelID != channelID {
|
||||
return model.BindingDTO{}, errs.ErrChannelMismatch
|
||||
}
|
||||
ch, err := repository.GetMessageChannel(ctx, channelID)
|
||||
if err != nil {
|
||||
if errors.Is(err, errs.ErrRecordNotFound) {
|
||||
return model.BindingDTO{}, errs.ErrCodeInvalid
|
||||
}
|
||||
return model.BindingDTO{}, err
|
||||
}
|
||||
if !ch.Enabled {
|
||||
return model.BindingDTO{}, errs.ErrChannelDisabled
|
||||
}
|
||||
|
||||
existing, err := repository.GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID)
|
||||
if err != nil && !errors.Is(err, errs.ErrRecordNotFound) {
|
||||
return model.BindingDTO{}, err
|
||||
}
|
||||
if err == nil && existing != nil {
|
||||
if existing.UserID != userID {
|
||||
return model.BindingDTO{}, errs.ErrPlatformAlreadyBound
|
||||
}
|
||||
_ = repository.DeletePairingCode(ctx, pairing.Code)
|
||||
return ToBindingDTO(existing, ch), nil
|
||||
}
|
||||
|
||||
row := &model.MessageBinding{
|
||||
UserID: userID,
|
||||
ChannelID: channelID,
|
||||
PlatformUserID: pairing.PlatformUserID,
|
||||
}
|
||||
if err := repository.CreateMessageBinding(ctx, row); err != nil {
|
||||
return model.BindingDTO{}, err
|
||||
}
|
||||
if err := repository.DeletePairingCode(ctx, pairing.Code); err != nil {
|
||||
return model.BindingDTO{}, err
|
||||
}
|
||||
return ToBindingDTO(row, ch), nil
|
||||
}
|
||||
|
||||
// ListEnabledPublicChannels returns the channels a user may bind to.
|
||||
func ListEnabledPublicChannels(ctx context.Context) ([]model.PublicChannelDTO, error) {
|
||||
rows, err := repository.ListEnabledMessageChannels(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]model.PublicChannelDTO, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
out = append(out, model.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) ([]model.BindingDTO, error) {
|
||||
rows, err := repository.ListBindingsByUser(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]model.BindingDTO, 0, len(rows))
|
||||
for i := range rows {
|
||||
ch, err := repository.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 := repository.GetMessageBinding(ctx, bindingID)
|
||||
if err != nil {
|
||||
if errors.Is(err, errs.ErrRecordNotFound) {
|
||||
return errs.ErrBindingNotFound
|
||||
}
|
||||
return err
|
||||
}
|
||||
if row.UserID != userID {
|
||||
return errs.ErrBindingForbidden
|
||||
}
|
||||
return repository.DeleteMessageBinding(ctx, bindingID)
|
||||
}
|
||||
|
||||
// ToBindingDTO projects a binding row and its optional channel onto the user DTO.
|
||||
func ToBindingDTO(row *model.MessageBinding, ch *model.MessageChannel) model.BindingDTO {
|
||||
dto := model.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