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

This commit is contained in:
ryan
2026-09-02 22:50:04 +08:00
parent 87e3bfd0e6
commit 8395dd5019
57 changed files with 2150 additions and 2120 deletions
@@ -0,0 +1,208 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dao provides database persistence and caching for the msg_gateway plugin.
package dao
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/idgen"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/model/entity"
"context"
"errors"
"sync"
"time"
"gorm.io/gorm"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
)
// SetDBServiceForTest injects a DBService for tests. Production wiring must use Apply.
func SetDBServiceForTest(s contracts.DBService) {
SetDBService(s)
}
// SetDBService sets the database service singleton.
func SetDBService(s contracts.DBService) {
dbMu.Lock()
defer dbMu.Unlock()
dbSvc = s
}
// GetDB resolves the persistence handle for the current call, preferring an
// explicitly injected *core.Context before falling back to the plugin singleton.
func GetDB(ctx context.Context) *gorm.DB {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
return s.DB(ctx)
}
}
dbMu.RLock()
s := dbSvc
dbMu.RUnlock()
if s != nil {
return s.DB(ctx)
}
return nil
}
// mapNotFound translates GORM's missing-row sentinel into the plugin-level
// consts.ErrRecordNotFound so the service and controller layers stay free of gorm imports.
func mapNotFound(err error) error {
if errors.Is(err, gorm.ErrRecordNotFound) {
return consts.ErrRecordNotFound
}
return err
}
// CreateMessageChannel inserts a channel row.
func CreateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error {
if ch.ID == 0 {
ch.ID = idgen.NextUint64ID()
}
return GetDB(ctx).Create(ch).Error
}
// UpdateMessageChannel saves a channel row.
func UpdateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error {
return GetDB(ctx).Save(ch).Error
}
// GetMessageChannel loads a channel by id.
func GetMessageChannel(ctx context.Context, id uint64) (*entity.MessageChannel, error) {
var ch entity.MessageChannel
if err := GetDB(ctx).Where("id = ?", id).First(&ch).Error; err != nil {
return nil, mapNotFound(err)
}
return &ch, nil
}
// ListMessageChannels returns all channels newest first.
func ListMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) {
var rows []entity.MessageChannel
if err := GetDB(ctx).Order("id DESC").Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// DeleteMessageChannel removes pairings, bindings, then the channel.
func DeleteMessageChannel(ctx context.Context, id uint64) error {
return GetDB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("channel_id = ?", id).Delete(&entity.MessagePairingCode{}).Error; err != nil {
return err
}
if err := tx.Where("channel_id = ?", id).Delete(&entity.MessageBinding{}).Error; err != nil {
return err
}
return tx.Delete(&entity.MessageChannel{}, id).Error
})
}
// CreateMessageBinding inserts a binding.
func CreateMessageBinding(ctx context.Context, b *entity.MessageBinding) error {
if b.ID == 0 {
b.ID = idgen.NextUint64ID()
}
return GetDB(ctx).Create(b).Error
}
// GetBindingByChannelPlatform finds a binding for a platform user on a channel.
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*entity.MessageBinding, error) {
var b entity.MessageBinding
err := GetDB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error
if err != nil {
return nil, mapNotFound(err)
}
return &b, nil
}
// ListBindingsByUser lists bindings for a Wavelet user.
func ListBindingsByUser(ctx context.Context, userID uint64) ([]entity.MessageBinding, error) {
var rows []entity.MessageBinding
if err := GetDB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// ListBindingsByChannel lists bindings on one messaging channel.
func ListBindingsByChannel(ctx context.Context, channelID uint64) ([]entity.MessageBinding, error) {
var rows []entity.MessageBinding
if err := GetDB(ctx).Where("channel_id = ?", channelID).Order("id DESC").Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// GetMessageBinding loads a binding by id.
func GetMessageBinding(ctx context.Context, id uint64) (*entity.MessageBinding, error) {
var b entity.MessageBinding
if err := GetDB(ctx).Where("id = ?", id).First(&b).Error; err != nil {
return nil, mapNotFound(err)
}
return &b, nil
}
// DeleteMessageBinding deletes a binding by id.
func DeleteMessageBinding(ctx context.Context, id uint64) error {
return GetDB(ctx).Delete(&entity.MessageBinding{}, id).Error
}
// UpsertPairingCode reuses an unexpired code for the same channel+platform user.
func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*entity.MessagePairingCode, error) {
var existing entity.MessagePairingCode
err := GetDB(ctx).
Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()).
First(&existing).Error
if err == nil {
return &existing, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
row := &entity.MessagePairingCode{
Code: code,
ChannelID: channelID,
PlatformUserID: platformUserID,
ExpiresAt: expiresAt,
}
if err := GetDB(ctx).Create(row).Error; err != nil {
return nil, err
}
return row, nil
}
// GetPairingCode loads a pairing code by normalized code string.
func GetPairingCode(ctx context.Context, code string) (*entity.MessagePairingCode, error) {
var row entity.MessagePairingCode
if err := GetDB(ctx).Where("code = ?", code).First(&row).Error; err != nil {
return nil, mapNotFound(err)
}
return &row, nil
}
// DeletePairingCode removes a pairing code.
func DeletePairingCode(ctx context.Context, code string) error {
return GetDB(ctx).Where("code = ?", code).Delete(&entity.MessagePairingCode{}).Error
}
// DeleteExpiredPairingCodes removes expired pairing rows.
func DeleteExpiredPairingCodes(ctx context.Context) error {
return GetDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&entity.MessagePairingCode{}).Error
}
// ListEnabledMessageChannels returns enabled channels.
func ListEnabledMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) {
var rows []entity.MessageChannel
if err := GetDB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
@@ -0,0 +1,358 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package dao
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/model/do"
"Wavelet/plugins/domain/msg_gateway/model/entity"
"context"
"errors"
"fmt"
"sync"
"time"
"gorm.io/gorm"
)
const (
activePushChannelCacheTTL = 24 * time.Hour
activePushEventCacheTTL = 24 * time.Hour
)
var (
cacheMu sync.RWMutex
cacheSvc contracts.CacheService
)
// SetCacheService sets the cache service singleton.
func SetCacheService(s contracts.CacheService) {
cacheMu.Lock()
defer cacheMu.Unlock()
cacheSvc = s
}
// GetCache resolves the cache service for the current call.
func GetCache(ctx context.Context) contracts.CacheService {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
return s
}
}
cacheMu.RLock()
s := cacheSvc
cacheMu.RUnlock()
return s
}
// ListPushChannelsRecord returns all push channels ordered by creation time descending.
func ListPushChannelsRecord(ctx context.Context) ([]entity.PushChannel, error) {
var channels []entity.PushChannel
if err := GetDB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
return nil, err
}
return channels, nil
}
// GetPushChannelByIDRecord loads a push channel by primary key.
func GetPushChannelByIDRecord(ctx context.Context, id uint64) (entity.PushChannel, error) {
var channel entity.PushChannel
if err := GetDB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
return entity.PushChannel{}, mapNotFound(err)
}
return channel, nil
}
// GetPushChannelByNameRecord loads a push channel by its unique name.
func GetPushChannelByNameRecord(ctx context.Context, name string) (*entity.PushChannel, error) {
var channel entity.PushChannel
if err := GetDB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
return nil, mapNotFound(err)
}
return &channel, nil
}
// CountPushChannelsByNameRecord returns how many channels share the given name.
func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, error) {
var count int64
if err := GetDB(ctx).Model(&entity.PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CreatePushChannelRecord persists a new channel and invalidates cache.
func CreatePushChannelRecord(ctx context.Context, channel *entity.PushChannel) error {
if err := GetDB(ctx).Create(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
// SavePushChannelRecord updates a channel and invalidates cache.
func SavePushChannelRecord(ctx context.Context, channel *entity.PushChannel) error {
if err := GetDB(ctx).Save(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
// DeletePushChannelRecord removes a channel and invalidates cache.
func DeletePushChannelRecord(ctx context.Context, channel *entity.PushChannel) error {
if err := GetDB(ctx).Delete(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
func getCachedOrQuery[T any](ctx context.Context, cacheKey string, ttl time.Duration, query func(db *gorm.DB, dest *T) error) (*T, error) {
var val T
if cache := GetCache(ctx); cache != nil {
if err := cache.Get(ctx, cacheKey, &val); err == nil {
return &val, nil
}
}
db := GetDB(ctx)
if err := query(db, &val); err != nil {
return nil, err
}
if cache := GetCache(ctx); cache != nil {
_ = cache.Set(ctx, cacheKey, val, ttl)
}
return &val, nil
}
// GetActivePushChannelByName loads an enabled push channel, preferring the cache layer.
func GetActivePushChannelByName(ctx context.Context, name string) (*entity.PushChannel, error) {
channel, err := getCachedOrQuery(ctx, "push:channel:active:"+name, activePushChannelCacheTTL, func(db *gorm.DB, dest *entity.PushChannel) error {
return db.Where("name = ? AND enabled = ?", name, true).First(dest).Error
})
if err != nil {
return nil, mapNotFound(err)
}
return channel, nil
}
// DeleteActivePushChannelCache drops the cached enabled-channel entry.
func DeleteActivePushChannelCache(ctx context.Context, name string) {
if cache := GetCache(ctx); cache != nil {
_ = cache.Delete(ctx, "push:channel:active:"+name)
}
}
// ListPushEventsRecord returns all push events ordered by creation time descending.
func ListPushEventsRecord(ctx context.Context) ([]entity.PushEvent, error) {
var events []entity.PushEvent
if err := GetDB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
return nil, err
}
return events, nil
}
// GetPushEventByIDRecord loads a push event by primary key.
func GetPushEventByIDRecord(ctx context.Context, id uint64) (entity.PushEvent, error) {
var event entity.PushEvent
if err := GetDB(ctx).First(&event, id).Error; err != nil {
return entity.PushEvent{}, mapNotFound(err)
}
return event, nil
}
// GetPushEventByKeyRecord loads a push event by event key.
func GetPushEventByKeyRecord(ctx context.Context, key string) (entity.PushEvent, error) {
var event entity.PushEvent
if err := GetDB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
return entity.PushEvent{}, mapNotFound(err)
}
return event, nil
}
// CountPushEventsByKeyRecord returns how many events use the given event key.
func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error) {
var count int64
if err := GetDB(ctx).Model(&entity.PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CreatePushEventRecord persists a new push event and invalidates cache.
func CreatePushEventRecord(ctx context.Context, event *entity.PushEvent) error {
if err := GetDB(ctx).Create(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// SavePushEventRecord updates a push event and invalidates cache.
func SavePushEventRecord(ctx context.Context, event *entity.PushEvent) error {
if err := GetDB(ctx).Save(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// UpdatePushEventEnabledRecord toggles the enabled flag for a push event.
func UpdatePushEventEnabledRecord(ctx context.Context, event *entity.PushEvent, enabled bool) error {
event.Enabled = enabled
if err := GetDB(ctx).Model(event).Update("enabled", enabled).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// DeletePushEventRecord removes a push event and invalidates cache.
func DeletePushEventRecord(ctx context.Context, event *entity.PushEvent) error {
if err := GetDB(ctx).Delete(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// ListActivePushEventsByTaskTypeRecord returns enabled events bound to a task type.
func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]entity.PushEvent, error) {
var events []entity.PushEvent
if err := GetDB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil {
return nil, err
}
return events, nil
}
// GetActivePushEventByKey loads an enabled push event, preferring the cache layer.
func GetActivePushEventByKey(ctx context.Context, key string) (*entity.PushEvent, error) {
event, err := getCachedOrQuery(ctx, "push:event:active:"+key, activePushEventCacheTTL, func(db *gorm.DB, dest *entity.PushEvent) error {
return db.Where("event_key = ? AND enabled = ?", key, true).First(dest).Error
})
if err != nil {
return nil, mapNotFound(err)
}
return event, nil
}
// DeleteActivePushEventCache drops the cached enabled-event entry.
func DeleteActivePushEventCache(ctx context.Context, key string) {
if cache := GetCache(ctx); cache != nil {
_ = cache.Delete(ctx, "push:event:active:"+key)
}
}
// ListPushHistoriesRecord returns paginated push history records.
func ListPushHistoriesRecord(ctx context.Context, filter do.PushHistoryListFilter) (int64, []entity.PushHistory, error) {
query := GetDB(ctx).Model(&entity.PushHistory{}).Order("created_at DESC")
if filter.EventKey != "" {
query = query.Where("event_key = ?", filter.EventKey)
}
if filter.Status != "" {
query = query.Where("status = ?", filter.Status)
}
var total int64
if err := query.Count(&total).Error; err != nil {
return 0, nil, err
}
var results []entity.PushHistory
offset := (filter.Page - 1) * filter.PageSize
if err := query.Offset(offset).Limit(filter.PageSize).Find(&results).Error; err != nil {
return 0, nil, err
}
return total, results, nil
}
// CreatePushHistoryRecord persists a push history audit record.
func CreatePushHistoryRecord(ctx context.Context, history *entity.PushHistory) error {
return GetDB(ctx).Create(history).Error
}
// PushHistoryQuery returns a scoped query builder for push histories.
func PushHistoryQuery(ctx context.Context) *gorm.DB {
return GetDB(ctx).Model(&entity.PushHistory{})
}
// smtpConfigKeys are the system-config rows backing the built-in email channel.
var smtpConfigKeys = []string{"smtp_host", "smtp_port", "smtp_username", "smtp_password"}
// LoadSMTPConfigRecord reads the SMTP settings in one query.
func LoadSMTPConfigRecord(ctx context.Context) (do.SMTPConfig, error) {
db := GetDB(ctx)
if db == nil {
return do.SMTPConfig{}, errors.New("database not available")
}
var rows []struct {
Key string
Value string
}
if err := db.Table("w_system_configs").
Select("key", "value").
Where("key IN ?", smtpConfigKeys).
Find(&rows).Error; err != nil {
return do.SMTPConfig{}, fmt.Errorf("read smtp system configs: %w", err)
}
var cfg do.SMTPConfig
for _, row := range rows {
switch row.Key {
case "smtp_host":
cfg.Host = row.Value
case "smtp_port":
cfg.Port = row.Value
case "smtp_username":
cfg.Username = row.Value
case "smtp_password":
cfg.Password = row.Value
}
}
return cfg, nil
}
// userLookupColumns allow-lists the columns FindUserByFieldRecord may filter on.
var userLookupColumns = map[string]struct{}{
"id": {},
"username": {},
}
// FindUserByFieldRecord is the user lookup fallback for when the UserService
// contract is not wired yet. field must be one of userLookupColumns.
func FindUserByFieldRecord(ctx context.Context, field string, value any) (*contracts.UserDTO, error) {
if _, ok := userLookupColumns[field]; !ok {
return nil, consts.ErrUnsupportedUserLookupField
}
db := GetDB(ctx)
if db == nil {
return nil, consts.ErrRecordNotFound
}
var user contracts.UserDTO
if err := db.Table("w_users").Where(field+" = ?", value).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
// FindFirstAdminUserRecord is the admin lookup fallback for when the UserService
// contract is not wired yet.
func FindFirstAdminUserRecord(ctx context.Context) (*contracts.UserDTO, error) {
db := GetDB(ctx)
if db == nil {
return nil, consts.ErrRecordNotFound
}
var adminUser contracts.UserDTO
if err := db.Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
return nil, err
}
return &adminUser, nil
}
@@ -0,0 +1,131 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package dao_test
import (
"Wavelet/pkg/testhelper"
"Wavelet/plugins/domain/msg_gateway/consts"
"Wavelet/plugins/domain/msg_gateway/dao"
"context"
"errors"
"testing"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
// stubDBService satisfies contracts.DBService over a test database handle.
type stubDBService struct{ db *gorm.DB }
func (s stubDBService) GORM() *gorm.DB { return s.db }
func (s stubDBService) DB(_ context.Context) *gorm.DB { return s.db }
func (s stubDBService) Named(_ string) *gorm.DB { return s.db }
// TestFindUserByFieldRecordRejectsUnlistedColumns pins the column allow-list. The
// lookup column is interpolated into SQL, so an unlisted name must be refused before
// any query is built rather than trusted because call sites happen to pass literals.
func TestFindUserByFieldRecordRejectsUnlistedColumns(t *testing.T) {
db, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
if err := db.Table("w_users").Create(map[string]any{"id": 77, "username": "seeded"}).Error; err != nil {
t.Fatalf("seed user failed: %v", err)
}
dao.SetDBServiceForTest(stubDBService{db: db})
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
ctx := context.Background()
user, err := dao.FindUserByFieldRecord(ctx, "username", "seeded")
if err != nil {
t.Fatalf("allowlisted lookup by username failed: %v", err)
}
if user.ID != 77 {
t.Errorf("allowlisted lookup returned ID %d, want 77", user.ID)
}
if _, err := dao.FindUserByFieldRecord(ctx, "id", uint64(77)); err != nil {
t.Errorf("allowlisted lookup by id failed: %v", err)
}
cases := []struct {
name string
field string
}{
{"tautology injection", `username = '' OR 1=1 --`},
{"stacked statement", "id; DROP TABLE w_users"},
{"column outside allow-list", "password"},
{"empty field", ""},
}
for _, tc := range cases {
if _, err := dao.FindUserByFieldRecord(ctx, tc.field, "seeded"); !errors.Is(err, consts.ErrUnsupportedUserLookupField) {
t.Errorf("%s: got err %v, want ErrUnsupportedUserLookupField", tc.name, err)
}
}
var remaining int64
if err := db.Table("w_users").Count(&remaining).Error; err != nil || remaining != 1 {
t.Fatalf("w_users damaged by rejected lookups: count=%d err=%v", remaining, err)
}
}
// smtpTestValues are the four system-config rows the built-in email channel reads.
var smtpTestValues = map[string]string{
"smtp_host": "mail.example.test",
"smtp_port": "465",
"smtp_username": "notify@example.test",
"smtp_password": "s3cret-value",
}
// TestLoadSMTPConfigRecordMapsEveryKey guards the single-query rewrite: every field
// must still be filled from its own row.
func TestLoadSMTPConfigRecordMapsEveryKey(t *testing.T) {
db, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
keys := make([]string, 0, len(smtpTestValues))
for key := range smtpTestValues {
keys = append(keys, key)
}
if err := db.Table("w_system_configs").Where("key IN ?", keys).Delete(map[string]any{}).Error; err != nil {
t.Fatalf("clear smtp rows: %v", err)
}
for _, key := range keys {
row := map[string]any{"key": key, "value": smtpTestValues[key], "type": "system"}
if err := db.Table("w_system_configs").Create(row).Error; err != nil {
t.Fatalf("seed %s: %v", key, err)
}
}
dao.SetDBServiceForTest(stubDBService{db: db})
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
cfg, err := dao.LoadSMTPConfigRecord(context.Background())
if err != nil {
t.Fatalf("LoadSMTPConfigRecord: %v", err)
}
if cfg.Host != smtpTestValues["smtp_host"] || cfg.Port != smtpTestValues["smtp_port"] ||
cfg.Username != smtpTestValues["smtp_username"] || cfg.Password != smtpTestValues["smtp_password"] {
t.Errorf("got %+v, want every SMTP field mapped from its own row", cfg)
}
}
// TestLoadSMTPConfigRecordSurfacesReadFailure pins the actual defect: a read that
// fails used to be discarded, returning four blank strings that callers could only
// interpret as "SMTP was never configured", so the notification was dropped silently.
func TestLoadSMTPConfigRecordSurfacesReadFailure(t *testing.T) {
bare, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open bare sqlite: %v", err)
}
dao.SetDBServiceForTest(stubDBService{db: bare})
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
if _, err := dao.LoadSMTPConfigRecord(context.Background()); err == nil {
t.Fatal("LoadSMTPConfigRecord returned nil error although the config table cannot be read")
}
}