mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 16:46:37 +08:00
refactor(msg_gateway): decouple bot gateway and push notification architecture
- Split shared monolithic consts into bot, push, and errs with typed sentinel errors - Restructure model layer into distinct bot and push subdomains - Refactor DAO layer to enforce single-owner principle and remove cross-table raw SQL queries - Decompose 1150+ line service/push.go into push_channel, push_event, push_trigger, push_worker, and push_template - Clean up controller layer with generic request handlers and parameter validation in controller/base.go - Streamline plugin.go to core Cordis lifecycle orchestration and remove re-export bloat - Verify all unit tests, race tests, Cordis architecture rules, and Swagger generation pass cleanly
This commit is contained in:
@@ -0,0 +1,160 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package dao
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// CreateMessageChannel inserts a channel row.
|
||||
func CreateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error {
|
||||
if ch.ID == 0 {
|
||||
ch.ID = idgen.NextUint64ID()
|
||||
}
|
||||
return GetDB(ctx).Create(ch).Error
|
||||
}
|
||||
|
||||
// UpdateMessageChannel saves a channel row.
|
||||
func UpdateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error {
|
||||
return GetDB(ctx).Save(ch).Error
|
||||
}
|
||||
|
||||
// GetMessageChannel loads a channel by id.
|
||||
func GetMessageChannel(ctx context.Context, id uint64) (*entity.MessageChannel, error) {
|
||||
var ch entity.MessageChannel
|
||||
if err := GetDB(ctx).Where("id = ?", id).First(&ch).Error; err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
}
|
||||
return &ch, nil
|
||||
}
|
||||
|
||||
// ListMessageChannels returns all channels newest first.
|
||||
func ListMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) {
|
||||
var rows []entity.MessageChannel
|
||||
if err := GetDB(ctx).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
// DeleteMessageChannel removes pairings, bindings, then the channel.
|
||||
func DeleteMessageChannel(ctx context.Context, id uint64) error {
|
||||
return GetDB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("channel_id = ?", id).Delete(&entity.MessagePairingCode{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("channel_id = ?", id).Delete(&entity.MessageBinding{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&entity.MessageChannel{}, id).Error
|
||||
})
|
||||
}
|
||||
|
||||
// CreateMessageBinding inserts a binding.
|
||||
func CreateMessageBinding(ctx context.Context, b *entity.MessageBinding) error {
|
||||
if b.ID == 0 {
|
||||
b.ID = idgen.NextUint64ID()
|
||||
}
|
||||
return GetDB(ctx).Create(b).Error
|
||||
}
|
||||
|
||||
// GetBindingByChannelPlatform finds a binding for a platform user on a channel.
|
||||
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*entity.MessageBinding, error) {
|
||||
var b entity.MessageBinding
|
||||
err := GetDB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error
|
||||
if err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
}
|
||||
return &b, nil
|
||||
}
|
||||
|
||||
// ListBindingsByUser lists bindings for a Wavelet user.
|
||||
func ListBindingsByUser(ctx context.Context, userID uint64) ([]entity.MessageBinding, error) {
|
||||
var rows []entity.MessageBinding
|
||||
if err := GetDB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
// ListBindingsByChannel lists bindings on one messaging channel.
|
||||
func ListBindingsByChannel(ctx context.Context, channelID uint64) ([]entity.MessageBinding, error) {
|
||||
var rows []entity.MessageBinding
|
||||
if err := GetDB(ctx).Where("channel_id = ?", channelID).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
// GetMessageBinding loads a binding by id.
|
||||
func GetMessageBinding(ctx context.Context, id uint64) (*entity.MessageBinding, error) {
|
||||
var b entity.MessageBinding
|
||||
if err := GetDB(ctx).Where("id = ?", id).First(&b).Error; err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
}
|
||||
return &b, nil
|
||||
}
|
||||
|
||||
// DeleteMessageBinding deletes a binding by id.
|
||||
func DeleteMessageBinding(ctx context.Context, id uint64) error {
|
||||
return GetDB(ctx).Delete(&entity.MessageBinding{}, id).Error
|
||||
}
|
||||
|
||||
// UpsertPairingCode reuses an unexpired code for the same channel+platform user.
|
||||
func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*entity.MessagePairingCode, error) {
|
||||
var existing entity.MessagePairingCode
|
||||
err := GetDB(ctx).
|
||||
Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()).
|
||||
First(&existing).Error
|
||||
if err == nil {
|
||||
return &existing, nil
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
row := &entity.MessagePairingCode{
|
||||
Code: code,
|
||||
ChannelID: channelID,
|
||||
PlatformUserID: platformUserID,
|
||||
ExpiresAt: expiresAt,
|
||||
}
|
||||
if err := GetDB(ctx).Create(row).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return row, nil
|
||||
}
|
||||
|
||||
// GetPairingCode loads a pairing code by normalized code string.
|
||||
func GetPairingCode(ctx context.Context, code string) (*entity.MessagePairingCode, error) {
|
||||
var row entity.MessagePairingCode
|
||||
if err := GetDB(ctx).Where("code = ?", code).First(&row).Error; err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
}
|
||||
return &row, nil
|
||||
}
|
||||
|
||||
// DeletePairingCode removes a pairing code.
|
||||
func DeletePairingCode(ctx context.Context, code string) error {
|
||||
return GetDB(ctx).Where("code = ?", code).Delete(&entity.MessagePairingCode{}).Error
|
||||
}
|
||||
|
||||
// DeleteExpiredPairingCodes removes expired pairing rows.
|
||||
func DeleteExpiredPairingCodes(ctx context.Context) error {
|
||||
return GetDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&entity.MessagePairingCode{}).Error
|
||||
}
|
||||
|
||||
// ListEnabledMessageChannels returns enabled channels.
|
||||
func ListEnabledMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) {
|
||||
var rows []entity.MessageChannel
|
||||
if err := GetDB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
@@ -0,0 +1,64 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package dao_test
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/msg_gateway/dao"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestBotDAO_ChannelAndBinding(t *testing.T) {
|
||||
_ = idgen.Init(1)
|
||||
db, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
require.NoError(t, db.AutoMigrate(&entity.MessageChannel{}, &entity.MessageBinding{}, &entity.MessagePairingCode{}))
|
||||
|
||||
dao.SetDBServiceForTest(stubDBService{db: db})
|
||||
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
ch := entity.MessageChannel{
|
||||
Name: "tg_bot",
|
||||
Type: "telegram",
|
||||
OwnerScope: "system",
|
||||
Credentials: "encrypted_token",
|
||||
Enabled: true,
|
||||
}
|
||||
require.NoError(t, dao.CreateMessageChannel(ctx, &ch))
|
||||
assert.NotZero(t, ch.ID)
|
||||
|
||||
code, err := dao.UpsertPairingCode(ctx, ch.ID, "tg_user_1", "ABCD1234", time.Now().Add(10*time.Minute))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "ABCD1234", code.Code)
|
||||
|
||||
// Reusing pairing code
|
||||
code2, err := dao.UpsertPairingCode(ctx, ch.ID, "tg_user_1", "XYZ9999", time.Now().Add(10*time.Minute))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "ABCD1234", code2.Code)
|
||||
|
||||
binding := entity.MessageBinding{
|
||||
UserID: 42,
|
||||
ChannelID: ch.ID,
|
||||
PlatformUserID: "tg_user_1",
|
||||
}
|
||||
require.NoError(t, dao.CreateMessageBinding(ctx, &binding))
|
||||
assert.NotZero(t, binding.ID)
|
||||
|
||||
bindings, err := dao.ListBindingsByUser(ctx, 42)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, bindings, 1)
|
||||
|
||||
require.NoError(t, dao.DeleteMessageChannel(ctx, ch.ID))
|
||||
_, err = dao.GetMessageChannel(ctx, ch.ID)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
@@ -7,27 +7,21 @@ package dao
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
cacheMu sync.RWMutex
|
||||
cacheSvc contracts.CacheService
|
||||
)
|
||||
|
||||
// SetDBServiceForTest injects a DBService for tests. Production wiring must use Apply.
|
||||
func SetDBServiceForTest(s contracts.DBService) {
|
||||
SetDBService(s)
|
||||
}
|
||||
|
||||
// SetDBService sets the database service singleton.
|
||||
func SetDBService(s contracts.DBService) {
|
||||
dbMu.Lock()
|
||||
@@ -35,8 +29,19 @@ func SetDBService(s contracts.DBService) {
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
// GetDB resolves the persistence handle for the current call, preferring an
|
||||
// explicitly injected *core.Context before falling back to the plugin singleton.
|
||||
// SetDBServiceForTest injects a DBService for tests.
|
||||
func SetDBServiceForTest(s contracts.DBService) {
|
||||
SetDBService(s)
|
||||
}
|
||||
|
||||
// SetCacheService sets the cache service singleton.
|
||||
func SetCacheService(s contracts.CacheService) {
|
||||
cacheMu.Lock()
|
||||
defer cacheMu.Unlock()
|
||||
cacheSvc = s
|
||||
}
|
||||
|
||||
// GetDB resolves the persistence handle for the current call.
|
||||
func GetDB(ctx context.Context) *gorm.DB {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||
@@ -52,6 +57,19 @@ func GetDB(ctx context.Context) *gorm.DB {
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetCache resolves the cache service for the current call.
|
||||
func GetCache(ctx context.Context) contracts.CacheService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
cacheMu.RLock()
|
||||
s := cacheSvc
|
||||
cacheMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// mapNotFound translates GORM's missing-row sentinel into the plugin-level
|
||||
// consts.ErrRecordNotFound so the service and controller layers stay free of gorm imports.
|
||||
func mapNotFound(err error) error {
|
||||
@@ -60,149 +78,3 @@ func mapNotFound(err error) error {
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// CreateMessageChannel inserts a channel row.
|
||||
func CreateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error {
|
||||
if ch.ID == 0 {
|
||||
ch.ID = idgen.NextUint64ID()
|
||||
}
|
||||
return GetDB(ctx).Create(ch).Error
|
||||
}
|
||||
|
||||
// UpdateMessageChannel saves a channel row.
|
||||
func UpdateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error {
|
||||
return GetDB(ctx).Save(ch).Error
|
||||
}
|
||||
|
||||
// GetMessageChannel loads a channel by id.
|
||||
func GetMessageChannel(ctx context.Context, id uint64) (*entity.MessageChannel, error) {
|
||||
var ch entity.MessageChannel
|
||||
if err := GetDB(ctx).Where("id = ?", id).First(&ch).Error; err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
}
|
||||
return &ch, nil
|
||||
}
|
||||
|
||||
// ListMessageChannels returns all channels newest first.
|
||||
func ListMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) {
|
||||
var rows []entity.MessageChannel
|
||||
if err := GetDB(ctx).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
// DeleteMessageChannel removes pairings, bindings, then the channel.
|
||||
func DeleteMessageChannel(ctx context.Context, id uint64) error {
|
||||
return GetDB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("channel_id = ?", id).Delete(&entity.MessagePairingCode{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("channel_id = ?", id).Delete(&entity.MessageBinding{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&entity.MessageChannel{}, id).Error
|
||||
})
|
||||
}
|
||||
|
||||
// CreateMessageBinding inserts a binding.
|
||||
func CreateMessageBinding(ctx context.Context, b *entity.MessageBinding) error {
|
||||
if b.ID == 0 {
|
||||
b.ID = idgen.NextUint64ID()
|
||||
}
|
||||
return GetDB(ctx).Create(b).Error
|
||||
}
|
||||
|
||||
// GetBindingByChannelPlatform finds a binding for a platform user on a channel.
|
||||
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*entity.MessageBinding, error) {
|
||||
var b entity.MessageBinding
|
||||
err := GetDB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error
|
||||
if err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
}
|
||||
return &b, nil
|
||||
}
|
||||
|
||||
// ListBindingsByUser lists bindings for a Wavelet user.
|
||||
func ListBindingsByUser(ctx context.Context, userID uint64) ([]entity.MessageBinding, error) {
|
||||
var rows []entity.MessageBinding
|
||||
if err := GetDB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
// ListBindingsByChannel lists bindings on one messaging channel.
|
||||
func ListBindingsByChannel(ctx context.Context, channelID uint64) ([]entity.MessageBinding, error) {
|
||||
var rows []entity.MessageBinding
|
||||
if err := GetDB(ctx).Where("channel_id = ?", channelID).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
// GetMessageBinding loads a binding by id.
|
||||
func GetMessageBinding(ctx context.Context, id uint64) (*entity.MessageBinding, error) {
|
||||
var b entity.MessageBinding
|
||||
if err := GetDB(ctx).Where("id = ?", id).First(&b).Error; err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
}
|
||||
return &b, nil
|
||||
}
|
||||
|
||||
// DeleteMessageBinding deletes a binding by id.
|
||||
func DeleteMessageBinding(ctx context.Context, id uint64) error {
|
||||
return GetDB(ctx).Delete(&entity.MessageBinding{}, id).Error
|
||||
}
|
||||
|
||||
// UpsertPairingCode reuses an unexpired code for the same channel+platform user.
|
||||
func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*entity.MessagePairingCode, error) {
|
||||
var existing entity.MessagePairingCode
|
||||
err := GetDB(ctx).
|
||||
Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()).
|
||||
First(&existing).Error
|
||||
if err == nil {
|
||||
return &existing, nil
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
row := &entity.MessagePairingCode{
|
||||
Code: code,
|
||||
ChannelID: channelID,
|
||||
PlatformUserID: platformUserID,
|
||||
ExpiresAt: expiresAt,
|
||||
}
|
||||
if err := GetDB(ctx).Create(row).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return row, nil
|
||||
}
|
||||
|
||||
// GetPairingCode loads a pairing code by normalized code string.
|
||||
func GetPairingCode(ctx context.Context, code string) (*entity.MessagePairingCode, error) {
|
||||
var row entity.MessagePairingCode
|
||||
if err := GetDB(ctx).Where("code = ?", code).First(&row).Error; err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
}
|
||||
return &row, nil
|
||||
}
|
||||
|
||||
// DeletePairingCode removes a pairing code.
|
||||
func DeletePairingCode(ctx context.Context, code string) error {
|
||||
return GetDB(ctx).Where("code = ?", code).Delete(&entity.MessagePairingCode{}).Error
|
||||
}
|
||||
|
||||
// DeleteExpiredPairingCodes removes expired pairing rows.
|
||||
func DeleteExpiredPairingCodes(ctx context.Context) error {
|
||||
return GetDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&entity.MessagePairingCode{}).Error
|
||||
}
|
||||
|
||||
// ListEnabledMessageChannels returns enabled channels.
|
||||
func ListEnabledMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) {
|
||||
var rows []entity.MessageChannel
|
||||
if err := GetDB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
@@ -4,15 +4,9 @@
|
||||
package dao
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
@@ -23,31 +17,6 @@ const (
|
||||
activePushEventCacheTTL = 24 * time.Hour
|
||||
)
|
||||
|
||||
var (
|
||||
cacheMu sync.RWMutex
|
||||
cacheSvc contracts.CacheService
|
||||
)
|
||||
|
||||
// SetCacheService sets the cache service singleton.
|
||||
func SetCacheService(s contracts.CacheService) {
|
||||
cacheMu.Lock()
|
||||
defer cacheMu.Unlock()
|
||||
cacheSvc = s
|
||||
}
|
||||
|
||||
// GetCache resolves the cache service for the current call.
|
||||
func GetCache(ctx context.Context) contracts.CacheService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
cacheMu.RLock()
|
||||
s := cacheSvc
|
||||
cacheMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// ListPushChannelsRecord returns all push channels ordered by creation time descending.
|
||||
func ListPushChannelsRecord(ctx context.Context) ([]entity.PushChannel, error) {
|
||||
var channels []entity.PushChannel
|
||||
@@ -112,14 +81,15 @@ func DeletePushChannelRecord(ctx context.Context, channel *entity.PushChannel) e
|
||||
}
|
||||
|
||||
func getCachedOrQuery[T any](ctx context.Context, cacheKey string, ttl time.Duration, query func(db *gorm.DB, dest *T) error) (*T, error) {
|
||||
var val T
|
||||
if cache := GetCache(ctx); cache != nil {
|
||||
var val T
|
||||
if err := cache.Get(ctx, cacheKey, &val); err == nil {
|
||||
return &val, nil
|
||||
}
|
||||
}
|
||||
|
||||
db := GetDB(ctx)
|
||||
var val T
|
||||
if err := query(db, &val); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -255,6 +225,9 @@ func ListPushHistoriesRecord(ctx context.Context, filter do.PushHistoryListFilte
|
||||
if filter.EventKey != "" {
|
||||
query = query.Where("event_key = ?", filter.EventKey)
|
||||
}
|
||||
if filter.Channel != "" {
|
||||
query = query.Where("channel = ?", filter.Channel)
|
||||
}
|
||||
if filter.Status != "" {
|
||||
query = query.Where("status = ?", filter.Status)
|
||||
}
|
||||
@@ -282,77 +255,3 @@ func CreatePushHistoryRecord(ctx context.Context, history *entity.PushHistory) e
|
||||
func PushHistoryQuery(ctx context.Context) *gorm.DB {
|
||||
return GetDB(ctx).Model(&entity.PushHistory{})
|
||||
}
|
||||
|
||||
// smtpConfigKeys are the system-config rows backing the built-in email channel.
|
||||
var smtpConfigKeys = []string{"smtp_host", "smtp_port", "smtp_username", "smtp_password"}
|
||||
|
||||
// LoadSMTPConfigRecord reads the SMTP settings in one query.
|
||||
func LoadSMTPConfigRecord(ctx context.Context) (do.SMTPConfig, error) {
|
||||
db := GetDB(ctx)
|
||||
if db == nil {
|
||||
return do.SMTPConfig{}, errors.New("database not available")
|
||||
}
|
||||
|
||||
var rows []struct {
|
||||
Key string
|
||||
Value string
|
||||
}
|
||||
if err := db.Table("w_system_configs").
|
||||
Select("key", "value").
|
||||
Where("key IN ?", smtpConfigKeys).
|
||||
Find(&rows).Error; err != nil {
|
||||
return do.SMTPConfig{}, fmt.Errorf("read smtp system configs: %w", err)
|
||||
}
|
||||
|
||||
var cfg do.SMTPConfig
|
||||
for _, row := range rows {
|
||||
switch row.Key {
|
||||
case "smtp_host":
|
||||
cfg.Host = row.Value
|
||||
case "smtp_port":
|
||||
cfg.Port = row.Value
|
||||
case "smtp_username":
|
||||
cfg.Username = row.Value
|
||||
case "smtp_password":
|
||||
cfg.Password = row.Value
|
||||
}
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
// userLookupColumns allow-lists the columns FindUserByFieldRecord may filter on.
|
||||
var userLookupColumns = map[string]struct{}{
|
||||
"id": {},
|
||||
"username": {},
|
||||
}
|
||||
|
||||
// FindUserByFieldRecord is the user lookup fallback for when the UserService
|
||||
// contract is not wired yet. field must be one of userLookupColumns.
|
||||
func FindUserByFieldRecord(ctx context.Context, field string, value any) (*contracts.UserDTO, error) {
|
||||
if _, ok := userLookupColumns[field]; !ok {
|
||||
return nil, consts.ErrUnsupportedUserLookupField
|
||||
}
|
||||
db := GetDB(ctx)
|
||||
if db == nil {
|
||||
return nil, consts.ErrRecordNotFound
|
||||
}
|
||||
var user contracts.UserDTO
|
||||
if err := db.Table("w_users").Where(field+" = ?", value).First(&user).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
// FindFirstAdminUserRecord is the admin lookup fallback for when the UserService
|
||||
// contract is not wired yet.
|
||||
func FindFirstAdminUserRecord(ctx context.Context) (*contracts.UserDTO, error) {
|
||||
db := GetDB(ctx)
|
||||
if db == nil {
|
||||
return nil, consts.ErrRecordNotFound
|
||||
}
|
||||
var adminUser contracts.UserDTO
|
||||
if err := db.Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &adminUser, nil
|
||||
}
|
||||
|
||||
@@ -4,128 +4,93 @@
|
||||
package dao_test
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"Wavelet/plugins/domain/msg_gateway/dao"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"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) 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 }
|
||||
|
||||
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) {
|
||||
func TestPushChannelDAO_CRUD(t *testing.T) {
|
||||
_ = idgen.Init(1)
|
||||
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)
|
||||
}
|
||||
require.NoError(t, db.AutoMigrate(&entity.PushChannel{}, &entity.PushEvent{}, &entity.PushHistory{}))
|
||||
|
||||
dao.SetDBServiceForTest(stubDBService{db: db})
|
||||
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
user, err := dao.FindUserByFieldRecord(ctx, "username", "seeded")
|
||||
if err != nil {
|
||||
t.Fatalf("allowlisted lookup by username failed: %v", err)
|
||||
}
|
||||
if user.ID != 77 {
|
||||
t.Errorf("allowlisted lookup returned ID %d, want 77", user.ID)
|
||||
}
|
||||
if _, err := dao.FindUserByFieldRecord(ctx, "id", uint64(77)); err != nil {
|
||||
t.Errorf("allowlisted lookup by id failed: %v", err)
|
||||
ch := entity.PushChannel{
|
||||
Name: "test_webhook",
|
||||
Type: "custom",
|
||||
URL: "https://example.com/hook",
|
||||
Enabled: true,
|
||||
}
|
||||
require.NoError(t, dao.CreatePushChannelRecord(ctx, &ch))
|
||||
assert.NotZero(t, ch.ID)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
field string
|
||||
}{
|
||||
{"tautology injection", `username = '' OR 1=1 --`},
|
||||
{"stacked statement", "id; DROP TABLE w_users"},
|
||||
{"column outside allow-list", "password"},
|
||||
{"empty field", ""},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
if _, err := dao.FindUserByFieldRecord(ctx, tc.field, "seeded"); !errors.Is(err, consts.ErrUnsupportedUserLookupField) {
|
||||
t.Errorf("%s: got err %v, want ErrUnsupportedUserLookupField", tc.name, err)
|
||||
}
|
||||
}
|
||||
loaded, err := dao.GetPushChannelByIDRecord(ctx, ch.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "test_webhook", loaded.Name)
|
||||
|
||||
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)
|
||||
}
|
||||
active, err := dao.GetActivePushChannelByName(ctx, "test_webhook")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, ch.ID, active.ID)
|
||||
|
||||
channels, err := dao.ListPushChannelsRecord(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, channels)
|
||||
|
||||
require.NoError(t, dao.DeletePushChannelRecord(ctx, &ch))
|
||||
_, err = dao.GetPushChannelByIDRecord(ctx, ch.ID)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
// smtpTestValues are the four system-config rows the built-in email channel reads.
|
||||
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) {
|
||||
func TestPushEventDAO_CRUD(t *testing.T) {
|
||||
_ = idgen.Init(1)
|
||||
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)
|
||||
}
|
||||
}
|
||||
require.NoError(t, db.AutoMigrate(&entity.PushChannel{}, &entity.PushEvent{}, &entity.PushHistory{}))
|
||||
|
||||
dao.SetDBServiceForTest(stubDBService{db: db})
|
||||
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
||||
|
||||
cfg, err := dao.LoadSMTPConfigRecord(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("LoadSMTPConfigRecord: %v", err)
|
||||
}
|
||||
if cfg.Host != smtpTestValues["smtp_host"] || cfg.Port != smtpTestValues["smtp_port"] ||
|
||||
cfg.Username != smtpTestValues["smtp_username"] || cfg.Password != smtpTestValues["smtp_password"] {
|
||||
t.Errorf("got %+v, want every SMTP field mapped from its own row", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadSMTPConfigRecordSurfacesReadFailure pins the actual defect: a read that
|
||||
// fails used to be discarded, returning four blank strings that callers could only
|
||||
// interpret as "SMTP was never configured", so the notification was dropped silently.
|
||||
func TestLoadSMTPConfigRecordSurfacesReadFailure(t *testing.T) {
|
||||
bare, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("open bare sqlite: %v", err)
|
||||
}
|
||||
|
||||
dao.SetDBServiceForTest(stubDBService{db: bare})
|
||||
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
||||
|
||||
if _, err := dao.LoadSMTPConfigRecord(context.Background()); err == nil {
|
||||
t.Fatal("LoadSMTPConfigRecord returned nil error although the config table cannot be read")
|
||||
ctx := context.Background()
|
||||
|
||||
ev := entity.PushEvent{
|
||||
EventKey: "test_event",
|
||||
Name: "测试事件",
|
||||
Channels: []string{"test_webhook"},
|
||||
Targets: []string{"admin"},
|
||||
Template: `{"title":"Hello"}`,
|
||||
Enabled: true,
|
||||
}
|
||||
require.NoError(t, dao.CreatePushEventRecord(ctx, &ev))
|
||||
assert.NotZero(t, ev.ID)
|
||||
|
||||
loaded, err := dao.GetPushEventByKeyRecord(ctx, "test_event")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "测试事件", loaded.Name)
|
||||
|
||||
require.NoError(t, dao.UpdatePushEventEnabledRecord(ctx, &ev, false))
|
||||
loadedDisabled, err := dao.GetPushEventByIDRecord(ctx, ev.ID)
|
||||
require.NoError(t, err)
|
||||
assert.False(t, loadedDisabled.Enabled)
|
||||
|
||||
require.NoError(t, dao.DeletePushEventRecord(ctx, &ev))
|
||||
_, err = dao.GetPushEventByIDRecord(ctx, ev.ID)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user