mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 16:46:37 +08:00
refactor(msg_gateway): restructure and rename message_gateway aligned with custom_example
This commit is contained in:
@@ -0,0 +1,311 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"Wavelet/plugins/domain/msg_gateway/dao"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/tencent-connect/botgo/token"
|
||||
)
|
||||
|
||||
const defaultTelegramAPI = "https://api.telegram.org"
|
||||
|
||||
// ListDefinitions returns the admin form schema of every supported channel type.
|
||||
func ListDefinitions() []do.Definition {
|
||||
return []do.Definition{
|
||||
{
|
||||
Type: consts.MessageChannelTypeTelegram,
|
||||
Fields: []do.Field{
|
||||
{Key: "token", Type: consts.TypePassword, Required: true},
|
||||
{Key: "api_base", Type: consts.TypeText, Required: false},
|
||||
},
|
||||
},
|
||||
{
|
||||
Type: consts.MessageChannelTypeQQ,
|
||||
Fields: []do.Field{
|
||||
{Key: "app_id", Type: consts.TypeText, Required: true},
|
||||
{Key: "client_secret", Type: "password", Required: true},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// CreateChannel validates the admin payload and persists an encrypted channel.
|
||||
func CreateChannel(ctx context.Context, req do.CreateChannelRequest) (do.ChannelDTO, error) {
|
||||
name := strings.TrimSpace(req.Name)
|
||||
if name == "" {
|
||||
return do.ChannelDTO{}, errors.New(consts.ErrNameRequired)
|
||||
}
|
||||
channelType := strings.TrimSpace(req.Type)
|
||||
if channelType != consts.MessageChannelTypeTelegram && channelType != consts.MessageChannelTypeQQ {
|
||||
return do.ChannelDTO{}, errors.New(consts.ErrTypeInvalid)
|
||||
}
|
||||
creds := req.Credentials
|
||||
if creds == nil {
|
||||
creds = map[string]string{}
|
||||
}
|
||||
if err := ValidateCredentials(channelType, creds, false); err != nil {
|
||||
return do.ChannelDTO{}, err
|
||||
}
|
||||
cipher, err := EncryptCredentials(creds)
|
||||
if err != nil {
|
||||
return do.ChannelDTO{}, err
|
||||
}
|
||||
extra := req.Extra
|
||||
if extra == nil {
|
||||
extra = map[string]string{}
|
||||
}
|
||||
enabled := true
|
||||
if req.Enabled != nil {
|
||||
enabled = *req.Enabled
|
||||
}
|
||||
row := &entity.MessageChannel{
|
||||
Name: name,
|
||||
Type: channelType,
|
||||
OwnerScope: consts.MessageOwnerScopeSystem,
|
||||
Enabled: enabled,
|
||||
Credentials: cipher,
|
||||
Extra: EncodeExtra(extra),
|
||||
}
|
||||
if err := dao.CreateMessageChannel(ctx, row); err != nil {
|
||||
return do.ChannelDTO{}, err
|
||||
}
|
||||
return ToDTO(row, creds, extra), nil
|
||||
}
|
||||
|
||||
// UpdateChannel patches a channel; empty secrets keep the stored ciphertext.
|
||||
func UpdateChannel(ctx context.Context, id uint64, req do.UpdateChannelRequest) (do.ChannelDTO, error) {
|
||||
row, err := dao.GetMessageChannel(ctx, id)
|
||||
if err != nil {
|
||||
if errors.Is(err, consts.ErrRecordNotFound) {
|
||||
return do.ChannelDTO{}, errors.New(consts.ErrChannelNotFound)
|
||||
}
|
||||
return do.ChannelDTO{}, err
|
||||
}
|
||||
creds, err := DecryptCredentials(row.Credentials)
|
||||
if err != nil {
|
||||
return do.ChannelDTO{}, err
|
||||
}
|
||||
extra := ParseExtra(row.Extra)
|
||||
|
||||
if name := strings.TrimSpace(req.Name); name != "" {
|
||||
row.Name = name
|
||||
}
|
||||
if req.Enabled != nil {
|
||||
row.Enabled = *req.Enabled
|
||||
}
|
||||
if req.Extra != nil {
|
||||
extra = req.Extra
|
||||
}
|
||||
if len(req.Credentials) > 0 {
|
||||
merged := make(map[string]string, len(creds))
|
||||
for k, v := range creds {
|
||||
merged[k] = v
|
||||
}
|
||||
for k, v := range req.Credentials {
|
||||
if strings.TrimSpace(v) == "" {
|
||||
continue
|
||||
}
|
||||
merged[k] = v
|
||||
}
|
||||
if err := ValidateCredentials(row.Type, merged, true); err != nil {
|
||||
return do.ChannelDTO{}, err
|
||||
}
|
||||
creds = merged
|
||||
}
|
||||
|
||||
cipher, err := EncryptCredentials(creds)
|
||||
if err != nil {
|
||||
return do.ChannelDTO{}, err
|
||||
}
|
||||
row.Credentials = cipher
|
||||
row.Extra = EncodeExtra(extra)
|
||||
if err := dao.UpdateMessageChannel(ctx, row); err != nil {
|
||||
return do.ChannelDTO{}, err
|
||||
}
|
||||
return ToDTO(row, creds, extra), nil
|
||||
}
|
||||
|
||||
// ListChannels returns every channel with secrets masked.
|
||||
func ListChannels(ctx context.Context) ([]do.ChannelDTO, error) {
|
||||
rows, err := dao.ListMessageChannels(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]do.ChannelDTO, 0, len(rows))
|
||||
for i := range rows {
|
||||
creds, _ := DecryptCredentials(rows[i].Credentials)
|
||||
extra := ParseExtra(rows[i].Extra)
|
||||
out = append(out, ToDTO(&rows[i], creds, extra))
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// DeleteChannel removes a channel together with its bindings and pairing codes.
|
||||
func DeleteChannel(ctx context.Context, id uint64) error {
|
||||
if _, err := dao.GetMessageChannel(ctx, id); err != nil {
|
||||
if errors.Is(err, consts.ErrRecordNotFound) {
|
||||
return errors.New(consts.ErrChannelNotFound)
|
||||
}
|
||||
return err
|
||||
}
|
||||
return dao.DeleteMessageChannel(ctx, id)
|
||||
}
|
||||
|
||||
// ProbeChannel verifies the stored credentials against the upstream platform.
|
||||
func ProbeChannel(ctx context.Context, id uint64) error {
|
||||
row, err := dao.GetMessageChannel(ctx, id)
|
||||
if err != nil {
|
||||
if errors.Is(err, consts.ErrRecordNotFound) {
|
||||
return errors.New(consts.ErrChannelNotFound)
|
||||
}
|
||||
return err
|
||||
}
|
||||
creds, err := DecryptCredentials(row.Credentials)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
switch row.Type {
|
||||
case consts.MessageChannelTypeTelegram:
|
||||
return ProbeTelegram(ctx, creds)
|
||||
case consts.MessageChannelTypeQQ:
|
||||
return ProbeQQ(ctx, creds)
|
||||
default:
|
||||
return errors.New(consts.ErrTypeInvalid)
|
||||
}
|
||||
}
|
||||
|
||||
// ProbeTelegram calls getMe to confirm the bot token is usable.
|
||||
func ProbeTelegram(ctx context.Context, creds map[string]string) error {
|
||||
tok := creds["token"]
|
||||
if strings.TrimSpace(tok) == "" {
|
||||
return errors.New(consts.ErrMissingTelegramToken)
|
||||
}
|
||||
base := creds["api_base"]
|
||||
base = strings.TrimRight(strings.TrimSpace(base), "/")
|
||||
if base == "" {
|
||||
base = defaultTelegramAPI
|
||||
}
|
||||
url := fmt.Sprintf("%s/bot%s/getMe", base, tok)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
client := &http.Client{Timeout: 10 * time.Second}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("%s (%d): %s", consts.ErrTelegramGetMeFailed, resp.StatusCode, string(body))
|
||||
}
|
||||
var res struct {
|
||||
OK bool `json:"ok"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &res); err != nil {
|
||||
return err
|
||||
}
|
||||
if !res.OK {
|
||||
return fmt.Errorf("%s: %s", consts.ErrTelegramNotOK, string(body))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ProbeQQ exchanges the app credentials for an access token.
|
||||
func ProbeQQ(_ context.Context, creds map[string]string) error {
|
||||
appID := strings.TrimSpace(creds["app_id"])
|
||||
secret := strings.TrimSpace(creds["client_secret"])
|
||||
if appID == "" || secret == "" {
|
||||
return errors.New(consts.ErrMissingQQCredentials)
|
||||
}
|
||||
credentials := &token.QQBotCredentials{
|
||||
AppID: appID,
|
||||
AppSecret: secret,
|
||||
}
|
||||
tokSrc := token.NewQQBotTokenSource(credentials)
|
||||
tok, err := tokSrc.Token()
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s: %w", consts.ErrQQTokenFetchFailed, err)
|
||||
}
|
||||
if tok == nil || tok.AccessToken == "" {
|
||||
return errors.New(consts.ErrQQEmptyToken)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateCredentials checks the admin submitted credentials for a channel type.
|
||||
func ValidateCredentials(t string, creds map[string]string, isUpdate bool) error {
|
||||
switch t {
|
||||
case consts.MessageChannelTypeTelegram:
|
||||
tok := creds["token"]
|
||||
if strings.TrimSpace(tok) == "" && !isUpdate {
|
||||
return errors.New(consts.ErrTelegramTokenRequired)
|
||||
}
|
||||
if base, ok := creds["api_base"]; ok && strings.TrimSpace(base) != "" {
|
||||
if !strings.HasPrefix(base, "http://") && !strings.HasPrefix(base, "https://") {
|
||||
return errors.New(consts.ErrAPIBaseInvalid)
|
||||
}
|
||||
}
|
||||
case consts.MessageChannelTypeQQ:
|
||||
appID := creds["app_id"]
|
||||
secret := creds["client_secret"]
|
||||
if (strings.TrimSpace(appID) == "" || strings.TrimSpace(secret) == "") && !isUpdate {
|
||||
return errors.New(consts.ErrQQCredentialsRequired)
|
||||
}
|
||||
default:
|
||||
return errors.New(consts.ErrTypeInvalid)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ToDTO projects a channel row onto the admin DTO with credentials masked.
|
||||
func ToDTO(row *entity.MessageChannel, creds, extra map[string]string) do.ChannelDTO {
|
||||
return do.ChannelDTO{
|
||||
ID: row.ID,
|
||||
Name: row.Name,
|
||||
Type: row.Type,
|
||||
OwnerScope: row.OwnerScope,
|
||||
OwnerID: row.OwnerID,
|
||||
Enabled: row.Enabled,
|
||||
Credentials: MaskCredentials(row.Type, creds),
|
||||
Extra: extra,
|
||||
}
|
||||
}
|
||||
|
||||
// MaskCredentials hides secret bearing credential entries.
|
||||
func MaskCredentials(_ string, in map[string]string) map[string]string {
|
||||
out := make(map[string]string, len(in))
|
||||
for k, v := range in {
|
||||
if k == "token" || k == "client_secret" {
|
||||
out[k] = MaskSecret(v)
|
||||
} else {
|
||||
out[k] = v
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
const minMaskSecretLength = 8
|
||||
|
||||
// MaskSecret keeps only a short visible prefix and suffix of a secret.
|
||||
func MaskSecret(s string) string {
|
||||
s = strings.TrimSpace(s)
|
||||
if len(s) <= minMaskSecretLength {
|
||||
return "******"
|
||||
}
|
||||
return s[:4] + "..." + s[len(s)-4:]
|
||||
}
|
||||
@@ -0,0 +1,191 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"Wavelet/plugins/domain/msg_gateway/dao"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
// TaskDispatchBotMsg is the queue pattern for bot downlink dispatch.
|
||||
TaskDispatchBotMsg = consts.TaskDispatchBotMsg
|
||||
// TaskTypeDispatchBotMsg is the admin type identifier for bot downlink dispatch.
|
||||
TaskTypeDispatchBotMsg = consts.TaskTypeDispatchBotMsg
|
||||
|
||||
taskQueueDefault = "default"
|
||||
taskParamTypeString = "string"
|
||||
paramNameText = "text"
|
||||
)
|
||||
|
||||
// BotDispatchMeta describes the bot downlink dispatch task.
|
||||
var BotDispatchMeta = contracts.TaskMetaDTO{
|
||||
Type: TaskTypeDispatchBotMsg,
|
||||
AsynqTask: TaskDispatchBotMsg,
|
||||
Name: "分发 Bot 消息",
|
||||
DisplayName: "分发 Bot 消息",
|
||||
Description: "向已绑定的平台账号异步下发 Bot 文本消息",
|
||||
Category: "messaging",
|
||||
Queue: taskQueueDefault,
|
||||
Retryable: true,
|
||||
Params: []contracts.TaskParamDTO{
|
||||
{Name: paramNameText, Label: "消息内容", Type: consts.TypeText, Required: true, Placeholder: "要发送的文本", Description: "下发给绑定用户的文本"},
|
||||
{Name: "channel_id", Label: "频道 ID", Type: taskParamTypeString, Required: false, Placeholder: "留空表示全部启用频道", Description: "仅向指定频道的绑定发送"},
|
||||
{Name: "user_id", Label: "用户 ID", Type: taskParamTypeString, Required: false, Placeholder: "留空表示频道下全部绑定", Description: "仅向指定 Wavelet 用户的绑定发送"},
|
||||
},
|
||||
}
|
||||
|
||||
type botDispatchPayload struct {
|
||||
Text string `json:"text"`
|
||||
ChannelID uint64 `json:"channel_id,string"`
|
||||
UserID uint64 `json:"user_id,string"`
|
||||
}
|
||||
|
||||
// BotDispatchHandler sends a text message through enabled bot channels.
|
||||
type BotDispatchHandler struct{}
|
||||
|
||||
// ValidatePayload requires a non-empty message body.
|
||||
func (h *BotDispatchHandler) ValidatePayload(payload []byte) ([]byte, error) {
|
||||
p, err := parseBotDispatchPayload(payload)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return json.Marshal(p)
|
||||
}
|
||||
|
||||
// Execute delivers the text to matching channel bindings.
|
||||
func (h *BotDispatchHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) {
|
||||
p, err := parseBotDispatchPayload(payload)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
channels, err := dao.ListEnabledMessageChannels(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if p.ChannelID != 0 {
|
||||
filtered := channels[:0]
|
||||
for i := range channels {
|
||||
if channels[i].ID == p.ChannelID {
|
||||
filtered = append(filtered, channels[i])
|
||||
}
|
||||
}
|
||||
channels = filtered
|
||||
if len(channels) == 0 {
|
||||
return nil, errors.New(consts.ErrChannelNotFound)
|
||||
}
|
||||
}
|
||||
|
||||
sent := 0
|
||||
failed := 0
|
||||
for i := range channels {
|
||||
n, ferr := dispatchOnChannel(ctx, &channels[i], p.UserID, p.Text)
|
||||
sent += n
|
||||
failed += ferr
|
||||
}
|
||||
msg := fmt.Sprintf("Bot 消息已尝试发送,成功 %d,失败 %d", sent, failed)
|
||||
if svc := GetTaskService(ctx); svc != nil {
|
||||
svc.AppendLog(ctx, "%s", msg)
|
||||
}
|
||||
if sent == 0 && failed > 0 {
|
||||
return nil, errors.New(msg)
|
||||
}
|
||||
return &contracts.TaskResultDTO{Message: msg}, nil
|
||||
}
|
||||
|
||||
func parseBotDispatchPayload(payload []byte) (botDispatchPayload, error) {
|
||||
var p botDispatchPayload
|
||||
if len(payload) > 0 {
|
||||
if err := json.Unmarshal(payload, &p); err != nil {
|
||||
return p, fmt.Errorf("%s: %w", consts.ErrInvalidJSONFormat, err)
|
||||
}
|
||||
}
|
||||
p.Text = strings.TrimSpace(p.Text)
|
||||
if p.Text == "" {
|
||||
return p, errors.New(consts.ErrBotDispatchTextRequired)
|
||||
}
|
||||
return p, nil
|
||||
}
|
||||
|
||||
func dispatchOnChannel(ctx context.Context, row *entity.MessageChannel, userID uint64, text string) (sent, failed int) {
|
||||
factory, ok := Lookup(row.Type)
|
||||
if !ok {
|
||||
logger.ErrorF(ctx, "bot dispatch: %s type=%s", consts.ErrBotChannelNotRegistered, row.Type)
|
||||
return 0, 1
|
||||
}
|
||||
cfg, err := channelConfigFromRow(row)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "bot dispatch: decode channel %d: %v", row.ID, err)
|
||||
return 0, 1
|
||||
}
|
||||
ch, err := factory(cfg, nil)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "bot dispatch: create adapter %d: %v", row.ID, err)
|
||||
return 0, 1
|
||||
}
|
||||
if err := ch.Connect(ctx); err != nil {
|
||||
logger.ErrorF(ctx, "bot dispatch: connect channel %d: %v", row.ID, err)
|
||||
return 0, 1
|
||||
}
|
||||
defer func() { _ = ch.Disconnect(ctx) }()
|
||||
|
||||
bindings, err := dao.ListBindingsByChannel(ctx, row.ID)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "bot dispatch: list bindings %d: %v", row.ID, err)
|
||||
return 0, 1
|
||||
}
|
||||
for i := range bindings {
|
||||
if userID != 0 && bindings[i].UserID != userID {
|
||||
continue
|
||||
}
|
||||
to := do.Recipient{
|
||||
ChatID: bindings[i].PlatformUserID,
|
||||
PlatformUserID: bindings[i].PlatformUserID,
|
||||
}
|
||||
if err := ch.Send(ctx, to, do.OutboundMessage{Text: text}); err != nil {
|
||||
logger.ErrorF(ctx, "bot dispatch: send channel=%d user=%d: %v", row.ID, bindings[i].UserID, err)
|
||||
failed++
|
||||
continue
|
||||
}
|
||||
sent++
|
||||
}
|
||||
return sent, failed
|
||||
}
|
||||
|
||||
func channelConfigFromRow(row *entity.MessageChannel) (do.ChannelConfig, error) {
|
||||
creds, err := DecryptCredentials(row.Credentials)
|
||||
if err != nil {
|
||||
return do.ChannelConfig{}, err
|
||||
}
|
||||
if creds == nil {
|
||||
creds = map[string]string{}
|
||||
}
|
||||
if creds["bot_token"] == "" && creds["token"] != "" {
|
||||
creds["bot_token"] = creds["token"]
|
||||
}
|
||||
if creds["app_secret"] == "" && creds["client_secret"] != "" {
|
||||
creds["app_secret"] = creds["client_secret"]
|
||||
}
|
||||
extra := ParseExtra(row.Extra)
|
||||
if extra["base_url"] == "" && creds["api_base"] != "" {
|
||||
extra["base_url"] = creds["api_base"]
|
||||
}
|
||||
return do.ChannelConfig{
|
||||
ID: row.ID,
|
||||
Type: row.Type,
|
||||
Name: row.Name,
|
||||
Credentials: creds,
|
||||
Extra: extra,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service_test
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/msg_gateway/dao"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"Wavelet/plugins/domain/msg_gateway/service"
|
||||
"context"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type dispatchTestDB struct{ db *gorm.DB }
|
||||
|
||||
func (m *dispatchTestDB) GORM() *gorm.DB { return m.db }
|
||||
func (m *dispatchTestDB) DB(ctx context.Context) *gorm.DB { return m.db.WithContext(ctx) }
|
||||
func (m *dispatchTestDB) Named(_ string) *gorm.DB { return m.db }
|
||||
|
||||
func TestBotDispatchValidatePayload(t *testing.T) {
|
||||
h := &service.BotDispatchHandler{}
|
||||
_, err := h.ValidatePayload([]byte(`{}`))
|
||||
require.Error(t, err)
|
||||
_, err = h.ValidatePayload([]byte(`{"text":"hello"}`))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestBotDispatchNoChannels(t *testing.T) {
|
||||
testDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "dispatch.db")), &gorm.Config{})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, testDB.AutoMigrate(&entity.MessageChannel{}, &entity.MessageBinding{}))
|
||||
dao.SetDBServiceForTest(&dispatchTestDB{db: testDB})
|
||||
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
||||
|
||||
h := &service.BotDispatchHandler{}
|
||||
res, err := h.Execute(context.Background(), []byte(`{"text":"hello"}`))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, res)
|
||||
assert.Contains(t, res.Message, "成功 0")
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,405 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package service implements domain business logic and channel runners for msg_gateway.
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"Wavelet/plugins/domain/msg_gateway/dao"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
// Handler processes one inbound message.
|
||||
type Handler func(ctx context.Context, msg do.InboundMessage) error
|
||||
|
||||
// Factory constructs a Channel from decrypted config.
|
||||
type Factory func(cfg do.ChannelConfig, onInbound Handler) (Channel, error)
|
||||
|
||||
// Channel is one connected messaging adapter.
|
||||
type Channel interface {
|
||||
Type() string
|
||||
Connect(ctx context.Context) error
|
||||
Disconnect(ctx context.Context) error
|
||||
Send(ctx context.Context, to do.Recipient, msg do.OutboundMessage) error
|
||||
Capabilities() do.Capability
|
||||
}
|
||||
|
||||
var (
|
||||
factoriesMu sync.RWMutex
|
||||
factories = map[string]Factory{}
|
||||
)
|
||||
|
||||
// Register stores a channel factory under typ.
|
||||
func Register(typ string, fn Factory) {
|
||||
factoriesMu.Lock()
|
||||
defer factoriesMu.Unlock()
|
||||
factories[typ] = fn
|
||||
}
|
||||
|
||||
// Lookup returns a previously registered factory.
|
||||
func Lookup(typ string) (Factory, bool) {
|
||||
factoriesMu.RLock()
|
||||
defer factoriesMu.RUnlock()
|
||||
fn, ok := factories[typ]
|
||||
return fn, ok
|
||||
}
|
||||
|
||||
// Re-exported constants.
|
||||
const (
|
||||
CodeAlphabet = consts.CodeAlphabet
|
||||
CodeLength = consts.CodeLength
|
||||
)
|
||||
|
||||
// GenerateCode returns an 8-character pairing code.
|
||||
func GenerateCode() (string, error) {
|
||||
buf := make([]byte, CodeLength)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
out := make([]byte, CodeLength)
|
||||
for i, b := range buf {
|
||||
out[i] = CodeAlphabet[int(b)%len(CodeAlphabet)]
|
||||
}
|
||||
return string(out), nil
|
||||
}
|
||||
|
||||
// NormalizeCode strips separators and uppercases.
|
||||
func NormalizeCode(s string) string {
|
||||
var b strings.Builder
|
||||
for _, r := range s {
|
||||
if r == '-' || unicode.IsSpace(r) {
|
||||
continue
|
||||
}
|
||||
b.WriteRune(unicode.ToUpper(r))
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// FormatCode renders ABCD-EFGH.
|
||||
func FormatCode(s string) string {
|
||||
s = NormalizeCode(s)
|
||||
if len(s) != CodeLength {
|
||||
return s
|
||||
}
|
||||
return s[:4] + "-" + s[4:]
|
||||
}
|
||||
|
||||
var (
|
||||
credentialSecretMu sync.RWMutex
|
||||
credentialSecret string
|
||||
)
|
||||
|
||||
// SetCredentialSecret sets the secret used to derive CredentialKey.
|
||||
func SetCredentialSecret(secret string) {
|
||||
credentialSecretMu.Lock()
|
||||
defer credentialSecretMu.Unlock()
|
||||
credentialSecret = secret
|
||||
}
|
||||
|
||||
// CredentialKey is AES-256 hex derived from the session secret.
|
||||
func CredentialKey() string {
|
||||
credentialSecretMu.RLock()
|
||||
secret := credentialSecret
|
||||
credentialSecretMu.RUnlock()
|
||||
sum := sha256.Sum256([]byte(secret))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// EncryptCredentials encrypts a credential map as JSON.
|
||||
func EncryptCredentials(creds map[string]string) (string, error) {
|
||||
if creds == nil {
|
||||
creds = map[string]string{}
|
||||
}
|
||||
raw, err := json.Marshal(creds)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return util.Encrypt(CredentialKey(), string(raw))
|
||||
}
|
||||
|
||||
// DecryptCredentials decrypts a credential map.
|
||||
func DecryptCredentials(ciphertext string) (map[string]string, error) {
|
||||
if ciphertext == "" {
|
||||
return map[string]string{}, nil
|
||||
}
|
||||
plain, err := util.Decrypt(CredentialKey(), ciphertext)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var out map[string]string
|
||||
if err := json.Unmarshal([]byte(plain), &out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if out == nil {
|
||||
out = map[string]string{}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// ParseExtra decodes optional extra JSON into a string map.
|
||||
func ParseExtra(raw string) map[string]string {
|
||||
if raw == "" {
|
||||
return map[string]string{}
|
||||
}
|
||||
var out map[string]string
|
||||
if err := json.Unmarshal([]byte(raw), &out); err != nil || out == nil {
|
||||
return map[string]string{}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// EncodeExtra encodes extra fields as JSON.
|
||||
func EncodeExtra(extra map[string]string) string {
|
||||
if extra == nil {
|
||||
return ""
|
||||
}
|
||||
raw, err := json.Marshal(extra)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return string(raw)
|
||||
}
|
||||
|
||||
// Runner manages lifecycle for long-lived channel adapters (WebSocket, long-polling, etc.).
|
||||
type Runner struct {
|
||||
mu sync.Mutex
|
||||
running bool
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
// GlobalRunner is the default global runner instance.
|
||||
var GlobalRunner = &Runner{}
|
||||
|
||||
// Start starts all background long-lived channel runners.
|
||||
func Start(ctx context.Context) error {
|
||||
GlobalRunner.mu.Lock()
|
||||
defer GlobalRunner.mu.Unlock()
|
||||
|
||||
if GlobalRunner.running {
|
||||
return nil
|
||||
}
|
||||
|
||||
runCtx, cancel := context.WithCancel(ctx)
|
||||
GlobalRunner.cancel = cancel
|
||||
GlobalRunner.running = true
|
||||
|
||||
logger.InfoF(runCtx, "[MessageGateway] Starting bot channel runners...")
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stop stops the channel runner.
|
||||
func Stop() {
|
||||
GlobalRunner.mu.Lock()
|
||||
defer GlobalRunner.mu.Unlock()
|
||||
|
||||
if !GlobalRunner.running {
|
||||
return
|
||||
}
|
||||
|
||||
if GlobalRunner.cancel != nil {
|
||||
GlobalRunner.cancel()
|
||||
}
|
||||
GlobalRunner.running = false
|
||||
}
|
||||
|
||||
// Cordis contract singletons consumed by service layer.
|
||||
var (
|
||||
cacheMu sync.RWMutex
|
||||
cacheSvc contracts.CacheService
|
||||
taskMu sync.RWMutex
|
||||
taskSvc contracts.TaskService
|
||||
userMu sync.RWMutex
|
||||
userSvc contracts.UserService
|
||||
)
|
||||
|
||||
// SetCacheService sets the cache service.
|
||||
func SetCacheService(s contracts.CacheService) {
|
||||
cacheMu.Lock()
|
||||
defer cacheMu.Unlock()
|
||||
cacheSvc = s
|
||||
}
|
||||
|
||||
// SetTaskService sets the task service.
|
||||
func SetTaskService(s contracts.TaskService) {
|
||||
taskMu.Lock()
|
||||
defer taskMu.Unlock()
|
||||
taskSvc = s
|
||||
}
|
||||
|
||||
// SetUserService sets the user service.
|
||||
func SetUserService(s contracts.UserService) {
|
||||
userMu.Lock()
|
||||
defer userMu.Unlock()
|
||||
userSvc = s
|
||||
}
|
||||
|
||||
// GetCache resolves the cache service for the context.
|
||||
func GetCache(ctx context.Context) contracts.CacheService {
|
||||
if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
cacheMu.RLock()
|
||||
s := cacheSvc
|
||||
cacheMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// GetTaskService returns the task service.
|
||||
func GetTaskService(ctx context.Context) contracts.TaskService {
|
||||
if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
taskMu.RLock()
|
||||
defer taskMu.RUnlock()
|
||||
return taskSvc
|
||||
}
|
||||
|
||||
// GetUserService resolves the user service for the context.
|
||||
func GetUserService(ctx context.Context) contracts.UserService {
|
||||
if s, err := core.InjectFrom[contracts.UserService](ctx); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
userMu.RLock()
|
||||
s := userSvc
|
||||
userMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// BindChannel consumes a pairing code and binds the platform identity to the user.
|
||||
func BindChannel(ctx context.Context, userID uint64, req do.BindRequest) (do.BindingDTO, error) {
|
||||
channelID, err := strconv.ParseUint(strings.TrimSpace(req.ChannelID), 10, 64)
|
||||
if err != nil || channelID == 0 {
|
||||
return do.BindingDTO{}, consts.ErrChannelIDRequired
|
||||
}
|
||||
code := NormalizeCode(req.Code)
|
||||
if code == "" {
|
||||
return do.BindingDTO{}, consts.ErrCodeInvalid
|
||||
}
|
||||
pairing, err := dao.GetPairingCode(ctx, code)
|
||||
if err != nil {
|
||||
if errors.Is(err, consts.ErrRecordNotFound) {
|
||||
return do.BindingDTO{}, consts.ErrCodeInvalid
|
||||
}
|
||||
return do.BindingDTO{}, err
|
||||
}
|
||||
if !pairing.ExpiresAt.After(time.Now()) {
|
||||
return do.BindingDTO{}, consts.ErrCodeInvalid
|
||||
}
|
||||
if pairing.ChannelID != channelID {
|
||||
return do.BindingDTO{}, consts.ErrChannelMismatch
|
||||
}
|
||||
ch, err := dao.GetMessageChannel(ctx, channelID)
|
||||
if err != nil {
|
||||
if errors.Is(err, consts.ErrRecordNotFound) {
|
||||
return do.BindingDTO{}, consts.ErrCodeInvalid
|
||||
}
|
||||
return do.BindingDTO{}, err
|
||||
}
|
||||
if !ch.Enabled {
|
||||
return do.BindingDTO{}, consts.ErrChannelDisabled
|
||||
}
|
||||
|
||||
existing, err := dao.GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID)
|
||||
if err != nil && !errors.Is(err, consts.ErrRecordNotFound) {
|
||||
return do.BindingDTO{}, err
|
||||
}
|
||||
if err == nil && existing != nil {
|
||||
if existing.UserID != userID {
|
||||
return do.BindingDTO{}, consts.ErrPlatformAlreadyBound
|
||||
}
|
||||
_ = dao.DeletePairingCode(ctx, pairing.Code)
|
||||
return ToBindingDTO(existing, ch), nil
|
||||
}
|
||||
|
||||
row := &entity.MessageBinding{
|
||||
UserID: userID,
|
||||
ChannelID: channelID,
|
||||
PlatformUserID: pairing.PlatformUserID,
|
||||
}
|
||||
if err := dao.CreateMessageBinding(ctx, row); err != nil {
|
||||
return do.BindingDTO{}, err
|
||||
}
|
||||
if err := dao.DeletePairingCode(ctx, pairing.Code); err != nil {
|
||||
return do.BindingDTO{}, err
|
||||
}
|
||||
return ToBindingDTO(row, ch), nil
|
||||
}
|
||||
|
||||
// ListEnabledPublicChannels returns the channels a user may bind to.
|
||||
func ListEnabledPublicChannels(ctx context.Context) ([]do.PublicChannelDTO, error) {
|
||||
rows, err := dao.ListEnabledMessageChannels(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]do.PublicChannelDTO, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
out = append(out, do.PublicChannelDTO{ID: row.ID, Name: row.Name, Type: row.Type})
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// ListUserBindings returns the binding rows of one user enriched with channel info.
|
||||
func ListUserBindings(ctx context.Context, userID uint64) ([]do.BindingDTO, error) {
|
||||
rows, err := dao.ListBindingsByUser(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]do.BindingDTO, 0, len(rows))
|
||||
for i := range rows {
|
||||
ch, err := dao.GetMessageChannel(ctx, rows[i].ChannelID)
|
||||
if err != nil {
|
||||
out = append(out, ToBindingDTO(&rows[i], nil))
|
||||
continue
|
||||
}
|
||||
out = append(out, ToBindingDTO(&rows[i], ch))
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// UnbindChannel removes a binding owned by the given user.
|
||||
func UnbindChannel(ctx context.Context, userID, bindingID uint64) error {
|
||||
row, err := dao.GetMessageBinding(ctx, bindingID)
|
||||
if err != nil {
|
||||
if errors.Is(err, consts.ErrRecordNotFound) {
|
||||
return consts.ErrBindingNotFound
|
||||
}
|
||||
return err
|
||||
}
|
||||
if row.UserID != userID {
|
||||
return consts.ErrBindingForbidden
|
||||
}
|
||||
return dao.DeleteMessageBinding(ctx, bindingID)
|
||||
}
|
||||
|
||||
// ToBindingDTO projects a binding row and its optional channel onto the user DTO.
|
||||
func ToBindingDTO(row *entity.MessageBinding, ch *entity.MessageChannel) do.BindingDTO {
|
||||
dto := do.BindingDTO{
|
||||
ID: row.ID,
|
||||
UserID: row.UserID,
|
||||
ChannelID: row.ChannelID,
|
||||
PlatformUserID: row.PlatformUserID,
|
||||
CreatedAt: row.CreatedAt,
|
||||
}
|
||||
if ch != nil {
|
||||
dto.ChannelName = ch.Name
|
||||
dto.ChannelType = ch.Type
|
||||
}
|
||||
return dto
|
||||
}
|
||||
Reference in New Issue
Block a user