refactor(plugins): complete physical encapsulation of auth, admin, message_gateway, and risk_control domain plugins

This commit is contained in:
ryan
2026-08-28 07:15:17 +08:00
parent e750fadacd
commit b259f35bb4
51 changed files with 9735 additions and 42 deletions
+1 -1
View File
@@ -9,13 +9,13 @@ import (
"errors"
"fmt"
"github.com/Rain-kl/Wavelet/internal/apps/risk_control"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control"
)
const (
+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 = "更新用户信息失败"
)
+105
View File
@@ -0,0 +1,105 @@
// 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/internal/infra/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
)
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, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := saveOrUpdateCacheConfig(ctx, model.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 repository.SaveOrUpdateSystemConfig(ctx, key, value)
}
+534
View File
@@ -0,0 +1,534 @@
// 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/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/apps/cap"
"github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/logger"
mail "github.com/Rain-kl/Wavelet/pkg/mail"
)
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 := repository.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 := repository.GetBoolByKey(ctx, model.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=[]model.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=model.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 == model.ConfigKeyLogDatabase || key == model.ConfigKeyLogDBMigration
}
func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) error {
if isProtectedConfigKey(req.Key) {
return errors.New(protectedConfigKeyMessage)
}
exists, err := repository.SystemConfigExists(ctx, req.Key)
if err != nil {
return err
}
if exists {
return errors.New(ConfigKeyExists)
}
config := model.SystemConfig{
Key: req.Key,
Value: req.Value,
Type: req.Type,
Visibility: req.Visibility,
Description: req.Description,
}
if err := repository.CreateSystemConfig(ctx, &config); err != nil {
return err
}
invalidateSystemConfigCaches(ctx, req.Key)
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
return nil
}
func listSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) {
return repository.ListAdminSystemConfigs(ctx, configType)
}
func getSystemConfig(ctx context.Context, key string) (model.SystemConfig, error) {
return repository.GetAdminSystemConfigByKey(ctx, key)
}
func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigRequest) error {
if isProtectedConfigKey(key) {
return errors.New(protectedConfigKeyMessage)
}
config, err := repository.GetAdminSystemConfigByKey(ctx, key)
if err != nil {
return err
}
var originalDriver objectstore.Driver
if key == model.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 != model.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 != model.ConfigKeyStorageConfig || originalDriver == "" {
return
}
var newCfg objectstore.Config
if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil {
return
}
if newCfg.Driver != originalDriver {
return
}
if err := repository.MarkFailedTaskExecutionsSucceededTx(
tx,
"storage:migrate",
"存储配置直接更新,故障迁移任务自动标记为已解决",
time.Now(),
); err != nil {
logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err)
}
}
func invalidateSystemConfigCaches(ctx context.Context, key string) {
if err := repository.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 == model.ConfigKeyStorageConfig {
upload.ResetAccessCaches()
upload.PublishAccessCacheInvalidation(ctx)
objectstore.ResetCache()
objectstore.PublishCacheInvalidation(ctx)
}
if key == model.ConfigKeyFileAccessWhitelist {
upload.ResetAccessCaches()
upload.PublishAccessCacheInvalidation(ctx)
}
if err := repository.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 := repository.GetSystemConfigByKey(c.Request.Context(), model.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 model.ConfigKeySMTPPassword:
return maskedConfigValue
case model.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).Model(&model.Upload{}).
Where("status != ?", model.UploadStatusDeleted).
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/internal/infra/config"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/shared/response"
)
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)
}
}
+676
View File
@@ -0,0 +1,676 @@
// 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/gin-gonic/gin"
"github.com/gorilla/websocket"
"github.com/Rain-kl/Wavelet/internal/apps/risk_control"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/util"
)
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 != "" {
userIDs, err := repository.ListUserIDsByUsernameContains(ctx, username)
if 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 })
if users, err := repository.ListUsersByIDs(ctx, userIDs); 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
})
users, errProfile := repository.ListUsersByIDs(ctx, userIDs)
if 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 := repository.GetSystemConfigByKey(ctx, model.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 = task.TaskMeta{
Type: TaskTypeLogDBSwitch,
AsynqTask: LogDBSwitchTask,
Name: "切换日志数据库",
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
Params: []task.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) (*task.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 {
task.AppendLog(ctx, "读取日志主库失败: %v", err)
return nil, err
}
task.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()
task.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
return &task.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 := repository.GetSystemConfigByKey(ctx, model.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 repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDBMigration, v)
}
func flipLogDatabase(ctx context.Context, target string) error {
return repository.SaveOrUpdateSystemConfig(ctx, model.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)
task.AppendLog(ctx, "已复制用户访问日志 %d 条", copied)
if len(rows) < copyBatchSize {
break
}
}
return nil
}
+236
View File
@@ -0,0 +1,236 @@
// 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/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
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, model.ConfigKeyLogRetentionDaysPostgres),
logDBNameSQLite: retentionOr(ctx, model.ConfigKeyLogRetentionDaysSQLite),
logDBNameClickHouse: retentionOr(ctx, model.ConfigKeyLogRetentionDaysClickHouse),
},
AvailableTargets: availableLogTargets(activeDB),
}))
}
func retentionOr(ctx context.Context, key string) int {
v, err := repository.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{}
}
+415
View File
@@ -0,0 +1,415 @@
// 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/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/infra/task/scheduler"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
// ListTaskTypes 获取支持的任务类型列表
// @Summary 获取支持的任务类型
// @Description 返回系统支持的所有可调度任务类型列表,包括任务名称、描述、是否支持时间范围等元数据,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]task.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(task.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 := task.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 := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
taskID, err := task.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 model.ListTaskExecutionsRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if req.TaskType != "" {
if meta := task.GetTaskMeta(req.TaskType); meta != nil {
req.TaskType = meta.AsynqTask
}
}
executions, total, err := repository.ListTaskExecutions(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=model.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 := repository.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 := task.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=[]model.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 := repository.ListSchedules(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=model.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 := task.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 := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
schedule := &model.Schedule{
Name: req.Name,
TaskType: req.TaskType,
Cron: req.Cron,
Payload: string(validated),
IsActive: *req.IsActive,
}
if err := repository.CreateSchedule(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
// 触发调度服务重载
if err := scheduler.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=model.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 := repository.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 := task.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 := task.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 := repository.UpdateSchedule(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
// 触发调度服务重载
if err := scheduler.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 := repository.DeleteSchedule(c.Request.Context(), id); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))
return
}
// 触发调度服务重载
if err := scheduler.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
c.JSON(http.StatusOK, response.OKNil())
}
+245
View File
@@ -0,0 +1,245 @@
// 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/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/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=[]model.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=model.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=model.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) (model.Template, error) {
exists, err := repository.TemplateExistsByKey(ctx, req.Key)
if err != nil {
return model.Template{}, err
}
if exists {
return model.Template{}, errors.New(TemplateKeyExists)
}
tmpl := model.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 model.Template{}, err
}
if err := repository.CreateTemplate(ctx, &tmpl); err != nil {
return model.Template{}, err
}
return tmpl, nil
}
func listTemplates(ctx context.Context) ([]model.Template, error) {
return repository.ListTemplates(ctx)
}
func getTemplate(ctx context.Context, key string) (model.Template, error) {
return repository.GetTemplateByKey(ctx, key)
}
func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest) (model.Template, error) {
tmpl, err := repository.GetTemplateByKey(ctx, key)
if err != nil {
return model.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 model.Template{}, err
}
if err := repository.SaveTemplate(ctx, &tmpl); err != nil {
return model.Template{}, err
}
return tmpl, nil
}
func deleteTemplate(ctx context.Context, key string) error {
tmpl, err := repository.GetTemplateByKey(ctx, key)
if err != nil {
return err
}
if tmpl.IsSystem {
return errors.New(SystemTemplateCannotDelete)
}
return repository.DeleteTemplate(ctx, &tmpl)
}
+697
View File
@@ -0,0 +1,697 @@
// 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/internal/buildinfo"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/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 := repository.GetSystemConfigByKey(ctx, model.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
}
+550
View File
@@ -0,0 +1,550 @@
// 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/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
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 model.User) 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, modelUsers, 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(modelUsers))
for _, modelUser := range modelUsers {
users = append(users, toUserResponse(modelUser))
}
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, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
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, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
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, []model.User, error) {
return repository.ListAdminUsers(ctx, repository.AdminUserListFilter{
UserID: req.UserID,
Username: strings.TrimSpace(req.Username),
Email: strings.TrimSpace(req.Email),
Page: req.Page,
PageSize: req.PageSize,
})
}
func getUserDetail(ctx context.Context, id uint64) (model.User, error) {
return repository.GetAdminUserDetail(ctx, id)
}
func updateUserStatus(ctx context.Context, id uint64, active bool) error {
flags, err := repository.GetUserAdminFlags(ctx, id)
if err != nil {
return err
}
if !active && flags.IsAdmin {
return errors.New(cannotDisable)
}
var tokens []model.AccessToken
if !active {
tokens, _ = repository.ListAccessTokensByUserID(ctx, id)
}
err = repository.UpdateUserActive(ctx, id, active)
if err == nil {
oauth.InvalidateCachedUser(ctx, id)
if !active {
for _, token := range tokens {
oauth.InvalidateCachedToken(ctx, token.TokenHash)
}
}
}
return err
}
func deleteUser(ctx context.Context, currentUserID, targetID uint64) error {
if currentUserID == targetID {
return errors.New(cannotDeleteSelf)
}
flags, err := repository.GetUserAdminFlags(ctx, targetID)
if err != nil {
return err
}
if flags.IsAdmin {
return errors.New(cannotDelete)
}
tokens, _ := repository.ListAccessTokensByUserID(ctx, targetID)
err = repository.DeleteUserWithRelations(ctx, targetID)
if err == nil {
oauth.InvalidateCachedUser(ctx, targetID)
for _, token := range tokens {
oauth.InvalidateCachedToken(ctx, token.TokenHash)
}
}
return err
}
func createUser(ctx context.Context, req createUserRequest) (model.User, 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 model.User{}, errors.New(usernameRequired)
}
if req.Email == "" {
return model.User{}, errors.New(emailRequired)
}
if len(req.Password) < minPasswordLength {
return model.User{}, errors.New(passwordTooShort)
}
count, err := repository.CountUsersByUsername(ctx, req.Username)
if err != nil {
return model.User{}, err
}
if count > 0 {
return model.User{}, errors.New(usernameExists)
}
emailCount, err := repository.CountUsersByEmail(ctx, req.Email)
if err != nil {
return model.User{}, err
}
if emailCount > 0 {
return model.User{}, errors.New(emailExists)
}
newUser := model.User{
ID: idgen.NextUint64ID(),
Username: req.Username,
Nickname: req.Nickname,
Email: req.Email,
IsActive: req.IsActive,
IsAdmin: req.IsAdmin,
LastLoginAt: time.Time{},
}
if newUser.Nickname == "" {
newUser.Nickname = req.Username
}
if err := newUser.SetEncryptedPassword(req.Password); err != nil {
return model.User{}, err
}
if err := repository.CreateUser(ctx, &newUser); err != nil {
return model.User{}, 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)
}
targetUser, err := repository.GetAdminUserDetail(ctx, param.ID)
if err != nil {
return err
}
// 不能撤销当前登录用户的管理员权限
if currentUserID == param.ID && !param.IsAdmin && targetUser.IsAdmin {
return errors.New(cannotRevokeSelfAdmin)
}
// 如果修改了邮箱,检查邮箱是否被其他用户占用
if targetUser.Email != param.Email {
count, err := repository.CountUsersByEmail(ctx, param.Email)
if 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 tokens []model.AccessToken
if needRevokeTokens {
tokens, _ = repository.ListAccessTokensByUserID(ctx, param.ID)
}
targetUser.Nickname = param.Nickname
if targetUser.Nickname == "" {
targetUser.Nickname = targetUser.Username
}
targetUser.Email = param.Email
targetUser.IsAdmin = param.IsAdmin
if param.Password != "" {
if err := targetUser.SetEncryptedPassword(param.Password); err != nil {
return err
}
}
err = repository.UpdateUser(ctx, &targetUser)
if err == nil {
oauth.InvalidateCachedUser(ctx, param.ID)
if needRevokeTokens {
for _, token := range tokens {
oauth.InvalidateCachedToken(ctx, token.TokenHash)
}
}
}
return err
}
+44
View File
@@ -0,0 +1,44 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/logger"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/gin-gonic/gin"
)
// LoginAdminRequired 返回管理员权限校验中间件
func LoginAdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
ctx, span := otel_trace.Start(c.Request.Context(), "LoginAdminRequired")
defer span.End()
user, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
if user == nil {
response.AbortNotFound(c, AdminRequired)
return
}
// 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限
if tokenAuth, _ := oauth.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth {
tokenAdmin, _ := oauth.GetFromContext[bool](c, oauth.TokenAdminKey)
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,73 @@
-- +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);
CREATE TABLE IF NOT EXISTS w_schedules (
id BIGINT PRIMARY KEY,
name VARCHAR(128) NOT NULL,
task_type VARCHAR(64) NOT NULL,
cron VARCHAR(64) NOT NULL,
payload TEXT,
is_active BOOLEAN NOT NULL DEFAULT TRUE,
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_schedules_is_active ON w_schedules (is_active);
CREATE TABLE IF NOT EXISTS w_task_executions (
id BIGINT PRIMARY KEY,
task_id VARCHAR(128) NOT NULL UNIQUE,
task_type VARCHAR(64) NOT NULL,
task_name VARCHAR(128),
status VARCHAR(32) NOT NULL,
retryable BOOLEAN NOT NULL DEFAULT FALSE,
max_retry INTEGER NOT NULL DEFAULT 0,
retry_count INTEGER NOT NULL DEFAULT 0,
log TEXT,
error_message TEXT,
result TEXT,
started_at TIMESTAMPTZ,
finished_at TIMESTAMPTZ,
duration BIGINT,
payload TEXT,
triggered_by VARCHAR(32) NOT NULL DEFAULT 'system',
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_task_executions_task_type ON w_task_executions (task_type);
CREATE INDEX IF NOT EXISTS idx_w_task_executions_status ON w_task_executions (status);
CREATE INDEX IF NOT EXISTS idx_w_task_executions_started_at ON w_task_executions (started_at);
CREATE INDEX IF NOT EXISTS idx_w_task_executions_created_at ON w_task_executions (created_at);
-- +goose StatementEnd
-- +goose Down
-- +goose StatementBegin
DROP TABLE IF EXISTS w_task_executions;
DROP TABLE IF EXISTS w_schedules;
DROP TABLE IF EXISTS w_templates;
DROP TABLE IF EXISTS w_system_configs;
-- +goose StatementEnd
@@ -0,0 +1,56 @@
-- +goose Up
-- +goose StatementBegin
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)
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'
);
-- +goose StatementEnd
+203
View File
@@ -0,0 +1,203 @@
// 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"
}
// 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 = "email"
}
}
// 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"
}
+50
View File
@@ -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/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())
}
+12
View File
@@ -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)
}
+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/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
)
// LogForAudit 将登录鉴权审计日志写入 Logger
func LogForAudit(ctx context.Context, user *model.User, 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)
}
}
+236
View File
@@ -0,0 +1,236 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"errors"
"fmt"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
)
func isOIDCLoginEnabled(ctx context.Context) bool {
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
if err != nil {
return true
}
return enabled
}
func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) {
name := strings.TrimSpace(strings.ToLower(sourceName))
if name == "" {
sources, err := repository.GetActiveAuthSourcesCached(ctx)
if err != nil {
return nil, err
}
if len(sources) == 0 {
return nil, errors.New(errNoActiveAuthSource)
}
return repository.GetAuthSourceByNameCached(ctx, sources[0].Name)
}
return repository.GetAuthSourceByNameCached(ctx, name)
}
func activeLoginSources(ctx context.Context) []AuthSourceView {
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
if err == nil && !enabled {
return nil
}
dbSources, err := repository.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) {
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
if err != nil || strings.TrimSpace(sc.Value) == "" {
return "", errors.New(errServerAddressMissing)
}
return strings.TrimRight(sc.Value, "/") + "/login", nil
}
func buildOAuthConfig(ctx context.Context, source *model.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 uniqueUsername(ctx context.Context, base string) (string, error) {
base = strings.TrimSpace(base)
if base == "" {
base = "user"
}
existingUsernames, err := repository.ListUsernamesMatchingBase(ctx, base)
if 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 buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, 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 := &model.OAuthUserInfo{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 *model.OAuthUserInfo) 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 *model.OAuthUserInfo) 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 *model.User, status string) OAuthCallbackResult {
result := OAuthCallbackResult{Status: status}
if user != nil {
info := BuildBasicUserInfo(user, false)
result.User = &info
}
return result
}
+252
View File
@@ -0,0 +1,252 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"fmt"
"strconv"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
"github.com/Rain-kl/Wavelet/pkg/util"
)
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"
)
var (
tokenRAM = ram.MustNew[string, *model.AccessToken](ram.Options{MaximumSize: 2048})
userRAM = ram.MustNew[uint64, *model.User](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 获取缓存的 AccessToken
func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken, error) {
ensureTokenCacheListener()
if val, ok := tokenRAM.GetIfPresent(tokenHash); ok {
return val, nil
}
if db.Redis != nil {
var token model.AccessToken
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 设置 AccessToken 缓存
func SetCachedToken(ctx context.Context, tokenHash string, token *model.AccessToken) {
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 获取缓存的 User
func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) {
ensureUserCacheListener()
if val, ok := userRAM.GetIfPresent(userID); ok {
return val, nil
}
if db.Redis != nil {
var u model.User
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 设置 User 缓存
func SetCachedUser(ctx context.Context, userID uint64, u *model.User) {
ensureUserCacheListener()
userRAM.Set(userID, u)
if db.Redis != nil {
key := userCacheKey(userID)
_ = db.SetJSON(ctx, key, u, userCacheTTL)
}
}
// InvalidateCachedUser 吊销/失效 User 缓存
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()
}
+125
View File
@@ -0,0 +1,125 @@
// 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/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
)
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 := &model.AccessToken{
ID: 123,
UserID: 456,
TokenHash: tokenHash,
Name: "test-token",
}
// 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 {
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 := &model.User{
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"
)
+37
View File
@@ -0,0 +1,37 @@
// 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 = "令牌无管理员权限"
)
+449
View File
@@ -0,0 +1,449 @@
// 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/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/listener"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/logger"
"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 := db.Redis.Set(ctx, db.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 *model.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 db.Redis == nil || sessionHash == "" {
return nil
}
key := db.PrefixedKey(fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash))
n, err := db.Redis.Incr(ctx, key).Result()
if err != nil {
return err
}
if n == 1 {
_ = db.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, shared.UnAuthorized)
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 := db.Redis.Set(ctx, db.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 := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
payloadRaw, err := db.Redis.Get(ctx, stateKey).Result()
if err != nil {
response.AbortBadRequest(c, errInvalidState)
return
}
_ = db.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, shared.UnAuthorized)
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 *model.AuthSource, userInfo *model.OAuthUserInfo) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, shared.UnAuthorized)
return
}
user, err := repository.GetUserByID(ctx, userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := repository.BindExternalAccount(ctx, &model.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()
_ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt)
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound")))
}
func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
var user model.User
account, err := repository.FindExternalAccount(ctx, source.ID, userInfo.Sub)
switch {
case err == nil:
loaded, loadErr := repository.GetUserByID(ctx, account.UserID)
if loadErr != nil {
response.AbortInternal(c, loadErr.Error())
return
}
user = loaded
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()
_ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt)
if err := SetLoginSession(ctx, c, &user); err != nil {
response.AbortInternal(c, err.Error())
return
}
SetCachedUser(ctx, user.ID, &user)
logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP())
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in")))
}
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) {
registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
if regErr != nil {
registrationEnabled = false
}
if !registrationEnabled {
c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind")))
return model.User{}, false
}
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
if uniqueErr != nil {
response.AbortInternal(c, uniqueErr.Error())
return model.User{}, false
}
userInfo.Username = username
var user model.User
if err := repository.CreateUserFromOAuth(ctx, &user, userInfo); err != nil {
response.AbortInternal(c, err.Error())
return model.User{}, false
}
if err := repository.BindExternalAccount(ctx, &model.ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
ExternalUsername: userInfo.Username,
Email: userInfo.Email,
}); err != nil {
response.AbortBadRequest(c, err.Error())
return model.User{}, 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, _ := GetFromContext[*model.User](c, UserObjKey)
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 := repository.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, shared.UnAuthorized)
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 := repository.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
+185
View File
@@ -0,0 +1,185 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"errors"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared"
"github.com/Rain-kl/Wavelet/internal/shared/response"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
// GetFromContext 从 Gin 请求上下文获取指定类型的值。
func GetFromContext[T any](c *gin.Context, key string) (T, bool) {
value, exists := c.Get(key)
if !exists {
var zero T
return zero, false
}
typed, ok := value.(T)
return typed, ok
}
// SetToContext 设置值到 Gin 请求上下文。
func SetToContext[T any](c *gin.Context, key string, value T) {
c.Set(key, value)
}
func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.AccessToken, error) {
tokenHash := model.HashToken(tokenStr)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, nil, err
}
tokenRecord = &dbToken
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || !user.IsActive {
dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, nil, err
}
user = &dbUser
SetCachedUser(ctx, tokenRecord.UserID, user)
}
return user, tokenRecord, nil
}
// GetUserFromRequest 校验 Access Token 或 Session 并返回用户对象,如果未登录或用户失效则返回 error
func GetUserFromRequest(c *gin.Context) (*model.User, 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")
}
SetToContext(c, TokenAuthKey, true)
SetToContext(c, TokenAdminKey, 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.IsActive {
dbUser, loadErr := repository.GetActiveUserByID(ctx, userID)
if loadErr != nil {
return nil, loadErr
}
user = &dbUser
SetCachedUser(ctx, userID, user)
}
// 密码哈希校验:当用户存在本地密码时,要求 Session 中的密码哈希必须与当前数据库中一致
if user.Password != "" {
session := sessions.Default(c)
sessionHash, _ := session.Get(PasswordHashKey).(string)
if sessionHash != user.Password {
return nil, errors.New("session expired due to password change")
}
}
SetToContext(c, TokenAuthKey, false)
SetToContext(c, TokenAdminKey, 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) {
ctx, span := otel_trace.Start(c.Request.Context(), "LoginRequired")
defer span.End()
user, err := GetUserFromRequest(c)
if err != nil {
response.AbortUnauthorized(c, shared.UnAuthorized)
return
}
LogForAudit(ctx, user, c)
SetToContext(c, UserObjKey, user)
c.Next()
}
}
// AdminRequired 校验管理员权限(支持 Session 和 Token 鉴权)
func AdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
ctx, span := otel_trace.Start(c.Request.Context(), "AdminRequired")
defer span.End()
user, err := GetUserFromRequest(c)
if err != nil {
response.AbortUnauthorized(c, shared.UnAuthorized)
return
}
isTokenAuth, _ := GetFromContext[bool](c, TokenAuthKey)
isTokenAdmin, _ := GetFromContext[bool](c, TokenAdminKey)
// 如果是通过 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(ctx, user, c)
SetToContext(c, UserObjKey, 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, _ := GetFromContext[bool](c, TokenAuthKey); tokenAuth {
response.AbortForbidden(c, ErrTokenAuthNotAllowed)
return
}
c.Next()
}
}
+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/internal/model"
)
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 将 User 模型转换为 BasicUserInfo
func BuildBasicUserInfo(user *model.User, 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
}
+12 -10
View File
@@ -1,3 +1,6 @@
// 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
@@ -7,7 +10,6 @@ import (
"github.com/Rain-kl/Wavelet/core"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/core/extpoints"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
)
//go:embed migrations/*.sql
@@ -81,16 +83,16 @@ func (p *Plugin) Apply(ctx *core.Context) error {
// 3. Register HTTP Routes
oauthGroup := ctx.Router().Group("/api/v1/oauth")
{
oauthGroup.GET("/sources", oauth.GetLoginSources)
oauthGroup.GET("/login", oauth.GetLoginURL)
oauthGroup.GET("/:source/authorize", oauth.Authorize)
oauthGroup.GET("/logout", oauth.Logout)
oauthGroup.POST("/callback", oauth.Callback)
oauthGroup.GET("/user-info", oauth.LoginRequired(), oauth.UserInfo)
oauthGroup.GET("/external-accounts", oauth.LoginRequired(), oauth.ListExternalAccounts)
oauthGroup.POST("/external-accounts/:id/delete", oauth.LoginRequired(), oauth.DeleteExternalAccount)
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", oauth.LoginRequired(), oauth.UserInfo)
ctx.Router().GET("/api/v1/user-info", LoginRequired(), UserInfo)
// 4. Register Settings Schemas
ctx.Settings().Register(extpoints.SettingSchema{
+81
View File
@@ -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)
}
+13 -12
View File
@@ -1,4 +1,6 @@
// Package auth provides authentication, OAuth, session management, and access token domain services.
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
@@ -7,8 +9,6 @@ import (
"sync"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/gin-gonic/gin"
@@ -44,21 +44,21 @@ func newAuthService() contracts.AuthService {
}
func (s *authServiceImpl) RequireAuthMiddleware() any {
return oauth.LoginRequired()
return LoginRequired()
}
func (s *authServiceImpl) RequireAdminMiddleware() any {
return admin.LoginAdminRequired()
return AdminRequired()
}
func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if ginCtx, ok := ctx.(*gin.Context); ok {
if u, ok := oauth.GetFromContext[*model.User](ginCtx, oauth.UserObjKey); ok && u != nil {
if u, ok := GetFromContext[*model.User](ginCtx, UserObjKey); ok && u != nil {
return toUserDTO(u), nil
}
}
if v := ctx.Value(oauth.UserObjKey); v != nil {
if v := ctx.Value(UserObjKey); v != nil {
if u, ok := v.(*model.User); ok && u != nil {
return toUserDTO(u), nil
}
@@ -76,25 +76,27 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
}
tokenHash := model.HashToken(token)
tokenRecord, err := oauth.GetCachedToken(ctx, tokenHash)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, err
}
tokenRecord = &dbToken
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := oauth.GetCachedUser(ctx, tokenRecord.UserID)
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || !user.IsActive {
dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, err
}
user = &dbUser
SetCachedUser(ctx, tokenRecord.UserID, user)
}
if user.Username == "system" {
if user.Username == SystemUsername {
return nil, errors.New("auth: system user token not allowed")
}
@@ -102,12 +104,11 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
}
func (s *authServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
// Session creation helper
return "", nil
}
func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error {
oauth.InvalidateCachedUser(ctx, userID)
InvalidateCachedUser(ctx, userID)
return nil
}
+140
View File
@@ -0,0 +1,140 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"crypto/sha256"
"encoding/hex"
"net/http"
"strings"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"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) uint64 {
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 *model.User, extras ...map[string]any) error {
session := sessions.Default(c)
session.Clear()
rotateSessionID(session)
session.Set(UserIDKey, user.ID)
session.Set(UserNameKey, user.Username)
session.Set(PasswordHashKey, user.Password)
if len(extras) > 0 {
for key, value := range extras[0] {
session.Set(key, value)
}
}
// 根据系统配置动态设置 Session 过期时间
maxAge := config.Config.App.SessionAge
isSessionCookie := false
ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
if 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
}
@@ -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/internal/shared/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(channelDefinitions()))
}
// 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,346 @@
// 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/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"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"`
Name string `json:"name"`
Fields []Field `json:"fields"`
}
// CreateChannelRequest is the admin create body.
type CreateChannelRequest struct {
Name string `json:"name"`
Type string `json:"type"`
Enabled *bool `json:"enabled"`
BotToken string `json:"bot_token"`
AppID string `json:"app_id"`
AppSecret string `json:"app_secret"`
BaseURL string `json:"base_url"`
PortalHost string `json:"portal_host"`
Sandbox string `json:"sandbox"`
}
// UpdateChannelRequest is the admin patch body.
type UpdateChannelRequest struct {
Name *string `json:"name"`
Enabled *bool `json:"enabled"`
BotToken string `json:"bot_token"`
AppID string `json:"app_id"`
AppSecret string `json:"app_secret"`
BaseURL *string `json:"base_url"`
PortalHost *string `json:"portal_host"`
Sandbox *string `json:"sandbox"`
}
// ChannelDTO is a list/detail view with secrets masked.
type ChannelDTO struct {
ID uint64 `json:"id,string"`
Name string `json:"name"`
Type string `json:"type"`
OwnerScope string `json:"owner_scope"`
Enabled bool `json:"enabled"`
BotToken string `json:"bot_token,omitempty"`
AppID string `json:"app_id,omitempty"`
AppSecret string `json:"app_secret,omitempty"`
BaseURL string `json:"base_url,omitempty"`
PortalHost string `json:"portal_host,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func channelDefinitions() []Definition {
return []Definition{
{
Type: model.MessageChannelTypeTelegram,
Name: "Telegram",
Fields: []Field{
{Key: "bot_token", Type: "password", Required: true},
{Key: "base_url", Type: "text"},
},
},
{
Type: model.MessageChannelTypeQQ,
Name: "QQ",
Fields: []Field{
{Key: "app_id", Required: true},
{Key: "app_secret", Type: "password", Required: true},
{Key: "portal_host", Type: "text"},
},
},
}
}
func createChannel(ctx context.Context, req CreateChannelRequest) (ChannelDTO, error) {
name := strings.TrimSpace(req.Name)
if name == "" {
return ChannelDTO{}, errors.New(errNameRequired)
}
typ := strings.TrimSpace(req.Type)
creds, extra, err := credentialsFromCreate(req)
if err != nil {
return ChannelDTO{}, err
}
cipher, err := EncryptCredentials(creds)
if err != nil {
return ChannelDTO{}, err
}
enabled := true
if req.Enabled != nil {
enabled = *req.Enabled
}
row := &model.MessageChannel{
Name: name,
Type: typ,
OwnerScope: model.MessageOwnerScopeSystem,
Enabled: enabled,
Credentials: cipher,
Extra: EncodeExtra(extra),
}
if err := repository.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 := repository.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 {
creds = map[string]string{}
}
extra := ParseExtra(row.Extra)
if req.Name != nil {
name := strings.TrimSpace(*req.Name)
if name == "" {
return ChannelDTO{}, errors.New(errNameRequired)
}
row.Name = name
}
if req.Enabled != nil {
row.Enabled = *req.Enabled
}
if token := strings.TrimSpace(req.BotToken); token != "" {
creds["bot_token"] = token
}
if appID := strings.TrimSpace(req.AppID); appID != "" {
creds["app_id"] = appID
}
if secret := strings.TrimSpace(req.AppSecret); secret != "" {
creds["app_secret"] = secret
}
if req.BaseURL != nil {
extra["base_url"] = strings.TrimSpace(*req.BaseURL)
}
if req.PortalHost != nil {
extra["portal_host"] = strings.TrimSpace(*req.PortalHost)
}
if req.Sandbox != nil {
extra["sandbox"] = strings.TrimSpace(*req.Sandbox)
}
if err := validateCredentials(row.Type, creds); err != nil {
return ChannelDTO{}, err
}
cipher, err := EncryptCredentials(creds)
if err != nil {
return ChannelDTO{}, err
}
row.Credentials = cipher
row.Extra = EncodeExtra(extra)
if err := repository.UpdateMessageChannel(ctx, row); err != nil {
return ChannelDTO{}, err
}
return toDTO(row, creds, extra), nil
}
func listChannels(ctx context.Context) ([]ChannelDTO, error) {
rows, err := repository.ListMessageChannels(ctx)
if err != nil {
return nil, err
}
out := make([]ChannelDTO, 0, len(rows))
for i := range rows {
creds, err := DecryptCredentials(rows[i].Credentials)
if err != nil {
creds = map[string]string{}
}
out = append(out, toDTO(&rows[i], creds, ParseExtra(rows[i].Extra)))
}
return out, nil
}
func deleteChannel(ctx context.Context, id uint64) error {
if _, err := repository.GetMessageChannel(ctx, id); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errChannelNotFound)
}
return err
}
return repository.DeleteMessageChannel(ctx, id)
}
func probeChannel(ctx context.Context, id uint64) error {
row, err := repository.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
}
extra := ParseExtra(row.Extra)
if err := probeCredentials(ctx, row.Type, creds, extra); err != nil {
return fmt.Errorf("%s: %w", errChannelProbeFailed, err)
}
return nil
}
func credentialsFromCreate(req CreateChannelRequest) (map[string]string, map[string]string, error) {
typ := strings.TrimSpace(req.Type)
creds := map[string]string{}
extra := map[string]string{}
switch typ {
case model.MessageChannelTypeTelegram:
creds["bot_token"] = strings.TrimSpace(req.BotToken)
if base := strings.TrimSpace(req.BaseURL); base != "" {
extra["base_url"] = base
}
case model.MessageChannelTypeQQ:
creds["app_id"] = strings.TrimSpace(req.AppID)
creds["app_secret"] = strings.TrimSpace(req.AppSecret)
if host := strings.TrimSpace(req.PortalHost); host != "" {
extra["portal_host"] = host
} else {
extra["portal_host"] = "q.qq.com"
}
if sandbox := strings.TrimSpace(req.Sandbox); sandbox != "" {
extra["sandbox"] = sandbox
}
default:
return nil, nil, errors.New(errTypeInvalid)
}
if err := validateCredentials(typ, creds); err != nil {
return nil, nil, err
}
return creds, extra, nil
}
func validateCredentials(typ string, creds map[string]string) error {
switch typ {
case model.MessageChannelTypeTelegram:
if strings.TrimSpace(creds["bot_token"]) == "" {
return errors.New(errTelegramTokenRequired)
}
case model.MessageChannelTypeQQ:
if strings.TrimSpace(creds["app_id"]) == "" || strings.TrimSpace(creds["app_secret"]) == "" {
return errors.New(errQQCredentialsRequired)
}
default:
return errors.New(errTypeInvalid)
}
return nil
}
func toDTO(row *model.MessageChannel, creds, extra map[string]string) ChannelDTO {
dto := ChannelDTO{
ID: row.ID,
Name: row.Name,
Type: row.Type,
OwnerScope: row.OwnerScope,
Enabled: row.Enabled,
CreatedAt: row.CreatedAt,
UpdatedAt: row.UpdatedAt,
}
if strings.TrimSpace(creds["bot_token"]) != "" {
dto.BotToken = maskedSecret
}
if id := strings.TrimSpace(creds["app_id"]); id != "" {
dto.AppID = id
}
if strings.TrimSpace(creds["app_secret"]) != "" {
dto.AppSecret = maskedSecret
}
dto.BaseURL = extra["base_url"]
dto.PortalHost = extra["portal_host"]
return dto
}
func probeCredentials(ctx context.Context, typ string, creds, extra map[string]string) error {
switch typ {
case model.MessageChannelTypeTelegram:
base := strings.TrimSpace(extra["base_url"])
if base == "" {
base = defaultTelegramAPI
}
url := strings.TrimRight(base, "/") + "/bot" + creds["bot_token"] + "/getMe"
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return err
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
const probeBodyLimit = 4096
body, _ := io.ReadAll(io.LimitReader(resp.Body, probeBodyLimit))
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("telegram getMe status %d", resp.StatusCode)
}
var parsed struct {
OK bool `json:"ok"`
}
if err := json.Unmarshal(body, &parsed); err != nil {
return err
}
if !parsed.OK {
return errors.New("telegram getMe returned ok=false")
}
return nil
case model.MessageChannelTypeQQ:
src := token.NewQQBotTokenSource(&token.QQBotCredentials{
AppID: creds["app_id"],
AppSecret: creds["app_secret"],
})
_, err := src.Token()
return err
default:
return errors.New(errTypeInvalid)
}
}
@@ -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/internal/listener"
)
// 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: "当管理员成功登录系统时触发此通知",
}
func handleAdminLogin(ctx context.Context, event listener.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)
listener.OnAdminLoggedIn(handleAdminLogin)
}
+26
View File
@@ -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 = "********"
)
+147
View File
@@ -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/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/gin-gonic/gin"
)
func currentUser(c *gin.Context) (*model.User, bool) {
return oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
}
// 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) {
mg := r.Group("/message-gateway", oauth.LoginRequired())
{
mg.GET("/channels", ListChannels)
mg.GET("/bindings", ListBindings)
mg.POST("/bindings", BindBinding)
mg.DELETE("/bindings/:id", UnbindBinding)
}
}
+158
View File
@@ -0,0 +1,158 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"errors"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
pkgmg "github.com/Rain-kl/Wavelet/pkg/message_gateway"
"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 := pkgmg.NormalizeCode(req.Code)
if code == "" {
return BindingDTO{}, errCodeInvalid
}
pairing, err := repository.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 := repository.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 := repository.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
}
_ = repository.DeletePairingCode(ctx, pairing.Code)
return toBindingDTO(existing, ch), nil
}
row := &model.MessageBinding{
UserID: userID,
ChannelID: channelID,
PlatformUserID: pairing.PlatformUserID,
}
if err := repository.CreateMessageBinding(ctx, row); err != nil {
return BindingDTO{}, err
}
if err := repository.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 := repository.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 := repository.ListBindingsByUser(ctx, userID)
if err != nil {
return nil, err
}
out := make([]BindingDTO, 0, len(rows))
for i := range rows {
ch, err := repository.GetMessageChannel(ctx, rows[i].ChannelID)
if err != nil {
out = append(out, toBindingDTO(&rows[i], nil))
continue
}
out = append(out, toBindingDTO(&rows[i], ch))
}
return out, nil
}
func unbindChannel(ctx context.Context, userID, bindingID uint64) error {
row, err := repository.GetMessageBinding(ctx, bindingID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errBindingNotFound
}
return err
}
if row.UserID != userID {
return errBindingForbidden
}
return repository.DeleteMessageBinding(ctx, bindingID)
}
func toBindingDTO(row *model.MessageBinding, ch *model.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
}
@@ -34,10 +34,59 @@ CREATE TABLE IF NOT EXISTS w_message_pairing_codes (
);
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;
+108 -15
View File
@@ -1,3 +1,6 @@
// 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
@@ -7,7 +10,7 @@ import (
"github.com/Rain-kl/Wavelet/core"
"github.com/Rain-kl/Wavelet/core/extpoints"
appgw "github.com/Rain-kl/Wavelet/internal/apps/message_gateway"
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/hibiken/asynq"
)
@@ -18,8 +21,18 @@ 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{}
type Plugin struct {
autoStartRunner bool
cancelRunner context.CancelFunc
}
// New creates a new message_gateway domain plugin.
func New(opts ...Option) *Plugin {
@@ -61,36 +74,100 @@ func (p *Plugin) Apply(ctx *core.Context) error {
// 1. Register migrations
ctx.Migrations().Register("message_gateway", mgMigrations)
// 2. Register HTTP Routes
// 2. Register User HTTP Routes
mgGroup := ctx.Router().Group("/api/v1/message-gateway", oauth.LoginRequired())
{
mgGroup.GET("/channels", appgw.ListChannels)
mgGroup.GET("/bindings", appgw.ListBindings)
mgGroup.POST("/bindings", appgw.BindBinding)
mgGroup.DELETE("/bindings/:id", appgw.UnbindBinding)
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", oauth.LoginRequired(), admin.LoginAdminRequired())
{
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", oauth.LoginRequired(), admin.LoginAdminRequired())
{
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{}
// 3. Register Asynq background tasks
ctx.Task().Register("message_gateway:push_notification", func(_ context.Context, _ *asynq.Task) error {
return nil
// 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
})
// 4. Register Cron Schedules
// 6. Register Cron Schedules
ctx.Schedule().RegisterCron("*/10 * * * *", "message_gateway:cleanup_pairing_codes", map[string]any{"action": "cleanup"})
// 5. Register EventBus listeners for decoupled push triggers
ctx.Events().On("notification:push", func(_ context.Context, _ PushNotificationEvent) error {
// Event triggered push handling
// 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
})
// 6. Register Settings Schemas
// 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,
@@ -106,5 +183,21 @@ func (p *Plugin) Apply(ctx *core.Context) error {
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,398 @@
// 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/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response"
pkgpush "github.com/Rain-kl/Wavelet/pkg/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 := model.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/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"github.com/Rain-kl/Wavelet/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 := repository.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 *model.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 *model.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 *model.PushEvent, msg NotificationMessage, flatBody map[string]any) {
for _, channelName := range event.Channels {
customChannel, err := repository.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 *model.PushEvent, channel *model.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 *model.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"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"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(), repository.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,515 @@
// 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/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"gorm.io/gorm"
)
type smtpConfig struct {
Host string
Port string
Username string
Password string
}
func loadSMTPConfig(ctx context.Context) smtpConfig {
host, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost)
port, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort)
user, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername)
pass, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword)
return smtpConfig{
Host: host.Value,
Port: port.Value,
Username: user.Value,
Password: pass.Value,
}
}
func syncBuiltInEvents(ctx context.Context) error {
for _, meta := range GetBuiltInEvents() {
_, err := repository.GetPushEventByKey(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 := model.PushEvent{
EventKey: meta.Key,
Name: meta.Name,
Channels: []string{},
Targets: []string{},
Template: defaultTemplateStr,
Enabled: false,
}
if err := repository.CreatePushEvent(ctx, &event); err != nil {
return err
}
} else if err != nil {
return err
}
}
return nil
}
func listPushEvents(ctx context.Context) ([]model.PushEvent, error) {
return repository.ListPushEvents(ctx)
}
func createPushEvent(ctx context.Context, req CreatePushEventRequest) (model.PushEvent, error) {
eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req)
if err != nil {
return model.PushEvent{}, err
}
count, err := repository.CountPushEventsByKey(ctx, eventKey)
if err != nil {
return model.PushEvent{}, err
}
if count > 0 {
return model.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 model.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 := model.PushEvent{
EventKey: eventKey,
Name: eventName,
TaskType: req.TaskType,
Channels: channels,
Targets: targets,
Template: templateStr,
Enabled: req.Enabled,
}
if err := event.Validate(); err != nil {
return model.PushEvent{}, err
}
if err := repository.CreatePushEvent(ctx, &event); err != nil {
return model.PushEvent{}, err
}
return event, nil
}
func deletePushEvent(ctx context.Context, id uint64) error {
event, err := repository.GetPushEventByID(ctx, id)
if err != nil {
return err
}
return repository.DeletePushEvent(ctx, &event)
}
func updatePushEvent(ctx context.Context, id uint64, req UpdatePushEventRequest) error {
event, err := repository.GetPushEventByID(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 repository.SavePushEvent(ctx, &event)
}
func togglePushEvent(ctx context.Context, id uint64) (bool, error) {
event, err := repository.GetPushEventByID(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 := repository.UpdatePushEventEnabled(ctx, &event, enabled); err != nil {
return false, err
}
return enabled, nil
}
func listPushHistories(ctx context.Context, filter repository.PushHistoryListFilter) (int64, []model.PushHistory, error) {
return repository.ListPushHistories(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) ([]model.PushChannel, error) {
return repository.ListPushChannels(ctx)
}
func createPushChannel(ctx context.Context, req CreatePushChannelRequest) (model.PushChannel, error) {
count, err := repository.CountPushChannelsByName(ctx, req.Name)
if err != nil {
return model.PushChannel{}, err
}
if count > 0 {
return model.PushChannel{}, errors.New("channel name already exists")
}
channel := model.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 model.PushChannel{}, err
}
if err := repository.CreatePushChannel(ctx, &channel); err != nil {
return model.PushChannel{}, err
}
return channel, nil
}
func updatePushChannel(ctx context.Context, id uint64, req UpdatePushChannelRequest) (model.PushChannel, error) {
channel, err := repository.GetPushChannelByID(ctx, id)
if err != nil {
return model.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 model.PushChannel{}, err
}
if err := repository.SavePushChannel(ctx, &channel); err != nil {
return model.PushChannel{}, err
}
return channel, nil
}
func deletePushChannel(ctx context.Context, id uint64) error {
channel, err := repository.GetPushChannelByID(ctx, id)
if err != nil {
return err
}
return repository.DeletePushChannel(ctx, &channel)
}
func loadChannelForTest(ctx context.Context, req TestPushChannelRequest) (string, string, string, string, error) {
if req.Name != "" {
channel, err := repository.GetPushChannelByName(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) ([]model.PushEvent, error) {
return repository.ListActivePushEventsByTaskType(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 {
if user, err := repository.GetUserByID(ctx, userID); err == nil {
return &user
}
}
if username := extractUsername(data); username != "" {
if user, err := repository.GetUserByUsername(ctx, username); 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 := model.PushHistory{
EventKey: req.EventKey,
Channel: req.Config.Channel,
Target: target,
Title: title,
Content: content,
Level: level,
Status: status,
ErrorMsg: errMsg,
}
return repository.CreatePushHistory(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) (model.User, bool) {
found := false
var user model.User
if id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
if u, err := repository.GetUserByID(ctx, id); err == nil {
user = u
found = true
}
}
if !found {
if u, err := repository.GetUserByUsername(ctx, resolved); err == nil {
user = u
found = true
}
}
return user, found
}
func resolveSystemTarget(ctx context.Context, resolved string, channel string) (string, bool) {
if resolved != "系统" && resolved != "system" && resolved != "0" {
return "", false
}
adminUser, err := repository.GetFirstAdminUser(ctx)
if 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) *model.User {
user := repository.GetSystemUser(ctx)
return &user
}
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 := task.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 = task.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,124 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"encoding/json"
"strconv"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
// RegisterTaskListeners subscribes push notification handlers to task completion events.
func RegisterTaskListeners() {
task.OnTaskCompleted(handleTaskCompleted)
}
func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, result *task.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/internal/infra/task"
"github.com/Rain-kl/Wavelet/pkg/push"
)
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 = task.TaskMeta{
Type: TaskTypeSendNotification,
AsynqTask: SendNotificationTask,
Name: "推送通知",
Description: "异步执行系统通知的多渠道派发与推送",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
Params: []task.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) (*task.TaskResult, error) {
var req SendPayload
if err := json.Unmarshal(payload, &req); err != nil {
task.AppendLog(ctx, "解析推送参数失败: %v", err)
return nil, fmt.Errorf("parse payload failed: %w", err)
}
task.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)
task.AppendLog(ctx, "推送失败: %v", errWrap)
if task.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 {
task.AppendLog(ctx, "消息推送失败 (标题: %s): %v", title, err)
if upstreamResp != "" {
task.AppendLog(ctx, "上游返回: %s", upstreamResp)
}
if task.IsFinalAttempt(ctx) {
h.recordHistory(ctx, req, "failed", err.Error())
}
return nil, fmt.Errorf("pusher.Send failed: %w", err)
}
task.AppendLog(ctx, "消息推送成功 (标题: %s, 内容摘要: %s)", title, content)
if upstreamResp != "" {
task.AppendLog(ctx, "上游返回: %s", upstreamResp)
}
h.recordHistory(ctx, req, "success", "")
return &task.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 {
task.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr)
}
}
+53
View File
@@ -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/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
}
+78
View File
@@ -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/internal/infra/config"
"github.com/Rain-kl/Wavelet/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)
}
+144
View File
@@ -0,0 +1,144 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package risk_control
import (
"context"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/batchwriter"
"github.com/Rain-kl/Wavelet/internal/model/analytics"
"github.com/Rain-kl/Wavelet/internal/platform/lifecycle"
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
var (
logWriterMu sync.RWMutex
logWriter *batchwriter.Writer[*analytics.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[*analytics.UserAccessLog](cfg, func(ctx context.Context, items []*analytics.UserAccessLog) error {
rows := make([]analytics.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[*analytics.UserAccessLog](func(item *analytics.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[*analytics.UserAccessLog](func(ctx context.Context, items []*analytics.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
lifecycle.OnShutdown("risk_control_log_writer", StopLogWriter)
}
// 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 *analytics.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[*analytics.UserAccessLog]) func() {
logWriterMu.Lock()
previous := logWriter
logWriter = writer
logWriterMu.Unlock()
return func() {
logWriterMu.Lock()
logWriter = previous
logWriterMu.Unlock()
}
}
func currentLogWriter() *batchwriter.Writer[*analytics.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:
}
}
}
+88
View File
@@ -0,0 +1,88 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package risk_control provides the access control, IP rate limiting, and telemetry risk analysis domain plugin for Cordis.
package risk_control
import (
"encoding/json"
"net/http"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/model/analytics"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/gin-gonic/gin"
)
// RiskControlMiddleware 全局日志采集中间件
func RiskControlMiddleware() gin.HandlerFunc {
return func(c *gin.Context) {
// 如果未启用 ClickHouse,直接放行
if config.Config == nil || !config.Config.ClickHouse.Enabled {
c.Next()
return
}
// 1. 限流背压检测(检测本地缓冲队列是否已满)
if IsBufferFull() {
response.AbortTooManyRequests(c, "系统繁忙,请稍后再试")
return
}
start := time.Now()
// 2. 执行后续请求(穿过业务处理和认证中间件)
c.Next()
// 3. 后置身份检查:仅记录通过认证的请求
userObj, exists := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
if !exists || userObj == nil {
return
}
// 4. 计算耗时并异步推送到缓冲队列
latency := time.Since(start).Milliseconds()
var headersStr string
if c.Request.Header != nil {
// 克隆 Header,避免污染原 HTTP 请求的 Header 对象
clonedHeaders := make(http.Header)
for k, v := range c.Request.Header {
clonedHeaders[k] = v
}
clonedHeaders.Del("Cookie")
if headersBytes, err := json.Marshal(clonedHeaders); err == nil {
headersStr = string(headersBytes)
}
}
const maxHTTPStatus = 999
status := c.Writer.Status()
if status < 0 {
status = 0
} else if status > maxHTTPStatus {
status = maxHTTPStatus
}
logItem := &analytics.UserAccessLog{
ID: idgen.NextUint64ID(),
UserID: userObj.ID, // 直接从 Context 获取已登录用户ID,避免数据库查询
Path: c.Request.URL.Path,
Method: c.Request.Method,
IP: c.ClientIP(),
UserAgent: c.Request.UserAgent(),
Headers: headersStr,
Status: int32(status),
Latency: latency,
CreatedAt: time.Now(),
}
// 非阻塞地推入缓存队列
QueueAccessLog(logItem)
}
}
@@ -0,0 +1,198 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package risk_control_test
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"sync"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/batchwriter"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/model/analytics"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
)
func newTestAccessLogWriter(t *testing.T, cfg batchwriter.Config) (*batchwriter.Writer[*analytics.UserAccessLog], func() []*analytics.UserAccessLog) {
t.Helper()
var (
mu sync.Mutex
captured []*analytics.UserAccessLog
)
writer, err := batchwriter.New(cfg, func(_ context.Context, items []*analytics.UserAccessLog) error {
mu.Lock()
captured = append(captured, items...)
mu.Unlock()
return nil
})
if err != nil {
t.Fatalf("batchwriter.New() error = %v", err)
}
writer.Start(context.Background())
restore := risk_control.SetLogWriterForTest(writer)
t.Cleanup(func() {
restore()
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_ = writer.Stop(stopCtx)
})
return writer, func() []*analytics.UserAccessLog {
mu.Lock()
defer mu.Unlock()
return append([]*analytics.UserAccessLog(nil), captured...)
}
}
func drainAccessLogWriter(t *testing.T, writer *batchwriter.Writer[*analytics.UserAccessLog]) {
t.Helper()
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := writer.Stop(stopCtx); err != nil {
t.Fatalf("writer.Stop() error = %v", err)
}
}
func TestRiskControlMiddleware(t *testing.T) {
gin.SetMode(gin.TestMode)
t.Run("ClickHouse disabled", func(t *testing.T) {
config.Config.ClickHouse.Enabled = false
defer func() { config.Config.ClickHouse.Enabled = false }()
r := testhelper.NewTestGinEngine(risk_control.RiskControlMiddleware())
r.GET("/test", func(c *gin.Context) {
c.String(http.StatusOK, "ok")
})
w := httptest.NewRecorder()
req, _ := http.NewRequest(http.MethodGet, "/test", nil)
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
assert.Equal(t, "ok", w.Body.String())
})
t.Run("ClickHouse enabled - Normal Authenticated Request", func(t *testing.T) {
config.Config.ClickHouse.Enabled = true
defer func() { config.Config.ClickHouse.Enabled = false }()
cfg := batchwriter.DefaultConfig()
cfg.MaxBatchSize = 100
cfg.FlushInterval = time.Hour
writer, getCaptured := newTestAccessLogWriter(t, cfg)
r := gin.New()
r.Use(func(c *gin.Context) {
user := &model.User{ID: 12345}
oauth.SetToContext(c, oauth.UserObjKey, user)
c.Next()
})
r.Use(risk_control.RiskControlMiddleware())
r.GET("/test", func(c *gin.Context) {
c.String(http.StatusOK, "ok")
})
w := httptest.NewRecorder()
req, _ := http.NewRequest(http.MethodGet, "/test", nil)
req.Header.Set("X-Test-Header", "hello")
req.Header.Set("Cookie", "session_id=abcdef123456")
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
assert.Equal(t, "ok", w.Body.String())
drainAccessLogWriter(t, writer)
captured := getCaptured()
if len(captured) != 1 {
t.Fatalf("captured access logs = %d, want 1", len(captured))
}
logItem := captured[0]
assert.Equal(t, uint64(12345), logItem.UserID)
assert.Equal(t, "/test", logItem.Path)
assert.Equal(t, http.MethodGet, logItem.Method)
assert.Equal(t, int32(http.StatusOK), logItem.Status)
assert.NotEmpty(t, logItem.Headers)
assert.Contains(t, logItem.Headers, "X-Test-Header")
assert.NotContains(t, logItem.Headers, "Cookie")
})
t.Run("ClickHouse enabled - Unauthenticated Request", func(t *testing.T) {
config.Config.ClickHouse.Enabled = true
defer func() { config.Config.ClickHouse.Enabled = false }()
cfg := batchwriter.DefaultConfig()
cfg.MaxBatchSize = 100
cfg.FlushInterval = time.Hour
writer, getCaptured := newTestAccessLogWriter(t, cfg)
r := testhelper.NewTestGinEngine(risk_control.RiskControlMiddleware())
r.GET("/test", func(c *gin.Context) {
c.String(http.StatusOK, "ok")
})
w := httptest.NewRecorder()
req, _ := http.NewRequest(http.MethodGet, "/test", nil)
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
assert.Equal(t, "ok", w.Body.String())
drainAccessLogWriter(t, writer)
if len(getCaptured()) != 0 {
t.Fatal("expected no log item for unauthenticated request")
}
})
t.Run("ClickHouse enabled - Buffer Full Rate Limiting", func(t *testing.T) {
config.Config.ClickHouse.Enabled = true
defer func() { config.Config.ClickHouse.Enabled = false }()
cfg := batchwriter.DefaultConfig()
cfg.QueueSize = 2
cfg.MaxBatchSize = 100
cfg.FlushInterval = time.Hour
writer, _ := newTestAccessLogWriter(t, cfg)
for range cfg.QueueSize {
writer.TryEnqueue(&analytics.UserAccessLog{})
}
if !risk_control.IsBufferFull() {
t.Fatal("IsBufferFull() = false, want true")
}
r := testhelper.NewTestGinEngine(risk_control.RiskControlMiddleware())
r.GET("/test", func(c *gin.Context) {
c.String(http.StatusOK, "ok")
})
w := httptest.NewRecorder()
req, _ := http.NewRequest(http.MethodGet, "/test", nil)
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusTooManyRequests, w.Code)
var resp map[string]interface{}
err := json.Unmarshal(w.Body.Bytes(), &resp)
assert.NoError(t, err)
assert.Contains(t, resp["error_msg"], "系统繁忙")
})
}
+6 -4
View File
@@ -1,3 +1,6 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package risk_control provides the access control, IP rate limiting, and telemetry risk analysis domain plugin for Cordis.
package risk_control
@@ -6,7 +9,6 @@ import (
"github.com/Rain-kl/Wavelet/core"
"github.com/Rain-kl/Wavelet/core/extpoints"
"github.com/Rain-kl/Wavelet/internal/apps/risk_control"
"github.com/gin-gonic/gin"
)
@@ -54,12 +56,12 @@ func (p *Plugin) Manifest() core.Manifest {
// Apply registers risk control middlewares, settings, and cleanup hooks into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
// 1. Initialize LogWriter if needed
risk_control.InitLogWriter(ctx.GoContext())
InitLogWriter(ctx.GoContext())
// 2. Register router middleware
mw := p.middleware
if mw == nil {
mw = risk_control.RiskControlMiddleware()
mw = RiskControlMiddleware()
}
ctx.Router().Use(mw)
@@ -81,7 +83,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
// 4. Register lifecycle disposal cleanup
ctx.OnDispose(func() error {
return risk_control.StopLogWriter(context.Background())
return StopLogWriter(context.Background())
})
return nil