diff --git a/internal/apps/admin/logs/tasks.go b/internal/apps/admin/logs/tasks.go index e03369dc..ed52b0cd 100644 --- a/internal/apps/admin/logs/tasks.go +++ b/internal/apps/admin/logs/tasks.go @@ -9,13 +9,13 @@ import ( "errors" "fmt" - "github.com/Rain-kl/Wavelet/internal/apps/risk_control" "github.com/Rain-kl/Wavelet/internal/infra/config" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/repository/logstore" "github.com/Rain-kl/Wavelet/pkg/logger" + "github.com/Rain-kl/Wavelet/plugins/domain/risk_control" ) const ( diff --git a/plugins/domain/admin/errs.go b/plugins/domain/admin/errs.go new file mode 100644 index 00000000..78e07657 --- /dev/null +++ b/plugins/domain/admin/errs.go @@ -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 = "更新用户信息失败" +) diff --git a/plugins/domain/admin/handlers_cache.go b/plugins/domain/admin/handlers_cache.go new file mode 100644 index 00000000..dbb49d7d --- /dev/null +++ b/plugins/domain/admin/handlers_cache.go @@ -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) +} diff --git a/plugins/domain/admin/handlers_config.go b/plugins/domain/admin/handlers_config.go new file mode 100644 index 00000000..a7ab4622 --- /dev/null +++ b/plugins/domain/admin/handlers_config.go @@ -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 := `

SMTP Mail Connection Test

+

If you received this message, your SMTP configuration is correct and mail sending is working properly.

+

Sent from Wavelet.

` + + 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 +} diff --git a/plugins/domain/admin/handlers_db.go b/plugins/domain/admin/handlers_db.go new file mode 100644 index 00000000..721100cc --- /dev/null +++ b/plugins/domain/admin/handlers_db.go @@ -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) + } +} diff --git a/plugins/domain/admin/handlers_logs.go b/plugins/domain/admin/handlers_logs.go new file mode 100644 index 00000000..06ea4f82 --- /dev/null +++ b/plugins/domain/admin/handlers_logs.go @@ -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 +} diff --git a/plugins/domain/admin/handlers_status.go b/plugins/domain/admin/handlers_status.go new file mode 100644 index 00000000..3aea8a81 --- /dev/null +++ b/plugins/domain/admin/handlers_status.go @@ -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{} +} diff --git a/plugins/domain/admin/handlers_tasks.go b/plugins/domain/admin/handlers_tasks.go new file mode 100644 index 00000000..b1d9fa7f --- /dev/null +++ b/plugins/domain/admin/handlers_tasks.go @@ -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()) +} diff --git a/plugins/domain/admin/handlers_templates.go b/plugins/domain/admin/handlers_templates.go new file mode 100644 index 00000000..759e6bf9 --- /dev/null +++ b/plugins/domain/admin/handlers_templates.go @@ -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) +} diff --git a/plugins/domain/admin/handlers_updater.go b/plugins/domain/admin/handlers_updater.go new file mode 100644 index 00000000..5febb1fd --- /dev/null +++ b/plugins/domain/admin/handlers_updater.go @@ -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 +} diff --git a/plugins/domain/admin/handlers_user.go b/plugins/domain/admin/handlers_user.go new file mode 100644 index 00000000..55449a07 --- /dev/null +++ b/plugins/domain/admin/handlers_user.go @@ -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 +} diff --git a/plugins/domain/admin/middlewares.go b/plugins/domain/admin/middlewares.go new file mode 100644 index 00000000..c4af729b --- /dev/null +++ b/plugins/domain/admin/middlewares.go @@ -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() + } +} diff --git a/plugins/domain/admin/migrations/20260827000001_create_admin_tables.sql b/plugins/domain/admin/migrations/20260827000001_create_admin_tables.sql new file mode 100644 index 00000000..d03b5868 --- /dev/null +++ b/plugins/domain/admin/migrations/20260827000001_create_admin_tables.sql @@ -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 diff --git a/plugins/domain/admin/migrations/20260827000002_seed_admin_data.sql b/plugins/domain/admin/migrations/20260827000002_seed_admin_data.sql new file mode 100644 index 00000000..d3cca2ea --- /dev/null +++ b/plugins/domain/admin/migrations/20260827000002_seed_admin_data.sql @@ -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 登录验证码', '

Wavelet 登录验证

您的登录验证码为:{{.Code}},5分钟内有效,请勿将验证码泄露给他人。

', '用户密码登录时发送的验证码邮件模板,支持变量:{{.Code}}', TRUE, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP), + (2, 'register_email', '注册验证码邮件', 'email', 'Wavelet 注册验证码', '

Wavelet 注册验证

您的注册验证码为:{{.Code}},5分钟内有效,请勿泄露给他人。

', '用户注册时发送的验证码邮件模板,支持变量:{{.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 diff --git a/plugins/domain/admin/models.go b/plugins/domain/admin/models.go new file mode 100644 index 00000000..718e7647 --- /dev/null +++ b/plugins/domain/admin/models.go @@ -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" +} diff --git a/plugins/domain/admin/restart_unix.go b/plugins/domain/admin/restart_unix.go new file mode 100644 index 00000000..fa272dca --- /dev/null +++ b/plugins/domain/admin/restart_unix.go @@ -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()) +} diff --git a/plugins/domain/admin/restart_windows.go b/plugins/domain/admin/restart_windows.go new file mode 100644 index 00000000..a8daae80 --- /dev/null +++ b/plugins/domain/admin/restart_windows.go @@ -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) +} diff --git a/plugins/domain/auth/audit.go b/plugins/domain/auth/audit.go new file mode 100644 index 00000000..f617cf3a --- /dev/null +++ b/plugins/domain/auth/audit.go @@ -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) + } +} diff --git a/plugins/domain/auth/auth_source_resolver.go b/plugins/domain/auth/auth_source_resolver.go new file mode 100644 index 00000000..642760f4 --- /dev/null +++ b/plugins/domain/auth/auth_source_resolver.go @@ -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 +} diff --git a/plugins/domain/auth/cache.go b/plugins/domain/auth/cache.go new file mode 100644 index 00000000..0780fa3a --- /dev/null +++ b/plugins/domain/auth/cache.go @@ -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() +} diff --git a/plugins/domain/auth/cache_test.go b/plugins/domain/auth/cache_test.go new file mode 100644 index 00000000..3383868a --- /dev/null +++ b/plugins/domain/auth/cache_test.go @@ -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") + } +} diff --git a/plugins/domain/auth/constants.go b/plugins/domain/auth/constants.go new file mode 100644 index 00000000..872cf164 --- /dev/null +++ b/plugins/domain/auth/constants.go @@ -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" +) diff --git a/plugins/domain/auth/errs.go b/plugins/domain/auth/errs.go new file mode 100644 index 00000000..37fca722 --- /dev/null +++ b/plugins/domain/auth/errs.go @@ -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 = "令牌无管理员权限" +) diff --git a/plugins/domain/auth/handlers.go b/plugins/domain/auth/handlers.go new file mode 100644 index 00000000..c96d8b44 --- /dev/null +++ b/plugins/domain/auth/handlers.go @@ -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()) +} diff --git a/plugins/domain/auth/middleware.go b/plugins/domain/auth/middleware.go new file mode 100644 index 00000000..478e7abc --- /dev/null +++ b/plugins/domain/auth/middleware.go @@ -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() + } +} diff --git a/plugins/domain/auth/models.go b/plugins/domain/auth/models.go new file mode 100644 index 00000000..2fa98b3d --- /dev/null +++ b/plugins/domain/auth/models.go @@ -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 +} diff --git a/plugins/domain/auth/plugin.go b/plugins/domain/auth/plugin.go index 6f236416..0f55f07f 100644 --- a/plugins/domain/auth/plugin.go +++ b/plugins/domain/auth/plugin.go @@ -1,3 +1,6 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + // Package auth provides the authentication, OAuth, session management, and access token domain plugin for Cordis. package auth @@ -7,7 +10,6 @@ import ( "github.com/Rain-kl/Wavelet/core" "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/core/extpoints" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" ) //go:embed migrations/*.sql @@ -81,16 +83,16 @@ func (p *Plugin) Apply(ctx *core.Context) error { // 3. Register HTTP Routes oauthGroup := ctx.Router().Group("/api/v1/oauth") { - oauthGroup.GET("/sources", oauth.GetLoginSources) - oauthGroup.GET("/login", oauth.GetLoginURL) - oauthGroup.GET("/:source/authorize", oauth.Authorize) - oauthGroup.GET("/logout", oauth.Logout) - oauthGroup.POST("/callback", oauth.Callback) - oauthGroup.GET("/user-info", oauth.LoginRequired(), oauth.UserInfo) - oauthGroup.GET("/external-accounts", oauth.LoginRequired(), oauth.ListExternalAccounts) - oauthGroup.POST("/external-accounts/:id/delete", oauth.LoginRequired(), oauth.DeleteExternalAccount) + oauthGroup.GET("/sources", GetLoginSources) + oauthGroup.GET("/login", GetLoginURL) + oauthGroup.GET("/:source/authorize", Authorize) + oauthGroup.GET("/logout", Logout) + oauthGroup.POST("/callback", Callback) + oauthGroup.GET("/user-info", LoginRequired(), UserInfo) + oauthGroup.GET("/external-accounts", LoginRequired(), ListExternalAccounts) + oauthGroup.POST("/external-accounts/:id/delete", LoginRequired(), DeleteExternalAccount) } - ctx.Router().GET("/api/v1/user-info", oauth.LoginRequired(), oauth.UserInfo) + ctx.Router().GET("/api/v1/user-info", LoginRequired(), UserInfo) // 4. Register Settings Schemas ctx.Settings().Register(extpoints.SettingSchema{ diff --git a/plugins/domain/auth/provider_cache.go b/plugins/domain/auth/provider_cache.go new file mode 100644 index 00000000..86465b7b --- /dev/null +++ b/plugins/domain/auth/provider_cache.go @@ -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) +} diff --git a/plugins/domain/auth/service.go b/plugins/domain/auth/service.go index 80de9f0a..8f34acc7 100644 --- a/plugins/domain/auth/service.go +++ b/plugins/domain/auth/service.go @@ -1,4 +1,6 @@ -// Package auth provides authentication, OAuth, session management, and access token domain services. +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + package auth import ( @@ -7,8 +9,6 @@ import ( "sync" "github.com/Rain-kl/Wavelet/core/contracts" - "github.com/Rain-kl/Wavelet/internal/apps/admin" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/gin-gonic/gin" @@ -44,21 +44,21 @@ func newAuthService() contracts.AuthService { } func (s *authServiceImpl) RequireAuthMiddleware() any { - return oauth.LoginRequired() + return LoginRequired() } func (s *authServiceImpl) RequireAdminMiddleware() any { - return admin.LoginAdminRequired() + return AdminRequired() } func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) { if ginCtx, ok := ctx.(*gin.Context); ok { - if u, ok := oauth.GetFromContext[*model.User](ginCtx, oauth.UserObjKey); ok && u != nil { + if u, ok := GetFromContext[*model.User](ginCtx, UserObjKey); ok && u != nil { return toUserDTO(u), nil } } - if v := ctx.Value(oauth.UserObjKey); v != nil { + if v := ctx.Value(UserObjKey); v != nil { if u, ok := v.(*model.User); ok && u != nil { return toUserDTO(u), nil } @@ -76,25 +76,27 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr } tokenHash := model.HashToken(token) - tokenRecord, err := oauth.GetCachedToken(ctx, tokenHash) + tokenRecord, err := GetCachedToken(ctx, tokenHash) if err != nil { dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash) if err != nil { return nil, err } tokenRecord = &dbToken + SetCachedToken(ctx, tokenHash, tokenRecord) } - user, err := oauth.GetCachedUser(ctx, tokenRecord.UserID) + user, err := GetCachedUser(ctx, tokenRecord.UserID) if err != nil || !user.IsActive { dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID) if err != nil { return nil, err } user = &dbUser + SetCachedUser(ctx, tokenRecord.UserID, user) } - if user.Username == "system" { + if user.Username == SystemUsername { return nil, errors.New("auth: system user token not allowed") } @@ -102,12 +104,11 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr } func (s *authServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) { - // Session creation helper return "", nil } func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error { - oauth.InvalidateCachedUser(ctx, userID) + InvalidateCachedUser(ctx, userID) return nil } diff --git a/plugins/domain/auth/session.go b/plugins/domain/auth/session.go new file mode 100644 index 00000000..d8fc3b88 --- /dev/null +++ b/plugins/domain/auth/session.go @@ -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 +} diff --git a/plugins/domain/message_gateway/admin_handlers.go b/plugins/domain/message_gateway/admin_handlers.go new file mode 100644 index 00000000..3e9ec63e --- /dev/null +++ b/plugins/domain/message_gateway/admin_handlers.go @@ -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) + } +} diff --git a/plugins/domain/message_gateway/admin_logics.go b/plugins/domain/message_gateway/admin_logics.go new file mode 100644 index 00000000..33b69a20 --- /dev/null +++ b/plugins/domain/message_gateway/admin_logics.go @@ -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) + } +} diff --git a/plugins/domain/message_gateway/custom_events_admin_login.go b/plugins/domain/message_gateway/custom_events_admin_login.go new file mode 100644 index 00000000..7bb244f4 --- /dev/null +++ b/plugins/domain/message_gateway/custom_events_admin_login.go @@ -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) +} diff --git a/plugins/domain/message_gateway/errs.go b/plugins/domain/message_gateway/errs.go new file mode 100644 index 00000000..2ed7a952 --- /dev/null +++ b/plugins/domain/message_gateway/errs.go @@ -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 = "********" +) diff --git a/plugins/domain/message_gateway/handlers.go b/plugins/domain/message_gateway/handlers.go new file mode 100644 index 00000000..096513bb --- /dev/null +++ b/plugins/domain/message_gateway/handlers.go @@ -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) + } +} diff --git a/plugins/domain/message_gateway/logics.go b/plugins/domain/message_gateway/logics.go new file mode 100644 index 00000000..514ed867 --- /dev/null +++ b/plugins/domain/message_gateway/logics.go @@ -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 +} diff --git a/plugins/domain/message_gateway/migrations/20260827000001_create_message_gateway_tables.sql b/plugins/domain/message_gateway/migrations/20260827000001_create_message_gateway_tables.sql index 36465fc4..4a7dc866 100644 --- a/plugins/domain/message_gateway/migrations/20260827000001_create_message_gateway_tables.sql +++ b/plugins/domain/message_gateway/migrations/20260827000001_create_message_gateway_tables.sql @@ -34,10 +34,59 @@ CREATE TABLE IF NOT EXISTS w_message_pairing_codes ( ); CREATE INDEX IF NOT EXISTS idx_w_message_pairing_lookup ON w_message_pairing_codes (channel_id, platform_user_id); + +CREATE TABLE IF NOT EXISTS w_push_events ( + id BIGINT PRIMARY KEY, + event_key VARCHAR(80) NOT NULL, + name VARCHAR(100) NOT NULL, + task_type VARCHAR(100) NOT NULL DEFAULT '', + channels TEXT NOT NULL DEFAULT '', + targets TEXT NOT NULL DEFAULT '', + template TEXT NOT NULL DEFAULT '', + enabled BOOLEAN NOT NULL DEFAULT FALSE, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_push_events_key ON w_push_events(event_key); +CREATE INDEX IF NOT EXISTS idx_w_push_events_enabled ON w_push_events(enabled); +CREATE INDEX IF NOT EXISTS idx_w_push_events_task_type ON w_push_events(task_type); + +CREATE TABLE IF NOT EXISTS w_push_channels ( + id BIGINT PRIMARY KEY, + name VARCHAR(80) NOT NULL, + description VARCHAR(255) NOT NULL DEFAULT '', + type VARCHAR(50) NOT NULL DEFAULT 'custom', + token VARCHAR(100) NOT NULL DEFAULT '', + url TEXT NOT NULL DEFAULT '', + other TEXT NOT NULL DEFAULT '', + enabled BOOLEAN NOT NULL DEFAULT TRUE, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_push_channels_name ON w_push_channels(name); +CREATE INDEX IF NOT EXISTS idx_w_push_channels_enabled ON w_push_channels(enabled); + +CREATE TABLE IF NOT EXISTS w_push_histories ( + id BIGINT PRIMARY KEY, + event_key VARCHAR(80) NOT NULL, + channel VARCHAR(50) NOT NULL, + target VARCHAR(255) NOT NULL, + title VARCHAR(255) NOT NULL, + content TEXT NOT NULL, + level VARCHAR(20) NOT NULL, + status VARCHAR(20) NOT NULL, + error_msg TEXT NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE INDEX IF NOT EXISTS idx_w_push_histories_event ON w_push_histories(event_key); +CREATE INDEX IF NOT EXISTS idx_w_push_histories_created ON w_push_histories(created_at); -- +goose StatementEnd -- +goose Down -- +goose StatementBegin +DROP TABLE IF EXISTS w_push_histories; +DROP TABLE IF EXISTS w_push_channels; +DROP TABLE IF EXISTS w_push_events; DROP TABLE IF EXISTS w_message_pairing_codes; DROP TABLE IF EXISTS w_message_bindings; DROP TABLE IF EXISTS w_message_channels; diff --git a/plugins/domain/message_gateway/plugin.go b/plugins/domain/message_gateway/plugin.go index 51840fd0..74916359 100644 --- a/plugins/domain/message_gateway/plugin.go +++ b/plugins/domain/message_gateway/plugin.go @@ -1,3 +1,6 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + // Package message_gateway provides the Bot gateway, multi-channel notification dispatching, and asynchronous push worker domain plugin for Cordis. package message_gateway @@ -7,7 +10,7 @@ import ( "github.com/Rain-kl/Wavelet/core" "github.com/Rain-kl/Wavelet/core/extpoints" - appgw "github.com/Rain-kl/Wavelet/internal/apps/message_gateway" + "github.com/Rain-kl/Wavelet/internal/apps/admin" "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/hibiken/asynq" ) @@ -18,8 +21,18 @@ var mgMigrations embed.FS // Option configures the message_gateway plugin. type Option func(*Plugin) +// WithAutoStartRunner enables automatic bot runner startup in the background. +func WithAutoStartRunner(enable bool) Option { + return func(p *Plugin) { + p.autoStartRunner = enable + } +} + // Plugin implements core.Plugin to provide Bot gateway and notification dispatch domain services. -type Plugin struct{} +type Plugin struct { + autoStartRunner bool + cancelRunner context.CancelFunc +} // New creates a new message_gateway domain plugin. func New(opts ...Option) *Plugin { @@ -61,36 +74,100 @@ func (p *Plugin) Apply(ctx *core.Context) error { // 1. Register migrations ctx.Migrations().Register("message_gateway", mgMigrations) - // 2. Register HTTP Routes + // 2. Register User HTTP Routes mgGroup := ctx.Router().Group("/api/v1/message-gateway", oauth.LoginRequired()) { - mgGroup.GET("/channels", appgw.ListChannels) - mgGroup.GET("/bindings", appgw.ListBindings) - mgGroup.POST("/bindings", appgw.BindBinding) - mgGroup.DELETE("/bindings/:id", appgw.UnbindBinding) + mgGroup.GET("/channels", ListChannels) + mgGroup.GET("/bindings", ListBindings) + mgGroup.POST("/bindings", BindBinding) + mgGroup.DELETE("/bindings/:id", UnbindBinding) + } + + // 3. Register Admin Message Gateway HTTP Routes + adminMgGroup := ctx.Router().Group("/api/v1/admin/message-gateway", oauth.LoginRequired(), admin.LoginAdminRequired()) + { + adminMgGroup.GET("/channels/definitions", ListAdminChannelDefinitions) + adminMgGroup.GET("/channels", ListAdminChannels) + adminMgGroup.POST("/channels", CreateAdminChannel) + adminMgGroup.PATCH("/channels/:id", UpdateAdminChannel) + adminMgGroup.DELETE("/channels/:id", DeleteAdminChannel) + adminMgGroup.POST("/channels/:id/test", TestAdminChannel) + } + + // 4. Register Admin Push HTTP Routes + adminPushGroup := ctx.Router().Group("/api/v1/admin/push", oauth.LoginRequired(), admin.LoginAdminRequired()) + { + events := adminPushGroup.Group("/events") + { + events.GET("", ListPushEvents) + events.GET("/builtin", ListBuiltInPushEvents) + events.POST("", CreatePushEvent) + events.PUT("/:id", UpdatePushEvent) + events.DELETE("/:id", DeletePushEvent) + events.POST("/:id/toggle", TogglePushEvent) + } + + adminPushGroup.GET("/histories", ListPushHistories) + adminPushGroup.POST("/test", TestPush) + + channels := adminPushGroup.Group("/channels") + { + channels.GET("/definitions", ListPushChannelDefinitions) + channels.GET("", ListPushChannels) + channels.POST("", CreatePushChannel) + channels.PUT("/:id", UpdatePushChannel) + channels.DELETE("/:id", DeletePushChannel) + channels.POST("/test", TestPushChannel) + } } const defaultTaskRetry = 3 + pushHandler := &PushHandler{} - // 3. Register Asynq background tasks - ctx.Task().Register("message_gateway:push_notification", func(_ context.Context, _ *asynq.Task) error { - return nil + // 5. Register Asynq background tasks + ctx.Task().Register("message_gateway:push_notification", func(c context.Context, t *asynq.Task) error { + _, err := pushHandler.Execute(c, t.Payload()) + return err + }, extpoints.WithTaskRetry(defaultTaskRetry)) + + ctx.Task().Register(SendNotificationTask, func(c context.Context, t *asynq.Task) error { + _, err := pushHandler.Execute(c, t.Payload()) + return err }, extpoints.WithTaskRetry(defaultTaskRetry)) ctx.Task().Register("message_gateway:dispatch_bot_msg", func(_ context.Context, _ *asynq.Task) error { return nil }) - // 4. Register Cron Schedules + // 6. Register Cron Schedules ctx.Schedule().RegisterCron("*/10 * * * *", "message_gateway:cleanup_pairing_codes", map[string]any{"action": "cleanup"}) - // 5. Register EventBus listeners for decoupled push triggers - ctx.Events().On("notification:push", func(_ context.Context, _ PushNotificationEvent) error { - // Event triggered push handling + // 7. Register EventBus listeners for decoupled push triggers + ctx.Events().On("notification:push", func(c context.Context, e PushNotificationEvent) error { + meta := EventMetadata{ + Key: "eventbus:" + e.Channel, + Name: e.Title, + DefaultTemplate: NotificationMessage{ + Title: e.Title, + Content: e.Content, + Level: defaultLevelInfo, + Ext: e.Metadata, + }, + Description: "EventBus triggered notification", + } + DefaultTrigger.Trigger(c, meta, map[string]any{ + "user.id": e.UserID, + "title": e.Title, + "content": e.Content, + }) return nil }) - // 6. Register Settings Schemas + // 8. Register built-in domain events and task listeners + RegisterCustomEvents() + RegisterTaskListeners() + + // 9. Register Settings Schemas ctx.Settings().Register(extpoints.SettingSchema{ Key: "message_gateway.pairing_code_expiry_minutes", Default: 15, @@ -106,5 +183,21 @@ func (p *Plugin) Apply(ctx *core.Context) error { Category: "messaging", }) + // 10. Optional runner start & lifecycle + if p.autoStartRunner { + runnerCtx, cancel := context.WithCancel(ctx.GoContext()) + p.cancelRunner = cancel + go func() { + _ = Start(runnerCtx) + }() + } + + ctx.OnDispose(func() error { + if p.cancelRunner != nil { + p.cancelRunner() + } + return nil + }) + return nil } diff --git a/plugins/domain/message_gateway/push_channels.go b/plugins/domain/message_gateway/push_channels.go new file mode 100644 index 00000000..9dcae0d9 --- /dev/null +++ b/plugins/domain/message_gateway/push_channels.go @@ -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 +} diff --git a/plugins/domain/message_gateway/push_constants.go b/plugins/domain/message_gateway/push_constants.go new file mode 100644 index 00000000..b9559547 --- /dev/null +++ b/plugins/domain/message_gateway/push_constants.go @@ -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" +) diff --git a/plugins/domain/message_gateway/push_events.go b/plugins/domain/message_gateway/push_events.go new file mode 100644 index 00000000..09631d6a --- /dev/null +++ b/plugins/domain/message_gateway/push_events.go @@ -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) +} diff --git a/plugins/domain/message_gateway/push_handlers.go b/plugins/domain/message_gateway/push_handlers.go new file mode 100644 index 00000000..f5fe7aab --- /dev/null +++ b/plugins/domain/message_gateway/push_handlers.go @@ -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()) +} diff --git a/plugins/domain/message_gateway/push_logics.go b/plugins/domain/message_gateway/push_logics.go new file mode 100644 index 00000000..4299cabc --- /dev/null +++ b/plugins/domain/message_gateway/push_logics.go @@ -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 + } + } +} diff --git a/plugins/domain/message_gateway/push_task_listener.go b/plugins/domain/message_gateway/push_task_listener.go new file mode 100644 index 00000000..1e45ae60 --- /dev/null +++ b/plugins/domain/message_gateway/push_task_listener.go @@ -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 "" +} diff --git a/plugins/domain/message_gateway/push_tasks.go b/plugins/domain/message_gateway/push_tasks.go new file mode 100644 index 00000000..8d91a072 --- /dev/null +++ b/plugins/domain/message_gateway/push_tasks.go @@ -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) + } +} diff --git a/plugins/domain/message_gateway/runner.go b/plugins/domain/message_gateway/runner.go new file mode 100644 index 00000000..daa0593c --- /dev/null +++ b/plugins/domain/message_gateway/runner.go @@ -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 +} diff --git a/plugins/domain/message_gateway/secret.go b/plugins/domain/message_gateway/secret.go new file mode 100644 index 00000000..62f0490e --- /dev/null +++ b/plugins/domain/message_gateway/secret.go @@ -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) +} diff --git a/plugins/domain/risk_control/logics.go b/plugins/domain/risk_control/logics.go new file mode 100644 index 00000000..aaa67407 --- /dev/null +++ b/plugins/domain/risk_control/logics.go @@ -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: + } + } +} diff --git a/plugins/domain/risk_control/middleware.go b/plugins/domain/risk_control/middleware.go new file mode 100644 index 00000000..b3769633 --- /dev/null +++ b/plugins/domain/risk_control/middleware.go @@ -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) + } +} diff --git a/plugins/domain/risk_control/middleware_test.go b/plugins/domain/risk_control/middleware_test.go new file mode 100644 index 00000000..eb41af2f --- /dev/null +++ b/plugins/domain/risk_control/middleware_test.go @@ -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"], "系统繁忙") + }) +} diff --git a/plugins/domain/risk_control/plugin.go b/plugins/domain/risk_control/plugin.go index 05055875..21b870ab 100644 --- a/plugins/domain/risk_control/plugin.go +++ b/plugins/domain/risk_control/plugin.go @@ -1,3 +1,6 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + // Package risk_control provides the access control, IP rate limiting, and telemetry risk analysis domain plugin for Cordis. package risk_control @@ -6,7 +9,6 @@ import ( "github.com/Rain-kl/Wavelet/core" "github.com/Rain-kl/Wavelet/core/extpoints" - "github.com/Rain-kl/Wavelet/internal/apps/risk_control" "github.com/gin-gonic/gin" ) @@ -54,12 +56,12 @@ func (p *Plugin) Manifest() core.Manifest { // Apply registers risk control middlewares, settings, and cleanup hooks into the Context. func (p *Plugin) Apply(ctx *core.Context) error { // 1. Initialize LogWriter if needed - risk_control.InitLogWriter(ctx.GoContext()) + InitLogWriter(ctx.GoContext()) // 2. Register router middleware mw := p.middleware if mw == nil { - mw = risk_control.RiskControlMiddleware() + mw = RiskControlMiddleware() } ctx.Router().Use(mw) @@ -81,7 +83,7 @@ func (p *Plugin) Apply(ctx *core.Context) error { // 4. Register lifecycle disposal cleanup ctx.OnDispose(func() error { - return risk_control.StopLogWriter(context.Background()) + return StopLogWriter(context.Background()) }) return nil