refactor(layout): consolidate backend codebase into backend/ package and clean root directory

- Moved cmd/, core/, plugins/, pkg/, downstream/, and main.go into backend/ directory
- Batch updated all Go source files to import github.com/Rain-kl/Wavelet/backend/...
- Updated Makefile, scripts/swagger.sh, architecture guards, and platform skills
- Passed all quality gates (100% tests, 0 lint issues, clean build)
This commit is contained in:
ryan
2026-08-28 12:56:02 +08:00
parent 33b38f8687
commit 43dc97e48c
319 changed files with 912 additions and 1031 deletions
+82
View File
@@ -0,0 +1,82 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
// 管理后台公共错误常量
const (
AdminRequired = "未经授权访问"
TokenAdminRequired = "该访问令牌没有管理员权限,无法访问管理端点" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
InvalidAuthSourceID = "认证源 ID 无效"
InvalidCursorParam = "无效的 cursor 参数"
InvalidTaskExecutionID = "无效的任务执行记录 ID"
)
// 系统配置错误消息常量
const (
SystemConfigNotFound = "系统配置不存在"
ConfigKeyRequired = "配置键不能为空"
ConfigValueRequired = "配置值不能为空"
ConfigKeyExists = "配置键已存在"
protectedConfigKeyMessage = "该配置项由系统任务管理,禁止手动修改"
StorageDriverSwitchRequiresMigration = "存在存量文件,请通过存储迁移任务切换存储引擎"
)
// 模板管理相关错误消息常量
const (
TemplateNotFound = "模板不存在"
TemplateKeyRequired = "模板标识符不能为空"
TemplateNameRequired = "模板名称不能为空"
TemplateContentRequired = "模板内容不能为空"
TemplateKeyExists = "模板标识符已存在"
SystemTemplateCannotDelete = "系统预置模板不可删除"
SystemTemplateCannotModifyKey = "系统预置模板不可修改标识符"
)
// 任务调度相关错误消息常量
const (
InvalidTaskType = "无效的任务类型"
InvalidTimeRange = "无效的时间范围"
TaskDispatchFailed = "任务下发失败"
UserIDRequired = "用户ID必填"
TaskNotFound = "任务执行记录不存在"
TaskNotRetryable = "该任务不支持重试"
TaskNotFailed = "只有失败的任务才能重试"
TaskMaxRetryExceeded = "已达到最大重试次数"
TaskRetryFailed = "任务重试失败"
InvalidCronExpression = "无效的 Cron 表达式"
ScheduleNotFound = "定时任务不存在"
ScheduleSaveFailed = "保存定时任务失败"
ScheduleDeleteFailed = "删除定时任务失败"
)
// 应用更新相关错误消息常量
const (
errInvalidRepository = "上游仓库地址无效"
errReleaseRequestFailed = "获取上游版本失败"
errReleaseResponseInvalid = "上游版本响应无效"
errNoCompatibleRelease = "未找到兼容的 Release"
errNoCompatibleAsset = "未找到当前系统对应的 Release 资产"
errDevelopmentBuild = "开发版本无法执行自动升级"
errAlreadyUpToDate = "当前已是最新版本"
errUpgradeAlreadyRunning = "已有升级任务正在执行"
errAutomaticUpgradeBlocked = "当前平台暂不支持自动替换二进制"
)
// 用户管理(管理员视角)错误消息常量
const (
userNotFound = "用户不存在"
cannotDisable = "不能禁用管理员账号"
cannotDelete = "不能删除管理员账号"
cannotDeleteSelf = "不能删除当前登录账号"
usernameRequired = "用户名不能为空"
emailRequired = "邮箱不能为空"
//nolint:gosec // error message, not hardcoded credentials
passwordTooShort = "密码长度不能少于 8 位"
usernameExists = "用户名已存在"
emailExists = "邮箱已被使用"
cannotRevokeSelfAdmin = "不能取消自身的管理员权限"
updateUserFailed = "更新用户状态失败"
deleteUserFailed = "删除用户失败"
updateUserInfoFailed = "更新用户信息失败"
)
@@ -0,0 +1,157 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/auth"
"github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
"github.com/gin-gonic/gin"
)
// ListAuthSources lists all configured authentication sources.
func ListAuthSources(c *gin.Context) {
var sources []auth.AuthSource
gormDB := database.DB(c.Request.Context())
if err := gormDB.Order("id ASC").Find(&sources).Error; err != nil {
response.AbortInternal(c, "获取认证源列表失败")
return
}
views := make([]auth.AuthSourceView, len(sources))
for i := range sources {
views[i] = auth.AuthSourceView{
ID: sources[i].ID,
Name: sources[i].Name,
Type: sources[i].Type,
DisplayName: sources[i].DisplayName,
IsActive: sources[i].IsActive,
IconURL: sources[i].IconURL,
ClientSecretConfigured: sources[i].ClientSecret != "",
}
}
c.JSON(http.StatusOK, response.OK(views))
}
// CreateAuthSource creates a new authentication source.
func CreateAuthSource(c *gin.Context) {
var source auth.AuthSource
if err := c.ShouldBindJSON(&source); err != nil {
response.AbortBadRequest(c, "无效的参数")
return
}
if err := source.Validate(); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
gormDB := database.DB(c.Request.Context())
if err := gormDB.Create(&source).Error; err != nil {
response.AbortBadRequest(c, "创建认证源失败: "+err.Error())
return
}
source.Sanitize()
c.JSON(http.StatusOK, response.OK(source))
}
// UpdateAuthSource updates an authentication source.
func UpdateAuthSource(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.AbortBadRequest(c, "无效的认证源 ID")
return
}
gormDB := database.DB(c.Request.Context())
var existing auth.AuthSource
if err := gormDB.First(&existing, id).Error; err != nil {
response.AbortNotFound(c, "认证源不存在")
return
}
var req auth.AuthSource
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, "无效的参数")
return
}
existing.DisplayName = req.DisplayName
existing.ClientID = req.ClientID
if req.ClientSecret != "" {
existing.ClientSecret = req.ClientSecret
}
existing.OpenIDDiscoveryURL = req.OpenIDDiscoveryURL
existing.Scopes = req.Scopes
existing.IconURL = req.IconURL
if err := existing.Validate(); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := gormDB.Save(&existing).Error; err != nil {
response.AbortInternal(c, "更新认证源失败")
return
}
existing.Sanitize()
c.JSON(http.StatusOK, response.OK(existing))
}
// ToggleAuthSource toggles the active state of an auth source.
func ToggleAuthSource(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.AbortBadRequest(c, "无效的认证源 ID")
return
}
gormDB := database.DB(c.Request.Context())
var existing auth.AuthSource
if err := gormDB.First(&existing, id).Error; err != nil {
response.AbortNotFound(c, "认证源不存在")
return
}
existing.IsActive = !existing.IsActive
if existing.IsActive {
if err := existing.Validate(); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
}
if err := gormDB.Model(&existing).Update("is_active", existing.IsActive).Error; err != nil {
response.AbortInternal(c, "切换认证源状态失败")
return
}
c.JSON(http.StatusOK, response.OK(gin.H{"is_active": existing.IsActive}))
}
// DeleteAuthSource deletes an authentication source.
func DeleteAuthSource(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.AbortBadRequest(c, "无效的认证源 ID")
return
}
gormDB := database.DB(c.Request.Context())
if err := gormDB.Delete(&auth.AuthSource{}, id).Error; err != nil {
response.AbortInternal(c, "删除认证源失败")
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,103 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"net/http"
"strconv"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/plugins/infra/storage/diskcache"
)
type updateCacheConfigRequest struct {
MaxSizeMB int64 `json:"max_size_mb" binding:"required,min=1"`
TTLMinutes int64 `json:"ttl_minutes" binding:"required,min=0"`
LRUEnabled bool `json:"lru_enabled"`
}
// GetCacheStatus 获取磁盘缓存状态与当前统计数据
// @Summary 获取缓存状态
// @Description 获取当前系统磁盘缓存的使用情况(已占用字节、Key 数量等)与策略配置
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=diskcache.Status} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/cache/status [get]
func GetCacheStatus(c *gin.Context) {
status := diskcache.GetGlobalCache().Status()
c.JSON(http.StatusOK, response.OK(status))
}
// UpdateCacheConfig 更新磁盘缓存策略配置
// @Summary 更新缓存配置
// @Description 更改磁盘缓存最大容量限制、文件生存时间(TTL)以及是否启用 LRU 淘汰淘汰算法,并进行热更新
// @Tags admin
// @Accept json
// @Produce json
// @Param request body updateCacheConfigRequest true "缓存配置请求体"
// @Security SessionCookie
// @Success 200 {object} response.Any "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/cache/config [post]
func UpdateCacheConfig(c *gin.Context) {
var req updateCacheConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
ctx := c.Request.Context()
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
response.AbortInternal(c, err.Error())
return
}
diskcache.GetGlobalCache().ReloadConfig(ctx)
c.JSON(http.StatusOK, response.OKNil())
}
// ClearCache 一键清空所有磁盘缓存数据
// @Summary 清空缓存
// @Description 清除系统磁盘缓存目录中的所有临时文件,并重置缓存容量和 Key 追踪数据
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any "清理成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/cache/clear [post]
func ClearCache(c *gin.Context) {
if err := diskcache.GetGlobalCache().Clear(); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
func saveOrUpdateCacheConfig(ctx context.Context, key, value string) error {
return SaveOrUpdateSystemConfig(ctx, key, value)
}
@@ -0,0 +1,533 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"time"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
mail "github.com/Rain-kl/Wavelet/backend/pkg/mail"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/cap"
cachepkg "github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
"github.com/Rain-kl/Wavelet/backend/plugins/infra/storage/objectstore"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
const maskedConfigValue = "******"
// CreateSystemConfigRequest 创建系统配置请求
type CreateSystemConfigRequest struct {
Key string `json:"key" binding:"required,max=64"`
Value string `json:"value" binding:"required"`
Type string `json:"type" binding:"required,oneof=system business"`
Visibility int `json:"visibility" binding:"oneof=0 1"`
Description string `json:"description" binding:"max=255"`
}
// UpdateSystemConfigRequest 更新系统配置请求
type UpdateSystemConfigRequest struct {
Value string `json:"value" binding:"required"`
Visibility *int `json:"visibility" binding:"omitempty,oneof=0 1"`
Description string `json:"description" binding:"max=255"`
}
// GetPublicConfig 获取公共配置
// @Summary 获取公共配置
// @Description 返回系统配置表中 visibility 为 1 的配置键值集合
// @Tags config
// @Accept json
// @Produce json
// @Success 200 {object} response.Any
// @Router /api/v1/config/public [get]
func GetPublicConfig(c *gin.Context) {
ctx := c.Request.Context()
configs, err := ListVisibleSystemConfigs(ctx)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
resp := make(map[string]string, len(configs))
for _, config := range configs {
resp[config.Key] = config.Value
}
c.JSON(http.StatusOK, response.OK(resp))
}
// GetRobotsTXT 动态生成 robots.txt
// @Summary 获取 robots.txt
// @Description 根据系统配置决定是否允许搜索引擎检索,并返回相应的 robots.txt 文件内容
// @Tags config
// @Produce text/plain
// @Success 200 {string} string "robots.txt 内容"
// @Router /robots.txt [get]
func GetRobotsTXT(c *gin.Context) {
ctx := c.Request.Context()
enabled, err := GetBoolByKey(ctx, ConfigKeySearchEngineIndexingEnabled)
content := "User-Agent: *\nDisallow: /\n"
if err == nil && enabled {
content = "User-Agent: *\nAllow: /\n"
}
c.Data(http.StatusOK, "text/plain; charset=utf-8", []byte(content))
}
// CreateSystemConfig 创建系统配置
// @Summary 创建系统配置
// @Description 创建一条新的系统配置项,配置键不可重复,同时将新配置同步到 Redis,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body CreateSystemConfigRequest true "创建请求参数"
// @Success 200 {object} response.Any{data=string} "创建成功"
// @Failure 400 {object} response.Any "参数错误或配置键已存在"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs [post]
func CreateSystemConfig(c *gin.Context) {
var req CreateSystemConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if isProtectedConfigKey(req.Key) {
response.AbortBadRequest(c, protectedConfigKeyMessage)
return
}
if err := createSystemConfig(c.Request.Context(), req); err != nil {
if err.Error() == ConfigKeyExists {
response.AbortBadRequest(c, ConfigKeyExists)
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// ListSystemConfigs 获取系统配置列表
// @Summary 获取系统配置列表
// @Description 返回所有系统配置列表,支持按配置类型(system/business)过滤,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param type query string false "配置类型(system/business)"
// @Success 200 {object} response.Any{data=[]SystemConfig} "系统配置列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs [get]
func ListSystemConfigs(c *gin.Context) {
configs, err := listSystemConfigs(c.Request.Context(), c.Query("type"))
if err != nil {
response.AbortInternal(c, err.Error())
return
}
for i := range configs {
configs[i].Value = maskSensitiveConfig(configs[i].Key, configs[i].Value)
}
c.JSON(http.StatusOK, response.OK(configs))
}
// GetSystemConfig 获取单个系统配置
// @Summary 获取单个系统配置
// @Description 根据配置键获取对应的系统配置详情,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param key path string true "配置键"
// @Success 200 {object} response.Any{data=SystemConfig} "系统配置详情"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "配置不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs/{key} [get]
func GetSystemConfig(c *gin.Context) {
config, err := getSystemConfig(c.Request.Context(), c.Param("key"))
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, SystemConfigNotFound)
} else {
response.AbortInternal(c, err.Error())
}
return
}
config.Value = maskSensitiveConfig(config.Key, config.Value)
c.JSON(http.StatusOK, response.OK(config))
}
// UpdateSystemConfig 更新系统配置
// @Summary 更新系统配置
// @Description 根据配置键更新对应的配置内容,同时将更新同步到 Redis,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param key path string true "配置键"
// @Param request body UpdateSystemConfigRequest true "更新请求参数"
// @Success 200 {object} response.Any{data=string} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "配置不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs/{key} [put]
func UpdateSystemConfig(c *gin.Context) {
var req UpdateSystemConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
key := c.Param("key")
if isProtectedConfigKey(key) {
response.AbortBadRequest(c, protectedConfigKeyMessage)
return
}
if err := updateSystemConfig(c.Request.Context(), key, req); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, SystemConfigNotFound)
return
}
if isStorageConfigValidationError(err) {
response.AbortBadRequest(c, err.Error())
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
func isProtectedConfigKey(key string) bool {
return key == ConfigKeyLogDatabase || key == ConfigKeyLogDBMigration
}
func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) error {
if isProtectedConfigKey(req.Key) {
return errors.New(protectedConfigKeyMessage)
}
exists, err := SystemConfigExists(ctx, req.Key)
if err != nil {
return err
}
if exists {
return errors.New(ConfigKeyExists)
}
config := SystemConfig{
Key: req.Key,
Value: req.Value,
Type: req.Type,
Visibility: req.Visibility,
Description: req.Description,
}
if err := CreateSystemConfigRecord(ctx, &config); err != nil {
return err
}
invalidateSystemConfigCaches(ctx, req.Key)
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
return nil
}
func listSystemConfigs(ctx context.Context, configType string) ([]SystemConfig, error) {
return ListAdminSystemConfigs(ctx, configType)
}
func getSystemConfig(ctx context.Context, key string) (SystemConfig, error) {
return GetAdminSystemConfigByKey(ctx, key)
}
func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigRequest) error {
if isProtectedConfigKey(key) {
return errors.New(protectedConfigKeyMessage)
}
config, err := GetAdminSystemConfigByKey(ctx, key)
if err != nil {
return err
}
var originalDriver objectstore.Driver
if key == ConfigKeyStorageConfig {
var currentCfg objectstore.Config
if err := json.Unmarshal([]byte(config.Value), &currentCfg); err == nil {
originalDriver = currentCfg.Driver
}
validatedVal, err := validateAndMergeStorageConfig(ctx, req.Value, config.Value)
if err != nil {
return err
}
req.Value = validatedVal
}
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
updates := map[string]any{
"description": req.Description,
}
if req.Visibility != nil {
updates["visibility"] = *req.Visibility
config.Visibility = *req.Visibility
}
if key != ConfigKeySMTPPassword || req.Value != maskedConfigValue {
updates["value"] = req.Value
config.Value = req.Value
}
if err := tx.Model(&config).Updates(updates).Error; err != nil {
return err
}
resolveStorageMigrationTasksOnDirectDriverUpdate(ctx, tx, key, originalDriver, req.Value)
return nil
}); err != nil {
return err
}
invalidateCachesAfterConfigUpdate(ctx, key)
return nil
}
func resolveStorageMigrationTasksOnDirectDriverUpdate(
ctx context.Context,
tx *gorm.DB,
key string,
originalDriver objectstore.Driver,
newValue string,
) {
if key != ConfigKeyStorageConfig || originalDriver == "" {
return
}
var newCfg objectstore.Config
if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil {
return
}
if newCfg.Driver != originalDriver {
return
}
if err := MarkFailedTaskExecutionsSucceededTx(
tx,
"storage:migrate",
"存储配置直接更新,故障迁移任务自动标记为已解决",
time.Now(),
); err != nil {
logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err)
}
}
func invalidateSystemConfigCaches(ctx context.Context, key string) {
if err := InvalidateSystemConfigCache(ctx, key); err != nil {
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
}
if cap.IsRuntimeConfigKey(key) {
cap.InvalidateRuntimeSettings()
}
}
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
invalidateSystemConfigCaches(ctx, key)
if key == ConfigKeyStorageConfig {
if cachepkg.Redis != nil {
_ = cachepkg.Redis.Publish(ctx, "upload:access_cache:invalidate", "reset").Err()
}
objectstore.ResetCache()
objectstore.PublishCacheInvalidation(ctx)
}
if key == ConfigKeyFileAccessWhitelist {
if cachepkg.Redis != nil {
_ = cachepkg.Redis.Publish(ctx, "upload:access_cache:invalidate", "reset").Err()
}
}
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
}
// TestSMTPRequest 测试 SMTP 配置请求
type TestSMTPRequest struct {
SMTPHost string `json:"smtp_host" binding:"required,max=255"`
SMTPPort int `json:"smtp_port" binding:"required"`
SMTPUsername string `json:"smtp_username" binding:"required,max=255"`
SMTPPassword string `json:"smtp_password" binding:"required,max=255"`
To string `json:"to" binding:"required,email"`
}
// TestSMTPResponse 测试 SMTP 配置响应
type TestSMTPResponse struct {
Success bool `json:"success"`
Log string `json:"log"`
Error string `json:"error"`
}
// TestSMTP 测试 SMTP 邮件发送
// @Summary 测试 SMTP 邮件发送
// @Description 使用传入的配置进行 SMTP 邮件发送测试,支持使用 ****** 占位符使用保存的数据库密码
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body TestSMTPRequest true "测试请求参数"
// @Success 200 {object} response.Any{data=TestSMTPResponse} "测试执行完毕"
// @Failure 400 {object} response.Any "参数错误"
// @Router /api/v1/admin/system-configs/smtp/test [post]
func TestSMTP(c *gin.Context) {
var req TestSMTPRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
password := req.SMTPPassword
if password == maskedConfigValue {
if sc, err := GetSystemConfigByKey(c.Request.Context(), ConfigKeySMTPPassword); err == nil {
password = sc.Value
}
}
cfg := mail.Config{
Host: req.SMTPHost,
Port: req.SMTPPort,
Username: req.SMTPUsername,
Password: password,
}
subject := "Wavelet SMTP Test Mail"
body := `<h3>SMTP Mail Connection Test</h3>
<p>If you received this message, your SMTP configuration is correct and mail sending is working properly.</p>
<p>Sent from Wavelet.</p>`
logs, err := mail.SendMailWithLog(c.Request.Context(), cfg, req.To, subject, body)
resp := TestSMTPResponse{
Success: err == nil,
Log: logs,
}
if err != nil {
resp.Error = err.Error()
}
c.JSON(http.StatusOK, response.OK(resp))
}
func isStorageConfigValidationError(err error) bool {
msg := err.Error()
return msg == StorageDriverSwitchRequiresMigration ||
strings.HasPrefix(msg, "解析") ||
strings.HasPrefix(msg, "验证") ||
strings.HasPrefix(msg, "初始化测试") ||
strings.HasPrefix(msg, "存储连通性") ||
strings.HasPrefix(msg, "序列化") ||
strings.HasPrefix(msg, "检查存量文件")
}
func maskSensitiveConfig(key, value string) string {
if value == "" {
return value
}
switch key {
case ConfigKeySMTPPassword:
return maskedConfigValue
case ConfigKeyStorageConfig:
var cfg objectstore.Config
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
masked := objectstore.MaskSecrets(cfg)
if val, err := json.Marshal(masked); err == nil {
return string(val)
}
}
}
return value
}
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
// and tests connectivity of the new storage configuration.
func validateAndMergeStorageConfig(ctx context.Context, value string, currentConfig string) (string, error) {
var currentCfg objectstore.Config
if err := json.Unmarshal([]byte(currentConfig), &currentCfg); err != nil {
return "", fmt.Errorf("解析当前存储配置失败: %w", err)
}
var newCfg objectstore.Config
if err := json.Unmarshal([]byte(value), &newCfg); err != nil {
return "", fmt.Errorf("解析目标存储配置失败: %w", err)
}
// 合并被掩码屏蔽的敏感信息,获取完整的真实配置
targetCfg := objectstore.MergeMaskedSecrets(newCfg, currentCfg)
if err := validateMergedStorageConfig(ctx, currentCfg, newCfg, targetCfg); err != nil {
return "", err
}
// 序列化为最终保存的真实明文配置,防止保存屏蔽的 ****** 字符
unmaskedVal, err := json.Marshal(targetCfg)
if err != nil {
return "", fmt.Errorf("序列化存储配置失败: %w", err)
}
return string(unmaskedVal), nil
}
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg objectstore.Config) error {
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
var uploadCount int64
if err := db.DB(ctx).Table("w_uploads").
Where("status != ?", "deleted").
Count(&uploadCount).Error; err != nil {
return fmt.Errorf("检查存量文件失败: %w", err)
}
if uploadCount > 0 {
return errors.New(StorageDriverSwitchRequiresMigration)
}
if err := validateDriverConfig(targetCfg, newCfg.Driver); err != nil {
return fmt.Errorf("验证目标存储配置参数失败: %w", err)
}
pendingCfg := targetCfg
pendingCfg.Driver = newCfg.Driver
return testStorageBackend(ctx, pendingCfg, newCfg.Driver)
}
if err := objectstore.ValidateConfig(targetCfg); err != nil {
return fmt.Errorf("验证存储配置参数失败: %w", err)
}
return testStorageBackend(ctx, targetCfg, targetCfg.Driver)
}
func validateDriverConfig(cfg objectstore.Config, driver objectstore.Driver) error {
cfg.Driver = driver
return objectstore.ValidateConfig(cfg)
}
func testStorageBackend(ctx context.Context, cfg objectstore.Config, driver objectstore.Driver) error {
testBackend, err := objectstore.NewBackend(ctx, cfg, driver)
if err != nil {
return fmt.Errorf("初始化测试存储实例失败: %w", err)
}
if err := testBackend.Test(ctx); err != nil {
return fmt.Errorf("存储连通性测试失败: %w", err)
}
return nil
}
+646
View File
@@ -0,0 +1,646 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"database/sql"
"fmt"
"log"
"math"
"net/http"
"os"
"os/exec"
"strings"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/backend/pkg/config"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
)
const (
binaryKB = 0
binaryMB = 1
binaryGB = 2
valueThreshold = 10
maxStringLength = 200
)
// DBOverviewResponse 数据库运行概览响应结构体
type DBOverviewResponse struct {
Type string `json:"type"`
Version string `json:"version"`
Name string `json:"name"`
Size string `json:"size"`
TableCount int64 `json:"table_count"`
Connections int64 `json:"connections"`
}
// GetTableDataRequest 分页拉取表数据请求结构体
type GetTableDataRequest struct {
Table string `form:"table" binding:"required"`
Page int `form:"page,default=1"`
PageSize int `form:"pageSize,default=10"`
}
// TableDataResponse 动态数据表响应结构体
type TableDataResponse struct {
Columns []string `json:"columns"`
Total int64 `json:"total"`
Results []map[string]interface{} `json:"results"`
}
// ExecuteSQLRequest 执行自定义 SQL 请求结构体
type ExecuteSQLRequest struct {
SQL string `json:"sql" binding:"required"`
}
// ExecuteSQLResponse 执行自定义 SQL 响应结构体
type ExecuteSQLResponse struct {
Type string `json:"type"` // "select" 或 "exec"
Columns []string `json:"columns,omitempty"`
Results []map[string]interface{} `json:"results,omitempty"`
AffectedRows int64 `json:"affected_rows"`
ExecutionTimeMs int64 `json:"execution_time_ms"`
}
// DatabaseInfoResponse 数据库信息响应结构体
type DatabaseInfoResponse struct {
Type string `json:"type"`
Name string `json:"name"`
Version string `json:"version"`
}
func formatBytes(bytes uint64) string {
const unit = 1024
if bytes < unit {
return fmt.Sprintf("%d B", bytes)
}
div, exp := int64(unit), 0
for n := bytes / unit; n >= unit; n /= unit {
div *= unit
exp++
}
value := float64(bytes) / float64(div)
var suffix string
switch exp {
case binaryKB:
suffix = "KiB"
case binaryMB:
suffix = "MiB"
case binaryGB:
suffix = "GiB"
default:
suffix = "TiB"
}
if value == math.Trunc(value) {
if value >= valueThreshold {
return fmt.Sprintf("%.0f %s", value, suffix)
}
return fmt.Sprintf("%.1f %s", value, suffix)
}
return fmt.Sprintf("%.1f %s", value, suffix)
}
const defaultSQLiteDBPath = "./data/wavelet.db"
func getSQLiteOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
name := config.Config.Database.SQLitePath
if name == "" {
name = defaultSQLiteDBPath
}
var version string
var ver string
if err := gormDB.Raw("SELECT sqlite_version()").Scan(&ver).Error; err == nil {
version = "SQLite " + ver
} else {
version = "SQLite"
}
var sizeStr string
if fi, err := os.Stat(name); err == nil {
size := fi.Size()
if size < 0 {
size = 0
}
sizeStr = formatBytes(uint64(size))
} else {
sizeStr = "0 B"
}
var tableCount int64
if err := gormDB.Raw("SELECT count(*) FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%'").Scan(&tableCount).Error; err != nil {
tableCount = 0
}
var connCount int64
if sqlDB, err := gormDB.DB(); err == nil {
connCount = int64(sqlDB.Stats().OpenConnections)
} else {
connCount = 1
}
return DBOverviewResponse{
Type: "sqlite",
Version: version,
Name: name,
Size: sizeStr,
TableCount: tableCount,
Connections: connCount,
}, nil
}
func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
name := config.Config.Database.Database
var version string
var ver string
if err := gormDB.Raw("SELECT version()").Scan(&ver).Error; err == nil {
version = ver
} else {
version = "PostgreSQL"
}
var sizeStr string
var sizeBytes sql.NullInt64
if err := gormDB.Raw("SELECT pg_database_size(current_database())").Scan(&sizeBytes).Error; err == nil && sizeBytes.Valid {
size := sizeBytes.Int64
if size < 0 {
size = 0
}
sizeStr = formatBytes(uint64(size))
} else {
sizeStr = "0 B"
}
var tableCount int64
if err := gormDB.Raw("SELECT count(*) FROM information_schema.tables WHERE table_schema = current_schema()").Scan(&tableCount).Error; err != nil {
tableCount = 0
}
var connCount int64
var pgc sql.NullInt64
if err := gormDB.Raw("SELECT count(*) FROM pg_stat_activity WHERE datname = current_database()").Scan(&pgc).Error; err == nil && pgc.Valid {
connCount = pgc.Int64
} else {
if sqlDB, err := gormDB.DB(); err == nil {
connCount = int64(sqlDB.Stats().OpenConnections)
} else {
connCount = 1
}
}
return DBOverviewResponse{
Type: "postgres",
Version: version,
Name: name,
Size: sizeStr,
TableCount: tableCount,
Connections: connCount,
}, nil
}
// GetDBOverview 获取数据库运行概览
// @Summary 获取数据库运行概览
// @Description 获取数据库类型、版本、名称、文件大小、表数量及当前连接数,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=DBOverviewResponse} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/db-manage/overview [get]
func GetDBOverview(c *gin.Context) {
gormDB := db.DB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
}
var overview DBOverviewResponse
var err error
if !config.Config.Database.Enabled {
overview, err = getSQLiteOverview(gormDB)
} else {
overview, err = getPostgresOverview(gormDB)
}
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(overview))
}
// ListDBTables 获取数据库所有表名
// @Summary 获取数据库所有表名
// @Description 返回当前数据库的所有用户自定义表名称列表,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]string} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/db-manage/tables [get]
func ListDBTables(c *gin.Context) {
gormDB := db.DB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
}
var tables []string
var err error
if !config.Config.Database.Enabled {
err = gormDB.Raw("SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%' ORDER BY name").Scan(&tables).Error
} else {
err = gormDB.Raw("SELECT table_name FROM information_schema.tables WHERE table_schema = current_schema() ORDER BY table_name").Scan(&tables).Error
}
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(tables))
}
// GetDBTableData 获取数据表 data
func GetDBTableData(c *gin.Context) {
var req GetTableDataRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
gormDB := db.DB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
}
// 安全转义表名并拼接
quotedTable := `"` + strings.ReplaceAll(req.Table, `"`, `""`) + `"`
var total int64
if err := gormDB.Raw("SELECT count(*) FROM " + quotedTable).Scan(&total).Error; err != nil {
response.AbortBadRequest(c, err.Error())
return
}
offset := (req.Page - 1) * req.PageSize
if offset < 0 {
offset = 0
}
limit := req.PageSize
if limit <= 0 {
limit = 10
}
rows, err := gormDB.Raw("SELECT * FROM "+quotedTable+" LIMIT ? OFFSET ?", limit, offset).Rows()
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
defer func() {
_ = rows.Close()
}()
cols, err := rows.Columns()
if err != nil {
response.AbortInternal(c, err.Error())
return
}
results, err := scanTableRows(rows, cols)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(TableDataResponse{
Columns: cols,
Total: total,
Results: results,
}))
}
func scanTableRows(rows *sql.Rows, cols []string) ([]map[string]interface{}, error) {
results := make([]map[string]interface{}, 0)
for rows.Next() {
columns := make([]interface{}, len(cols))
columnPointers := make([]interface{}, len(cols))
for i := range columns {
columnPointers[i] = &columns[i]
}
if err := rows.Scan(columnPointers...); err != nil {
return nil, err
}
rowMap := make(map[string]interface{})
for i, colName := range cols {
val := columns[i]
if b, ok := val.([]byte); ok {
strVal := string(b)
runes := []rune(strVal)
if len(runes) > maxStringLength {
strVal = string(runes[:maxStringLength]) + "..."
}
rowMap[colName] = strVal
} else if str, ok := val.(string); ok {
runes := []rune(str)
if len(runes) > maxStringLength {
str = string(runes[:maxStringLength]) + "..."
}
rowMap[colName] = str
} else {
rowMap[colName] = val
}
}
results = append(results, rowMap)
}
return results, nil
}
func executeSQLQuery(gormDB *gorm.DB, sqlStr string, startTime time.Time) (ExecuteSQLResponse, error) {
rows, err := gormDB.Raw(sqlStr).Rows()
if err != nil {
return ExecuteSQLResponse{}, err
}
defer func() {
_ = rows.Close()
}()
cols, err := rows.Columns()
if err != nil {
return ExecuteSQLResponse{}, err
}
results := make([]map[string]interface{}, 0)
for rows.Next() {
columns := make([]interface{}, len(cols))
columnPointers := make([]interface{}, len(cols))
for i := range columns {
columnPointers[i] = &columns[i]
}
if err := rows.Scan(columnPointers...); err != nil {
return ExecuteSQLResponse{}, err
}
rowMap := make(map[string]interface{})
for i, colName := range cols {
val := columns[i]
if b, ok := val.([]byte); ok {
rowMap[colName] = string(b)
} else {
rowMap[colName] = val
}
}
results = append(results, rowMap)
}
executionTime := time.Since(startTime).Milliseconds()
return ExecuteSQLResponse{
Type: "select",
Columns: cols,
Results: results,
AffectedRows: int64(len(results)),
ExecutionTimeMs: executionTime,
}, nil
}
func executeSQLMutation(gormDB *gorm.DB, sqlStr string, startTime time.Time) (ExecuteSQLResponse, error) {
tx := gormDB.Exec(sqlStr)
if tx.Error != nil {
return ExecuteSQLResponse{}, tx.Error
}
executionTime := time.Since(startTime).Milliseconds()
return ExecuteSQLResponse{
Type: "exec",
AffectedRows: tx.RowsAffected,
ExecutionTimeMs: executionTime,
}, nil
}
// ExecuteSQL 执行 SQL 查询
// @Summary 执行 SQL 查询
// @Description 在当前数据库中执行任意自定义 SQL,如果是查询语句将返回格式化后的列与数据集,否则返回受影响行数,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body ExecuteSQLRequest true "SQL 请求参数"
// @Success 200 {object} response.Any{data=ExecuteSQLResponse} "执行完毕"
// @Failure 400 {object} response.Any "SQL 语句错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/db-manage/query [post]
func ExecuteSQL(c *gin.Context) {
var req ExecuteSQLRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
gormDB := db.DB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
}
trimmedSQL := strings.TrimSpace(req.SQL)
if trimmedSQL == "" {
response.AbortBadRequest(c, "SQL 语句不能为空")
return
}
startTime := time.Now()
isQuery := false
lowerSQL := strings.ToLower(trimmedSQL)
queryKeywords := []string{"select", "show", "explain", "describe", "pragma"}
for _, kw := range queryKeywords {
if strings.HasPrefix(lowerSQL, kw) {
isQuery = true
break
}
}
var resp ExecuteSQLResponse
var err error
if isQuery {
resp, err = executeSQLQuery(gormDB, trimmedSQL, startTime)
} else {
resp, err = executeSQLMutation(gormDB, trimmedSQL, startTime)
}
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
func getSQLiteInfo(ctx context.Context) DatabaseInfoResponse {
info := DatabaseInfoResponse{
Type: "sqlite",
Name: config.Config.Database.SQLitePath,
Version: "SQLite",
}
if info.Name == "" {
info.Name = "./data/wavelet.db"
}
gormDB := db.DB(ctx)
if gormDB == nil {
return info
}
var ver string
if err := gormDB.Raw("SELECT sqlite_version()").Scan(&ver).Error; err == nil && ver != "" {
info.Version = "SQLite " + ver
}
return info
}
func getPostgresInfo(ctx context.Context) DatabaseInfoResponse {
info := DatabaseInfoResponse{
Type: "postgres",
Name: config.Config.Database.Database,
Version: "PostgreSQL",
}
gormDB := db.DB(ctx)
if gormDB == nil {
return info
}
var ver string
if err := gormDB.Raw("SELECT version()").Scan(&ver).Error; err == nil && ver != "" {
info.Version = ver
}
return info
}
// GetDatabaseInfo 获取当前数据库类型及版本信息
// @Summary 获取数据库信息
// @Description 返回当前使用的数据库类型(sqlite/postgres)、名称/路径及版本字符串,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=DatabaseInfoResponse} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/db-info [get]
func GetDatabaseInfo(c *gin.Context) {
var info DatabaseInfoResponse
if !config.Config.Database.Enabled {
info = getSQLiteInfo(c.Request.Context())
} else {
info = getPostgresInfo(c.Request.Context())
}
c.JSON(http.StatusOK, response.OK(info))
}
// ExportDatabase 导出数据库
// @Summary 导出数据库
// @Description SQLite 时直接下载 .db 文件;PostgreSQL 时执行 pg_dump 并流式下载 .sql 文件,需要管理员权限
// @Tags admin
// @Produce application/octet-stream
// @Security SessionCookie
// @Success 200 {file} binary "数据库文件"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "导出失败"
// @Router /api/v1/admin/db-export [get]
func ExportDatabase(c *gin.Context) {
if !config.Config.Database.Enabled {
exportSQLite(c)
} else {
exportPostgres(c)
}
}
func exportSQLite(c *gin.Context) {
path := config.Config.Database.SQLitePath
if path == "" {
path = defaultSQLiteDBPath
}
//nolint:gosec // export db file path is trusted
f, err := os.Open(path)
if err != nil {
response.AbortInternal(c, "无法打开数据库文件: "+err.Error())
return
}
defer func() {
if closeErr := f.Close(); closeErr != nil {
_ = closeErr
}
}()
fi, err := f.Stat()
if err != nil {
response.AbortInternal(c, "无法读取数据库文件信息: "+err.Error())
return
}
c.Header("Content-Disposition", `attachment; filename="wavelet.db"`)
c.Header("Content-Type", "application/octet-stream")
c.Header("Content-Length", fmt.Sprintf("%d", fi.Size()))
c.Status(http.StatusOK)
http.ServeContent(c.Writer, c.Request, "wavelet.db", fi.ModTime(), f)
}
func exportPostgres(c *gin.Context) {
dbCfg := config.Config.Database
pgDumpPath, err := exec.LookPath("pg_dump")
if err != nil {
response.AbortInternal(c, "pg_dump 不可用,请确保服务器已安装 PostgreSQL 客户端工具")
return
}
args := []string{
"--no-password",
"-h", dbCfg.Host,
"-p", fmt.Sprintf("%d", dbCfg.Port),
"-U", dbCfg.Username,
dbCfg.Database,
}
//nolint:gosec // pg_dump args are constructed from validated db config
cmd := exec.CommandContext(c.Request.Context(), pgDumpPath, args...)
if dbCfg.Password != "" {
cmd.Env = append(os.Environ(), "PGPASSWORD="+dbCfg.Password)
} else {
cmd.Env = os.Environ()
}
fileName := fmt.Sprintf("wavelet_%s.sql", time.Now().Format("20060102_150405"))
c.Header("Content-Disposition", `attachment; filename="`+fileName+`"`)
c.Header("Content-Type", "application/octet-stream")
c.Status(http.StatusOK)
cmd.Stdout = c.Writer
cmd.Stderr = nil
if err := cmd.Run(); err != nil {
log.Printf("[db-export] pg_dump failed: %v\n", err)
}
}
@@ -0,0 +1,685 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/url"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/backend/pkg/config"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/risk_control"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/risk_control/logstore"
"github.com/Rain-kl/Wavelet/backend/plugins/drivers/driver_asynq_worker"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
)
const (
defaultLimit = 200
maxLimit = 500
maxPageSize = 100
hoursInDay = 24
analyticsDays = 7
topActiveLimit = 10
)
// logsResponse 历史日志查询响应
type logsResponse struct {
Lines []logger.LogEntry `json:"lines"`
HasMore bool `json:"has_more"`
NextCursor int `json:"next_cursor"` // 用于加载更早日志的 cursor
}
// GetLogs 获取历史日志
// @Summary 获取系统日志
// @Description 分页获取系统历史日志,cursor=0 获取最新日志,cursor>0 获取更早日志
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param cursor query int false "日志游标,0=获取最新" default(0)
// @Param limit query int false "每页条数" default(200)
// @Success 200 {object} response.Any{data=logsResponse} "日志列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/logs [get]
func GetLogs(c *gin.Context) {
cursorStr := c.DefaultQuery("cursor", "0")
limitStr := c.DefaultQuery("limit", "200")
var cursor, limit int
if _, err := parsePositiveInt(cursorStr, &cursor); err != nil {
response.AbortWithError(c, http.StatusBadRequest, InvalidCursorParam)
return
}
if _, err := parsePositiveInt(limitStr, &limit); err != nil || limit <= 0 {
limit = defaultLimit
}
if limit > maxLimit {
limit = maxLimit
}
entries, hasMore := logger.GlobalRingBuffer.Query(cursor, limit)
resp := logsResponse{
Lines: entries,
HasMore: hasMore,
}
if len(entries) > 0 {
resp.NextCursor = entries[0].Index
}
c.JSON(http.StatusOK, response.OK(resp))
}
// wsMessage WebSocket 消息格式
type wsMessage struct {
Type string `json:"type"` // "log" | "error"
Data json.RawMessage `json:"data"`
}
// HandleLogWebSocket WebSocket 端点,实时推送系统日志
// @Summary 系统日志实时推送
// @Description 通过 WebSocket 实时推送系统日志,需要管理员权限
// @Tags admin
// @Router /api/v1/admin/logs/ws [get]
func HandleLogWebSocket(c *gin.Context) {
upgrader := getUpgrader()
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
return
}
defer func() { _ = conn.Close() }()
ch := logger.GlobalRingBuffer.Subscribe()
defer logger.GlobalRingBuffer.Unsubscribe(ch)
done := make(chan struct{})
util.Go(func() {
defer close(done)
for {
_, _, err := conn.ReadMessage()
if err != nil {
return
}
}
})
for {
select {
case <-done:
return
case entry, ok := <-ch:
if !ok {
return
}
data, _ := json.Marshal(entry)
msg := wsMessage{Type: "log", Data: data}
payload, _ := json.Marshal(msg)
if err := conn.WriteMessage(1, payload); err != nil {
return
}
}
}
}
// accessLogItem 访问日志单条数据
type accessLogItem struct {
ID uint64 `json:"id,string"`
UserID uint64 `json:"user_id,string"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Path string `json:"path"`
Method string `json:"method"`
IP string `json:"ip"`
UserAgent string `json:"user_agent"`
Headers string `json:"headers"`
Status int32 `json:"status"`
Latency int64 `json:"latency"`
CreatedAt string `json:"created_at"`
}
// accessLogsResponse 访问日志查询响应
type accessLogsResponse struct {
Total uint64 `json:"total"`
List []accessLogItem `json:"list"`
}
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (logstore.AccessLogFilter, error) {
filter := logstore.AccessLogFilter{}
username := c.Query("username")
if username != "" {
var userIDs []uint64
if err := db.DB(ctx).Table("w_users").
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
Pluck("id", &userIDs).Error; err != nil {
return filter, fmt.Errorf("查询用户信息失败: %w", err)
}
filter.UserIDs = userIDs
}
if path := c.Query("path"); path != "" {
filter.Path = path
}
if startTime := c.Query("start_time"); startTime != "" {
if t, err := parseAccessLogTime(startTime); err == nil {
filter.StartTime = &t
}
}
if endTime := c.Query("end_time"); endTime != "" {
if t, err := parseAccessLogTime(endTime); err == nil {
filter.EndTime = &t
}
}
return filter, nil
}
func parseAccessLogTime(value string) (time.Time, error) {
if t, err := time.Parse(time.RFC3339, value); err == nil {
return t, nil
}
return time.Parse("2006-01-02 15:04:05", value)
}
func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
if len(list) == 0 {
return
}
userIDs := make([]uint64, 0, len(list))
seen := make(map[uint64]struct{}, len(list))
for _, item := range list {
if _, ok := seen[item.UserID]; ok {
continue
}
seen[item.UserID] = struct{}{}
userIDs = append(userIDs, item.UserID)
}
userMap := make(map[uint64]struct{ Username, Nickname string })
var users []struct {
ID uint64
Username string
Nickname string
}
if err := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err == nil {
for _, u := range users {
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
}
}
for i := range list {
if info, ok := userMap[list[i].UserID]; ok {
list[i].Username = info.Username
list[i].Nickname = info.Nickname
}
}
}
// GetAccessLogs 获取 ClickHouse 异步采集的访问日志
// @Summary 获取用户访问日志
// @Description 分页并按照用户、接口路径、时间范围等维度检索用户访问日志列表(需要管理员权限)
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param page query int false "页码" default(1)
// @Param page_size query int false "每页条数" default(20)
// @Param username query string false "用户名模糊搜索"
// @Param path query string false "接口路径模糊搜索"
// @Param start_time query string false "起始时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)"
// @Param end_time query string false "结束时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)"
// @Success 200 {object} response.Any{data=accessLogsResponse} "访问日志列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/logs/access [get]
func GetAccessLogs(c *gin.Context) {
ctx := c.Request.Context()
store, err := logstore.Active(ctx)
if err != nil {
response.AbortInternal(c, "日志存储初始化失败")
return
}
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
if page < 1 {
page = 1
}
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
if pageSize < 1 {
pageSize = 20
}
if pageSize > maxPageSize {
pageSize = maxPageSize
}
filter, err := buildAccessLogFilter(ctx, c)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
return
}
if filter.UserIDs != nil && len(filter.UserIDs) == 0 {
c.JSON(http.StatusOK, response.OK(accessLogsResponse{Total: 0, List: []accessLogItem{}}))
return
}
logs, total, err := store.UserAccessLogs.List(ctx, filter, page, pageSize)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
return
}
if total == 0 {
c.JSON(http.StatusOK, response.OK(accessLogsResponse{Total: 0, List: []accessLogItem{}}))
return
}
list := make([]accessLogItem, len(logs))
for i, logItem := range logs {
list[i] = accessLogItem{
ID: logItem.ID,
UserID: logItem.UserID,
Path: logItem.Path,
Method: logItem.Method,
IP: logItem.IP,
UserAgent: logItem.UserAgent,
Headers: logItem.Headers,
Status: logItem.Status,
Latency: logItem.Latency,
CreatedAt: logItem.CreatedAt.Format(time.RFC3339),
}
}
enrichAccessLogsWithUsers(ctx, list)
c.JSON(http.StatusOK, response.OK(accessLogsResponse{
Total: total,
List: list,
}))
}
// trendItem 趋势图数据点
type trendItem struct {
Date string `json:"date"`
Count uint64 `json:"count"`
}
// browserItem 浏览器占比排行
type browserItem struct {
Browser string `json:"browser"`
Count uint64 `json:"count"`
}
// topUserItem 活跃用户数据
type topUserItem struct {
UserID uint64 `json:"user_id,string"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Count uint64 `json:"count"`
}
// logsAnalyticsResponse 访问日志数据分析结果
type logsAnalyticsResponse struct {
Trend []trendItem `json:"trend"`
Browsers []browserItem `json:"browsers"`
TopUsers []topUserItem `json:"top_users"`
}
// GetLogsAnalytics 获取 ClickHouse 访问日志图表聚合指标
// @Summary 获取访问日志分析数据
// @Description 聚合统计最近 7 天的每日访问趋势、浏览器分布以及前 10 名最活跃用户排行(需要管理员权限)
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=logsAnalyticsResponse} "分析统计数据"
// @Failure 500 {object} response.Any "内部错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/logs/analytics [get]
func GetLogsAnalytics(c *gin.Context) {
ctx := c.Request.Context()
store, err := logstore.Active(ctx)
if err != nil {
response.AbortInternal(c, "日志存储初始化失败")
return
}
startTime := time.Now().AddDate(0, 0, -(analyticsDays - 1)).Truncate(hoursInDay * time.Hour)
trendPoints, err := store.UserAccessLogs.GetDailyTrend(ctx, analyticsDays)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, "查询访问趋势失败: "+err.Error())
return
}
trendList := make([]trendItem, len(trendPoints))
for i, point := range trendPoints {
trendList[i] = trendItem{
Date: point.Date,
Count: point.Count,
}
}
browserPoints, err := store.UserAccessLogs.GetBrowserDistribution(ctx, startTime)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, "查询浏览器分布失败: "+err.Error())
return
}
browserList := make([]browserItem, len(browserPoints))
for i, point := range browserPoints {
browserList[i] = browserItem{
Browser: point.Browser,
Count: point.Count,
}
}
topUserPoints, err := store.UserAccessLogs.GetTopActiveUsers(ctx, startTime, topActiveLimit)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, "查询活跃用户失败: "+err.Error())
return
}
topUsers := make([]topUserItem, len(topUserPoints))
userIDs := make([]uint64, len(topUserPoints))
for i, point := range topUserPoints {
topUsers[i] = topUserItem{
UserID: point.UserID,
Count: point.Count,
}
userIDs[i] = point.UserID
}
if len(userIDs) > 0 {
userProfileMap := make(map[uint64]struct {
Username string
Nickname string
})
var users []struct {
ID uint64
Username string
Nickname string
}
if errProfile := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; errProfile == nil {
for _, u := range users {
userProfileMap[u.ID] = struct {
Username string
Nickname string
}{
Username: u.Username,
Nickname: u.Nickname,
}
}
}
for i := range topUsers {
if profile, ok := userProfileMap[topUsers[i].UserID]; ok {
topUsers[i].Username = profile.Username
topUsers[i].Nickname = profile.Nickname
}
}
}
c.JSON(http.StatusOK, response.OK(logsAnalyticsResponse{
Trend: trendList,
Browsers: browserList,
TopUsers: topUsers,
}))
}
func getUpgrader() *websocket.Upgrader {
return &websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool {
origin := r.Header.Get("Origin")
if origin == "" {
return true
}
// 1. 同源检查 (Same-origin check)
u, err := url.Parse(origin)
if err == nil && strings.EqualFold(u.Host, r.Host) {
return true
}
// 2. 检查配置的允许跨域 Origin (Check allowed origins in system config)
ctx := r.Context()
if sc, err := GetSystemConfigByKey(ctx, ConfigKeyServerAddress); err == nil && sc.Value != "" {
originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/")
allowedOrigins := strings.Split(sc.Value, ",")
for _, allowed := range allowedOrigins {
allowed = strings.TrimRight(strings.TrimSpace(allowed), "/")
if allowed != "" && strings.EqualFold(allowed, originToCheck) {
return true
}
}
}
return false
},
}
}
func parsePositiveInt(s string, result *int) (bool, error) {
if s == "" {
*result = 0
return true, nil
}
n, err := strconv.Atoi(s)
if err != nil || n < 0 {
return false, err
}
*result = n
return true, nil
}
// Log DB Switch Task
const (
// LogDBSwitchTask 切换日志数据库任务标识。
LogDBSwitchTask = "logs:db_switch"
// TaskTypeLogDBSwitch 管理端任务类型。
TaskTypeLogDBSwitch = "logs_db_switch"
copyBatchSize = 1000
targetPostgres = "postgres"
targetSQLite = "sqlite"
targetClickHouse = "clickhouse"
)
// LogDBSwitchMeta 描述切换日志数据库任务。
var LogDBSwitchMeta = driver_asynq_worker.TaskMeta{
Type: TaskTypeLogDBSwitch,
AsynqTask: LogDBSwitchTask,
Name: "切换日志数据库",
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
SupportsTime: false,
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
Queue: driver_asynq_worker.QueueDefault,
Retryable: true,
Params: []driver_asynq_worker.TaskParam{
{Name: "target", Label: "目标日志库", Type: "string", Required: true,
Placeholder: "postgres|sqlite|clickhouse", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse"},
},
}
type logDBSwitchPayload struct {
Target string `json:"target"`
}
// LogDBSwitchHandler 切换日志数据库任务处理器。
type LogDBSwitchHandler struct{}
// ValidatePayload 校验并规范化参数。
func (h *LogDBSwitchHandler) ValidatePayload(payload []byte) ([]byte, error) {
var p logDBSwitchPayload
if err := json.Unmarshal(payload, &p); err != nil {
return nil, fmt.Errorf("参数解析失败: %w", err)
}
p.Target = normalizeTarget(p.Target)
if !validTarget(p.Target) {
return nil, fmt.Errorf("目标日志库不合法: %s", p.Target)
}
out, err := json.Marshal(p)
if err != nil {
return nil, err
}
return out, nil
}
func normalizeTarget(v string) string {
switch v {
case targetPostgres, "postgresql":
return targetPostgres
case targetSQLite, "sqlite3":
return targetSQLite
case targetClickHouse, "ch":
return targetClickHouse
}
return v
}
func validTarget(v string) bool {
return v == targetPostgres || v == targetSQLite || v == targetClickHouse
}
// Execute 执行迁移。
func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
var p logDBSwitchPayload
if err := json.Unmarshal(payload, &p); err != nil {
return nil, fmt.Errorf("参数解析失败: %w", err)
}
p.Target = normalizeTarget(p.Target)
if err := validateSwitch(ctx, p.Target); err != nil {
return nil, err
}
source, err := currentLogDatabase(ctx)
if err != nil {
driver_asynq_worker.AppendLog(ctx, "读取日志主库失败: %v", err)
return nil, err
}
driver_asynq_worker.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
if err := setMigrationFlag(ctx, "migrating"); err != nil {
return nil, err
}
defer func() {
if err := setMigrationFlag(ctx, ""); err != nil {
logger.ErrorF(ctx, "清除日志迁移冻结标记失败: %v", err)
}
}()
if err := risk_control.Drain(ctx); err != nil {
return nil, fmt.Errorf("排空日志写入队列失败: %w", err)
}
src, err := logstore.Active(ctx)
if err != nil {
return nil, err
}
dst, err := logstore.BuildForMigration(ctx, p.Target)
if err != nil {
return nil, err
}
if _, err := dst.UserAccessLogs.DeleteAll(ctx); err != nil {
return nil, fmt.Errorf("清空目标用户访问日志失败: %w", err)
}
from, to, err := src.UserAccessLogs.MigrationRange(ctx)
if err != nil {
return nil, fmt.Errorf("读取源库时间范围失败: %w", err)
}
if !from.IsZero() && !to.IsZero() {
if err := dst.UserAccessLogs.EnsurePartitions(ctx, from, to); err != nil {
return nil, fmt.Errorf("预建目标分区失败: %w", err)
}
}
if err := copyUserAccessLogs(ctx, src, dst); err != nil {
return nil, err
}
if err := flipLogDatabase(ctx, p.Target); err != nil {
return nil, err
}
logstore.InvalidateCache()
driver_asynq_worker.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
return &driver_asynq_worker.TaskResult{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
}
func validateSwitch(ctx context.Context, target string) error {
source, err := currentLogDatabase(ctx)
if err != nil {
return err
}
if source == target {
return errors.New("目标日志库与当前日志库相同,无需迁移")
}
switch target {
case targetClickHouse:
if !config.Config.ClickHouse.Enabled {
return errors.New("ClickHouse 未启用,无法迁移到 ClickHouse")
}
case targetPostgres:
if !config.Config.Database.Enabled {
return errors.New("PostgreSQL 未启用(当前主库为 SQLite),无法迁移到 PostgreSQL")
}
case targetSQLite:
if config.Config.Database.Enabled {
return errors.New("当前主库为 PostgreSQL,日志库不能设置为 SQLite")
}
}
return nil
}
func currentLogDatabase(ctx context.Context) (string, error) {
cfg, err := GetSystemConfigByKey(ctx, ConfigKeyLogDatabase)
if err != nil {
return "", fmt.Errorf("读取日志主库失败: %w", err)
}
if cfg.Value == "" {
return "", errors.New("日志主库配置为空")
}
return cfg.Value, nil
}
func setMigrationFlag(ctx context.Context, v string) error {
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDBMigration, v)
}
func flipLogDatabase(ctx context.Context, target string) error {
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDatabase, target)
}
func copyUserAccessLogs(ctx context.Context, src, dst *logstore.Store) error {
var afterID uint64
var copied int
for {
rows, err := src.UserAccessLogs.ListForMigration(ctx, afterID, copyBatchSize)
if err != nil {
return fmt.Errorf("读取源用户访问日志失败: %w", err)
}
if len(rows) == 0 {
break
}
if err := dst.UserAccessLogs.BatchInsert(ctx, rows); err != nil {
return fmt.Errorf("写入目标用户访问日志失败: %w", err)
}
afterID = rows[len(rows)-1].ID
copied += len(rows)
driver_asynq_worker.AppendLog(ctx, "已复制用户访问日志 %d 条", copied)
if len(rows) < copyBatchSize {
break
}
}
return nil
}
@@ -0,0 +1,234 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"errors"
"fmt"
"math"
"net/http"
"runtime"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/backend/pkg/config"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/risk_control/logstore"
)
var startTime = time.Now()
const (
minutesInHour = 60
secondsInMinute = 60
nanosPerSecond = 1e9
logDBNamePostgres = "postgres"
logDBNameSQLite = "sqlite"
logDBNameClickHouse = "clickhouse"
defaultLogRetentionDays = 30
)
// SystemStatusResponse 系统状态响应结构体
type SystemStatusResponse struct {
Uptime string `json:"uptime"`
NumGoroutine int `json:"num_goroutine"`
Alloc string `json:"alloc"`
TotalAlloc string `json:"total_alloc"`
Sys string `json:"sys"`
Lookups uint64 `json:"lookups"`
Mallocs uint64 `json:"mallocs"`
Frees uint64 `json:"frees"`
HeapAlloc string `json:"heap_alloc"`
HeapSys string `json:"heap_sys"`
HeapIdle string `json:"heap_idle"`
HeapInuse string `json:"heap_inuse"`
HeapReleased string `json:"heap_released"`
HeapObjects uint64 `json:"heap_objects"`
StackInuse string `json:"stack_inuse"`
StackSys string `json:"stack_sys"`
MSpanInuse string `json:"mspan_inuse"`
MSpanSys string `json:"mspan_sys"`
MCacheInuse string `json:"mcache_inuse"`
MCacheSys string `json:"mcache_sys"`
BuckHashSys string `json:"buck_hash_sys"`
GCSys string `json:"gc_sys"`
OtherSys string `json:"other_sys"`
NextGC string `json:"next_gc"`
LastGCTime string `json:"last_gc_time"`
PauseTotalNs string `json:"pause_total_ns"`
LastPause string `json:"last_pause"`
NumGC uint32 `json:"num_gc"`
}
func formatDuration(d time.Duration) string {
days := int(d.Hours()) / hoursInDay
hours := int(d.Hours()) % hoursInDay
minutes := int(d.Minutes()) % minutesInHour
seconds := int(d.Seconds()) % secondsInMinute
var res string
if days > 0 {
res += fmt.Sprintf("%d天", days)
}
if hours > 0 {
res += fmt.Sprintf("%d小时", hours)
}
if minutes > 0 {
res += fmt.Sprintf("%d分钟", minutes)
}
if seconds > 0 || res == "" {
res += fmt.Sprintf("%d秒钟", seconds)
}
return res
}
// GetSystemStatus 获取系统状态信息
// @Summary 获取系统状态信息
// @Description 获取后端服务运行状态、Goroutine、内存指标等详细统计数据,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=SystemStatusResponse} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/status [get]
func GetSystemStatus(c *gin.Context) {
var m runtime.MemStats
runtime.ReadMemStats(&m)
uptime := formatDuration(time.Since(startTime))
numGoroutine := runtime.NumGoroutine()
var lastGCTime string
switch {
case m.LastGC > 0 && m.LastGC <= math.MaxInt64:
lastGCTime = formatDuration(time.Since(time.Unix(0, int64(m.LastGC))))
case m.LastGC > 0:
lastGCTime = "未知"
default:
lastGCTime = "无"
}
var lastPause string
if m.NumGC > 0 {
lastPause = fmt.Sprintf("%.3fs", float64(m.PauseNs[(m.NumGC-1)%256])/nanosPerSecond)
} else {
lastPause = "0.000s"
}
res := SystemStatusResponse{
Uptime: uptime,
NumGoroutine: numGoroutine,
Alloc: formatBytes(m.Alloc),
TotalAlloc: formatBytes(m.TotalAlloc),
Sys: formatBytes(m.Sys),
Lookups: m.Lookups,
Mallocs: m.Mallocs,
Frees: m.Frees,
HeapAlloc: formatBytes(m.HeapAlloc),
HeapSys: formatBytes(m.HeapSys),
HeapIdle: formatBytes(m.HeapIdle),
HeapInuse: formatBytes(m.HeapInuse),
HeapReleased: formatBytes(m.HeapReleased),
HeapObjects: m.HeapObjects,
StackInuse: formatBytes(m.StackInuse),
StackSys: formatBytes(m.StackSys),
MSpanInuse: formatBytes(m.MSpanInuse),
MSpanSys: formatBytes(m.MSpanSys),
MCacheInuse: formatBytes(m.MCacheInuse),
MCacheSys: formatBytes(m.MCacheSys),
BuckHashSys: formatBytes(m.BuckHashSys),
GCSys: formatBytes(m.GCSys),
OtherSys: formatBytes(m.OtherSys),
NextGC: formatBytes(m.NextGC),
LastGCTime: lastGCTime,
PauseTotalNs: fmt.Sprintf("%.1fs", float64(m.PauseTotalNs)/nanosPerSecond),
LastPause: lastPause,
NumGC: m.NumGC,
}
c.JSON(http.StatusOK, response.OK(res))
}
// LogDatabaseStatus 日志库状态。
type LogDatabaseStatus struct {
ActiveDatabase string `json:"active_database"`
Migration string `json:"migration"`
RetentionDays map[string]int `json:"retention_days"`
AvailableTargets []string `json:"available_targets"`
}
// GetLogDatabaseStatus 返回当前日志库状态。
// @Summary 获取日志数据库状态
// @Description 返回当前日志主库、迁移状态、各库保留天数与合法迁移目标,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=LogDatabaseStatus} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/status/log-database [get]
func GetLogDatabaseStatus(c *gin.Context) {
ctx := c.Request.Context()
store, err := logstore.Active(ctx)
if err != nil {
logger.ErrorF(ctx, "获取日志存储实例失败: %v", err)
response.AbortInternal(c, "日志存储初始化失败")
return
}
activeDB, err := store.Status.ActiveDatabase(ctx)
if err != nil {
logger.ErrorF(ctx, "获取日志库状态失败: %v", err)
response.AbortInternal(c, "获取日志库状态失败")
return
}
migration := "idle"
if logstore.Migrating(ctx) {
migration = "migrating"
}
c.JSON(http.StatusOK, response.OK(LogDatabaseStatus{
ActiveDatabase: activeDB,
Migration: migration,
RetentionDays: map[string]int{
logDBNamePostgres: retentionOr(ctx, ConfigKeyLogRetentionDaysPostgres),
logDBNameSQLite: retentionOr(ctx, ConfigKeyLogRetentionDaysSQLite),
logDBNameClickHouse: retentionOr(ctx, ConfigKeyLogRetentionDaysClickHouse),
},
AvailableTargets: availableLogTargets(activeDB),
}))
}
func retentionOr(ctx context.Context, key string) int {
v, err := GetIntByKey(ctx, key)
if err != nil {
if !errors.Is(err, gorm.ErrRecordNotFound) {
logger.ErrorF(ctx, "读取日志保留天数配置失败 key=%s: %v", key, err)
}
return defaultLogRetentionDays
}
if v < 1 {
return defaultLogRetentionDays
}
return v
}
func availableLogTargets(active string) []string {
if active == logDBNameClickHouse {
if config.Config.Database.Enabled {
return []string{logDBNamePostgres}
}
return []string{logDBNameSQLite}
}
if config.Config.ClickHouse.Enabled {
return []string{logDBNameClickHouse}
}
return []string{}
}
@@ -0,0 +1,413 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/robfig/cron/v3"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/plugins/drivers/driver_asynq_cron"
"github.com/Rain-kl/Wavelet/backend/plugins/drivers/driver_asynq_worker"
)
// ListTaskTypes 获取支持的任务类型列表
// @Summary 获取支持的任务类型
// @Description 返回系统支持的所有可调度任务类型列表,包括任务名称、描述、是否支持时间范围等元数据,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]driver_asynq_worker.TaskMeta} "任务类型列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/types [get]
func ListTaskTypes(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(driver_asynq_worker.GetDispatchableTasks()))
}
// DispatchTaskRequest 下发任务请求
type DispatchTaskRequest struct {
TaskType string `json:"task_type" binding:"required"`
StartTime *time.Time `json:"start_time"`
EndTime *time.Time `json:"end_time"`
UserID *uint64 `json:"user_id"`
Payload string `json:"payload"`
}
// DispatchTask 下发任务
// @Summary 下发异步任务
// @Description 手动触发指定类型的异步任务,支持指定时间范围和用户,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body DispatchTaskRequest true "任务请求参数"
// @Success 200 {object} response.Any{data=string} "任务已入队"
// @Failure 400 {object} response.Any "任务类型不存在或参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "任务入队失败"
// @Router /api/v1/admin/tasks/dispatch [post]
func DispatchTask(c *gin.Context) {
var req DispatchTaskRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
if meta == nil {
response.AbortBadRequest(c, InvalidTaskType)
return
}
var payloadBytes []byte
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
taskID, err := driver_asynq_worker.DispatchTask(c.Request.Context(), req.TaskType, validated, "manual")
if err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err))
return
}
c.JSON(http.StatusOK, response.OK(taskID))
}
// ListTaskExecutions 查询任务执行记录列表
// @Summary 查询任务执行记录
// @Description 分页查询任务执行记录,支持按状态和任务类型筛选,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param status query string false "状态筛选 (pending/running/succeeded/failed)"
// @Param task_type query string false "任务类型筛选"
// @Param page query int false "页码" default(1)
// @Param page_size query int false "每页条数" default(20)
// @Success 200 {object} response.Any{data=object} "任务执行记录列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/executions [get]
func ListTaskExecutions(c *gin.Context) {
var req ListTaskExecutionsRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if req.TaskType != "" {
if meta := driver_asynq_worker.GetTaskMeta(req.TaskType); meta != nil {
req.TaskType = meta.AsynqTask
}
}
executions, total, err := ListTaskExecutionRecords(c.Request.Context(), req)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(gin.H{
"items": executions,
"total": total,
"page": req.Page,
"page_size": req.PageSize,
}))
}
// GetTaskExecution 查询单条任务执行详情
// @Summary 查询任务执行详情
// @Description 根据 ID 查询任务执行记录详情,包含完整执行日志,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "任务执行记录 ID"
// @Success 200 {object} response.Any{data=TaskExecution} "任务执行详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "记录不存在"
// @Router /api/v1/admin/tasks/executions/{id} [get]
func GetTaskExecution(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, InvalidTaskExecutionID)
return
}
execution, err := GetTaskExecutionByID(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, TaskNotFound)
return
}
c.JSON(http.StatusOK, response.OK(execution))
}
// RetryTask 重试失败的任务
// @Summary 重试失败任务
// @Description 重新下发一条失败的任务,创建新的执行记录,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "任务执行记录 ID"
// @Success 200 {object} response.Any{data=string} "新任务的 TaskID"
// @Failure 400 {object} response.Any "任务不支持重试或参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "记录不存在"
// @Failure 500 {object} response.Any "重试失败"
// @Router /api/v1/admin/tasks/executions/{id}/retry [post]
func RetryTask(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, InvalidTaskExecutionID)
return
}
newTaskID, err := driver_asynq_worker.RetryTask(c.Request.Context(), id)
if err != nil {
errMsg := err.Error()
switch {
case strings.Contains(errMsg, "不存在"):
response.AbortNotFound(c, errMsg)
case strings.Contains(errMsg, "只有失败的任务") || strings.Contains(errMsg, "不支持重试") || strings.Contains(errMsg, "已达到最大重试"):
response.AbortBadRequest(c, errMsg)
default:
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskRetryFailed, err))
}
return
}
c.JSON(http.StatusOK, response.OK(newTaskID))
}
// ListSchedules 获取定时任务列表
// @Summary 获取定时任务列表
// @Description 返回系统所有的定时任务配置列表,包括名称、关联的异步任务类型、Cron 表达式和启用状态,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]Schedule} "定时任务列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/schedules [get]
func ListSchedules(c *gin.Context) {
schedules, err := ListSchedulesRecord(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(schedules))
}
// CreateScheduleRequest 创建定时任务请求
type CreateScheduleRequest struct {
Name string `json:"name" binding:"required"`
TaskType string `json:"task_type" binding:"required"`
Cron string `json:"cron" binding:"required"`
Payload string `json:"payload"`
IsActive *bool `json:"is_active" binding:"required"`
}
// CreateSchedule 创建定时任务
// @Summary 创建定时任务
// @Description 新增一个动态定时任务配置,关联已有的异步任务,配置 Cron 表达式和执行参数,并触发调度器热加载,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body CreateScheduleRequest true "创建定时任务请求参数"
// @Success 200 {object} response.Any{data=Schedule} "创建成功的定时任务信息"
// @Failure 400 {object} response.Any "Cron 表达式无效、异步任务类型不存在或参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "保存定时任务失败"
// @Router /api/v1/admin/tasks/schedules [post]
func CreateSchedule(c *gin.Context) {
var req CreateScheduleRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
// 校验 Cron 表达式
if _, err := cron.ParseStandard(req.Cron); err != nil {
response.AbortBadRequest(c, InvalidCronExpression)
return
}
// 校验关联的异步任务类型
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
if meta == nil {
response.AbortBadRequest(c, InvalidTaskType)
return
}
// 校验并规范化 Payload
var payloadBytes []byte
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
schedule := &Schedule{
Name: req.Name,
TaskType: req.TaskType,
Cron: req.Cron,
Payload: string(validated),
IsActive: *req.IsActive,
}
if err := CreateScheduleRecord(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
// 触发调度服务重载
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
c.JSON(http.StatusOK, response.OK(schedule))
}
// UpdateScheduleRequest 修改定时任务请求
type UpdateScheduleRequest struct {
Name string `json:"name" binding:"required"`
TaskType string `json:"task_type" binding:"required"`
Cron string `json:"cron" binding:"required"`
Payload string `json:"payload"`
IsActive *bool `json:"is_active" binding:"required"`
}
// UpdateSchedule 修改定时任务
// @Summary 修改定时任务
// @Description 修改一个定时任务的配置(名称、Cron 表达式、异步任务参数和是否启用等),并触发调度器热加载,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "定时任务 ID"
// @Param request body UpdateScheduleRequest true "修改定时任务请求参数"
// @Success 200 {object} response.Any{data=Schedule} "修改后的定时任务信息"
// @Failure 400 {object} response.Any "Cron 表达式无效、参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "定时任务不存在"
// @Failure 500 {object} response.Any "修改定时任务失败"
// @Router /api/v1/admin/tasks/schedules/{id} [put]
func UpdateSchedule(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "无效的定时任务ID")
return
}
var req UpdateScheduleRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
schedule, err := GetScheduleByID(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, ScheduleNotFound)
return
}
// 校验 Cron 表达式
if _, err := cron.ParseStandard(req.Cron); err != nil {
response.AbortBadRequest(c, InvalidCronExpression)
return
}
// 校验关联的异步任务类型
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
if meta == nil {
response.AbortBadRequest(c, InvalidTaskType)
return
}
// 校验并规范化 Payload
var payloadBytes []byte
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
schedule.Name = req.Name
schedule.TaskType = req.TaskType
schedule.Cron = req.Cron
schedule.Payload = string(validated)
schedule.IsActive = *req.IsActive
if err := UpdateScheduleRecord(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
// 触发调度服务重载
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
c.JSON(http.StatusOK, response.OK(schedule))
}
// DeleteSchedule 删除定时任务
// @Summary 删除定时任务
// @Description 删除指定的定时任务配置,并触发调度器热加载,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "定时任务 ID"
// @Success 200 {object} response.Any{data=string} "删除结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "删除定时任务失败"
// @Router /api/v1/admin/tasks/schedules/{id} [delete]
func DeleteSchedule(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "无效的定时任务ID")
return
}
if err := DeleteScheduleRecord(c.Request.Context(), id); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))
return
}
// 触发调度服务重载
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,243 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"errors"
"net/http"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
)
// CreateTemplateRequest 创建模板请求
type CreateTemplateRequest struct {
Key string `json:"key" binding:"required,max=80"`
Name string `json:"name" binding:"required,max=100"`
Type string `json:"type" binding:"required,max=20"`
Subject string `json:"subject" binding:"max=255"`
Content string `json:"content" binding:"required"`
Description string `json:"description" binding:"max=255"`
}
// UpdateTemplateRequest 更新模板请求
type UpdateTemplateRequest struct {
Name string `json:"name" binding:"required,max=100"`
Type string `json:"type" binding:"required,max=20"`
Subject string `json:"subject" binding:"max=255"`
Content string `json:"content" binding:"required"`
Description string `json:"description" binding:"max=255"`
}
func abortTemplateLogicError(c *gin.Context, err error) bool {
if err == nil {
return false
}
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, TemplateNotFound)
return true
}
msg := err.Error()
switch msg {
case TemplateKeyExists, SystemTemplateCannotDelete:
response.AbortBadRequest(c, msg)
return true
}
response.AbortInternal(c, msg)
return true
}
// CreateTemplate 创建模板
// @Summary 创建模板
// @Description 创建一条新的自定义通知模板,模板标识符(Key)不可重复,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body CreateTemplateRequest true "创建请求参数"
// @Success 200 {object} response.Any{data=string} "创建成功"
// @Failure 400 {object} response.Any "参数错误或模板标识符已存在"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates [post]
func CreateTemplate(c *gin.Context) {
var req CreateTemplateRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
tmpl, err := createTemplate(c.Request.Context(), req)
if abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(tmpl))
}
// ListTemplates 获取模板列表
// @Summary 获取模板列表
// @Description 返回所有通知模板列表,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]Template} "模板列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates [get]
func ListTemplates(c *gin.Context) {
templates, err := listTemplates(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(templates))
}
// GetTemplate 获取单个模板
// @Summary 获取单个模板
// @Description 根据模板标识符获取对应的模板详情,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Success 200 {object} response.Any{data=Template} "模板详情"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "模板不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates/{key} [get]
func GetTemplate(c *gin.Context) {
tmpl, err := getTemplate(c.Request.Context(), c.Param("key"))
if abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(tmpl))
}
// UpdateTemplate 更新模板
// @Summary 更新模板
// @Description 根据模板标识符更新对应的模板内容,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Param request body UpdateTemplateRequest true "更新请求参数"
// @Success 200 {object} response.Any{data=Template} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "模板不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates/{key} [put]
func UpdateTemplate(c *gin.Context) {
var req UpdateTemplateRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
tmpl, err := updateTemplate(c.Request.Context(), c.Param("key"), req)
if abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(tmpl))
}
// DeleteTemplate 删除模板
// @Summary 删除模板
// @Description 根据模板标识符删除对应模板,系统预置模板不可删除,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Failure 400 {object} response.Any "不可删除系统模板"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "模板不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates/{key} [delete]
func DeleteTemplate(c *gin.Context) {
if err := deleteTemplate(c.Request.Context(), c.Param("key")); abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
func createTemplate(ctx context.Context, req CreateTemplateRequest) (Template, error) {
exists, err := TemplateExistsByKey(ctx, req.Key)
if err != nil {
return Template{}, err
}
if exists {
return Template{}, errors.New(TemplateKeyExists)
}
tmpl := Template{
Key: req.Key,
Name: req.Name,
Type: req.Type,
Subject: req.Subject,
Content: req.Content,
Description: req.Description,
IsSystem: false,
}
if err := tmpl.Validate(); err != nil {
return Template{}, err
}
if err := CreateTemplateRecord(ctx, &tmpl); err != nil {
return Template{}, err
}
return tmpl, nil
}
func listTemplates(ctx context.Context) ([]Template, error) {
return ListTemplatesRecord(ctx)
}
func getTemplate(ctx context.Context, key string) (Template, error) {
return GetTemplateByKey(ctx, key)
}
func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest) (Template, error) {
tmpl, err := GetTemplateByKey(ctx, key)
if err != nil {
return Template{}, err
}
tmpl.Name = req.Name
tmpl.Type = req.Type
tmpl.Subject = req.Subject
tmpl.Content = req.Content
tmpl.Description = req.Description
if err := tmpl.Validate(); err != nil {
return Template{}, err
}
if err := SaveTemplateRecord(ctx, &tmpl); err != nil {
return Template{}, err
}
return tmpl, nil
}
func deleteTemplate(ctx context.Context, key string) error {
tmpl, err := GetTemplateByKey(ctx, key)
if err != nil {
return err
}
if tmpl.IsSystem {
return errors.New(SystemTemplateCannotDelete)
}
return DeleteTemplateRecord(ctx, &tmpl)
}
@@ -0,0 +1,695 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"archive/tar"
"archive/zip"
"compress/gzip"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"os"
"path/filepath"
"runtime"
"strings"
"sync"
"time"
"github.com/gin-gonic/gin"
"golang.org/x/mod/semver"
"github.com/Rain-kl/Wavelet/backend/pkg/buildinfo"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
)
const (
githubAPIBaseURL = "https://api.github.com"
maxArchiveSize = int64(1024 * 1024 * 1024)
maxReleaseSize = int64(4 * 1024 * 1024)
repositoryParts = 2
windowsOS = "windows"
archiveFileMode = 0o600
stagedBinaryMode = 0o700
)
type releaseAsset struct {
Name string `json:"name"`
BrowserDownloadURL string `json:"browser_download_url"`
Size int64 `json:"size"`
State string `json:"state"`
}
type githubRelease struct {
TagName string `json:"tag_name"`
Name string `json:"name"`
Body string `json:"body"`
HTMLURL string `json:"html_url"`
Draft bool `json:"draft"`
Prerelease bool `json:"prerelease"`
Published time.Time `json:"published_at"`
Assets []releaseAsset `json:"assets"`
}
// UpdaterStatus describes the current build and the newest compatible upstream release.
type UpdaterStatus struct {
CurrentVersion string `json:"current_version"`
BuildTime string `json:"build_time"`
LatestVersion string `json:"latest_version"`
UpdateAvailable bool `json:"update_available"`
CanUpgrade bool `json:"can_upgrade"`
Prerelease bool `json:"prerelease"`
ReleaseName string `json:"release_name"`
ReleaseNotes string `json:"release_notes"`
ReleaseURL string `json:"release_url"`
PublishedAt string `json:"published_at"`
UpstreamRepository string `json:"upstream_repository"`
AssetName string `json:"asset_name"`
Platform string `json:"platform"`
}
type releaseClient interface {
Do(req *http.Request) (*http.Response, error)
}
type updaterManager struct {
client releaseClient
mu sync.Mutex
upgrading bool
}
var defaultUpdaterManager = &updaterManager{
client: &http.Client{Timeout: 10 * time.Minute},
}
// GetUpdateStatus 获取应用更新状态
// @Summary 获取应用更新状态
// @Description 从系统配置指定的 GitHub 上游仓库查询最新兼容 Release,并与当前服务版本比较
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=UpdaterStatus} "更新状态"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "查询失败"
// @Router /api/v1/admin/update [get]
func GetUpdateStatus(c *gin.Context) {
status, _, err := defaultUpdaterManager.status(c.Request.Context())
if err != nil {
logger.ErrorF(c.Request.Context(), "[Updater] check release failed: %v", err)
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(status))
}
// ApplyUpdate 下载并应用应用更新
// @Summary 下载并应用应用更新
// @Description 下载当前平台对应的 GitHub Actions Release 资产,替换当前二进制并重启进程
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any "升级已准备并即将重启"
// @Failure 400 {object} response.Any "当前版本不可升级"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "升级准备失败"
// @Router /api/v1/admin/update/apply [post]
func ApplyUpdate(c *gin.Context) {
executable, stagedBinary, err := defaultUpdaterManager.prepareUpgrade(c.Request.Context())
if err != nil {
logger.ErrorF(c.Request.Context(), "[Updater] prepare upgrade failed: %v", err)
response.AbortBadRequest(c, err.Error())
return
}
logger.InfoF(c.Request.Context(), "[Updater] upgrade prepared; restarting with %s", stagedBinary)
c.JSON(http.StatusOK, response.OKNil())
util.Go(func() {
time.Sleep(time.Second)
if err := replaceAndRestart(executable, stagedBinary); err != nil {
defaultUpdaterManager.finishUpgrade()
logger.ErrorF(context.Background(), "[Updater] replace and restart failed: %v", err)
}
})
}
func normalizeVersion(version string) string {
version = strings.TrimSpace(version)
if version == "" || version == "dev" {
return ""
}
if !strings.HasPrefix(version, "v") {
version = "v" + version
}
if !semver.IsValid(version) {
return ""
}
return version
}
func parseRepository(raw string) (string, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return "", errors.New(errInvalidRepository)
}
if !strings.Contains(raw, "://") {
repo := strings.TrimSuffix(strings.Trim(raw, "/"), ".git")
if len(strings.Split(repo, "/")) == repositoryParts {
return repo, nil
}
return "", errors.New(errInvalidRepository)
}
parsed, err := url.Parse(raw)
if err != nil || !strings.EqualFold(parsed.Hostname(), "github.com") {
return "", errors.New(errInvalidRepository)
}
repo := strings.TrimSuffix(strings.Trim(parsed.Path, "/"), ".git")
if len(strings.Split(repo, "/")) != repositoryParts {
return "", errors.New(errInvalidRepository)
}
return repo, nil
}
func expectedAssetName(tag string) string {
extension := "tar.gz"
if runtime.GOOS == windowsOS {
extension = "zip"
}
return fmt.Sprintf("wavelet_%s_%s_%s.%s", tag, runtime.GOOS, runtime.GOARCH, extension)
}
func expectedAssetNames(repository, tag string) []string {
names := []string{expectedAssetName(tag)}
if parts := strings.Split(repository, "/"); len(parts) == repositoryParts {
repoName := parts[1]
if repoName != "wavelet" {
extension := "tar.gz"
if runtime.GOOS == windowsOS {
extension = "zip"
}
names = append(names, fmt.Sprintf("%s_%s_%s_%s.%s", repoName, tag, runtime.GOOS, runtime.GOARCH, extension))
}
}
return names
}
func selectLatestRelease(repository string, releases []githubRelease) (githubRelease, releaseAsset, error) {
var selected githubRelease
var selectedAsset releaseAsset
selectedVersion := ""
for _, release := range releases {
version := normalizeVersion(release.TagName)
if release.Draft || version == "" {
continue
}
expectedNames := expectedAssetNames(repository, release.TagName)
for _, asset := range release.Assets {
matched := false
for _, name := range expectedNames {
if asset.Name == name {
matched = true
break
}
}
if !matched || asset.BrowserDownloadURL == "" || asset.State != "uploaded" {
continue
}
if selectedVersion == "" || semver.Compare(version, selectedVersion) > 0 {
selected = release
selectedAsset = asset
selectedVersion = version
}
}
}
if selectedVersion == "" {
return githubRelease{}, releaseAsset{}, errors.New(errNoCompatibleRelease)
}
return selected, selectedAsset, nil
}
func (m *updaterManager) fetchRelease(ctx context.Context, repository string) (githubRelease, releaseAsset, error) {
req, err := http.NewRequestWithContext(
ctx,
http.MethodGet,
fmt.Sprintf("%s/repos/%s/releases?per_page=30", githubAPIBaseURL, repository),
nil,
)
if err != nil {
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseRequestFailed, err)
}
req.Header.Set("Accept", "application/vnd.github+json")
req.Header.Set("User-Agent", "Wavelet-Updater")
req.Header.Set("X-GitHub-Api-Version", "2022-11-28")
resp, err := m.client.Do(req)
if err != nil {
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseRequestFailed, err)
}
defer func() {
_ = resp.Body.Close()
}()
if resp.StatusCode != http.StatusOK {
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: HTTP %d", errReleaseRequestFailed, resp.StatusCode)
}
var releases []githubRelease
decoder := json.NewDecoder(io.LimitReader(resp.Body, maxReleaseSize))
if err := decoder.Decode(&releases); err != nil {
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseResponseInvalid, err)
}
release, asset, err := selectLatestRelease(repository, releases)
if err != nil {
return githubRelease{}, releaseAsset{}, err
}
logger.InfoF(ctx, "[Updater] Selected latest compatible release: %s (Asset: %s)", release.TagName, asset.Name)
return release, asset, nil
}
func loadRepository(ctx context.Context) (string, error) {
config, err := GetSystemConfigByKey(ctx, ConfigKeyUpdateUpstreamRepository)
if err != nil {
return "", fmt.Errorf("%s: %w", errInvalidRepository, err)
}
return parseRepository(config.Value)
}
func (m *updaterManager) status(ctx context.Context) (UpdaterStatus, releaseAsset, error) {
upstreamRepo, err := loadRepository(ctx)
if err != nil {
return UpdaterStatus{}, releaseAsset{}, err
}
release, asset, err := m.fetchRelease(ctx, upstreamRepo)
if err != nil {
return UpdaterStatus{}, releaseAsset{}, err
}
currentVersion := normalizeVersion(buildinfo.Version)
latestVersion := normalizeVersion(release.TagName)
updateAvailable := currentVersion != "" && semver.Compare(latestVersion, currentVersion) > 0
logger.InfoF(ctx, "[Updater] Check update complete. current: %s, latest: %s, update_available: %t", buildinfo.Version, release.TagName, updateAvailable)
return UpdaterStatus{
CurrentVersion: buildinfo.Version,
BuildTime: buildinfo.BuildTime,
LatestVersion: release.TagName,
UpdateAvailable: updateAvailable,
CanUpgrade: updateAvailable && runtime.GOOS != windowsOS,
Prerelease: release.Prerelease,
ReleaseName: release.Name,
ReleaseNotes: release.Body,
ReleaseURL: release.HTMLURL,
PublishedAt: release.Published.Format(time.RFC3339),
UpstreamRepository: upstreamRepo,
AssetName: asset.Name,
Platform: runtime.GOOS + "/" + runtime.GOARCH,
}, asset, nil
}
func downloadArchive(ctx context.Context, client releaseClient, asset releaseAsset, destination string) error {
if asset.Size <= 0 || asset.Size > maxArchiveSize {
return fmt.Errorf("release 资产大小无效: %d", asset.Size)
}
logger.InfoF(ctx, "[Updater] Downloading release asset: %s", asset.Name)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, asset.BrowserDownloadURL, nil)
if err != nil {
return fmt.Errorf("创建升级下载请求失败: %w", err)
}
req.Header.Set("User-Agent", "Wavelet-Updater")
resp, err := client.Do(req)
if err != nil {
return fmt.Errorf("下载升级资产失败: %w", err)
}
defer func() {
_ = resp.Body.Close()
}()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("下载升级资产失败: HTTP %d", resp.StatusCode)
}
//nolint:gosec // updater download destination is validated
file, err := os.OpenFile(destination, os.O_CREATE|os.O_EXCL|os.O_WRONLY, archiveFileMode)
if err != nil {
return fmt.Errorf("创建升级归档失败: %w", err)
}
written, err := io.Copy(file, io.LimitReader(resp.Body, maxArchiveSize+1))
if err != nil {
_ = file.Close()
return fmt.Errorf("写入升级归档失败: %w", err)
}
if err := file.Close(); err != nil {
return fmt.Errorf("关闭升级归档失败: %w", err)
}
if written > maxArchiveSize || written != asset.Size {
return fmt.Errorf("升级归档大小不匹配: got %d, want %d", written, asset.Size)
}
logger.InfoF(ctx, "[Updater] Successfully downloaded release asset to %s", destination)
return nil
}
func safeArchivePath(destination, name string) (string, error) {
cleanName := filepath.Clean(name)
if filepath.IsAbs(cleanName) || cleanName == "." || strings.HasPrefix(cleanName, ".."+string(filepath.Separator)) {
return "", fmt.Errorf("归档包含非法路径: %s", name)
}
target := filepath.Join(destination, cleanName)
relative, err := filepath.Rel(destination, target)
if err != nil || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) {
return "", fmt.Errorf("归档路径越界: %s", name)
}
return target, nil
}
func matchBinaryName(name string, candidates []string) bool {
for _, candidate := range candidates {
if runtime.GOOS == windowsOS {
if strings.EqualFold(name, candidate) {
return true
}
} else {
if name == candidate {
return true
}
}
}
return false
}
func getCandidateBinaryNames(executable string, repository string) []string {
execName := filepath.Base(executable)
names := []string{execName}
addName := func(base string) {
name := base
if runtime.GOOS == windowsOS && !strings.HasSuffix(strings.ToLower(name), ".exe") {
name += ".exe"
}
for _, existing := range names {
if existing == name {
return
}
}
names = append(names, name)
}
if parts := strings.Split(repository, "/"); len(parts) == repositoryParts {
addName(parts[1])
}
addName("wavelet")
return names
}
func isLikelyBinary(name string, isDir bool, mode os.FileMode) bool {
if isDir {
return false
}
base := strings.ToLower(filepath.Base(name))
exclusions := []string{
"license", "licence", "copying", "notice", "readme", "changelog",
}
for _, excl := range exclusions {
if strings.HasPrefix(base, excl) {
return false
}
}
if runtime.GOOS == windowsOS {
return filepath.Ext(base) == ".exe"
}
return (mode.Perm()&0111 != 0) || (filepath.Ext(base) == "")
}
func findBinaryInTarGz(archivePath string, candidates []string) (string, error) {
//nolint:gosec // updater archivePath is verified
file, err := os.Open(archivePath)
if err != nil {
return "", err
}
defer func() {
_ = file.Close()
}()
gzipReader, err := gzip.NewReader(file)
if err != nil {
return "", err
}
defer func() {
_ = gzipReader.Close()
}()
reader := tar.NewReader(gzipReader)
var binaries []string
for {
header, err := reader.Next()
if errors.Is(err, io.EOF) {
break
}
if err != nil {
return "", err
}
if header.Typeflag == tar.TypeReg && isLikelyBinary(header.Name, false, header.FileInfo().Mode()) {
binaries = append(binaries, header.Name)
}
}
if len(binaries) == 1 {
return binaries[0], nil
}
for _, name := range binaries {
if matchBinaryName(filepath.Base(name), candidates) {
return name, nil
}
}
return "", errors.New(errNoCompatibleAsset)
}
func findBinaryInZip(archivePath string, candidates []string) (string, error) {
reader, err := zip.OpenReader(archivePath)
if err != nil {
return "", err
}
defer func() {
_ = reader.Close()
}()
var binaries []string
for _, file := range reader.File {
if !file.FileInfo().IsDir() && isLikelyBinary(file.Name, false, file.FileInfo().Mode()) {
binaries = append(binaries, file.Name)
}
}
if len(binaries) == 1 {
return binaries[0], nil
}
for _, name := range binaries {
if matchBinaryName(filepath.Base(name), candidates) {
return name, nil
}
}
return "", errors.New(errNoCompatibleAsset)
}
func extractTarGz(ctx context.Context, archivePath, destination, targetName string, candidates []string) (string, error) {
binaryPathInArchive, err := findBinaryInTarGz(archivePath, candidates)
if err != nil {
return "", err
}
logger.InfoF(ctx, "[Updater] Extracting tar.gz archive: %s (extracting: %s)", archivePath, binaryPathInArchive)
//nolint:gosec // updater archivePath is verified
file, err := os.Open(archivePath)
if err != nil {
return "", err
}
defer func() {
_ = file.Close()
}()
gzipReader, err := gzip.NewReader(file)
if err != nil {
return "", err
}
defer func() {
_ = gzipReader.Close()
}()
reader := tar.NewReader(gzipReader)
for {
header, err := reader.Next()
if errors.Is(err, io.EOF) {
break
}
if err != nil {
return "", err
}
if header.Name != binaryPathInArchive {
continue
}
target, err := safeArchivePath(destination, targetName)
if err != nil {
return "", err
}
//nolint:gosec // updater destination is sanitized
output, err := os.OpenFile(target, os.O_CREATE|os.O_EXCL|os.O_WRONLY, stagedBinaryMode)
if err != nil {
return "", err
}
written, copyErr := io.Copy(output, io.LimitReader(reader, maxArchiveSize+1))
closeErr := output.Close()
if copyErr != nil {
return "", copyErr
}
if closeErr != nil {
return "", closeErr
}
if written > maxArchiveSize {
return "", errors.New("解压后的程序文件超过大小限制")
}
logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target)
return target, nil
}
return "", errors.New(errNoCompatibleAsset)
}
func extractZip(ctx context.Context, archivePath, destination, targetName string, candidates []string) (string, error) {
binaryPathInArchive, err := findBinaryInZip(archivePath, candidates)
if err != nil {
return "", err
}
logger.InfoF(ctx, "[Updater] Extracting zip archive: %s (extracting: %s)", archivePath, binaryPathInArchive)
reader, err := zip.OpenReader(archivePath)
if err != nil {
return "", err
}
defer func() {
_ = reader.Close()
}()
for _, file := range reader.File {
if file.Name != binaryPathInArchive {
continue
}
target, err := safeArchivePath(destination, targetName)
if err != nil {
return "", err
}
input, err := file.Open()
if err != nil {
return "", err
}
//nolint:gosec // updater extraction target is safe
output, err := os.OpenFile(target, os.O_CREATE|os.O_EXCL|os.O_WRONLY, stagedBinaryMode)
if err != nil {
return "", err
}
written, copyErr := io.Copy(output, io.LimitReader(input, maxArchiveSize+1))
inputCloseErr := input.Close()
outputCloseErr := output.Close()
if copyErr != nil {
return "", copyErr
}
if inputCloseErr != nil {
return "", inputCloseErr
}
if outputCloseErr != nil {
return "", outputCloseErr
}
if written > maxArchiveSize {
return "", errors.New("解压后的程序文件超过大小限制")
}
logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target)
return target, nil
}
return "", errors.New(errNoCompatibleAsset)
}
func (m *updaterManager) prepareUpgrade(ctx context.Context) (string, string, error) {
if runtime.GOOS == windowsOS {
return "", "", errors.New(errAutomaticUpgradeBlocked)
}
if normalizeVersion(buildinfo.Version) == "" {
return "", "", errors.New(errDevelopmentBuild)
}
m.mu.Lock()
defer m.mu.Unlock()
if m.upgrading {
return "", "", errors.New(errUpgradeAlreadyRunning)
}
status, asset, err := m.status(ctx)
if err != nil {
return "", "", err
}
if !status.UpdateAvailable {
return "", "", errors.New(errAlreadyUpToDate)
}
logger.InfoF(ctx, "[Updater] Preparing upgrade. current: %s, latest: %s", status.CurrentVersion, status.LatestVersion)
executable, err := os.Executable()
if err != nil {
return "", "", fmt.Errorf("定位当前程序失败: %w", err)
}
executable, err = filepath.EvalSymlinks(executable)
if err != nil {
return "", "", fmt.Errorf("解析当前程序路径失败: %w", err)
}
tempDir, err := os.MkdirTemp(filepath.Dir(executable), ".wavelet-update-*")
if err != nil {
return "", "", fmt.Errorf("创建升级目录失败: %w", err)
}
archivePath := filepath.Join(tempDir, asset.Name)
if err := downloadArchive(ctx, m.client, asset, archivePath); err != nil {
_ = os.RemoveAll(tempDir)
return "", "", err
}
targetName := filepath.Base(executable)
candidates := getCandidateBinaryNames(executable, status.UpstreamRepository)
var stagedBinary string
if strings.HasSuffix(asset.Name, ".zip") {
stagedBinary, err = extractZip(ctx, archivePath, tempDir, targetName, candidates)
} else {
stagedBinary, err = extractTarGz(ctx, archivePath, tempDir, targetName, candidates)
}
if err != nil {
_ = os.RemoveAll(tempDir)
return "", "", fmt.Errorf("解压升级资产失败: %w", err)
}
logger.InfoF(ctx, "[Updater] Staged binary successfully prepared: %s", stagedBinary)
m.upgrading = true
return executable, stagedBinary, nil
}
func (m *updaterManager) finishUpgrade() {
m.mu.Lock()
defer m.mu.Unlock()
m.upgrading = false
}
@@ -0,0 +1,615 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"errors"
"net/http"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/pkg/idgen"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/auth"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
)
const minPasswordLength = 8
// listUsersRequest 用户列表查询请求
type listUsersRequest struct {
Page int `form:"page" binding:"min=1"`
PageSize int `form:"page_size" binding:"min=1,max=100"`
UserID *uint64 `form:"user_id" binding:"omitempty,gt=0"`
Username string `form:"username"`
Email string `form:"email"`
}
type userResponse struct {
ID uint64 `json:"id,string"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Email string `json:"email"`
AvatarURL string `json:"avatar_url"`
IsActive bool `json:"is_active"`
IsAdmin bool `json:"is_admin"`
Bio string `json:"bio"`
Phone string `json:"phone"`
Gender string `json:"gender"`
Website string `json:"website"`
Location string `json:"location"`
LastLoginAt time.Time `json:"last_login_at"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// listUsersResponse 用户列表响应
type listUsersResponse struct {
Users []userResponse `json:"users"`
Total int64 `json:"total"`
}
func parseUserID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil || id == 0 {
response.AbortBadRequest(c, userNotFound)
return 0, false
}
return id, true
}
func toUserResponse(u *contracts.UserDTO) userResponse {
if u == nil {
return userResponse{}
}
return userResponse{
ID: u.ID,
Username: u.Username,
Nickname: u.Nickname,
Email: u.Email,
AvatarURL: u.AvatarURL,
IsActive: u.IsActive,
IsAdmin: u.IsAdmin,
Bio: u.Bio,
Phone: u.Phone,
Gender: u.Gender,
Website: u.Website,
Location: u.Location,
LastLoginAt: u.LastLoginAt,
CreatedAt: u.CreatedAt,
UpdatedAt: u.UpdatedAt,
}
}
func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbiddenMsgs, badRequestMsgs []string) bool {
if err == nil {
return false
}
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, notFoundMsg)
return true
}
msg := err.Error()
for _, m := range badRequestMsgs {
if msg == m {
response.AbortBadRequest(c, msg)
return true
}
}
for _, m := range forbiddenMsgs {
if msg == m {
response.AbortForbidden(c, msg)
return true
}
}
logger.ErrorF(c.Request.Context(), "Admin user error: %v", err)
response.AbortInternal(c, "内部服务器错误")
return true
}
// ListUsers 获取用户列表
// @Summary 获取用户列表
// @Description 分页返回用户列表,支持按用户 ID 和用户名筛选,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param request query listUsersRequest true "查询参数"
// @Success 200 {object} response.Any{data=listUsersResponse} "用户列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users [get]
func ListUsers(c *gin.Context) {
var req listUsersRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
total, dtos, err := listUsers(c.Request.Context(), req)
if err != nil {
logger.ErrorF(c.Request.Context(), "List admin users failed: %v", err)
response.AbortInternal(c, "获取用户列表失败")
return
}
users := make([]userResponse, 0, len(dtos))
for _, dto := range dtos {
users = append(users, toUserResponse(dto))
}
c.JSON(http.StatusOK, response.OK(listUsersResponse{
Users: users,
Total: total,
}))
}
// GetUser 获取用户详情
// @Summary 获取用户详情
// @Description 返回指定用户的完整个人资料和系统状态,需要管理员权限,不返回密码等敏感字段
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Success 200 {object} response.Any{data=userResponse} "用户详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "用户不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id} [get]
func GetUser(c *gin.Context) {
id, ok := parseUserID(c)
if !ok {
return
}
targetUser, err := getUserDetail(c.Request.Context(), id)
if abortUserLogicError(c, err, userNotFound, nil, nil) {
return
}
c.JSON(http.StatusOK, response.OK(toUserResponse(targetUser)))
}
// updateUserStatusRequest 更新用户状态请求
type updateUserStatusRequest struct {
IsActive bool `json:"is_active"`
}
// UpdateUserStatus 更新用户状态(启用/禁用)
// @Summary 更新用户状态
// @Description 启用或禁用指定用户,管理员账号无法被禁用,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Param request body updateUserStatusRequest true "状态参数"
// @Success 200 {object} response.Any{data=string} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限或尝试禁用管理员"
// @Failure 404 {object} response.Any "用户不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id}/status [put]
func UpdateUserStatus(c *gin.Context) {
var req updateUserStatusRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
id, ok := parseUserID(c)
if !ok {
return
}
if err := updateUserStatus(c.Request.Context(), id, req.IsActive); err != nil {
if abortUserLogicError(c, err, userNotFound, []string{cannotDisable}, nil) {
return
}
response.AbortInternal(c, updateUserFailed)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// DeleteUser 删除用户
// @Summary 删除用户
// @Description 删除指定非管理员用户,需要管理员权限,不能删除当前登录用户
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限、尝试删除管理员或当前用户"
// @Failure 404 {object} response.Any "用户不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id} [delete]
func DeleteUser(c *gin.Context) {
id, ok := parseUserID(c)
if !ok {
return
}
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if currUser == nil {
response.AbortUnauthorized(c, AdminRequired)
return
}
if err := deleteUser(c.Request.Context(), currUser.ID, id); err != nil {
if abortUserLogicError(c, err, userNotFound, []string{cannotDelete, cannotDeleteSelf}, nil) {
return
}
response.AbortInternal(c, deleteUserFailed)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// createUserRequest 创建用户请求
type createUserRequest struct {
Username string `json:"username" binding:"required,min=3,max=64"`
Password string `json:"password" binding:"required,min=8,max=64"`
Nickname string `json:"nickname" binding:"omitempty,max=64"`
Email string `json:"email" binding:"required,email,max=255"`
IsActive bool `json:"is_active"`
IsAdmin bool `json:"is_admin"`
}
// CreateUser 创建用户
// @Summary 创建用户
// @Description 创建一个本地密码登录的新用户,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body createUserRequest true "创建用户参数"
// @Success 200 {object} response.Any{data=userResponse} "创建成功"
// @Failure 400 {object} response.Any "参数错误或用户名已存在"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users [post]
func CreateUser(c *gin.Context) {
var req createUserRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
newUser, err := createUser(c.Request.Context(), req)
if abortUserLogicError(c, err, "", nil, []string{usernameRequired, emailRequired, passwordTooShort, usernameExists, emailExists}) {
return
}
c.JSON(http.StatusOK, response.OK(toUserResponse(newUser)))
}
// updateUserRequest 更新用户信息请求
type updateUserRequest struct {
Nickname string `json:"nickname" binding:"max=64"`
Email string `json:"email" binding:"required,email,max=255"`
IsAdmin bool `json:"is_admin"`
Password string `json:"password" binding:"omitempty,min=8,max=64"`
}
// UpdateUser 更新用户信息
// @Summary 更新用户信息
// @Description 更新指定用户的昵称、邮箱、管理员权限,并可选重置密码,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Param request body updateUserRequest true "更新参数"
// @Success 200 {object} response.Any{data=string} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限或尝试修改自身权限"
// @Failure 404 {object} response.Any "用户不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id} [put]
func UpdateUser(c *gin.Context) {
var req updateUserRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
id, ok := parseUserID(c)
if !ok {
return
}
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if currUser == nil {
response.AbortUnauthorized(c, AdminRequired)
return
}
err := updateUser(c.Request.Context(), currUser.ID, updateUserParam{
ID: id,
Nickname: req.Nickname,
Email: req.Email,
IsAdmin: req.IsAdmin,
Password: req.Password,
})
if err != nil {
if abortUserLogicError(c, err, userNotFound, []string{cannotRevokeSelfAdmin}, []string{emailRequired, emailExists, passwordTooShort}) {
return
}
response.AbortInternal(c, updateUserInfoFailed)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
func listUsers(ctx context.Context, req listUsersRequest) (int64, []*contracts.UserDTO, error) {
query := db.DB(ctx).Table("w_users")
if req.UserID != nil {
query = query.Where("id = ?", *req.UserID)
}
if req.Username != "" {
query = query.Where("username LIKE ? ESCAPE '\\'", util.EscapeLike(req.Username)+"%")
}
if req.Email != "" {
query = query.Where("email LIKE ? ESCAPE '\\'", util.EscapeLike(req.Email)+"%")
}
var total int64
if err := query.Count(&total).Error; err != nil {
return 0, nil, err
}
var users []*contracts.UserDTO
offset := (req.Page - 1) * req.PageSize
if err := query.
Select("id, username, nickname, email, avatar_url, is_active, is_admin, last_login_at, created_at, updated_at").
Order("id ASC").
Offset(offset).
Limit(req.PageSize).
Find(&users).Error; err != nil {
return 0, nil, err
}
return total, users, nil
}
func getUserDetail(ctx context.Context, id uint64) (*contracts.UserDTO, error) {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").
Select("id, username, nickname, email, avatar_url, is_active, is_admin, bio, phone, gender, website, location, last_login_at, created_at, updated_at").
Where("id = ?", id).
First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
func updateUserStatus(ctx context.Context, id uint64, active bool) error {
var flags struct {
ID uint64
IsAdmin bool
}
if err := db.DB(ctx).Table("w_users").Select("id, is_admin").Where("id = ?", id).First(&flags).Error; err != nil {
return err
}
if !active && flags.IsAdmin {
return errors.New(cannotDisable)
}
var tokenHashes []string
if !active {
_ = db.DB(ctx).Table("w_access_tokens").Where("user_id = ?", id).Pluck("token_hash", &tokenHashes).Error
}
err := db.DB(ctx).Table("w_users").Where("id = ?", id).Update("is_active", active).Error
if err == nil {
auth.InvalidateCachedUser(ctx, id)
if !active {
for _, hash := range tokenHashes {
auth.InvalidateCachedToken(ctx, hash)
}
}
}
return err
}
func deleteUser(ctx context.Context, currentUserID, targetID uint64) error {
if currentUserID == targetID {
return errors.New(cannotDeleteSelf)
}
var flags struct {
ID uint64
IsAdmin bool
}
if err := db.DB(ctx).Table("w_users").Select("id, is_admin").Where("id = ?", targetID).First(&flags).Error; err != nil {
return err
}
if flags.IsAdmin {
return errors.New(cannotDelete)
}
var tokenHashes []string
_ = db.DB(ctx).Table("w_access_tokens").Where("user_id = ?", targetID).Pluck("token_hash", &tokenHashes).Error
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Table("w_access_tokens").Where("user_id = ?", targetID).Delete(map[string]any{}).Error; err != nil {
return err
}
if err := tx.Table("w_external_accounts").Where("user_id = ?", targetID).Delete(map[string]any{}).Error; err != nil {
return err
}
return tx.Table("w_users").Where("id = ?", targetID).Delete(map[string]any{}).Error
})
if err == nil {
auth.InvalidateCachedUser(ctx, targetID)
for _, hash := range tokenHashes {
auth.InvalidateCachedToken(ctx, hash)
}
}
return err
}
func createUser(ctx context.Context, req createUserRequest) (*contracts.UserDTO, error) {
req.Username = strings.TrimSpace(req.Username)
req.Nickname = strings.TrimSpace(req.Nickname)
req.Password = strings.TrimSpace(req.Password)
req.Email = strings.TrimSpace(req.Email)
if req.Username == "" {
return nil, errors.New(usernameRequired)
}
if req.Email == "" {
return nil, errors.New(emailRequired)
}
if len(req.Password) < minPasswordLength {
return nil, errors.New(passwordTooShort)
}
var count int64
if err := db.DB(ctx).Table("w_users").Where("username = ?", req.Username).Count(&count).Error; err != nil {
return nil, err
}
if count > 0 {
return nil, errors.New(usernameExists)
}
var emailCount int64
if err := db.DB(ctx).Table("w_users").Where("email = ?", req.Email).Count(&emailCount).Error; err != nil {
return nil, err
}
if emailCount > 0 {
return nil, errors.New(emailExists)
}
hash, err := util.HashPassword(req.Password)
if err != nil {
return nil, err
}
if req.Nickname == "" {
req.Nickname = req.Username
}
now := time.Now()
newUser := contracts.UserDTO{
ID: idgen.NextUint64ID(),
Username: req.Username,
Nickname: req.Nickname,
Email: req.Email,
IsActive: req.IsActive,
IsAdmin: req.IsAdmin,
CreatedAt: now,
UpdatedAt: now,
}
row := map[string]any{
"id": newUser.ID,
"username": newUser.Username,
"password": hash,
"nickname": newUser.Nickname,
"email": newUser.Email,
"is_active": newUser.IsActive,
"is_admin": newUser.IsAdmin,
"created_at": now,
"updated_at": now,
}
if err := db.DB(ctx).Table("w_users").Create(row).Error; err != nil {
return nil, err
}
return &newUser, nil
}
type updateUserParam struct {
ID uint64
Nickname string
Email string
IsAdmin bool
Password string
}
func updateUser(ctx context.Context, currentUserID uint64, param updateUserParam) error {
param.Nickname = strings.TrimSpace(param.Nickname)
param.Email = strings.TrimSpace(param.Email)
param.Password = strings.TrimSpace(param.Password)
if param.Email == "" {
return errors.New(emailRequired)
}
var targetUser contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ?", param.ID).First(&targetUser).Error; err != nil {
return err
}
if currentUserID == param.ID && !param.IsAdmin && targetUser.IsAdmin {
return errors.New(cannotRevokeSelfAdmin)
}
if targetUser.Email != param.Email {
var count int64
if err := db.DB(ctx).Table("w_users").Where("email = ? AND id != ?", param.Email, param.ID).Count(&count).Error; err != nil {
return err
}
if count > 0 {
return errors.New(emailExists)
}
}
if param.Password != "" && len(param.Password) < minPasswordLength {
return errors.New(passwordTooShort)
}
needRevokeTokens := (param.Password != "") || (targetUser.IsAdmin && !param.IsAdmin)
var tokenHashes []string
if needRevokeTokens {
_ = db.DB(ctx).Table("w_access_tokens").Where("user_id = ?", param.ID).Pluck("token_hash", &tokenHashes).Error
}
if param.Nickname == "" {
param.Nickname = targetUser.Username
}
updates := map[string]any{
"nickname": param.Nickname,
"email": param.Email,
"is_admin": param.IsAdmin,
"updated_at": time.Now(),
}
if param.Password != "" {
hash, err := util.HashPassword(param.Password)
if err != nil {
return err
}
updates["password"] = hash
}
err := db.DB(ctx).Table("w_users").Where("id = ?", param.ID).Updates(updates).Error
if err == nil {
auth.InvalidateCachedUser(ctx, param.ID)
if needRevokeTokens {
for _, hash := range tokenHashes {
auth.InvalidateCachedToken(ctx, hash)
}
}
}
return err
}
@@ -0,0 +1,44 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/pkg/trace"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
"github.com/gin-gonic/gin"
)
// LoginAdminRequired 返回管理员权限校验中间件
func LoginAdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
ctx, span := trace.Start(c.Request.Context(), "LoginAdminRequired")
defer span.End()
user, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if user == nil {
response.AbortNotFound(c, AdminRequired)
return
}
// 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限
if tokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
tokenAdmin, _ := util.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
if !tokenAdmin {
response.AbortNotFound(c, TokenAdminRequired)
return
}
}
if !user.IsAdmin {
response.AbortNotFound(c, AdminRequired)
return
}
logger.InfoF(ctx, "[LoginAdminRequired] %d %s", user.ID, user.Username)
c.Next()
}
}
@@ -0,0 +1,87 @@
-- +goose Up
-- +goose StatementBegin
CREATE TABLE IF NOT EXISTS w_system_configs (
key VARCHAR(64) PRIMARY KEY,
value TEXT NOT NULL,
type VARCHAR(32) NOT NULL DEFAULT 'system',
visibility INTEGER NOT NULL DEFAULT 0,
description VARCHAR(255),
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE TABLE IF NOT EXISTS w_templates (
id BIGINT PRIMARY KEY,
key VARCHAR(80) NOT NULL UNIQUE,
name VARCHAR(100) NOT NULL,
type VARCHAR(20) NOT NULL DEFAULT 'email',
subject VARCHAR(255),
content TEXT NOT NULL,
description VARCHAR(255),
is_system BOOLEAN NOT NULL DEFAULT FALSE,
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_templates_is_system ON w_templates (is_system);
CREATE INDEX IF NOT EXISTS idx_w_templates_created_at ON w_templates (created_at);
CREATE INDEX IF NOT EXISTS idx_w_templates_updated_at ON w_templates (updated_at);
-- Seed system configs (all default platform configs)
INSERT INTO w_system_configs (key, value, type, visibility, description, created_at, updated_at) VALUES
('cap_login_enabled', 'false', 'system', 1, '是否启用登录人机验证(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_auto_solve', 'true', 'system', 1, '打开页面后是否自动开始计算,关闭则需用户手动点击触发', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_challenge_count', '1', 'system', 0, '客户端需求解的 PoW 难题总数,默认 1,推荐 1~5', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_challenge_size', '32', 'system', 0, '人机验证盐值长度', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_challenge_difficulty', '4', 'system', 0, '人机验证 PoW 难度(目标前缀长度)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_challenge_ttl_seconds', '600', 'system', 0, '人机验证难题有效时间(秒)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_token_ttl_seconds', '1200', 'system', 0, '人机验证兑换凭证有效时间(秒)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('server_address', '', 'system', 0, '服务器地址(用于跨域源控制,不设定则允许任意源)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('smtp_host', '', 'system', 0, 'SMTP 服务器地址(例如 smtp.example.com)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('smtp_port', '587', 'system', 0, 'SMTP 端口(例如 587 或 465)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('smtp_username', '', 'system', 0, 'SMTP 账户(如 sender@example.com)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('smtp_password', '', 'system', 0, 'SMTP 访问凭证(授权码/密码)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('upload_allowed_extensions', 'jpg,png,webp', 'system', 1, '允许上传的图片扩展名(逗号分隔)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('site_name', 'Wavelet', 'system', 1, '系统平台的展示名称', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('password_login_enabled', 'true', 'system', 1, '是否允许使用账号密码登录', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('registration_enabled', 'true', 'system', 1, '控制普通用户是否可以自主注册(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('password_register_enabled', 'true', 'system', 1, '是否允许通过密码创建本地账号', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('oidc_login_enabled', 'true', 'system', 1, '是否允许使用第三方 OIDC 认证源登录', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('max_api_keys_per_user', '5', 'business', 1, '限制每个普通用户可以创建的 API Key 最大数量', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('email_login_verification_enabled', 'false', 'system', 1, '是否开启邮箱登录验证(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('email_register_verification_enabled', 'false', 'system', 1, '是否开启邮箱注册验证(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('menu_display_config', '{}', 'system', 1, '目录显示配置(JSON 字符串,格式为 {url: enabled})', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('search_engine_indexing_enabled', 'false', 'system', 1, '是否允许搜索引擎爬取/检索该站点(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('update_upstream_repository', 'Rain-kl/Wavelet', 'system', 0, 'GitHub Actions Release 上游仓库(owner/repo 或 GitHub 仓库地址)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('storage_config', '{"driver":"local","local":{"root":"."},"s3":{"region":"us-east-1"},"r2":{"region":"auto"},"minio":{"region":"us-east-1","path_style":true},"oss":{},"webdav":{}}', 'system', 0, '文件存储驱动及连接配置(JSON)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('disk_cache_max_size_mb', '1024', 'system', 0, '磁盘缓存最大空间大小(MB)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('disk_cache_ttl_minutes', '1440', 'system', 0, '磁盘缓存默认有效期(分钟)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('disk_cache_lru_enabled', 'true', 'system', 0, '是否启用 LRU 淘汰机制', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('file_access_whitelist', '["avatar"]', 'system', 0, '免登录访问的文件业务类型白名单 (JSON 数组)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('login_session_ttl_hours', '168', 'system', 0, '登录会话过期时间(小时)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('log_database', '', 'system', 0, '当前日志主库(postgres/sqlite/clickhouse),由切换任务写入', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('log_db_migration', '', 'system', 0, '日志库迁移冻结标记(空或 migrating)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
ON CONFLICT (key) DO NOTHING;
INSERT INTO w_templates (id, key, name, type, subject, content, description, is_system, created_at, updated_at) VALUES
(1, 'login_email', '登录验证码邮件', 'email', 'Wavelet 登录验证码', '<h3>Wavelet 登录验证</h3><p>您的登录验证码为:<strong>{{.Code}}</strong>,5分钟内有效,请勿将验证码泄露给他人。</p>', '用户密码登录时发送的验证码邮件模板,支持变量:{{.Code}}', TRUE, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
(2, 'register_email', '注册验证码邮件', 'email', 'Wavelet 注册验证码', '<h3>Wavelet 注册验证</h3><p>您的注册验证码为:<strong>{{.Code}}</strong>,5分钟内有效,请勿泄露给他人。</p>', '用户注册时发送的验证码邮件模板,支持变量:{{.Code}}', TRUE, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
ON CONFLICT (key) DO NOTHING;
-- +goose StatementEnd
-- +goose Down
-- +goose StatementBegin
DELETE FROM w_templates WHERE key IN ('login_email', 'register_email');
DELETE FROM w_system_configs WHERE key IN (
'cap_login_enabled', 'cap_auto_solve', 'cap_challenge_count', 'cap_challenge_size',
'cap_challenge_difficulty', 'cap_challenge_ttl_seconds', 'cap_token_ttl_seconds',
'server_address', 'smtp_host', 'smtp_port', 'smtp_username', 'smtp_password',
'upload_allowed_extensions', 'site_name', 'password_login_enabled', 'registration_enabled',
'password_register_enabled', 'oidc_login_enabled', 'max_api_keys_per_user',
'email_login_verification_enabled', 'email_register_verification_enabled',
'menu_display_config', 'search_engine_indexing_enabled', 'update_upstream_repository',
'storage_config', 'disk_cache_max_size_mb', 'disk_cache_ttl_minutes', 'disk_cache_lru_enabled',
'file_access_whitelist', 'login_session_ttl_hours', 'log_database', 'log_db_migration'
);
DROP TABLE IF EXISTS w_templates;
DROP TABLE IF EXISTS w_system_configs;
-- +goose StatementEnd
+222
View File
@@ -0,0 +1,222 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"bytes"
"errors"
"strings"
"text/template"
"time"
)
// 配置键常量 - 所有系统配置的 key 定义
const (
ConfigKeyUploadAllowedExtensions = "upload_allowed_extensions" // 允许上传的文件扩展名,逗号分隔
ConfigKeySiteName = "site_name" // 站点名称
ConfigKeyPasswordLoginEnabled = "password_login_enabled" // 是否允许密码登录
ConfigKeyRegistrationEnabled = "registration_enabled" // 是否允许注册
ConfigKeyPasswordRegisterEnabled = "password_register_enabled" // 是否允许密码注册
ConfigKeyOIDCLoginEnabled = "oidc_login_enabled" // 是否允许 OIDC 登录
ConfigKeyMaxAPIKeysPerUser = "max_api_keys_per_user" //nolint:gosec // false positive: config key name, not credentials
ConfigKeyCapLoginEnabled = "cap_login_enabled" // 是否启用登录人机验证
ConfigKeyCapAutoSolve = "cap_auto_solve" // 打开页面后是否自动开始计算(false 则需用户手动点击)
ConfigKeyCapChallengeCount = "cap_challenge_count" // 客户端需求解的 PoW 难题总数,默认 1,推荐 1~5
ConfigKeyCapChallengeSize = "cap_challenge_size" // 人机验证盐值长度
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty" // 人机验证 PoW 难度(目标前缀长度)
ConfigKeyCapChallengeTTL = "cap_challenge_ttl_seconds" // 人机验证难题有效时间(秒)
ConfigKeyCapTokenTTL = "cap_token_ttl_seconds" //nolint:gosec // false positive: config key name, not credentials
ConfigKeyServerAddress = "server_address" // 服务器地址
ConfigKeySMTPHost = "smtp_host" // SMTP 服务器地址
ConfigKeySMTPPort = "smtp_port" // SMTP 端口
ConfigKeySMTPUsername = "smtp_username" // SMTP 账户
ConfigKeySMTPPassword = "smtp_password" // SMTP 访问凭证
ConfigKeyEmailLoginVerificationEnabled = "email_login_verification_enabled" // 是否启用邮箱登录验证
ConfigKeyEmailRegisterVerificationEnabled = "email_register_verification_enabled" // 是否启用邮箱注册验证
ConfigKeyMenuDisplayConfig = "menu_display_config" // 目录显示配置 (JSON 字符串)
ConfigKeySearchEngineIndexingEnabled = "search_engine_indexing_enabled" // 是否允许搜索引擎检索
ConfigKeyFileAccessWhitelist = "file_access_whitelist" // 免登录访问的文件业务类型白名单 (JSON 数组格式)
ConfigKeyDiskCacheMaxSizeMB = "disk_cache_max_size_mb" // 磁盘缓存最大空间大小 (MB)
ConfigKeyDiskCacheTTLMinutes = "disk_cache_ttl_minutes" // 磁盘缓存默认有效期 (分钟)
ConfigKeyDiskCacheLRUEnabled = "disk_cache_lru_enabled" // 是否启用 LRU 淘汰机制
ConfigKeyLoginSessionTTLHours = "login_session_ttl_hours" // 登录会话过期时间 (小时)
ConfigKeyUpdateUpstreamRepository = "update_upstream_repository" // GitHub Actions Release 上游仓库
ConfigKeyStorageConfig = "storage_config" // 文件存储配置 (JSON)
ConfigKeyLogDatabase = "log_database" // 当前日志主库(postgres/sqlite/clickhouse),受保护
ConfigKeyLogDBMigration = "log_db_migration" // 日志库迁移冻结标记(空/migrating),受保护
ConfigKeyLogRetentionDaysPostgres = "log_retention_days_postgres" // PostgreSQL 用户访问日志保留天数
ConfigKeyLogRetentionDaysSQLite = "log_retention_days_sqlite" // SQLite 用户访问日志保留天数
ConfigKeyLogRetentionDaysClickHouse = "log_retention_days_clickhouse" // ClickHouse 用户访问日志保留天数
)
const (
// ConfigVisibilityHidden 表示配置不通过公共配置接口暴露
ConfigVisibilityHidden = 0
// ConfigVisibilityVisible 表示配置通过公共配置接口暴露
ConfigVisibilityVisible = 1
)
// SystemConfig 系统配置实体
type SystemConfig struct {
Key string `json:"key" gorm:"primaryKey;size:64;not null"`
Value string `json:"value" gorm:"type:text;not null"`
Type string `json:"type" gorm:"size:32;not null;default:'system'"`
Visibility int `json:"visibility" gorm:"not null;default:0"`
Description string `json:"description" gorm:"size:255"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName 表名
func (SystemConfig) TableName() string {
return "w_system_configs"
}
// Template 邮件/消息模板实体
type Template struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
Key string `json:"key" gorm:"uniqueIndex;size:80;not null"`
Name string `json:"name" gorm:"size:100;not null"`
Type string `json:"type" gorm:"size:20;not null;default:'email'"`
Subject string `json:"subject" gorm:"size:255"`
Content string `json:"content" gorm:"type:text;not null"`
Description string `json:"description" gorm:"size:255"`
IsSystem bool `json:"is_system" 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 表名
func (Template) TableName() string {
return "w_templates"
}
// TemplateTypeEmail 邮件模板类型
const TemplateTypeEmail = "email"
// Normalize 规范化模板字段
func (t *Template) Normalize() {
t.Key = strings.TrimSpace(t.Key)
t.Name = strings.TrimSpace(t.Name)
t.Type = strings.ToLower(strings.TrimSpace(t.Type))
t.Subject = strings.TrimSpace(t.Subject)
t.Content = strings.TrimSpace(t.Content)
t.Description = strings.TrimSpace(t.Description)
if t.Type == "" {
t.Type = TemplateTypeEmail
}
}
// Validate 校验模板必填字段
func (t *Template) Validate() error {
t.Normalize()
if t.Key == "" {
return errors.New(TemplateKeyRequired)
}
if t.Name == "" {
return errors.New(TemplateNameRequired)
}
if t.Content == "" {
return errors.New(TemplateContentRequired)
}
return nil
}
// Render 渲染模板的 Subject 和 Content
func (t *Template) Render(data any) (string, string, error) {
var subject string
if t.Subject != "" {
tmplSubject, err := template.New(t.Key + "_subject").Parse(t.Subject)
if err != nil {
return "", "", err
}
var subBuf bytes.Buffer
if err := tmplSubject.Execute(&subBuf, data); err != nil {
return "", "", err
}
subject = subBuf.String()
}
tmplContent, err := template.New(t.Key + "_content").Parse(t.Content)
if err != nil {
return "", "", err
}
var bodyBuf bytes.Buffer
if err := tmplContent.Execute(&bodyBuf, data); err != nil {
return "", "", err
}
return subject, bodyBuf.String(), nil
}
// Schedule 定时任务配置表
type Schedule struct {
ID uint64 `json:"id,string" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:128;not null"`
TaskType string `json:"task_type" gorm:"size:64;not null"`
Cron string `json:"cron" gorm:"size:64;not null"`
Payload string `json:"payload" gorm:"type:text"`
IsActive bool `json:"is_active" gorm:"not null;default:true"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (Schedule) TableName() string {
return "w_schedules"
}
// TaskExecutionStatus 任务执行状态
type TaskExecutionStatus string
// 任务执行状态
const (
TaskExecutionStatusPending TaskExecutionStatus = "pending"
TaskExecutionStatusRunning TaskExecutionStatus = "running"
TaskExecutionStatusSucceeded TaskExecutionStatus = "succeeded"
TaskExecutionStatusFailed TaskExecutionStatus = "failed"
)
// TaskExecution 任务执行记录
type TaskExecution struct {
ID uint64 `json:"id,string" gorm:"primaryKey"`
TaskID string `json:"task_id" gorm:"size:128;uniqueIndex;not null"`
TaskType string `json:"task_type" gorm:"size:64;index;not null"`
TaskName string `json:"task_name" gorm:"size:128"`
Status TaskExecutionStatus `json:"status" gorm:"size:32;index;not null"`
Retryable bool `json:"retryable" gorm:"not null;default:false"`
MaxRetry int `json:"max_retry" gorm:"not null;default:0"`
RetryCount int `json:"retry_count" gorm:"not null;default:0"`
Log string `json:"log" gorm:"type:text"`
ErrorMessage string `json:"error_message" gorm:"type:text"`
Result string `json:"result" gorm:"type:text"`
StartedAt *time.Time `json:"started_at" gorm:"index"`
FinishedAt *time.Time `json:"finished_at"`
Duration int64 `json:"duration" gorm:"comment:耗时毫秒"`
Payload string `json:"payload" gorm:"type:text"`
TriggeredBy string `json:"triggered_by" gorm:"size:32;not null;default:system"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (TaskExecution) TableName() string {
return "w_task_executions"
}
// ListTaskExecutionsRequest 分页查询任务执行记录请求参数
type ListTaskExecutionsRequest struct {
Page int `form:"page"`
PageSize int `form:"page_size"`
Status string `form:"status"`
TaskType string `form:"task_type"`
TaskTypes string `form:"task_types"`
TaskTypePrefix string `form:"task_type_prefix"`
}
// TaskExecutionCleanupStats 任务日志清理结果统计
type TaskExecutionCleanupStats struct {
HighFrequencyDeleted int64 `json:"high_frequency_deleted"`
LowFrequencyDeleted int64 `json:"low_frequency_deleted"`
}
+202
View File
@@ -0,0 +1,202 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package admin provides the system management console, diagnostics, audit logging, and configuration hot-reloading domain plugin for Cordis.
package admin
import (
"context"
"embed"
"github.com/Rain-kl/Wavelet/backend/core"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/core/extpoints"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
)
//go:embed migrations/*.sql
var adminMigrations embed.FS
// Option configures the admin plugin.
type Option func(*Plugin)
// Plugin implements core.Plugin to provide system administration and management APIs.
type Plugin struct{}
// New creates a new admin 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 admin domain plugin.
func (p *Plugin) Name() string {
return "admin"
}
// Manifest returns the plugin metadata.
func (p *Plugin) Manifest() core.Manifest {
return core.Manifest{
Name: "admin",
Version: "1.0.0",
Description: "System administration console, diagnostic monitoring, and configuration hot-reload plugin",
Author: "Wavelet Team",
}
}
// Apply registers admin routes, tasks, schedules, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
// 0. Resolve auth service for middleware (via IoC, not direct import)
var loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
var adminMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
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
}
}
// 0a. Register migrations
ctx.Migrations().Register("admin", adminMigrations)
// 1. Register Admin HTTP Routes
adminRouter := ctx.Router().Group("/api/v1/admin", loginMW, adminMW)
{
// Status & Diagnostics
adminRouter.GET("/status", GetSystemStatus)
adminRouter.GET("/status/log-database", GetLogDatabaseStatus)
adminRouter.GET("/db-info", GetDatabaseInfo)
adminRouter.GET("/db-export", ExportDatabase)
// DB Management
dbGroup := adminRouter.Group("/db-manage")
{
dbGroup.GET("/overview", GetDBOverview)
dbGroup.GET("/tables", ListDBTables)
dbGroup.GET("/table-data", GetDBTableData)
dbGroup.POST("/query", ExecuteSQL)
}
// Cache Management
cacheGroup := adminRouter.Group("/cache")
{
cacheGroup.GET("/status", GetCacheStatus)
cacheGroup.POST("/config", UpdateCacheConfig)
cacheGroup.POST("/clear", ClearCache)
}
// Updater
updateGroup := adminRouter.Group("/update")
{
updateGroup.GET("", GetUpdateStatus)
updateGroup.POST("/apply", ApplyUpdate)
}
// Logs
logsGroup := adminRouter.Group("/logs")
{
logsGroup.GET("", GetLogs)
logsGroup.GET("/access", GetAccessLogs)
logsGroup.GET("/analytics", GetLogsAnalytics)
logsGroup.GET("/ws", HandleLogWebSocket)
}
// Users
usersGroup := adminRouter.Group("/users")
{
usersGroup.GET("", ListUsers)
usersGroup.POST("", CreateUser)
usersGroup.GET("/:id", GetUser)
usersGroup.PUT("/:id/status", UpdateUserStatus)
usersGroup.PUT("/:id", UpdateUser)
usersGroup.DELETE("/:id", DeleteUser)
}
// Auth Sources
authSourcesGroup := adminRouter.Group("/auth-sources")
{
authSourcesGroup.GET("", ListAuthSources)
authSourcesGroup.POST("", CreateAuthSource)
authSourcesGroup.PUT("/:id", UpdateAuthSource)
authSourcesGroup.PUT("/:id/toggle", ToggleAuthSource)
authSourcesGroup.DELETE("/:id", DeleteAuthSource)
}
// System Configs
configGroup := adminRouter.Group("/system-configs")
{
configGroup.GET("", ListSystemConfigs)
configGroup.POST("", CreateSystemConfig)
configGroup.POST("/smtp/test", TestSMTP)
keyGroup := configGroup.Group("/:key")
{
keyGroup.GET("", GetSystemConfig)
keyGroup.PUT("", UpdateSystemConfig)
}
}
// Templates
templateGroup := adminRouter.Group("/templates")
{
templateGroup.GET("", ListTemplates)
templateGroup.POST("", CreateTemplate)
keyGroup := templateGroup.Group("/:key")
{
keyGroup.GET("", GetTemplate)
keyGroup.PUT("", UpdateTemplate)
keyGroup.DELETE("", DeleteTemplate)
}
}
// Tasks
taskGroup := adminRouter.Group("/tasks")
{
taskGroup.GET("/types", ListTaskTypes)
taskGroup.POST("/dispatch", DispatchTask)
executions := taskGroup.Group("/executions")
{
executions.GET("", ListTaskExecutions)
executions.GET("/:id", GetTaskExecution)
executions.POST("/:id/retry", RetryTask)
}
schedules := taskGroup.Group("/schedules")
{
schedules.GET("", ListSchedules)
schedules.POST("", CreateSchedule)
schedules.PUT("/:id", UpdateSchedule)
schedules.DELETE("/:id", DeleteSchedule)
}
}
}
// 2. Register Background Tasks
ctx.Task().Register("admin:system_cleanup", func(_ context.Context, _ *asynq.Task) error {
return nil
}, extpoints.WithTaskRetry(1))
// 3. Register Cron Schedules
ctx.Schedule().RegisterCron("0 4 * * *", "admin:system_cleanup", map[string]string{"type": "daily"})
// 4. Register Settings Schemas
ctx.Settings().Register(extpoints.SettingSchema{
Key: "admin.system_cleanup_cron",
Default: "0 4 * * *",
Description: "Cron expression for nightly system logs and expired tokens cleanup",
Type: "string",
Category: "maintenance",
})
return nil
}
@@ -0,0 +1,41 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin_test
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/Rain-kl/Wavelet/backend/core"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/admin"
)
func TestAdminPluginUnit(t *testing.T) {
ctx := core.NewContext(context.Background())
p := admin.New()
assert.Equal(t, "admin", p.Name())
assert.Equal(t, "1.0.0", p.Manifest().Version)
require.NoError(t, p.Apply(ctx))
// Verify routes
routes := ctx.Router().Routes()
assert.NotEmpty(t, routes)
// Verify tasks
_, ok := ctx.Tasks().Get("admin:system_cleanup")
require.True(t, ok)
// Verify schedules
sched, ok := ctx.Schedules().Get("admin:system_cleanup")
require.True(t, ok)
assert.Equal(t, "0 4 * * *", sched.Spec)
// Verify settings
setting, ok := ctx.Settings().Get("admin.system_cleanup_cron")
require.True(t, ok)
assert.Equal(t, "0 4 * * *", setting.Default)
}
+705
View File
@@ -0,0 +1,705 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/redis/go-redis/v9"
"github.com/shopspring/decimal"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/backend/pkg/cache/ram"
"github.com/Rain-kl/Wavelet/backend/pkg/idgen"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
cachepkg "github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
)
const (
configTypeSystem = "system"
errDatabaseNotInitialized = "database not initialized"
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
taskExecutionLogRedisKeyPrefix = "task:execution:log:"
taskExecutionLogExpiration = 24 * time.Hour
taskExecutionLogMaxLines = 1000
)
// PreheatSystemConfigs loads all system configs from database.
func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
database := db.DB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var configs []SystemConfig
if err := database.Find(&configs).Error; err != nil {
return nil, err
}
return configs, nil
}
// PreheatSystemConfigByKey loads a single config key from database.
func PreheatSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
database := db.DB(ctx)
if database == nil {
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
}
var sc SystemConfig
if err := database.Where("key = ?", key).First(&sc).Error; err != nil {
return SystemConfig{}, err
}
return sc, nil
}
// GetSystemConfigByGroup queries a configuration by Type and Key.
func GetSystemConfigByGroup(ctx context.Context, configType string, key string) (SystemConfig, error) {
ensureSystemConfigCacheListener()
if item, ok := ram.Get(configType, key); ok {
var sc SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
return sc, nil
}
}
database := db.DB(ctx)
if database == nil {
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
}
var sc SystemConfig
if err := database.Where("key = ?", key).First(&sc).Error; err != nil {
return SystemConfig{}, err
}
valBytes, err := json.Marshal(sc)
if err == nil {
ram.Set(ram.CacheItem{
Key: sc.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(sc.Key),
})
}
return sc, nil
}
// GetSystemConfigByKey queries config by key.
func GetSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
return GetSystemConfigByGroup(ctx, ConfigCacheType, key)
}
// ListSystemConfigsByKeys loads multiple config keys.
func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]SystemConfig, error) {
if len(keys) == 0 {
return map[string]SystemConfig{}, nil
}
ensureSystemConfigCacheListener()
result := make(map[string]SystemConfig, len(keys))
missing := make([]string, 0, len(keys))
for _, key := range keys {
if item, ok := ram.Get(ConfigCacheType, key); ok {
var sc SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
result[key] = sc
continue
}
}
missing = append(missing, key)
}
if len(missing) == 0 {
return result, nil
}
database := db.DB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var configs []SystemConfig
if err := database.Where("key IN ?", missing).Find(&configs).Error; err != nil {
return nil, err
}
for i := range configs {
valBytes, err := json.Marshal(configs[i])
if err == nil {
ram.Set(ram.CacheItem{
Key: configs[i].Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(configs[i].Key),
})
}
result[configs[i].Key] = configs[i]
}
return result, nil
}
// InvalidateVisibleSystemConfigsCache clears the cached public config list.
func InvalidateVisibleSystemConfigsCache(ctx context.Context) error {
return InvalidateAllSystemConfigCaches(ctx)
}
// ListVisibleSystemConfigs queries visible configs using local cache store.
func ListVisibleSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
ensureSystemConfigCacheListener()
items := ram.GetTypeItems(ConfigCacheType)
if len(items) > 0 {
var list []SystemConfig
for _, item := range items {
var sc SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
if sc.Visibility == ConfigVisibilityVisible {
list = append(list, sc)
}
}
}
return list, nil
}
database := db.DB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var configs []SystemConfig
if err := database.Where("visibility = ?", ConfigVisibilityVisible).Find(&configs).Error; err != nil {
return nil, err
}
for _, cfg := range configs {
valBytes, err := json.Marshal(cfg)
if err == nil {
ram.Set(ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(cfg.Key),
})
}
}
return configs, nil
}
// GetIntByKey queries config and converts to int.
func GetIntByKey(ctx context.Context, key string) (int, error) {
sc, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return 0, err
}
value, err := strconv.Atoi(sc.Value)
if err != nil {
return 0, fmt.Errorf(errConfigIntParseFailed, key, sc.Value, err)
}
return value, nil
}
// GetDecimalByKey queries config and converts to decimal.Decimal.
func GetDecimalByKey(ctx context.Context, key string, precision int32) (decimal.Decimal, error) {
sc, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return decimal.Zero, err
}
value, err := decimal.NewFromString(sc.Value)
if err != nil {
return decimal.Zero, fmt.Errorf(errConfigDecimalParseFailed, key, sc.Value, err)
}
return value.Truncate(precision), nil
}
// GetBoolByKey queries config and converts to bool.
func GetBoolByKey(ctx context.Context, key string) (bool, error) {
sc, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return false, err
}
value, err := strconv.ParseBool(sc.Value)
if err != nil {
return false, fmt.Errorf(errConfigBoolParseFailed, key, sc.Value, err)
}
return value, nil
}
// GetMenuDisplayConfig queries and parses menu config.
func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
sc, err := GetSystemConfigByKey(ctx, ConfigKeyMenuDisplayConfig)
if err != nil {
return nil, err
}
config := make(map[string]bool)
if sc.Value == "" || sc.Value == "{}" {
return config, nil
}
if err := json.Unmarshal([]byte(sc.Value), &config); err != nil {
return nil, fmt.Errorf(errParseMenuDisplayConfigFailed, err)
}
return config, nil
}
// ListAdminSystemConfigs returns all configs, optionally filtered by type.
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]SystemConfig, error) {
query := db.DB(ctx).Order("created_at DESC")
if configType != "" {
query = query.Where("type = ?", configType)
}
var configs []SystemConfig
if err := query.Find(&configs).Error; err != nil {
return nil, err
}
return configs, nil
}
// GetAdminSystemConfigByKey loads a config directly from DB.
func GetAdminSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
var config SystemConfig
if err := db.DB(ctx).Where("key = ?", key).First(&config).Error; err != nil {
return SystemConfig{}, err
}
return config, nil
}
// SystemConfigExists reports whether a config key already exists.
func SystemConfigExists(ctx context.Context, key string) (bool, error) {
var existing SystemConfig
err := db.DB(ctx).Where("key = ?", key).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
if err != nil {
return false, err
}
return true, nil
}
// CreateSystemConfigRecord persists a new system config row.
func CreateSystemConfigRecord(ctx context.Context, config *SystemConfig) error {
return db.DB(ctx).Create(config).Error
}
// UpdateSystemConfigFields applies partial updates to a system config row.
func UpdateSystemConfigFields(ctx context.Context, config *SystemConfig, updates map[string]any) error {
return db.DB(ctx).Model(config).Updates(updates).Error
}
// SaveOrUpdateSystemConfig creates or updates a config row and invalidates cache.
func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
var sc SystemConfig
err := db.DB(ctx).Where("key = ?", key).First(&sc).Error
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
if errors.Is(err, gorm.ErrRecordNotFound) {
sc = SystemConfig{
Key: key,
Value: value,
Type: configTypeSystem,
Visibility: ConfigVisibilityHidden,
}
if err := db.DB(ctx).Create(&sc).Error; err != nil {
return err
}
} else {
sc.Value = value
if err := db.DB(ctx).Save(&sc).Error; err != nil {
return err
}
}
return InvalidateSystemConfigCache(ctx, key)
}
// ListTemplatesRecord returns all templates ordered by system flag and creation time.
func ListTemplatesRecord(ctx context.Context) ([]Template, error) {
var templates []Template
if err := db.DB(ctx).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
return nil, err
}
return templates, nil
}
// GetTemplateByKey loads a template by its key.
func GetTemplateByKey(ctx context.Context, key string) (Template, error) {
var tmpl Template
if err := db.DB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil {
return Template{}, err
}
return tmpl, nil
}
// TemplateExistsByKey reports whether a template key is already taken.
func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
var existing Template
err := db.DB(ctx).Where("key = ?", key).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
if err != nil {
return false, err
}
return true, nil
}
// CreateTemplateRecord persists a new template.
func CreateTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Create(tmpl).Error
}
// SaveTemplateRecord updates an existing template.
func SaveTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Save(tmpl).Error
}
// DeleteTemplateRecord removes a template record.
func DeleteTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Delete(tmpl).Error
}
// CreateScheduleRecord 创建定时任务
func CreateScheduleRecord(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Create(schedule).Error
}
// UpdateScheduleRecord 更新定时任务
func UpdateScheduleRecord(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Save(schedule).Error
}
// DeleteScheduleRecord 删除定时任务
func DeleteScheduleRecord(ctx context.Context, id uint64) error {
return db.DB(ctx).Delete(&Schedule{}, id).Error
}
// GetScheduleByID 根据 ID 获取定时任务
func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
var schedule Schedule
if err := db.DB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
return nil, err
}
return &schedule, nil
}
// ListSchedulesRecord 获取所有定时任务
func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := db.DB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
// ListActiveSchedules 获取所有启用的定时任务
func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := db.DB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
// CreateTaskExecutionRecord 创建任务执行记录
func CreateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
execution.ID = idgen.NextUint64ID()
return db.DB(ctx).Create(execution).Error
}
// UpdateTaskExecutionRecord 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
func UpdateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
return db.DB(ctx).Omit("log").Save(execution).Error
}
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) {
var execution TaskExecution
if err := db.DB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetTaskExecutionByID 根据 ID 获取执行记录
func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) {
var execution TaskExecution
if err := db.DB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetLatestTaskExecutionByTaskType returns the most recent execution for a task type.
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*TaskExecution, bool, error) {
var execution TaskExecution
err := db.DB(ctx).
Where("task_type = ?", taskType).
Order("id DESC").
First(&execution).Error
if err == nil {
if loadErr := loadTaskExecutionLog(ctx, &execution); loadErr != nil {
return nil, false, loadErr
}
return &execution, true, nil
}
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, nil
}
return nil, false, err
}
// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
if cachepkg.Redis == nil {
return errors.New("redis client is not initialized")
}
now := time.Now().Format("15:04:05")
line := fmt.Sprintf("[%s] %s\n", now, logLine)
key := taskExecutionLogRedisKey(taskID)
_, err := cachepkg.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
pipe.RPush(ctx, key, line)
pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1)
pipe.Expire(ctx, key, taskExecutionLogExpiration)
return nil
})
if err != nil {
return fmt.Errorf("append task execution log to redis: %w", err)
}
return nil
}
// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
if cachepkg.Redis == nil {
return errors.New("redis client is not initialized")
}
key := taskExecutionLogRedisKey(taskID)
logLines, err := cachepkg.Redis.LRange(ctx, key, 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
}
if len(logLines) == 0 {
return nil
}
logText := strings.Join(logLines, "")
result := db.DB(ctx).Model(&TaskExecution{}).
Where("task_id = ?", taskID).
Update("log", logText)
if result.Error != nil {
return fmt.Errorf("persist task execution log: %w", result.Error)
}
if result.RowsAffected == 0 {
return fmt.Errorf("persist task execution log: task %q not found", taskID)
}
if err := cachepkg.Redis.Del(ctx, key).Err(); err != nil {
return fmt.Errorf("delete persisted task execution log from redis: %w", err)
}
return nil
}
// ListTaskExecutionRecords 分页查询任务执行记录
func ListTaskExecutionRecords(ctx context.Context, req ListTaskExecutionsRequest) ([]TaskExecution, int64, error) {
if req.Page <= 0 {
req.Page = 1
}
if req.PageSize <= 0 {
req.PageSize = 20
}
query := db.DB(ctx).Model(&TaskExecution{})
if req.Status != "" {
query = query.Where("status = ?", req.Status)
}
if req.TaskType != "" {
query = query.Where("task_type = ?", req.TaskType)
} else if types := parseTaskTypesFilter(req.TaskTypes); len(types) > 0 {
query = query.Where("task_type IN ?", types)
} else if req.TaskTypePrefix != "" {
query = query.Where("task_type LIKE ? ESCAPE '\\'", util.EscapeLike(req.TaskTypePrefix)+"%")
}
var total int64
if err := query.Count(&total).Error; err != nil {
return nil, 0, err
}
var executions []TaskExecution
offset := (req.Page - 1) * req.PageSize
if err := query.Order("id DESC").Offset(offset).Limit(req.PageSize).Find(&executions).Error; err != nil {
return nil, 0, err
}
if err := loadTaskExecutionLogs(ctx, executions); err != nil {
return nil, 0, err
}
return executions, total, nil
}
func parseTaskTypesFilter(raw string) []string {
if strings.TrimSpace(raw) == "" {
return nil
}
parts := strings.Split(raw, ",")
out := make([]string, 0, len(parts))
for _, part := range parts {
part = strings.TrimSpace(part)
if part != "" {
out = append(out, part)
}
}
return out
}
// MarkFailedTaskExecutionsSucceededTx marks failed executions of a task type as succeeded within a transaction.
func MarkFailedTaskExecutionsSucceededTx(
tx *gorm.DB,
taskType string,
result string,
finishedAt time.Time,
) error {
return tx.Model(&TaskExecution{}).
Where("task_type = ? AND status = ?", taskType, TaskExecutionStatusFailed).
Updates(map[string]any{
"status": TaskExecutionStatusSucceeded,
"result": result,
"finished_at": finishedAt,
}).Error
}
// CleanupTaskExecutionLogs removes finished task execution logs according to frequency-based retention.
func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecutionCleanupStats, error) {
const (
frequencyWindowDays = 30
highFrequencyThreshold = frequencyWindowDays
)
frequencyWindowStart := now.AddDate(0, 0, -frequencyWindowDays)
highFrequencyCutoff := now.AddDate(0, 0, -3)
lowFrequencyCutoff := now.AddDate(0, 0, -30)
terminalStatuses := []TaskExecutionStatus{TaskExecutionStatusSucceeded, TaskExecutionStatusFailed}
var highFrequencyTaskTypes []string
if err := db.DB(ctx).
Model(&TaskExecution{}).
Select("task_type").
Where("created_at >= ?", frequencyWindowStart).
Group("task_type").
Having("COUNT(*) > ?", highFrequencyThreshold).
Pluck("task_type", &highFrequencyTaskTypes).Error; err != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("query high-frequency task types: %w", err)
}
var highFrequencyDeleted int64
if len(highFrequencyTaskTypes) > 0 {
highFrequencyResult := db.DB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", highFrequencyCutoff).
Where("task_type IN ?", highFrequencyTaskTypes).
Delete(&TaskExecution{})
if highFrequencyResult.Error != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("delete high-frequency task execution logs: %w", highFrequencyResult.Error)
}
highFrequencyDeleted = highFrequencyResult.RowsAffected
}
lowFrequencyQuery := db.DB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", lowFrequencyCutoff)
if len(highFrequencyTaskTypes) > 0 {
lowFrequencyQuery = lowFrequencyQuery.Where("task_type NOT IN ?", highFrequencyTaskTypes)
}
lowFrequencyResult := lowFrequencyQuery.Delete(&TaskExecution{})
if lowFrequencyResult.Error != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("delete low-frequency task execution logs: %w", lowFrequencyResult.Error)
}
return TaskExecutionCleanupStats{
HighFrequencyDeleted: highFrequencyDeleted,
LowFrequencyDeleted: lowFrequencyResult.RowsAffected,
}, nil
}
func taskExecutionLogRedisKey(taskID string) string {
return cachepkg.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
}
func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
if cachepkg.Redis == nil {
return nil
}
logLines, err := cachepkg.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
}
if len(logLines) == 0 {
return nil
}
execution.Log = strings.Join(logLines, "")
return nil
}
func loadTaskExecutionLogs(ctx context.Context, executions []TaskExecution) error {
if cachepkg.Redis == nil || len(executions) == 0 {
return nil
}
commands := make([]*redis.StringSliceCmd, len(executions))
_, err := cachepkg.Redis.Pipelined(ctx, func(pipe redis.Pipeliner) error {
for i := range executions {
commands[i] = pipe.LRange(ctx, taskExecutionLogRedisKey(executions[i].TaskID), 0, -1)
}
return nil
})
if err != nil {
return fmt.Errorf("get task execution logs from redis: %w", err)
}
for i := range executions {
logLines := commands[i].Val()
if len(logLines) > 0 {
executions[i].Log = strings.Join(logLines, "")
}
}
return nil
}
@@ -0,0 +1,50 @@
//go:build !windows
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"fmt"
"os"
"path/filepath"
"syscall"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
)
const installedBinaryMode = 0o755
func replaceAndRestart(executable, stagedBinary string) error {
ctx := context.Background()
logger.InfoF(ctx, "[Updater] Swapping executable: %s -> %s", executable, stagedBinary)
backup := executable + ".old"
if err := os.Remove(backup); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("删除旧备份失败: %w", err)
}
if err := os.Rename(executable, backup); err != nil {
return fmt.Errorf("备份当前程序失败: %w", err)
}
if err := os.Rename(stagedBinary, executable); err != nil {
_ = os.Rename(backup, executable)
return fmt.Errorf("替换当前程序失败: %w", err)
}
if err := os.Chmod(executable, installedBinaryMode); err != nil {
_ = os.Remove(executable)
_ = os.Rename(backup, executable)
return fmt.Errorf("设置程序执行权限失败: %w", err)
}
stagingDir := filepath.Dir(stagedBinary)
_ = os.RemoveAll(stagingDir)
logger.InfoF(ctx, "[Updater] Executing syscall.Exec to restart service: %s %v", executable, os.Args)
//nolint:gosec // restart process via exec with same binary and args
return syscall.Exec(executable, os.Args, os.Environ())
}
@@ -0,0 +1,12 @@
//go:build windows
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import "errors"
func replaceAndRestart(_, _ string) error {
return errors.New(errAutomaticUpgradeBlocked)
}
@@ -0,0 +1,208 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"encoding/json"
"errors"
"sync"
"time"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/backend/pkg/cache/ram"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
cachepkg "github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
)
const (
// SystemConfigBroadcastChannel broadcasts system config cache updates across nodes.
SystemConfigBroadcastChannel = "system:config_broadcast"
// SystemConfigInvalidationChannel is kept as an alias for backward compatibility.
SystemConfigInvalidationChannel = SystemConfigBroadcastChannel
// SystemConfigRedisHashKey is kept for backward compatibility in tests.
SystemConfigRedisHashKey = "system:system_configs"
// SystemConfigVisibleListRedisKey is kept for backward compatibility in tests.
SystemConfigVisibleListRedisKey = "system:visible_configs"
// ConfigCacheType is the cache type for all system configs.
ConfigCacheType = "config"
)
type systemConfigBroadcastMessage struct {
Type string `json:"type"`
Key string `json:"key"`
}
// ConfigLoader loads configuration data from the database.
type ConfigLoader struct{}
// LoadAll loads all system configs from database as CacheItems.
func (ConfigLoader) LoadAll(ctx context.Context, configType string) ([]ram.CacheItem, error) {
configs, err := PreheatSystemConfigs(ctx)
if err != nil {
return nil, err
}
items := make([]ram.CacheItem, len(configs))
for i, cfg := range configs {
valBytes, err := json.Marshal(cfg)
if err != nil {
return nil, err
}
items[i] = ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(cfg.Key),
}
}
return items, nil
}
// LoadOne loads a single system config from database as a CacheItem.
func (ConfigLoader) LoadOne(ctx context.Context, configType string, key string) (ram.CacheItem, error) {
cfg, err := PreheatSystemConfigByKey(ctx, key)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return ram.CacheItem{}, ram.ErrNotFound
}
return ram.CacheItem{}, err
}
valBytes, err := json.Marshal(cfg)
if err != nil {
return ram.CacheItem{}, err
}
return ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(cfg.Key),
}, nil
}
// PreloadSystemConfigs warms the in-memory RAM cache from database on startup.
func PreloadSystemConfigs(ctx context.Context) error {
return ram.Refresh(ctx, ConfigCacheType, "", ConfigLoader{})
}
var (
systemConfigListenerOnce sync.Once
systemConfigListenerCtx context.Context
systemConfigListenerCancel context.CancelFunc
systemConfigListenerDone chan struct{}
)
func ensureSystemConfigCacheListener() {
systemConfigListenerOnce.Do(startSystemConfigCacheInvalidationListener)
}
func startSystemConfigCacheInvalidationListener() {
if cachepkg.Redis == nil {
return
}
systemConfigListenerCtx, systemConfigListenerCancel = context.WithCancel(context.Background())
systemConfigListenerDone = make(chan struct{})
redisClient := cachepkg.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 cachepkg.Redis 竞争
util.Go(func() {
listenerCtx := systemConfigListenerCtx
defer close(systemConfigListenerDone)
pubsub := redisClient.Subscribe(listenerCtx, SystemConfigBroadcastChannel)
defer func() {
_ = pubsub.Close()
}()
util.Go(func() {
<-listenerCtx.Done()
_ = pubsub.Close()
})
for msg := range pubsub.Channel() {
var payload systemConfigBroadcastMessage
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
ram.UpdateTypeItems(ConfigCacheType, nil)
continue
}
key := payload.Key
if key == "*" || key == "" {
ram.UpdateTypeItems(payload.Type, nil)
} else {
ram.Delete(payload.Type, key)
}
}
})
}
// StopSystemConfigCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard.
func StopSystemConfigCacheListener() {
if systemConfigListenerCancel != nil {
systemConfigListenerCancel()
if systemConfigListenerDone != nil {
<-systemConfigListenerDone
}
systemConfigListenerCancel = nil
systemConfigListenerDone = nil
}
systemConfigListenerOnce = sync.Once{}
}
func determineTTL(_ string) time.Duration {
// Program-determined TTL: -1 means never expire for all configs by default
return -1
}
// InvalidateSystemConfigCache triggers a broadcast to refresh the cache for key.
func InvalidateSystemConfigCache(ctx context.Context, key string) error {
ensureSystemConfigCacheListener()
// Invalidate local cache synchronously first
ram.Delete(ConfigCacheType, key)
// Broadcast to other nodes and clean legacy Redis cache key
if cachepkg.Redis != nil {
_ = cachepkg.HDel(ctx, SystemConfigRedisHashKey, key)
publishSystemConfigBroadcast(ctx, ConfigCacheType, key)
}
return nil
}
// InvalidateAllSystemConfigCaches triggers a broadcast to refresh the entire config cache.
func InvalidateAllSystemConfigCaches(ctx context.Context) error {
ensureSystemConfigCacheListener()
// Invalidate all items of type ConfigCacheType synchronously first
ram.UpdateTypeItems(ConfigCacheType, nil)
// Broadcast to other nodes and clean legacy Redis cache keys
if cachepkg.Redis != nil {
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(SystemConfigRedisHashKey), cachepkg.PrefixedKey(SystemConfigVisibleListRedisKey)).Err()
publishSystemConfigBroadcast(ctx, ConfigCacheType, "*")
}
return nil
}
func publishSystemConfigBroadcast(ctx context.Context, configType string, key string) {
if cachepkg.Redis == nil {
return
}
payload, err := json.Marshal(systemConfigBroadcastMessage{Type: configType, Key: key})
if err != nil {
return
}
_ = cachepkg.Redis.Publish(ctx, SystemConfigBroadcastChannel, payload).Err()
}
// ResetSystemConfigRAMCacheForTest clears only the process-local RAM cache.
func ResetSystemConfigRAMCacheForTest() {
ram.ResetForTest()
}
@@ -0,0 +1,158 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/glebarez/sqlite"
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
"github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
)
func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
if err != nil {
t.Fatalf("gorm.Open(sqlite) error = %v", err)
}
if err := sqliteDB.AutoMigrate(&SystemConfig{}); err != nil {
t.Fatalf("AutoMigrate(SystemConfig) error = %v", err)
}
siteConfig := SystemConfig{
Key: ConfigKeySiteName,
Value: "Wavelet",
Type: "system",
Description: "系统平台的展示名称",
}
if err := sqliteDB.Create(&siteConfig).Error; err != nil {
t.Fatalf("Create(site_name) error = %v", err)
}
mr, err := miniredis.Run()
if err != nil {
t.Fatalf("miniredis.Run() error = %v", err)
}
redisClient := redis.NewClient(&redis.Options{
Addr: mr.Addr(),
MaintNotificationsConfig: &maintnotifications.Config{
Mode: maintnotifications.ModeDisabled,
},
})
previousRedis := cache.Redis
database.SetDB(sqliteDB)
cache.Redis = redisClient
cleanup := func() {
StopSystemConfigCacheListener()
ResetSystemConfigRAMCacheForTest()
database.SetDB(nil)
cache.Redis = previousRedis
_ = redisClient.Close()
mr.Close()
}
return sqliteDB, cleanup
}
func TestListSystemConfigsByKeys_EmptyKeys(t *testing.T) {
result, err := ListSystemConfigsByKeys(context.Background(), nil)
if err != nil {
t.Fatalf("ListSystemConfigsByKeys(nil) error = %v", err)
}
if len(result) != 0 {
t.Fatalf("ListSystemConfigsByKeys(nil) = %#v, want empty map", result)
}
}
func TestListSystemConfigsByKeys_LoadsFromRAMCache(t *testing.T) {
dbConn, cleanup := setupSystemConfigTest(t)
defer cleanup()
ctx := context.Background()
ResetSystemConfigRAMCacheForTest()
// Initial load
warm, err := GetSystemConfigByKey(ctx, ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey(site_name) warm error = %v", err)
}
if warm.Value != "Wavelet" {
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", warm.Value, "Wavelet")
}
// Update DB directly
if err := dbConn.Model(&SystemConfig{}).
Where("key = ?", ConfigKeySiteName).
Update("value", "db_only_value").Error; err != nil {
t.Fatalf("Update(site_name) error = %v", err)
}
// Fetch via ListSystemConfigsByKeys should serve from local store (meaning the old value "Wavelet")
configs, err := ListSystemConfigsByKeys(ctx, []string{ConfigKeySiteName})
if err != nil {
t.Fatalf("ListSystemConfigsByKeys(site_name) error = %v", err)
}
sc, ok := configs[ConfigKeySiteName]
if !ok {
t.Fatal("ListSystemConfigsByKeys(site_name) missing site_name entry")
}
if sc.Value != "Wavelet" {
t.Fatalf("ListSystemConfigsByKeys(site_name).Value = %q, want cached value %q", sc.Value, "Wavelet")
}
}
func TestGetSystemConfigByGroupAndInvalidation(t *testing.T) {
dbConn, cleanup := setupSystemConfigTest(t)
defer cleanup()
ctx := context.Background()
ResetSystemConfigRAMCacheForTest()
// Get via specific group/type
cfg, err := GetSystemConfigByGroup(ctx, ConfigCacheType, ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByGroup error = %v", err)
}
if cfg.Value != "Wavelet" {
t.Fatalf("value = %q, want %q", cfg.Value, "Wavelet")
}
// Direct DB update
if err := dbConn.Model(&SystemConfig{}).
Where("key = ?", ConfigKeySiteName).
Update("value", "new_site_name").Error; err != nil {
t.Fatalf("DB Update error = %v", err)
}
// Invalidate
if err := InvalidateSystemConfigCache(ctx, ConfigKeySiteName); err != nil {
t.Fatalf("InvalidateSystemConfigCache error = %v", err)
}
// Wait for broadcast execution
time.Sleep(100 * time.Millisecond)
// Fetch again
updated, err := GetSystemConfigByKey(ctx, ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey error = %v", err)
}
if updated.Value != "new_site_name" {
t.Fatalf("value = %q, want %q", updated.Value, "new_site_name")
}
}
+38
View File
@@ -0,0 +1,38 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"encoding/json"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
"github.com/gin-gonic/gin"
)
// LogForAudit 将登录鉴权审计日志写入 Logger
func LogForAudit(ctx context.Context, user *contracts.UserDTO, c *gin.Context) {
if user == nil || c == nil {
return
}
auditLog := loginRequiredAuditLog{
UserID: user.ID,
Username: user.Username,
ClientIP: c.ClientIP(),
Method: c.Request.Method,
Path: c.Request.URL.Path,
RequestURI: c.Request.RequestURI,
UserAgent: c.Request.UserAgent(),
Referer: c.Request.Referer(),
}
auditJSON, err := json.Marshal(auditLog)
if err != nil {
logger.ErrorF(ctx, "[LoginRequiredAudit] marshal failed: %v", err)
logger.DebugF(ctx, "[LoginRequiredAudit] %s %d %s", c.ClientIP(), user.ID, user.Username)
} else {
logger.DebugF(ctx, "[LoginRequiredAudit] %s", auditJSON)
}
}
@@ -0,0 +1,219 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"errors"
"fmt"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
)
func isOIDCLoginEnabled(ctx context.Context) bool {
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "oidc_login_enabled").Pluck("value", &val).Error; err != nil || val == "" {
return true
}
b, err := strconv.ParseBool(val)
if err != nil {
return true
}
return b
}
func resolveAuthSource(ctx context.Context, sourceName string) (*AuthSource, error) {
name := strings.TrimSpace(strings.ToLower(sourceName))
if name == "" {
sources, err := GetActiveAuthSourcesCached(ctx)
if err != nil {
return nil, err
}
if len(sources) == 0 {
return nil, errors.New(errNoActiveAuthSource)
}
src, err := GetAuthSourceByNameCached(ctx, sources[0].Name)
if err != nil {
return nil, err
}
return src, nil
}
src, err := GetAuthSourceByNameCached(ctx, name)
if err != nil {
return nil, err
}
return src, nil
}
func activeLoginSources(ctx context.Context) []AuthSourceView {
if !isOIDCLoginEnabled(ctx) {
return nil
}
dbSources, err := GetActiveAuthSourcesCached(ctx)
if err != nil {
return nil
}
sources := make([]AuthSourceView, 0, len(dbSources))
for _, source := range dbSources {
sources = append(sources, AuthSourceView{
ID: source.ID,
Name: source.Name,
Type: source.Type,
DisplayName: source.DisplayName,
IsActive: source.IsActive,
IconURL: source.IconURL,
ClientSecretConfigured: source.ClientSecretConfigured,
})
}
return sources
}
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || strings.TrimSpace(val) == "" {
return "", errors.New(errServerAddressMissing)
}
return strings.TrimRight(val, "/") + "/login", nil
}
func buildOAuthConfig(ctx context.Context, source *AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
if source == nil {
return nil, nil, errors.New(errAuthSourceRequired)
}
if source.OpenIDDiscoveryURL == "" {
return nil, nil, errors.New(errDiscoveryURLRequired)
}
// Clean the issuer URL
issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/")
issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration")
issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server")
provider, err := globalOIDCProviderCache.get(ctx, issuer)
if err != nil {
return nil, nil, err
}
verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID})
scopes := strings.Fields(source.Scopes)
if len(scopes) == 0 {
scopes = []string{oidc.ScopeOpenID, "profile", "email"}
}
if !containsScope(scopes, oidc.ScopeOpenID) {
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
}
return &oauth2.Config{
ClientID: source.ClientID,
ClientSecret: source.ClientSecret,
RedirectURL: redirectURL,
Scopes: scopes,
Endpoint: provider.Endpoint(),
}, verifier, nil
}
func containsScope(scopes []string, scope string) bool {
for _, item := range scopes {
if item == scope {
return true
}
}
return false
}
func buildOAuthUserInfo(ctx context.Context, source *AuthSource, code string, nonce string, redirectURL string) (*contracts.OAuthUserInfoDTO, error) {
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
if err != nil {
return nil, err
}
token, err := authConfig.Exchange(ctx, code)
if err != nil {
return nil, err
}
userInfo := &contracts.OAuthUserInfoDTO{Active: true}
if verifier != nil {
if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil {
return nil, verifyErr
}
}
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
userInfo.Username = userInfo.PreferredUsername
}
if userInfo.Username == "" && userInfo.Email != "" {
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
}
if userInfo.Username == "" && userInfo.Sub != "" {
userInfo.Username = userInfo.Sub
}
if userInfo.Name == "" {
userInfo.Name = userInfo.Username
}
return userInfo, nil
}
func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *contracts.OAuthUserInfoDTO) error {
rawIDToken, ok := token.Extra("id_token").(string)
if !ok {
return nil
}
idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
if verifyErr != nil {
return fmt.Errorf(errIDTokenVerifyFailedFormat, errIDTokenVerifyFailed, verifyErr)
}
if nonce != "" && idToken.Nonce != nonce {
return errors.New(errNonceMismatch)
}
if claimsErr := idToken.Claims(userInfo); claimsErr != nil {
return claimsErr
}
return nil
}
func normalizeOAuthUserInfo(userInfo *contracts.OAuthUserInfoDTO) error {
userInfo.Username = strings.TrimSpace(userInfo.Username)
userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername)
userInfo.Email = strings.TrimSpace(userInfo.Email)
userInfo.Name = strings.TrimSpace(userInfo.Name)
userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL)
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
userInfo.Username = userInfo.PreferredUsername
}
if userInfo.Username == "" && userInfo.Email != "" {
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
}
if userInfo.Username == "" && userInfo.Sub != "" {
userInfo.Username = userInfo.Sub
}
if userInfo.Username == "" {
return errors.New(errUsernameFromSourceFailed)
}
if userInfo.Name == "" {
userInfo.Name = userInfo.Username
}
if !userInfo.Active {
userInfo.Active = true
}
return nil
}
func buildCallbackResult(user *contracts.UserDTO, status string) OAuthCallbackResult {
result := OAuthCallbackResult{Status: status}
if user != nil {
info := BuildBasicUserInfo(user, false)
result.User = &info
}
return result
}
+259
View File
@@ -0,0 +1,259 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"fmt"
"strconv"
"sync"
"time"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/pkg/cache/ram"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
)
const (
tokenCacheTTL = 5 * time.Minute
userCacheTTL = 5 * time.Minute
//nolint:gosec // This is a Redis Pub/Sub channel name, not a credential
oauthTokenInvalidationChannel = "oauth:token_invalidation"
oauthUserInvalidationChannel = "oauth:user_invalidation"
)
// CachedToken represents the minimal cached representation of an access token.
type CachedToken struct {
ID uint64 `json:"id"`
UserID uint64 `json:"user_id"`
IsAdmin bool `json:"is_admin"`
}
var (
tokenRAM = ram.MustNew[string, *CachedToken](ram.Options{MaximumSize: 2048})
userRAM = ram.MustNew[uint64, *contracts.UserDTO](ram.Options{MaximumSize: 2048})
tokenListenerOnce sync.Once
tokenListenerCtx context.Context
tokenListenerCancel context.CancelFunc
tokenListenerDone chan struct{}
userListenerOnce sync.Once
userListenerCtx context.Context
userListenerCancel context.CancelFunc
userListenerDone chan struct{}
)
func tokenCacheKey(tokenHash string) string {
return "oauth:token:" + tokenHash
}
func userCacheKey(userID uint64) string {
return fmt.Sprintf("oauth:user:%d", userID)
}
func ensureTokenCacheListener() {
if db.Redis == nil {
return
}
tokenListenerOnce.Do(startTokenCacheInvalidationListener)
}
func startTokenCacheInvalidationListener() {
tokenListenerCtx, tokenListenerCancel = context.WithCancel(context.Background())
tokenListenerDone = make(chan struct{})
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
util.Go(func() {
listenerCtx := tokenListenerCtx
defer close(tokenListenerDone)
pubsub := redisClient.Subscribe(listenerCtx, oauthTokenInvalidationChannel)
defer func() {
_ = pubsub.Close()
}()
util.Go(func() {
<-listenerCtx.Done()
_ = pubsub.Close()
})
for msg := range pubsub.Channel() {
tokenHash := msg.Payload
if tokenHash == "" || tokenHash == "*" || tokenHash == "reset" {
tokenRAM.InvalidateAll()
} else {
tokenRAM.Invalidate(tokenHash)
}
}
})
}
func publishTokenRAMInvalidation(ctx context.Context, tokenHash string) {
if db.Redis == nil {
return
}
_ = db.Redis.Publish(ctx, oauthTokenInvalidationChannel, tokenHash).Err()
}
func ensureUserCacheListener() {
if db.Redis == nil {
return
}
userListenerOnce.Do(startUserCacheInvalidationListener)
}
func startUserCacheInvalidationListener() {
userListenerCtx, userListenerCancel = context.WithCancel(context.Background())
userListenerDone = make(chan struct{})
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
util.Go(func() {
listenerCtx := userListenerCtx
defer close(userListenerDone)
pubsub := redisClient.Subscribe(listenerCtx, oauthUserInvalidationChannel)
defer func() {
_ = pubsub.Close()
}()
util.Go(func() {
<-listenerCtx.Done()
_ = pubsub.Close()
})
for msg := range pubsub.Channel() {
userIDStr := msg.Payload
if userIDStr == "" || userIDStr == "*" || userIDStr == "reset" {
userRAM.InvalidateAll()
} else if userID, err := strconv.ParseUint(userIDStr, 10, 64); err == nil {
userRAM.Invalidate(userID)
}
}
})
}
func publishUserRAMInvalidation(ctx context.Context, userID uint64) {
if db.Redis == nil {
return
}
_ = db.Redis.Publish(ctx, oauthUserInvalidationChannel, strconv.FormatUint(userID, 10)).Err()
}
// GetCachedToken 获取缓存的 Token
func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) {
ensureTokenCacheListener()
if val, ok := tokenRAM.GetIfPresent(tokenHash); ok {
return val, nil
}
if db.Redis != nil {
var token CachedToken
key := tokenCacheKey(tokenHash)
if err := db.GetJSON(ctx, key, &token); err == nil {
// Write back to local cache
tokenRAM.Set(tokenHash, &token)
return &token, nil
}
}
return nil, fmt.Errorf("cache miss")
}
// SetCachedToken 设置 Token 缓存
func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) {
ensureTokenCacheListener()
tokenRAM.Set(tokenHash, token)
if db.Redis != nil {
key := tokenCacheKey(tokenHash)
_ = db.SetJSON(ctx, key, token, tokenCacheTTL)
}
}
// InvalidateCachedToken 吊销/删除 token 缓存
func InvalidateCachedToken(ctx context.Context, tokenHash string) {
ensureTokenCacheListener()
tokenRAM.Invalidate(tokenHash)
if db.Redis != nil {
key := tokenCacheKey(tokenHash)
_ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err()
publishTokenRAMInvalidation(ctx, tokenHash)
}
}
// GetCachedUser 获取缓存的 UserDTO
func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
ensureUserCacheListener()
if val, ok := userRAM.GetIfPresent(userID); ok {
return val, nil
}
if db.Redis != nil {
var u contracts.UserDTO
key := userCacheKey(userID)
if err := db.GetJSON(ctx, key, &u); err == nil {
// Write back to local cache
userRAM.Set(userID, &u)
return &u, nil
}
}
return nil, fmt.Errorf("cache miss")
}
// SetCachedUser 设置 UserDTO 缓存
func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) {
ensureUserCacheListener()
userRAM.Set(userID, u)
if db.Redis != nil {
key := userCacheKey(userID)
_ = db.SetJSON(ctx, key, u, userCacheTTL)
}
}
// InvalidateCachedUser 吊销/失效 UserDTO 缓存
func InvalidateCachedUser(ctx context.Context, userID uint64) {
ensureUserCacheListener()
userRAM.Invalidate(userID)
if db.Redis != nil {
key := userCacheKey(userID)
_ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err()
publishUserRAMInvalidation(ctx, userID)
}
}
// StopAuthCacheListener stops both token and user Redis Pub/Sub subscription listeners and resets the sync.Once guards.
func StopAuthCacheListener() {
if tokenListenerCancel != nil {
tokenListenerCancel()
if tokenListenerDone != nil {
<-tokenListenerDone
}
tokenListenerCancel = nil
tokenListenerDone = nil
}
tokenListenerOnce = sync.Once{}
if userListenerCancel != nil {
userListenerCancel()
if userListenerDone != nil {
<-userListenerDone
}
userListenerCancel = nil
userListenerDone = nil
}
userListenerOnce = sync.Once{}
}
// ResetAuthRAMCacheForTest clears only the process-local RAM cache.
func ResetAuthRAMCacheForTest() {
tokenRAM.InvalidateAll()
userRAM.InvalidateAll()
}
+124
View File
@@ -0,0 +1,124 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth_test
import (
"context"
"testing"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/auth"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
)
func setupOauthCacheTest(t *testing.T) (*miniredis.Miniredis, func()) {
t.Helper()
miniRedis, err := miniredis.Run()
if err != nil {
t.Fatalf("failed to start miniredis: %v", err)
}
db.Redis = redis.NewClient(&redis.Options{
Addr: miniRedis.Addr(),
MaintNotificationsConfig: &maintnotifications.Config{
Mode: maintnotifications.ModeDisabled,
},
})
auth.ResetAuthRAMCacheForTest()
cleanup := func() {
auth.StopAuthCacheListener()
auth.ResetAuthRAMCacheForTest()
_ = db.Redis.Close()
miniRedis.Close()
db.Redis = nil
}
return miniRedis, cleanup
}
func TestTokenCache_GetSetInvalidate(t *testing.T) {
_, cleanup := setupOauthCacheTest(t)
defer cleanup()
ctx := context.Background()
tokenHash := "test-token-hash"
token := &auth.CachedToken{
ID: 123,
UserID: 456,
IsAdmin: true,
}
// 1. Get from empty cache -> miss
_, err := auth.GetCachedToken(ctx, tokenHash)
if err == nil {
t.Fatal("expected cache miss for un-cached token")
}
// 2. Set to cache
auth.SetCachedToken(ctx, tokenHash, token)
// 3. Get from cache -> hit
cached, err := auth.GetCachedToken(ctx, tokenHash)
if err != nil {
t.Fatalf("GetCachedToken() failed: %v", err)
}
if cached.ID != token.ID || cached.UserID != token.UserID || cached.IsAdmin != token.IsAdmin {
t.Fatalf("expected cached token %+v, got %+v", token, cached)
}
// 4. Invalidate cache
auth.InvalidateCachedToken(ctx, tokenHash)
// 5. Get from cache -> miss
_, err = auth.GetCachedToken(ctx, tokenHash)
if err == nil {
t.Fatal("expected cache miss after invalidation")
}
}
func TestUserCache_GetSetInvalidate(t *testing.T) {
_, cleanup := setupOauthCacheTest(t)
defer cleanup()
ctx := context.Background()
userID := uint64(789)
user := &contracts.UserDTO{
ID: userID,
Username: "testuser",
Email: "test@example.com",
}
// 1. Get from empty cache -> miss
_, err := auth.GetCachedUser(ctx, userID)
if err == nil {
t.Fatal("expected cache miss for un-cached user")
}
// 2. Set to cache
auth.SetCachedUser(ctx, userID, user)
// 3. Get from cache -> hit
cached, err := auth.GetCachedUser(ctx, userID)
if err != nil {
t.Fatalf("GetCachedUser() failed: %v", err)
}
if cached.ID != user.ID || cached.Username != user.Username {
t.Fatalf("expected cached user %+v, got %+v", user, cached)
}
// 4. Invalidate cache
auth.InvalidateCachedUser(ctx, userID)
// 5. Get from cache -> miss
_, err = auth.GetCachedUser(ctx, userID)
if err == nil {
t.Fatal("expected cache miss after invalidation")
}
}
+40
View File
@@ -0,0 +1,40 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"time"
)
// Session and Context Keys
const (
UserNameKey = "username"
UserIDKey = "user_id"
UserObjKey = "user_obj"
TokenAuthKey = "token_auth" // 标记当前请求是否通过 Access Token 鉴权
TokenAdminKey = "token_admin" // Access Token 本身是否具有管理员权限
SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: this is a session key, not hardcoded credentials
PasswordHashKey = "password_hash"
SystemUsername = "system"
)
// OAuth State Cache Keys and Expirations
const (
OAuthStateCacheKeyFormat = "oauth:state:%s"
OAuthStateCacheKeyExpiration = 10 * time.Minute
oauthStateLimitKeyFormat = "oauth:state:limit:%s"
oauthStateLimitMax = 10
)
// OAuth Purpose Constants
const (
OAuthPurposeLogin = "login"
OAuthPurposeBind = "bind"
)
// Auth Source Types
const (
AuthSourceTypeOIDC = "oidc"
)
+39
View File
@@ -0,0 +1,39 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// OAuth and Auth error messages
const (
errInvalidState = "非法登录请求"
errIDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errIDTokenVerifyFailedFormat = "%s: %w"
errNonceMismatch = "nonce 不匹配,可能存在重放攻击"
errNoActiveAuthSource = "未配置可用认证源"
errServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试"
errAuthSourceRequired = "认证源不能为空"
errDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
errUsernameGenerateFailed = "无法生成可用用户名"
errUsernameFromSourceFailed = "无法从认证源获取用户名"
errAuthSourceDisabled = "认证源未启用"
errInvalidExternalAccountBindingID = "绑定记录 ID 无效"
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errOAuthStateRateLimited = "请求授权过于频繁,请稍后重试"
errAuthSourceNameRequired = "认证源名称不能为空"
errAuthSourceNameInvalid = "认证源名称格式不正确"
errAuthSourceTypeUnsupported = "不支持的认证源类型"
errAuthSourceDiscoveryURLRequired = "Discovery URL 不能为空"
//nolint:gosec // error message, not hardcoded credentials
errAuthSourceClientCredentialsRequired = "启用认证源时必须配置 Client ID 和 Client Secret"
errAuthSourceIDRequired = "认证源 ID 不能为空"
errUserIDRequired = "用户 ID 不能为空"
errExternalAccountBindingIncomplete = "外部帐号绑定信息不完整"
errExternalAccountAlreadyBoundToAnother = "该外部帐号已被其他用户绑定"
errExternalAccountBindingIDRequired = "外部帐号绑定记录 ID 不能为空"
errAdminRequired = "无权访问"
//nolint:gosec // error message, not hardcoded credentials
errTokenAdminRequired = "令牌无管理员权限"
errBannedAccount = "账号已被封禁"
errUnAuthorized = "未登录"
)
+492
View File
@@ -0,0 +1,492 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"errors"
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/pkg/idgen"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
cachepkg "github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
"github.com/coreos/go-oidc/v3/oidc"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"gorm.io/gorm"
)
// GetLoginSources 获取可用登录源列表
func GetLoginSources(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(activeLoginSources(c.Request.Context())))
}
// GetLoginURL 获取登录授权地址
func GetLoginURL(c *gin.Context) {
ctx := c.Request.Context()
if !isOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
source, err := resolveAuthSource(ctx, c.Query("source"))
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
session := sessions.Default(c)
token, isNew := ensureSessionToken(session)
if isNew {
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
userID := GetUserIDFromSession(session)
sessionHash := hashSessionToken(token)
if err := reserveOAuthStateSlot(ctx, sessionHash); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
state := uuid.NewString()
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
SourceName: source.Name,
Purpose: OAuthPurposeLogin,
UserID: userID,
SessionHash: sessionHash,
})
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := cachepkg.Redis.Set(ctx, cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
response.AbortInternal(c, err.Error())
return
}
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
}
func buildAuthorizeURL(ctx context.Context, source *AuthSource, state string) (string, error) {
redirectURL, err := getFrontendLoginRedirectURL(ctx)
if err != nil {
return "", err
}
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
if err != nil {
return "", err
}
if verifier != nil {
return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil
}
return authConfig.AuthCodeURL(state), nil
}
func reserveOAuthStateSlot(ctx context.Context, sessionHash string) error {
if cachepkg.Redis == nil || sessionHash == "" {
return nil
}
key := cachepkg.PrefixedKey(fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash))
n, err := cachepkg.Redis.Incr(ctx, key).Result()
if err != nil {
return err
}
if n == 1 {
_ = cachepkg.Redis.Expire(ctx, key, OAuthStateCacheKeyExpiration).Err()
}
if n > oauthStateLimitMax {
return errors.New(errOAuthStateRateLimited)
}
return nil
}
// Authorize 发起指定认证源授权
func Authorize(c *gin.Context) {
ctx := c.Request.Context()
if !isOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
source, err := resolveAuthSource(ctx, c.Param("source"))
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose")))
if purpose != OAuthPurposeBind {
purpose = OAuthPurposeLogin
}
session := sessions.Default(c)
userID := GetUserIDFromSession(session)
if purpose == OAuthPurposeBind && userID == 0 {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
token, isNew := ensureSessionToken(session)
if isNew {
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
sessionHash := hashSessionToken(token)
if err := reserveOAuthStateSlot(ctx, sessionHash); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
state := uuid.NewString()
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
SourceName: source.Name,
Purpose: purpose,
UserID: userID,
SessionHash: sessionHash,
})
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := cachepkg.Redis.Set(ctx, cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
response.AbortInternal(c, err.Error())
return
}
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
}
// Callback OAuth 回调处理
func Callback(c *gin.Context) {
var req CallbackRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
ctx := c.Request.Context()
stateKey := cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
payloadRaw, err := cachepkg.Redis.Get(ctx, stateKey).Result()
if err != nil {
response.AbortBadRequest(c, errInvalidState)
return
}
_ = cachepkg.Redis.Del(ctx, stateKey)
payload, err := decodeOAuthStatePayload(payloadRaw)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
session := sessions.Default(c)
currentUserID := GetUserIDFromSession(session)
if payload.Purpose == OAuthPurposeBind && currentUserID == 0 {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
token, ok := session.Get(SessionTokenKey).(string)
if !ok || token == "" {
response.AbortBadRequest(c, "invalid session context")
return
}
if hashSessionToken(token) != payload.SessionHash {
response.AbortBadRequest(c, "session mismatch for oauth state")
return
}
if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID {
response.AbortBadRequest(c, "user context mismatch for oauth binding")
return
}
if !isOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
source, err := resolveAuthSource(ctx, payload.SourceName)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
redirectURL, err := getFrontendLoginRedirectURL(ctx)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
userInfo, err := buildOAuthUserInfo(ctx, source, req.Code, req.State, redirectURL)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := normalizeOAuthUserInfo(userInfo); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if userInfo.Sub == "" {
userInfo.Sub = userInfo.Username
}
if payload.Purpose == OAuthPurposeBind {
handleCallbackBind(ctx, c, source, userInfo)
return
}
handleCallbackLogin(ctx, c, source, userInfo)
}
func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := BindExternalAccount(ctx, &ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
ExternalUsername: userInfo.Username,
Email: userInfo.Email,
}); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
user.LastLoginAt = time.Now()
_ = db.DB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound")))
}
func handleCallbackLogin(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
var user contracts.UserDTO
account, err := FindExternalAccount(ctx, source.ID, userInfo.Sub)
switch {
case err == nil:
if loadErr := db.DB(ctx).Table("w_users").Where("id = ?", account.UserID).First(&user).Error; loadErr != nil {
response.AbortInternal(c, loadErr.Error())
return
}
case errors.Is(err, gorm.ErrRecordNotFound):
newUser, ok := handleCallbackRegister(ctx, c, source, userInfo)
if !ok {
return
}
user = newUser
default:
response.AbortInternal(c, err.Error())
return
}
user.LastLoginAt = time.Now()
_ = db.DB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
if err := SetLoginSession(ctx, c, &user); err != nil {
response.AbortInternal(c, err.Error())
return
}
SetCachedUser(ctx, user.ID, &user)
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in")))
}
func uniqueUsername(ctx context.Context, base string) (string, error) {
base = strings.TrimSpace(base)
if base == "" {
base = "user"
}
var existingUsernames []string
if err := db.DB(ctx).Table("w_users").
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
Pluck("username", &existingUsernames).Error; err != nil {
return "", err
}
exists := make(map[string]bool, len(existingUsernames))
for _, u := range existingUsernames {
exists[strings.ToLower(u)] = true
}
if !exists[strings.ToLower(base)] {
return base, nil
}
for i := 1; i <= 1000; i++ {
candidate := fmt.Sprintf("%s-%d", base, i)
if !exists[strings.ToLower(candidate)] {
return candidate, nil
}
}
return "", errors.New(errUsernameGenerateFailed)
}
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) (contracts.UserDTO, bool) {
registrationEnabled := true
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "registration_enabled").Pluck("value", &val).Error; err == nil && val != "" {
if b, err := strconv.ParseBool(val); err == nil {
registrationEnabled = b
}
}
if !registrationEnabled {
c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind")))
return contracts.UserDTO{}, false
}
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
if uniqueErr != nil {
response.AbortInternal(c, uniqueErr.Error())
return contracts.UserDTO{}, false
}
userInfo.Username = username
now := time.Now()
user := contracts.UserDTO{
ID: idgen.NextUint64ID(),
Username: userInfo.Username,
Nickname: userInfo.Name,
Email: userInfo.Email,
AvatarURL: userInfo.AvatarURL,
IsActive: userInfo.Active,
LastLoginAt: now,
CreatedAt: now,
UpdatedAt: now,
}
if err := db.DB(ctx).Table("w_users").Create(&user).Error; err != nil {
response.AbortInternal(c, err.Error())
return contracts.UserDTO{}, false
}
if err := BindExternalAccount(ctx, &ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
ExternalUsername: userInfo.Username,
Email: userInfo.Email,
}); err != nil {
response.AbortBadRequest(c, err.Error())
return contracts.UserDTO{}, false
}
logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
return user, true
}
// UserInfo 获取当前登录用户信息
func UserInfo(c *gin.Context) {
user, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
session := sessions.Default(c)
needChange := session.Get("need_change_password") == true
c.JSON(
http.StatusOK,
response.OK(BuildBasicUserInfo(user, needChange)),
)
}
// Logout 退出登录
func Logout(c *gin.Context) {
session := sessions.Default(c)
userID := session.Get(UserIDKey)
username := session.Get(UserNameKey)
if userID != nil {
logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP())
if id := ParseUserID(userID); id > 0 {
InvalidateCachedUser(c.Request.Context(), id)
}
}
session.Options(GetSessionOptions(-1))
session.Clear()
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// ListExternalAccounts 获取当前用户的外部帐号绑定列表
func ListExternalAccounts(c *gin.Context) {
userID := GetUserIDFromContext(c)
accounts, err := ListExternalAccountsByUserID(c.Request.Context(), userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(accounts))
}
// DeleteExternalAccount 解除外部帐号绑定
func DeleteExternalAccount(c *gin.Context) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
rawID := strings.TrimSpace(c.Param("id"))
id, err := strconv.ParseUint(rawID, 10, 64)
if err != nil || id == 0 {
response.AbortBadRequest(c, errInvalidExternalAccountBindingID)
return
}
if err := UnbindExternalAccount(c.Request.Context(), id, userID); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
+176
View File
@@ -0,0 +1,176 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/pkg/trace"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
"github.com/gin-gonic/gin"
)
func hashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *CachedToken, error) {
tokenHash := hashToken(tokenStr)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err == nil {
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err == nil && user != nil && user.IsActive {
return user, tokenRecord, nil
}
}
var tokenRow struct {
ID uint64
UserID uint64
IsAdmin bool
}
if err := db.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
return nil, nil, err
}
tokenRecord = &CachedToken{
ID: tokenRow.ID,
UserID: tokenRow.UserID,
IsAdmin: tokenRow.IsAdmin,
}
SetCachedToken(ctx, tokenHash, tokenRecord)
var userRow contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRow.UserID, true).First(&userRow).Error; err != nil {
return nil, nil, err
}
SetCachedUser(ctx, userRow.ID, &userRow)
return &userRow, tokenRecord, nil
}
// GetUserFromRequest 校验 Access Token 或 Session 并返回用户对象,如果未登录或用户失效则返回 error
func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
ctx := c.Request.Context()
// Check token in headers
tokenStr := c.GetHeader("X-Access-Token")
if tokenStr == "" {
authHeader := c.GetHeader("Authorization")
if len(authHeader) > 7 && authHeader[:7] == "Bearer " {
tokenStr = authHeader[7:]
}
}
// 优先使用 Access Token 鉴权
if tokenStr != "" {
if user, tokenRecord, err := getUserByToken(ctx, tokenStr); err == nil {
if user.Username == SystemUsername {
return nil, errors.New("system user is not allowed to login")
}
util.SetToContext(c, contracts.AuthTokenAuthKey, true)
util.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin)
return user, nil
}
}
// 降级使用 Session 鉴权
userID := GetUserIDFromContext(c)
if userID <= 0 {
return nil, errors.New("unauthorized")
}
user, err := GetCachedUser(ctx, userID)
if err != nil || user == nil || !user.IsActive {
var dbUser contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&dbUser).Error; err != nil {
return nil, err
}
user = &dbUser
SetCachedUser(ctx, userID, user)
}
util.SetToContext(c, contracts.AuthTokenAuthKey, false)
util.SetToContext(c, contracts.AuthTokenAdminKey, false)
if user.Username == "system" {
return nil, errors.New("system user is not allowed to login")
}
return user, nil
}
// LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session
func LoginRequired() gin.HandlerFunc {
return func(c *gin.Context) {
_, span := trace.Start(c.Request.Context(), "LoginRequired")
defer span.End()
user, err := GetUserFromRequest(c)
if err != nil {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
LogForAudit(c.Request.Context(), user, c)
util.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next()
}
}
// AdminRequired 校验管理员权限(支持 Session 和 Token 鉴权)
func AdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
_, span := trace.Start(c.Request.Context(), "AdminRequired")
defer span.End()
user, err := GetUserFromRequest(c)
if err != nil {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
isTokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey)
isTokenAdmin, _ := util.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
// 如果是通过 Token 鉴权,要求该 Token 具备管理员权限或者用户本身为管理员
if isTokenAuth && !isTokenAdmin && !user.IsAdmin {
response.AbortNotFound(c, errTokenAdminRequired)
return
}
// 如果是通过 Session 鉴权,直接检查用户的 is_admin 属性
if !isTokenAuth && !user.IsAdmin {
response.AbortNotFound(c, errAdminRequired)
return
}
LogForAudit(c.Request.Context(), user, c)
util.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next()
}
}
// LoginAdminRequired is an alias for AdminRequired.
func LoginAdminRequired() gin.HandlerFunc {
return AdminRequired()
}
// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点
func DisallowTokenAuth() gin.HandlerFunc {
return func(c *gin.Context) {
if tokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
response.AbortForbidden(c, ErrTokenAuthNotAllowed)
return
}
c.Next()
}
}
@@ -0,0 +1,58 @@
-- +goose Up
-- +goose StatementBegin
CREATE TABLE IF NOT EXISTS w_auth_sources (
id BIGINT PRIMARY KEY,
name VARCHAR(80) NOT NULL UNIQUE,
type VARCHAR(20) NOT NULL,
display_name VARCHAR(100),
is_active BOOLEAN NOT NULL DEFAULT FALSE,
client_id VARCHAR(255),
client_secret VARCHAR(1024),
openid_discovery_url VARCHAR(1024),
scopes VARCHAR(255),
icon_url VARCHAR(1024),
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_auth_sources_is_active ON w_auth_sources (is_active);
CREATE TABLE IF NOT EXISTS w_external_accounts (
id BIGINT PRIMARY KEY,
auth_source_id BIGINT,
user_id BIGINT NOT NULL,
external_id VARCHAR(255) NOT NULL,
external_username VARCHAR(255),
email VARCHAR(255),
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_external_accounts_auth_source_id ON w_external_accounts (auth_source_id);
CREATE INDEX IF NOT EXISTS idx_w_external_accounts_user_id ON w_external_accounts (user_id);
CREATE UNIQUE INDEX IF NOT EXISTS idx_w_external_accounts_source_external ON w_external_accounts (auth_source_id, external_id);
CREATE TABLE IF NOT EXISTS w_access_tokens (
id BIGINT PRIMARY KEY,
user_id BIGINT NOT NULL,
token_hash VARCHAR(64) NOT NULL UNIQUE,
name VARCHAR(128) NOT NULL,
description VARCHAR(255),
is_admin BOOLEAN DEFAULT FALSE,
expires_at TIMESTAMPTZ,
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_access_tokens_user_id ON w_access_tokens (user_id);
-- Seed: login session TTL
INSERT INTO w_system_configs (key, value, type, visibility, description, created_at, updated_at)
VALUES ('login_session_ttl_hours', '0', 'system', 0, '登录会话过期时间 (小时,0表示浏览器关闭后自动退出,-1表示永不过期)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
ON CONFLICT (key) DO NOTHING;
-- +goose StatementEnd
-- +goose Down
-- +goose StatementBegin
DELETE FROM w_system_configs WHERE key = 'login_session_ttl_hours';
DROP TABLE IF EXISTS w_access_tokens;
DROP TABLE IF EXISTS w_external_accounts;
DROP TABLE IF EXISTS w_auth_sources;
-- +goose StatementEnd
+243
View File
@@ -0,0 +1,243 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"encoding/json"
"errors"
"regexp"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
)
var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`)
// AuthSource 认证源实体
//
//nolint:revive // auth.AuthSource is standard domain entity name
type AuthSource struct {
ID uint64 `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"uniqueIndex;size:80;not null"`
Type string `json:"type" gorm:"size:20;not null"`
DisplayName string `json:"display_name" gorm:"size:100"`
IsActive bool `json:"is_active" gorm:"index;not null;default:false"`
ClientID string `json:"client_id" gorm:"size:255"`
ClientSecret string `json:"-" gorm:"size:1024"`
OpenIDDiscoveryURL string `json:"openid_discovery_url" gorm:"column:openid_discovery_url;size:1024"`
Scopes string `json:"scopes" gorm:"size:255"`
IconURL string `json:"icon_url" gorm:"size:1024"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
ClientSecretConfigured bool `json:"client_secret_configured" gorm:"-"`
}
// TableName 表名
func (AuthSource) TableName() string {
return "w_auth_sources"
}
// Normalize 对认证源字段进行标准化处理
func (source *AuthSource) Normalize() {
source.Type = strings.ToLower(strings.TrimSpace(source.Type))
source.Name = strings.TrimSpace(source.Name)
source.DisplayName = strings.TrimSpace(source.DisplayName)
source.ClientID = strings.TrimSpace(source.ClientID)
source.ClientSecret = strings.TrimSpace(source.ClientSecret)
source.OpenIDDiscoveryURL = strings.TrimSpace(source.OpenIDDiscoveryURL)
source.Scopes = strings.TrimSpace(source.Scopes)
source.IconURL = strings.TrimSpace(source.IconURL)
if source.DisplayName == "" {
source.DisplayName = source.Name
}
if source.Type == AuthSourceTypeOIDC && source.Scopes == "" {
source.Scopes = "openid profile email"
}
}
// Validate 校验认证源字段合法性
func (source *AuthSource) Validate() error {
source.Normalize()
if source.Name == "" {
return errors.New(errAuthSourceNameRequired)
}
if !authSourceNamePattern.MatchString(source.Name) {
return errors.New(errAuthSourceNameInvalid)
}
if source.Type != AuthSourceTypeOIDC {
return errors.New(errAuthSourceTypeUnsupported)
}
if source.OpenIDDiscoveryURL == "" {
//nolint:staticcheck // descriptive error constant
return errors.New(errAuthSourceDiscoveryURLRequired)
}
if source.IsActive && (source.ClientID == "" || source.ClientSecret == "") {
return errors.New(errAuthSourceClientCredentialsRequired)
}
return nil
}
// Sanitize 脱敏处理,将 ClientSecret 清空并设置 ClientSecretConfigured 标志
func (source *AuthSource) Sanitize() {
source.ClientSecretConfigured = source.ClientSecret != ""
source.ClientSecret = ""
}
// ExternalAccount 外部账号绑定实体
type ExternalAccount struct {
ID uint64 `json:"id" gorm:"primaryKey"`
AuthSourceID uint64 `json:"auth_source_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:1;index"`
UserID uint64 `json:"user_id" gorm:"index;not null"`
ExternalID string `json:"external_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:2;size:255;not null"`
ExternalUsername string `json:"external_username" gorm:"size:255"`
Email string `json:"email" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// TableName 表名
func (ExternalAccount) TableName() string {
return "w_external_accounts"
}
// ExternalAccountView 外部帐号绑定视图(脱敏展示用)
type ExternalAccountView struct {
ID uint64 `json:"id"`
AuthSourceID uint64 `json:"auth_source_id"`
AuthSourceName string `json:"auth_source_name"`
AuthSourceType string `json:"auth_source_type"`
AuthSourceLabel string `json:"auth_source_label"`
ExternalUsername string `json:"external_username"`
Email string `json:"email"`
CreatedAt time.Time `json:"created_at"`
}
// AuthSourceView 登录源展示信息
//
//nolint:revive // auth.AuthSourceView is standard domain presentation struct
type AuthSourceView struct {
ID uint64 `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
DisplayName string `json:"display_name"`
IsActive bool `json:"is_active"`
IconURL string `json:"icon_url"`
ClientSecretConfigured bool `json:"client_secret_configured"`
}
// OAuthAuthorizeResponse 授权 URL 响应
type OAuthAuthorizeResponse struct {
AuthorizeURL string `json:"authorize_url"`
}
// OAuthCallbackResult 回调处理结果
type OAuthCallbackResult struct {
Status string `json:"status"`
User *BasicUserInfo `json:"user,omitempty"`
}
// CallbackRequest OAuth 回调请求参数
type CallbackRequest struct {
State string `json:"state" binding:"required"`
Code string `json:"code" binding:"required"`
}
// BasicUserInfo 用户基本信息结构体
type BasicUserInfo struct {
ID uint64 `json:"id"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Email string `json:"email"`
AvatarURL string `json:"avatar_url"`
IsAdmin bool `json:"is_admin"`
NeedChangePassword bool `json:"need_change_password"`
Bio string `json:"bio"`
Phone string `json:"phone"`
Gender string `json:"gender"`
Website string `json:"website"`
Location string `json:"location"`
}
// BuildBasicUserInfo 将 UserDTO 转换为 BasicUserInfo
func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo {
if user == nil {
return BasicUserInfo{}
}
return BasicUserInfo{
ID: user.ID,
Username: user.Username,
Nickname: user.Nickname,
Email: user.Email,
AvatarURL: user.AvatarURL,
IsAdmin: user.IsAdmin,
NeedChangePassword: needChange,
Bio: user.Bio,
Phone: user.Phone,
Gender: user.Gender,
Website: user.Website,
Location: user.Location,
}
}
type oauthStatePayload struct {
SourceName string `json:"source_name"`
Purpose string `json:"purpose"`
UserID uint64 `json:"user_id,omitempty"`
SessionHash string `json:"session_hash"`
}
func encodeOAuthStatePayload(payload oauthStatePayload) (string, error) {
data, err := json.Marshal(payload)
if err != nil {
return "", err
}
return string(data), nil
}
func decodeOAuthStatePayload(value string) (oauthStatePayload, error) {
var payload oauthStatePayload
if err := json.Unmarshal([]byte(value), &payload); err != nil {
return oauthStatePayload{}, err
}
return payload, nil
}
type loginRequiredAuditLog struct {
UserID uint64 `json:"user_id"`
Username string `json:"username"`
ClientIP string `json:"client_ip"`
Method string `json:"method"`
Path string `json:"path"`
RequestURI string `json:"request_uri"`
UserAgent string `json:"user_agent"`
Referer string `json:"referer"`
}
// ParseUserID parses a string or float64 user ID representation.
func ParseUserID(v any) uint64 {
switch val := v.(type) {
case uint64:
return val
case int64:
if val > 0 {
return uint64(val)
}
case int:
if val > 0 {
return uint64(val)
}
case float64:
if val > 0 {
return uint64(val)
}
case string:
if id, err := strconv.ParseUint(val, 10, 64); err == nil {
return id
}
}
return 0
}
+114
View File
@@ -0,0 +1,114 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package auth provides the authentication, OAuth, session management, and access token domain plugin for Cordis.
package auth
import (
"embed"
"github.com/Rain-kl/Wavelet/backend/core"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/core/extpoints"
)
//go:embed migrations/*.sql
var authMigrations embed.FS
// Option configures the auth plugin.
type Option func(*Plugin)
// WithAuthService sets a custom AuthService implementation.
func WithAuthService(svc contracts.AuthService) Option {
return func(p *Plugin) {
p.authSvc = svc
}
}
// WithAuthRegistry sets a custom AuthRegistry implementation.
func WithAuthRegistry(reg contracts.AuthRegistry) Option {
return func(p *Plugin) {
p.authRegistry = reg
}
}
// Plugin implements core.Plugin to provide authentication and OAuth domain services.
type Plugin struct {
authSvc contracts.AuthService
authRegistry contracts.AuthRegistry
}
// New creates a new auth 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 auth domain plugin.
func (p *Plugin) Name() string {
return "auth"
}
// Manifest returns the plugin metadata.
func (p *Plugin) Manifest() core.Manifest {
return core.Manifest{
Name: "auth",
Version: "1.0.0",
Description: "Authentication, OAuth, Session and Passkey domain plugin",
Author: "Wavelet Team",
}
}
// Apply registers the auth migrations, services, routes, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
// 1. Register migrations
ctx.Migrations().Register("auth", authMigrations)
// 2. Initialize and provide AuthService & AuthRegistry
if p.authSvc == nil {
p.authSvc = newAuthService()
}
if p.authRegistry == nil {
p.authRegistry = newAuthRegistry()
}
core.Provide[contracts.AuthService](ctx, p.authSvc)
core.Provide[contracts.AuthRegistry](ctx, p.authRegistry)
// 3. Register HTTP Routes
oauthGroup := ctx.Router().Group("/api/v1/oauth")
{
oauthGroup.GET("/sources", GetLoginSources)
oauthGroup.GET("/login", GetLoginURL)
oauthGroup.GET("/:source/authorize", Authorize)
oauthGroup.GET("/logout", Logout)
oauthGroup.POST("/callback", Callback)
oauthGroup.GET("/user-info", LoginRequired(), UserInfo)
oauthGroup.GET("/external-accounts", LoginRequired(), ListExternalAccounts)
oauthGroup.POST("/external-accounts/:id/delete", LoginRequired(), DeleteExternalAccount)
}
ctx.Router().GET("/api/v1/user-info", LoginRequired(), UserInfo)
// 4. Register Settings Schemas
ctx.Settings().Register(extpoints.SettingSchema{
Key: "auth.session_age",
Default: 86400 * 7,
Description: "Default session lifetime in seconds",
Type: "integer",
Category: "security",
})
ctx.Settings().Register(extpoints.SettingSchema{
Key: "auth.login_rate_limit_max_attempts",
Default: 5,
Description: "Max login failure attempts before temporary IP lock",
Type: "integer",
Category: "security",
})
return nil
}
+140
View File
@@ -0,0 +1,140 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth_test
import (
"context"
"crypto/sha256"
"encoding/hex"
"path/filepath"
"testing"
"time"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/backend/core"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/auth"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
)
type testUser struct {
ID uint64 `gorm:"primaryKey"`
Username string
IsActive bool
LastLoginAt time.Time
}
func (testUser) TableName() string { return "w_users" }
type testAccessToken struct {
ID uint64 `gorm:"primaryKey"`
UserID uint64
TokenHash string
Name string
IsAdmin bool
}
func (testAccessToken) TableName() string { return "w_access_tokens" }
func hashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func setupTestDB(t *testing.T) *gorm.DB {
t.Helper()
dbPath := filepath.Join(t.TempDir(), "auth_test.db")
testDB, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, testDB.AutoMigrate(
&testUser{},
&testAccessToken{},
&auth.AuthSource{},
&auth.ExternalAccount{},
))
db.SetDB(testDB)
return testDB
}
type mockProvider struct{}
func (m *mockProvider) Name() string { return "custom" }
func (m *mockProvider) GetAuthURL(state string) string {
return "https://custom.com/auth?state=" + state
}
func (m *mockProvider) ExchangeCode(ctx context.Context, code string) (*contracts.OAuthUserInfoDTO, error) {
return &contracts.OAuthUserInfoDTO{
ID: 555,
Username: "custom_user",
Email: "custom@example.com",
}, nil
}
func TestAuthPluginUnit(t *testing.T) {
ctx := core.NewContext(context.Background())
testDB := setupTestDB(t)
p := auth.New()
assert.Equal(t, "auth", p.Name())
assert.Equal(t, "1.0.0", p.Manifest().Version)
require.NoError(t, p.Apply(ctx))
// Test AuthService injection
authSvc, err := core.Inject[contracts.AuthService](ctx)
require.NoError(t, err)
assert.NotNil(t, authSvc.RequireAuthMiddleware())
assert.NotNil(t, authSvc.RequireAdminMiddleware())
// Test AuthRegistry injection
authReg, err := core.Inject[contracts.AuthRegistry](ctx)
require.NoError(t, err)
authReg.RegisterOAuthProvider("custom", &mockProvider{})
prov, ok := authReg.GetOAuthProvider("custom")
require.True(t, ok)
assert.Equal(t, "custom", prov.Name())
// Test User Token Verification with dummy token
user := testUser{
ID: 101,
Username: "token_user",
IsActive: true,
}
require.NoError(t, testDB.Create(&user).Error)
tokenStr := "test-secret-token-123456"
tokenHash := hashToken(tokenStr)
tokenRecord := testAccessToken{
ID: 201,
UserID: user.ID,
TokenHash: tokenHash,
Name: "test-token",
IsAdmin: false,
}
require.NoError(t, testDB.Create(&tokenRecord).Error)
userDTO, err := authSvc.VerifyToken(context.Background(), tokenStr)
require.NoError(t, err)
assert.Equal(t, user.ID, userDTO.ID)
assert.Equal(t, "token_user", userDTO.Username)
// Empty token fails
_, err = authSvc.VerifyToken(context.Background(), "")
assert.Error(t, err)
// Revoke sessions
require.NoError(t, authSvc.RevokeUserSessions(context.Background(), user.ID))
// GetCurrentUser from context
userCtx := context.WithValue(context.Background(), contracts.AuthUserObjKey, userDTO)
current, err := authSvc.GetCurrentUser(userCtx)
require.NoError(t, err)
assert.Equal(t, user.ID, current.ID)
}
@@ -0,0 +1,81 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"net/http"
"sync"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
"golang.org/x/sync/singleflight"
)
// oidcProviderCache 进程级 OIDC provider 缓存。
type oidcProviderCache struct {
mu sync.RWMutex
entries map[string]*oidc.Provider // key: normalized issuer URL
sfGroup singleflight.Group
}
// globalOIDCProviderCache 是包级单例缓存,与进程同生命周期。
var globalOIDCProviderCache = &oidcProviderCache{
entries: make(map[string]*oidc.Provider),
}
// discoveryContext 从请求 ctx 提取 HTTP 客户端,并绑定到不可取消的 Background ctx。
func discoveryContext(ctx context.Context) context.Context {
bg := context.Background()
if client, ok := ctx.Value(oauth2.HTTPClient).(*http.Client); ok && client != nil {
bg = oidc.ClientContext(bg, client)
}
return bg
}
// get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。
func (c *oidcProviderCache) get(ctx context.Context, issuer string) (*oidc.Provider, error) {
c.mu.RLock()
if p, ok := c.entries[issuer]; ok {
c.mu.RUnlock()
return p, nil
}
c.mu.RUnlock()
discCtx := discoveryContext(ctx)
v, err, _ := c.sfGroup.Do(issuer, func() (any, error) {
c.mu.RLock()
if p, ok := c.entries[issuer]; ok {
c.mu.RUnlock()
return p, nil
}
c.mu.RUnlock()
p, err := oidc.NewProvider(discCtx, issuer)
if err != nil {
return nil, err
}
c.mu.Lock()
c.entries[issuer] = p
c.mu.Unlock()
return p, nil
})
if err != nil {
return nil, err
}
return v.(*oidc.Provider), nil //nolint:forcetypeassert
}
// invalidate 从缓存中移除指定 issuer 对应的 provider。
func (c *oidcProviderCache) invalidate(issuer string) {
c.mu.Lock()
delete(c.entries, issuer)
c.mu.Unlock()
}
// InvalidateOIDCProviderCache 从进程级缓存中清除指定 issuer 的 provider 条目。
func InvalidateOIDCProviderCache(issuer string) {
globalOIDCProviderCache.invalidate(issuer)
}
+76
View File
@@ -0,0 +1,76 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
)
// GetAuthSourceByID 根据 ID 获取认证源
func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
var src AuthSource
if err := db.DB(ctx).First(&src, id).Error; err != nil {
return nil, err
}
return &src, nil
}
// GetAuthSourceByName 根据名称获取认证源
func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) {
var src AuthSource
if err := db.DB(ctx).Where("name = ?", name).First(&src).Error; err != nil {
return nil, err
}
return &src, nil
}
// ListActiveAuthSources 获取所有启用的认证源
func ListActiveAuthSources(ctx context.Context) ([]AuthSource, error) {
var sources []AuthSource
if err := db.DB(ctx).Where("is_active = ?", true).Order("id ASC").Find(&sources).Error; err != nil {
return nil, err
}
return sources, nil
}
// GetActiveAuthSourcesCached 获取所有启用的认证源(带缓存或直接查询)
func GetActiveAuthSourcesCached(ctx context.Context) ([]AuthSource, error) {
return ListActiveAuthSources(ctx)
}
// GetAuthSourceByNameCached 根据名称获取认证源(带缓存或直接查询)
func GetAuthSourceByNameCached(ctx context.Context, name string) (*AuthSource, error) {
return GetAuthSourceByName(ctx, name)
}
// FindExternalAccount 查询指定认证源的外部账号绑定
func FindExternalAccount(ctx context.Context, authSourceID uint64, externalID string) (*ExternalAccount, error) {
var account ExternalAccount
if err := db.DB(ctx).Where("auth_source_id = ? AND external_id = ?", authSourceID, externalID).First(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
// BindExternalAccount 绑定外部账号
func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
return db.DB(ctx).Create(account).Error
}
// ListExternalAccountsByUserID 获取用户绑定的所有外部账号
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccount, error) {
var accounts []ExternalAccount
if err := db.DB(ctx).Where("user_id = ?", userID).Find(&accounts).Error; err != nil {
return nil, err
}
return accounts, nil
}
// UnbindExternalAccount 解绑外部账号
func UnbindExternalAccount(ctx context.Context, id uint64, userID uint64) error {
return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
}
+145
View File
@@ -0,0 +1,145 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"errors"
"sync"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
"github.com/gin-gonic/gin"
)
type authServiceImpl struct{}
func newAuthService() contracts.AuthService {
return &authServiceImpl{}
}
func (s *authServiceImpl) RequireAuthMiddleware() any {
return LoginRequired()
}
func (s *authServiceImpl) RequireAdminMiddleware() any {
return AdminRequired()
}
func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if ginCtx, ok := ctx.(*gin.Context); ok {
if u, ok := util.GetFromContext[*contracts.UserDTO](ginCtx, contracts.AuthUserObjKey); ok && u != nil {
return u, nil
}
}
if v := ctx.Value(contracts.AuthUserObjKey); v != nil {
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
return u, nil
}
}
return nil, errors.New("auth: user not found in context")
}
func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contracts.UserDTO, error) {
if token == "" {
return nil, errors.New("auth: empty token")
}
tokenHash := hashToken(token)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
var tokenRow struct {
ID uint64
UserID uint64
IsAdmin bool
}
if err := db.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
return nil, err
}
tokenRecord = &CachedToken{
ID: tokenRow.ID,
UserID: tokenRow.UserID,
IsAdmin: tokenRow.IsAdmin,
}
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
var dbUser contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil {
return nil, err
}
user = &dbUser
SetCachedUser(ctx, tokenRecord.UserID, user)
}
if user.Username == SystemUsername {
return nil, errors.New("auth: system user token not allowed")
}
return user, nil
}
func (s *authServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
return "", nil
}
func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error {
InvalidateCachedUser(ctx, userID)
return nil
}
func (s *authServiceImpl) GetCurrentUserID(ctx context.Context) (uint64, error) {
if ginCtx, ok := ctx.(*gin.Context); ok {
return GetUserIDFromContext(ginCtx), nil
}
return 0, errors.New("auth: user not found in context")
}
func (s *authServiceImpl) RevokeToken(ctx context.Context, tokenHash string) error {
InvalidateCachedToken(ctx, tokenHash)
return nil
}
func (s *authServiceImpl) DisallowTokenAuthMiddleware() any {
return DisallowTokenAuth()
}
type authRegistryImpl struct {
mu sync.RWMutex
providers map[string]contracts.OAuthProvider
}
func newAuthRegistry() contracts.AuthRegistry {
return &authRegistryImpl{
providers: make(map[string]contracts.OAuthProvider),
}
}
func (r *authRegistryImpl) RegisterOAuthProvider(name string, provider contracts.OAuthProvider) {
r.mu.Lock()
defer r.mu.Unlock()
r.providers[name] = provider
}
func (r *authRegistryImpl) GetOAuthProvider(name string) (contracts.OAuthProvider, bool) {
r.mu.RLock()
defer r.mu.RUnlock()
p, ok := r.providers[name]
return p, ok
}
func (r *authRegistryImpl) ListOAuthProviders() []string {
r.mu.RLock()
defer r.mu.RUnlock()
res := make([]string, 0, len(r.providers))
for name := range r.providers {
res = append(res, name)
}
return res
}
+145
View File
@@ -0,0 +1,145 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"crypto/sha256"
"encoding/hex"
"net/http"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/pkg/config"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
gsessions "github.com/gorilla/sessions"
)
// GetSessionOptions 根据配置构建 Session 选项
func GetSessionOptions(maxAge int) sessions.Options {
return sessions.Options{
Path: "/",
Domain: config.Config.App.SessionDomain,
MaxAge: maxAge,
HttpOnly: config.Config.App.SessionHTTPOnly,
Secure: config.Config.App.SessionSecure,
SameSite: http.SameSiteLaxMode,
}
}
// StripCookieMaxAgeAndExpires 从 Set-Cookie 响应头中移除 Max-Age 和 Expires,从而使其成为浏览器会话 Cookie
func StripCookieMaxAgeAndExpires(header http.Header, cookieName string) {
headers := header["Set-Cookie"]
if len(headers) == 0 {
return
}
newHeaders := make([]string, 0, len(headers))
for _, h := range headers {
if strings.HasPrefix(h, cookieName+"=") {
parts := strings.Split(h, ";")
newParts := make([]string, 0, len(parts))
for _, p := range parts {
trimmed := strings.TrimSpace(p)
lower := strings.ToLower(trimmed)
if strings.HasPrefix(lower, "max-age=") || strings.HasPrefix(lower, "expires=") {
continue
}
newParts = append(newParts, p)
}
newHeaders = append(newHeaders, strings.Join(newParts, ";"))
} else {
newHeaders = append(newHeaders, h)
}
}
header["Set-Cookie"] = newHeaders
}
// GetUserIDFromSession 从 Session 中提取用户 ID
func GetUserIDFromSession(s sessions.Session) uint64 {
val := s.Get(UserIDKey)
return ParseUserID(val)
}
// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID
func GetUserIDFromContext(c *gin.Context) (uid uint64) {
defer func() {
_ = recover()
}()
session := sessions.Default(c)
return GetUserIDFromSession(session)
}
func ensureSessionToken(s sessions.Session) (string, bool) {
token, ok := s.Get(SessionTokenKey).(string)
if !ok || token == "" {
token = uuid.NewString()
s.Set(SessionTokenKey, token)
return token, true
}
return token, false
}
func hashSessionToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func rotateSessionID(s sessions.Session) {
if inner, ok := s.(interface{ Session() *gsessions.Session }); ok {
if sess := inner.Session(); sess != nil {
sess.ID = ""
}
}
}
// SetLoginSession writes the authenticated user into a freshly rotated session.
func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDTO, extras ...map[string]any) error {
session := sessions.Default(c)
session.Clear()
rotateSessionID(session)
session.Set(UserIDKey, user.ID)
session.Set(UserNameKey, user.Username)
if len(extras) > 0 {
for key, value := range extras[0] {
session.Set(key, value)
}
}
// 根据系统配置动态设置 Session 过期时间
maxAge := config.Config.App.SessionAge
isSessionCookie := false
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "login_session_ttl_hours").Pluck("value", &val).Error; err == nil && val != "" {
if ttlHours, err := strconv.Atoi(val); err == nil {
switch {
case ttlHours == -1:
// 永不过期,设置为 10 年
maxAge = 10 * 365 * 24 * 3600
case ttlHours > 0:
maxAge = ttlHours * 3600
case ttlHours == 0:
isSessionCookie = true
}
}
}
session.Options(GetSessionOptions(maxAge))
if err := session.Save(); err != nil {
return err
}
if isSessionCookie {
StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName)
}
return nil
}
+10
View File
@@ -0,0 +1,10 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package cap 提供人机验证中间件
package cap
const (
errCapTokenMissing = "验证码验证失败,缺少验证码凭证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errCapTokenInvalidOrExpired = "验证码校验失败或已过期,请重试" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
)
+101
View File
@@ -0,0 +1,101 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"net/http"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/cap/pow"
"github.com/gin-gonic/gin"
)
// ChallengeResponse is a local type alias for the pow.ChallengeResponse struct
type ChallengeResponse = pow.ChallengeResponse
type challengeRequest struct {
Scope string `json:"scope" form:"scope"`
}
type redeemRequest struct {
Token string `json:"token" binding:"required"`
Solutions []int `json:"solutions" binding:"required"`
Scope string `json:"scope" form:"scope"`
}
// Challenge 生成 PoW 人机验证难题
// @Summary 生成人机验证难题
// @Description 客户端获取 PoW 难题和签名的 JWT Token,并在后台计算。
// @Tags cap
// @Accept json
// @Produce json
// @Param request body challengeRequest false "可选范围限制参数"
// @Success 200 {object} response.Any{data=cap.ChallengeResponse} "成功返回 PoW 难题"
// @Failure 500 {object} response.Any "内部服务错误"
// @Router /api/cap/challenge [post]
func Challenge(c *gin.Context) {
var req challengeRequest
_ = c.ShouldBind(&req) // 允许不传 body,默认使用 login scope
if req.Scope == "" {
req.Scope = "login"
}
mgr := GetDefaultManager()
if mgr == nil {
response.AbortInternal(c, "captcha is not configured")
return
}
resp, err := mgr.Generate(c.Request.Context(), req.Scope)
if err != nil {
logger.ErrorF(c.Request.Context(), "Generate cap challenge failed: %v", err)
response.AbortInternal(c, "生成验证难题失败,请稍后再试")
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
// Redeem 提交 PoW 解答并兑换一次性凭证 Token
// @Summary 校验人机验证解答
// @Description 提交 PoW 解答进行核销,成功后返回一次性 X-Cap-Token 凭证
// @Tags cap
// @Accept json
// @Produce json
// @Param request body redeemRequest true "难题 Token 与解答 solutions 数组"
// @Success 200 {object} response.Any{data=cap.RedeemResponse} "核销成功,返回 X-Cap-Token"
// @Failure 400 {object} response.Any "参数错误或核销失败"
// @Failure 500 {object} response.Any "内部服务错误"
// @Router /api/cap/redeem [post]
func Redeem(c *gin.Context) {
var req redeemRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, "无效的参数")
return
}
if req.Scope == "" {
req.Scope = "login"
}
mgr := GetDefaultManager()
if mgr == nil {
response.AbortInternal(c, "captcha is not configured")
return
}
resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
if err != nil {
logger.ErrorF(c.Request.Context(), "Redeem cap solutions failed: %v", err)
response.AbortInternal(c, "校验验证解答失败,请稍后再试")
return
}
if !resp.Success {
response.AbortBadRequest(c, resp.Error)
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
+199
View File
@@ -0,0 +1,199 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package cap provides CAPTCHA and proof-of-work (PoW) verification services.
package cap
import (
"context"
"crypto/sha256"
"encoding/hex"
"strconv"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/backend/pkg/config"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/cap/pow"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
)
const (
redeemTokenIDLength = 8 // 兑换 Token ID 字节长度
redeemVerTokenLength = 15 // 兑换验证 Token 字节长度
tokenPartsCount = 2 // 兑换 Token 由两部分组成
valuePartsCount = 2 // 存储值由 scope 和过期时间组成
)
// Manager orchestrates challenge generation and solution validation.
type Manager struct {
secret []byte
store pow.Store
}
// NewManager creates a new CAPTCHA Manager.
func NewManager(secret []byte, store pow.Store) *Manager {
return &Manager{
secret: secret,
store: store,
}
}
// Generate creates a challenge response.
func (m *Manager) Generate(ctx context.Context, scope string) (*pow.ChallengeResponse, error) {
settings, err := CurrentSettings(ctx)
if err != nil {
return nil, err
}
challengeConfig := pow.ChallengeConfig{
Count: settings.ChallengeCount,
Size: settings.ChallengeSize,
Difficulty: settings.ChallengeDifficulty,
Expires: settings.ChallengeTTL,
}
return pow.GenerateChallenge(m.secret, challengeConfig, scope)
}
// RedeemResponse is returned to the client on redeem.
type RedeemResponse struct {
Success bool `json:"success"`
Token string `json:"token,omitempty"`
Expires int64 `json:"expires,omitempty"`
Error string `json:"error,omitempty"`
}
// Redeem verifies PoW solutions and returns a one-time redeem token.
func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*RedeemResponse, error) {
sigHex := pow.JwtSigHex(token)
if sigHex == "" {
return &RedeemResponse{Success: false, Error: "invalid_token"}, nil
}
nonceKey := "cap:nonce:" + sigHex
payload, err := pow.VerifyChallengeSolutions(token, solutions, m.secret, scope)
if err != nil {
return &RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // validation errors are returned as response, not system errors
}
now := time.Now().UnixNano() / int64(time.Millisecond)
nonceTTL := time.Duration(payload.Expires-now) * time.Millisecond
if nonceTTL < time.Second {
nonceTTL = time.Second
}
set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL)
if err != nil {
return &RedeemResponse{Success: false, Error: "nonce_store_error"}, err
}
if !set {
return &RedeemResponse{Success: false, Error: "already_redeemed"}, nil
}
settings, err := CurrentSettings(ctx)
if err != nil {
return &RedeemResponse{Success: false, Error: "settings_load_error"}, err
}
id := pow.RandomHex(redeemTokenIDLength)
verToken := pow.RandomHex(redeemVerTokenLength)
verHashBytes := sha256.Sum256([]byte(verToken))
verHashHex := hex.EncodeToString(verHashBytes[:])
tokenKey := "cap:token:" + id + ":" + verHashHex
tokenExpires := time.Now().Add(settings.TokenTTL)
storeVal := strconv.FormatInt(tokenExpires.UnixNano(), 10) + "|" + scope
if err := m.store.Set(ctx, tokenKey, storeVal, settings.TokenTTL); err != nil {
return &RedeemResponse{Success: false, Error: "token_store_error"}, err
}
return &RedeemResponse{
Success: true,
Token: id + ":" + verToken,
Expires: tokenExpires.UnixNano() / int64(time.Millisecond),
}, nil
}
// VerifyToken validates and consumes the redeem token (single-use).
func (m *Manager) VerifyToken(ctx context.Context, token string, expectedScope string) (bool, error) {
if token == "" {
return false, nil
}
parts := strings.Split(token, ":")
if len(parts) != tokenPartsCount {
return false, nil
}
id := parts[0]
verToken := parts[1]
verHashBytes := sha256.Sum256([]byte(verToken))
verHashHex := hex.EncodeToString(verHashBytes[:])
tokenKey := "cap:token:" + id + ":" + verHashHex
val, exists, err := sGetAndDelete(ctx, m.store, tokenKey)
if err != nil {
return false, err
}
if !exists {
return false, nil
}
valParts := strings.Split(val, "|")
if len(valParts) != valuePartsCount {
return false, nil
}
expNano, err := strconv.ParseInt(valParts[0], 10, 64)
if err != nil {
return false, nil //nolint:nilerr // invalid format is treated as validation failure
}
tokenScope := valParts[1]
if expectedScope != "" && tokenScope != expectedScope {
return false, nil
}
if time.Now().UnixNano() > expNano {
return false, nil
}
return true, nil
}
func sGetAndDelete(ctx context.Context, store pow.Store, key string) (string, bool, error) {
if store == nil {
return "", false, nil
}
return store.GetAndDelete(ctx, key)
}
var (
defaultManager *Manager
once sync.Once
)
// GetDefaultManager yields the global singleton CAPTCHA manager.
func GetDefaultManager() *Manager {
once.Do(func() {
var secret []byte
if config.Config != nil && strings.TrimSpace(config.Config.App.SessionSecret) != "" {
secret = []byte(config.Config.App.SessionSecret)
}
if len(secret) == 0 {
return
}
var store pow.Store
if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil {
store = pow.NewRedisStore(db.Redis)
} else {
store = pow.NewMemoryStore(1 * time.Minute)
}
defaultManager = NewManager(secret, store)
})
return defaultManager
}
+38
View File
@@ -0,0 +1,38 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
)
// VerifyMiddleware returns a Gin middleware that checks and consumes the X-Cap-Token header.
func VerifyMiddleware(mgr *Manager, scope string) gin.HandlerFunc {
return func(c *gin.Context) {
if !ProtectionEnabled(c.Request.Context()) {
c.Next()
return
}
if mgr == nil {
response.AbortUnauthorized(c, errCapTokenInvalidOrExpired)
return
}
token := c.GetHeader("X-Cap-Token")
if token == "" {
response.AbortUnauthorized(c, errCapTokenMissing)
return
}
valid, err := mgr.VerifyToken(c.Request.Context(), token, scope)
if err != nil || !valid {
response.AbortUnauthorized(c, errCapTokenInvalidOrExpired)
return
}
c.Next()
}
}
+62
View File
@@ -0,0 +1,62 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package cap provides the proof-of-work (PoW) CAPTCHA verification domain plugin for Cordis.
package cap
import (
"github.com/Rain-kl/Wavelet/backend/core"
"github.com/Rain-kl/Wavelet/backend/core/extpoints"
)
// Plugin implements core.Plugin to provide CAPTCHA generation, validation, and route protection.
type Plugin struct{}
// New creates a new cap domain plugin.
func New() *Plugin {
return &Plugin{}
}
// Name returns the unique identifier for the cap domain plugin.
func (p *Plugin) Name() string {
return "cap"
}
// Manifest returns the plugin metadata.
func (p *Plugin) Manifest() core.Manifest {
return core.Manifest{
Name: "cap",
Version: "1.0.0",
Description: "Proof-of-work CAPTCHA challenge and verification domain plugin",
Author: "Wavelet Team",
}
}
// Apply registers the cap routes and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
// Register HTTP Routes
capGroup := ctx.Router().Group("/api/v1/cap")
{
capGroup.GET("/challenge", Challenge)
capGroup.POST("/challenge", Challenge)
capGroup.POST("/redeem", Redeem)
}
// Register Settings Schemas
ctx.Settings().Register(extpoints.SettingSchema{
Key: "cap.login_enabled",
Default: false,
Description: "Whether to require CAPTCHA verification for user login",
Type: "boolean",
Category: "security",
})
ctx.Settings().Register(extpoints.SettingSchema{
Key: "cap.challenge_count",
Default: 1,
Description: "Number of PoW puzzle challenges to solve",
Type: "integer",
Category: "security",
})
return nil
}
+259
View File
@@ -0,0 +1,259 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package pow provides proof-of-work challenge generation and verification.
package pow
import (
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"encoding/json"
"errors"
"strconv"
"strings"
"time"
)
const (
jwtHeaderB64 = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9"
jwtPartsCount = 3 // JWT 三段结构
defaultChallengeCount = 50 // 默认 PoW 难题数
defaultChallengeSize = 32 // 默认盐值长度
defaultDifficulty = 4 // 默认难度
defaultNonceLength = 25 // 随机 Nonce 字节长度
defaultExpires = 10 * time.Minute // 默认过期时间
)
// ChallengeConfig holds parameters for the PoW challenge
type ChallengeConfig struct {
Count int // Number of puzzles (c)
Size int // Salt length (s)
Difficulty int // Difficulty prefix length (d)
Expires time.Duration // Challenge TTL
}
// ChallengeResponse is returned to the client
type ChallengeResponse struct {
Challenge struct {
C int `json:"c"`
S int `json:"s"`
D int `json:"d"`
} `json:"challenge"`
Token string `json:"token"`
Expires int64 `json:"expires"` // ms timestamp
}
// ChallengePayload represents the signed JWT payload
type ChallengePayload struct {
Nonce string `json:"n"`
Count int `json:"c"`
Size int `json:"s"`
Difficulty int `json:"d"`
Expires int64 `json:"exp"` // ms timestamp
IssuedAt int64 `json:"iat"` // ms timestamp
Scope string `json:"sk,omitempty"`
}
// RedeemRequest payload sent by client
type RedeemRequest struct {
Token string `json:"token"`
Solutions []int `json:"solutions"`
}
// RedeemResponse returned to client after verification
type RedeemResponse struct {
Success bool `json:"success"`
Token string `json:"token,omitempty"`
Expires int64 `json:"expires,omitempty"`
Error string `json:"error,omitempty"`
}
func b64urlEncode(data []byte) string {
return base64.RawURLEncoding.EncodeToString(data)
}
func b64urlDecode(str string) ([]byte, error) {
return base64.RawURLEncoding.DecodeString(str)
}
// RandomHex generates a cryptographically secure random hexadecimal string of the specified byte length.
func RandomHex(byteLen int) string {
bytes := make([]byte, byteLen)
if _, err := rand.Read(bytes); err != nil {
panic(err)
}
return hex.EncodeToString(bytes)
}
func jwtSign(payload []byte, secret []byte) string {
body := b64urlEncode(payload)
sigInput := jwtHeaderB64 + "." + body
mac := hmac.New(sha256.New, secret)
mac.Write([]byte(sigInput))
sig := mac.Sum(nil)
return sigInput + "." + b64urlEncode(sig)
}
func jwtVerify(token string, secret []byte) ([]byte, error) {
parts := strings.Split(token, ".")
if len(parts) != jwtPartsCount {
return nil, errors.New(errInvalidTokenFormat)
}
if parts[0] != jwtHeaderB64 {
return nil, errors.New(errInvalidHeader)
}
sigInput := parts[0] + "." + parts[1]
mac := hmac.New(sha256.New, secret)
mac.Write([]byte(sigInput))
expectedSig := mac.Sum(nil)
actualSig, err := b64urlDecode(parts[2])
if err != nil {
return nil, err
}
if !hmac.Equal(expectedSig, actualSig) {
return nil, errors.New(errSignatureMismatch)
}
payload, err := b64urlDecode(parts[1])
if err != nil {
return nil, err
}
return payload, nil
}
// JwtSigHex extracts the signature part of a JWT token and returns it as a hexadecimal string.
func JwtSigHex(token string) string {
parts := strings.Split(token, ".")
if len(parts) != jwtPartsCount {
return ""
}
sigBytes, err := b64urlDecode(parts[2])
if err != nil {
return ""
}
return hex.EncodeToString(sigBytes)
}
// GenerateChallenge produces a new challenge and signed token
func GenerateChallenge(secret []byte, conf ChallengeConfig, scope string) (*ChallengeResponse, error) {
if conf.Count <= 0 {
conf.Count = defaultChallengeCount
}
if conf.Size <= 0 {
conf.Size = defaultChallengeSize
}
if conf.Difficulty <= 0 {
conf.Difficulty = defaultDifficulty
}
if conf.Expires <= 0 {
conf.Expires = defaultExpires
}
now := time.Now().UnixNano() / int64(time.Millisecond)
expires := now + int64(conf.Expires/time.Millisecond)
payload := ChallengePayload{
Nonce: RandomHex(defaultNonceLength),
Count: conf.Count,
Size: conf.Size,
Difficulty: conf.Difficulty,
Expires: expires,
IssuedAt: now,
Scope: scope,
}
payloadBytes, err := json.Marshal(payload)
if err != nil {
return nil, err
}
token := jwtSign(payloadBytes, secret)
resp := &ChallengeResponse{
Token: token,
Expires: expires,
}
resp.Challenge.C = conf.Count
resp.Challenge.S = conf.Size
resp.Challenge.D = conf.Difficulty
return resp, nil
}
// VerifyChallengeSolutions verifies client submitted solutions
func VerifyChallengeSolutions(token string, solutions []int, secret []byte, expectedScope string) (*ChallengePayload, error) {
payloadBytes, err := jwtVerify(token, secret)
if err != nil {
return nil, errors.New(errInvalidToken)
}
var payload ChallengePayload
if err := json.Unmarshal(payloadBytes, &payload); err != nil {
return nil, errors.New(errInvalidToken)
}
if expectedScope != "" && payload.Scope != expectedScope {
return nil, errors.New(errScopeMismatch)
}
now := time.Now().UnixNano() / int64(time.Millisecond)
if payload.Expires < now {
return nil, errors.New(errExpired)
}
if len(solutions) != payload.Count {
return nil, errors.New(errInvalidSolutions)
}
tokenFnv := fnv1a(token)
for i := 0; i < payload.Count; i++ {
idxStr := strconv.Itoa(i + 1)
saltSeed := fnv1aResume(tokenFnv, idxStr)
targetSeed := fnv1aResume(saltSeed, "d")
salt := prngFromHash(saltSeed, payload.Size)
target := prngFromHash(targetSeed, payload.Difficulty)
hashInput := salt + strconv.Itoa(solutions[i])
hashBytes := sha256.Sum256([]byte(hashInput))
hashHex := hex.EncodeToString(hashBytes[:])
if !strings.HasPrefix(hashHex, target) {
return nil, errors.New(errInvalidSolution)
}
}
return &payload, nil
}
// Solve is a utility function to solve a challenge (mainly used for tests and reference implementation)
func Solve(token string, count, size, difficulty int) []int {
solutions := make([]int, count)
tokenFnv := fnv1a(token)
for i := 0; i < count; i++ {
idxStr := strconv.Itoa(i + 1)
saltSeed := fnv1aResume(tokenFnv, idxStr)
targetSeed := fnv1aResume(saltSeed, "d")
salt := prngFromHash(saltSeed, size)
target := prngFromHash(targetSeed, difficulty)
for nonce := 0; nonce < 1000000; nonce++ {
hashInput := salt + strconv.Itoa(nonce)
hashBytes := sha256.Sum256([]byte(hashInput))
hashHex := hex.EncodeToString(hashBytes[:])
if strings.HasPrefix(hashHex, target) {
solutions[i] = nonce
break
}
}
}
return solutions
}
+15
View File
@@ -0,0 +1,15 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pow
const (
errInvalidTokenFormat = "invalid token format"
errInvalidHeader = "invalid header"
errSignatureMismatch = "signature mismatch"
errInvalidToken = "invalid_token"
errScopeMismatch = "scope_mismatch"
errExpired = "expired"
errInvalidSolutions = "invalid_solutions"
errInvalidSolution = "invalid_solution"
)
+111
View File
@@ -0,0 +1,111 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pow
import (
"context"
"testing"
"time"
)
func TestPowChallengeFlow(t *testing.T) {
secret := []byte("test-secret-key-1234567890123456")
conf := ChallengeConfig{
Count: 2,
Size: 16,
Difficulty: 1,
Expires: 1 * time.Minute,
}
scope := "login"
resp, err := GenerateChallenge(secret, conf, scope)
if err != nil {
t.Fatalf("GenerateChallenge failed: %v", err)
}
if resp.Token == "" {
t.Fatal("expected non-empty token")
}
if resp.Challenge.C != 2 {
t.Fatalf("expected count 2, got %d", resp.Challenge.C)
}
sigHex := JwtSigHex(resp.Token)
if sigHex == "" {
t.Fatal("expected non-empty sigHex")
}
solutions := Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
if len(solutions) != 2 {
t.Fatalf("expected 2 solutions, got %d", len(solutions))
}
payload, err := VerifyChallengeSolutions(resp.Token, solutions, secret, scope)
if err != nil {
t.Fatalf("VerifyChallengeSolutions failed: %v", err)
}
if payload.Scope != scope {
t.Fatalf("expected scope %s, got %s", scope, payload.Scope)
}
// Scope mismatch test
_, err = VerifyChallengeSolutions(resp.Token, solutions, secret, "other_scope")
if err == nil {
t.Fatal("expected scope mismatch error")
}
// Invalid solutions test
_, err = VerifyChallengeSolutions(resp.Token, []int{9999999, 9999999}, secret, scope)
if err == nil {
t.Fatal("expected invalid solution error")
}
// Invalid token test
_, err = VerifyChallengeSolutions("invalid.jwt.token", solutions, secret, scope)
if err == nil {
t.Fatal("expected invalid token error")
}
}
func TestMemoryStore(t *testing.T) {
ctx := context.Background()
store := NewMemoryStore(100 * time.Millisecond)
// Set and Get
err := store.Set(ctx, "k1", "v1", 200*time.Millisecond)
if err != nil {
t.Fatalf("Set failed: %v", err)
}
val, ok, err := store.Get(ctx, "k1")
if err != nil || !ok || val != "v1" {
t.Fatalf("Get failed: val=%s, ok=%v, err=%v", val, ok, err)
}
// SetNX
set, err := store.SetNX(ctx, "k1", "v2", 200*time.Millisecond)
if err != nil || set {
t.Fatalf("SetNX should have failed because key exists: set=%v, err=%v", set, err)
}
set, err = store.SetNX(ctx, "k2", "v2", 200*time.Millisecond)
if err != nil || !set {
t.Fatalf("SetNX should have succeeded: set=%v, err=%v", set, err)
}
// GetAndDelete
val, ok, err = store.GetAndDelete(ctx, "k2")
if err != nil || !ok || val != "v2" {
t.Fatalf("GetAndDelete failed: val=%s, ok=%v, err=%v", val, ok, err)
}
_, ok, _ = store.Get(ctx, "k2")
if ok {
t.Fatal("k2 should be deleted")
}
// Delete
_ = store.Delete(ctx, "k1")
_, ok, _ = store.Get(ctx, "k1")
if ok {
t.Fatal("k1 should be deleted")
}
}
+56
View File
@@ -0,0 +1,56 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pow
import (
"fmt"
"strings"
)
// fnv1a returns the 32-bit FNV-1a hash of a string
//
//nolint:mnd // FNV-1a 算法位移常量
func fnv1a(str string) uint32 {
var hash uint32 = 2166136261
for i := 0; i < len(str); i++ {
hash ^= uint32(str[i])
hash += (hash << 1) + (hash << 4) + (hash << 7) + (hash << 8) + (hash << 24)
}
return hash
}
// fnv1aResume resumes FNV-1a hashing from a given state
//
//nolint:mnd // FNV-1a 算法位移常量
func fnv1aResume(state uint32, str string) uint32 {
h := state
for i := 0; i < len(str); i++ {
h ^= uint32(str[i])
h += (hashShift(h))
}
return h
}
// hashShift computes FNV-1a mix additions
//
//nolint:mnd
func hashShift(h uint32) uint32 {
return (h << 1) + (h << 4) + (h << 7) + (h << 8) + (h << 24)
}
// prngFromHash generates a hex string of specified length using an initial hash state
//
//nolint:mnd // xorshift 算法位移常量
func prngFromHash(initialHash uint32, length int) string {
state := initialHash
var result strings.Builder
for result.Len() < length {
state ^= state << 13
state ^= state >> 17
state ^= state << 5
hexStr := fmt.Sprintf("%08x", state)
result.WriteString(hexStr)
}
return result.String()[:length]
}
+186
View File
@@ -0,0 +1,186 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pow
import (
"context"
"sync"
"time"
"github.com/redis/go-redis/v9"
)
// Store defines the storage interface for challenge nonces and verification tokens
type Store interface {
Get(ctx context.Context, key string) (string, bool, error)
Set(ctx context.Context, key string, val string, ttl time.Duration) error
Delete(ctx context.Context, key string) error
// SetNX atomically sets key=val with the given TTL only when the key does not
// exist yet. It returns true when the key was actually written (i.e. this
// caller "won" the race), and false when the key already existed.
SetNX(ctx context.Context, key string, val string, ttl time.Duration) (bool, error)
// GetAndDelete atomically retrieves the value of key and removes it in a
// single operation. Returns ("", false, nil) when the key does not exist.
GetAndDelete(ctx context.Context, key string) (string, bool, error)
}
type memoryItem struct {
value string
expiresAt time.Time
}
// MemoryStore is a thread-safe in-memory implementation of Store
type MemoryStore struct {
items map[string]memoryItem
mu sync.Mutex // unified write-lock; promotes to exclusive for all ops
}
// NewMemoryStore creates and initializes a new MemoryStore
func NewMemoryStore(cleanupInterval time.Duration) *MemoryStore {
store := &MemoryStore{
items: make(map[string]memoryItem),
}
if cleanupInterval > 0 {
go store.startCleanupLoop(cleanupInterval)
}
return store
}
// Get 从 MemoryStore 获取指定 key 的值
func (s *MemoryStore) Get(_ context.Context, key string) (string, bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
return s.getLocked(key)
}
// getLocked is the internal helper – caller must hold s.mu.
func (s *MemoryStore) getLocked(key string) (string, bool, error) {
item, found := s.items[key]
if !found {
return "", false, nil
}
if time.Now().After(item.expiresAt) {
delete(s.items, key)
return "", false, nil
}
return item.value, true, nil
}
// Set 向 MemoryStore 写入指定 key 的值
func (s *MemoryStore) Set(_ context.Context, key string, val string, ttl time.Duration) error {
s.mu.Lock()
defer s.mu.Unlock()
s.items[key] = memoryItem{
value: val,
expiresAt: time.Now().Add(ttl),
}
return nil
}
// Delete 从 MemoryStore 删除指定 key
func (s *MemoryStore) Delete(_ context.Context, key string) error {
s.mu.Lock()
defer s.mu.Unlock()
delete(s.items, key)
return nil
}
// SetNX atomically sets key only when it is absent (or expired).
// Returns true if the key was written by this call.
func (s *MemoryStore) SetNX(_ context.Context, key string, val string, ttl time.Duration) (bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
_, exists, _ := s.getLocked(key)
if exists {
return false, nil
}
s.items[key] = memoryItem{
value: val,
expiresAt: time.Now().Add(ttl),
}
return true, nil
}
// GetAndDelete atomically retrieves and removes key in one critical section.
func (s *MemoryStore) GetAndDelete(_ context.Context, key string) (string, bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
val, exists, err := s.getLocked(key)
if err != nil || !exists {
return "", false, err
}
delete(s.items, key)
return val, true, nil
}
func (s *MemoryStore) startCleanupLoop(interval time.Duration) {
ticker := time.NewTicker(interval)
for range ticker.C {
s.cleanupExpired()
}
}
func (s *MemoryStore) cleanupExpired() {
now := time.Now()
s.mu.Lock()
defer s.mu.Unlock()
for k, v := range s.items {
if now.After(v.expiresAt) {
delete(s.items, k)
}
}
}
// RedisStore is a GORM-compatible/standalone Redis-backed implementation of Store
type RedisStore struct {
client redis.UniversalClient
}
// NewRedisStore creates a new RedisStore wrapping a redis.UniversalClient
func NewRedisStore(client redis.UniversalClient) *RedisStore {
return &RedisStore{
client: client,
}
}
// Get 从 RedisStore 获取指定 key 的值
func (s *RedisStore) Get(ctx context.Context, key string) (string, bool, error) {
val, err := s.client.Get(ctx, key).Result()
if err == redis.Nil {
return "", false, nil
}
if err != nil {
return "", false, err
}
return val, true, nil
}
// Set 向 RedisStore 写入指定 key 的值
func (s *RedisStore) Set(ctx context.Context, key string, val string, ttl time.Duration) error {
return s.client.Set(ctx, key, val, ttl).Err()
}
// Delete 从 RedisStore 删除指定 key
func (s *RedisStore) Delete(ctx context.Context, key string) error {
return s.client.Del(ctx, key).Err()
}
// SetNX wraps Redis SET NX – returns true only when the key was newly created.
func (s *RedisStore) SetNX(ctx context.Context, key string, val string, ttl time.Duration) (bool, error) {
return s.client.SetNX(ctx, key, val, ttl).Result()
}
// GetAndDelete wraps Redis GETDEL (available since Redis 6.2).
func (s *RedisStore) GetAndDelete(ctx context.Context, key string) (string, bool, error) {
val, err := s.client.GetDel(ctx, key).Result()
if err == redis.Nil {
return "", false, nil
}
if err != nil {
return "", false, err
}
return val, true, nil
}
@@ -0,0 +1,236 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"context"
"encoding/json"
"errors"
"strconv"
"sync"
"sync/atomic"
"time"
"golang.org/x/sync/singleflight"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
cachepkg "github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
database "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
)
const (
defaultChallengeCount = 1
defaultChallengeSize = 32
defaultChallengeDifficulty = 4
defaultChallengeTTL = 10 * time.Minute
defaultTokenTTL = 20 * time.Minute
)
// RuntimeSettings is the parsed CAPTCHA runtime configuration loaded from system_configs.
type RuntimeSettings struct {
LoginEnabled bool
ChallengeCount int
ChallengeSize int
ChallengeDifficulty int
ChallengeTTL time.Duration
TokenTTL time.Duration
}
// CAP 动态配置键常量
const (
ConfigKeyCapLoginEnabled = "cap_login_enabled"
ConfigKeyCapChallengeCount = "cap_challenge_count"
ConfigKeyCapChallengeSize = "cap_challenge_size"
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty"
ConfigKeyCapChallengeTTL = "cap_challenge_ttl"
// ConfigKeyCapTokenTTL 验证码 Token 过期时间键
// #nosec G101
ConfigKeyCapTokenTTL = "cap_token_ttl"
)
var runtimeConfigKeys = []string{
ConfigKeyCapLoginEnabled,
ConfigKeyCapChallengeCount,
ConfigKeyCapChallengeSize,
ConfigKeyCapChallengeDifficulty,
ConfigKeyCapChallengeTTL,
ConfigKeyCapTokenTTL,
}
var runtimeConfigKeySet = func() map[string]struct{} {
set := make(map[string]struct{}, len(runtimeConfigKeys))
for _, key := range runtimeConfigKeys {
set[key] = struct{}{}
}
return set
}()
type runtimeSettingsStore struct {
snapshot atomic.Pointer[RuntimeSettings]
loadGroup singleflight.Group
listenerOnce sync.Once
}
var settingsStore = &runtimeSettingsStore{}
// IsRuntimeConfigKey reports whether a system config key affects CAPTCHA runtime settings.
func IsRuntimeConfigKey(key string) bool {
_, ok := runtimeConfigKeySet[key]
return ok
}
// CurrentSettings returns the cached CAPTCHA runtime settings snapshot.
func CurrentSettings(ctx context.Context) (RuntimeSettings, error) {
return settingsStore.current(ctx)
}
// ProtectionEnabled reports whether CAPTCHA verification is required for protected routes.
func ProtectionEnabled(ctx context.Context) bool {
settings, err := CurrentSettings(ctx)
if err != nil {
return false
}
return settings.LoginEnabled
}
// InvalidateRuntimeSettings drops the in-process CAPTCHA settings snapshot.
func InvalidateRuntimeSettings() {
settingsStore.snapshot.Store(nil)
}
// ResetRuntimeSettingsForTest clears the CAPTCHA runtime snapshot.
func ResetRuntimeSettingsForTest() {
InvalidateRuntimeSettings()
}
// InstallTestRuntimeSettings installs a fixed snapshot for unit tests.
func InstallTestRuntimeSettings(settings RuntimeSettings) func() {
snapshot := settings
settingsStore.snapshot.Store(&snapshot)
return InvalidateRuntimeSettings
}
func (s *runtimeSettingsStore) current(ctx context.Context) (RuntimeSettings, error) {
s.ensureInvalidationListener()
if snapshot := s.snapshot.Load(); snapshot != nil {
return *snapshot, nil
}
loaded, err, _ := s.loadGroup.Do("cap-runtime-settings", func() (any, error) {
if snapshot := s.snapshot.Load(); snapshot != nil {
return *snapshot, nil
}
settings, loadErr := loadRuntimeSettings(ctx)
if loadErr != nil {
return RuntimeSettings{}, loadErr
}
s.snapshot.Store(&settings)
return settings, nil
})
if err != nil {
return RuntimeSettings{}, err
}
settings, ok := loaded.(RuntimeSettings)
if !ok {
return RuntimeSettings{}, errors.New("cap runtime settings loader returned unexpected type")
}
return settings, nil
}
func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
type configRecord struct {
Key string `gorm:"column:key"`
Value string `gorm:"column:value"`
}
var records []configRecord
if err := database.DB(ctx).Table("w_system_configs").Where("key IN ?", runtimeConfigKeys).Find(&records).Error; err != nil {
return RuntimeSettings{}, err
}
configs := make(map[string]string, len(records))
for _, r := range records {
configs[r.Key] = r.Value
}
return parseRuntimeSettings(configs), nil
}
func parseRuntimeSettings(configs map[string]string) RuntimeSettings {
settings := RuntimeSettings{
ChallengeCount: defaultChallengeCount,
ChallengeSize: defaultChallengeSize,
ChallengeDifficulty: defaultChallengeDifficulty,
ChallengeTTL: defaultChallengeTTL,
TokenTTL: defaultTokenTTL,
}
if val, ok := configs[ConfigKeyCapLoginEnabled]; ok {
if enabled, err := strconv.ParseBool(val); err == nil {
settings.LoginEnabled = enabled
}
}
if val, ok := configs[ConfigKeyCapChallengeCount]; ok {
if count, err := strconv.Atoi(val); err == nil && count > 0 {
settings.ChallengeCount = count
}
}
if val, ok := configs[ConfigKeyCapChallengeSize]; ok {
if size, err := strconv.Atoi(val); err == nil && size > 0 {
settings.ChallengeSize = size
}
}
if val, ok := configs[ConfigKeyCapChallengeDifficulty]; ok {
if diff, err := strconv.Atoi(val); err == nil && diff > 0 {
settings.ChallengeDifficulty = diff
}
}
if val, ok := configs[ConfigKeyCapChallengeTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.ChallengeTTL = time.Duration(ttlSeconds) * time.Second
}
}
if val, ok := configs[ConfigKeyCapTokenTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.TokenTTL = time.Duration(ttlSeconds) * time.Second
}
}
return settings
}
func (s *runtimeSettingsStore) ensureInvalidationListener() {
s.listenerOnce.Do(startRuntimeSettingsInvalidationListener)
}
// SystemConfigInvalidationChannel 系统配置失效广播通道
const SystemConfigInvalidationChannel = "system_config:invalidation"
func startRuntimeSettingsInvalidationListener() {
rdb := cachepkg.Redis
if rdb == nil {
return
}
util.Go(func() {
pubsub := rdb.Subscribe(context.Background(), SystemConfigInvalidationChannel)
defer func() {
_ = pubsub.Close()
}()
for msg := range pubsub.Channel() {
var payload struct {
Key string `json:"key"`
}
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
InvalidateRuntimeSettings()
continue
}
if payload.Key == "" || payload.Key == "*" || IsRuntimeConfigKey(payload.Key) {
InvalidateRuntimeSettings()
}
}
})
}
+444
View File
@@ -0,0 +1,444 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package domain_test
import (
"context"
"io/fs"
"path/filepath"
"testing"
"github.com/alicebob/miniredis/v2"
"github.com/glebarez/sqlite"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/backend/core"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/admin"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/auth"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/risk_control"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/user"
"github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
"github.com/Rain-kl/Wavelet/backend/plugins/infra/logger"
"github.com/Rain-kl/Wavelet/backend/plugins/infra/storage"
)
func setupTestDB(t *testing.T) *gorm.DB {
t.Helper()
dbPath := filepath.Join(t.TempDir(), "domain_test.db")
testDB, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, testDB.AutoMigrate(
&user.User{},
&user.AccessToken{},
&auth.AuthSource{},
&auth.ExternalAccount{},
&message_gateway.MessageChannel{},
&message_gateway.MessageBinding{},
&message_gateway.MessagePairingCode{},
&admin.SystemConfig{},
&message_gateway.PushChannel{},
&message_gateway.PushEvent{},
&message_gateway.PushHistory{},
))
db.SetDB(testDB)
return testDB
}
type mockOAuthProvider struct {
name string
}
func (m *mockOAuthProvider) Name() string {
return m.name
}
func (m *mockOAuthProvider) GetAuthURL(state string) string {
return "https://auth.example.com/auth?state=" + state
}
func (m *mockOAuthProvider) ExchangeCode(ctx context.Context, code string) (*contracts.OAuthUserInfoDTO, error) {
return &contracts.OAuthUserInfoDTO{
ID: 1001,
Username: "mock_user",
Email: "mock@example.com",
Active: true,
}, nil
}
func TestAuthPlugin(t *testing.T) {
ctx := core.NewContext(context.Background())
testDB := setupTestDB(t)
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
require.NoError(t, cache.New().Apply(ctx))
require.NoError(t, logger.New().Apply(ctx))
p := auth.New()
assert.Equal(t, "auth", p.Name())
assert.Equal(t, "auth", p.Manifest().Name)
require.NoError(t, p.Apply(ctx))
// 1. Verify migrations registered
entry, ok := ctx.Migrations().Get("auth")
require.True(t, ok)
assert.Equal(t, "auth", entry.PluginID)
entries, err := fs.ReadDir(entry.FS, entry.Dir)
require.NoError(t, err)
assert.NotEmpty(t, entries)
// 2. Verify AuthService
authSvc, err := core.Inject[contracts.AuthService](ctx)
require.NoError(t, err)
require.NotNil(t, authSvc)
assert.NotNil(t, authSvc.RequireAuthMiddleware())
assert.NotNil(t, authSvc.RequireAdminMiddleware())
// 3. Verify AuthRegistry
authReg, err := core.Inject[contracts.AuthRegistry](ctx)
require.NoError(t, err)
require.NotNil(t, authReg)
mockProv := &mockOAuthProvider{name: "github"}
authReg.RegisterOAuthProvider("github", mockProv)
retrieved, ok := authReg.GetOAuthProvider("github")
require.True(t, ok)
assert.Equal(t, "github", retrieved.Name())
assert.Contains(t, authReg.ListOAuthProviders(), "github")
// 4. Verify Routes
routes := ctx.Router().Routes()
var hasSources, hasLogin, hasUserInfo bool
for _, r := range routes {
if r.Path == "/api/v1/oauth/sources" {
hasSources = true
}
if r.Path == "/api/v1/oauth/login" {
hasLogin = true
}
if r.Path == "/api/v1/user-info" {
hasUserInfo = true
}
}
assert.True(t, hasSources)
assert.True(t, hasLogin)
assert.True(t, hasUserInfo)
// 5. Verify Settings
schema, ok := ctx.Settings().Get("auth.session_age")
require.True(t, ok)
assert.Equal(t, 86400*7, schema.Default)
}
func TestUserPlugin(t *testing.T) {
ctx := core.NewContext(context.Background())
testDB := setupTestDB(t)
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
require.NoError(t, cache.New().Apply(ctx))
require.NoError(t, logger.New().Apply(ctx))
p := user.New()
assert.Equal(t, "user", p.Name())
assert.Equal(t, "user", p.Manifest().Name)
require.NoError(t, p.Apply(ctx))
// 1. Verify migrations
entry, ok := ctx.Migrations().Get("user")
require.True(t, ok)
assert.Equal(t, "user", entry.PluginID)
// 2. Verify UserService
userSvc, err := core.Inject[contracts.UserService](ctx)
require.NoError(t, err)
require.NotNil(t, userSvc)
testCtx := context.Background()
// 3. Create user
created, err := userSvc.CreateUser(testCtx, contracts.CreateUserRequest{
Username: "bob",
Password: "SecurePassword123!",
Nickname: "Bob Builder",
Email: "bob@example.com",
IsAdmin: false,
})
require.NoError(t, err)
require.NotNil(t, created)
assert.Equal(t, "bob", created.Username)
assert.Equal(t, "Bob Builder", created.Nickname)
assert.Equal(t, "bob@example.com", created.Email)
assert.False(t, created.IsAdmin)
// 4. Query user
byID, err := userSvc.GetUserByID(testCtx, created.ID)
require.NoError(t, err)
assert.Equal(t, "bob", byID.Username)
byUsername, err := userSvc.GetUserByUsername(testCtx, "bob")
require.NoError(t, err)
assert.Equal(t, created.ID, byUsername.ID)
byEmail, err := userSvc.GetUserByEmail(testCtx, "bob@example.com")
require.NoError(t, err)
assert.Equal(t, created.ID, byEmail.ID)
// 5. Password verification and update
assert.True(t, userSvc.VerifyPassword(testCtx, created.ID, "SecurePassword123!"))
assert.False(t, userSvc.VerifyPassword(testCtx, created.ID, "WrongPass"))
require.NoError(t, userSvc.UpdatePassword(testCtx, created.ID, "SecurePassword123!", "NewSecurePassword456!"))
assert.True(t, userSvc.VerifyPassword(testCtx, created.ID, "NewSecurePassword456!"))
// 6. Update Profile
newBio := "I build things"
newPhone := "13800138000"
updated, err := userSvc.UpdateProfile(testCtx, created.ID, contracts.UpdateUserProfileRequest{
Bio: &newBio,
Phone: &newPhone,
})
require.NoError(t, err)
assert.Equal(t, newBio, updated.Bio)
assert.Equal(t, newPhone, updated.Phone)
// 7. Update Last Login
require.NoError(t, userSvc.UpdateLastLogin(testCtx, created.ID, "127.0.0.1"))
// 8. Admin operations: SetUserActive, SetUserAdmin, ListUsers
require.NoError(t, userSvc.SetUserAdmin(testCtx, created.ID, true))
reloaded, err := userSvc.GetUserByID(testCtx, created.ID)
require.NoError(t, err)
assert.True(t, reloaded.IsAdmin)
require.NoError(t, userSvc.SetUserActive(testCtx, created.ID, false))
reloadedBanned, err := userSvc.GetUserByID(testCtx, created.ID)
require.NoError(t, err)
assert.False(t, reloadedBanned.IsActive)
list, total, err := userSvc.ListUsers(testCtx, 1, 10, "bob")
require.NoError(t, err)
assert.Equal(t, int64(1), total)
assert.Len(t, list, 1)
assert.Equal(t, "bob", list[0].Username)
// 9. Tasks & Schedules
taskDef, ok := ctx.Tasks().Get("user:send_email_code")
require.True(t, ok)
assert.Equal(t, 3, taskDef.Retry)
schedDef, ok := ctx.Schedules().Get("user:daily_audit")
require.True(t, ok)
assert.Equal(t, "0 3 * * *", schedDef.Spec)
// 10. Settings
sReg, ok := ctx.Settings().Get("user.registration_enabled")
require.True(t, ok)
assert.Equal(t, true, sReg.Default)
}
func TestMessageGatewayPlugin(t *testing.T) {
ctx := core.NewContext(context.Background())
testDB := setupTestDB(t)
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
require.NoError(t, cache.New().Apply(ctx))
require.NoError(t, logger.New().Apply(ctx))
p := message_gateway.New()
assert.Equal(t, "message_gateway", p.Name())
assert.Equal(t, "message_gateway", p.Manifest().Name)
require.NoError(t, p.Apply(ctx))
// 1. Migrations
entry, ok := ctx.Migrations().Get("message_gateway")
require.True(t, ok)
assert.Equal(t, "message_gateway", entry.PluginID)
// 2. Routes
routes := ctx.Router().Routes()
var hasChannels, hasBindings bool
for _, r := range routes {
if r.Path == "/api/v1/message-gateway/channels" {
hasChannels = true
}
if r.Path == "/api/v1/message-gateway/bindings" {
hasBindings = true
}
}
assert.True(t, hasChannels)
assert.True(t, hasBindings)
// 3. Tasks & Schedules
taskDef, ok := ctx.Tasks().Get("message_gateway:push_notification")
require.True(t, ok)
assert.Equal(t, 3, taskDef.Retry)
schedDef, ok := ctx.Schedules().Get("message_gateway:cleanup_pairing_codes")
require.True(t, ok)
assert.Equal(t, "*/10 * * * *", schedDef.Spec)
// 4. EventBus Trigger
var receivedEvent message_gateway.PushNotificationEvent
var eventFired bool
ctx.Events().On("notification:push", func(c context.Context, e message_gateway.PushNotificationEvent) error {
eventFired = true
receivedEvent = e
return nil
})
err := ctx.Events().Emit(context.Background(), "notification:push", message_gateway.PushNotificationEvent{
UserID: 99,
Channel: "telegram",
Title: "System Alert",
Content: "Disk 85% full",
})
require.NoError(t, err)
assert.True(t, eventFired)
assert.Equal(t, uint64(99), receivedEvent.UserID)
assert.Equal(t, "telegram", receivedEvent.Channel)
assert.Equal(t, "System Alert", receivedEvent.Title)
// 5. Settings
schema, ok := ctx.Settings().Get("message_gateway.pairing_code_expiry_minutes")
require.True(t, ok)
assert.Equal(t, 15, schema.Default)
}
func TestRiskControlPlugin(t *testing.T) {
ctx := core.NewContext(context.Background())
p := risk_control.New()
assert.Equal(t, "risk_control", p.Name())
assert.Equal(t, "risk_control", p.Manifest().Name)
require.NoError(t, p.Apply(ctx))
// 1. Middleware registered on Router
mws := ctx.Router().Middlewares()
assert.NotEmpty(t, mws)
// 2. Settings
schema, ok := ctx.Settings().Get("risk_control.ip_rate_limit_per_minute")
require.True(t, ok)
assert.Equal(t, 60, schema.Default)
// 3. Disposal cleanup
require.NoError(t, ctx.Dispose())
}
func TestAdminPlugin(t *testing.T) {
ctx := core.NewContext(context.Background())
testDB := setupTestDB(t)
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
require.NoError(t, cache.New().Apply(ctx))
require.NoError(t, logger.New().Apply(ctx))
p := admin.New()
assert.Equal(t, "admin", p.Name())
assert.Equal(t, "admin", p.Manifest().Name)
require.NoError(t, p.Apply(ctx))
// 1. Admin Routes
routes := ctx.Router().Routes()
var hasStatus, hasDBOverview, hasUsers, hasTasks, hasConfigs bool
for _, r := range routes {
if r.Path == "/api/v1/admin/status" {
hasStatus = true
}
if r.Path == "/api/v1/admin/db-manage/overview" {
hasDBOverview = true
}
if r.Path == "/api/v1/admin/users" {
hasUsers = true
}
if r.Path == "/api/v1/admin/tasks/types" {
hasTasks = true
}
if r.Path == "/api/v1/admin/system-configs" {
hasConfigs = true
}
}
assert.True(t, hasStatus)
assert.True(t, hasDBOverview)
assert.True(t, hasUsers)
assert.True(t, hasTasks)
assert.True(t, hasConfigs)
// 2. Task & Schedule
_, ok := ctx.Tasks().Get("admin:system_cleanup")
require.True(t, ok)
sched, ok := ctx.Schedules().Get("admin:system_cleanup")
require.True(t, ok)
assert.Equal(t, "0 4 * * *", sched.Spec)
// 3. Settings
schema, ok := ctx.Settings().Get("admin.system_cleanup_cron")
require.True(t, ok)
assert.Equal(t, "0 4 * * *", schema.Default)
}
func TestAllDomainPluginsCombined(t *testing.T) {
mr, err := miniredis.Run()
require.NoError(t, err)
defer mr.Close()
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
defer func() { _ = rdb.Close() }()
ctx := core.NewContext(context.Background())
testDB := setupTestDB(t)
// Apply Infra plugins
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
require.NoError(t, cache.New(cache.WithRedis(rdb)).Apply(ctx))
require.NoError(t, logger.New().Apply(ctx))
require.NoError(t, storage.New().Apply(ctx))
// Apply Domain plugins
require.NoError(t, auth.New().Apply(ctx))
require.NoError(t, user.New().Apply(ctx))
require.NoError(t, message_gateway.New().Apply(ctx))
require.NoError(t, risk_control.New().Apply(ctx))
require.NoError(t, admin.New().Apply(ctx))
// Verify cross-plugin service injection via Using3
var resolved bool
err = core.Using3(ctx, func(authSvc contracts.AuthService, userSvc contracts.UserService, authReg contracts.AuthRegistry) {
resolved = true
assert.NotNil(t, authSvc)
assert.NotNil(t, userSvc)
assert.NotNil(t, authReg)
})
require.NoError(t, err)
assert.True(t, resolved)
// Verify all migration entries
allMigrations := ctx.Migrations().Entries()
assert.GreaterOrEqual(t, len(allMigrations), 3)
// Verify total routes registered
allRoutes := ctx.Router().Routes()
assert.GreaterOrEqual(t, len(allRoutes), 20)
// Verify total tasks registered
allTasks := ctx.Tasks().Tasks()
assert.GreaterOrEqual(t, len(allTasks), 4)
// Verify total schedules registered
allSchedules := ctx.Schedules().Schedules()
assert.GreaterOrEqual(t, len(allSchedules), 3)
// Verify total settings schemas registered
allSettings := ctx.Settings().Schemas()
assert.GreaterOrEqual(t, len(allSettings), 7)
// Clean shutdown
require.NoError(t, ctx.Dispose())
}
@@ -0,0 +1,170 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"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=[]Definition}
// @Router /api/v1/admin/message-gateway/channels/definitions [get]
func ListAdminChannelDefinitions(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(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=[]ChannelDTO}
// @Router /api/v1/admin/message-gateway/channels [get]
func ListAdminChannels(c *gin.Context) {
rows, err := listChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(rows))
}
// 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 CreateChannelRequest true "create body"
// @Success 200 {object} response.Any{data=ChannelDTO}
// @Failure 400 {object} response.Any
// @Router /api/v1/admin/message-gateway/channels [post]
func CreateAdminChannel(c *gin.Context) {
var req CreateChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := createChannel(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(dto))
}
// 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 UpdateChannelRequest true "update body"
// @Success 200 {object} response.Any{data=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) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return
}
var req UpdateChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := updateChannel(c.Request.Context(), id, req)
if err != nil {
if err.Error() == errChannelNotFound {
response.AbortNotFound(c, err.Error())
return
}
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(dto))
}
// 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, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return
}
if err := deleteChannel(c.Request.Context(), id); err != nil {
if err.Error() == errChannelNotFound {
response.AbortNotFound(c, err.Error())
return
}
response.AbortInternal(c, err.Error())
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, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return
}
if err := probeChannel(c.Request.Context(), id); err != nil {
if err.Error() == errChannelNotFound {
response.AbortNotFound(c, err.Error())
return
}
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// RegisterAdminRoutes mounts admin message-gateway APIs under /admin.
func RegisterAdminRoutes(adminRouter *gin.RouterGroup) {
g := adminRouter.Group("/message-gateway")
{
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)
}
}
@@ -0,0 +1,338 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strings"
"time"
"github.com/tencent-connect/botgo/token"
"gorm.io/gorm"
)
const defaultTelegramAPI = "https://api.telegram.org"
// 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"`
}
func listDefinitions() []Definition {
return []Definition{
{
Type: MessageChannelTypeTelegram,
Fields: []Field{
{Key: "token", Type: "password", Required: true},
{Key: "api_base", Type: "text", Required: false},
},
},
{
Type: MessageChannelTypeQQ,
Fields: []Field{
{Key: "app_id", Type: "text", Required: true},
{Key: "client_secret", Type: "password", Required: true},
},
},
}
}
func createChannel(ctx context.Context, req CreateChannelRequest) (ChannelDTO, error) {
name := strings.TrimSpace(req.Name)
if name == "" {
return ChannelDTO{}, errors.New(errNameRequired)
}
channelType := strings.TrimSpace(req.Type)
if channelType != MessageChannelTypeTelegram && channelType != MessageChannelTypeQQ {
return ChannelDTO{}, errors.New(errTypeInvalid)
}
creds := req.Credentials
if creds == nil {
creds = map[string]string{}
}
if err := validateCredentials(channelType, creds, false); err != nil {
return ChannelDTO{}, err
}
cipher, err := EncryptCredentials(creds)
if err != nil {
return ChannelDTO{}, err
}
extra := req.Extra
if extra == nil {
extra = map[string]string{}
}
enabled := true
if req.Enabled != nil {
enabled = *req.Enabled
}
row := &MessageChannel{
Name: name,
Type: channelType,
OwnerScope: MessageOwnerScopeSystem,
Enabled: enabled,
Credentials: cipher,
Extra: EncodeExtra(extra),
}
if err := CreateMessageChannel(ctx, row); err != nil {
return ChannelDTO{}, err
}
return toDTO(row, creds, extra), nil
}
func updateChannel(ctx context.Context, id uint64, req UpdateChannelRequest) (ChannelDTO, error) {
row, err := GetMessageChannel(ctx, id)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return ChannelDTO{}, errors.New(errChannelNotFound)
}
return ChannelDTO{}, err
}
creds, err := DecryptCredentials(row.Credentials)
if err != nil {
return 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 ChannelDTO{}, err
}
creds = merged
}
cipher, err := EncryptCredentials(creds)
if err != nil {
return ChannelDTO{}, err
}
row.Credentials = cipher
row.Extra = EncodeExtra(extra)
if err := UpdateMessageChannel(ctx, row); err != nil {
return ChannelDTO{}, err
}
return toDTO(row, creds, extra), nil
}
func listChannels(ctx context.Context) ([]ChannelDTO, error) {
rows, err := ListMessageChannels(ctx)
if err != nil {
return nil, err
}
out := make([]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
}
func deleteChannel(ctx context.Context, id uint64) error {
if _, err := GetMessageChannel(ctx, id); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errChannelNotFound)
}
return err
}
return DeleteMessageChannel(ctx, id)
}
func probeChannel(ctx context.Context, id uint64) error {
row, err := GetMessageChannel(ctx, id)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errChannelNotFound)
}
return err
}
creds, err := DecryptCredentials(row.Credentials)
if err != nil {
return err
}
switch row.Type {
case MessageChannelTypeTelegram:
return probeTelegram(ctx, creds)
case MessageChannelTypeQQ:
return probeQQ(ctx, creds)
default:
return errors.New(errTypeInvalid)
}
}
func probeTelegram(ctx context.Context, creds map[string]string) error {
tok := creds["token"]
if strings.TrimSpace(tok) == "" {
return errors.New("missing telegram bot token")
}
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("telegram getMe failed (%d): %s", 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("telegram returned ok=false: %s", string(body))
}
return nil
}
func probeQQ(_ context.Context, creds map[string]string) error {
appID := strings.TrimSpace(creds["app_id"])
secret := strings.TrimSpace(creds["app_secret"])
if appID == "" || secret == "" {
return errors.New("missing qq app_id or app_secret")
}
credentials := &token.QQBotCredentials{
AppID: appID,
AppSecret: secret,
}
tokSrc := token.NewQQBotTokenSource(credentials)
tok, err := tokSrc.Token()
if err != nil {
return fmt.Errorf("qq token fetch failed: %w", err)
}
if tok == nil || tok.AccessToken == "" {
return errors.New("qq returned empty access token")
}
return nil
}
func validateCredentials(t string, creds map[string]string, isUpdate bool) error {
switch t {
case MessageChannelTypeTelegram:
tok := creds["token"]
if strings.TrimSpace(tok) == "" && !isUpdate {
return errors.New(errTelegramTokenRequired)
}
if base, ok := creds["api_base"]; ok && strings.TrimSpace(base) != "" {
if !strings.HasPrefix(base, "http://") && !strings.HasPrefix(base, "https://") {
return errors.New("api_base must start with http:// or https://")
}
}
case MessageChannelTypeQQ:
appID := creds["app_id"]
secret := creds["client_secret"]
if (strings.TrimSpace(appID) == "" || strings.TrimSpace(secret) == "") && !isUpdate {
return errors.New(errQQCredentialsRequired)
}
default:
return errors.New(errTypeInvalid)
}
return nil
}
func toDTO(row *MessageChannel, creds, extra map[string]string) ChannelDTO {
return 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,
}
}
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
func maskSecret(s string) string {
s = strings.TrimSpace(s)
if len(s) <= minMaskSecretLength {
return "******"
}
return s[:4] + "..." + s[len(s)-4:]
}
@@ -0,0 +1,21 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import "context"
// Handler processes one inbound message.
type Handler func(ctx context.Context, msg InboundMessage) error
// Factory constructs a Channel from decrypted config.
type Factory func(cfg 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 Recipient, msg OutboundMessage) error
Capabilities() Capability
}
@@ -0,0 +1,161 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package qq implements the official QQ Bot C2C adapter.
package qq
import (
"context"
"fmt"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway"
"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 message_gateway.ChannelConfig
onInbound message_gateway.Handler
api openapi.OpenAPI
tokenSrc oauth2.TokenSource
cancel context.CancelFunc
mu sync.Mutex
disconnected bool
}
// New constructs a QQ adapter.
func New(cfg message_gateway.ChannelConfig, onInbound message_gateway.Handler) (message_gateway.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 message_gateway.ChannelTypeQQ }
// Capabilities reports C2C text/media support.
func (a *Adapter) Capabilities() message_gateway.Capability {
return message_gateway.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()
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 message_gateway.Recipient, msg message_gateway.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, message_gateway.InboundMessage{
ChannelID: a.cfg.ID,
ChannelType: message_gateway.ChannelTypeQQ,
PlatformUserID: ev.UserID,
ChatID: ev.UserID,
MessageID: ev.MessageID,
Text: ev.Text,
})
}
@@ -0,0 +1,42 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package qq
import (
"context"
"testing"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway"
)
func TestHandleEvent_DropsNonC2C(t *testing.T) {
var got int
a := &Adapter{onInbound: func(ctx context.Context, msg message_gateway.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 message_gateway.InboundMessage
a := &Adapter{cfg: message_gateway.ChannelConfig{ID: 3}, onInbound: func(ctx context.Context, msg message_gateway.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(message_gateway.ChannelConfig{}, nil)
if err == nil {
t.Fatal("expected error")
}
}
@@ -0,0 +1,153 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package telegram implements the Telegram private-chat adapter.
package telegram
import (
"context"
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway"
tele "gopkg.in/telebot.v4"
)
// Adapter is a Telegram private-chat channel.
type Adapter struct {
cfg message_gateway.ChannelConfig
onInbound message_gateway.Handler
bot *tele.Bot
}
// New constructs a Telegram adapter. Call message_gateway.Register from the runner.
func New(cfg message_gateway.ChannelConfig, onInbound message_gateway.Handler) (message_gateway.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 message_gateway.ChannelTypeTelegram }
// Capabilities reports private-chat media support.
func (a *Adapter) Capabilities() message_gateway.Capability {
return message_gateway.Capability{Text: true, Image: true, File: true, Reply: true}
}
// Connect starts long polling.
func (a *Adapter) Connect(ctx context.Context) error {
pref := tele.Settings{
Token: a.cfg.Credentials["bot_token"],
Poller: &tele.LongPoller{Timeout: 10},
}
if base := strings.TrimSpace(a.cfg.Extra["base_url"]); base != "" {
pref.URL = strings.TrimSuffix(base, "/")
}
bot, err := tele.NewBot(pref)
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
})
go bot.Start()
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 message_gateway.Recipient, msg message_gateway.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 := message_gateway.InboundMessage{
ChannelID: a.cfg.ID,
ChannelType: message_gateway.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 {
msg.Attachments = a.downloadMedia(m)
}
_ = a.onInbound(ctx, msg)
}
func (a *Adapter) downloadMedia(m *tele.Message) []message_gateway.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 []message_gateway.Attachment{{Error: err.Error()}}
}
out := make([]message_gateway.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, message_gateway.Attachment{FileName: names[i], Error: err.Error()})
continue
}
out = append(out, message_gateway.Attachment{Path: path, FileName: names[i]})
}
return out
}
@@ -0,0 +1,56 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package telegram
import (
"context"
"testing"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway"
tele "gopkg.in/telebot.v4"
)
func TestHandleUpdate_DropsGroups(t *testing.T) {
var got int
a := &Adapter{onInbound: func(ctx context.Context, msg message_gateway.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 message_gateway.InboundMessage
a := &Adapter{
cfg: message_gateway.ChannelConfig{ID: 7, Type: "telegram"},
onInbound: func(ctx context.Context, msg message_gateway.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(message_gateway.ChannelConfig{}, nil)
if err == nil {
t.Fatal("expected error")
}
}
@@ -0,0 +1,42 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"time"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
)
// AdminLogin is the metadata definition for the admin login event.
var AdminLogin = EventMetadata{
Key: "admin_login",
Name: "管理员登录",
DefaultTemplate: NotificationMessage{
Title: "管理员登录提醒",
Content: "管理员 {{user.username}} 于 {{time}} 从 IP {{ip}} 登录系统。",
Level: "INFO",
},
Description: "当管理员成功登录系统时触发此通知",
}
// HandleAdminLoggedIn 处理管理员登录事件并触发通知
func HandleAdminLoggedIn(ctx context.Context, event contracts.AdminLoggedIn) {
if event.User == nil {
return
}
body := map[string]any{
"user": event.User,
"ip": event.IP,
"time": time.Now().Format("2006-01-02 15:04:05"),
}
DefaultTrigger.Trigger(ctx, AdminLogin, body)
}
// RegisterCustomEvents registers default domain push notification events.
func RegisterCustomEvents() {
RegisterBuiltInEvent(AdminLogin)
}
@@ -0,0 +1,26 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import "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")
)
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 = "********"
)
@@ -0,0 +1,62 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package message_gateway defines channel adapters, pairing codes, and inbound types.
package message_gateway
// ChannelTypeTelegram is the Telegram private-chat adapter type.
const ChannelTypeTelegram = "telegram"
// ChannelTypeQQ is the official QQ Bot C2C adapter type.
const ChannelTypeQQ = "qq"
// 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
}
@@ -0,0 +1,147 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"errors"
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
"github.com/gin-gonic/gin"
)
func currentUser(c *gin.Context) (*contracts.UserDTO, bool) {
return util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
}
// 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=[]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, "login required")
return
}
rows, err := 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=[]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, "login required")
return
}
rows, err := 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 BindRequest true "bind body"
// @Success 200 {object} response.Any{data=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, "login required")
return
}
var req BindRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := bindChannel(c.Request.Context(), user.ID, req)
if err != nil {
if errors.Is(err, 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, "login required")
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid binding id")
return
}
if err := unbindChannel(c.Request.Context(), user.ID, id); err != nil {
if errors.Is(err, errBindingNotFound) {
response.AbortNotFound(c, err.Error())
return
}
if errors.Is(err, errBindingForbidden) {
response.AbortForbidden(c, err.Error())
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// RegisterUserRoutes mounts user-facing message gateway endpoints.
func RegisterUserRoutes(r *gin.RouterGroup, 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)
}
}
@@ -0,0 +1,155 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"errors"
"strconv"
"strings"
"time"
"gorm.io/gorm"
)
// 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"`
}
func bindChannel(ctx context.Context, userID uint64, req BindRequest) (BindingDTO, error) {
channelID, err := strconv.ParseUint(strings.TrimSpace(req.ChannelID), 10, 64)
if err != nil || channelID == 0 {
return BindingDTO{}, errChannelIDRequired
}
code := NormalizeCode(req.Code)
if code == "" {
return BindingDTO{}, errCodeInvalid
}
pairing, err := GetPairingCode(ctx, code)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return BindingDTO{}, errCodeInvalid
}
return BindingDTO{}, err
}
if !pairing.ExpiresAt.After(time.Now()) {
return BindingDTO{}, errCodeInvalid
}
if pairing.ChannelID != channelID {
return BindingDTO{}, errChannelMismatch
}
ch, err := GetMessageChannel(ctx, channelID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return BindingDTO{}, errCodeInvalid
}
return BindingDTO{}, err
}
if !ch.Enabled {
return BindingDTO{}, errChannelDisabled
}
existing, err := GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID)
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return BindingDTO{}, err
}
if err == nil && existing != nil {
if existing.UserID != userID {
return BindingDTO{}, errPlatformAlreadyBound
}
_ = DeletePairingCode(ctx, pairing.Code)
return toBindingDTO(existing, ch), nil
}
row := &MessageBinding{
UserID: userID,
ChannelID: channelID,
PlatformUserID: pairing.PlatformUserID,
}
if err := CreateMessageBinding(ctx, row); err != nil {
return BindingDTO{}, err
}
if err := DeletePairingCode(ctx, pairing.Code); err != nil {
return BindingDTO{}, err
}
return toBindingDTO(row, ch), nil
}
// 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"`
}
func listEnabledPublicChannels(ctx context.Context) ([]PublicChannelDTO, error) {
rows, err := ListEnabledMessageChannels(ctx)
if err != nil {
return nil, err
}
out := make([]PublicChannelDTO, 0, len(rows))
for _, row := range rows {
out = append(out, PublicChannelDTO{ID: row.ID, Name: row.Name, Type: row.Type})
}
return out, nil
}
func listUserBindings(ctx context.Context, userID uint64) ([]BindingDTO, error) {
rows, err := ListBindingsByUser(ctx, userID)
if err != nil {
return nil, err
}
out := make([]BindingDTO, 0, len(rows))
for i := range rows {
ch, err := 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
}
func unbindChannel(ctx context.Context, userID, bindingID uint64) error {
row, err := GetMessageBinding(ctx, bindingID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errBindingNotFound
}
return err
}
if row.UserID != userID {
return errBindingForbidden
}
return DeleteMessageBinding(ctx, bindingID)
}
func toBindingDTO(row *MessageBinding, ch *MessageChannel) BindingDTO {
dto := 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
}
@@ -0,0 +1,30 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway_test
import (
"context"
"testing"
"time"
"github.com/Rain-kl/Wavelet/backend/pkg/testhelper"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway"
)
func TestUpsertPairingCode_ReusesUnexpired(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer 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,165 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"errors"
"strings"
"time"
)
// Message channel and push channel constants.
const (
MessageChannelTypeTelegram = "telegram"
MessageChannelTypeQQ = "qq"
MessageOwnerScopeSystem = "system"
TypeCustom = "custom"
TypeEmail = "email"
TypeTelegram = "telegram"
)
// 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("channel name is required")
}
c.Type = strings.TrimSpace(c.Type)
if c.Type == "" {
return errors.New("channel type is required")
}
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("event_key is required")
}
e.Name = strings.TrimSpace(e.Name)
if e.Name == "" {
return errors.New("name is required")
}
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"
}
// 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
}
@@ -0,0 +1,50 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"crypto/rand"
"strings"
"unicode"
)
// 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:]
}
@@ -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,215 @@
// 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 (
"context"
"embed"
"github.com/Rain-kl/Wavelet/backend/core"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
"github.com/Rain-kl/Wavelet/backend/core/extpoints"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
)
//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"
}
// 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",
}
}
// 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"`
}
// Apply registers message_gateway migrations, routes, tasks, schedules, events, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
// 0. Resolve auth service for middleware (via IoC, not direct import)
var loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
var adminMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
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
mgGroup := ctx.Router().Group("/api/v1/message-gateway", loginMW)
{
mgGroup.GET("/channels", ListChannels)
mgGroup.GET("/bindings", ListBindings)
mgGroup.POST("/bindings", BindBinding)
mgGroup.DELETE("/bindings/:id", UnbindBinding)
}
// 3. Register Admin Message Gateway HTTP Routes
adminMgGroup := ctx.Router().Group("/api/v1/admin/message-gateway", loginMW, adminMW)
{
adminMgGroup.GET("/channels/definitions", ListAdminChannelDefinitions)
adminMgGroup.GET("/channels", ListAdminChannels)
adminMgGroup.POST("/channels", CreateAdminChannel)
adminMgGroup.PATCH("/channels/:id", UpdateAdminChannel)
adminMgGroup.DELETE("/channels/:id", DeleteAdminChannel)
adminMgGroup.POST("/channels/:id/test", TestAdminChannel)
}
// 4. Register Admin Push HTTP Routes
adminPushGroup := ctx.Router().Group("/api/v1/admin/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)
}
}
const defaultTaskRetry = 3
pushHandler := &PushHandler{}
// 5. Register Asynq background tasks
ctx.Task().Register("message_gateway:push_notification", func(c context.Context, t *asynq.Task) error {
_, err := pushHandler.Execute(c, t.Payload())
return err
}, extpoints.WithTaskRetry(defaultTaskRetry))
ctx.Task().Register(SendNotificationTask, func(c context.Context, t *asynq.Task) error {
_, err := pushHandler.Execute(c, t.Payload())
return err
}, extpoints.WithTaskRetry(defaultTaskRetry))
ctx.Task().Register("message_gateway:dispatch_bot_msg", func(_ context.Context, _ *asynq.Task) error {
return nil
})
// 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 PushNotificationEvent) error {
meta := EventMetadata{
Key: "eventbus:" + e.Channel,
Name: e.Title,
DefaultTemplate: NotificationMessage{
Title: e.Title,
Content: e.Content,
Level: defaultLevelInfo,
Ext: e.Metadata,
},
Description: "EventBus triggered notification",
}
DefaultTrigger.Trigger(c, meta, map[string]any{
"user.id": e.UserID,
"title": e.Title,
"content": e.Content,
})
return nil
})
// 8. Register built-in domain events and task listeners
RegisterCustomEvents()
RegisterTaskListeners()
// 9. 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",
})
// 10. Optional runner start & lifecycle
if p.autoStartRunner {
runnerCtx, cancel := context.WithCancel(ctx.GoContext())
p.cancelRunner = cancel
go func() {
_ = Start(runnerCtx)
}()
}
ctx.OnDispose(func() error {
if p.cancelRunner != nil {
p.cancelRunner()
}
return nil
})
return nil
}
@@ -0,0 +1,46 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway_test
import (
"context"
"io/fs"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/Rain-kl/Wavelet/backend/core"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway"
)
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)
}
@@ -0,0 +1,98 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strings"
"github.com/Rain-kl/Wavelet/backend/pkg/httppool"
)
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 { //nolint:mnd
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,119 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"errors"
"fmt"
"net"
"net/smtp"
"strings"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
)
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 := defaultTitle
if t, ok := body["title"].(string); ok && t != "" {
title = t
}
content := ""
if c, ok := body["content"].(string); ok && c != "" {
content = c
} else {
// 自动格式化 map
var parts []string
for k, v := range body {
parts = append(parts, fmt.Sprintf("<p><b>%s</b>: %v</p>", k, v))
}
content = strings.Join(parts, "")
}
// 邮件头和体
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,275 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"bytes"
"context"
"crypto/hmac"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/backend/pkg/httppool"
)
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 := defaultTitle
if t, ok := body["title"].(string); ok && t != "" {
title = t
}
content := ""
if c, ok := body["content"].(string); ok && c != "" {
content = c
} else {
// 兜底:如果连 content 都没有,把 body 里的所有值拼成 markdown
var parts []string
for k, v := range body {
parts = append(parts, fmt.Sprintf("**%s**: %v", k, v))
}
content = strings.Join(parts, "\n")
}
level := levelInfo
if l, ok := body["level"].(string); ok && l != "" {
level = strings.ToUpper(l)
}
// 根据级别确定飞书卡片头部的背景色模板
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,158 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"github.com/Rain-kl/Wavelet/backend/pkg/httppool"
)
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 消息发送
//
//nolint:cyclop
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 := defaultTitle
if t, ok := body["title"].(string); ok && t != "" {
title = t
}
content := ""
if c, ok := body["content"].(string); ok && c != "" {
content = c
} else {
var parts []string
for k, v := range body {
parts = append(parts, fmt.Sprintf("<b>%s</b>: %v", k, v))
}
content = strings.Join(parts, "\n")
}
level := levelInfo
if l, ok := body["level"].(string); ok && l != "" {
level = strings.ToUpper(l)
}
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: %v)", 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, "&", "&amp;")
s = strings.ReplaceAll(s, "<", "&lt;")
s = strings.ReplaceAll(s, ">", "&gt;")
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,78 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"encoding/json"
"fmt"
"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)
}
}
@@ -0,0 +1,75 @@
// 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)
})
}
}
@@ -0,0 +1,397 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"encoding/json"
"errors"
"net/http"
"strconv"
"strings"
"sync"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
pkgpush "github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway/push"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
const (
// 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"
)
// 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"`
}
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{},
})
}
// ListPushChannelDefinitions returns channel definitions.
func ListPushChannelDefinitions(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(ListPushDefinitions()))
}
// ListPushChannels lists configured push channels.
func ListPushChannels(c *gin.Context) {
channels, err := listPushChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(channels))
}
// 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"`
}
// CreatePushChannel creates a push channel.
func CreatePushChannel(c *gin.Context) {
var req CreatePushChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
channel, err := createPushChannel(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(channel))
}
// 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"`
}
// UpdatePushChannel updates a push channel.
func UpdatePushChannel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return
}
var req UpdatePushChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
channel, err := updatePushChannel(c.Request.Context(), id, req)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "channel not found")
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(channel))
}
// DeletePushChannel deletes a push channel.
func DeletePushChannel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return
}
if err := deletePushChannel(c.Request.Context(), id); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "channel not found")
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// 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"`
}
// TestPushChannel tests connectivity of a push channel.
func TestPushChannel(c *gin.Context) {
var req TestPushChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
ctx := c.Request.Context()
url, token, other, channelType, err := loadChannelForTest(ctx, req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if channelType == channelEmail {
url, token, other = resolveSMTPConfig(ctx, url, token, other)
}
tempChannel := PushChannel{
Name: "test_temp",
URL: url,
Token: token,
Other: other,
Type: channelType,
Enabled: true,
}
if err := tempChannel.Validate(); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
url = tempChannel.URL
var config pkgpush.Config
var renderedJSON string
switch channelType {
case channelLark:
config = pkgpush.Config{Channel: channelLark, URL: url, Secret: token}
renderedJSON = other
case channelEmail:
config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other}
case channelTelegram:
config = pkgpush.Config{Channel: channelTelegram, URL: url, Secret: token, Key: other}
default:
config = pkgpush.Config{Channel: channelCustom, URL: url}
customPushReq := CustomPushRequest{
Title: "通道测试通知",
Content: "这是一条来自系统的消息通道连通性测试消息。",
Description: "系统通道测试",
URL: "https://example.com",
To: req.Target,
}
renderedJSON = renderCustomPayload(other, customPushReq)
}
payload := SendPayload{
EventKey: "test_channel",
Config: config,
Target: req.Target,
Body: NotificationMessage{
Title: "通道测试通知",
Content: "这是一条来自系统的消息通道连通性测试消息。",
Level: defaultLevelInfo,
},
Template: renderedJSON,
}
if err := enqueuePushTask(ctx, payload); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// 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"`
}
func escapeJSONString(s string) string {
b, _ := json.Marshal(s)
const minJSONLen = 2
if len(b) >= minJSONLen {
return string(b[1 : len(b)-1])
}
return s
}
func renderCustomPayload(template string, req CustomPushRequest) string {
result := template
result = strings.ReplaceAll(result, "$title", escapeJSONString(req.Title))
result = strings.ReplaceAll(result, "$description", escapeJSONString(req.Description))
result = strings.ReplaceAll(result, "$content", escapeJSONString(req.Content))
result = strings.ReplaceAll(result, "$url", escapeJSONString(req.URL))
result = strings.ReplaceAll(result, "$to", escapeJSONString(req.To))
return result
}
@@ -0,0 +1,15 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
const (
channelCustom = "custom"
channelEmail = "email"
channelLark = "lark"
channelTelegram = "telegram"
defaultLevelInfo = "INFO"
keyTitle = "title"
keyContent = "content"
keyLevel = "level"
)
@@ -0,0 +1,274 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"encoding/json"
"errors"
"sync"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
pkgpush "github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway/push"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
"gorm.io/gorm"
)
// 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 异步投递推送载荷 (供 task/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"`
}
var (
builtInEventsMu sync.RWMutex
// BuiltInEvents lists all built-in events defined in custom_events.
BuiltInEvents []EventMetadata
)
// RegisterBuiltInEvent registers a built-in event definition.
func RegisterBuiltInEvent(meta EventMetadata) {
builtInEventsMu.Lock()
defer builtInEventsMu.Unlock()
for i, e := range BuiltInEvents {
if e.Key == meta.Key {
BuiltInEvents[i] = meta
return
}
}
BuiltInEvents = append(BuiltInEvents, meta)
}
// GetBuiltInEvents returns a copy of registered built-in events.
func GetBuiltInEvents() []EventMetadata {
builtInEventsMu.RLock()
defer builtInEventsMu.RUnlock()
out := make([]EventMetadata, len(BuiltInEvents))
copy(out, BuiltInEvents)
return out
}
// EventTrigger represents the unified event trigger class.
type EventTrigger struct{}
// DefaultTrigger is the singleton instance of EventTrigger.
var DefaultTrigger = &EventTrigger{}
// Trigger receives event metadata and processes the event notification dispatch asynchronously.
//
//nolint:contextcheck
func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map[string]any) {
asyncCtx := context.WithoutCancel(ctx)
util.Go(func() {
if body == nil {
body = make(map[string]any)
}
if _, hasUser := body["user"]; !hasUser || body["user"] == nil {
body["user"] = getSystemUser(asyncCtx)
}
eventPtr, err := GetActivePushEventByKey(asyncCtx, meta.Key)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return
}
logger.ErrorF(asyncCtx, "push_event_trigger: failed to get active event %s: %v", meta.Key, err)
return
}
event := *eventPtr
if len(event.Channels) == 0 {
return
}
flatBody := getFlatBody(body)
msg, _ := t.buildMessage(&event, meta, flatBody, body)
t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody)
})
}
func (t *EventTrigger) buildMessage(event *PushEvent, meta EventMetadata, flatBody map[string]any, body map[string]any) (NotificationMessage, string) {
var msg NotificationMessage
renderedTemplate := ""
templateSource := event.Template
if templateSource != "" {
var err error
msg, renderedTemplate, err = t.parseCustomTemplate(event, templateSource, flatBody)
if err != nil {
msg.Title = event.Name
msg.Content = renderedTemplate
msg.Level = defaultLevelInfo
}
} else {
msg = t.parseDefaultTemplate(meta, flatBody)
}
if msg.Ext == nil {
msg.Ext = make(map[string]any)
}
for k, v := range body {
if k == keyTitle || k == keyContent || k == keyLevel {
continue
}
if _, exists := msg.Ext[k]; !exists {
msg.Ext[k] = v
}
}
return msg, renderedTemplate
}
func (t *EventTrigger) parseCustomTemplate(event *PushEvent, templateSource string, flatBody map[string]any) (NotificationMessage, string, error) {
var msg NotificationMessage
renderedTemplate := pkgpush.ParseTemplate(templateSource, flatBody)
var tMap map[string]any
if err := json.Unmarshal([]byte(renderedTemplate), &tMap); err != nil {
return msg, renderedTemplate, err
}
if title, ok := tMap[keyTitle].(string); ok && title != "" {
msg.Title = title
} else {
msg.Title = event.Name
}
delete(tMap, keyTitle)
if content, ok := tMap[keyContent].(string); ok && content != "" {
msg.Content = content
} else {
msg.Content = renderedTemplate
}
delete(tMap, keyContent)
if level, ok := tMap[keyLevel].(string); ok && level != "" {
msg.Level = level
} else {
msg.Level = defaultLevelInfo
}
delete(tMap, keyLevel)
msg.Ext = tMap
return msg, renderedTemplate, nil
}
func (t *EventTrigger) parseDefaultTemplate(meta EventMetadata, flatBody map[string]any) NotificationMessage {
var msg NotificationMessage
msg.Title = pkgpush.ParseTemplate(meta.DefaultTemplate.Title, flatBody)
msg.Content = pkgpush.ParseTemplate(meta.DefaultTemplate.Content, flatBody)
msg.Level = pkgpush.ParseTemplate(meta.DefaultTemplate.Level, flatBody)
if meta.DefaultTemplate.Ext != nil {
msg.Ext = make(map[string]any)
for k, v := range meta.DefaultTemplate.Ext {
if strVal, ok := v.(string); ok {
msg.Ext[k] = pkgpush.ParseTemplate(strVal, flatBody)
} else {
msg.Ext[k] = v
}
}
}
return msg
}
func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta EventMetadata, event *PushEvent, msg NotificationMessage, flatBody map[string]any) {
for _, channelName := range event.Channels {
customChannel, err := GetActivePushChannelByName(ctx, channelName)
if err == nil {
t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody)
continue
}
logger.WarnF(ctx, "push_event_trigger: channel %q not found in DB or disabled: %v", channelName, err)
}
}
func (t *EventTrigger) enqueueCustomPushChannelTasks(ctx context.Context, meta EventMetadata, event *PushEvent, channel *PushChannel, msg NotificationMessage, flatBody map[string]any) {
if len(event.Targets) == 0 {
t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, "", msg)
return
}
for _, target := range event.Targets {
resolvedTarget := resolveTarget(ctx, target, flatBody, channel.Name)
t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, resolvedTarget, msg)
}
}
func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, meta EventMetadata, channel *PushChannel, target string, msg NotificationMessage) {
var config pkgpush.Config
var renderedTemplate string
switch channel.Type {
case channelLark:
config = pkgpush.Config{Channel: channelLark, URL: channel.URL, Secret: channel.Token}
renderedTemplate = channel.Other
case channelEmail:
url, token, other := resolveSMTPConfig(ctx, channel.URL, channel.Token, channel.Other)
config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other}
case channelTelegram:
config = pkgpush.Config{Channel: channelTelegram, URL: channel.URL, Secret: channel.Token, Key: channel.Other}
default:
config = pkgpush.Config{Channel: channelCustom, URL: channel.URL}
customPushReq := CustomPushRequest{
Title: msg.Title,
Content: msg.Content,
Description: meta.Description,
To: target,
}
if urlVal, ok := msg.Ext["url"].(string); ok {
customPushReq.URL = urlVal
}
renderedTemplate = renderCustomPayload(channel.Other, customPushReq)
}
payload := SendPayload{
EventKey: meta.Key,
Config: config,
Target: target,
Body: msg,
Template: renderedTemplate,
}
if err := enqueuePushTask(ctx, payload); err != nil {
logger.ErrorF(ctx, "push_event_trigger: enqueuePushTask failed for %s channel %s -> %s: %v", channel.Type, channel.Name, target, err)
}
}
// SyncEvents automatically registers/updates built-in events in the database.
func SyncEvents(ctx context.Context) error {
return syncBuiltInEvents(ctx)
}
@@ -0,0 +1,197 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"errors"
"fmt"
"net/http"
"strconv"
pkgpush "github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway/push"
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
// 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"`
}
// 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"`
}
// TestPushRequest is the request body for testing push config.
type TestPushRequest struct {
Config pkgpush.Config `json:"config" binding:"required"`
Target string `json:"target"`
}
// ListPushEvents lists configured push events.
func ListPushEvents(c *gin.Context) {
ctx := c.Request.Context()
events, err := 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(GetBuiltInEvents()))
}
// CreatePushEvent creates a new push event configuration.
func CreatePushEvent(c *gin.Context) {
var req CreatePushEventRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
event, err := createPushEvent(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(event))
}
// DeletePushEvent deletes a push event configuration by ID.
func DeletePushEvent(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
return
}
if err := deletePushEvent(c.Request.Context(), id); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// UpdatePushEvent updates an existing push event.
func UpdatePushEvent(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
return
}
var req UpdatePushEventRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := updatePushEvent(c.Request.Context(), id, req); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
return
}
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// TogglePushEvent toggles the enabled state of a push event.
func TogglePushEvent(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
return
}
enabled, err := togglePushEvent(c.Request.Context(), id)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
return
}
response.AbortBadRequest(c, err.Error())
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 := listPushHistories(c.Request.Context(), 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 TestPushRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
pusher, err := pkgpush.GetPusher(req.Config.Channel)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := pusher.ValidateConfig(req.Config); err != nil {
response.AbortBadRequest(c, fmt.Sprintf("validation failed: %v", err))
return
}
applySMTPFallbackToPushConfig(c.Request.Context(), &req.Config)
testBody := map[string]any{
keyTitle: "测试通道推送",
keyContent: "当您收到这条消息,说明当前渠道连通性测试通过。",
keyLevel: defaultLevelInfo,
}
if _, err := pusher.Send(c.Request.Context(), req.Config, req.Target, testBody, "", nil); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,518 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/backend/core/contracts"
pkgpush "github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway/push"
"github.com/Rain-kl/Wavelet/backend/plugins/drivers/driver_asynq_worker"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
"gorm.io/gorm"
)
type smtpConfig struct {
Host string
Port string
Username string
Password string
}
func loadSMTPConfig(ctx context.Context) smtpConfig {
var cfg smtpConfig
var host, port, user, pass string
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_host").Pluck("value", &host).Error
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_port").Pluck("value", &port).Error
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_username").Pluck("value", &user).Error
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_password").Pluck("value", &pass).Error
cfg.Host = host
cfg.Port = port
cfg.Username = user
cfg.Password = pass
return cfg
}
func syncBuiltInEvents(ctx context.Context) error {
for _, meta := range GetBuiltInEvents() {
_, err := GetPushEventByKeyRecord(ctx, meta.Key)
if errors.Is(err, gorm.ErrRecordNotFound) {
var defaultTemplateStr string
if defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate); err == nil {
defaultTemplateStr = string(defaultTemplateBytes)
}
event := PushEvent{
EventKey: meta.Key,
Name: meta.Name,
Channels: []string{},
Targets: []string{},
Template: defaultTemplateStr,
Enabled: false,
}
if err := CreatePushEventRecord(ctx, &event); err != nil {
return err
}
} else if err != nil {
return err
}
}
return nil
}
func listPushEvents(ctx context.Context) ([]PushEvent, error) {
return ListPushEventsRecord(ctx)
}
func createPushEvent(ctx context.Context, req CreatePushEventRequest) (PushEvent, error) {
eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req)
if err != nil {
return PushEvent{}, err
}
count, err := CountPushEventsByKeyRecord(ctx, eventKey)
if err != nil {
return PushEvent{}, err
}
if count > 0 {
return PushEvent{}, errors.New("this notification event is already configured")
}
templateStr := strings.TrimSpace(req.Template)
if templateStr == "" {
templateStr = string(defaultTemplateBytes)
} else {
var tempMap map[string]any
if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil {
return PushEvent{}, errors.New("custom template is not a valid JSON format")
}
}
channels := req.Channels
if channels == nil {
channels = []string{}
}
targets := req.Targets
if targets == nil {
targets = []string{}
}
event := PushEvent{
EventKey: eventKey,
Name: eventName,
TaskType: req.TaskType,
Channels: channels,
Targets: targets,
Template: templateStr,
Enabled: req.Enabled,
}
if err := event.Validate(); err != nil {
return PushEvent{}, err
}
if err := CreatePushEventRecord(ctx, &event); err != nil {
return PushEvent{}, err
}
return event, nil
}
func deletePushEvent(ctx context.Context, id uint64) error {
event, err := GetPushEventByIDRecord(ctx, id)
if err != nil {
return err
}
return DeletePushEventRecord(ctx, &event)
}
func updatePushEvent(ctx context.Context, id uint64, req UpdatePushEventRequest) error {
event, err := GetPushEventByIDRecord(ctx, id)
if err != nil {
return err
}
event.Channels = req.Channels
event.Targets = req.Targets
event.Template = req.Template
event.Enabled = req.Enabled
if err := event.Validate(); err != nil {
return err
}
return SavePushEventRecord(ctx, &event)
}
func togglePushEvent(ctx context.Context, id uint64) (bool, error) {
event, err := GetPushEventByIDRecord(ctx, id)
if err != nil {
return false, err
}
enabled := !event.Enabled
if enabled && len(event.Channels) == 0 {
return false, errors.New("cannot enable event without any push channels configured")
}
if err := UpdatePushEventEnabledRecord(ctx, &event, enabled); err != nil {
return false, err
}
return enabled, nil
}
func listPushHistories(ctx context.Context, filter PushHistoryListFilter) (int64, []PushHistory, error) {
return ListPushHistoriesRecord(ctx, filter)
}
func applySMTPFallbackToPushConfig(ctx context.Context, cfg *pkgpush.Config) {
if cfg.Channel != channelEmail || (cfg.URL != "" && cfg.Key != "") {
return
}
smtp := loadSMTPConfig(ctx)
if smtp.Host == "" || smtp.Username == "" {
return
}
port := smtp.Port
if port == "" {
port = "587"
}
cfg.URL = smtp.Host + ":" + port
cfg.Key = smtp.Username
cfg.Secret = smtp.Password
}
func listPushChannels(ctx context.Context) ([]PushChannel, error) {
return ListPushChannelsRecord(ctx)
}
func createPushChannel(ctx context.Context, req CreatePushChannelRequest) (PushChannel, error) {
count, err := CountPushChannelsByNameRecord(ctx, req.Name)
if err != nil {
return PushChannel{}, err
}
if count > 0 {
return PushChannel{}, errors.New("channel name already exists")
}
channel := PushChannel{
Name: req.Name,
Description: req.Description,
Type: req.Type,
Token: req.Token,
URL: req.URL,
Other: req.Other,
Enabled: req.Enabled,
}
if err := channel.Validate(); err != nil {
return PushChannel{}, err
}
if err := CreatePushChannelRecord(ctx, &channel); err != nil {
return PushChannel{}, err
}
return channel, nil
}
func updatePushChannel(ctx context.Context, id uint64, req UpdatePushChannelRequest) (PushChannel, error) {
channel, err := GetPushChannelByIDRecord(ctx, id)
if err != nil {
return PushChannel{}, err
}
channel.Description = req.Description
channel.Type = req.Type
channel.Token = req.Token
channel.URL = req.URL
channel.Other = req.Other
channel.Enabled = req.Enabled
if err := channel.Validate(); err != nil {
return PushChannel{}, err
}
if err := SavePushChannelRecord(ctx, &channel); err != nil {
return PushChannel{}, err
}
return channel, nil
}
func deletePushChannel(ctx context.Context, id uint64) error {
channel, err := GetPushChannelByIDRecord(ctx, id)
if err != nil {
return err
}
return DeletePushChannelRecord(ctx, &channel)
}
func loadChannelForTest(ctx context.Context, req TestPushChannelRequest) (string, string, string, string, error) {
if req.Name != "" {
channel, err := GetPushChannelByNameRecord(ctx, req.Name)
if err != nil {
return "", "", "", "", errors.New("channel not found")
}
return channel.URL, channel.Token, channel.Other, channel.Type, nil
}
return req.URL, req.Token, req.Other, req.Type, nil
}
func listActivePushEventsByTaskType(ctx context.Context, taskType string) ([]PushEvent, error) {
return ListActivePushEventsByTaskTypeRecord(ctx, taskType)
}
func loadUserFromPayload(ctx context.Context, data map[string]any) any {
if u, exists := data["user"]; exists && u != nil {
return u
}
if userID, ok := extractUserID(data); ok && userID > 0 {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err == nil {
return &user
}
}
if username := extractUsername(data); username != "" {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("username = ?", username).First(&user).Error; err == nil {
return &user
}
}
return nil
}
func recordPushHistory(ctx context.Context, req SendPayload, status, errMsg string) error {
title := req.Body.Title
content := req.Body.Content
level := req.Body.Level
if title == "" {
title = "系统通知"
}
if level == "" {
level = defaultLevelInfo
}
target := req.Target
if target == "" {
if req.Config.URL != "" {
target = req.Config.URL
const maxTargetLen = 50
const truncatedLen = 47
if len(target) > maxTargetLen {
target = target[:truncatedLen] + "..."
}
} else {
target = "default"
}
}
history := PushHistory{
EventKey: req.EventKey,
Channel: req.Config.Channel,
Target: target,
Title: title,
Content: content,
Level: level,
Status: status,
ErrorMsg: errMsg,
}
return CreatePushHistoryRecord(ctx, &history)
}
func resolveTarget(ctx context.Context, target string, flatBody map[string]any, channel string) string {
target = strings.TrimSpace(target)
if target == "" {
return ""
}
resolved := resolveDynamicKeyword(target, flatBody)
if strings.Contains(resolved, "@") {
return resolved
}
if val, matched := resolveSystemTarget(ctx, resolved, channel); matched {
return val
}
user, found := resolveTargetUser(ctx, resolved, channel)
if !found {
return resolved
}
if channel == channelEmail && user.Email != "" {
return user.Email
}
if channel != channelEmail && user.Username != "" {
return user.Username
}
return resolved
}
func resolveDynamicKeyword(target string, flatBody map[string]any) string {
switch target {
case "user.id", "id":
if val, ok := flatBody["user.id"]; ok {
return fmt.Sprintf("%v", val)
}
if val, ok := flatBody["id"]; ok {
return fmt.Sprintf("%v", val)
}
case "user.username", "username":
if val, ok := flatBody["user.username"]; ok {
return fmt.Sprintf("%v", val)
}
if val, ok := flatBody["username"]; ok {
return fmt.Sprintf("%v", val)
}
case "user.email", channelEmail:
if val, ok := flatBody["user.email"]; ok {
return fmt.Sprintf("%v", val)
}
if val, ok := flatBody["email"]; ok {
return fmt.Sprintf("%v", val)
}
}
return target
}
func resolveTargetUser(ctx context.Context, resolved string, _ string) (contracts.UserDTO, bool) {
var user contracts.UserDTO
if id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
if err := db.DB(ctx).Table("w_users").Where("id = ?", id).First(&user).Error; err == nil {
return user, true
}
}
if err := db.DB(ctx).Table("w_users").Where("username = ?", resolved).First(&user).Error; err == nil {
return user, true
}
return user, false
}
func resolveSystemTarget(ctx context.Context, resolved string, channel string) (string, bool) {
if resolved != "系统" && resolved != "system" && resolved != "0" {
return "", false
}
var adminUser contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
return resolved, true
}
if channel == channelEmail && adminUser.Email != "" {
return adminUser.Email, true
}
if channel != channelEmail && adminUser.Username != "" {
return adminUser.Username, true
}
return resolved, true
}
func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, string, string) {
if url != "" && token != "" {
return url, token, other
}
smtp := loadSMTPConfig(ctx)
if smtp.Host == "" || smtp.Username == "" {
return url, token, other
}
port := smtp.Port
if port == "" {
port = "587"
}
if url == "" {
url = smtp.Host + ":" + port
}
if token == "" {
token = smtp.Username
}
if other == "" {
other = smtp.Password
}
return url, token, other
}
func getSystemUser(ctx context.Context) *contracts.UserDTO {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&user).Error; err == nil {
return &user
}
return &contracts.UserDTO{
Username: "system",
Nickname: "系统管理员",
}
}
func findBuiltInEvent(key string) (EventMetadata, bool) {
for _, meta := range GetBuiltInEvents() {
if meta.Key == key {
return meta, true
}
}
return EventMetadata{}, false
}
func getEventInfo(req CreatePushEventRequest) (string, string, []byte, error) {
if req.TaskType != "" {
meta := driver_asynq_worker.GetTaskMetaByAsynqTask(req.TaskType)
if meta == nil {
return "", "", nil, errors.New("unsupported task type")
}
eventKey := "task_completed:" + req.TaskType
eventName := "任务完成: " + meta.Name
defaultTemplate := NotificationMessage{
Title: "任务完成: " + meta.Name,
Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。",
Level: defaultLevelInfo,
}
defaultTemplateBytes, err := json.Marshal(defaultTemplate)
if err != nil {
return "", "", nil, err
}
return eventKey, eventName, defaultTemplateBytes, nil
}
if req.EventKey == "" {
return "", "", nil, errors.New("either event_key or task_type must be provided")
}
meta, found := findBuiltInEvent(req.EventKey)
if !found {
return "", "", nil, errors.New("unsupported built-in event key")
}
defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate)
if err != nil {
return "", "", nil, err
}
return req.EventKey, meta.Name, defaultTemplateBytes, nil
}
func enqueuePushTask(ctx context.Context, payload SendPayload) error {
payloadBytes, err := json.Marshal(payload)
if err != nil {
return err
}
_, err = driver_asynq_worker.DispatchTask(ctx, "send_notification", payloadBytes, "system")
return err
}
func getFlatBody(body map[string]any) map[string]any {
jsonBytes, err := json.Marshal(body)
if err != nil {
return body
}
var jsonMap map[string]any
if err := json.Unmarshal(jsonBytes, &jsonMap); err != nil {
return body
}
flatResult := make(map[string]any)
flattenMap("", jsonMap, flatResult)
return flatResult
}
func flattenMap(prefix string, m map[string]any, result map[string]any) {
for k, v := range m {
key := k
if prefix != "" {
key = prefix + "." + k
}
if nestedMap, ok := v.(map[string]any); ok {
flattenMap(key, nestedMap, result)
} else {
result[key] = v
}
}
}
@@ -0,0 +1,123 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"encoding/json"
"strconv"
"time"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
"github.com/Rain-kl/Wavelet/backend/plugins/drivers/driver_asynq_worker"
)
// RegisterTaskListeners subscribes push notification handlers to task completion events.
func RegisterTaskListeners() {
driver_asynq_worker.OnTaskCompleted(handleTaskCompleted)
}
func handleTaskCompleted(ctx context.Context, execution *driver_asynq_worker.TaskExecution, result *driver_asynq_worker.TaskResult, execErr error) {
events, err := listActivePushEventsByTaskType(ctx, execution.TaskType)
if err != nil {
logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", execution.TaskType, err)
return
}
if len(events) == 0 {
return
}
body := map[string]any{
"task_id": execution.TaskID,
"task_name": execution.TaskName,
"task_type": execution.TaskType,
"task_status": string(execution.Status),
"task_duration": execution.Duration,
"time": time.Now().Format("2006-01-02 15:04:05"),
}
if execErr != nil {
body["task_error"] = execErr.Error()
} else {
body["task_error"] = ""
}
if result != nil {
body["task_result"] = result.Message
} else {
body["task_result"] = ""
}
var payloadMap map[string]any
if execution.Payload != "" {
if err := json.Unmarshal([]byte(execution.Payload), &payloadMap); err == nil {
body["payload"] = payloadMap
extractUserFromMap(ctx, payloadMap, body)
}
}
if result != nil && result.Detail != "" {
var detailMap map[string]any
if err := json.Unmarshal([]byte(result.Detail), &detailMap); err == nil {
body["detail"] = detailMap
extractUserFromMap(ctx, detailMap, body)
}
}
for _, event := range events {
meta := EventMetadata{
Key: event.EventKey,
Name: event.Name,
Description: "异步任务执行完毕触发的自动通知",
}
DefaultTrigger.Trigger(ctx, meta, body)
}
}
func extractUserFromMap(ctx context.Context, data map[string]any, body map[string]any) {
if u, exists := body["user"]; exists && u != nil {
return
}
if user := loadUserFromPayload(ctx, data); user != nil {
body["user"] = user
}
}
func extractUserID(data map[string]any) (uint64, bool) {
for _, k := range []string{"user_id", "userId", "uid"} {
val, ok := data[k]
if !ok || val == nil {
continue
}
switch v := val.(type) {
case float64:
if v >= 0 {
return uint64(v), true
}
case int:
if v >= 0 {
return uint64(v), true
}
case int64:
if v >= 0 {
return uint64(v), true
}
case uint64:
return v, true
case string:
if id, err := strconv.ParseUint(v, 10, 64); err == nil {
return id, true
}
}
}
return 0, false
}
func extractUsername(data map[string]any) string {
for _, k := range []string{"username", "user_name"} {
if val, ok := data[k]; ok && val != nil {
if s, ok := val.(string); ok && s != "" {
return s
}
}
}
return ""
}
@@ -0,0 +1,123 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"encoding/json"
"errors"
"fmt"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway/push"
"github.com/Rain-kl/Wavelet/backend/plugins/drivers/driver_asynq_worker"
)
const (
// SendNotificationTask is the asynq task name for push notification.
SendNotificationTask = "push:send"
// TaskTypeSendNotification is the admin task manager type identifier.
TaskTypeSendNotification = "send_notification"
)
// SendNotificationMeta represents the task metadata.
var SendNotificationMeta = driver_asynq_worker.TaskMeta{
Type: TaskTypeSendNotification,
AsynqTask: SendNotificationTask,
Name: "推送通知",
Description: "异步执行系统通知的多渠道派发与推送",
SupportsTime: false,
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
Queue: driver_asynq_worker.QueueDefault,
Retryable: true,
Params: []driver_asynq_worker.TaskParam{
{
Name: "event_key",
Label: "事件标识",
Type: "string",
Required: true,
Placeholder: "admin_login",
},
{
Name: "target",
Label: "目标接收者",
Type: "string",
Required: false,
},
},
}
// PushHandler handles asynchronous notification sending.
type PushHandler struct{}
// ValidatePayload validates and normalizes push parameters.
func (h *PushHandler) ValidatePayload(payload []byte) ([]byte, error) {
if len(payload) == 0 {
return nil, errors.New("payload is required")
}
var req SendPayload
if err := json.Unmarshal(payload, &req); err != nil {
return nil, fmt.Errorf("invalid json format: %w", err)
}
if req.Config.Channel == "" {
return nil, errors.New("channel type is required")
}
return json.Marshal(req)
}
// Execute performs the push send and logs delivery history audit.
func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
var req SendPayload
if err := json.Unmarshal(payload, &req); err != nil {
driver_asynq_worker.AppendLog(ctx, "解析推送参数失败: %v", err)
return nil, fmt.Errorf("parse payload failed: %w", err)
}
driver_asynq_worker.AppendLog(ctx, "开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target)
pusher, err := push.GetPusher(req.Config.Channel)
if err != nil {
errWrap := fmt.Errorf("get pusher failed: %w", err)
driver_asynq_worker.AppendLog(ctx, "推送失败: %v", errWrap)
if driver_asynq_worker.IsFinalAttempt(ctx) {
h.recordHistory(ctx, req, "failed", errWrap.Error())
}
return nil, errWrap
}
flatBody := req.Body.Flatten()
upstreamResp, err := pusher.Send(ctx, req.Config, req.Target, flatBody, req.Template, nil)
title := req.Body.Title
content := req.Body.Content
if err != nil {
driver_asynq_worker.AppendLog(ctx, "消息推送失败 (标题: %s): %v", title, err)
if upstreamResp != "" {
driver_asynq_worker.AppendLog(ctx, "上游返回: %s", upstreamResp)
}
if driver_asynq_worker.IsFinalAttempt(ctx) {
h.recordHistory(ctx, req, "failed", err.Error())
}
return nil, fmt.Errorf("pusher.Send failed: %w", err)
}
driver_asynq_worker.AppendLog(ctx, "消息推送成功 (标题: %s, 内容摘要: %s)", title, content)
if upstreamResp != "" {
driver_asynq_worker.AppendLog(ctx, "上游返回: %s", upstreamResp)
}
h.recordHistory(ctx, req, "success", "")
return &driver_asynq_worker.TaskResult{
Message: fmt.Sprintf("推送成功: [%s] -> %s", req.Config.Channel, req.Target),
}, nil
}
func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status string, errMsg string) {
if dbErr := recordPushHistory(ctx, req, status, errMsg); dbErr != nil {
driver_asynq_worker.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr)
}
}
@@ -0,0 +1,26 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import "sync"
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
}
@@ -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,393 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"errors"
"time"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/backend/pkg/idgen"
cachepkg "github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
)
const (
activePushChannelCacheTTL = 24 * time.Hour
activePushEventCacheTTL = 24 * time.Hour
)
// CreateMessageChannel inserts a channel row.
func CreateMessageChannel(ctx context.Context, ch *MessageChannel) error {
if ch.ID == 0 {
ch.ID = idgen.NextUint64ID()
}
return db.DB(ctx).Create(ch).Error
}
// UpdateMessageChannel saves a channel row.
func UpdateMessageChannel(ctx context.Context, ch *MessageChannel) error {
return db.DB(ctx).Save(ch).Error
}
// GetMessageChannel loads a channel by id.
func GetMessageChannel(ctx context.Context, id uint64) (*MessageChannel, error) {
var ch MessageChannel
if err := db.DB(ctx).Where("id = ?", id).First(&ch).Error; err != nil {
return nil, err
}
return &ch, nil
}
// ListMessageChannels returns all channels newest first.
func ListMessageChannels(ctx context.Context) ([]MessageChannel, error) {
var rows []MessageChannel
if err := db.DB(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 db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("channel_id = ?", id).Delete(&MessagePairingCode{}).Error; err != nil {
return err
}
if err := tx.Where("channel_id = ?", id).Delete(&MessageBinding{}).Error; err != nil {
return err
}
return tx.Delete(&MessageChannel{}, id).Error
})
}
// CreateMessageBinding inserts a binding.
func CreateMessageBinding(ctx context.Context, b *MessageBinding) error {
if b.ID == 0 {
b.ID = idgen.NextUint64ID()
}
return db.DB(ctx).Create(b).Error
}
// GetBindingByChannelPlatform finds a binding for a platform user on a channel.
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*MessageBinding, error) {
var b MessageBinding
err := db.DB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error
if err != nil {
return nil, err
}
return &b, nil
}
// ListBindingsByUser lists bindings for a Wavelet user.
func ListBindingsByUser(ctx context.Context, userID uint64) ([]MessageBinding, error) {
var rows []MessageBinding
if err := db.DB(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) (*MessageBinding, error) {
var b MessageBinding
if err := db.DB(ctx).Where("id = ?", id).First(&b).Error; err != nil {
return nil, err
}
return &b, nil
}
// DeleteMessageBinding deletes a binding by id.
func DeleteMessageBinding(ctx context.Context, id uint64) error {
return db.DB(ctx).Delete(&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) (*MessagePairingCode, error) {
var existing MessagePairingCode
err := db.DB(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 := &MessagePairingCode{
Code: code,
ChannelID: channelID,
PlatformUserID: platformUserID,
ExpiresAt: expiresAt,
}
if err := db.DB(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) (*MessagePairingCode, error) {
var row MessagePairingCode
if err := db.DB(ctx).Where("code = ?", code).First(&row).Error; err != nil {
return nil, err
}
return &row, nil
}
// DeletePairingCode removes a pairing code.
func DeletePairingCode(ctx context.Context, code string) error {
return db.DB(ctx).Where("code = ?", code).Delete(&MessagePairingCode{}).Error
}
// DeleteExpiredPairingCodes removes expired pairing rows.
func DeleteExpiredPairingCodes(ctx context.Context) error {
return db.DB(ctx).Where("expires_at <= ?", time.Now()).Delete(&MessagePairingCode{}).Error
}
// ListEnabledMessageChannels returns enabled channels.
func ListEnabledMessageChannels(ctx context.Context) ([]MessageChannel, error) {
var rows []MessageChannel
if err := db.DB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// ListPushChannelsRecord returns all push channels ordered by creation time descending.
func ListPushChannelsRecord(ctx context.Context) ([]PushChannel, error) {
var channels []PushChannel
if err := db.DB(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) (PushChannel, error) {
var channel PushChannel
if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
return PushChannel{}, err
}
return channel, nil
}
// GetPushChannelByNameRecord 根据名称获取消息通道。
func GetPushChannelByNameRecord(ctx context.Context, name string) (*PushChannel, error) {
var channel PushChannel
if err := db.DB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
return nil, 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 := db.DB(ctx).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 *PushChannel) error {
if err := db.DB(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 *PushChannel) error {
if err := db.DB(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 *PushChannel) error {
if err := db.DB(ctx).Delete(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
// GetActivePushChannelByName 根据名称获取启用的消息通道 (优先从 Redis 缓存获取)。
func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel, error) {
cacheKey := "push:channel:active:" + name
var channel PushChannel
if cachepkg.Redis != nil {
if err := cachepkg.GetJSON(ctx, cacheKey, &channel); err == nil {
return &channel, nil
}
}
if err := db.DB(ctx).Where("name = ? AND enabled = ?", name, true).First(&channel).Error; err != nil {
return nil, err
}
if cachepkg.Redis != nil {
_ = cachepkg.SetJSON(ctx, cacheKey, channel, activePushChannelCacheTTL)
}
return &channel, nil
}
// DeleteActivePushChannelCache 清理启用消息通道的缓存。
func DeleteActivePushChannelCache(ctx context.Context, name string) {
if cachepkg.Redis != nil {
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey("push:channel:active:"+name)).Err()
}
}
// ListPushEventsRecord returns all push events ordered by creation time descending.
func ListPushEventsRecord(ctx context.Context) ([]PushEvent, error) {
var events []PushEvent
if err := db.DB(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) (PushEvent, error) {
var event PushEvent
if err := db.DB(ctx).First(&event, id).Error; err != nil {
return PushEvent{}, err
}
return event, nil
}
// GetPushEventByKeyRecord loads a push event by event key.
func GetPushEventByKeyRecord(ctx context.Context, key string) (PushEvent, error) {
var event PushEvent
if err := db.DB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
return PushEvent{}, 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 := db.DB(ctx).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 *PushEvent) error {
if err := db.DB(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 *PushEvent) error {
if err := db.DB(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 *PushEvent, enabled bool) error {
event.Enabled = enabled
if err := db.DB(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 *PushEvent) error {
if err := db.DB(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) ([]PushEvent, error) {
var events []PushEvent
if err := db.DB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil {
return nil, err
}
return events, nil
}
// GetActivePushEventByKey 获取启用的通知事件 (优先从 Redis 缓存获取)。
func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error) {
cacheKey := "push:event:active:" + key
var event PushEvent
if cachepkg.Redis != nil {
if err := cachepkg.GetJSON(ctx, cacheKey, &event); err == nil {
return &event, nil
}
}
if err := db.DB(ctx).Where("event_key = ? AND enabled = ?", key, true).First(&event).Error; err != nil {
return nil, err
}
if cachepkg.Redis != nil {
_ = cachepkg.SetJSON(ctx, cacheKey, event, activePushEventCacheTTL)
}
return &event, nil
}
// DeleteActivePushEventCache 清理启用通知事件的缓存。
func DeleteActivePushEventCache(ctx context.Context, key string) {
if cachepkg.Redis != nil {
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey("push:event:active:"+key)).Err()
}
}
// ListPushHistoriesRecord returns paginated push history records.
func ListPushHistoriesRecord(ctx context.Context, filter PushHistoryListFilter) (int64, []PushHistory, error) {
query := db.DB(ctx).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 []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 *PushHistory) error {
return db.DB(ctx).Create(history).Error
}
// PushHistoryQuery returns a scoped query builder for push histories.
func PushHistoryQuery(ctx context.Context) *gorm.DB {
return db.DB(ctx).Model(&PushHistory{})
}
@@ -0,0 +1,53 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"sync"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
)
// 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
}
@@ -0,0 +1,78 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"github.com/Rain-kl/Wavelet/backend/pkg/config"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
)
// CredentialKey is AES-256 hex derived from the session secret.
func CredentialKey() string {
secret := ""
if config.Config != nil {
secret = config.Config.App.SessionSecret
}
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)
}
@@ -0,0 +1,141 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package risk_control
import (
"context"
"sync"
"time"
"github.com/Rain-kl/Wavelet/backend/pkg/batchwriter"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
"github.com/Rain-kl/Wavelet/backend/plugins/domain/risk_control/logstore"
)
var (
logWriterMu sync.RWMutex
logWriter *batchwriter.Writer[*logstore.UserAccessLog]
)
// InitLogWriter initializes the access-log batch writer for the active log database.
func InitLogWriter(ctx context.Context) {
logWriterMu.Lock()
defer logWriterMu.Unlock()
if logWriter != nil {
return
}
cfg := batchwriter.DefaultConfig()
writer, err := batchwriter.New[*logstore.UserAccessLog](cfg, func(ctx context.Context, items []*logstore.UserAccessLog) error {
rows := make([]logstore.UserAccessLog, 0, len(items))
for _, item := range items {
if item == nil {
continue
}
rows = append(rows, *item)
}
store, err := logstore.Active(ctx)
if err != nil {
return err
}
return store.UserAccessLogs.BatchInsert(ctx, rows)
},
batchwriter.WithDropHandler[*logstore.UserAccessLog](func(item *logstore.UserAccessLog) {
path := ""
if item != nil {
path = item.Path
}
logger.WarnF(context.Background(), "[RiskControl] Log queue full, dropping log item for path: %s", path)
}),
batchwriter.WithFlushErrorHandler[*logstore.UserAccessLog](func(ctx context.Context, items []*logstore.UserAccessLog, err error) {
logger.ErrorF(ctx, "[RiskControl] flush access-log batch failed (batch=%d): %v", len(items), err)
}),
)
if err != nil {
logger.ErrorF(ctx, "[RiskControl] init log writer failed: %v", err)
return
}
writer.Start(ctx)
logWriter = writer
}
// StopLogWriter stops the ClickHouse access-log batch writer and drains pending logs.
func StopLogWriter(ctx context.Context) error {
writer := currentLogWriter()
if writer == nil {
return nil
}
return writer.Stop(ctx)
}
// IsBufferFull reports whether the access-log queue has no remaining capacity.
func IsBufferFull() bool {
writer := currentLogWriter()
if writer == nil {
return false
}
return writer.IsFull()
}
// QueueAccessLog enqueues an access log without blocking.
func QueueAccessLog(logItem *logstore.UserAccessLog) {
writer := currentLogWriter()
if writer == nil || logItem == nil {
return
}
writer.TryEnqueue(logItem)
}
// SetLogWriterForTest swaps the access-log writer for unit tests.
func SetLogWriterForTest(writer *batchwriter.Writer[*logstore.UserAccessLog]) func() {
logWriterMu.Lock()
previous := logWriter
logWriter = writer
logWriterMu.Unlock()
return func() {
logWriterMu.Lock()
logWriter = previous
logWriterMu.Unlock()
}
}
func currentLogWriter() *batchwriter.Writer[*logstore.UserAccessLog] {
logWriterMu.RLock()
defer logWriterMu.RUnlock()
return logWriter
}
const drainPollInterval = 50 * time.Millisecond
// Drain waits until the in-memory access-log queue has been empty for one flush interval.
func Drain(ctx context.Context) error {
writer := currentLogWriter()
if writer == nil {
return nil
}
quietPeriod := batchwriter.DefaultConfig().FlushInterval
if quietPeriod <= 0 {
quietPeriod = time.Second
}
ticker := time.NewTicker(drainPollInterval)
defer ticker.Stop()
var quietSince time.Time
for {
if writer.Len() == 0 {
if quietSince.IsZero() {
quietSince = time.Now()
} else if time.Since(quietSince) >= quietPeriod {
return nil
}
} else {
quietSince = time.Time{}
}
select {
case <-ctx.Done():
return ctx.Err()
case <-ticker.C:
}
}
}
@@ -0,0 +1,119 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package logstore provides data access for analytics tables.
package logstore
import (
"context"
"fmt"
"time"
"github.com/Rain-kl/Wavelet/backend/pkg/util"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
"gorm.io/gorm"
)
// CountAccessLogs returns the number of access logs matching filter.
func CountAccessLogs(ctx context.Context, filter AccessLogFilter) (uint64, error) {
ch := db.ChDB(ctx)
if ch == nil {
return 0, fmt.Errorf("clickhouse gorm connection is not initialized")
}
var count int64
query := applyFilter(ch.Model(&UserAccessLog{}), filter)
if err := query.Count(&count).Error; err != nil {
return 0, fmt.Errorf("count access logs: %w", err)
}
return safeUint64Count(count), nil
}
// ListAccessLogs returns paginated access logs and the total match count.
func ListAccessLogs(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]UserAccessLog, uint64, error) {
ch := db.ChDB(ctx)
if ch == nil {
return nil, 0, fmt.Errorf("clickhouse gorm connection is not initialized")
}
if filter.UserIDs != nil && len(filter.UserIDs) == 0 {
return []UserAccessLog{}, 0, nil
}
var total int64
baseQuery := applyFilter(ch.Model(&UserAccessLog{}), filter)
if err := baseQuery.Count(&total).Error; err != nil {
return nil, 0, fmt.Errorf("count access logs: %w", err)
}
if total == 0 {
return []UserAccessLog{}, 0, nil
}
if page < 1 {
page = 1
}
if pageSize < 1 {
pageSize = 20
}
offset := (page - 1) * pageSize
var logs []UserAccessLog
err := applyFilter(ch.Model(&UserAccessLog{}), filter).
Order("created_at DESC, id DESC").
Limit(pageSize).
Offset(offset).
Find(&logs).Error
if err != nil {
return nil, 0, fmt.Errorf("list access logs: %w", err)
}
return logs, safeUint64Count(total), nil
}
// DeleteAllUserAccessLogs hard-deletes all user access logs via TRUNCATE.
func DeleteAllUserAccessLogs(ctx context.Context) (int64, error) {
if db.ChConn == nil {
return 0, fmt.Errorf("clickhouse connection is not initialized")
}
if err := db.ChConn.Exec(ctx, "TRUNCATE TABLE "+UserAccessLog{}.TableName()); err != nil {
return 0, fmt.Errorf("truncate user access logs: %w", err)
}
return 0, nil
}
// DeleteUserAccessLogsBefore deletes user access logs older than cutoff.
func DeleteUserAccessLogsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
if db.ChConn == nil {
return 0, fmt.Errorf("clickhouse connection is not initialized")
}
if err := db.ChConn.Exec(ctx, "ALTER TABLE "+UserAccessLog{}.TableName()+" DELETE WHERE created_at < ?", cutoff); err != nil {
return 0, fmt.Errorf("delete expired user access logs: %w", err)
}
return 0, nil
}
func safeUint64Count(count int64) uint64 {
if count < 0 {
return 0
}
return uint64(count)
}
func applyFilter(query *gorm.DB, filter AccessLogFilter) *gorm.DB {
if filter.UserIDs != nil {
if len(filter.UserIDs) == 0 {
return query.Where("1 = 0")
}
query = query.Where("user_id IN ?", filter.UserIDs)
}
if filter.Path != "" {
query = query.Where("path LIKE ?", "%"+util.EscapeLike(filter.Path)+"%")
}
if filter.StartTime != nil {
query = query.Where("created_at >= ?", *filter.StartTime)
}
if filter.EndTime != nil {
query = query.Where("created_at <= ?", *filter.EndTime)
}
return query
}
@@ -0,0 +1,35 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import "time"
// AccessLogFilter scopes user access log queries.
type AccessLogFilter struct {
// UserIDs filters by user IDs. nil means no user filter; an empty slice means no matches.
UserIDs []uint64
Path string
// StartTime filters created_at >= StartTime when non-nil.
StartTime *time.Time
// EndTime filters created_at <= EndTime when non-nil.
EndTime *time.Time
}
// DailyTrend is a single day's access count.
type DailyTrend struct {
Date string
Count uint64
}
// BrowserShare is a browser group's share of access logs.
type BrowserShare struct {
Browser string
Count uint64
}
// TopUser is an active user ranked by access count.
type TopUser struct {
UserID uint64
Count uint64
}
@@ -0,0 +1,140 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"fmt"
"sort"
"time"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
)
const hoursInDay = 24
// GetDailyTrend returns per-day access counts for the last days days (inclusive of today).
func GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
if days < 1 {
days = 7
}
ch := db.ChDB(ctx)
if ch == nil {
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
}
startTime := time.Now().AddDate(0, 0, -(days - 1)).Truncate(hoursInDay * time.Hour)
tableName := UserAccessLog{}.TableName()
query := fmt.Sprintf(`
SELECT toDate(created_at) AS date, count() AS count
FROM %s
WHERE created_at >= ?
GROUP BY date
ORDER BY date ASC
`, tableName)
type trendRow struct {
Date time.Time
Count uint64
}
var rows []trendRow
if err := ch.Raw(query, startTime).Scan(&rows).Error; err != nil {
return nil, fmt.Errorf("get daily trend: %w", err)
}
trendMap := make(map[string]uint64, days)
for i := 0; i < days; i++ {
dateStr := time.Now().AddDate(0, 0, -i).Format("2006-01-02")
trendMap[dateStr] = 0
}
for _, row := range rows {
dateStr := row.Date.Format("2006-01-02")
trendMap[dateStr] = row.Count
}
result := make([]DailyTrend, 0, days)
for i := days - 1; i >= 0; i-- {
dateStr := time.Now().AddDate(0, 0, -i).Format("2006-01-02")
result = append(result, DailyTrend{
Date: dateStr,
Count: trendMap[dateStr],
})
}
return result, nil
}
// GetBrowserDistribution returns browser-grouped access counts since startTime.
func GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error) {
ch := db.ChDB(ctx)
if ch == nil {
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
}
tableName := UserAccessLog{}.TableName()
query := fmt.Sprintf(`
SELECT user_agent, count() AS count
FROM %s
WHERE created_at >= ?
GROUP BY user_agent
`, tableName)
type uaRow struct {
UserAgent string
Count uint64
}
var rows []uaRow
if err := ch.Raw(query, startTime).Scan(&rows).Error; err != nil {
return nil, fmt.Errorf("get browser distribution: %w", err)
}
browserCounts := make(map[string]uint64)
for _, row := range rows {
browser := ParseBrowserName(row.UserAgent)
browserCounts[browser] += row.Count
}
result := make([]BrowserShare, 0, len(browserCounts))
for browser, count := range browserCounts {
result = append(result, BrowserShare{
Browser: browser,
Count: count,
})
}
sort.Slice(result, func(i, j int) bool {
return result[i].Count > result[j].Count
})
return result, nil
}
// GetTopActiveUsers returns the most active users since startTime.
func GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]TopUser, error) {
if limit < 1 {
limit = 10
}
ch := db.ChDB(ctx)
if ch == nil {
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
}
tableName := UserAccessLog{}.TableName()
query := fmt.Sprintf(`
SELECT user_id, count() AS count
FROM %s
WHERE created_at >= ? AND user_id > 0
GROUP BY user_id
ORDER BY count DESC
LIMIT ?
`, tableName)
var users []TopUser
if err := ch.Raw(query, startTime, limit).Scan(&users).Error; err != nil {
return nil, fmt.Errorf("get top active users: %w", err)
}
return users, nil
}
@@ -0,0 +1,215 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"io"
"testing"
"time"
"github.com/ClickHouse/clickhouse-go/v2/lib/column"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
)
func setupChGormDB(t *testing.T) *gorm.DB {
t.Helper()
gormDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, gormDB.AutoMigrate(&UserAccessLog{}))
db.SetChDBForTest(gormDB)
return gormDB
}
func TestParseBrowserName(t *testing.T) {
tests := []struct {
name string
ua string
want string
}{
{name: "chrome", ua: "Mozilla/5.0 Chrome/120.0.0.0", want: "Chrome"},
{name: "firefox", ua: "Mozilla/5.0 Firefox/121.0", want: "Firefox"},
{name: "safari", ua: "Mozilla/5.0 Safari/605.1.15", want: "Safari"},
{name: "edge", ua: "Mozilla/5.0 Edg/120.0.0.0", want: "Edge"},
{name: "wechat", ua: "MicroMessenger/8.0", want: "WeChat"},
{name: "postman", ua: "PostmanRuntime/7.36.0", want: "Postman"},
{name: "other", ua: "curl/8.0", want: "Other"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
assert.Equal(t, tt.want, ParseBrowserName(tt.ua))
})
}
}
func TestCountAccessLogs_EmptyUserIDs(t *testing.T) {
setupChGormDB(t)
t.Cleanup(func() { db.SetChDBForTest(nil) })
count, err := CountAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}})
require.NoError(t, err)
assert.Equal(t, uint64(0), count)
}
func TestListAccessLogs_EmptyUserIDs(t *testing.T) {
setupChGormDB(t)
t.Cleanup(func() { db.SetChDBForTest(nil) })
logs, total, err := ListAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}}, 1, 20)
require.NoError(t, err)
assert.Equal(t, uint64(0), total)
assert.Empty(t, logs)
}
func TestListAccessLogs_WithFilters(t *testing.T) {
gormDB := setupChGormDB(t)
t.Cleanup(func() { db.SetChDBForTest(nil) })
now := time.Now().UTC().Truncate(time.Second)
logs := []UserAccessLog{
{ID: 1, UserID: 10, Path: "/api/v1/users", Method: "GET", Status: 200, CreatedAt: now},
{ID: 2, UserID: 20, Path: "/api/v1/admin/logs", Method: "GET", Status: 200, CreatedAt: now},
{ID: 3, UserID: 10, Path: "/api/v1/other", Method: "POST", Status: 201, CreatedAt: now},
}
require.NoError(t, gormDB.Create(&logs).Error)
start := now.Add(-time.Hour)
filter := AccessLogFilter{
UserIDs: []uint64{10},
Path: "users",
StartTime: &start,
}
count, err := CountAccessLogs(context.Background(), filter)
require.NoError(t, err)
assert.Equal(t, uint64(1), count)
result, total, err := ListAccessLogs(context.Background(), filter, 1, 10)
require.NoError(t, err)
assert.Equal(t, uint64(1), total)
require.Len(t, result, 1)
assert.Equal(t, uint64(1), result[0].ID)
assert.Equal(t, "/api/v1/users", result[0].Path)
}
func TestBatchInsert_Empty(t *testing.T) {
err := BatchInsert(context.Background(), nil)
require.NoError(t, err)
}
func TestBatchInsert_UsesModelBatchSQL(t *testing.T) {
ctx := context.Background()
mockBatch := &mockBatch{}
mockConn := &mockConn{
batch: mockBatch,
batchQuery: UserAccessLog{}.BatchInsertSQL(),
}
db.SetChConnForTest(mockConn)
t.Cleanup(func() { db.SetChConnForTest(nil) })
createdAt := time.Now().UTC()
err := BatchInsert(ctx, []UserAccessLog{
{
ID: 1,
UserID: 42,
Path: "/api/v1/test",
Method: "GET",
IP: "127.0.0.1",
UserAgent: "test-agent",
Headers: "{}",
Status: 200,
Latency: 12,
CreatedAt: createdAt,
},
})
require.NoError(t, err)
assert.True(t, mockConn.prepareCalled)
assert.Equal(t, UserAccessLog{}.BatchInsertSQL(), mockConn.preparedQuery)
assert.True(t, mockBatch.sendCalled)
require.Len(t, mockBatch.rows, 1)
assert.Equal(t, uint64(42), mockBatch.rows[0][1])
}
type mockConn struct {
batch driver.Batch
batchQuery string
prepareCalled bool
preparedQuery string
}
func (m *mockConn) Contributors() []string { return nil }
func (m *mockConn) ServerVersion() (*driver.ServerVersion, error) { return nil, nil }
func (m *mockConn) Select(_ context.Context, _ any, _ string, _ ...any) error { return nil }
func (m *mockConn) Query(_ context.Context, _ string, _ ...any) (driver.Rows, error) {
return nil, nil
}
func (m *mockConn) QueryRow(_ context.Context, _ string, _ ...any) driver.Row { return nil }
func (m *mockConn) PrepareBatch(_ context.Context, query string, _ ...driver.PrepareBatchOption) (driver.Batch, error) {
m.prepareCalled = true
m.preparedQuery = query
return m.batch, nil
}
func (m *mockConn) Exec(_ context.Context, _ string, _ ...any) error { return nil }
func (m *mockConn) AsyncInsert(_ context.Context, _ string, _ bool, _ ...any) error { return nil }
func (m *mockConn) InsertFormat(_ context.Context, _ string, _ string, _ io.Reader) error { return nil }
func (m *mockConn) QueryFormat(_ context.Context, _ string, _ string, _ ...any) (io.ReadCloser, error) {
return nil, nil
}
func (m *mockConn) Ping(_ context.Context) error { return nil }
func (m *mockConn) Stats() driver.Stats { return driver.Stats{} }
func (m *mockConn) Close() error { return nil }
type mockBatch struct {
rows [][]any
sendCalled bool
}
func (m *mockBatch) Abort() error { return nil }
func (m *mockBatch) Append(v ...any) error {
m.rows = append(m.rows, v)
return nil
}
func (m *mockBatch) AppendStruct(_ any) error { return nil }
func (m *mockBatch) Column(_ int) driver.BatchColumn { return nil }
func (m *mockBatch) Flush() error { return nil }
func (m *mockBatch) Send() error {
m.sendCalled = true
return nil
}
func (m *mockBatch) IsSent() bool { return m.sendCalled }
func (m *mockBatch) Rows() int { return len(m.rows) }
func (m *mockBatch) Columns() []column.Interface { return nil }
func (m *mockBatch) Close() error { return nil }
@@ -0,0 +1,48 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"fmt"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
)
// BatchInsert writes access logs to ClickHouse using the native batch API.
func BatchInsert(ctx context.Context, logs []UserAccessLog) error {
if len(logs) == 0 {
return nil
}
if db.ChConn == nil {
return fmt.Errorf("clickhouse connection is not initialized")
}
batch, err := db.ChConn.PrepareBatch(ctx, UserAccessLog{}.BatchInsertSQL())
if err != nil {
return fmt.Errorf("prepare clickhouse batch: %w", err)
}
for _, logItem := range logs {
if err := batch.Append(
logItem.ID,
logItem.UserID,
logItem.Path,
logItem.Method,
logItem.IP,
logItem.UserAgent,
logItem.Headers,
logItem.Status,
logItem.Latency,
logItem.CreatedAt,
); err != nil {
return fmt.Errorf("append access log to batch: %w", err)
}
}
if err := batch.Send(); err != nil {
return fmt.Errorf("send clickhouse batch: %w", err)
}
return nil
}
@@ -0,0 +1,30 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import "strings"
// ParseBrowserName performs lightweight User-Agent browser identification.
func ParseBrowserName(ua string) string {
uaLower := strings.ToLower(ua)
if strings.Contains(uaLower, "micromessenger") {
return "WeChat"
}
if strings.Contains(uaLower, "postman") {
return "Postman"
}
if strings.Contains(uaLower, "edg/") || strings.Contains(uaLower, "edge") {
return "Edge"
}
if strings.Contains(uaLower, "firefox") {
return "Firefox"
}
if strings.Contains(uaLower, "chrome") {
return "Chrome"
}
if strings.Contains(uaLower, "safari") {
return "Safari"
}
return "Other"
}
@@ -0,0 +1,99 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"errors"
"fmt"
"strconv"
"time"
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
)
const (
defaultLogRetentionDays = 30
partitionLeadMonths = 2
userAccessLogTable = "w_user_access_logs"
)
// CleanupSummary 汇总本次清理结果。
type CleanupSummary struct {
ActiveDatabase string `json:"active_database"`
RetentionDays int `json:"retention_days"`
Deleted int64 `json:"deleted"`
}
// CleanupExpired 按当前日志库保留天数删除过期用户访问日志,并预建 PG 分区。
func CleanupExpired(ctx context.Context) (CleanupSummary, error) {
active, err := ActiveDatabase(ctx)
if err != nil {
return CleanupSummary{}, err
}
days := retentionDaysForDatabase(ctx, active)
summary := CleanupSummary{ActiveDatabase: active, RetentionDays: days}
store, err := Active(ctx)
if err != nil {
return summary, err
}
now := time.Now().UTC()
if err := store.UserAccessLogs.EnsurePartitions(ctx, now, now.AddDate(0, partitionLeadMonths, 0)); err != nil {
logger.WarnF(ctx, "logstore: ensure partitions during cleanup failed: %v", err)
}
cutoff := now.AddDate(0, 0, -days)
// 先 DROP 完全过期的整月分区,再对边界月逐行 DeleteBefore。
if err := store.UserAccessLogs.DropExpiredPartitions(ctx, cutoff); err != nil {
return summary, fmt.Errorf("drop expired partitions: %w", err)
}
deleted, err := store.UserAccessLogs.DeleteBefore(ctx, cutoff)
if err != nil {
return summary, fmt.Errorf("delete expired user access logs: %w", err)
}
summary.Deleted = deleted
if err := store.UserAccessLogs.DropEmptyPartitions(ctx, now); err != nil {
logger.WarnF(ctx, "drop empty log partitions failed: %v", err)
}
return summary, nil
}
func retentionDaysForDatabase(ctx context.Context, dbName string) int {
key := "log_retention_days_postgres"
switch dbName {
case dbNameSQLite:
key = "log_retention_days_sqlite"
case dbNameClickHouse:
key = "log_retention_days_clickhouse"
}
v, err := getConfig(ctx, key)
if err != nil {
if !errors.Is(err, errConfigReaderNotWired) {
logger.ErrorF(ctx, "读取日志保留天数配置失败(key=%s),回退默认 %d 天: %v", key, defaultLogRetentionDays, err)
}
return defaultLogRetentionDays
}
days, perr := strconv.Atoi(v)
if perr != nil || days <= 0 {
logger.ErrorF(ctx, "日志保留天数配置非法(key=%s, value=%q),回退默认 %d 天", key, v, defaultLogRetentionDays)
return defaultLogRetentionDays
}
return days
}
func partitionStatementsRange(from, to time.Time) []string {
var out []string
start := time.Date(from.Year(), from.Month(), 1, 0, 0, 0, 0, time.UTC)
end := time.Date(to.Year(), to.Month(), 1, 0, 0, 0, 0, time.UTC).AddDate(0, 1, 0)
for ; start.Before(end); start = start.AddDate(0, 1, 0) {
monthEnd := start.AddDate(0, 1, 0)
suffix := start.Format("200601")
fromDay := start.Format("2006-01-02")
toDay := monthEnd.Format("2006-01-02")
out = append(out, fmt.Sprintf(
"CREATE TABLE IF NOT EXISTS %s_%s PARTITION OF %s FOR VALUES FROM ('%s') TO ('%s')",
userAccessLogTable, suffix, userAccessLogTable, fromDay, toDay))
}
return out
}
@@ -0,0 +1,152 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"fmt"
"time"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
)
type clickhouseUserAccessLogStore struct {
skipFreeze bool
}
func newClickHouseUserAccessLogStore() *clickhouseUserAccessLogStore {
return &clickhouseUserAccessLogStore{}
}
var (
_ UserAccessLogStore = (*clickhouseUserAccessLogStore)(nil)
_ StatusStore = (*clickhouseUserAccessLogStore)(nil)
)
func (s *clickhouseUserAccessLogStore) ActiveDatabase(_ context.Context) (string, error) {
return dbNameClickHouse, nil
}
func (s *clickhouseUserAccessLogStore) ensureWritable(ctx context.Context) error {
if !s.skipFreeze && Migrating(ctx) {
return ErrMigrating
}
return nil
}
func (s *clickhouseUserAccessLogStore) BatchInsert(ctx context.Context, logs []UserAccessLog) error {
if len(logs) == 0 {
return nil
}
if err := s.ensureWritable(ctx); err != nil {
return err
}
return BatchInsert(ctx, logs)
}
func (s *clickhouseUserAccessLogStore) DeleteAll(ctx context.Context) (int64, error) {
if err := s.ensureWritable(ctx); err != nil {
return 0, err
}
return DeleteAllUserAccessLogs(ctx)
}
func (s *clickhouseUserAccessLogStore) DeleteBefore(ctx context.Context, cutoff time.Time) (int64, error) {
if err := s.ensureWritable(ctx); err != nil {
return 0, err
}
return DeleteUserAccessLogsBefore(ctx, cutoff)
}
func (s *clickhouseUserAccessLogStore) Count(ctx context.Context, filter AccessLogFilter) (uint64, error) {
return CountAccessLogs(ctx, filter)
}
func (s *clickhouseUserAccessLogStore) List(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]UserAccessLog, uint64, error) {
return ListAccessLogs(ctx, filter, page, pageSize)
}
func (s *clickhouseUserAccessLogStore) GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
return GetDailyTrend(ctx, days)
}
func (s *clickhouseUserAccessLogStore) GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error) {
return GetBrowserDistribution(ctx, startTime)
}
func (s *clickhouseUserAccessLogStore) GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]TopUser, error) {
return GetTopActiveUsers(ctx, startTime, limit)
}
func (s *clickhouseUserAccessLogStore) EnsurePartitions(_ context.Context, _, _ time.Time) error {
return nil
}
func (s *clickhouseUserAccessLogStore) DropEmptyPartitions(_ context.Context, _ time.Time) error {
return nil
}
func (s *clickhouseUserAccessLogStore) DropExpiredPartitions(_ context.Context, _ time.Time) error {
return nil
}
func (s *clickhouseUserAccessLogStore) MigrationRange(ctx context.Context) (time.Time, time.Time, error) {
if db.ChConn == nil {
return time.Time{}, time.Time{}, fmt.Errorf("clickhouse connection is not initialized")
}
table := UserAccessLog{}.TableName()
var minTime, maxTime *time.Time
if err := db.ChConn.QueryRow(ctx, "SELECT min(created_at), max(created_at) FROM "+table).Scan(&minTime, &maxTime); err != nil {
return time.Time{}, time.Time{}, fmt.Errorf("query migration range %s: %w", table, err)
}
if minTime == nil || maxTime == nil {
return time.Time{}, time.Time{}, nil
}
return minTime.UTC(), maxTime.UTC(), nil
}
func (s *clickhouseUserAccessLogStore) ListForMigration(ctx context.Context, afterID uint64, limit int) ([]UserAccessLog, error) {
if db.ChConn == nil {
return nil, fmt.Errorf("clickhouse connection is not initialized")
}
if limit <= 0 {
limit = migrationPageSize
}
table := UserAccessLog{}.TableName()
columns := UserAccessLog{}.InsertColumns()
rows, err := db.ChConn.Query(ctx, fmt.Sprintf(
"SELECT %s FROM %s WHERE id > ? ORDER BY id ASC LIMIT ?",
columns, table,
), afterID, limit)
if err != nil {
return nil, fmt.Errorf("list user access logs for migration: %w", err)
}
defer func() { _ = rows.Close() }()
return scanUserAccessLogs(rows)
}
func scanUserAccessLogs(rows driver.Rows) ([]UserAccessLog, error) {
var result []UserAccessLog
for rows.Next() {
var item UserAccessLog
if err := rows.Scan(
&item.ID,
&item.UserID,
&item.Path,
&item.Method,
&item.IP,
&item.UserAgent,
&item.Headers,
&item.Status,
&item.Latency,
&item.CreatedAt,
); err != nil {
return nil, fmt.Errorf("scan user access log row: %w", err)
}
item.CreatedAt = item.CreatedAt.UTC()
result = append(result, item)
}
return result, nil
}
@@ -0,0 +1,325 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"errors"
"fmt"
"sort"
"strings"
"time"
"github.com/Rain-kl/Wavelet/backend/pkg/idgen"
"gorm.io/gorm"
)
const (
insertBatchSize = 500
migrationPageSize = 100
defaultPageSize = 20
defaultTopN = 10
topUserAgents = 100
dayDuration = 24 * time.Hour
)
type gormLogStore struct {
db *gorm.DB
skipFreeze bool
}
func newGormStore(db *gorm.DB) *gormLogStore { return &gormLogStore{db: db} }
type userAccessLogGormStore struct {
*gormLogStore
}
func newUserAccessLogGormStore(db *gorm.DB) *userAccessLogGormStore {
return &userAccessLogGormStore{gormLogStore: newGormStore(db)}
}
var (
_ UserAccessLogStore = (*userAccessLogGormStore)(nil)
_ StatusStore = (*userAccessLogGormStore)(nil)
)
func (s *gormLogStore) ActiveDatabase(_ context.Context) (string, error) {
if isPostgresDialect(s.db) {
return dbNamePostgres, nil
}
return dbNameSQLite, nil
}
func (s *gormLogStore) ensureWritable(ctx context.Context) error {
if !s.skipFreeze && Migrating(ctx) {
return ErrMigrating
}
return nil
}
func (s *userAccessLogGormStore) BatchInsert(ctx context.Context, logs []UserAccessLog) error {
if len(logs) == 0 {
return nil
}
if err := s.ensureWritable(ctx); err != nil {
return err
}
for i := range logs {
if logs[i].ID == 0 {
logs[i].ID = idgen.NextUint64ID()
}
}
return s.db.WithContext(ctx).CreateInBatches(logs, insertBatchSize).Error
}
func (s *userAccessLogGormStore) DeleteAll(ctx context.Context) (int64, error) {
if err := s.ensureWritable(ctx); err != nil {
return 0, err
}
res := s.db.WithContext(ctx).Where("1 = 1").Delete(&UserAccessLog{})
return res.RowsAffected, res.Error
}
func (s *userAccessLogGormStore) DeleteBefore(ctx context.Context, cutoff time.Time) (int64, error) {
if err := s.ensureWritable(ctx); err != nil {
return 0, err
}
res := s.db.WithContext(ctx).Where("created_at < ?", cutoff).Delete(&UserAccessLog{})
if res.Error != nil && isMissingRelation(res.Error) {
return 0, nil
}
return res.RowsAffected, res.Error
}
func (s *userAccessLogGormStore) ListForMigration(ctx context.Context, afterID uint64, limit int) ([]UserAccessLog, error) {
var rows []UserAccessLog
q := s.db.WithContext(ctx).Model(&UserAccessLog{}).
Where("id > ?", afterID).
Order("id ASC").
Limit(limitOr(limit, migrationPageSize))
if err := q.Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
func (s *userAccessLogGormStore) MigrationRange(ctx context.Context) (time.Time, time.Time, error) {
return gormMigrationRange(ctx, s.db, "created_at", UserAccessLog{}, func(v *UserAccessLog) time.Time {
return v.CreatedAt
})
}
func (s *userAccessLogGormStore) Count(ctx context.Context, filter AccessLogFilter) (uint64, error) {
where, args, ok := buildUserAccessLogWhere(filter)
if !ok {
return 0, nil
}
var total int64
if err := s.db.WithContext(ctx).Model(&UserAccessLog{}).Where(where, args...).Count(&total).Error; err != nil {
return 0, err
}
return countToUint64(total), nil
}
func (s *userAccessLogGormStore) List(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]UserAccessLog, uint64, error) {
where, args, ok := buildUserAccessLogWhere(filter)
if !ok {
return []UserAccessLog{}, 0, nil
}
var total int64
if err := s.db.WithContext(ctx).Model(&UserAccessLog{}).Where(where, args...).Count(&total).Error; err != nil {
return nil, 0, err
}
if total == 0 {
return []UserAccessLog{}, 0, nil
}
var rows []UserAccessLog
q := s.db.WithContext(ctx).Where(where, args...).Order("created_at DESC, id DESC")
if err := q.Limit(limitOr(pageSize, defaultPageSize)).Offset(offsetOf(page, pageSize)).Find(&rows).Error; err != nil {
return nil, 0, err
}
return rows, countToUint64(total), nil
}
func buildUserAccessLogWhere(filter AccessLogFilter) (string, []any, bool) {
if filter.UserIDs != nil && len(filter.UserIDs) == 0 {
return "", nil, false
}
var parts []string
var args []any
if filter.UserIDs != nil {
parts = append(parts, "user_id IN ?")
args = append(args, filter.UserIDs)
}
if trimmed := strings.TrimSpace(filter.Path); trimmed != "" {
parts = append(parts, "path LIKE ?")
args = append(args, "%"+trimmed+"%")
}
if filter.StartTime != nil {
parts = append(parts, "created_at >= ?")
args = append(args, *filter.StartTime)
}
if filter.EndTime != nil {
parts = append(parts, "created_at <= ?")
args = append(args, *filter.EndTime)
}
if len(parts) == 0 {
return "1 = 1", args, true
}
return strings.Join(parts, " AND "), args, true
}
func (s *userAccessLogGormStore) GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
if days <= 0 {
days = 7
}
start := time.Now().AddDate(0, 0, -(days - 1)).Truncate(dayDuration)
type row struct {
Date string
Cnt uint64
}
var rows []row
err := s.db.WithContext(ctx).Model(&UserAccessLog{}).
Select(dailyTrendDateSQL(s.db)+" AS date, COUNT(*) AS cnt").
Where("created_at >= ?", start).
Group("date").Order("date ASC").Scan(&rows).Error
if err != nil {
return nil, err
}
counts := make(map[string]uint64, len(rows))
for _, r := range rows {
counts[r.Date] = r.Cnt
}
out := make([]DailyTrend, 0, days)
for i := 0; i < days; i++ {
d := start.AddDate(0, 0, i).Format("2006-01-02")
out = append(out, DailyTrend{Date: d, Count: counts[d]})
}
return out, nil
}
func (s *userAccessLogGormStore) GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error) {
type row struct {
UserAgent string
Cnt uint64
}
var rows []row
err := s.db.WithContext(ctx).Model(&UserAccessLog{}).
Select("user_agent, COUNT(*) AS cnt").
Where("created_at >= ?", startTime).
Group("user_agent").Order("cnt DESC").Limit(topUserAgents).Scan(&rows).Error
if err != nil {
return nil, err
}
counts := make(map[string]uint64)
for _, r := range rows {
counts[ParseBrowserName(r.UserAgent)] += r.Cnt
}
out := make([]BrowserShare, 0, len(counts))
for label, count := range counts {
out = append(out, BrowserShare{Browser: label, Count: count})
}
sort.Slice(out, func(i, j int) bool { return out[i].Count > out[j].Count })
return out, nil
}
func (s *userAccessLogGormStore) GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]TopUser, error) {
type row struct {
UserID uint64
Cnt uint64
}
var rows []row
err := s.db.WithContext(ctx).Model(&UserAccessLog{}).
Select("user_id, COUNT(*) AS cnt").
Where("user_id <> 0 AND created_at >= ?", startTime).
Group("user_id").Order("cnt DESC").Limit(limitOr(limit, defaultTopN)).Scan(&rows).Error
if err != nil {
return nil, err
}
out := make([]TopUser, len(rows))
for i, r := range rows {
out[i] = TopUser{UserID: r.UserID, Count: r.Cnt}
}
return out, nil
}
func (s *userAccessLogGormStore) EnsurePartitions(ctx context.Context, from, to time.Time) error {
if !isPostgresDialect(s.db) {
return nil
}
for _, sql := range partitionStatementsRange(from, to) {
if err := s.db.WithContext(ctx).Exec(sql).Error; err != nil {
return fmt.Errorf("ensure partition: %w", err)
}
}
return nil
}
func gormMigrationRange[T any](
ctx context.Context,
gdb *gorm.DB,
column string,
model T,
timeOf func(*T) time.Time,
) (time.Time, time.Time, error) {
var first, last T
found := false
for _, order := range []string{"ASC", "DESC"} {
out := &first
if order == "DESC" {
out = &last
}
res := gdb.WithContext(ctx).Model(model).Order(column + " " + order).Limit(1).Take(out)
if res.Error != nil && !errors.Is(res.Error, gorm.ErrRecordNotFound) {
return time.Time{}, time.Time{}, fmt.Errorf("query migration range %s: %w", column, res.Error)
}
if res.Error == nil {
found = true
}
}
if !found {
return time.Time{}, time.Time{}, nil
}
return timeOf(&first).UTC(), timeOf(&last).UTC(), nil
}
func limitOr(v, def int) int {
if v <= 0 {
return def
}
return v
}
func offsetOf(page, pageSize int) int {
if page < 1 {
page = 1
}
return (page - 1) * limitOr(pageSize, defaultPageSize)
}
func countToUint64(v int64) uint64 {
if v < 0 {
return 0
}
return uint64(v)
}
func isPostgresDialect(db *gorm.DB) bool {
return db != nil && db.Dialector != nil && db.Name() == "postgres"
}
func dailyTrendDateSQL(db *gorm.DB) string {
if isPostgresDialect(db) {
return "to_char(created_at, 'YYYY-MM-DD')"
}
return "strftime('%Y-%m-%d', created_at)"
}
func isMissingRelation(err error) bool {
if err == nil {
return false
}
msg := strings.ToLower(err.Error())
return strings.Contains(msg, "no such table") || strings.Contains(msg, "does not exist")
}

Some files were not shown because too many files have changed in this diff Show More