mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-11 17:56:37 +08:00
refactor(plugins): complete physical encapsulation of auth, admin, message_gateway, and risk_control domain plugins
This commit is contained in:
@@ -9,13 +9,13 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"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/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/infra/task"
|
"github.com/Rain-kl/Wavelet/internal/infra/task"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
|
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
|
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
|
|||||||
@@ -0,0 +1,82 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package admin
|
||||||
|
|
||||||
|
// 管理后台公共错误常量
|
||||||
|
const (
|
||||||
|
AdminRequired = "未经授权访问"
|
||||||
|
TokenAdminRequired = "该访问令牌没有管理员权限,无法访问管理端点" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||||
|
InvalidAuthSourceID = "认证源 ID 无效"
|
||||||
|
InvalidCursorParam = "无效的 cursor 参数"
|
||||||
|
InvalidTaskExecutionID = "无效的任务执行记录 ID"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 系统配置错误消息常量
|
||||||
|
const (
|
||||||
|
SystemConfigNotFound = "系统配置不存在"
|
||||||
|
ConfigKeyRequired = "配置键不能为空"
|
||||||
|
ConfigValueRequired = "配置值不能为空"
|
||||||
|
ConfigKeyExists = "配置键已存在"
|
||||||
|
protectedConfigKeyMessage = "该配置项由系统任务管理,禁止手动修改"
|
||||||
|
StorageDriverSwitchRequiresMigration = "存在存量文件,请通过存储迁移任务切换存储引擎"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 模板管理相关错误消息常量
|
||||||
|
const (
|
||||||
|
TemplateNotFound = "模板不存在"
|
||||||
|
TemplateKeyRequired = "模板标识符不能为空"
|
||||||
|
TemplateNameRequired = "模板名称不能为空"
|
||||||
|
TemplateContentRequired = "模板内容不能为空"
|
||||||
|
TemplateKeyExists = "模板标识符已存在"
|
||||||
|
SystemTemplateCannotDelete = "系统预置模板不可删除"
|
||||||
|
SystemTemplateCannotModifyKey = "系统预置模板不可修改标识符"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 任务调度相关错误消息常量
|
||||||
|
const (
|
||||||
|
InvalidTaskType = "无效的任务类型"
|
||||||
|
InvalidTimeRange = "无效的时间范围"
|
||||||
|
TaskDispatchFailed = "任务下发失败"
|
||||||
|
UserIDRequired = "用户ID必填"
|
||||||
|
TaskNotFound = "任务执行记录不存在"
|
||||||
|
TaskNotRetryable = "该任务不支持重试"
|
||||||
|
TaskNotFailed = "只有失败的任务才能重试"
|
||||||
|
TaskMaxRetryExceeded = "已达到最大重试次数"
|
||||||
|
TaskRetryFailed = "任务重试失败"
|
||||||
|
InvalidCronExpression = "无效的 Cron 表达式"
|
||||||
|
ScheduleNotFound = "定时任务不存在"
|
||||||
|
ScheduleSaveFailed = "保存定时任务失败"
|
||||||
|
ScheduleDeleteFailed = "删除定时任务失败"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 应用更新相关错误消息常量
|
||||||
|
const (
|
||||||
|
errInvalidRepository = "上游仓库地址无效"
|
||||||
|
errReleaseRequestFailed = "获取上游版本失败"
|
||||||
|
errReleaseResponseInvalid = "上游版本响应无效"
|
||||||
|
errNoCompatibleRelease = "未找到兼容的 Release"
|
||||||
|
errNoCompatibleAsset = "未找到当前系统对应的 Release 资产"
|
||||||
|
errDevelopmentBuild = "开发版本无法执行自动升级"
|
||||||
|
errAlreadyUpToDate = "当前已是最新版本"
|
||||||
|
errUpgradeAlreadyRunning = "已有升级任务正在执行"
|
||||||
|
errAutomaticUpgradeBlocked = "当前平台暂不支持自动替换二进制"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 用户管理(管理员视角)错误消息常量
|
||||||
|
const (
|
||||||
|
userNotFound = "用户不存在"
|
||||||
|
cannotDisable = "不能禁用管理员账号"
|
||||||
|
cannotDelete = "不能删除管理员账号"
|
||||||
|
cannotDeleteSelf = "不能删除当前登录账号"
|
||||||
|
usernameRequired = "用户名不能为空"
|
||||||
|
emailRequired = "邮箱不能为空"
|
||||||
|
//nolint:gosec // error message, not hardcoded credentials
|
||||||
|
passwordTooShort = "密码长度不能少于 8 位"
|
||||||
|
usernameExists = "用户名已存在"
|
||||||
|
emailExists = "邮箱已被使用"
|
||||||
|
cannotRevokeSelfAdmin = "不能取消自身的管理员权限"
|
||||||
|
updateUserFailed = "更新用户状态失败"
|
||||||
|
deleteUserFailed = "删除用户失败"
|
||||||
|
updateUserInfoFailed = "更新用户信息失败"
|
||||||
|
)
|
||||||
@@ -0,0 +1,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)
|
||||||
|
}
|
||||||
@@ -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), ¤tCfg); 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), ¤tCfg); 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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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{}
|
||||||
|
}
|
||||||
@@ -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())
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
@@ -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"
|
||||||
|
}
|
||||||
@@ -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())
|
||||||
|
}
|
||||||
@@ -0,0 +1,12 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package admin
|
||||||
|
|
||||||
|
import "errors"
|
||||||
|
|
||||||
|
func replaceAndRestart(_, _ string) error {
|
||||||
|
return errors.New(errAutomaticUpgradeBlocked)
|
||||||
|
}
|
||||||
@@ -0,0 +1,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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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()
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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"
|
||||||
|
)
|
||||||
@@ -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 = "令牌无管理员权限"
|
||||||
|
)
|
||||||
@@ -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())
|
||||||
|
}
|
||||||
@@ -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()
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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 provides the authentication, OAuth, session management, and access token domain plugin for Cordis.
|
||||||
package auth
|
package auth
|
||||||
|
|
||||||
@@ -7,7 +10,6 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/core"
|
"github.com/Rain-kl/Wavelet/core"
|
||||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||||
"github.com/Rain-kl/Wavelet/core/extpoints"
|
"github.com/Rain-kl/Wavelet/core/extpoints"
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
//go:embed migrations/*.sql
|
//go:embed migrations/*.sql
|
||||||
@@ -81,16 +83,16 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
// 3. Register HTTP Routes
|
// 3. Register HTTP Routes
|
||||||
oauthGroup := ctx.Router().Group("/api/v1/oauth")
|
oauthGroup := ctx.Router().Group("/api/v1/oauth")
|
||||||
{
|
{
|
||||||
oauthGroup.GET("/sources", oauth.GetLoginSources)
|
oauthGroup.GET("/sources", GetLoginSources)
|
||||||
oauthGroup.GET("/login", oauth.GetLoginURL)
|
oauthGroup.GET("/login", GetLoginURL)
|
||||||
oauthGroup.GET("/:source/authorize", oauth.Authorize)
|
oauthGroup.GET("/:source/authorize", Authorize)
|
||||||
oauthGroup.GET("/logout", oauth.Logout)
|
oauthGroup.GET("/logout", Logout)
|
||||||
oauthGroup.POST("/callback", oauth.Callback)
|
oauthGroup.POST("/callback", Callback)
|
||||||
oauthGroup.GET("/user-info", oauth.LoginRequired(), oauth.UserInfo)
|
oauthGroup.GET("/user-info", LoginRequired(), UserInfo)
|
||||||
oauthGroup.GET("/external-accounts", oauth.LoginRequired(), oauth.ListExternalAccounts)
|
oauthGroup.GET("/external-accounts", LoginRequired(), ListExternalAccounts)
|
||||||
oauthGroup.POST("/external-accounts/:id/delete", oauth.LoginRequired(), oauth.DeleteExternalAccount)
|
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
|
// 4. Register Settings Schemas
|
||||||
ctx.Settings().Register(extpoints.SettingSchema{
|
ctx.Settings().Register(extpoints.SettingSchema{
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
|
package auth
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -7,8 +9,6 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
"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/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
@@ -44,21 +44,21 @@ func newAuthService() contracts.AuthService {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *authServiceImpl) RequireAuthMiddleware() any {
|
func (s *authServiceImpl) RequireAuthMiddleware() any {
|
||||||
return oauth.LoginRequired()
|
return LoginRequired()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *authServiceImpl) RequireAdminMiddleware() any {
|
func (s *authServiceImpl) RequireAdminMiddleware() any {
|
||||||
return admin.LoginAdminRequired()
|
return AdminRequired()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
|
func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
|
||||||
if ginCtx, ok := ctx.(*gin.Context); ok {
|
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
|
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 {
|
if u, ok := v.(*model.User); ok && u != nil {
|
||||||
return toUserDTO(u), nil
|
return toUserDTO(u), nil
|
||||||
}
|
}
|
||||||
@@ -76,25 +76,27 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
|
|||||||
}
|
}
|
||||||
|
|
||||||
tokenHash := model.HashToken(token)
|
tokenHash := model.HashToken(token)
|
||||||
tokenRecord, err := oauth.GetCachedToken(ctx, tokenHash)
|
tokenRecord, err := GetCachedToken(ctx, tokenHash)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash)
|
dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
tokenRecord = &dbToken
|
tokenRecord = &dbToken
|
||||||
|
SetCachedToken(ctx, tokenHash, tokenRecord)
|
||||||
}
|
}
|
||||||
|
|
||||||
user, err := oauth.GetCachedUser(ctx, tokenRecord.UserID)
|
user, err := GetCachedUser(ctx, tokenRecord.UserID)
|
||||||
if err != nil || !user.IsActive {
|
if err != nil || !user.IsActive {
|
||||||
dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID)
|
dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
user = &dbUser
|
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")
|
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) {
|
func (s *authServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
|
||||||
// Session creation helper
|
|
||||||
return "", nil
|
return "", nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error {
|
func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error {
|
||||||
oauth.InvalidateCachedUser(ctx, userID)
|
InvalidateCachedUser(ctx, userID)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -0,0 +1,26 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package message_gateway
|
||||||
|
|
||||||
|
import "errors"
|
||||||
|
|
||||||
|
var (
|
||||||
|
errCodeInvalid = errors.New("invalid or expired pairing code")
|
||||||
|
errChannelMismatch = errors.New("pairing code does not match channel")
|
||||||
|
errPlatformAlreadyBound = errors.New("this platform account is already bound")
|
||||||
|
errBindingNotFound = errors.New("binding not found")
|
||||||
|
errBindingForbidden = errors.New("cannot unbind another user's binding")
|
||||||
|
errChannelIDRequired = errors.New("channel_id is required")
|
||||||
|
errChannelDisabled = errors.New("channel is not enabled")
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
errNameRequired = "name is required"
|
||||||
|
errTypeInvalid = "type must be telegram or qq"
|
||||||
|
errTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text
|
||||||
|
errQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text
|
||||||
|
errChannelNotFound = "channel not found"
|
||||||
|
errChannelProbeFailed = "channel probe failed"
|
||||||
|
maskedSecret = "********"
|
||||||
|
)
|
||||||
@@ -0,0 +1,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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
+49
@@ -34,10 +34,59 @@ CREATE TABLE IF NOT EXISTS w_message_pairing_codes (
|
|||||||
);
|
);
|
||||||
CREATE INDEX IF NOT EXISTS idx_w_message_pairing_lookup
|
CREATE INDEX IF NOT EXISTS idx_w_message_pairing_lookup
|
||||||
ON w_message_pairing_codes (channel_id, platform_user_id);
|
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 StatementEnd
|
||||||
|
|
||||||
-- +goose Down
|
-- +goose Down
|
||||||
-- +goose StatementBegin
|
-- +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_pairing_codes;
|
||||||
DROP TABLE IF EXISTS w_message_bindings;
|
DROP TABLE IF EXISTS w_message_bindings;
|
||||||
DROP TABLE IF EXISTS w_message_channels;
|
DROP TABLE IF EXISTS w_message_channels;
|
||||||
|
|||||||
@@ -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 provides the Bot gateway, multi-channel notification dispatching, and asynchronous push worker domain plugin for Cordis.
|
||||||
package message_gateway
|
package message_gateway
|
||||||
|
|
||||||
@@ -7,7 +10,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/core"
|
"github.com/Rain-kl/Wavelet/core"
|
||||||
"github.com/Rain-kl/Wavelet/core/extpoints"
|
"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/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/hibiken/asynq"
|
"github.com/hibiken/asynq"
|
||||||
)
|
)
|
||||||
@@ -18,8 +21,18 @@ var mgMigrations embed.FS
|
|||||||
// Option configures the message_gateway plugin.
|
// Option configures the message_gateway plugin.
|
||||||
type Option func(*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.
|
// 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.
|
// New creates a new message_gateway domain plugin.
|
||||||
func New(opts ...Option) *Plugin {
|
func New(opts ...Option) *Plugin {
|
||||||
@@ -61,36 +74,100 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
// 1. Register migrations
|
// 1. Register migrations
|
||||||
ctx.Migrations().Register("message_gateway", mgMigrations)
|
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 := ctx.Router().Group("/api/v1/message-gateway", oauth.LoginRequired())
|
||||||
{
|
{
|
||||||
mgGroup.GET("/channels", appgw.ListChannels)
|
mgGroup.GET("/channels", ListChannels)
|
||||||
mgGroup.GET("/bindings", appgw.ListBindings)
|
mgGroup.GET("/bindings", ListBindings)
|
||||||
mgGroup.POST("/bindings", appgw.BindBinding)
|
mgGroup.POST("/bindings", BindBinding)
|
||||||
mgGroup.DELETE("/bindings/:id", appgw.UnbindBinding)
|
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
|
const defaultTaskRetry = 3
|
||||||
|
pushHandler := &PushHandler{}
|
||||||
|
|
||||||
// 3. Register Asynq background tasks
|
// 5. Register Asynq background tasks
|
||||||
ctx.Task().Register("message_gateway:push_notification", func(_ context.Context, _ *asynq.Task) error {
|
ctx.Task().Register("message_gateway:push_notification", func(c context.Context, t *asynq.Task) error {
|
||||||
return nil
|
_, 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))
|
}, extpoints.WithTaskRetry(defaultTaskRetry))
|
||||||
|
|
||||||
ctx.Task().Register("message_gateway:dispatch_bot_msg", func(_ context.Context, _ *asynq.Task) error {
|
ctx.Task().Register("message_gateway:dispatch_bot_msg", func(_ context.Context, _ *asynq.Task) error {
|
||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
|
|
||||||
// 4. Register Cron Schedules
|
// 6. Register Cron Schedules
|
||||||
ctx.Schedule().RegisterCron("*/10 * * * *", "message_gateway:cleanup_pairing_codes", map[string]any{"action": "cleanup"})
|
ctx.Schedule().RegisterCron("*/10 * * * *", "message_gateway:cleanup_pairing_codes", map[string]any{"action": "cleanup"})
|
||||||
|
|
||||||
// 5. Register EventBus listeners for decoupled push triggers
|
// 7. Register EventBus listeners for decoupled push triggers
|
||||||
ctx.Events().On("notification:push", func(_ context.Context, _ PushNotificationEvent) error {
|
ctx.Events().On("notification:push", func(c context.Context, e PushNotificationEvent) error {
|
||||||
// Event triggered push handling
|
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
|
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{
|
ctx.Settings().Register(extpoints.SettingSchema{
|
||||||
Key: "message_gateway.pairing_code_expiry_minutes",
|
Key: "message_gateway.pairing_code_expiry_minutes",
|
||||||
Default: 15,
|
Default: 15,
|
||||||
@@ -106,5 +183,21 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
Category: "messaging",
|
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
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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"], "系统繁忙")
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -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 provides the access control, IP rate limiting, and telemetry risk analysis domain plugin for Cordis.
|
||||||
package risk_control
|
package risk_control
|
||||||
|
|
||||||
@@ -6,7 +9,6 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/core"
|
"github.com/Rain-kl/Wavelet/core"
|
||||||
"github.com/Rain-kl/Wavelet/core/extpoints"
|
"github.com/Rain-kl/Wavelet/core/extpoints"
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/risk_control"
|
|
||||||
"github.com/gin-gonic/gin"
|
"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.
|
// Apply registers risk control middlewares, settings, and cleanup hooks into the Context.
|
||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||||
// 1. Initialize LogWriter if needed
|
// 1. Initialize LogWriter if needed
|
||||||
risk_control.InitLogWriter(ctx.GoContext())
|
InitLogWriter(ctx.GoContext())
|
||||||
|
|
||||||
// 2. Register router middleware
|
// 2. Register router middleware
|
||||||
mw := p.middleware
|
mw := p.middleware
|
||||||
if mw == nil {
|
if mw == nil {
|
||||||
mw = risk_control.RiskControlMiddleware()
|
mw = RiskControlMiddleware()
|
||||||
}
|
}
|
||||||
ctx.Router().Use(mw)
|
ctx.Router().Use(mw)
|
||||||
|
|
||||||
@@ -81,7 +83,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
|
|
||||||
// 4. Register lifecycle disposal cleanup
|
// 4. Register lifecycle disposal cleanup
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
return risk_control.StopLogWriter(context.Background())
|
return StopLogWriter(context.Background())
|
||||||
})
|
})
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
Reference in New Issue
Block a user