refactor(plugins): restructure admin and message_gateway into standard layered sub-packages

This commit is contained in:
ryan
2026-08-28 22:33:26 +08:00
parent 85b383a4e0
commit f4975d6732
128 changed files with 12050 additions and 9901 deletions
+26 -18
View File
@@ -16,37 +16,45 @@ description: "Wavelet 项目专用:当新增或修改业务 API、Handler、
### 插件目录推荐结构 (`backend/plugins/domain/<name>/` 或下游 `custom_plugins/<name>/`)
#### 模式 1:扁平自包含分层(适用于简单业务逻辑 / 推荐默认)
#### 模式 1:极简单文件自包含(适用于极简微型插件 / 单一实体 / <500行)
```text
backend/plugins/domain/order/
backend/plugins/domain/demo/
├── plugin.go # 插件入口:实现 core.Plugin,通过 ctx.Router() 挂载路由
├── handlers.go # HTTP 控制器:参数校验、上下文提取、调用 Service、信封响应
├── service.go # 业务服务层:纯 Go 逻辑,仅依赖 context.Context
├── repository.go # 数据库访问层:GORM 查询、SQL 防注入与转义
├── handlers.go # HTTP 控制器单文件:参数校验、上下文提取、调用 Service、信封响应
├── service.go # 业务服务层单文件:纯 Go 逻辑,仅依赖 context.Context
├── repository.go # 数据库访问层单文件:GORM 查询、SQL 防注入与转义
├── models.go # GORM 数据实体定义(自带表前缀)与 DTO
├── errs.go # 模块内错误常量定义(camelCase 字符串)
└── migrations/ # 专属嵌入式 Goose SQL 迁移脚本
└── 20260827000001_create_orders_table.sql
└── 20260827000001_create_demo_table.sql
```
> ⚠️ **严禁**:当需要拆分多个 Handler/Service 文件时,**严禁在根目录平铺 `handlers_*.go`、`service_*.go`、`repository_*.go` 等前缀文件**,必须立即采用模式 2(独立子包分层)。
#### 模式 2:严格子包分层架构(适用于复杂业务逻辑 / 多聚合根 / 大代码量)
#### 模式 2:标准独立子包分层架构(适用于标准/中大型业务插件 / 官方推荐标准)
```text
backend/plugins/domain/order/
├── plugin.go # 插件根入口:实现 core.Plugin,装配各子包并向 Cordis 注册
├── controller/ # package controller:HTTP 控制器与路由声明
│ ├── http.go
│ └── router.go
│
├── handler/ # package handler:HTTP 控制器与路由声明(或 controller/)
│ ├── router.go # 路由组声明与中间件挂载
│ └── order.go # 订单 Handler(直接以业务命名,禁止 handlers_order.go)
│
├── service/ # package service:业务逻辑层(用例编排、事件发布)
│ ├── service.go
│ └── service_impl.go
│ ├── service.go # Service 接口与组装
│ └── order.go # 订单业务用例实现(直接以业务命名,禁止 service_order.go)
│
├── repository/ # package repository:数据持久化访问层 (DAL)
│ ├── repository.go
│ └── repository_impl.go
├── model/ # package model:纯数据实体与 DTO(无外部依赖)
│ ├── entity.go
│ └── dto.go
├── errs/ # package errs:错误常量与错误码
│ ├── repository.go # 仓储抽象与通用工厂
│ └── order.go # 订单仓储实现(直接以业务命名,禁止 repository_order.go)
│
├── model/ # package model (或 models/):纯数据实体与 DTO(无外部依赖)
│ ├── entity.go # 数据库映射实体 (TableName() 带插件专属前缀)
│ ├── dto.go # 请求与响应 DTO
│ └── events.go # 领域事件定义
│
├── errs/ # package errs:错误常量与错误码 (或根目录 errs.go)
│ └── errs.go
│
└── migrations/ # 专属嵌入式 Goose SQL 迁移脚本
└── 20260827000001_create_orders_table.sql
```
+1
View File
@@ -66,3 +66,4 @@ s3_cache
/backend/plugins/domain/upload/filesrv/uploads/
/backend/plugins/domain/upload/task/uploads/
/backend/data/
/backend/plugins/drivers/driver_http/dist/
+2 -2
View File
@@ -102,8 +102,8 @@ Strong success criteria let you loop independently. Weak criteria ("make it work
- 所有业务功能与驱动实现均以插件形式存在(`backend/plugins/drivers/`、`backend/plugins/infra/`、`backend/plugins/domain/` 或下游 `backend/downstream/`)。
- 每个插件实现 `core.Plugin`(`Name() string` 与 `Apply(ctx *core.Context) error`)。
- **分层模式选型**:
- **模式 1(扁平自包含分层,简单业务默认)**:单 package 结构(`plugin.go`, `handlers.go`, `service.go`, `repository.go`, `models.go`, `errs.go`, `migrations/`)。
- **模式 2(严格子包物理分层,复杂业务使用)**:多 package 物理隔离(`plugin.go`, `controller/`, `service/`, `repository/`, `model/`, `errs/`, `migrations/`),严格约束 `controller -> service -> repository -> model` 单向依赖。
- **模式 1(极简单文件分层,微型插件)**:单 package 极简结构(仅单文件 `plugin.go`, `handlers.go`, `service.go`, `repository.go`, `models.go`, `errs.go`, `migrations/`)。
- **模式 2(标准独立子包分层,推荐标准)**:多 package 物理隔离(`plugin.go`, `handler/`, `service/`, `repository/`, `model/`, `errs/`, `migrations/`)。**严禁在根包平铺 `handlers_*`、`service_*`、`repository_*` 等前缀文件**,子包内文件直接按业务命名(如 `user.go`, `config.go`),严格约束 `handler -> service -> repository -> model` 单向依赖。
- **插件通信与依赖隔离**:
- **严禁跨包 import internal/私有实现**:插件之间严禁直接 import 对方具体实现包代码。
- **单向服务契约调用**:调用方仅面向 `backend/core/contracts` 编程,在 `Apply` 中通过 `core.Provide[contracts.XxxService](ctx, svc)` 注册服务,通过 `core.Inject[contracts.XxxService](ctx)` 或 `ctx.Using(func(svc contracts.XxxService) { ... })` 声明式解析。
+983 -983
View File
File diff suppressed because it is too large Load Diff
+983 -983
View File
File diff suppressed because it is too large Load Diff
+684 -684
View File
File diff suppressed because it is too large Load Diff
-82
View File
@@ -1,82 +0,0 @@
// 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 = "更新用户信息失败"
)
+246
View File
@@ -0,0 +1,246 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package errs defines error constants, sentinels, and error helpers for the admin domain.
package errs
import (
"errors"
"strings"
)
// 管理后台公共错误常量
const (
AdminRequired = "未经授权访问"
TokenAdminRequired = "该访问令牌没有管理员权限,无法访问管理端点" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
InvalidAuthSourceID = "认证源 ID 无效"
ErrInvalidAuthSourceID = "无效的认证源 ID"
InvalidCursorParam = "无效的 cursor 参数"
InvalidTaskExecutionID = "无效的任务执行记录 ID"
InvalidParams = "无效的参数"
InternalServerError = "内部服务器错误"
InvalidScheduleID = "无效的定时任务ID"
)
// 依赖服务未就绪错误常量
const (
DatabaseNotInitialized = "数据库未初始化"
ErrDatabaseServiceNotAvailable = "database service not available"
ErrDatabaseNotInitialized = "database not initialized"
ErrCacheServiceNotInitialized = "cache service is not initialized"
UserServiceUnavailable = "用户服务未就绪"
AuthServiceUnavailable = "认证服务未就绪"
TaskServiceUnavailable = "task service not available"
LogStoreUnavailable = "日志存储服务未初始化"
)
// 系统配置错误消息常量
const (
SystemConfigNotFound = "系统配置不存在"
ConfigKeyRequired = "配置键不能为空"
ConfigValueRequired = "配置值不能为空"
ConfigKeyExists = "配置键已存在"
ProtectedConfigKeyMessage = "该配置项由系统任务管理,禁止手动修改"
StorageDriverSwitchRequiresMigration = "存在存量文件,请通过存储迁移任务切换存储引擎"
ErrConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
ErrConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
ErrConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
ErrParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
ErrCheckExistingUploadsFailed = "检查存量文件失败: %w"
ErrParseCurrentStorageConfigFailed = "解析当前存储配置失败: %w"
ErrParseTargetStorageConfigFailed = "解析目标存储配置失败: %w"
ErrSerializeStorageConfigFailed = "序列化存储配置失败: %w"
ErrAutoResolveMigrationTaskFailed = "自动更新迁移任务状态失败: %v"
StorageMigrationTaskType = "storage:migrate"
StorageDriverResolvedResult = "存储配置直接更新,故障迁移任务自动标记为已解决"
)
// 存储配置校验错误前缀,用于区分参数校验失败与内部错误。
var storageValidationErrPrefixes = []string{
"解析", "验证", "初始化测试", "存储连通性", "序列化", "检查存量文件",
}
// 模板管理相关错误消息常量
const (
TemplateNotFound = "模板不存在"
TemplateKeyRequired = "模板标识符不能为空"
TemplateNameRequired = "模板名称不能为空"
TemplateContentRequired = "模板内容不能为空"
TemplateKeyExists = "模板标识符已存在"
SystemTemplateCannotDelete = "系统预置模板不可删除"
SystemTemplateCannotModifyKey = "系统预置模板不可修改标识符"
)
// 任务调度相关错误消息常量
const (
InvalidTaskType = "无效的任务类型"
InvalidTimeRange = "无效的时间范围"
TaskDispatchFailed = "任务下发失败"
UserIDRequired = "用户ID必填"
TaskNotFound = "任务执行记录不存在"
TaskNotRetryable = "该任务不支持重试"
TaskNotFailed = "只有失败的任务才能重试"
TaskMaxRetryExceeded = "已达到最大重试次数"
TaskRetryFailed = "任务重试失败"
ScheduleSaveFailed = "保存定时任务失败"
ScheduleDeleteFailed = "删除定时任务失败"
InvalidCronExpression = "无效的 Cron 表达式"
ScheduleNotFound = "定时任务不存在"
// 任务契约实现返回的远端错误文案,用于状态码归类。
RemoteTaskNotFoundMsg = "不存在"
RemoteTaskNotFailedMsg = "只有失败的任务"
RemoteTaskNotRetryableMsg = "不支持重试"
RemoteTaskMaxRetryMsg = "已达到最大重试"
)
// 数据库管理相关错误消息常量
const (
InvalidSQLStatement = "SQL 语句不能为空"
ErrOpenDatabaseFileFailed = "无法打开数据库文件"
ErrReadDatabaseFileInfoFailed = "无法读取数据库文件信息"
ErrPgDumpUnavailable = "pg_dump 不可用,请确保服务器已安装 PostgreSQL 客户端工具"
)
// 访问日志相关错误消息常量
const (
ErrQueryUserFailed = "查询用户信息失败: %w"
ErrQueryAccessTrendFailed = "查询访问趋势失败: "
)
// 日志库切换相关错误消息常量
const (
ErrReadLogDatabaseFailed = "读取日志主库失败: %w"
ErrLogDatabaseEmpty = "日志主库配置为空"
ErrSameLogTarget = "目标日志库与当前日志库相同,无需迁移"
ErrClickHouseNotEnabled = "ClickHouse 未启用,无法迁移到 ClickHouse"
ErrPostgresNotEnabled = "PostgreSQL 未启用(当前主库为 SQLite),无法迁移到 PostgreSQL"
ErrSQLiteNotAllowedAsLogDB = "当前主库为 PostgreSQL,日志库不能设置为 SQLite"
)
// 应用更新相关错误消息常量
const (
ErrInvalidRepository = "上游仓库地址无效"
ErrReleaseRequestFailed = "获取上游版本失败"
ErrReleaseResponseInvalid = "上游版本响应无效"
ErrNoCompatibleRelease = "未找到兼容的 Release"
ErrNoCompatibleAsset = "未找到当前系统对应的 Release 资产"
ErrDevelopmentBuild = "开发版本无法执行自动升级"
ErrAlreadyUpToDate = "当前已是最新版本"
ErrUpgradeAlreadyRunning = "已有升级任务正在执行"
ErrAutomaticUpgradeBlocked = "当前平台暂不支持自动替换二进制"
ErrReleaseAssetSizeInvalid = "release 资产大小无效: %d"
ErrCreateUpgradeRequestFailed = "创建升级下载请求失败: %w"
ErrDownloadUpgradeAssetFailed = "下载升级资产失败: %w"
ErrUpgradeAssetHTTPFailed = "下载升级资产失败: HTTP %d"
ErrCreateUpgradeArchiveFailed = "创建升级归档失败: %w"
ErrWriteUpgradeArchiveFailed = "写入升级归档失败: %w"
ErrCloseUpgradeArchiveFailed = "关闭升级归档失败: %w"
ErrUpgradeArchiveSizeMismatch = "升级归档大小不匹配: got %d, want %d"
ErrArchiveContainsIllegalPath = "归档包含非法路径: %s"
ErrArchivePathOutOfDestination = "归档路径越界: %s"
ErrExtractedBinaryTooLarge = "解压后的程序文件超过大小限制"
ErrLocateExecutableFailed = "定位当前程序失败: %w"
ErrResolveExecutablePathFailed = "解析当前程序路径失败: %w"
ErrCreateUpgradeDirFailed = "创建升级目录失败: %w"
ErrExtractUpgradeAssetFailed = "解压升级资产失败: %w"
)
// 用户管理(管理员视角)错误消息常量
const (
UserNotFound = "用户不存在"
CannotDisable = "不能禁用管理员账号"
CannotDelete = "不能删除管理员账号"
CannotDeleteSelf = "不能删除当前登录账号"
UsernameRequired = "用户名不能为空"
EmailRequired = "邮箱不能为空"
PasswordTooShort = "密码长度不能少于 8 位" //nolint:gosec // error message, not hardcoded credentials
UsernameExists = "用户名已存在"
EmailExists = "邮箱已被使用"
CannotRevokeSelfAdmin = "不能取消自身的管理员权限"
UpdateUserFailed = "更新用户状态失败"
DeleteUserFailed = "删除用户失败"
UpdateUserInfoFailed = "更新用户信息失败"
ListAdminUsersFailed = "获取用户列表失败"
)
// 认证源管理相关错误消息常量
const (
ListAuthSourcesFailed = "获取认证源列表失败"
CreateAuthSourceFailed = "创建认证源失败: "
ToggleAuthSourceFailed = "切换认证源状态失败: "
DeleteAuthSourceFailed = "删除认证源失败: "
)
// 层边界哨兵错误:Service/Repository 层返回、Handler 层据以选择信封状态码。
var (
// ErrDatabaseUninitialized 表示数据库服务尚未注入。
ErrDatabaseUninitialized = errors.New(DatabaseNotInitialized)
// ErrSystemConfigNotFound 表示系统配置键不存在。
ErrSystemConfigNotFound = errors.New(SystemConfigNotFound)
// ErrConfigKeyExists 表示系统配置键已存在。
ErrConfigKeyExists = errors.New(ConfigKeyExists)
// ErrProtectedConfigKey 表示配置键由系统任务托管,禁止手动修改。
ErrProtectedConfigKey = errors.New(ProtectedConfigKeyMessage)
// ErrTemplateNotFound 表示模板标识符不存在。
ErrTemplateNotFound = errors.New(TemplateNotFound)
// ErrTemplateKeyExists 表示模板标识符已被占用。
ErrTemplateKeyExists = errors.New(TemplateKeyExists)
// ErrSystemTemplateCannotDelete 表示系统预置模板不可删除。
ErrSystemTemplateCannotDelete = errors.New(SystemTemplateCannotDelete)
// ErrUserNotFound 表示目标用户不存在。
ErrUserNotFound = errors.New(UserNotFound)
// ErrUserServiceUnavailable 表示用户契约服务尚未注入。
ErrUserServiceUnavailable = errors.New(UserServiceUnavailable)
// ErrAuthServiceUnavailable 表示认证契约服务尚未注入。
ErrAuthServiceUnavailable = errors.New(AuthServiceUnavailable)
// ErrTaskServiceUnavailable 表示任务契约服务尚未注入。
ErrTaskServiceUnavailable = errors.New(TaskServiceUnavailable)
// ErrLogStoreUnavailable 表示日志分析契约服务尚未注入。
ErrLogStoreUnavailable = errors.New(LogStoreUnavailable)
// ErrScheduleNotFound 表示定时任务不存在。
ErrScheduleNotFound = errors.New(ScheduleNotFound)
// ErrInvalidCronExpression 表示 Cron 表达式无法解析。
ErrInvalidCronExpression = errors.New(InvalidCronExpression)
// ErrInvalidTaskType 表示任务类型未在任务注册表中声明。
ErrInvalidTaskType = errors.New(InvalidTaskType)
)
// InvalidInputError marks a failure caused by caller supplied content (an unusable SQL
// statement, a rejected task payload, ...). It carries no HTTP semantics; the handler
// layer decides how such errors surface to the client.
type InvalidInputError struct {
Msg string
}
func (e *InvalidInputError) Error() string { return e.Msg }
// NewInvalidInputError builds an invalid input failure preserving the original message.
func NewInvalidInputError(msg string) error {
return &InvalidInputError{Msg: msg}
}
// AsInvalidInput reports whether err was caused by rejected caller input.
func AsInvalidInput(err error) (string, bool) {
var target *InvalidInputError
if errors.As(err, &target) {
return target.Msg, true
}
return "", false
}
// IsStorageConfigValidationError 判定错误是否属于存储配置参数校验失败。
func IsStorageConfigValidationError(err error) bool {
if err == nil {
return false
}
msg := err.Error()
if msg == StorageDriverSwitchRequiresMigration {
return true
}
for _, prefix := range storageValidationErrPrefixes {
if strings.HasPrefix(msg, prefix) {
return true
}
}
return false
}
@@ -0,0 +1,106 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"Wavelet/core/contracts"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/service"
"net/http"
"strconv"
"github.com/gin-gonic/gin"
)
// ListAuthSources lists all configured authentication sources.
func ListAuthSources(c *gin.Context) {
views, err := service.ListAuthSources(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(views))
}
// CreateAuthSource creates a new authentication source.
func CreateAuthSource(c *gin.Context) {
var source contracts.AuthSourceDTO
if err := c.ShouldBindJSON(&source); err != nil {
response.AbortBadRequest(c, errs.InvalidParams)
return
}
created, err := service.CreateAuthSource(c.Request.Context(), source)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(created))
}
// UpdateAuthSource updates an authentication source.
func UpdateAuthSource(c *gin.Context) {
id, ok := parseAuthSourceID(c)
if !ok {
return
}
var req contracts.AuthSourceDTO
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, errs.InvalidParams)
return
}
updated, err := service.UpdateAuthSource(c.Request.Context(), id, req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(updated))
}
// ToggleAuthSource toggles the active state of an auth source.
func ToggleAuthSource(c *gin.Context) {
id, ok := parseAuthSourceID(c)
if !ok {
return
}
toggled, err := service.ToggleAuthSource(c.Request.Context(), id)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(gin.H{"is_active": toggled.IsActive}))
}
// DeleteAuthSource deletes an authentication source.
func DeleteAuthSource(c *gin.Context) {
id, ok := parseAuthSourceID(c)
if !ok {
return
}
if err := service.DeleteAuthSource(c.Request.Context(), id); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// parseAuthSourceID reads the numeric auth source path parameter.
func parseAuthSourceID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, errs.ErrInvalidAuthSourceID)
return 0, false
}
return id, true
}
@@ -1,25 +1,17 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
package handler
import (
"Wavelet/pkg/response"
"context"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/service"
"net/http"
"strconv"
"github.com/gin-gonic/gin"
pkgcache "Wavelet/pkg/cache/disk"
)
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 数量等)与策略配置
@@ -32,8 +24,7 @@ type updateCacheConfigRequest struct {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/cache/status [get]
func GetCacheStatus(c *gin.Context) {
status := pkgcache.Default().Status()
c.JSON(http.StatusOK, response.OK(status))
c.JSON(http.StatusOK, response.OK(service.DiskCacheStatus()))
}
// UpdateCacheConfig 更新磁盘缓存策略配置
@@ -42,7 +33,7 @@ func GetCacheStatus(c *gin.Context) {
// @Tags admin
// @Accept json
// @Produce json
// @Param request body updateCacheConfigRequest true "缓存配置请求体"
// @Param request body model.UpdateCacheConfigRequest true "缓存配置请求体"
// @Security SessionCookie
// @Success 200 {object} response.Any "更新成功"
// @Failure 400 {object} response.Any "参数错误"
@@ -51,31 +42,17 @@ func GetCacheStatus(c *gin.Context) {
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/cache/config [post]
func UpdateCacheConfig(c *gin.Context) {
var req updateCacheConfigRequest
var req model.UpdateCacheConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
ctx := c.Request.Context()
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
if err := service.UpdateDiskCachePolicy(c.Request.Context(), req); err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
response.AbortInternal(c, err.Error())
return
}
pkgcache.Default().UpdatePolicy(req.MaxSizeMB, req.TTLMinutes, req.LRUEnabled)
c.JSON(http.StatusOK, response.OKNil())
}
@@ -91,13 +68,9 @@ func UpdateCacheConfig(c *gin.Context) {
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/cache/clear [post]
func ClearCache(c *gin.Context) {
if err := pkgcache.Default().Clear(); err != nil {
if err := service.ClearDiskCache(); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
func saveOrUpdateCacheConfig(ctx context.Context, key, value string) error {
return SaveOrUpdateSystemConfig(ctx, key, value)
}
@@ -0,0 +1,187 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"Wavelet/pkg/response"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/service"
"errors"
"net/http"
"github.com/gin-gonic/gin"
)
// 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) {
resp, err := service.PublicSystemConfigs(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
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) {
c.Data(http.StatusOK, "text/plain; charset=utf-8", []byte(service.RobotsTxtBody(c.Request.Context())))
}
// CreateSystemConfig 创建系统配置
// @Summary 创建系统配置
// @Description 创建一条新的系统配置项,配置键不可重复,同时将新配置同步到 Redis,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body model.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 model.CreateSystemConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := service.CreateAdminSystemConfig(c.Request.Context(), req); err != nil {
if errors.Is(err, errs.ErrProtectedConfigKey) || errors.Is(err, errs.ErrConfigKeyExists) {
response.AbortBadRequest(c, err.Error())
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 := service.ListAdminSystemConfigs(c.Request.Context(), c.Query("type"))
if err != nil {
response.AbortInternal(c, err.Error())
return
}
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 := service.GetAdminSystemConfig(c.Request.Context(), c.Param("key"))
if err != nil {
if errors.Is(err, errs.ErrSystemConfigNotFound) {
response.AbortNotFound(c, errs.SystemConfigNotFound)
return
}
response.AbortInternal(c, err.Error())
return
}
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 model.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 model.UpdateSystemConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
key := c.Param("key")
if err := service.UpdateAdminSystemConfig(c.Request.Context(), key, req); err != nil {
if errors.Is(err, errs.ErrSystemConfigNotFound) {
response.AbortNotFound(c, errs.SystemConfigNotFound)
return
}
if errors.Is(err, errs.ErrProtectedConfigKey) || errs.IsStorageConfigValidationError(err) {
response.AbortBadRequest(c, err.Error())
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// TestSMTP 测试 SMTP 邮件发送
// @Summary 测试 SMTP 邮件发送
// @Description 使用传入的配置进行 SMTP 邮件发送测试,支持使用 ****** 占位符使用保存的数据库密码
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body model.TestSMTPRequest true "测试请求参数"
// @Success 200 {object} response.Any{data=model.TestSMTPResponse} "测试执行完毕"
// @Failure 400 {object} response.Any "参数错误"
// @Router /api/v1/admin/system-configs/smtp/test [post]
func TestSMTP(c *gin.Context) {
var req model.TestSMTPRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(service.TestSMTP(c.Request.Context(), req)))
}
+190
View File
@@ -0,0 +1,190 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"Wavelet/pkg/config"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/service"
"fmt"
"net/http"
"strings"
"github.com/gin-gonic/gin"
)
// GetDBOverview 获取数据库运行概览
// @Summary 获取数据库运行概览
// @Description 获取数据库类型、版本、名称、文件大小、表数量及当前连接数,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=model.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) {
overview, err := service.DatabaseOverview(c.Request.Context())
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) {
tables, err := service.DatabaseTableNames(c.Request.Context())
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 model.GetTableDataRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
data, err := service.DatabaseTableData(c.Request.Context(), req)
if err != nil {
if msg, ok := errs.AsInvalidInput(err); ok {
response.AbortBadRequest(c, msg)
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(data))
}
// ExecuteSQL 执行 SQL 查询
// @Summary 执行 SQL 查询
// @Description 在当前数据库中执行任意自定义 SQL,如果是查询语句将返回格式化后的列与数据集,否则返回受影响行数,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body model.ExecuteSQLRequest true "SQL 请求参数"
// @Success 200 {object} response.Any{data=model.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 model.ExecuteSQLRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
trimmedSQL := strings.TrimSpace(req.SQL)
if trimmedSQL == "" {
response.AbortBadRequest(c, errs.InvalidSQLStatement)
return
}
resp, err := service.ExecuteCustomSQL(c.Request.Context(), trimmedSQL)
if err != nil {
if err == errs.ErrDatabaseUninitialized {
response.AbortInternal(c, err.Error())
return
}
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
// GetDatabaseInfo 获取当前数据库类型及版本信息
// @Summary 获取数据库信息
// @Description 返回当前使用的数据库类型(sqlite/postgres)、名称/路径及版本字符串,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=model.DatabaseInfoResponse} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/db-info [get]
func GetDatabaseInfo(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(service.DatabaseInfo(c.Request.Context())))
}
// 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) {
f, fi, err := service.OpenSQLiteExportFile()
if err != nil {
response.AbortInternal(c, err.Error())
return
}
defer func() {
_ = f.Close()
}()
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) {
cmd, fileName, err := service.NewPgDumpCommand(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
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 {
logger.ErrorF(c.Request.Context(), "[db-export] pg_dump failed: %v", err)
}
}
@@ -0,0 +1,197 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/service"
"encoding/json"
"net/http"
"strconv"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
)
const (
defaultLimit = 200
maxLimit = 500
maxPageSize = 100
)
// wsMessage WebSocket 消息格式
type wsMessage struct {
Type string `json:"type"` // "log" | "error"
Data json.RawMessage `json:"data"`
}
// 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=model.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, errs.InvalidCursorParam)
return
}
if _, err := parsePositiveInt(limitStr, &limit); err != nil || limit <= 0 {
limit = defaultLimit
}
if limit > maxLimit {
limit = maxLimit
}
c.JSON(http.StatusOK, response.OK(service.RecentSystemLogs(cursor, limit)))
}
// 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
}
}
}
}
// 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=model.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) {
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
}
resp, err := service.AccessLogs(c.Request.Context(), model.AccessLogQuery{
Username: c.Query("username"),
Path: c.Query("path"),
StartTime: c.Query("start_time"),
EndTime: c.Query("end_time"),
Page: page,
PageSize: pageSize,
})
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
// GetLogsAnalytics 获取 ClickHouse 访问日志图表聚合指标
// @Summary 获取访问日志分析数据
// @Description 聚合统计最近 7 天的每日访问趋势、浏览器分布以及前 10 名最活跃用户排行(需要管理员权限)
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=model.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) {
resp, err := service.AccessLogAnalytics(c.Request.Context())
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
func getUpgrader() *websocket.Upgrader {
return &websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool {
return service.IsAllowedLogOrigin(r.Context(), r.Header.Get("Origin"), r.Host)
},
}
}
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
}
@@ -1,7 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
package handler
import (
"Wavelet/core/contracts"
@@ -9,6 +9,7 @@ import (
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/trace"
"Wavelet/plugins/domain/admin/errs"
"github.com/gin-gonic/gin"
)
@@ -21,7 +22,7 @@ func LoginAdminRequired() gin.HandlerFunc {
user, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if user == nil {
response.AbortNotFound(c, AdminRequired)
response.AbortNotFound(c, errs.AdminRequired)
return
}
@@ -29,13 +30,13 @@ func LoginAdminRequired() gin.HandlerFunc {
if tokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
tokenAdmin, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
if !tokenAdmin {
response.AbortNotFound(c, TokenAdminRequired)
response.AbortNotFound(c, errs.TokenAdminRequired)
return
}
}
if !user.IsAdmin {
response.AbortNotFound(c, AdminRequired)
response.AbortNotFound(c, errs.AdminRequired)
return
}
@@ -0,0 +1,122 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package handler provides HTTP routing and handlers for the admin domain.
package handler
import (
"Wavelet/core/extpoints"
)
// RegisterRoutes mounts all admin console endpoints under the provided admin router group.
func RegisterRoutes(adminRouter extpoints.RouterExtension) {
// Status & Diagnostics
adminRouter.GET("/status", GetSystemStatus)
adminRouter.GET("/status/log-database", GetLogDatabaseStatus)
adminRouter.GET("/db-info", GetDatabaseInfo)
adminRouter.GET("/db-export", ExportDatabase)
// DB Management
dbGroup := adminRouter.Group("/db-manage")
{
dbGroup.GET("/overview", GetDBOverview)
dbGroup.GET("/tables", ListDBTables)
dbGroup.GET("/table-data", GetDBTableData)
dbGroup.POST("/query", ExecuteSQL)
}
// Cache Management
cacheGroup := adminRouter.Group("/cache")
{
cacheGroup.GET("/status", GetCacheStatus)
cacheGroup.POST("/config", UpdateCacheConfig)
cacheGroup.POST("/clear", ClearCache)
}
// Updater
updateGroup := adminRouter.Group("/update")
{
updateGroup.GET("", GetUpdateStatus)
updateGroup.POST("/apply", ApplyUpdate)
}
// Logs
logsGroup := adminRouter.Group("/logs")
{
logsGroup.GET("", GetLogs)
logsGroup.GET("/access", GetAccessLogs)
logsGroup.GET("/analytics", GetLogsAnalytics)
logsGroup.GET("/ws", HandleLogWebSocket)
}
// Users
usersGroup := adminRouter.Group("/users")
{
usersGroup.GET("", ListUsers)
usersGroup.POST("", CreateUser)
usersGroup.GET("/:id", GetUser)
usersGroup.PUT("/:id/status", UpdateUserStatus)
usersGroup.PUT("/:id", UpdateUser)
usersGroup.DELETE("/:id", DeleteUser)
}
// Auth Sources
authSourcesGroup := adminRouter.Group("/auth-sources")
{
authSourcesGroup.GET("", ListAuthSources)
authSourcesGroup.POST("", CreateAuthSource)
authSourcesGroup.PUT("/:id", UpdateAuthSource)
authSourcesGroup.PUT("/:id/toggle", ToggleAuthSource)
authSourcesGroup.DELETE("/:id", DeleteAuthSource)
}
// System Configs
configGroup := adminRouter.Group("/system-configs")
{
configGroup.GET("", ListSystemConfigs)
configGroup.POST("", CreateSystemConfig)
configGroup.POST("/smtp/test", TestSMTP)
keyGroup := configGroup.Group("/:key")
{
keyGroup.GET("", GetSystemConfig)
keyGroup.PUT("", UpdateSystemConfig)
}
}
// Templates
templateGroup := adminRouter.Group("/templates")
{
templateGroup.GET("", ListTemplates)
templateGroup.POST("", CreateTemplate)
keyGroup := templateGroup.Group("/:key")
{
keyGroup.GET("", GetTemplate)
keyGroup.PUT("", UpdateTemplate)
keyGroup.DELETE("", DeleteTemplate)
}
}
// Tasks
taskGroup := adminRouter.Group("/tasks")
{
taskGroup.GET("/types", ListTaskTypes)
taskGroup.POST("/dispatch", DispatchTask)
executions := taskGroup.Group("/executions")
{
executions.GET("", ListTaskExecutions)
executions.GET("/:id", GetTaskExecution)
executions.POST("/:id/retry", RetryTask)
}
schedules := taskGroup.Group("/schedules")
{
schedules.GET("", ListSchedules)
schedules.POST("", CreateSchedule)
schedules.PUT("/:id", UpdateSchedule)
schedules.DELETE("/:id", DeleteSchedule)
}
}
}
@@ -0,0 +1,41 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"Wavelet/pkg/response"
"Wavelet/plugins/domain/admin/service"
"net/http"
"github.com/gin-gonic/gin"
)
// GetSystemStatus 获取系统状态信息
// @Summary 获取系统状态信息
// @Description 获取后端服务运行状态、Goroutine、内存指标等详细统计数据,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=model.SystemStatusResponse} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/status [get]
func GetSystemStatus(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(service.CollectSystemStatus()))
}
// GetLogDatabaseStatus 返回当前日志库状态。
// @Summary 获取日志数据库状态
// @Description 返回当前日志主库、迁移状态、各库保留天数与合法迁移目标,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=model.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) {
c.JSON(http.StatusOK, response.OK(service.LogDatabaseStatus(c.Request.Context())))
}
@@ -1,22 +1,42 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
package handler
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"fmt"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/service"
"errors"
"net/http"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/robfig/cron/v3"
)
// abortTaskLogicError maps a task service failure onto the unified response envelope.
func abortTaskLogicError(c *gin.Context, err error) bool {
if err == nil {
return false
}
if errors.Is(err, errs.ErrTaskServiceUnavailable) || errors.Is(err, errs.ErrScheduleNotFound) {
response.AbortInternal(c, err.Error())
return true
}
msg := err.Error()
if errors.Is(err, errs.ErrInvalidCronExpression) || errors.Is(err, errs.ErrInvalidTaskType) {
response.AbortBadRequest(c, msg)
return true
}
if text, ok := errs.AsInvalidInput(err); ok {
response.AbortBadRequest(c, text)
return true
}
response.AbortInternal(c, msg)
return true
}
// ListTaskTypes 获取支持的任务类型列表
// @Summary 获取支持的任务类型
// @Description 返回系统支持的所有可调度任务类型列表,包括任务名称、描述、是否支持时间范围等元数据,需要管理员权限
@@ -28,21 +48,7 @@ import (
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/types [get]
func ListTaskTypes(c *gin.Context) {
taskSvc := GetTaskService()
if taskSvc == nil {
c.JSON(http.StatusOK, response.OK([]contracts.TaskMetaDTO{}))
return
}
c.JSON(http.StatusOK, response.OK(taskSvc.ListTasks()))
}
// 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"`
c.JSON(http.StatusOK, response.OK(service.ListTaskTypes()))
}
// DispatchTask 下发任务
@@ -52,7 +58,7 @@ type DispatchTaskRequest struct {
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body DispatchTaskRequest true "任务请求参数"
// @Param request body model.DispatchTaskRequest true "任务请求参数"
// @Success 200 {object} response.Any{data=string} "任务已入队"
// @Failure 400 {object} response.Any "任务类型不存在或参数错误"
// @Failure 401 {object} response.Any "未登录"
@@ -60,38 +66,26 @@ type DispatchTaskRequest struct {
// @Failure 500 {object} response.Any "任务入队失败"
// @Router /api/v1/admin/tasks/dispatch [post]
func DispatchTask(c *gin.Context) {
var req DispatchTaskRequest
var req model.DispatchTaskRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
taskSvc := GetTaskService()
if taskSvc == nil {
response.AbortInternal(c, "task service not available")
return
}
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
if !ok {
response.AbortBadRequest(c, InvalidTaskType)
return
}
var payloadBytes []byte
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := taskSvc.ValidatePayload(meta.Name, payloadBytes)
taskID, err := service.DispatchTask(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
taskID, err := taskSvc.Dispatch(c.Request.Context(), req.TaskType, validated, "manual")
if err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err))
switch {
case errors.Is(err, errs.ErrTaskServiceUnavailable):
response.AbortInternal(c, err.Error())
case errors.Is(err, errs.ErrInvalidTaskType):
response.AbortBadRequest(c, errs.InvalidTaskType)
default:
if text, ok := errs.AsInvalidInput(err); ok {
response.AbortBadRequest(c, text)
return
}
response.AbortInternal(c, err.Error())
}
return
}
@@ -113,22 +107,13 @@ func DispatchTask(c *gin.Context) {
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/executions [get]
func ListTaskExecutions(c *gin.Context) {
var req ListTaskExecutionsRequest
var req model.ListTaskExecutionsRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if req.TaskType != "" {
taskSvc := GetTaskService()
if taskSvc != nil {
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
req.TaskType = meta.Name
}
}
}
executions, total, err := ListTaskExecutionRecords(c.Request.Context(), req)
executions, total, err := service.ListTaskExecutions(c.Request.Context(), req)
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -149,22 +134,22 @@ func ListTaskExecutions(c *gin.Context) {
// @Produce json
// @Security SessionCookie
// @Param id path int true "任务执行记录 ID"
// @Success 200 {object} response.Any{data=TaskExecution} "任务执行详情"
// @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)
id, err := parseUintParam(c, errs.InvalidTaskExecutionID)
if err != nil {
response.AbortBadRequest(c, InvalidTaskExecutionID)
response.AbortBadRequest(c, err.Error())
return
}
execution, err := GetTaskExecutionByID(c.Request.Context(), id)
execution, err := service.TaskExecution(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, TaskNotFound)
response.AbortNotFound(c, errs.TaskNotFound)
return
}
@@ -186,28 +171,23 @@ func GetTaskExecution(c *gin.Context) {
// @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)
id, err := parseUintParam(c, errs.InvalidTaskExecutionID)
if err != nil {
response.AbortBadRequest(c, InvalidTaskExecutionID)
response.AbortBadRequest(c, err.Error())
return
}
taskSvc := GetTaskService()
if taskSvc == nil {
response.AbortInternal(c, "task service not available")
return
}
newTaskID, err := taskSvc.Retry(c.Request.Context(), id)
newTaskID, err := service.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)
case errors.Is(err, errs.ErrTaskServiceUnavailable):
response.AbortInternal(c, err.Error())
case service.IsRetryMissingError(err):
response.AbortNotFound(c, err.Error())
case service.IsRetryConflictError(err):
response.AbortBadRequest(c, err.Error())
default:
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskRetryFailed, err))
response.AbortInternal(c, err.Error())
}
return
}
@@ -221,12 +201,12 @@ func RetryTask(c *gin.Context) {
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]Schedule} "定时任务列表"
// @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 := ListSchedulesRecord(c.Request.Context())
schedules, err := service.ListSchedules(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -234,15 +214,6 @@ func ListSchedules(c *gin.Context) {
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 表达式和执行参数,并触发调度器热加载,需要管理员权限
@@ -250,80 +221,28 @@ type CreateScheduleRequest struct {
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body CreateScheduleRequest true "创建定时任务请求参数"
// @Success 200 {object} response.Any{data=Schedule} "创建成功的定时任务信息"
// @Param request body model.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
var req model.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)
schedule, err := service.CreateSchedule(c.Request.Context(), req)
if abortTaskLogicError(c, err) {
return
}
taskSvc := GetTaskService()
if taskSvc == nil {
response.AbortInternal(c, "task service not available")
return
}
// 校验关联的异步任务类型
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
if !ok {
response.AbortBadRequest(c, InvalidTaskType)
return
}
// 校验并规范化 Payload
var payloadBytes []byte
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := taskSvc.ValidatePayload(meta.Name, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
schedule := &Schedule{
Name: req.Name,
TaskType: req.TaskType,
Cron: req.Cron,
Payload: string(validated),
IsActive: *req.IsActive,
}
if err := CreateScheduleRecord(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
// 触发调度服务重载
if err := taskSvc.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 表达式、异步任务参数和是否启用等),并触发调度器热加载,需要管理员权限
@@ -332,8 +251,8 @@ type UpdateScheduleRequest struct {
// @Produce json
// @Security SessionCookie
// @Param id path int true "定时任务 ID"
// @Param request body UpdateScheduleRequest true "修改定时任务请求参数"
// @Success 200 {object} response.Any{data=Schedule} "修改后的定时任务信息"
// @Param request body model.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 "无管理员权限"
@@ -341,70 +260,23 @@ type UpdateScheduleRequest struct {
// @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)
id, err := parseUintParam(c, errs.InvalidScheduleID)
if err != nil {
response.AbortBadRequest(c, "无效的定时任务ID")
response.AbortBadRequest(c, err.Error())
return
}
var req UpdateScheduleRequest
var req model.UpdateScheduleRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
schedule, err := GetScheduleByID(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, ScheduleNotFound)
schedule, err := service.UpdateSchedule(c.Request.Context(), id, req)
if abortTaskLogicError(c, err) {
return
}
// 校验 Cron 表达式
if _, err := cron.ParseStandard(req.Cron); err != nil {
response.AbortBadRequest(c, InvalidCronExpression)
return
}
taskSvc := GetTaskService()
if taskSvc == nil {
response.AbortInternal(c, "task service not available")
return
}
// 校验关联的异步任务类型
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
if !ok {
response.AbortBadRequest(c, InvalidTaskType)
return
}
// 校验并规范化 Payload
var payloadBytes []byte
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := taskSvc.ValidatePayload(meta.Name, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
schedule.Name = req.Name
schedule.TaskType = req.TaskType
schedule.Cron = req.Cron
schedule.Payload = string(validated)
schedule.IsActive = *req.IsActive
if err := UpdateScheduleRecord(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
// 触发调度服务重载
if err := taskSvc.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
c.JSON(http.StatusOK, response.OK(schedule))
}
@@ -422,24 +294,25 @@ func UpdateSchedule(c *gin.Context) {
// @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)
id, err := parseUintParam(c, errs.InvalidScheduleID)
if err != nil {
response.AbortBadRequest(c, "无效的定时任务ID")
response.AbortBadRequest(c, err.Error())
return
}
if err := DeleteScheduleRecord(c.Request.Context(), id); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))
if err := service.DeleteSchedule(c.Request.Context(), id); err != nil {
response.AbortInternal(c, err.Error())
return
}
// 触发调度服务重载
taskSvc := GetTaskService()
if taskSvc != nil {
if err := taskSvc.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
}
c.JSON(http.StatusOK, response.OKNil())
}
// parseUintParam reads a positive numeric path parameter.
func parseUintParam(c *gin.Context, invalidMsg string) (uint64, error) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
return 0, errors.New(invalidMsg)
}
return id, nil
}
@@ -1,48 +1,31 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
package handler
import (
"Wavelet/pkg/response"
"context"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/service"
"errors"
"net/http"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
// 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"`
}
// abortTemplateLogicError maps the template service outcome onto the response envelope.
func abortTemplateLogicError(c *gin.Context, err error) bool {
if err == nil {
return false
}
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, TemplateNotFound)
if errors.Is(err, errs.ErrTemplateNotFound) {
response.AbortNotFound(c, errs.TemplateNotFound)
return true
}
msg := err.Error()
switch msg {
case TemplateKeyExists, SystemTemplateCannotDelete:
case errs.TemplateKeyExists, errs.SystemTemplateCannotDelete:
response.AbortBadRequest(c, msg)
return true
}
@@ -57,7 +40,7 @@ func abortTemplateLogicError(c *gin.Context, err error) bool {
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body CreateTemplateRequest true "创建请求参数"
// @Param request body model.CreateTemplateRequest true "创建请求参数"
// @Success 200 {object} response.Any{data=string} "创建成功"
// @Failure 400 {object} response.Any "参数错误或模板标识符已存在"
// @Failure 401 {object} response.Any "未登录"
@@ -65,13 +48,13 @@ func abortTemplateLogicError(c *gin.Context, err error) bool {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates [post]
func CreateTemplate(c *gin.Context) {
var req CreateTemplateRequest
var req model.CreateTemplateRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
tmpl, err := createTemplate(c.Request.Context(), req)
tmpl, err := service.CreateTemplate(c.Request.Context(), req)
if abortTemplateLogicError(c, err) {
return
}
@@ -85,13 +68,13 @@ func CreateTemplate(c *gin.Context) {
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]Template} "模板列表"
// @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())
templates, err := service.ListTemplates(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -107,14 +90,14 @@ func ListTemplates(c *gin.Context) {
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Success 200 {object} response.Any{data=Template} "模板详情"
// @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"))
tmpl, err := service.GetTemplate(c.Request.Context(), c.Param("key"))
if abortTemplateLogicError(c, err) {
return
}
@@ -130,8 +113,8 @@ func GetTemplate(c *gin.Context) {
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Param request body UpdateTemplateRequest true "更新请求参数"
// @Success 200 {object} response.Any{data=Template} "更新成功"
// @Param request body model.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 "无管理员权限"
@@ -139,13 +122,13 @@ func GetTemplate(c *gin.Context) {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates/{key} [put]
func UpdateTemplate(c *gin.Context) {
var req UpdateTemplateRequest
var req model.UpdateTemplateRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
tmpl, err := updateTemplate(c.Request.Context(), c.Param("key"), req)
tmpl, err := service.UpdateTemplate(c.Request.Context(), c.Param("key"), req)
if abortTemplateLogicError(c, err) {
return
}
@@ -168,75 +151,9 @@ func UpdateTemplate(c *gin.Context) {
// @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) {
if err := service.DeleteTemplate(c.Request.Context(), c.Param("key")); abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
func createTemplate(ctx context.Context, req CreateTemplateRequest) (Template, error) {
exists, err := TemplateExistsByKey(ctx, req.Key)
if err != nil {
return Template{}, err
}
if exists {
return Template{}, errors.New(TemplateKeyExists)
}
tmpl := Template{
Key: req.Key,
Name: req.Name,
Type: req.Type,
Subject: req.Subject,
Content: req.Content,
Description: req.Description,
IsSystem: false,
}
if err := tmpl.Validate(); err != nil {
return Template{}, err
}
if err := CreateTemplateRecord(ctx, &tmpl); err != nil {
return Template{}, err
}
return tmpl, nil
}
func listTemplates(ctx context.Context) ([]Template, error) {
return ListTemplatesRecord(ctx)
}
func getTemplate(ctx context.Context, key string) (Template, error) {
return GetTemplateByKey(ctx, key)
}
func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest) (Template, error) {
tmpl, err := GetTemplateByKey(ctx, key)
if err != nil {
return Template{}, err
}
tmpl.Name = req.Name
tmpl.Type = req.Type
tmpl.Subject = req.Subject
tmpl.Content = req.Content
tmpl.Description = req.Description
if err := tmpl.Validate(); err != nil {
return Template{}, err
}
if err := SaveTemplateRecord(ctx, &tmpl); err != nil {
return Template{}, err
}
return tmpl, nil
}
func deleteTemplate(ctx context.Context, key string) error {
tmpl, err := GetTemplateByKey(ctx, key)
if err != nil {
return err
}
if tmpl.IsSystem {
return errors.New(SystemTemplateCannotDelete)
}
return DeleteTemplateRecord(ctx, &tmpl)
}
@@ -0,0 +1,69 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/admin/service"
"context"
"net/http"
"time"
"github.com/gin-gonic/gin"
)
// GetUpdateStatus 获取应用更新状态
// @Summary 获取应用更新状态
// @Description 从系统配置指定的 GitHub 上游仓库查询最新兼容 Release,并与当前服务版本比较
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=model.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 := service.GetUpdateStatus(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 := service.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 := service.ReplaceAndRestart(executable, stagedBinary); err != nil {
service.DefaultUpdaterManager.FinishUpgrade()
logger.ErrorF(context.Background(), "[Updater] replace and restart failed: %v", err)
}
})
}
@@ -1,92 +1,42 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
package handler
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/service"
"errors"
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
// 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)
response.AbortBadRequest(c, errs.UserNotFound)
return 0, false
}
return id, true
}
func toUserResponse(u *contracts.UserDTO) userResponse {
if u == nil {
return userResponse{}
}
return userResponse{
ID: u.ID,
Username: u.Username,
Nickname: u.Nickname,
Email: u.Email,
AvatarURL: u.AvatarURL,
IsActive: u.IsActive,
IsAdmin: u.IsAdmin,
Bio: u.Bio,
Phone: u.Phone,
Gender: u.Gender,
Website: u.Website,
Location: u.Location,
LastLoginAt: u.LastLoginAt,
CreatedAt: u.CreatedAt,
UpdatedAt: u.UpdatedAt,
}
}
// abortUserLogicError maps the user service outcome onto the unified response envelope.
func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbiddenMsgs, badRequestMsgs []string) bool {
if err == nil {
return false
}
if errors.Is(err, gorm.ErrRecordNotFound) {
if errors.Is(err, errs.ErrUserServiceUnavailable) {
response.AbortInternal(c, err.Error())
return true
}
if errors.Is(err, errs.ErrUserNotFound) {
response.AbortNotFound(c, notFoundMsg)
return true
}
@@ -104,7 +54,7 @@ func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbidde
}
}
logger.ErrorF(c.Request.Context(), "Admin user error: %v", err)
response.AbortInternal(c, "内部服务器错误")
response.AbortInternal(c, errs.InternalServerError)
return true
}
@@ -114,27 +64,21 @@ func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbidde
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param request query listUsersRequest true "查询参数"
// @Success 200 {object} response.Any{data=listUsersResponse} "用户列表"
// @Param request query model.ListUsersRequest true "查询参数"
// @Success 200 {object} response.Any{data=model.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
var req model.ListUsersRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
}
total, dtos, err := userSvc.AdminListUsers(c.Request.Context(), contracts.AdminListUsersFilter{
total, dtos, err := service.AdminListUsers(c.Request.Context(), contracts.AdminListUsersFilter{
Page: req.Page,
PageSize: req.PageSize,
UserID: req.UserID,
@@ -142,17 +86,16 @@ func ListUsers(c *gin.Context) {
Email: req.Email,
})
if err != nil {
logger.ErrorF(c.Request.Context(), "List admin users failed: %v", err)
response.AbortInternal(c, "获取用户列表失败")
response.AbortInternal(c, err.Error())
return
}
users := make([]userResponse, 0, len(dtos))
users := make([]model.UserResponse, 0, len(dtos))
for _, dto := range dtos {
users = append(users, toUserResponse(dto))
users = append(users, service.ToUserResponse(dto))
}
c.JSON(http.StatusOK, response.OK(listUsersResponse{
c.JSON(http.StatusOK, response.OK(model.ListUsersResponse{
Users: users,
Total: total,
}))
@@ -165,7 +108,7 @@ func ListUsers(c *gin.Context) {
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Success 200 {object} response.Any{data=userResponse} "用户详情"
// @Success 200 {object} response.Any{data=model.UserResponse} "用户详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
@@ -178,23 +121,12 @@ func GetUser(c *gin.Context) {
return
}
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
targetUser, err := service.AdminGetUser(c.Request.Context(), id)
if abortUserLogicError(c, err, errs.UserNotFound, nil, nil) {
return
}
targetUser, err := userSvc.AdminGetUser(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"`
c.JSON(http.StatusOK, response.OK(service.ToUserResponse(targetUser)))
}
// UpdateUserStatus 更新用户状态(启用/禁用)
@@ -205,7 +137,7 @@ type updateUserStatusRequest struct {
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Param request body updateUserStatusRequest true "状态参数"
// @Param request body model.UpdateUserStatusRequest true "状态参数"
// @Success 200 {object} response.Any{data=string} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
@@ -214,7 +146,7 @@ type updateUserStatusRequest struct {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id}/status [put]
func UpdateUserStatus(c *gin.Context) {
var req updateUserStatusRequest
var req model.UpdateUserStatusRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
@@ -225,17 +157,11 @@ func UpdateUserStatus(c *gin.Context) {
return
}
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
}
if err := userSvc.AdminUpdateUserStatus(c.Request.Context(), id, req.IsActive); err != nil {
if abortUserLogicError(c, err, userNotFound, []string{cannotDisable}, nil) {
if err := service.AdminUpdateUserStatus(c.Request.Context(), id, req.IsActive); err != nil {
if abortUserLogicError(c, err, errs.UserNotFound, []string{errs.CannotDisable}, nil) {
return
}
response.AbortInternal(c, updateUserFailed)
response.AbortInternal(c, errs.UpdateUserFailed)
return
}
@@ -264,37 +190,21 @@ func DeleteUser(c *gin.Context) {
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if currUser == nil {
response.AbortUnauthorized(c, AdminRequired)
response.AbortUnauthorized(c, errs.AdminRequired)
return
}
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
}
if err := userSvc.AdminDeleteUser(c.Request.Context(), currUser.ID, id); err != nil {
if abortUserLogicError(c, err, userNotFound, []string{cannotDelete, cannotDeleteSelf}, nil) {
if err := service.AdminDeleteUser(c.Request.Context(), currUser.ID, id); err != nil {
if abortUserLogicError(c, err, errs.UserNotFound, []string{errs.CannotDelete, errs.CannotDeleteSelf}, nil) {
return
}
response.AbortInternal(c, deleteUserFailed)
response.AbortInternal(c, errs.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 创建一个本地密码登录的新用户,需要管理员权限
@@ -302,27 +212,21 @@ type createUserRequest struct {
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body createUserRequest true "创建用户参数"
// @Success 200 {object} response.Any{data=userResponse} "创建成功"
// @Param request body model.CreateUserRequest true "创建用户参数"
// @Success 200 {object} response.Any{data=model.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
var req model.CreateUserRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
}
newUser, err := userSvc.AdminCreateUser(c.Request.Context(), contracts.AdminCreateUserRequest{
newUser, err := service.AdminCreateUser(c.Request.Context(), contracts.AdminCreateUserRequest{
Username: req.Username,
Password: req.Password,
Nickname: req.Nickname,
@@ -330,19 +234,11 @@ func CreateUser(c *gin.Context) {
IsActive: req.IsActive,
IsAdmin: req.IsAdmin,
})
if abortUserLogicError(c, err, "", nil, []string{usernameRequired, emailRequired, passwordTooShort, usernameExists, emailExists}) {
if abortUserLogicError(c, err, "", nil, []string{errs.UsernameRequired, errs.EmailRequired, errs.PasswordTooShort, errs.UsernameExists, errs.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"`
c.JSON(http.StatusOK, response.OK(service.ToUserResponse(newUser)))
}
// UpdateUser 更新用户信息
@@ -353,7 +249,7 @@ type updateUserRequest struct {
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Param request body updateUserRequest true "更新参数"
// @Param request body model.UpdateUserRequest true "更新参数"
// @Success 200 {object} response.Any{data=string} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
@@ -362,7 +258,7 @@ type updateUserRequest struct {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id} [put]
func UpdateUser(c *gin.Context) {
var req updateUserRequest
var req model.UpdateUserRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
@@ -375,17 +271,11 @@ func UpdateUser(c *gin.Context) {
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if currUser == nil {
response.AbortUnauthorized(c, AdminRequired)
response.AbortUnauthorized(c, errs.AdminRequired)
return
}
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
}
err := userSvc.AdminUpdateUser(c.Request.Context(), currUser.ID, contracts.AdminUpdateUserRequest{
err := service.AdminUpdateUser(c.Request.Context(), currUser.ID, contracts.AdminUpdateUserRequest{
ID: id,
Nickname: req.Nickname,
Email: req.Email,
@@ -393,10 +283,10 @@ func UpdateUser(c *gin.Context) {
Password: req.Password,
})
if err != nil {
if abortUserLogicError(c, err, userNotFound, []string{cannotRevokeSelfAdmin}, []string{emailRequired, emailExists, passwordTooShort}) {
if abortUserLogicError(c, err, errs.UserNotFound, []string{errs.CannotRevokeSelfAdmin}, []string{errs.EmailRequired, errs.EmailExists, errs.PasswordTooShort}) {
return
}
response.AbortInternal(c, updateUserInfoFailed)
response.AbortInternal(c, errs.UpdateUserInfoFailed)
return
}
@@ -1,130 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"Wavelet/core/contracts"
"Wavelet/pkg/response"
"net/http"
"strconv"
"github.com/gin-gonic/gin"
)
// ListAuthSources lists all configured authentication sources.
func ListAuthSources(c *gin.Context) {
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
}
views, err := authSvc.ListAuthSources(c.Request.Context())
if err != nil {
response.AbortInternal(c, "获取认证源列表失败")
return
}
c.JSON(http.StatusOK, response.OK(views))
}
// CreateAuthSource creates a new authentication source.
func CreateAuthSource(c *gin.Context) {
var source contracts.AuthSourceDTO
if err := c.ShouldBindJSON(&source); err != nil {
response.AbortBadRequest(c, "无效的参数")
return
}
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
}
created, err := authSvc.CreateAuthSource(c.Request.Context(), source)
if err != nil {
response.AbortBadRequest(c, "创建认证源失败: "+err.Error())
return
}
c.JSON(http.StatusOK, response.OK(created))
}
// UpdateAuthSource updates an authentication source.
func UpdateAuthSource(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.AbortBadRequest(c, "无效的认证源 ID")
return
}
var req contracts.AuthSourceDTO
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, "无效的参数")
return
}
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
}
updated, err := authSvc.UpdateAuthSource(c.Request.Context(), id, req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(updated))
}
// ToggleAuthSource toggles the active state of an auth source.
func ToggleAuthSource(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.AbortBadRequest(c, "无效的认证源 ID")
return
}
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
}
toggled, err := authSvc.ToggleAuthSource(c.Request.Context(), id)
if err != nil {
response.AbortInternal(c, "切换认证源状态失败: "+err.Error())
return
}
c.JSON(http.StatusOK, response.OK(gin.H{"is_active": toggled.IsActive}))
}
// DeleteAuthSource deletes an authentication source.
func DeleteAuthSource(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.AbortBadRequest(c, "无效的认证源 ID")
return
}
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
}
if err := authSvc.DeleteAuthSource(c.Request.Context(), id); err != nil {
response.AbortInternal(c, "删除认证源失败: "+err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -1,534 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
mail "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 := ListVisibleSystemConfigs(ctx)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
resp := make(map[string]string, len(configs))
for _, config := range configs {
resp[config.Key] = config.Value
}
c.JSON(http.StatusOK, response.OK(resp))
}
// GetRobotsTXT 动态生成 robots.txt
// @Summary 获取 robots.txt
// @Description 根据系统配置决定是否允许搜索引擎检索,并返回相应的 robots.txt 文件内容
// @Tags config
// @Produce text/plain
// @Success 200 {string} string "robots.txt 内容"
// @Router /robots.txt [get]
func GetRobotsTXT(c *gin.Context) {
ctx := c.Request.Context()
enabled, err := GetBoolByKey(ctx, ConfigKeySearchEngineIndexingEnabled)
content := "User-Agent: *\nDisallow: /\n"
if err == nil && enabled {
content = "User-Agent: *\nAllow: /\n"
}
c.Data(http.StatusOK, "text/plain; charset=utf-8", []byte(content))
}
// CreateSystemConfig 创建系统配置
// @Summary 创建系统配置
// @Description 创建一条新的系统配置项,配置键不可重复,同时将新配置同步到 Redis,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body CreateSystemConfigRequest true "创建请求参数"
// @Success 200 {object} response.Any{data=string} "创建成功"
// @Failure 400 {object} response.Any "参数错误或配置键已存在"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs [post]
func CreateSystemConfig(c *gin.Context) {
var req CreateSystemConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if isProtectedConfigKey(req.Key) {
response.AbortBadRequest(c, protectedConfigKeyMessage)
return
}
if err := createSystemConfig(c.Request.Context(), req); err != nil {
if err.Error() == ConfigKeyExists {
response.AbortBadRequest(c, ConfigKeyExists)
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// ListSystemConfigs 获取系统配置列表
// @Summary 获取系统配置列表
// @Description 返回所有系统配置列表,支持按配置类型(system/business)过滤,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param type query string false "配置类型(system/business)"
// @Success 200 {object} response.Any{data=[]SystemConfig} "系统配置列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs [get]
func ListSystemConfigs(c *gin.Context) {
configs, err := listSystemConfigs(c.Request.Context(), c.Query("type"))
if err != nil {
response.AbortInternal(c, err.Error())
return
}
for i := range configs {
configs[i].Value = maskSensitiveConfig(configs[i].Key, configs[i].Value)
}
c.JSON(http.StatusOK, response.OK(configs))
}
// GetSystemConfig 获取单个系统配置
// @Summary 获取单个系统配置
// @Description 根据配置键获取对应的系统配置详情,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param key path string true "配置键"
// @Success 200 {object} response.Any{data=SystemConfig} "系统配置详情"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "配置不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs/{key} [get]
func GetSystemConfig(c *gin.Context) {
config, err := getSystemConfig(c.Request.Context(), c.Param("key"))
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, SystemConfigNotFound)
} else {
response.AbortInternal(c, err.Error())
}
return
}
config.Value = maskSensitiveConfig(config.Key, config.Value)
c.JSON(http.StatusOK, response.OK(config))
}
// UpdateSystemConfig 更新系统配置
// @Summary 更新系统配置
// @Description 根据配置键更新对应的配置内容,同时将更新同步到 Redis,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param key path string true "配置键"
// @Param request body UpdateSystemConfigRequest true "更新请求参数"
// @Success 200 {object} response.Any{data=string} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "配置不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs/{key} [put]
func UpdateSystemConfig(c *gin.Context) {
var req UpdateSystemConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
key := c.Param("key")
if isProtectedConfigKey(key) {
response.AbortBadRequest(c, protectedConfigKeyMessage)
return
}
if err := updateSystemConfig(c.Request.Context(), key, req); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, SystemConfigNotFound)
return
}
if isStorageConfigValidationError(err) {
response.AbortBadRequest(c, err.Error())
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
func isProtectedConfigKey(key string) bool {
return key == ConfigKeyLogDatabase || key == ConfigKeyLogDBMigration
}
func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) error {
if isProtectedConfigKey(req.Key) {
return errors.New(protectedConfigKeyMessage)
}
exists, err := SystemConfigExists(ctx, req.Key)
if err != nil {
return err
}
if exists {
return errors.New(ConfigKeyExists)
}
config := SystemConfig{
Key: req.Key,
Value: req.Value,
Type: req.Type,
Visibility: req.Visibility,
Description: req.Description,
}
if err := CreateSystemConfigRecord(ctx, &config); err != nil {
return err
}
invalidateSystemConfigCaches(ctx, req.Key)
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
return nil
}
func listSystemConfigs(ctx context.Context, configType string) ([]SystemConfig, error) {
return ListAdminSystemConfigs(ctx, configType)
}
func getSystemConfig(ctx context.Context, key string) (SystemConfig, error) {
return GetAdminSystemConfigByKey(ctx, key)
}
func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigRequest) error {
if isProtectedConfigKey(key) {
return errors.New(protectedConfigKeyMessage)
}
config, err := GetAdminSystemConfigByKey(ctx, key)
if err != nil {
return err
}
var originalDriver contracts.StorageDriver
if key == ConfigKeyStorageConfig {
var currentCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(config.Value), &currentCfg); err == nil {
originalDriver = currentCfg.Driver
}
validatedVal, err := validateAndMergeStorageConfig(ctx, req.Value, config.Value)
if err != nil {
return err
}
req.Value = validatedVal
}
gormDB := GetDB(ctx)
if gormDB == nil {
return errors.New("database service not available")
}
if err := gormDB.Transaction(func(tx *gorm.DB) error {
updates := map[string]any{
"description": req.Description,
}
if req.Visibility != nil {
updates["visibility"] = *req.Visibility
config.Visibility = *req.Visibility
}
if key != ConfigKeySMTPPassword || req.Value != maskedConfigValue {
updates["value"] = req.Value
config.Value = req.Value
}
if err := tx.Model(&config).Updates(updates).Error; err != nil {
return err
}
resolveStorageMigrationTasksOnDirectDriverUpdate(ctx, tx, key, originalDriver, req.Value)
return nil
}); err != nil {
return err
}
invalidateCachesAfterConfigUpdate(ctx, key)
return nil
}
func resolveStorageMigrationTasksOnDirectDriverUpdate(
ctx context.Context,
tx *gorm.DB,
key string,
originalDriver contracts.StorageDriver,
newValue string,
) {
if key != ConfigKeyStorageConfig || originalDriver == "" {
return
}
var newCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil {
return
}
if newCfg.Driver != originalDriver {
return
}
if err := MarkFailedTaskExecutionsSucceededTx(
tx,
"storage:migrate",
"存储配置直接更新,故障迁移任务自动标记为已解决",
time.Now(),
); err != nil {
logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err)
}
}
func invalidateSystemConfigCaches(ctx context.Context, key string) {
if err := InvalidateSystemConfigCache(ctx, key); err != nil {
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
}
_ = EmitEvent(ctx, contracts.EventTopicConfigChanged, contracts.ConfigChangedEvent{Key: key})
}
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
invalidateSystemConfigCaches(ctx, key)
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
}
// TestSMTPRequest 测试 SMTP 配置请求
type TestSMTPRequest struct {
SMTPHost string `json:"smtp_host" binding:"required,max=255"`
SMTPPort int `json:"smtp_port" binding:"required"`
SMTPUsername string `json:"smtp_username" binding:"required,max=255"`
SMTPPassword string `json:"smtp_password" binding:"required,max=255"`
To string `json:"to" binding:"required,email"`
}
// TestSMTPResponse 测试 SMTP 配置响应
type TestSMTPResponse struct {
Success bool `json:"success"`
Log string `json:"log"`
Error string `json:"error"`
}
// TestSMTP 测试 SMTP 邮件发送
// @Summary 测试 SMTP 邮件发送
// @Description 使用传入的配置进行 SMTP 邮件发送测试,支持使用 ****** 占位符使用保存的数据库密码
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body TestSMTPRequest true "测试请求参数"
// @Success 200 {object} response.Any{data=TestSMTPResponse} "测试执行完毕"
// @Failure 400 {object} response.Any "参数错误"
// @Router /api/v1/admin/system-configs/smtp/test [post]
func TestSMTP(c *gin.Context) {
var req TestSMTPRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
password := req.SMTPPassword
if password == maskedConfigValue {
if sc, err := GetSystemConfigByKey(c.Request.Context(), ConfigKeySMTPPassword); err == nil {
password = sc.Value
}
}
cfg := mail.Config{
Host: req.SMTPHost,
Port: req.SMTPPort,
Username: req.SMTPUsername,
Password: password,
}
subject := "Wavelet SMTP Test Mail"
body := `<h3>SMTP Mail Connection Test</h3>
<p>If you received this message, your SMTP configuration is correct and mail sending is working properly.</p>
<p>Sent from Wavelet.</p>`
logs, err := mail.SendMailWithLog(c.Request.Context(), cfg, req.To, subject, body)
resp := TestSMTPResponse{
Success: err == nil,
Log: logs,
}
if err != nil {
resp.Error = err.Error()
}
c.JSON(http.StatusOK, response.OK(resp))
}
func isStorageConfigValidationError(err error) bool {
msg := err.Error()
return msg == StorageDriverSwitchRequiresMigration ||
strings.HasPrefix(msg, "解析") ||
strings.HasPrefix(msg, "验证") ||
strings.HasPrefix(msg, "初始化测试") ||
strings.HasPrefix(msg, "存储连通性") ||
strings.HasPrefix(msg, "序列化") ||
strings.HasPrefix(msg, "检查存量文件")
}
func maskSensitiveConfig(key, value string) string {
if value == "" {
return value
}
switch key {
case ConfigKeySMTPPassword:
return maskedConfigValue
case ConfigKeyStorageConfig:
return maskStorageConfig(value)
}
return value
}
func maskStorageConfig(value string) string {
var cfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(value), &cfg); err != nil {
return value
}
if cfg.S3.SecretAccessKey != "" {
cfg.S3.SecretAccessKey = maskedConfigValue
}
if cfg.R2.SecretAccessKey != "" {
cfg.R2.SecretAccessKey = maskedConfigValue
}
if cfg.MinIO.SecretAccessKey != "" {
cfg.MinIO.SecretAccessKey = maskedConfigValue
}
if cfg.OSS.SecretAccessKey != "" {
cfg.OSS.SecretAccessKey = maskedConfigValue
}
if cfg.WebDAV.Password != "" {
cfg.WebDAV.Password = maskedConfigValue
}
val, err := json.Marshal(cfg)
if err != nil {
return value
}
return string(val)
}
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
// and tests connectivity of the new storage configuration.
func validateAndMergeStorageConfig(ctx context.Context, value, currentConfig string) (string, error) {
var currentCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(currentConfig), &currentCfg); err != nil {
return "", fmt.Errorf("解析当前存储配置失败: %w", err)
}
var newCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(value), &newCfg); err != nil {
return "", fmt.Errorf("解析目标存储配置失败: %w", err)
}
// 合并被掩码屏蔽的敏感信息,获取完整的真实配置
targetCfg := newCfg
if targetCfg.S3.SecretAccessKey == maskedConfigValue {
targetCfg.S3.SecretAccessKey = currentCfg.S3.SecretAccessKey
}
if targetCfg.R2.SecretAccessKey == maskedConfigValue {
targetCfg.R2.SecretAccessKey = currentCfg.R2.SecretAccessKey
}
if targetCfg.MinIO.SecretAccessKey == maskedConfigValue {
targetCfg.MinIO.SecretAccessKey = currentCfg.MinIO.SecretAccessKey
}
if targetCfg.OSS.SecretAccessKey == maskedConfigValue {
targetCfg.OSS.SecretAccessKey = currentCfg.OSS.SecretAccessKey
}
if targetCfg.WebDAV.Password == maskedConfigValue {
targetCfg.WebDAV.Password = currentCfg.WebDAV.Password
}
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, _ contracts.StorageConfigDTO) error {
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
var uploadCount int64
gormDB := GetDB(ctx)
if gormDB != nil {
if err := gormDB.Table("w_uploads").
Where("status != ?", "deleted").
Count(&uploadCount).Error; err != nil {
return fmt.Errorf("检查存量文件失败: %w", err)
}
}
if uploadCount > 0 {
return errors.New(StorageDriverSwitchRequiresMigration)
}
}
return nil
}
-643
View File
@@ -1,643 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"Wavelet/pkg/config"
"Wavelet/pkg/response"
"context"
"database/sql"
"fmt"
"log"
"math"
"net/http"
"os"
"os/exec"
"strings"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
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: logDBNameSQLite,
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 := GetDB(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 := GetDB(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 := GetDB(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 := GetDB(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: logDBNameSQLite,
Name: config.Config.Database.SQLitePath,
Version: "SQLite",
}
if info.Name == "" {
info.Name = "./data/wavelet.db"
}
gormDB := GetDB(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 := GetDB(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)
}
}
@@ -1,612 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"Wavelet/core/contracts"
"Wavelet/pkg/config"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/url"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
)
const (
defaultLimit = 200
maxLimit = 500
maxPageSize = 100
hoursInDay = 24
analyticsDays = 7
topActiveLimit = 10
)
// logsResponse 历史日志查询响应
type logsResponse struct {
Lines []logger.LogEntry `json:"lines"`
HasMore bool `json:"has_more"`
NextCursor int `json:"next_cursor"` // 用于加载更早日志的 cursor
}
// GetLogs 获取历史日志
// @Summary 获取系统日志
// @Description 分页获取系统历史日志,cursor=0 获取最新日志,cursor>0 获取更早日志
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param cursor query int false "日志游标,0=获取最新" default(0)
// @Param limit query int false "每页条数" default(200)
// @Success 200 {object} response.Any{data=logsResponse} "日志列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/logs [get]
func GetLogs(c *gin.Context) {
cursorStr := c.DefaultQuery("cursor", "0")
limitStr := c.DefaultQuery("limit", "200")
var cursor, limit int
if _, err := parsePositiveInt(cursorStr, &cursor); err != nil {
response.AbortWithError(c, http.StatusBadRequest, InvalidCursorParam)
return
}
if _, err := parsePositiveInt(limitStr, &limit); err != nil || limit <= 0 {
limit = defaultLimit
}
if limit > maxLimit {
limit = maxLimit
}
entries, hasMore := logger.GlobalRingBuffer.Query(cursor, limit)
resp := logsResponse{
Lines: entries,
HasMore: hasMore,
}
if len(entries) > 0 {
resp.NextCursor = entries[0].Index
}
c.JSON(http.StatusOK, response.OK(resp))
}
// wsMessage WebSocket 消息格式
type wsMessage struct {
Type string `json:"type"` // "log" | "error"
Data json.RawMessage `json:"data"`
}
// HandleLogWebSocket WebSocket 端点,实时推送系统日志
// @Summary 系统日志实时推送
// @Description 通过 WebSocket 实时推送系统日志,需要管理员权限
// @Tags admin
// @Router /api/v1/admin/logs/ws [get]
func HandleLogWebSocket(c *gin.Context) {
upgrader := getUpgrader()
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
return
}
defer func() { _ = conn.Close() }()
ch := logger.GlobalRingBuffer.Subscribe()
defer logger.GlobalRingBuffer.Unsubscribe(ch)
done := make(chan struct{})
util.Go(func() {
defer close(done)
for {
_, _, err := conn.ReadMessage()
if err != nil {
return
}
}
})
for {
select {
case <-done:
return
case entry, ok := <-ch:
if !ok {
return
}
data, _ := json.Marshal(entry)
msg := wsMessage{Type: "log", Data: data}
payload, _ := json.Marshal(msg)
if err := conn.WriteMessage(1, payload); err != nil {
return
}
}
}
}
// accessLogItem 访问日志单条数据
type accessLogItem struct {
ID uint64 `json:"id,string"`
TraceID string `json:"trace_id"`
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"`
}
const userQueryMaxLimit = 100
func findUserIDsByUsername(ctx context.Context, username string) ([]uint64, error) {
if userSvc := GetUserService(ctx); userSvc != nil {
users, _, err := userSvc.ListUsers(ctx, 1, userQueryMaxLimit, username)
if err != nil {
return nil, fmt.Errorf("查询用户信息失败: %w", err)
}
ids := make([]uint64, 0, len(users))
for _, u := range users {
ids = append(ids, u.ID)
}
return ids, nil
}
gormDB := GetDB(ctx)
if gormDB == nil {
return nil, nil
}
var ids []uint64
if err := gormDB.Table("w_users").
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
Pluck("id", &ids).Error; err != nil {
return nil, fmt.Errorf("查询用户信息失败: %w", err)
}
return ids, nil
}
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (contracts.AccessLogFilterDTO, error) {
filter := contracts.AccessLogFilterDTO{}
username := c.Query("username")
if username != "" {
userIDs, err := findUserIDsByUsername(ctx, username)
if err != nil {
return filter, 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 userSvc := GetUserService(ctx); userSvc != nil {
for _, uid := range userIDs {
if u, err := userSvc.GetUserByID(ctx, uid); err == nil && u != nil {
userMap[uid] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
}
}
} else if gormDB := GetDB(ctx); gormDB != nil {
var users []struct {
ID uint64
Username string
Nickname string
}
if err := gormDB.Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err == nil {
for _, u := range users {
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
}
}
}
for i := range list {
if info, ok := userMap[list[i].UserID]; ok {
list[i].Username = info.Username
list[i].Nickname = info.Nickname
}
}
}
// GetAccessLogs 获取 ClickHouse 异步采集的访问日志
// @Summary 获取用户访问日志
// @Description 分页并按照用户、接口路径、时间范围等维度检索用户访问日志列表(需要管理员权限)
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param page query int false "页码" default(1)
// @Param page_size query int false "每页条数" default(20)
// @Param username query string false "用户名模糊搜索"
// @Param path query string false "接口路径模糊搜索"
// @Param start_time query string false "起始时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)"
// @Param end_time query string false "结束时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)"
// @Success 200 {object} response.Any{data=accessLogsResponse} "访问日志列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/logs/access [get]
func GetAccessLogs(c *gin.Context) {
ctx := c.Request.Context()
rc := GetRiskControlService()
if rc == 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 := rc.QueryAccessLogs(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,
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()
rc := GetRiskControlService()
if rc == nil {
response.AbortInternal(c, "日志存储服务未初始化")
return
}
stats, err := rc.QueryAccessLogStats(ctx, analyticsDays)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, "查询访问趋势失败: "+err.Error())
return
}
trendList := make([]trendItem, len(stats))
for i, st := range stats {
trendList[i] = trendItem{
Date: st.Date,
Count: st.PV,
}
}
browserList := []browserItem{}
topUsers := []topUserItem{}
c.JSON(http.StatusOK, response.OK(logsAnalyticsResponse{
Trend: trendList,
Browsers: browserList,
TopUsers: topUsers,
}))
}
func getUpgrader() *websocket.Upgrader {
return &websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool {
origin := r.Header.Get("Origin")
if origin == "" {
return true
}
// 1. 同源检查 (Same-origin check)
u, err := url.Parse(origin)
if err == nil && strings.EqualFold(u.Host, r.Host) {
return true
}
// 2. 检查配置的允许跨域 Origin (Check allowed origins in system config)
ctx := r.Context()
if sc, err := GetSystemConfigByKey(ctx, ConfigKeyServerAddress); err == nil && sc.Value != "" {
originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/")
allowedOrigins := strings.Split(sc.Value, ",")
for _, allowed := range allowedOrigins {
allowed = strings.TrimRight(strings.TrimSpace(allowed), "/")
if allowed != "" && strings.EqualFold(allowed, originToCheck) {
return true
}
}
}
return false
},
}
}
func parsePositiveInt(s string, result *int) (bool, error) {
if s == "" {
*result = 0
return true, nil
}
n, err := strconv.Atoi(s)
if err != nil || n < 0 {
return false, err
}
*result = n
return true, nil
}
// Log DB Switch Task
const (
// LogDBSwitchTask 切换日志数据库任务标识。
LogDBSwitchTask = "logs:db_switch"
// TaskTypeLogDBSwitch 管理端任务类型。
TaskTypeLogDBSwitch = "logs_db_switch"
copyBatchSize = 1000
targetPostgres = "postgres"
targetSQLite = "sqlite"
targetClickHouse = "clickhouse"
)
// LogDBSwitchMeta 描述切换日志数据库任务。
var LogDBSwitchMeta = contracts.TaskMetaDTO{
Name: LogDBSwitchTask,
DisplayName: "切换日志数据库",
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
MaxRetry: 3,
Queue: "default",
Params: []contracts.TaskParamDTO{
{Name: "target", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse", Type: "string", Required: true},
},
}
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) (*contracts.TaskResultDTO, 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 {
return nil, err
}
taskSvc := GetTaskService()
if taskSvc != nil {
taskSvc.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)
}
}()
rc := GetRiskControlService()
if rc != nil {
if err := rc.SwitchLogEngine(ctx, p.Target); err != nil {
return nil, err
}
}
if err := flipLogDatabase(ctx, p.Target); err != nil {
return nil, err
}
if taskSvc != nil {
taskSvc.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
}
return &contracts.TaskResultDTO{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
}
func validateSwitch(ctx context.Context, target string) error {
source, err := currentLogDatabase(ctx)
if err != nil {
return err
}
if source == target {
return errors.New("目标日志库与当前日志库相同,无需迁移")
}
switch target {
case targetClickHouse:
if !config.Config.ClickHouse.Enabled {
return errors.New("ClickHouse 未启用,无法迁移到 ClickHouse")
}
case targetPostgres:
if !config.Config.Database.Enabled {
return errors.New("PostgreSQL 未启用(当前主库为 SQLite),无法迁移到 PostgreSQL")
}
case targetSQLite:
if config.Config.Database.Enabled {
return errors.New("当前主库为 PostgreSQL,日志库不能设置为 SQLite")
}
}
return nil
}
func currentLogDatabase(ctx context.Context) (string, error) {
cfg, err := GetSystemConfigByKey(ctx, ConfigKeyLogDatabase)
if err != nil {
return "", fmt.Errorf("读取日志主库失败: %w", err)
}
if cfg.Value == "" {
return "", errors.New("日志主库配置为空")
}
return cfg.Value, nil
}
func setMigrationFlag(ctx context.Context, v string) error {
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDBMigration, v)
}
func flipLogDatabase(ctx context.Context, target string) error {
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDatabase, target)
}
@@ -1,223 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"Wavelet/pkg/config"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"context"
"errors"
"fmt"
"math"
"net/http"
"runtime"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
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()
activeDB := logDBNameSQLite
migration := "idle"
if rc := GetRiskControlService(); rc != nil {
activeDB = rc.ActiveLogEngine(ctx)
if rc.IsLogEngineMigrating(ctx) {
migration = "migrating"
}
}
c.JSON(http.StatusOK, response.OK(LogDatabaseStatus{
ActiveDatabase: activeDB,
Migration: migration,
RetentionDays: map[string]int{
logDBNamePostgres: retentionOr(ctx, ConfigKeyLogRetentionDaysPostgres),
logDBNameSQLite: retentionOr(ctx, ConfigKeyLogRetentionDaysSQLite),
logDBNameClickHouse: retentionOr(ctx, ConfigKeyLogRetentionDaysClickHouse),
},
AvailableTargets: availableLogTargets(activeDB),
}))
}
func retentionOr(ctx context.Context, key string) int {
v, err := GetIntByKey(ctx, key)
if err != nil {
if !errors.Is(err, gorm.ErrRecordNotFound) {
logger.ErrorF(ctx, "读取日志保留天数配置失败 key=%s: %v", key, err)
}
return defaultLogRetentionDays
}
if v < 1 {
return defaultLogRetentionDays
}
return v
}
func availableLogTargets(active string) []string {
if active == logDBNameClickHouse {
if config.Config.Database.Enabled {
return []string{logDBNamePostgres}
}
return []string{logDBNameSQLite}
}
if config.Config.ClickHouse.Enabled {
return []string{logDBNameClickHouse}
}
return []string{}
}
+382
View File
@@ -0,0 +1,382 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"Wavelet/pkg/logger"
"fmt"
"math"
"time"
)
const (
binaryKB = 0
binaryMB = 1
binaryGB = 2
valueThreshold = 10
maxStringLength = 200
)
// FormatBytes renders a byte count using binary units.
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)
}
// TruncateDisplayValue caps oversized cell values before they reach the console UI.
func TruncateDisplayValue(value string) string {
runes := []rune(value)
if len(runes) > maxStringLength {
return string(runes[:maxStringLength]) + "..."
}
return value
}
// 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"`
}
// 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"`
}
// 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"`
}
// 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"`
}
// 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"`
}
// UpdateCacheConfigRequest 磁盘缓存策略更新请求
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"`
}
// 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"`
}
// 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"`
}
// 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"`
}
// 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"`
}
// ListTaskExecutionsRequest 分页查询任务执行记录请求参数
type ListTaskExecutionsRequest struct {
Page int `form:"page"`
PageSize int `form:"page_size"`
Status string `form:"status"`
TaskType string `form:"task_type"`
TaskTypes string `form:"task_types"`
TaskTypePrefix string `form:"task_type_prefix"`
}
// 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"`
}
// UserResponse 用户资料响应
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"`
}
// UpdateUserStatusRequest 更新用户状态请求
type UpdateUserStatusRequest struct {
IsActive bool `json:"is_active"`
}
// 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"`
}
// 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"`
}
// LogsResponse 历史日志查询响应
type LogsResponse struct {
Lines []logger.LogEntry `json:"lines"`
HasMore bool `json:"has_more"`
NextCursor int `json:"next_cursor"` // 用于加载更早日志的 cursor
}
// AccessLogItem 访问日志单条数据
type AccessLogItem struct {
ID uint64 `json:"id,string"`
TraceID string `json:"trace_id"`
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"`
}
// 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"`
}
// AccessLogQuery carries the raw console filters before they become a contract filter.
type AccessLogQuery struct {
Username string
Path string
StartTime string
EndTime string
Page int
PageSize int
}
// 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"`
}
@@ -1,9 +1,11 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
// Package model contains database entities and data transfer objects for the admin domain.
package model
import (
"Wavelet/plugins/domain/admin/errs"
"bytes"
"errors"
"strings"
@@ -112,13 +114,13 @@ func (t *Template) Normalize() {
func (t *Template) Validate() error {
t.Normalize()
if t.Key == "" {
return errors.New(TemplateKeyRequired)
return errors.New(errs.TemplateKeyRequired)
}
if t.Name == "" {
return errors.New(TemplateNameRequired)
return errors.New(errs.TemplateNameRequired)
}
if t.Content == "" {
return errors.New(TemplateContentRequired)
return errors.New(errs.TemplateContentRequired)
}
return nil
}
@@ -205,16 +207,6 @@ func (TaskExecution) TableName() string {
return "w_task_executions"
}
// ListTaskExecutionsRequest 分页查询任务执行记录请求参数
type ListTaskExecutionsRequest struct {
Page int `form:"page"`
PageSize int `form:"page_size"`
Status string `form:"status"`
TaskType string `form:"task_type"`
TaskTypes string `form:"task_types"`
TaskTypePrefix string `form:"task_type_prefix"`
}
// TaskExecutionCleanupStats 任务日志清理结果统计
type TaskExecutionCleanupStats struct {
HighFrequencyDeleted int64 `json:"high_frequency_deleted"`
+25 -129
View File
@@ -8,6 +8,9 @@ import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/plugins/domain/admin/handler"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/service"
"context"
"embed"
"reflect"
@@ -16,6 +19,9 @@ import (
"github.com/hibiken/asynq"
)
// SystemConfig aliases model.SystemConfig for external compatibility.
type SystemConfig = model.SystemConfig
//go:embed migrations/*/*.sql
var adminMigrations embed.FS
@@ -65,64 +71,64 @@ func (p *Plugin) Manifest() core.Manifest {
func (p *Plugin) Apply(ctx *core.Context) error {
// 0. Bind Services reactively
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
SetDBService(db)
service.SetDBService(db)
} else {
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
SetDBService(db)
service.SetDBService(db)
})
}
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
SetCacheService(cache)
service.SetCacheService(cache)
} else {
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
SetCacheService(cache)
service.SetCacheService(cache)
})
}
if user, err := core.Inject[contracts.UserService](ctx); err == nil && user != nil {
SetUserService(user)
service.SetUserService(user)
} else {
core.When[contracts.UserService](ctx, func(user contracts.UserService) {
SetUserService(user)
service.SetUserService(user)
})
}
if auth, err := core.Inject[contracts.AuthService](ctx); err == nil && auth != nil {
SetAuthService(auth)
service.SetAuthService(auth)
} else {
core.When[contracts.AuthService](ctx, func(auth contracts.AuthService) {
SetAuthService(auth)
service.SetAuthService(auth)
})
}
if task, err := core.Inject[contracts.TaskService](ctx); err == nil && task != nil {
SetTaskService(task)
service.SetTaskService(task)
} else {
core.When[contracts.TaskService](ctx, func(task contracts.TaskService) {
SetTaskService(task)
service.SetTaskService(task)
})
}
if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil {
SetStorageService(storage)
service.SetStorageService(storage)
} else {
core.When[contracts.StorageService](ctx, func(storage contracts.StorageService) {
SetStorageService(storage)
service.SetStorageService(storage)
})
}
if rc, err := core.Inject[contracts.RiskControlService](ctx); err == nil && rc != nil {
SetRiskControlService(rc)
service.SetRiskControlService(rc)
} else {
core.When[contracts.RiskControlService](ctx, func(rc contracts.RiskControlService) {
SetRiskControlService(rc)
service.SetRiskControlService(rc)
})
}
SetEventEmitter(ctx.Events().Emit)
service.SetEventEmitter(ctx.Events().Emit)
ctx.OnDispose(func() error {
ResetServices()
service.ResetServices()
return nil
})
// 0a. Dynamic Auth Middlewares
var loginMW gin.HandlerFunc = func(c *gin.Context) {
if authSvc := GetAuthService(c.Request.Context()); authSvc != nil {
if authSvc := service.GetAuthService(c.Request.Context()); authSvc != nil {
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
mw(c)
return
@@ -131,7 +137,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
c.Next()
}
var adminMW gin.HandlerFunc = func(c *gin.Context) {
if authSvc := GetAuthService(c.Request.Context()); authSvc != nil {
if authSvc := service.GetAuthService(c.Request.Context()); authSvc != nil {
if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok {
mw(c)
return
@@ -145,117 +151,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
// 1. Register Admin HTTP Routes
adminRouter := ctx.Router().Group("/api/v1/admin", loginMW, adminMW)
{
// Status & Diagnostics
adminRouter.GET("/status", GetSystemStatus)
adminRouter.GET("/status/log-database", GetLogDatabaseStatus)
adminRouter.GET("/db-info", GetDatabaseInfo)
adminRouter.GET("/db-export", ExportDatabase)
// DB Management
dbGroup := adminRouter.Group("/db-manage")
{
dbGroup.GET("/overview", GetDBOverview)
dbGroup.GET("/tables", ListDBTables)
dbGroup.GET("/table-data", GetDBTableData)
dbGroup.POST("/query", ExecuteSQL)
}
// Cache Management
cacheGroup := adminRouter.Group("/cache")
{
cacheGroup.GET("/status", GetCacheStatus)
cacheGroup.POST("/config", UpdateCacheConfig)
cacheGroup.POST("/clear", ClearCache)
}
// Updater
updateGroup := adminRouter.Group("/update")
{
updateGroup.GET("", GetUpdateStatus)
updateGroup.POST("/apply", ApplyUpdate)
}
// Logs
logsGroup := adminRouter.Group("/logs")
{
logsGroup.GET("", GetLogs)
logsGroup.GET("/access", GetAccessLogs)
logsGroup.GET("/analytics", GetLogsAnalytics)
logsGroup.GET("/ws", HandleLogWebSocket)
}
// Users
usersGroup := adminRouter.Group("/users")
{
usersGroup.GET("", ListUsers)
usersGroup.POST("", CreateUser)
usersGroup.GET("/:id", GetUser)
usersGroup.PUT("/:id/status", UpdateUserStatus)
usersGroup.PUT("/:id", UpdateUser)
usersGroup.DELETE("/:id", DeleteUser)
}
// Auth Sources
authSourcesGroup := adminRouter.Group("/auth-sources")
{
authSourcesGroup.GET("", ListAuthSources)
authSourcesGroup.POST("", CreateAuthSource)
authSourcesGroup.PUT("/:id", UpdateAuthSource)
authSourcesGroup.PUT("/:id/toggle", ToggleAuthSource)
authSourcesGroup.DELETE("/:id", DeleteAuthSource)
}
// System Configs
configGroup := adminRouter.Group("/system-configs")
{
configGroup.GET("", ListSystemConfigs)
configGroup.POST("", CreateSystemConfig)
configGroup.POST("/smtp/test", TestSMTP)
keyGroup := configGroup.Group("/:key")
{
keyGroup.GET("", GetSystemConfig)
keyGroup.PUT("", UpdateSystemConfig)
}
}
// Templates
templateGroup := adminRouter.Group("/templates")
{
templateGroup.GET("", ListTemplates)
templateGroup.POST("", CreateTemplate)
keyGroup := templateGroup.Group("/:key")
{
keyGroup.GET("", GetTemplate)
keyGroup.PUT("", UpdateTemplate)
keyGroup.DELETE("", DeleteTemplate)
}
}
// Tasks
taskGroup := adminRouter.Group("/tasks")
{
taskGroup.GET("/types", ListTaskTypes)
taskGroup.POST("/dispatch", DispatchTask)
executions := taskGroup.Group("/executions")
{
executions.GET("", ListTaskExecutions)
executions.GET("/:id", GetTaskExecution)
executions.POST("/:id/retry", RetryTask)
}
schedules := taskGroup.Group("/schedules")
{
schedules.GET("", ListSchedules)
schedules.POST("", CreateSchedule)
schedules.PUT("/:id", UpdateSchedule)
schedules.DELETE("/:id", DeleteSchedule)
}
}
}
handler.RegisterRoutes(adminRouter)
// 2. Register Background Tasks
ctx.Task().Register("admin:system_cleanup", func(_ context.Context, _ *asynq.Task) error {
-689
View File
@@ -1,689 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"Wavelet/core/contracts"
"Wavelet/pkg/cache/ram"
"Wavelet/pkg/idgen"
"Wavelet/pkg/util"
"context"
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/shopspring/decimal"
"gorm.io/gorm"
)
const (
configTypeSystem = "system"
errDatabaseNotInitialized = "database not initialized"
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
taskExecutionLogRedisKeyPrefix = "task:execution:log:"
taskExecutionLogExpiration = 24 * time.Hour
taskExecutionLogMaxLines = 1000
)
// PreheatSystemConfigs loads all system configs from database.
func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
database := GetDB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var configs []SystemConfig
if err := database.Find(&configs).Error; err != nil {
return nil, err
}
return configs, nil
}
// PreheatSystemConfigByKey loads a single config key from database.
func PreheatSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
database := GetDB(ctx)
if database == nil {
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
}
var sc SystemConfig
if err := database.Where("key = ?", key).First(&sc).Error; err != nil {
return SystemConfig{}, err
}
return sc, nil
}
// GetSystemConfigByGroup queries a configuration by Type and Key.
func GetSystemConfigByGroup(ctx context.Context, configType, key string) (SystemConfig, error) {
ensureSystemConfigCacheListener()
if item, ok := ram.Get(configType, key); ok {
var sc SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
return sc, nil
}
}
database := GetDB(ctx)
if database == nil {
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
}
var sc SystemConfig
if err := database.Where("key = ?", key).First(&sc).Error; err != nil {
return SystemConfig{}, err
}
valBytes, err := json.Marshal(sc)
if err == nil {
ram.Set(ram.CacheItem{
Key: sc.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(sc.Key),
})
}
return sc, nil
}
// GetSystemConfigByKey queries config by key.
func GetSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
return GetSystemConfigByGroup(ctx, ConfigCacheType, key)
}
// ListSystemConfigsByKeys loads multiple config keys.
func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]SystemConfig, error) {
if len(keys) == 0 {
return map[string]SystemConfig{}, nil
}
ensureSystemConfigCacheListener()
result := make(map[string]SystemConfig, len(keys))
missing := make([]string, 0, len(keys))
for _, key := range keys {
if item, ok := ram.Get(ConfigCacheType, key); ok {
var sc SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
result[key] = sc
continue
}
}
missing = append(missing, key)
}
if len(missing) == 0 {
return result, nil
}
database := GetDB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var configs []SystemConfig
if err := database.Where("key IN ?", missing).Find(&configs).Error; err != nil {
return nil, err
}
for i := range configs {
valBytes, err := json.Marshal(configs[i])
if err == nil {
ram.Set(ram.CacheItem{
Key: configs[i].Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(configs[i].Key),
})
}
result[configs[i].Key] = configs[i]
}
return result, nil
}
// InvalidateVisibleSystemConfigsCache clears the cached public config list.
func InvalidateVisibleSystemConfigsCache(ctx context.Context) error {
return InvalidateAllSystemConfigCaches(ctx)
}
// ListVisibleSystemConfigs queries visible configs using local cache store.
func ListVisibleSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
ensureSystemConfigCacheListener()
items := ram.GetTypeItems(ConfigCacheType)
if len(items) > 0 {
var list []SystemConfig
for _, item := range items {
var sc SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
if sc.Visibility == ConfigVisibilityVisible {
list = append(list, sc)
}
}
}
return list, nil
}
database := GetDB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var configs []SystemConfig
if err := database.Where("visibility = ?", ConfigVisibilityVisible).Find(&configs).Error; err != nil {
return nil, err
}
for _, cfg := range configs {
valBytes, err := json.Marshal(cfg)
if err == nil {
ram.Set(ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(cfg.Key),
})
}
}
return configs, nil
}
// GetIntByKey queries config and converts to int.
func GetIntByKey(ctx context.Context, key string) (int, error) {
sc, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return 0, err
}
value, err := strconv.Atoi(sc.Value)
if err != nil {
return 0, fmt.Errorf(errConfigIntParseFailed, key, sc.Value, err)
}
return value, nil
}
// GetDecimalByKey queries config and converts to decimal.Decimal.
func GetDecimalByKey(ctx context.Context, key string, precision int32) (decimal.Decimal, error) {
sc, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return decimal.Zero, err
}
value, err := decimal.NewFromString(sc.Value)
if err != nil {
return decimal.Zero, fmt.Errorf(errConfigDecimalParseFailed, key, sc.Value, err)
}
return value.Truncate(precision), nil
}
// GetBoolByKey queries config and converts to bool.
func GetBoolByKey(ctx context.Context, key string) (bool, error) {
sc, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return false, err
}
value, err := strconv.ParseBool(sc.Value)
if err != nil {
return false, fmt.Errorf(errConfigBoolParseFailed, key, sc.Value, err)
}
return value, nil
}
// GetMenuDisplayConfig queries and parses menu config.
func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
sc, err := GetSystemConfigByKey(ctx, ConfigKeyMenuDisplayConfig)
if err != nil {
return nil, err
}
config := make(map[string]bool)
if sc.Value == "" || sc.Value == "{}" {
return config, nil
}
if err := json.Unmarshal([]byte(sc.Value), &config); err != nil {
return nil, fmt.Errorf(errParseMenuDisplayConfigFailed, err)
}
return config, nil
}
// ListAdminSystemConfigs returns all configs, optionally filtered by type.
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]SystemConfig, error) {
query := GetDB(ctx).Order("created_at DESC")
if configType != "" {
query = query.Where("type = ?", configType)
}
var configs []SystemConfig
if err := query.Find(&configs).Error; err != nil {
return nil, err
}
return configs, nil
}
// GetAdminSystemConfigByKey loads a config directly from DB.
func GetAdminSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
var config SystemConfig
if err := GetDB(ctx).Where("key = ?", key).First(&config).Error; err != nil {
return SystemConfig{}, err
}
return config, nil
}
// SystemConfigExists reports whether a config key already exists.
func SystemConfigExists(ctx context.Context, key string) (bool, error) {
var existing SystemConfig
err := GetDB(ctx).Where("key = ?", key).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
if err != nil {
return false, err
}
return true, nil
}
// CreateSystemConfigRecord persists a new system config row.
func CreateSystemConfigRecord(ctx context.Context, config *SystemConfig) error {
return GetDB(ctx).Create(config).Error
}
// UpdateSystemConfigFields applies partial updates to a system config row.
func UpdateSystemConfigFields(ctx context.Context, config *SystemConfig, updates map[string]any) error {
return GetDB(ctx).Model(config).Updates(updates).Error
}
// SaveOrUpdateSystemConfig creates or updates a config row and invalidates cache.
func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
var sc SystemConfig
err := GetDB(ctx).Where("key = ?", key).First(&sc).Error
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
if errors.Is(err, gorm.ErrRecordNotFound) {
sc = SystemConfig{
Key: key,
Value: value,
Type: configTypeSystem,
Visibility: ConfigVisibilityHidden,
}
if err := GetDB(ctx).Create(&sc).Error; err != nil {
return err
}
} else {
sc.Value = value
if err := GetDB(ctx).Save(&sc).Error; err != nil {
return err
}
}
return InvalidateSystemConfigCache(ctx, key)
}
// ListTemplatesRecord returns all templates ordered by system flag and creation time.
func ListTemplatesRecord(ctx context.Context) ([]Template, error) {
var templates []Template
if err := GetDB(ctx).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
return nil, err
}
return templates, nil
}
// GetTemplateByKey loads a template by its key.
func GetTemplateByKey(ctx context.Context, key string) (Template, error) {
var tmpl Template
if err := GetDB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil {
return Template{}, err
}
return tmpl, nil
}
// TemplateExistsByKey reports whether a template key is already taken.
func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
var existing Template
err := GetDB(ctx).Where("key = ?", key).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
if err != nil {
return false, err
}
return true, nil
}
// CreateTemplateRecord persists a new template.
func CreateTemplateRecord(ctx context.Context, tmpl *Template) error {
return GetDB(ctx).Create(tmpl).Error
}
// SaveTemplateRecord updates an existing template.
func SaveTemplateRecord(ctx context.Context, tmpl *Template) error {
return GetDB(ctx).Save(tmpl).Error
}
// DeleteTemplateRecord removes a template record.
func DeleteTemplateRecord(ctx context.Context, tmpl *Template) error {
return GetDB(ctx).Delete(tmpl).Error
}
// CreateScheduleRecord 创建定时任务
func CreateScheduleRecord(ctx context.Context, schedule *Schedule) error {
return GetDB(ctx).Create(schedule).Error
}
// UpdateScheduleRecord 更新定时任务
func UpdateScheduleRecord(ctx context.Context, schedule *Schedule) error {
return GetDB(ctx).Save(schedule).Error
}
// DeleteScheduleRecord 删除定时任务
func DeleteScheduleRecord(ctx context.Context, id uint64) error {
return GetDB(ctx).Delete(&Schedule{}, id).Error
}
// GetScheduleByID 根据 ID 获取定时任务
func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
var schedule Schedule
if err := GetDB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
return nil, err
}
return &schedule, nil
}
// ListSchedulesRecord 获取所有定时任务
func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := GetDB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
// ListActiveSchedules 获取所有启用的定时任务
func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := GetDB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
// CreateTaskExecutionRecord 创建任务执行记录
func CreateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
execution.ID = idgen.NextUint64ID()
return GetDB(ctx).Create(execution).Error
}
// UpdateTaskExecutionRecord 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
func UpdateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
return GetDB(ctx).Omit("log").Save(execution).Error
}
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) {
var execution TaskExecution
if err := GetDB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetTaskExecutionByID 根据 ID 获取执行记录
func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) {
var execution TaskExecution
if err := GetDB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetLatestTaskExecutionByTaskType returns the most recent execution for a task type.
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*TaskExecution, bool, error) {
var execution TaskExecution
err := GetDB(ctx).
Where("task_type = ?", taskType).
Order("id DESC").
First(&execution).Error
if err == nil {
if loadErr := loadTaskExecutionLog(ctx, &execution); loadErr != nil {
return nil, false, loadErr
}
return &execution, true, nil
}
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, nil
}
return nil, false, err
}
// AppendTaskExecutionLog 将日志追加到缓冲,任务完成后再持久化到数据库。
func AppendTaskExecutionLog(ctx context.Context, taskID, logLine string) error {
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return errors.New("cache service is not initialized")
}
now := time.Now().Format("15:04:05")
line := fmt.Sprintf("[%s] %s\n", now, logLine)
key := taskExecutionLogRedisKey(taskID)
var existing string
_ = cacheSvc.Get(ctx, key, &existing)
return cacheSvc.Set(ctx, key, existing+line, taskExecutionLogExpiration)
}
// FlushTaskExecutionLog 将缓冲中的完整任务日志写入数据库,并在成功后清理缓存。
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return errors.New("cache service is not initialized")
}
key := taskExecutionLogRedisKey(taskID)
var logText string
if err := cacheSvc.Get(ctx, key, &logText); err != nil {
// 缓存未命中属于正常情况(任务无输出),其余错误必须上抛,
// 否则缓冲日志会被静默丢弃并误报持久化成功。
if !errors.Is(err, contracts.ErrCacheMiss) {
return fmt.Errorf("load buffered task execution log: %w", err)
}
return nil
}
if logText == "" {
return nil
}
gormDB := GetDB(ctx)
if gormDB == nil {
return errors.New(errDatabaseNotInitialized)
}
result := gormDB.Model(&TaskExecution{}).
Where("task_id = ?", taskID).
Update("log", logText)
if result.Error != nil {
return fmt.Errorf("persist task execution log: %w", result.Error)
}
if result.RowsAffected == 0 {
return fmt.Errorf("persist task execution log: task %q not found", taskID)
}
_ = cacheSvc.Delete(ctx, key)
return nil
}
// ListTaskExecutionRecords 分页查询任务执行记录
func ListTaskExecutionRecords(ctx context.Context, req ListTaskExecutionsRequest) ([]TaskExecution, int64, error) {
if req.Page <= 0 {
req.Page = 1
}
if req.PageSize <= 0 {
req.PageSize = 20
}
query := GetDB(ctx).Model(&TaskExecution{})
if req.Status != "" {
query = query.Where("status = ?", req.Status)
}
if req.TaskType != "" {
query = query.Where("task_type = ?", req.TaskType)
} else if types := parseTaskTypesFilter(req.TaskTypes); len(types) > 0 {
query = query.Where("task_type IN ?", types)
} else if req.TaskTypePrefix != "" {
query = query.Where("task_type LIKE ? ESCAPE '\\'", util.EscapeLike(req.TaskTypePrefix)+"%")
}
var total int64
if err := query.Count(&total).Error; err != nil {
return nil, 0, err
}
var executions []TaskExecution
offset := (req.Page - 1) * req.PageSize
if err := query.Order("id DESC").Offset(offset).Limit(req.PageSize).Find(&executions).Error; err != nil {
return nil, 0, err
}
if err := loadTaskExecutionLogs(ctx, executions); err != nil {
return nil, 0, err
}
return executions, total, nil
}
func parseTaskTypesFilter(raw string) []string {
if strings.TrimSpace(raw) == "" {
return nil
}
parts := strings.Split(raw, ",")
out := make([]string, 0, len(parts))
for _, part := range parts {
part = strings.TrimSpace(part)
if part != "" {
out = append(out, part)
}
}
return out
}
// MarkFailedTaskExecutionsSucceededTx marks failed executions of a task type as succeeded within a transaction.
func MarkFailedTaskExecutionsSucceededTx(
tx *gorm.DB,
taskType string,
result string,
finishedAt time.Time,
) error {
return tx.Model(&TaskExecution{}).
Where("task_type = ? AND status = ?", taskType, TaskExecutionStatusFailed).
Updates(map[string]any{
"status": TaskExecutionStatusSucceeded,
"result": result,
"finished_at": finishedAt,
}).Error
}
// CleanupTaskExecutionLogs removes finished task execution logs according to frequency-based retention.
func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecutionCleanupStats, error) {
const (
frequencyWindowDays = 30
highFrequencyThreshold = frequencyWindowDays
)
frequencyWindowStart := now.AddDate(0, 0, -frequencyWindowDays)
highFrequencyCutoff := now.AddDate(0, 0, -3)
lowFrequencyCutoff := now.AddDate(0, 0, -30)
terminalStatuses := []TaskExecutionStatus{TaskExecutionStatusSucceeded, TaskExecutionStatusFailed}
var highFrequencyTaskTypes []string
if err := GetDB(ctx).
Model(&TaskExecution{}).
Select("task_type").
Where("created_at >= ?", frequencyWindowStart).
Group("task_type").
Having("COUNT(*) > ?", highFrequencyThreshold).
Pluck("task_type", &highFrequencyTaskTypes).Error; err != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("query high-frequency task types: %w", err)
}
var highFrequencyDeleted int64
if len(highFrequencyTaskTypes) > 0 {
highFrequencyResult := GetDB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", highFrequencyCutoff).
Where("task_type IN ?", highFrequencyTaskTypes).
Delete(&TaskExecution{})
if highFrequencyResult.Error != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("delete high-frequency task execution logs: %w", highFrequencyResult.Error)
}
highFrequencyDeleted = highFrequencyResult.RowsAffected
}
lowFrequencyQuery := GetDB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", lowFrequencyCutoff)
if len(highFrequencyTaskTypes) > 0 {
lowFrequencyQuery = lowFrequencyQuery.Where("task_type NOT IN ?", highFrequencyTaskTypes)
}
lowFrequencyResult := lowFrequencyQuery.Delete(&TaskExecution{})
if lowFrequencyResult.Error != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("delete low-frequency task execution logs: %w", lowFrequencyResult.Error)
}
return TaskExecutionCleanupStats{
HighFrequencyDeleted: highFrequencyDeleted,
LowFrequencyDeleted: lowFrequencyResult.RowsAffected,
}, nil
}
func taskExecutionLogRedisKey(taskID string) string {
return taskExecutionLogRedisKeyPrefix + taskID
}
func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return nil
}
var logText string
if err := cacheSvc.Get(ctx, taskExecutionLogRedisKey(execution.TaskID), &logText); err == nil && logText != "" {
execution.Log = logText
}
return nil
}
func loadTaskExecutionLogs(ctx context.Context, executions []TaskExecution) error {
cacheSvc := GetCache(ctx)
if cacheSvc == nil || len(executions) == 0 {
return nil
}
for i := range executions {
var logText string
if err := cacheSvc.Get(ctx, taskExecutionLogRedisKey(executions[i].TaskID), &logText); err == nil && logText != "" {
executions[i].Log = logText
}
}
return nil
}
@@ -1,10 +1,11 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
package repository
import (
"Wavelet/pkg/cache/ram"
"Wavelet/plugins/domain/admin/model"
"context"
"encoding/json"
"errors"
@@ -77,9 +78,9 @@ func (ConfigLoader) LoadOne(ctx context.Context, configType, key string) (ram.Ca
}
// GetCachedSystemConfig retrieves a single system config with RAM L1 fallback to DB.
func GetCachedSystemConfig(ctx context.Context, key string) (*SystemConfig, error) {
func GetCachedSystemConfig(ctx context.Context, key string) (*model.SystemConfig, error) {
if item, ok := ram.Get(ConfigCacheType, key); ok {
var cfg SystemConfig
var cfg model.SystemConfig
if err := json.Unmarshal([]byte(item.Value), &cfg); err == nil {
return &cfg, nil
}
@@ -0,0 +1,371 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"Wavelet/pkg/config"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"context"
"database/sql"
"errors"
"fmt"
"os"
"os/exec"
"strings"
"time"
)
const (
defaultSQLiteDBPath = "./data/wavelet.db"
logDBNameSQLite = "sqlite"
)
// sqliteDatabasePath resolves the effective SQLite file path from configuration.
func sqliteDatabasePath() string {
name := config.Config.Database.SQLitePath
if name == "" {
name = defaultSQLiteDBPath
}
return name
}
// QuoteTableName escapes a raw identifier for use inside a quoted SQL fragment.
func QuoteTableName(table string) string {
return `"` + strings.ReplaceAll(table, `"`, `""`) + `"`
}
// GetSQLiteOverview collects the SQLite runtime overview.
func GetSQLiteOverview(ctx context.Context) (model.DBOverviewResponse, error) {
gormDB := GetDB(ctx)
if gormDB == nil {
return model.DBOverviewResponse{}, errs.ErrDatabaseUninitialized
}
name := sqliteDatabasePath()
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 = model.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 model.DBOverviewResponse{
Type: logDBNameSQLite,
Version: version,
Name: name,
Size: sizeStr,
TableCount: tableCount,
Connections: connCount,
}, nil
}
// GetPostgresOverview collects the PostgreSQL runtime overview.
func GetPostgresOverview(ctx context.Context) (model.DBOverviewResponse, error) {
gormDB := GetDB(ctx)
if gormDB == nil {
return model.DBOverviewResponse{}, errs.ErrDatabaseUninitialized
}
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 = model.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 model.DBOverviewResponse{
Type: "postgres",
Version: version,
Name: name,
Size: sizeStr,
TableCount: tableCount,
Connections: connCount,
}, nil
}
// ListDatabaseTableNames returns every user table of the active database.
func ListDatabaseTableNames(ctx context.Context) ([]string, error) {
gormDB := GetDB(ctx)
if gormDB == nil {
return nil, errs.ErrDatabaseUninitialized
}
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 {
return nil, err
}
return tables, nil
}
// CountDatabaseTableRows counts the rows of the quoted table.
func CountDatabaseTableRows(ctx context.Context, quotedTable string) (int64, error) {
gormDB := GetDB(ctx)
if gormDB == nil {
return 0, errs.ErrDatabaseUninitialized
}
var total int64
if err := gormDB.Raw("SELECT count(*) FROM " + quotedTable).Scan(&total).Error; err != nil {
return 0, errs.NewInvalidInputError(err.Error())
}
return total, nil
}
// QueryDatabaseTableRows loads one page of raw rows from the quoted table.
func QueryDatabaseTableRows(
ctx context.Context,
quotedTable string,
limit int,
offset int,
) ([]string, []map[string]any, error) {
gormDB := GetDB(ctx)
if gormDB == nil {
return nil, nil, errs.ErrDatabaseUninitialized
}
rows, err := gormDB.Raw("SELECT * FROM "+quotedTable+" LIMIT ? OFFSET ?", limit, offset).Rows()
if err != nil {
return nil, nil, errs.NewInvalidInputError(err.Error())
}
defer func() {
_ = rows.Close()
}()
cols, err := rows.Columns()
if err != nil {
return nil, nil, err
}
results, err := scanTableRows(rows, cols)
if err != nil {
return nil, nil, err
}
return cols, results, nil
}
// RunSelectSQL executes an arbitrary select-like statement.
func RunSelectSQL(ctx context.Context, sqlStr string) ([]string, []map[string]any, error) {
gormDB := GetDB(ctx)
if gormDB == nil {
return nil, nil, errs.ErrDatabaseUninitialized
}
rows, err := gormDB.Raw(sqlStr).Rows()
if err != nil {
return nil, nil, errs.NewInvalidInputError(err.Error())
}
defer func() {
_ = rows.Close()
}()
cols, err := rows.Columns()
if err != nil {
return nil, nil, err
}
results, err := scanTableRows(rows, cols)
if err != nil {
return nil, nil, err
}
return cols, results, nil
}
// RunMutationSQL executes a non-query statement and reports affected rows.
func RunMutationSQL(ctx context.Context, sqlStr string) (int64, error) {
gormDB := GetDB(ctx)
if gormDB == nil {
return 0, errs.ErrDatabaseUninitialized
}
tx := gormDB.Exec(sqlStr)
if tx.Error != nil {
return 0, errs.NewInvalidInputError(tx.Error.Error())
}
return tx.RowsAffected, nil
}
// scanTableRows decodes every row of the result set into a column keyed map.
func scanTableRows(rows *sql.Rows, cols []string) ([]map[string]any, error) {
results := make([]map[string]any, 0)
for rows.Next() {
row, err := scanRowAsMap(rows, cols)
if err != nil {
return nil, err
}
results = append(results, row)
}
return results, nil
}
// scanRowAsMap decodes a single row, normalising driver byte slices to strings.
func scanRowAsMap(rows *sql.Rows, cols []string) (map[string]any, error) {
columns := make([]any, len(cols))
columnPointers := make([]any, len(cols))
for i := range columns {
columnPointers[i] = &columns[i]
}
if err := rows.Scan(columnPointers...); err != nil {
return nil, err
}
rowMap := make(map[string]any)
for i, colName := range cols {
val := columns[i]
if b, ok := val.([]byte); ok {
rowMap[colName] = string(b)
continue
}
rowMap[colName] = val
}
return rowMap, nil
}
// GetSQLiteInfo collects the SQLite type/name/version triple.
func GetSQLiteInfo(ctx context.Context) model.DatabaseInfoResponse {
info := model.DatabaseInfoResponse{
Type: logDBNameSQLite,
Name: config.Config.Database.SQLitePath,
Version: "SQLite",
}
if info.Name == "" {
info.Name = defaultSQLiteDBPath
}
gormDB := GetDB(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
}
// GetPostgresInfo collects the PostgreSQL type/name/version triple.
func GetPostgresInfo(ctx context.Context) model.DatabaseInfoResponse {
info := model.DatabaseInfoResponse{
Type: "postgres",
Name: config.Config.Database.Database,
Version: "PostgreSQL",
}
gormDB := GetDB(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
}
// OpenSQLiteExportFile opens the active SQLite database file together with its stat info.
func OpenSQLiteExportFile() (*os.File, os.FileInfo, error) {
//nolint:gosec // export db file path is trusted
f, err := os.Open(sqliteDatabasePath())
if err != nil {
return nil, nil, fmt.Errorf("%s: %w", errs.ErrOpenDatabaseFileFailed, err)
}
fi, err := f.Stat()
if err != nil {
_ = f.Close()
return nil, nil, fmt.Errorf("%s: %w", errs.ErrReadDatabaseFileInfoFailed, err)
}
return f, fi, nil
}
// NewPgDumpCommand builds the streaming pg_dump command for the active database.
func NewPgDumpCommand(ctx context.Context) (*exec.Cmd, string, error) {
dbCfg := config.Config.Database
pgDumpPath, err := exec.LookPath("pg_dump")
if err != nil {
return nil, "", errors.New(errs.ErrPgDumpUnavailable)
}
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(ctx, 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"))
return cmd, fileName, nil
}
@@ -1,7 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
package repository_test
import (
"context"
@@ -19,6 +19,8 @@ import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
cacheplugin "Wavelet/plugins/infra/cache"
)
@@ -44,10 +46,9 @@ func newFlushLogTestCache(t *testing.T) (contracts.CacheService, *miniredis.Mini
svc, err := core.Inject[contracts.CacheService](ctx)
require.NoError(t, err)
prev := cacheService
SetCacheService(svc)
repository.SetCacheService(svc)
cleanup := func() {
SetCacheService(prev)
repository.SetCacheService(nil)
_ = rdb.Close()
mr.Close()
}
@@ -64,12 +65,12 @@ func TestFlushTaskExecutionLogPropagatesCacheError(t *testing.T) {
const taskID = "flush-err-task"
// 先缓冲一行日志
require.NoError(t, AppendTaskExecutionLog(ctx, taskID, "step-1 ok"))
require.NoError(t, repository.AppendTaskExecutionLog(ctx, taskID, "step-1 ok"))
// 关闭 miniredis 模拟缓存基础设施故障(读取出错而非未命中)
mr.Close()
err := FlushTaskExecutionLog(ctx, taskID)
err := repository.FlushTaskExecutionLog(ctx, taskID)
assert.Error(t, err, "缓存故障时必须返回错误,防止缓冲日志被静默丢弃")
}
@@ -79,7 +80,7 @@ func TestFlushTaskExecutionLogCacheMissIsNoop(t *testing.T) {
defer cleanup()
ctx := context.Background()
assert.NoError(t, FlushTaskExecutionLog(ctx, "missing-task"))
assert.NoError(t, repository.FlushTaskExecutionLog(ctx, "missing-task"))
}
// TestFlushTaskExecutionLogPersistsAndClears 验证正常路径:缓冲日志写入执行记录后清理缓存。
@@ -89,25 +90,25 @@ func TestFlushTaskExecutionLogPersistsAndClears(t *testing.T) {
ctx := context.Background()
const taskID = "flush-ok-task"
require.NoError(t, AppendTaskExecutionLog(ctx, taskID, "done"))
require.NoError(t, repository.AppendTaskExecutionLog(ctx, taskID, "done"))
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&TaskExecution{}))
SetDBService(stubDBService{db: sqliteDB})
defer SetDBService(nil)
require.NoError(t, sqliteDB.AutoMigrate(&model.TaskExecution{}))
repository.SetDBService(stubDBService{db: sqliteDB})
defer repository.SetDBService(nil)
gormDB := sqliteDB
exec := &TaskExecution{TaskID: taskID, TaskType: "upload:test", TaskName: "t", Status: TaskExecutionStatusSucceeded}
exec := &model.TaskExecution{TaskID: taskID, TaskType: "upload:test", TaskName: "t", Status: model.TaskExecutionStatusSucceeded}
require.NoError(t, gormDB.Create(exec).Error)
require.NoError(t, FlushTaskExecutionLog(ctx, taskID))
require.NoError(t, repository.FlushTaskExecutionLog(ctx, taskID))
var got TaskExecution
var got model.TaskExecution
require.NoError(t, gormDB.First(&got, exec.ID).Error)
assert.Contains(t, got.Log, "done")
// 缓存中的缓冲日志应已被清理
var buf string
err = svc.Get(ctx, taskExecutionLogRedisKey(taskID), &buf)
err = svc.Get(ctx, repository.TaskExecutionLogRedisKey(taskID), &buf)
assert.True(t, errors.Is(err, contracts.ErrCacheMiss), "flush 后缓存应清空, got %v", err)
}
@@ -0,0 +1,53 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"Wavelet/pkg/util"
"context"
)
// UserDisplayName is the minimal user projection needed to decorate access log rows.
type UserDisplayName struct {
Username string
Nickname string
}
// SearchUserIDsByUsername is the database fallback used when the user contract is absent.
func SearchUserIDsByUsername(ctx context.Context, username string) ([]uint64, error) {
gormDB := GetDB(ctx)
if gormDB == nil {
return nil, nil
}
var ids []uint64
if err := gormDB.Table("w_users").
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
Pluck("id", &ids).Error; err != nil {
return nil, err
}
return ids, nil
}
// LoadUserDisplayNames resolves usernames and nicknames for the given ids.
func LoadUserDisplayNames(ctx context.Context, userIDs []uint64) (map[uint64]UserDisplayName, error) {
result := make(map[uint64]UserDisplayName, len(userIDs))
gormDB := GetDB(ctx)
if gormDB == nil || len(userIDs) == 0 {
return result, nil
}
var users []struct {
ID uint64
Username string
Nickname string
}
if err := gormDB.Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err != nil {
return nil, err
}
for _, u := range users {
result[u.ID] = UserDisplayName{Username: u.Username, Nickname: u.Nickname}
}
return result, nil
}
@@ -0,0 +1,421 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package repository provides persistence operations for the admin domain.
package repository
import (
"Wavelet/core/contracts"
"Wavelet/pkg/cache/ram"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"context"
"encoding/json"
"errors"
"fmt"
"strconv"
"sync"
"time"
"github.com/shopspring/decimal"
"gorm.io/gorm"
)
const (
configTypeSystem = "system"
)
var (
repoMu sync.RWMutex
dbService contracts.DBService
cacheService contracts.CacheService
)
// SetDBService injects the DBService contract.
func SetDBService(s contracts.DBService) {
repoMu.Lock()
defer repoMu.Unlock()
dbService = s
}
// SetCacheService injects the CacheService contract.
func SetCacheService(s contracts.CacheService) {
repoMu.Lock()
defer repoMu.Unlock()
cacheService = s
}
// ResetServices clears injected persistence services.
func ResetServices() {
repoMu.Lock()
defer repoMu.Unlock()
dbService = nil
cacheService = nil
}
// GetDB returns the GORM DB instance bound to the context if available.
func GetDB(ctx context.Context) *gorm.DB {
repoMu.RLock()
defer repoMu.RUnlock()
if dbService == nil {
return nil
}
return dbService.DB(ctx)
}
// GetCache returns the unified CacheService instance.
func GetCache(_ context.Context) contracts.CacheService {
repoMu.RLock()
defer repoMu.RUnlock()
return cacheService
}
// PreheatSystemConfigs loads all system configs from database.
func PreheatSystemConfigs(ctx context.Context) ([]model.SystemConfig, error) {
database := GetDB(ctx)
if database == nil {
return nil, errors.New(errs.ErrDatabaseNotInitialized)
}
var configs []model.SystemConfig
if err := database.Find(&configs).Error; err != nil {
return nil, err
}
return configs, nil
}
// PreheatSystemConfigByKey loads a single config key from database.
func PreheatSystemConfigByKey(ctx context.Context, key string) (model.SystemConfig, error) {
database := GetDB(ctx)
if database == nil {
return model.SystemConfig{}, errors.New(errs.ErrDatabaseNotInitialized)
}
var sc model.SystemConfig
if err := database.Where("key = ?", key).First(&sc).Error; err != nil {
return model.SystemConfig{}, err
}
return sc, nil
}
// GetSystemConfigByGroup queries a configuration by Type and Key.
func GetSystemConfigByGroup(ctx context.Context, configType, key string) (model.SystemConfig, error) {
ensureSystemConfigCacheListener()
if item, ok := ram.Get(configType, key); ok {
var sc model.SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
return sc, nil
}
}
database := GetDB(ctx)
if database == nil {
return model.SystemConfig{}, errors.New(errs.ErrDatabaseNotInitialized)
}
var sc model.SystemConfig
if err := database.Where("key = ?", key).First(&sc).Error; err != nil {
return model.SystemConfig{}, err
}
valBytes, err := json.Marshal(sc)
if err == nil {
ram.Set(ram.CacheItem{
Key: sc.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(sc.Key),
})
}
return sc, nil
}
// GetSystemConfigByKey queries config by key.
func GetSystemConfigByKey(ctx context.Context, key string) (model.SystemConfig, error) {
return GetSystemConfigByGroup(ctx, ConfigCacheType, key)
}
// ListSystemConfigsByKeys loads multiple config keys.
func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]model.SystemConfig, error) {
if len(keys) == 0 {
return map[string]model.SystemConfig{}, nil
}
ensureSystemConfigCacheListener()
result := make(map[string]model.SystemConfig, len(keys))
missing := make([]string, 0, len(keys))
for _, key := range keys {
if item, ok := ram.Get(ConfigCacheType, key); ok {
var sc model.SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
result[key] = sc
continue
}
}
missing = append(missing, key)
}
if len(missing) == 0 {
return result, nil
}
database := GetDB(ctx)
if database == nil {
return nil, errors.New(errs.ErrDatabaseNotInitialized)
}
var configs []model.SystemConfig
if err := database.Where("key IN ?", missing).Find(&configs).Error; err != nil {
return nil, err
}
for i := range configs {
valBytes, err := json.Marshal(configs[i])
if err == nil {
ram.Set(ram.CacheItem{
Key: configs[i].Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(configs[i].Key),
})
}
result[configs[i].Key] = configs[i]
}
return result, nil
}
// InvalidateVisibleSystemConfigsCache clears the cached public config list.
func InvalidateVisibleSystemConfigsCache(ctx context.Context) error {
return InvalidateAllSystemConfigCaches(ctx)
}
// ListVisibleSystemConfigs queries visible configs using local cache store.
func ListVisibleSystemConfigs(ctx context.Context) ([]model.SystemConfig, error) {
ensureSystemConfigCacheListener()
items := ram.GetTypeItems(ConfigCacheType)
if len(items) > 0 {
var list []model.SystemConfig
for _, item := range items {
var sc model.SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
if sc.Visibility == model.ConfigVisibilityVisible {
list = append(list, sc)
}
}
}
return list, nil
}
database := GetDB(ctx)
if database == nil {
return nil, errors.New(errs.ErrDatabaseNotInitialized)
}
var configs []model.SystemConfig
if err := database.Where("visibility = ?", model.ConfigVisibilityVisible).Find(&configs).Error; err != nil {
return nil, err
}
for _, cfg := range configs {
valBytes, err := json.Marshal(cfg)
if err == nil {
ram.Set(ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(cfg.Key),
})
}
}
return configs, nil
}
// GetIntByKey queries config and converts to int.
func GetIntByKey(ctx context.Context, key string) (int, error) {
sc, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return 0, err
}
value, err := strconv.Atoi(sc.Value)
if err != nil {
return 0, fmt.Errorf(errs.ErrConfigIntParseFailed, key, sc.Value, err)
}
return value, nil
}
// GetDecimalByKey queries config and converts to decimal.Decimal.
func GetDecimalByKey(ctx context.Context, key string, precision int32) (decimal.Decimal, error) {
sc, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return decimal.Zero, err
}
value, err := decimal.NewFromString(sc.Value)
if err != nil {
return decimal.Zero, fmt.Errorf(errs.ErrConfigDecimalParseFailed, key, sc.Value, err)
}
return value.Truncate(precision), nil
}
// GetBoolByKey queries config and converts to bool.
func GetBoolByKey(ctx context.Context, key string) (bool, error) {
sc, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return false, err
}
value, err := strconv.ParseBool(sc.Value)
if err != nil {
return false, fmt.Errorf(errs.ErrConfigBoolParseFailed, key, sc.Value, err)
}
return value, nil
}
// GetMenuDisplayConfig queries and parses menu config.
func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
sc, err := GetSystemConfigByKey(ctx, model.ConfigKeyMenuDisplayConfig)
if err != nil {
return nil, err
}
config := make(map[string]bool)
if sc.Value == "" || sc.Value == "{}" {
return config, nil
}
if err := json.Unmarshal([]byte(sc.Value), &config); err != nil {
return nil, fmt.Errorf(errs.ErrParseMenuDisplayConfigFailed, err)
}
return config, nil
}
// ListAdminSystemConfigs returns all configs, optionally filtered by type.
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) {
query := GetDB(ctx).Order("created_at DESC")
if configType != "" {
query = query.Where("type = ?", configType)
}
var configs []model.SystemConfig
if err := query.Find(&configs).Error; err != nil {
return nil, err
}
return configs, nil
}
// GetAdminSystemConfigByKey loads a config directly from DB.
func GetAdminSystemConfigByKey(ctx context.Context, key string) (model.SystemConfig, error) {
var config model.SystemConfig
if err := GetDB(ctx).Where("key = ?", key).First(&config).Error; err != nil {
return model.SystemConfig{}, err
}
return config, nil
}
// SystemConfigExists reports whether a config key already exists.
func SystemConfigExists(ctx context.Context, key string) (bool, error) {
var existing model.SystemConfig
err := GetDB(ctx).Where("key = ?", key).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
if err != nil {
return false, err
}
return true, nil
}
// CreateSystemConfigRecord persists a new system config row.
func CreateSystemConfigRecord(ctx context.Context, config *model.SystemConfig) error {
return GetDB(ctx).Create(config).Error
}
// UpdateSystemConfigFields applies partial updates to a system config row.
func UpdateSystemConfigFields(ctx context.Context, config *model.SystemConfig, updates map[string]any) error {
return GetDB(ctx).Model(config).Updates(updates).Error
}
// UpdateSystemConfigTx applies the config row updates inside a transaction and, when
// resolveTaskType is not empty, marks that task type's failed executions as succeeded
// within the same transaction.
func UpdateSystemConfigTx(
ctx context.Context,
config *model.SystemConfig,
updates map[string]any,
resolveTaskType string,
resolveResult string,
) error {
database := GetDB(ctx)
if database == nil {
return errors.New(errs.ErrDatabaseServiceNotAvailable)
}
return database.Transaction(func(tx *gorm.DB) error {
if err := tx.Model(config).Updates(updates).Error; err != nil {
return err
}
if resolveTaskType == "" {
return nil
}
if err := MarkFailedTaskExecutionsSucceededTx(tx, resolveTaskType, resolveResult, time.Now()); err != nil {
logger.ErrorF(ctx, errs.ErrAutoResolveMigrationTaskFailed, err)
}
return nil
})
}
// SaveOrUpdateSystemConfig creates or updates a config row and invalidates cache.
func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
var sc model.SystemConfig
err := GetDB(ctx).Where("key = ?", key).First(&sc).Error
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
if errors.Is(err, gorm.ErrRecordNotFound) {
sc = model.SystemConfig{
Key: key,
Value: value,
Type: configTypeSystem,
Visibility: model.ConfigVisibilityHidden,
}
if err := GetDB(ctx).Create(&sc).Error; err != nil {
return err
}
} else {
sc.Value = value
if err := GetDB(ctx).Save(&sc).Error; err != nil {
return err
}
}
return InvalidateSystemConfigCache(ctx, key)
}
// CountActiveUploads counts non-deleted rows of the storage upload table. A missing
// database handle yields zero, matching the pre-refactor guard behaviour.
func CountActiveUploads(ctx context.Context) (int64, error) {
gormDB := GetDB(ctx)
if gormDB == nil {
return 0, nil
}
var uploadCount int64
if err := gormDB.Table("w_uploads").
Where("status != ?", "deleted").
Count(&uploadCount).Error; err != nil {
return 0, fmt.Errorf(errs.ErrCheckExistingUploadsFailed, err)
}
return uploadCount, nil
}
@@ -0,0 +1,331 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"Wavelet/core/contracts"
"Wavelet/pkg/idgen"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"context"
"errors"
"fmt"
"strings"
"time"
"gorm.io/gorm"
)
const (
taskExecutionLogRedisKeyPrefix = "task:execution:log:"
taskExecutionLogExpiration = 24 * time.Hour
)
// CreateScheduleRecord 创建定时任务
func CreateScheduleRecord(ctx context.Context, schedule *model.Schedule) error {
return GetDB(ctx).Create(schedule).Error
}
// UpdateScheduleRecord 更新定时任务
func UpdateScheduleRecord(ctx context.Context, schedule *model.Schedule) error {
return GetDB(ctx).Save(schedule).Error
}
// DeleteScheduleRecord 删除定时任务
func DeleteScheduleRecord(ctx context.Context, id uint64) error {
return GetDB(ctx).Delete(&model.Schedule{}, id).Error
}
// GetScheduleByID 根据 ID 获取定时任务
func GetScheduleByID(ctx context.Context, id uint64) (*model.Schedule, error) {
var schedule model.Schedule
if err := GetDB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
return nil, err
}
return &schedule, nil
}
// ListSchedulesRecord 获取所有定时任务
func ListSchedulesRecord(ctx context.Context) ([]model.Schedule, error) {
var schedules []model.Schedule
if err := GetDB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
// ListActiveSchedules 获取所有启用的定时任务
func ListActiveSchedules(ctx context.Context) ([]model.Schedule, error) {
var schedules []model.Schedule
if err := GetDB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
// CreateTaskExecutionRecord 创建任务执行记录
func CreateTaskExecutionRecord(ctx context.Context, execution *model.TaskExecution) error {
execution.ID = idgen.NextUint64ID()
return GetDB(ctx).Create(execution).Error
}
// UpdateTaskExecutionRecord 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
func UpdateTaskExecutionRecord(ctx context.Context, execution *model.TaskExecution) error {
return GetDB(ctx).Omit("log").Save(execution).Error
}
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*model.TaskExecution, error) {
var execution model.TaskExecution
if err := GetDB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetTaskExecutionByID 根据 ID 获取执行记录
func GetTaskExecutionByID(ctx context.Context, id uint64) (*model.TaskExecution, error) {
var execution model.TaskExecution
if err := GetDB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetLatestTaskExecutionByTaskType returns the most recent execution for a task type.
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*model.TaskExecution, bool, error) {
var execution model.TaskExecution
err := GetDB(ctx).
Where("task_type = ?", taskType).
Order("id DESC").
First(&execution).Error
if err == nil {
if loadErr := loadTaskExecutionLog(ctx, &execution); loadErr != nil {
return nil, false, loadErr
}
return &execution, true, nil
}
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, nil
}
return nil, false, err
}
// AppendTaskExecutionLog 将日志追加到缓冲,任务完成后再持久化到数据库。
func AppendTaskExecutionLog(ctx context.Context, taskID, logLine string) error {
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return errors.New(errs.ErrCacheServiceNotInitialized)
}
now := time.Now().Format("15:04:05")
line := fmt.Sprintf("[%s] %s\n", now, logLine)
key := TaskExecutionLogRedisKey(taskID)
var existing string
_ = cacheSvc.Get(ctx, key, &existing)
return cacheSvc.Set(ctx, key, existing+line, taskExecutionLogExpiration)
}
// FlushTaskExecutionLog 将缓冲中的完整任务日志写入数据库,并在成功后清理缓存。
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return errors.New(errs.ErrCacheServiceNotInitialized)
}
key := TaskExecutionLogRedisKey(taskID)
var logText string
if err := cacheSvc.Get(ctx, key, &logText); err != nil {
// 缓存未命中属于正常情况(任务无输出),其余错误必须上抛,
// 否则缓冲日志会被静默丢弃并误报持久化成功。
if !errors.Is(err, contracts.ErrCacheMiss) {
return fmt.Errorf("load buffered task execution log: %w", err)
}
return nil
}
if logText == "" {
return nil
}
gormDB := GetDB(ctx)
if gormDB == nil {
return errors.New(errs.ErrDatabaseNotInitialized)
}
result := gormDB.Model(&model.TaskExecution{}).
Where("task_id = ?", taskID).
Update("log", logText)
if result.Error != nil {
return fmt.Errorf("persist task execution log: %w", result.Error)
}
if result.RowsAffected == 0 {
return fmt.Errorf("persist task execution log: task %q not found", taskID)
}
_ = cacheSvc.Delete(ctx, key)
return nil
}
// ListTaskExecutionRecords 分页查询任务执行记录
func ListTaskExecutionRecords(ctx context.Context, req model.ListTaskExecutionsRequest) ([]model.TaskExecution, int64, error) {
if req.Page <= 0 {
req.Page = 1
}
if req.PageSize <= 0 {
req.PageSize = 20
}
query := GetDB(ctx).Model(&model.TaskExecution{})
if req.Status != "" {
query = query.Where("status = ?", req.Status)
}
if req.TaskType != "" {
query = query.Where("task_type = ?", req.TaskType)
} else if types := parseTaskTypesFilter(req.TaskTypes); len(types) > 0 {
query = query.Where("task_type IN ?", types)
} else if req.TaskTypePrefix != "" {
query = query.Where("task_type LIKE ? ESCAPE '\\'", util.EscapeLike(req.TaskTypePrefix)+"%")
}
var total int64
if err := query.Count(&total).Error; err != nil {
return nil, 0, err
}
var executions []model.TaskExecution
offset := (req.Page - 1) * req.PageSize
if err := query.Order("id DESC").Offset(offset).Limit(req.PageSize).Find(&executions).Error; err != nil {
return nil, 0, err
}
if err := loadTaskExecutionLogs(ctx, executions); err != nil {
return nil, 0, err
}
return executions, total, nil
}
func parseTaskTypesFilter(raw string) []string {
if strings.TrimSpace(raw) == "" {
return nil
}
parts := strings.Split(raw, ",")
out := make([]string, 0, len(parts))
for _, part := range parts {
part = strings.TrimSpace(part)
if part != "" {
out = append(out, part)
}
}
return out
}
// MarkFailedTaskExecutionsSucceededTx marks failed executions of a task type as succeeded within a transaction.
func MarkFailedTaskExecutionsSucceededTx(
tx *gorm.DB,
taskType string,
result string,
finishedAt time.Time,
) error {
return tx.Model(&model.TaskExecution{}).
Where("task_type = ? AND status = ?", taskType, model.TaskExecutionStatusFailed).
Updates(map[string]any{
"status": model.TaskExecutionStatusSucceeded,
"result": result,
"finished_at": finishedAt,
}).Error
}
// CleanupTaskExecutionLogs removes finished task execution logs according to frequency-based retention.
func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (model.TaskExecutionCleanupStats, error) {
const (
frequencyWindowDays = 30
highFrequencyThreshold = frequencyWindowDays
)
frequencyWindowStart := now.AddDate(0, 0, -frequencyWindowDays)
highFrequencyCutoff := now.AddDate(0, 0, -3)
lowFrequencyCutoff := now.AddDate(0, 0, -30)
terminalStatuses := []model.TaskExecutionStatus{model.TaskExecutionStatusSucceeded, model.TaskExecutionStatusFailed}
var highFrequencyTaskTypes []string
if err := GetDB(ctx).
Model(&model.TaskExecution{}).
Select("task_type").
Where("created_at >= ?", frequencyWindowStart).
Group("task_type").
Having("COUNT(*) > ?", highFrequencyThreshold).
Pluck("task_type", &highFrequencyTaskTypes).Error; err != nil {
return model.TaskExecutionCleanupStats{}, fmt.Errorf("query high-frequency task types: %w", err)
}
var highFrequencyDeleted int64
if len(highFrequencyTaskTypes) > 0 {
highFrequencyResult := GetDB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", highFrequencyCutoff).
Where("task_type IN ?", highFrequencyTaskTypes).
Delete(&model.TaskExecution{})
if highFrequencyResult.Error != nil {
return model.TaskExecutionCleanupStats{}, fmt.Errorf("delete high-frequency task execution logs: %w", highFrequencyResult.Error)
}
highFrequencyDeleted = highFrequencyResult.RowsAffected
}
lowFrequencyQuery := GetDB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", lowFrequencyCutoff)
if len(highFrequencyTaskTypes) > 0 {
lowFrequencyQuery = lowFrequencyQuery.Where("task_type NOT IN ?", highFrequencyTaskTypes)
}
lowFrequencyResult := lowFrequencyQuery.Delete(&model.TaskExecution{})
if lowFrequencyResult.Error != nil {
return model.TaskExecutionCleanupStats{}, fmt.Errorf("delete low-frequency task execution logs: %w", lowFrequencyResult.Error)
}
return model.TaskExecutionCleanupStats{
HighFrequencyDeleted: highFrequencyDeleted,
LowFrequencyDeleted: lowFrequencyResult.RowsAffected,
}, nil
}
// TaskExecutionLogRedisKey builds the Redis key for task execution logs.
func TaskExecutionLogRedisKey(taskID string) string {
return taskExecutionLogRedisKeyPrefix + taskID
}
func loadTaskExecutionLog(ctx context.Context, execution *model.TaskExecution) error {
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return nil
}
var logText string
if err := cacheSvc.Get(ctx, TaskExecutionLogRedisKey(execution.TaskID), &logText); err == nil && logText != "" {
execution.Log = logText
}
return nil
}
func loadTaskExecutionLogs(ctx context.Context, executions []model.TaskExecution) error {
cacheSvc := GetCache(ctx)
if cacheSvc == nil || len(executions) == 0 {
return nil
}
for i := range executions {
var logText string
if err := cacheSvc.Get(ctx, TaskExecutionLogRedisKey(executions[i].TaskID), &logText); err == nil && logText != "" {
executions[i].Log = logText
}
}
return nil
}
@@ -0,0 +1,58 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"Wavelet/plugins/domain/admin/model"
"context"
"errors"
"gorm.io/gorm"
)
// ListTemplatesRecord returns all templates ordered by system flag and creation time.
func ListTemplatesRecord(ctx context.Context) ([]model.Template, error) {
var templates []model.Template
if err := GetDB(ctx).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
return nil, err
}
return templates, nil
}
// GetTemplateByKey loads a template by its key.
func GetTemplateByKey(ctx context.Context, key string) (model.Template, error) {
var tmpl model.Template
if err := GetDB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil {
return model.Template{}, err
}
return tmpl, nil
}
// TemplateExistsByKey reports whether a template key is already taken.
func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
var existing model.Template
err := GetDB(ctx).Where("key = ?", key).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
if err != nil {
return false, err
}
return true, nil
}
// CreateTemplateRecord persists a new template.
func CreateTemplateRecord(ctx context.Context, tmpl *model.Template) error {
return GetDB(ctx).Create(tmpl).Error
}
// SaveTemplateRecord updates an existing template.
func SaveTemplateRecord(ctx context.Context, tmpl *model.Template) error {
return GetDB(ctx).Save(tmpl).Error
}
// DeleteTemplateRecord removes a template record.
func DeleteTemplateRecord(ctx context.Context, tmpl *model.Template) error {
return GetDB(ctx).Delete(tmpl).Error
}
@@ -1,12 +0,0 @@
//go:build windows
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import "errors"
func replaceAndRestart(_, _ string) error {
return errors.New(errAutomaticUpgradeBlocked)
}
@@ -0,0 +1,87 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/admin/errs"
"context"
"errors"
"fmt"
)
// ListAuthSources returns every configured authentication source.
func ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
authSvc, err := requireAuthService(ctx)
if err != nil {
return nil, err
}
views, err := authSvc.ListAuthSources(ctx)
if err != nil {
logger.ErrorF(ctx, "List auth sources failed: %v", err)
return nil, errors.New(errs.ListAuthSourcesFailed)
}
return views, nil
}
// CreateAuthSource registers a new authentication source.
func CreateAuthSource(ctx context.Context, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
authSvc, err := requireAuthService(ctx)
if err != nil {
return nil, err
}
created, err := authSvc.CreateAuthSource(ctx, source)
if err != nil {
return nil, fmt.Errorf("%s%w", errs.CreateAuthSourceFailed, err)
}
return created, nil
}
// UpdateAuthSource rewrites an existing authentication source.
func UpdateAuthSource(
ctx context.Context,
id uint64,
source contracts.AuthSourceDTO,
) (*contracts.AuthSourceDTO, error) {
authSvc, err := requireAuthService(ctx)
if err != nil {
return nil, err
}
updated, err := authSvc.UpdateAuthSource(ctx, id, source)
if err != nil {
return nil, err
}
return updated, nil
}
// ToggleAuthSource flips the active state of an authentication source.
func ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
authSvc, err := requireAuthService(ctx)
if err != nil {
return nil, err
}
toggled, err := authSvc.ToggleAuthSource(ctx, id)
if err != nil {
return nil, fmt.Errorf("%s%w", errs.ToggleAuthSourceFailed, err)
}
return toggled, nil
}
// DeleteAuthSource removes an authentication source.
func DeleteAuthSource(ctx context.Context, id uint64) error {
authSvc, err := requireAuthService(ctx)
if err != nil {
return err
}
if err := authSvc.DeleteAuthSource(ctx, id); err != nil {
return fmt.Errorf("%s%w", errs.DeleteAuthSourceFailed, err)
}
return nil
}
@@ -0,0 +1,45 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"context"
"strconv"
pkgcache "Wavelet/pkg/cache/disk"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
)
// DiskCacheStatus reports the disk cache usage counters.
func DiskCacheStatus() pkgcache.Status {
return pkgcache.Default().Status()
}
// ClearDiskCache purges every cached object and resets the tracking counters.
func ClearDiskCache() error {
return pkgcache.Default().Clear()
}
// UpdateDiskCachePolicy persists the disk cache settings and applies them hot.
func UpdateDiskCachePolicy(ctx context.Context, req model.UpdateCacheConfigRequest) error {
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
return err
}
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
return err
}
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
return err
}
pkgcache.Default().UpdatePolicy(req.MaxSizeMB, req.TTLMinutes, req.LRUEnabled)
return nil
}
func saveOrUpdateCacheConfig(ctx context.Context, key, value string) error {
return repository.SaveOrUpdateSystemConfig(ctx, key, value)
}
@@ -0,0 +1,299 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
mail "Wavelet/pkg/mail"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
"context"
"encoding/json"
"errors"
"fmt"
)
const maskedConfigValue = "******"
// PublicSystemConfigs returns the key/value map exposed to unauthenticated clients.
func PublicSystemConfigs(ctx context.Context) (map[string]string, error) {
configs, err := repository.ListVisibleSystemConfigs(ctx)
if err != nil {
return nil, err
}
resp := make(map[string]string, len(configs))
for _, config := range configs {
resp[config.Key] = config.Value
}
return resp, nil
}
// ListAdminSystemConfigs returns every config, optionally filtered by type, with secrets masked.
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) {
configs, err := repository.ListAdminSystemConfigs(ctx, configType)
if err != nil {
return nil, err
}
for i := range configs {
configs[i].Value = MaskSensitiveConfig(configs[i].Key, configs[i].Value)
}
return configs, nil
}
// GetAdminSystemConfig loads a single config with its secrets masked.
func GetAdminSystemConfig(ctx context.Context, key string) (model.SystemConfig, error) {
config, err := repository.GetAdminSystemConfigByKey(ctx, key)
if err != nil {
return model.SystemConfig{}, translateNotFound(err, errs.ErrSystemConfigNotFound)
}
config.Value = MaskSensitiveConfig(config.Key, config.Value)
return config, nil
}
// CreateAdminSystemConfig persists a new config key and refreshes the cache layer.
func CreateAdminSystemConfig(ctx context.Context, req model.CreateSystemConfigRequest) error {
if isProtectedConfigKey(req.Key) {
return errs.ErrProtectedConfigKey
}
exists, err := repository.SystemConfigExists(ctx, req.Key)
if err != nil {
return err
}
if exists {
return errs.ErrConfigKeyExists
}
config := model.SystemConfig{
Key: req.Key,
Value: req.Value,
Type: req.Type,
Visibility: req.Visibility,
Description: req.Description,
}
if err := repository.CreateSystemConfigRecord(ctx, &config); err != nil {
return err
}
invalidateSystemConfigCaches(ctx, req.Key)
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
return nil
}
// UpdateAdminSystemConfig applies an update to a protected-aware config key inside a transaction.
func UpdateAdminSystemConfig(ctx context.Context, key string, req model.UpdateSystemConfigRequest) error {
if isProtectedConfigKey(key) {
return errs.ErrProtectedConfigKey
}
config, err := repository.GetAdminSystemConfigByKey(ctx, key)
if err != nil {
return translateNotFound(err, errs.ErrSystemConfigNotFound)
}
var originalDriver contracts.StorageDriver
resolveTaskType := ""
resolveResult := ""
if key == model.ConfigKeyStorageConfig {
var currentCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(config.Value), &currentCfg); err == nil {
originalDriver = currentCfg.Driver
}
validatedVal, err := validateAndMergeStorageConfig(ctx, req.Value, config.Value)
if err != nil {
return err
}
req.Value = validatedVal
var newCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(req.Value), &newCfg); err == nil {
resolveTaskType, resolveResult = storageMigrationResolutionTask(originalDriver, newCfg.Driver)
}
}
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 := repository.UpdateSystemConfigTx(ctx, &config, updates, resolveTaskType, resolveResult); err != nil {
return err
}
invalidateCachesAfterConfigUpdate(ctx, key)
return nil
}
// storageMigrationResolutionTask reports the failed-task resolution that a direct storage
// config rewrite implies. An empty task type means nothing has to be resolved.
func storageMigrationResolutionTask(
originalDriver contracts.StorageDriver,
newDriver contracts.StorageDriver,
) (string, string) {
if originalDriver == "" || newDriver != originalDriver {
return "", ""
}
return errs.StorageMigrationTaskType, errs.StorageDriverResolvedResult
}
func isProtectedConfigKey(key string) bool {
return key == model.ConfigKeyLogDatabase || key == model.ConfigKeyLogDBMigration
}
func invalidateSystemConfigCaches(ctx context.Context, key string) {
if err := repository.InvalidateSystemConfigCache(ctx, key); err != nil {
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
}
_ = EmitEvent(ctx, contracts.EventTopicConfigChanged, contracts.ConfigChangedEvent{Key: key})
}
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
invalidateSystemConfigCaches(ctx, key)
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
}
// TestSMTP sends a probe mail, resolving a masked password from the stored config.
func TestSMTP(ctx context.Context, req model.TestSMTPRequest) model.TestSMTPResponse {
password := req.SMTPPassword
if password == maskedConfigValue {
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword); err == nil {
password = sc.Value
}
}
cfg := mail.Config{
Host: req.SMTPHost,
Port: req.SMTPPort,
Username: req.SMTPUsername,
Password: password,
}
subject := "Wavelet SMTP Test Mail"
body := `<h3>SMTP Mail Connection Test</h3>
<p>If you received this message, your SMTP configuration is correct and mail sending is working properly.</p>
<p>Sent from Wavelet.</p>`
logs, err := mail.SendMailWithLog(ctx, cfg, req.To, subject, body)
resp := model.TestSMTPResponse{
Success: err == nil,
Log: logs,
}
if err != nil {
resp.Error = err.Error()
}
return resp
}
// MaskSensitiveConfig masks secret config values before exposing to clients.
func MaskSensitiveConfig(key, value string) string {
if value == "" {
return value
}
switch key {
case model.ConfigKeySMTPPassword:
return maskedConfigValue
case model.ConfigKeyStorageConfig:
return maskStorageConfig(value)
}
return value
}
func maskStorageConfig(value string) string {
var cfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(value), &cfg); err != nil {
return value
}
if cfg.S3.SecretAccessKey != "" {
cfg.S3.SecretAccessKey = maskedConfigValue
}
if cfg.R2.SecretAccessKey != "" {
cfg.R2.SecretAccessKey = maskedConfigValue
}
if cfg.MinIO.SecretAccessKey != "" {
cfg.MinIO.SecretAccessKey = maskedConfigValue
}
if cfg.OSS.SecretAccessKey != "" {
cfg.OSS.SecretAccessKey = maskedConfigValue
}
if cfg.WebDAV.Password != "" {
cfg.WebDAV.Password = maskedConfigValue
}
val, err := json.Marshal(cfg)
if err != nil {
return value
}
return string(val)
}
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
// and tests connectivity of the new storage configuration.
func validateAndMergeStorageConfig(ctx context.Context, value, currentConfig string) (string, error) {
var currentCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(currentConfig), &currentCfg); err != nil {
return "", fmt.Errorf(errs.ErrParseCurrentStorageConfigFailed, err)
}
var newCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(value), &newCfg); err != nil {
return "", fmt.Errorf(errs.ErrParseTargetStorageConfigFailed, err)
}
// 合并被掩码屏蔽的敏感信息,获取完整的真实配置
targetCfg := newCfg
if targetCfg.S3.SecretAccessKey == maskedConfigValue {
targetCfg.S3.SecretAccessKey = currentCfg.S3.SecretAccessKey
}
if targetCfg.R2.SecretAccessKey == maskedConfigValue {
targetCfg.R2.SecretAccessKey = currentCfg.R2.SecretAccessKey
}
if targetCfg.MinIO.SecretAccessKey == maskedConfigValue {
targetCfg.MinIO.SecretAccessKey = currentCfg.MinIO.SecretAccessKey
}
if targetCfg.OSS.SecretAccessKey == maskedConfigValue {
targetCfg.OSS.SecretAccessKey = currentCfg.OSS.SecretAccessKey
}
if targetCfg.WebDAV.Password == maskedConfigValue {
targetCfg.WebDAV.Password = currentCfg.WebDAV.Password
}
if err := validateMergedStorageConfig(ctx, currentCfg, newCfg, targetCfg); err != nil {
return "", err
}
// 序列化为最终保存的真实明文配置,防止保存屏蔽的 ****** 字符
unmaskedVal, err := json.Marshal(targetCfg)
if err != nil {
return "", fmt.Errorf(errs.ErrSerializeStorageConfigFailed, err)
}
return string(unmaskedVal), nil
}
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, _ contracts.StorageConfigDTO) error {
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
uploadCount, err := repository.CountActiveUploads(ctx)
if err != nil {
return err
}
if uploadCount > 0 {
return errors.New(errs.StorageDriverSwitchRequiresMigration)
}
}
return nil
}
+131
View File
@@ -0,0 +1,131 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/pkg/config"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
"context"
"os"
"os/exec"
"strings"
"time"
)
// selectSQLKeywords marks statements that return a result set instead of a row count.
var selectSQLKeywords = []string{"select", "show", "explain", "describe", "pragma"}
// DatabaseOverview collects the runtime overview of the active database.
func DatabaseOverview(ctx context.Context) (model.DBOverviewResponse, error) {
if !config.Config.Database.Enabled {
return repository.GetSQLiteOverview(ctx)
}
return repository.GetPostgresOverview(ctx)
}
// DatabaseTableNames returns every user table of the active database.
func DatabaseTableNames(ctx context.Context) ([]string, error) {
return repository.ListDatabaseTableNames(ctx)
}
// DatabaseTableData loads one page of a table with its column layout and total row count.
func DatabaseTableData(ctx context.Context, req model.GetTableDataRequest) (model.TableDataResponse, error) {
quotedTable := repository.QuoteTableName(req.Table)
total, err := repository.CountDatabaseTableRows(ctx, quotedTable)
if err != nil {
return model.TableDataResponse{}, err
}
offset := (req.Page - 1) * req.PageSize
if offset < 0 {
offset = 0
}
limit := req.PageSize
if limit <= 0 {
limit = 10
}
cols, results, err := repository.QueryDatabaseTableRows(ctx, quotedTable, limit, offset)
if err != nil {
return model.TableDataResponse{}, err
}
return model.TableDataResponse{
Columns: cols,
Total: total,
Results: truncateCellValues(results),
}, nil
}
// truncateCellValues caps oversized string cells before they reach the console grid.
func truncateCellValues(rows []map[string]any) []map[string]any {
for _, row := range rows {
for column, value := range row {
if str, ok := value.(string); ok {
row[column] = model.TruncateDisplayValue(str)
}
}
}
return rows
}
// ExecuteCustomSQL runs an arbitrary statement issued from the console SQL runner.
func ExecuteCustomSQL(ctx context.Context, trimmedSQL string) (model.ExecuteSQLResponse, error) {
startTime := time.Now()
if isSelectStatement(trimmedSQL) {
cols, results, err := repository.RunSelectSQL(ctx, trimmedSQL)
if err != nil {
return model.ExecuteSQLResponse{}, err
}
return model.ExecuteSQLResponse{
Type: "select",
Columns: cols,
Results: results,
AffectedRows: int64(len(results)),
ExecutionTimeMs: time.Since(startTime).Milliseconds(),
}, nil
}
affectedRows, err := repository.RunMutationSQL(ctx, trimmedSQL)
if err != nil {
return model.ExecuteSQLResponse{}, err
}
return model.ExecuteSQLResponse{
Type: "exec",
AffectedRows: affectedRows,
ExecutionTimeMs: time.Since(startTime).Milliseconds(),
}, nil
}
// isSelectStatement reports whether the statement yields a result set.
func isSelectStatement(trimmedSQL string) bool {
lowerSQL := strings.ToLower(trimmedSQL)
for _, kw := range selectSQLKeywords {
if strings.HasPrefix(lowerSQL, kw) {
return true
}
}
return false
}
// DatabaseInfo returns the active database type, name and version.
func DatabaseInfo(ctx context.Context) model.DatabaseInfoResponse {
if !config.Config.Database.Enabled {
return repository.GetSQLiteInfo(ctx)
}
return repository.GetPostgresInfo(ctx)
}
// OpenSQLiteExportFile opens the active SQLite database file together with its stat info.
func OpenSQLiteExportFile() (*os.File, os.FileInfo, error) {
return repository.OpenSQLiteExportFile()
}
// NewPgDumpCommand builds the streaming pg_dump command for the active database.
func NewPgDumpCommand(ctx context.Context) (*exec.Cmd, string, error) {
return repository.NewPgDumpCommand(ctx)
}
+238
View File
@@ -0,0 +1,238 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
"context"
"fmt"
"net/url"
"strings"
"time"
)
const (
analyticsDays = 7
denyingRobotsFile = "User-Agent: *\nDisallow: /\n"
allowingRobotsFile = "User-Agent: *\nAllow: /\n"
)
// RecentSystemLogs reads a page of the process log ring buffer.
func RecentSystemLogs(cursor, limit int) model.LogsResponse {
entries, hasMore := logger.GlobalRingBuffer.Query(cursor, limit)
resp := model.LogsResponse{
Lines: entries,
HasMore: hasMore,
}
if len(entries) > 0 {
resp.NextCursor = entries[0].Index
}
return resp
}
// RobotsTxtBody resolves the robots.txt payload from the indexing setting.
func RobotsTxtBody(ctx context.Context) string {
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeySearchEngineIndexingEnabled)
if err == nil && enabled {
return allowingRobotsFile
}
return denyingRobotsFile
}
// IsAllowedLogOrigin reports whether a WebSocket handshake origin may subscribe to logs.
func IsAllowedLogOrigin(ctx context.Context, origin, host string) bool {
if origin == "" {
return true
}
// 1. 同源检查 (Same-origin check)
u, err := url.Parse(origin)
if err == nil && strings.EqualFold(u.Host, host) {
return true
}
// 2. 检查配置的允许跨域 Origin (Check allowed origins in system config)
sc, cfgErr := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
if cfgErr != nil || sc.Value == "" {
return false
}
originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/")
for _, allowed := range strings.Split(sc.Value, ",") {
allowed = strings.TrimRight(strings.TrimSpace(allowed), "/")
if allowed != "" && strings.EqualFold(allowed, originToCheck) {
return true
}
}
return false
}
// AccessLogs queries the analytical access log store and decorates rows with user names.
func AccessLogs(ctx context.Context, q model.AccessLogQuery) (model.AccessLogsResponse, error) {
rc := GetRiskControlService()
if rc == nil {
return model.AccessLogsResponse{}, errs.ErrLogStoreUnavailable
}
filter, err := buildAccessLogFilter(ctx, q)
if err != nil {
return model.AccessLogsResponse{}, err
}
if filter.UserIDs != nil && len(filter.UserIDs) == 0 {
return model.AccessLogsResponse{Total: 0, List: []model.AccessLogItem{}}, nil
}
logs, total, err := rc.QueryAccessLogs(ctx, filter, q.Page, q.PageSize)
if err != nil {
return model.AccessLogsResponse{}, err
}
if total == 0 {
return model.AccessLogsResponse{Total: 0, List: []model.AccessLogItem{}}, nil
}
list := make([]model.AccessLogItem, len(logs))
for i, logItem := range logs {
list[i] = model.AccessLogItem{
ID: logItem.ID,
UserID: logItem.UserID,
Path: logItem.Path,
Method: logItem.Method,
IP: logItem.IP,
UserAgent: logItem.UserAgent,
Status: logItem.Status,
Latency: logItem.Latency,
CreatedAt: logItem.CreatedAt.Format(time.RFC3339),
}
}
enrichAccessLogsWithUsers(ctx, list)
return model.AccessLogsResponse{Total: total, List: list}, nil
}
// AccessLogAnalytics aggregates the daily trend of the access log store.
func AccessLogAnalytics(ctx context.Context) (model.LogsAnalyticsResponse, error) {
rc := GetRiskControlService()
if rc == nil {
return model.LogsAnalyticsResponse{}, errs.ErrLogStoreUnavailable
}
stats, err := rc.QueryAccessLogStats(ctx, analyticsDays)
if err != nil {
return model.LogsAnalyticsResponse{}, fmt.Errorf("%s%w", errs.ErrQueryAccessTrendFailed, err)
}
trendList := make([]model.TrendItem, len(stats))
for i, st := range stats {
trendList[i] = model.TrendItem{
Date: st.Date,
Count: st.PV,
}
}
return model.LogsAnalyticsResponse{
Trend: trendList,
Browsers: []model.BrowserItem{},
TopUsers: []model.TopUserItem{},
}, nil
}
// findUserIDsByUsername resolves the user id filter behind a username search term.
func findUserIDsByUsername(ctx context.Context, username string) ([]uint64, error) {
if userSvc := GetUserService(ctx); userSvc != nil {
users, _, err := userSvc.ListUsers(ctx, 1, userQueryMaxLimit, username)
if err != nil {
return nil, fmt.Errorf(errs.ErrQueryUserFailed, err)
}
ids := make([]uint64, 0, len(users))
for _, u := range users {
ids = append(ids, u.ID)
}
return ids, nil
}
ids, err := repository.SearchUserIDsByUsername(ctx, username)
if err != nil {
return nil, fmt.Errorf(errs.ErrQueryUserFailed, err)
}
return ids, nil
}
const userQueryMaxLimit = 100
func buildAccessLogFilter(ctx context.Context, q model.AccessLogQuery) (contracts.AccessLogFilterDTO, error) {
filter := contracts.AccessLogFilterDTO{}
if q.Username != "" {
userIDs, err := findUserIDsByUsername(ctx, q.Username)
if err != nil {
return filter, err
}
filter.UserIDs = userIDs
}
if q.Path != "" {
filter.Path = q.Path
}
if q.StartTime != "" {
if t, err := parseAccessLogTime(q.StartTime); err == nil {
filter.StartTime = &t
}
}
if q.EndTime != "" {
if t, err := parseAccessLogTime(q.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)
}
// enrichAccessLogsWithUsers attaches usernames and nicknames to access log rows.
func enrichAccessLogsWithUsers(ctx context.Context, list []model.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]repository.UserDisplayName, len(userIDs))
if userSvc := GetUserService(ctx); userSvc != nil {
for _, uid := range userIDs {
if u, err := userSvc.GetUserByID(ctx, uid); err == nil && u != nil {
userMap[uid] = repository.UserDisplayName{Username: u.Username, Nickname: u.Nickname}
}
}
} else if names, err := repository.LoadUserDisplayNames(ctx, userIDs); err == nil {
userMap = names
}
for i := range list {
if info, ok := userMap[list[i].UserID]; ok {
list[i].Username = info.Username
list[i].Nickname = info.Nickname
}
}
}
@@ -0,0 +1,174 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/core/contracts"
"Wavelet/pkg/config"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
"context"
"encoding/json"
"errors"
"fmt"
)
const (
// LogDBSwitchTask 切换日志数据库任务标识。
LogDBSwitchTask = "logs:db_switch"
// TaskTypeLogDBSwitch 管理端任务类型。
TaskTypeLogDBSwitch = "logs_db_switch"
targetPostgres = "postgres"
targetSQLite = "sqlite"
targetClickHouse = "clickhouse"
errParseTaskPayloadFailed = "参数解析失败: %w"
errInvalidLogTarget = "目标日志库不合法: %s"
)
// LogDBSwitchMeta 描述切换日志数据库任务。
var LogDBSwitchMeta = contracts.TaskMetaDTO{
Name: LogDBSwitchTask,
DisplayName: "切换日志数据库",
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
MaxRetry: 3,
Queue: "default",
Params: []contracts.TaskParamDTO{
{Name: "target", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse", Type: "string", Required: true},
},
}
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(errParseTaskPayloadFailed, err)
}
p.Target = normalizeTarget(p.Target)
if !validTarget(p.Target) {
return nil, fmt.Errorf(errInvalidLogTarget, 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) (*contracts.TaskResultDTO, error) {
var p logDBSwitchPayload
if err := json.Unmarshal(payload, &p); err != nil {
return nil, fmt.Errorf(errParseTaskPayloadFailed, err)
}
p.Target = normalizeTarget(p.Target)
if err := validateSwitch(ctx, p.Target); err != nil {
return nil, err
}
source, err := currentLogDatabase(ctx)
if err != nil {
return nil, err
}
taskSvc := GetTaskService()
if taskSvc != nil {
taskSvc.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
}
if err := setMigrationFlag(ctx, logMigrationInProgress); err != nil {
return nil, err
}
defer func() {
if err := setMigrationFlag(ctx, ""); err != nil {
logger.ErrorF(ctx, "清除日志迁移冻结标记失败: %v", err)
}
}()
rc := GetRiskControlService()
if rc != nil {
if err := rc.SwitchLogEngine(ctx, p.Target); err != nil {
return nil, err
}
}
if err := flipLogDatabase(ctx, p.Target); err != nil {
return nil, err
}
if taskSvc != nil {
taskSvc.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
}
return &contracts.TaskResultDTO{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(errs.ErrSameLogTarget)
}
switch target {
case targetClickHouse:
if !config.Config.ClickHouse.Enabled {
return errors.New(errs.ErrClickHouseNotEnabled)
}
case targetPostgres:
if !config.Config.Database.Enabled {
return errors.New(errs.ErrPostgresNotEnabled)
}
case targetSQLite:
if config.Config.Database.Enabled {
return errors.New(errs.ErrSQLiteNotAllowedAsLogDB)
}
}
return nil
}
func currentLogDatabase(ctx context.Context) (string, error) {
cfg, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyLogDatabase)
if err != nil {
return "", fmt.Errorf(errs.ErrReadLogDatabaseFailed, err)
}
if cfg.Value == "" {
return "", errors.New(errs.ErrLogDatabaseEmpty)
}
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)
}
@@ -3,7 +3,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
package service
import (
"Wavelet/pkg/logger"
@@ -16,7 +16,8 @@ import (
const installedBinaryMode = 0o755
func replaceAndRestart(executable, stagedBinary string) error {
// ReplaceAndRestart replaces the current executable binary with the staged binary and restarts via syscall.Exec.
func ReplaceAndRestart(executable, stagedBinary string) error {
ctx := context.Background()
logger.InfoF(ctx, "[Updater] Swapping executable: %s -> %s", executable, stagedBinary)
backup := executable + ".old"
@@ -0,0 +1,16 @@
//go:build windows
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/plugins/domain/admin/errs"
"errors"
)
// ReplaceAndRestart is blocked on Windows.
func ReplaceAndRestart(_, _ string) error {
return errors.New(errs.ErrAutomaticUpgradeBlocked)
}
@@ -1,11 +1,15 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
// Package service provides business logic and orchestration for the admin domain.
package service
import (
"Wavelet/core/contracts"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/repository"
"context"
"errors"
"sync"
"gorm.io/gorm"
@@ -28,6 +32,7 @@ func SetDBService(s contracts.DBService) {
servicesMu.Lock()
defer servicesMu.Unlock()
dbService = s
repository.SetDBService(s)
}
// SetCacheService injects the CacheService contract.
@@ -35,6 +40,7 @@ func SetCacheService(s contracts.CacheService) {
servicesMu.Lock()
defer servicesMu.Unlock()
cacheService = s
repository.SetCacheService(s)
}
// SetUserService injects the UserService contract.
@@ -101,6 +107,7 @@ func ResetServices() {
storageSvc = nil
riskControlService = nil
eventEmitter = nil
repository.ResetServices()
}
// GetDB returns the GORM DB instance bound to the context if available.
@@ -113,24 +120,21 @@ func GetDB(ctx context.Context) *gorm.DB {
return dbService.DB(ctx)
}
// GetCache returns the unified CacheService instance. ctx is kept for
// signature symmetry with the other context-aware accessors.
// GetCache returns the unified CacheService instance.
func GetCache(_ context.Context) contracts.CacheService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return cacheService
}
// GetUserService returns the UserService instance. ctx is kept for
// signature symmetry with the other context-aware accessors.
// GetUserService returns the UserService instance.
func GetUserService(_ context.Context) contracts.UserService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return userService
}
// GetAuthService returns the AuthService instance. ctx is kept for
// signature symmetry with the other context-aware accessors.
// GetAuthService returns the AuthService instance.
func GetAuthService(_ context.Context) contracts.AuthService {
servicesMu.RLock()
defer servicesMu.RUnlock()
@@ -157,3 +161,44 @@ func GetRiskControlService() contracts.RiskControlService {
defer servicesMu.RUnlock()
return riskControlService
}
// translateNotFound collapses the persistence layer's record-not-found sentinel into
// the plugin's own domain error so that no layer above the repository has to import gorm.
func translateNotFound(err error, notFound error) error {
if errors.Is(err, gorm.ErrRecordNotFound) {
return notFound
}
return err
}
// isRecordMissing reports whether err originates from a missing persistence row.
func isRecordMissing(err error) bool {
return errors.Is(err, gorm.ErrRecordNotFound)
}
// requireUserService resolves the injected user contract service.
func requireUserService(ctx context.Context) (contracts.UserService, error) {
userSvc := GetUserService(ctx)
if userSvc == nil {
return nil, errs.ErrUserServiceUnavailable
}
return userSvc, nil
}
// requireAuthService resolves the injected auth contract service.
func requireAuthService(ctx context.Context) (contracts.AuthService, error) {
authSvc := GetAuthService(ctx)
if authSvc == nil {
return nil, errs.ErrAuthServiceUnavailable
}
return authSvc, nil
}
// requireTaskService resolves the injected task contract service.
func requireTaskService() (contracts.TaskService, error) {
taskSvc := GetTaskService()
if taskSvc == nil {
return nil, errs.ErrTaskServiceUnavailable
}
return taskSvc, nil
}
@@ -0,0 +1,164 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/pkg/config"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
"context"
"fmt"
"math"
"runtime"
"time"
)
var startTime = time.Now()
const (
hoursInDay = 24
minutesInHour = 60
secondsInMinute = 60
nanosPerSecond = 1e9
logDBNamePostgres = "postgres"
logDBNameSQLite = "sqlite"
logDBNameClickHouse = "clickhouse"
defaultLogRetentionDays = 30
logMigrationIdle = "idle"
logMigrationInProgress = "migrating"
unknownGCLabel = "未知"
noGCLabel = "无"
)
// CollectSystemStatus samples the Go runtime counters for the console status page.
func CollectSystemStatus() model.SystemStatusResponse {
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 = unknownGCLabel
default:
lastGCTime = noGCLabel
}
var lastPause string
if m.NumGC > 0 {
lastPause = fmt.Sprintf("%.3fs", float64(m.PauseNs[(m.NumGC-1)%256])/nanosPerSecond)
} else {
lastPause = "0.000s"
}
return model.SystemStatusResponse{
Uptime: uptime,
NumGoroutine: numGoroutine,
Alloc: model.FormatBytes(m.Alloc),
TotalAlloc: model.FormatBytes(m.TotalAlloc),
Sys: model.FormatBytes(m.Sys),
Lookups: m.Lookups,
Mallocs: m.Mallocs,
Frees: m.Frees,
HeapAlloc: model.FormatBytes(m.HeapAlloc),
HeapSys: model.FormatBytes(m.HeapSys),
HeapIdle: model.FormatBytes(m.HeapIdle),
HeapInuse: model.FormatBytes(m.HeapInuse),
HeapReleased: model.FormatBytes(m.HeapReleased),
HeapObjects: m.HeapObjects,
StackInuse: model.FormatBytes(m.StackInuse),
StackSys: model.FormatBytes(m.StackSys),
MSpanInuse: model.FormatBytes(m.MSpanInuse),
MSpanSys: model.FormatBytes(m.MSpanSys),
MCacheInuse: model.FormatBytes(m.MCacheInuse),
MCacheSys: model.FormatBytes(m.MCacheSys),
BuckHashSys: model.FormatBytes(m.BuckHashSys),
GCSys: model.FormatBytes(m.GCSys),
OtherSys: model.FormatBytes(m.OtherSys),
NextGC: model.FormatBytes(m.NextGC),
LastGCTime: lastGCTime,
PauseTotalNs: fmt.Sprintf("%.1fs", float64(m.PauseTotalNs)/nanosPerSecond),
LastPause: lastPause,
NumGC: m.NumGC,
}
}
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
}
// LogDatabaseStatus reports the active log engine, migration freeze state and retention.
func LogDatabaseStatus(ctx context.Context) model.LogDatabaseStatus {
activeDB := logDBNameSQLite
migration := logMigrationIdle
if rc := GetRiskControlService(); rc != nil {
activeDB = rc.ActiveLogEngine(ctx)
if rc.IsLogEngineMigrating(ctx) {
migration = logMigrationInProgress
}
}
return model.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 !isRecordMissing(err) {
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{}
}
@@ -1,9 +1,12 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
package service_test
import (
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
"Wavelet/plugins/domain/admin/service"
"context"
"testing"
"time"
@@ -41,12 +44,12 @@ func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
if err != nil {
t.Fatalf("gorm.Open(sqlite) error = %v", err)
}
if err := sqliteDB.AutoMigrate(&SystemConfig{}); err != nil {
if err := sqliteDB.AutoMigrate(&model.SystemConfig{}); err != nil {
t.Fatalf("AutoMigrate(SystemConfig) error = %v", err)
}
siteConfig := SystemConfig{
Key: ConfigKeySiteName,
siteConfig := model.SystemConfig{
Key: model.ConfigKeySiteName,
Value: "Wavelet",
Type: "system",
Description: "系统平台的展示名称",
@@ -55,19 +58,19 @@ func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
t.Fatalf("Create(site_name) error = %v", err)
}
SetDBService(&testDBService{db: sqliteDB})
service.SetDBService(&testDBService{db: sqliteDB})
cleanup := func() {
StopSystemConfigCacheListener()
ResetSystemConfigRAMCacheForTest()
ResetServices()
repository.StopSystemConfigCacheListener()
repository.ResetSystemConfigRAMCacheForTest()
service.ResetServices()
}
return sqliteDB, cleanup
}
func TestListSystemConfigsByKeys_EmptyKeys(t *testing.T) {
result, err := ListSystemConfigsByKeys(context.Background(), nil)
result, err := repository.ListSystemConfigsByKeys(context.Background(), nil)
if err != nil {
t.Fatalf("ListSystemConfigsByKeys(nil) error = %v", err)
}
@@ -81,10 +84,10 @@ func TestListSystemConfigsByKeys_LoadsFromRAMCache(t *testing.T) {
defer cleanup()
ctx := context.Background()
ResetSystemConfigRAMCacheForTest()
repository.ResetSystemConfigRAMCacheForTest()
// Initial load
warm, err := GetSystemConfigByKey(ctx, ConfigKeySiteName)
warm, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey(site_name) warm error = %v", err)
}
@@ -93,19 +96,19 @@ func TestListSystemConfigsByKeys_LoadsFromRAMCache(t *testing.T) {
}
// Update DB directly
if err := dbConn.Model(&SystemConfig{}).
Where("key = ?", ConfigKeySiteName).
if err := dbConn.Model(&model.SystemConfig{}).
Where("key = ?", model.ConfigKeySiteName).
Update("value", "db_only_value").Error; err != nil {
t.Fatalf("Update(site_name) error = %v", err)
}
// Fetch via ListSystemConfigsByKeys should serve from local store (meaning the old value "Wavelet")
configs, err := ListSystemConfigsByKeys(ctx, []string{ConfigKeySiteName})
configs, err := repository.ListSystemConfigsByKeys(ctx, []string{model.ConfigKeySiteName})
if err != nil {
t.Fatalf("ListSystemConfigsByKeys(site_name) error = %v", err)
}
sc, ok := configs[ConfigKeySiteName]
sc, ok := configs[model.ConfigKeySiteName]
if !ok {
t.Fatal("ListSystemConfigsByKeys(site_name) missing site_name entry")
}
@@ -119,10 +122,10 @@ func TestGetSystemConfigByGroupAndInvalidation(t *testing.T) {
defer cleanup()
ctx := context.Background()
ResetSystemConfigRAMCacheForTest()
repository.ResetSystemConfigRAMCacheForTest()
// Get via specific group/type
cfg, err := GetSystemConfigByGroup(ctx, ConfigCacheType, ConfigKeySiteName)
cfg, err := repository.GetSystemConfigByGroup(ctx, repository.ConfigCacheType, model.ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByGroup error = %v", err)
}
@@ -131,14 +134,14 @@ func TestGetSystemConfigByGroupAndInvalidation(t *testing.T) {
}
// Direct DB update
if err := dbConn.Model(&SystemConfig{}).
Where("key = ?", ConfigKeySiteName).
if err := dbConn.Model(&model.SystemConfig{}).
Where("key = ?", model.ConfigKeySiteName).
Update("value", "new_site_name").Error; err != nil {
t.Fatalf("DB Update error = %v", err)
}
// Invalidate
if err := InvalidateSystemConfigCache(ctx, ConfigKeySiteName); err != nil {
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
t.Fatalf("InvalidateSystemConfigCache error = %v", err)
}
@@ -146,7 +149,7 @@ func TestGetSystemConfigByGroupAndInvalidation(t *testing.T) {
time.Sleep(100 * time.Millisecond)
// Fetch again
updated, err := GetSystemConfigByKey(ctx, ConfigKeySiteName)
updated, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey error = %v", err)
}
@@ -0,0 +1,218 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
"context"
"fmt"
"strings"
"github.com/robfig/cron/v3"
)
// ListTaskTypes returns every dispatchable task type declared in the task registry.
func ListTaskTypes() []contracts.TaskMetaDTO {
taskSvc := GetTaskService()
if taskSvc == nil {
return []contracts.TaskMetaDTO{}
}
return taskSvc.ListTasks()
}
// DispatchTask validates and enqueues a manual task run, returning the new task id.
func DispatchTask(ctx context.Context, req model.DispatchTaskRequest) (string, error) {
taskSvc, err := requireTaskService()
if err != nil {
return "", err
}
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
if !ok {
return "", errs.ErrInvalidTaskType
}
validated, err := validateTaskPayload(taskSvc, meta.Name, req.Payload)
if err != nil {
return "", err
}
taskID, err := taskSvc.Dispatch(ctx, req.TaskType, validated, "manual")
if err != nil {
return "", fmt.Errorf("%s: %w", errs.TaskDispatchFailed, err)
}
return taskID, nil
}
// validateTaskPayload normalises an optional raw payload through the task registry.
func validateTaskPayload(taskSvc contracts.TaskService, name, payload string) ([]byte, error) {
var payloadBytes []byte
if strings.TrimSpace(payload) != "" {
payloadBytes = []byte(payload)
}
validated, err := taskSvc.ValidatePayload(name, payloadBytes)
if err != nil {
return nil, errs.NewInvalidInputError(err.Error())
}
return validated, nil
}
// ListTaskExecutions pages task execution records for the console.
func ListTaskExecutions(
ctx context.Context,
req model.ListTaskExecutionsRequest,
) ([]model.TaskExecution, int64, error) {
if req.TaskType != "" {
if taskSvc := GetTaskService(); taskSvc != nil {
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
req.TaskType = meta.Name
}
}
}
executions, total, err := repository.ListTaskExecutionRecords(ctx, req)
if err != nil {
return nil, 0, err
}
return executions, total, nil
}
// TaskExecution loads a single execution record including its buffered log.
func TaskExecution(ctx context.Context, id uint64) (*model.TaskExecution, error) {
return repository.GetTaskExecutionByID(ctx, id)
}
// RetryTask re-dispatches a failed execution as a new task run.
func RetryTask(ctx context.Context, id uint64) (string, error) {
taskSvc, err := requireTaskService()
if err != nil {
return "", err
}
newTaskID, err := taskSvc.Retry(ctx, id)
if err != nil {
return "", err
}
return newTaskID, nil
}
// IsRetryConflictError reports whether the task registry rejected the retry request
// because of the record state rather than an infrastructure failure.
func IsRetryConflictError(err error) bool {
msg := err.Error()
return strings.Contains(msg, errs.RemoteTaskNotFailedMsg) ||
strings.Contains(msg, errs.RemoteTaskNotRetryableMsg) ||
strings.Contains(msg, errs.RemoteTaskMaxRetryMsg)
}
// IsRetryMissingError reports whether the referenced execution record is absent.
func IsRetryMissingError(err error) bool {
return strings.Contains(err.Error(), errs.RemoteTaskNotFoundMsg)
}
// ListSchedules returns every dynamic schedule definition.
func ListSchedules(ctx context.Context) ([]model.Schedule, error) {
return repository.ListSchedulesRecord(ctx)
}
// CreateSchedule validates a schedule definition, persists it and reloads the scheduler.
func CreateSchedule(ctx context.Context, req model.CreateScheduleRequest) (*model.Schedule, error) {
if _, err := cron.ParseStandard(req.Cron); err != nil {
return nil, errs.ErrInvalidCronExpression
}
taskSvc, err := requireTaskService()
if err != nil {
return nil, err
}
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
if !ok {
return nil, errs.ErrInvalidTaskType
}
validated, err := validateTaskPayload(taskSvc, meta.Name, req.Payload)
if err != nil {
return nil, err
}
schedule := &model.Schedule{
Name: req.Name,
TaskType: req.TaskType,
Cron: req.Cron,
Payload: string(validated),
IsActive: *req.IsActive,
}
if err := repository.CreateScheduleRecord(ctx, schedule); err != nil {
return nil, fmt.Errorf("%s: %w", errs.ScheduleSaveFailed, err)
}
reloadScheduler(ctx, taskSvc)
return schedule, nil
}
// UpdateSchedule rewrites an existing schedule definition and reloads the scheduler.
func UpdateSchedule(ctx context.Context, id uint64, req model.UpdateScheduleRequest) (*model.Schedule, error) {
schedule, err := repository.GetScheduleByID(ctx, id)
if err != nil {
return nil, errs.ErrScheduleNotFound
}
if _, err := cron.ParseStandard(req.Cron); err != nil {
return nil, errs.ErrInvalidCronExpression
}
taskSvc, err := requireTaskService()
if err != nil {
return nil, err
}
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
if !ok {
return nil, errs.ErrInvalidTaskType
}
validated, err := validateTaskPayload(taskSvc, meta.Name, req.Payload)
if err != nil {
return nil, err
}
schedule.Name = req.Name
schedule.TaskType = req.TaskType
schedule.Cron = req.Cron
schedule.Payload = string(validated)
schedule.IsActive = *req.IsActive
if err := repository.UpdateScheduleRecord(ctx, schedule); err != nil {
return nil, fmt.Errorf("%s: %w", errs.ScheduleSaveFailed, err)
}
reloadScheduler(ctx, taskSvc)
return schedule, nil
}
// DeleteSchedule removes a schedule definition and reloads the scheduler.
func DeleteSchedule(ctx context.Context, id uint64) error {
if err := repository.DeleteScheduleRecord(ctx, id); err != nil {
return fmt.Errorf("%s: %w", errs.ScheduleDeleteFailed, err)
}
if taskSvc := GetTaskService(); taskSvc != nil {
reloadScheduler(ctx, taskSvc)
}
return nil
}
// reloadScheduler triggers the hot reload, degrading gracefully when the scheduler rejects it.
func reloadScheduler(ctx context.Context, taskSvc contracts.TaskService) {
if err := taskSvc.ReloadScheduler(); err != nil {
logger.ErrorF(ctx, "[TaskAdmin] 重载调度器失败: %v", err)
}
}
@@ -0,0 +1,86 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
"context"
)
// CreateTemplate persists a new notification template after key collision and field checks.
func CreateTemplate(ctx context.Context, req model.CreateTemplateRequest) (model.Template, error) {
exists, err := repository.TemplateExistsByKey(ctx, req.Key)
if err != nil {
return model.Template{}, err
}
if exists {
return model.Template{}, errs.ErrTemplateKeyExists
}
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.CreateTemplateRecord(ctx, &tmpl); err != nil {
return model.Template{}, err
}
return tmpl, nil
}
// ListTemplates returns every notification template.
func ListTemplates(ctx context.Context) ([]model.Template, error) {
return repository.ListTemplatesRecord(ctx)
}
// GetTemplate loads a template by its identifier.
func GetTemplate(ctx context.Context, key string) (model.Template, error) {
tmpl, err := repository.GetTemplateByKey(ctx, key)
if err != nil {
return model.Template{}, translateNotFound(err, errs.ErrTemplateNotFound)
}
return tmpl, nil
}
// UpdateTemplate rewrites the mutable fields of an existing template.
func UpdateTemplate(ctx context.Context, key string, req model.UpdateTemplateRequest) (model.Template, error) {
tmpl, err := GetTemplate(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.SaveTemplateRecord(ctx, &tmpl); err != nil {
return model.Template{}, err
}
return tmpl, nil
}
// DeleteTemplate removes a custom template; system presets are protected.
func DeleteTemplate(ctx context.Context, key string) error {
tmpl, err := GetTemplate(ctx, key)
if err != nil {
return err
}
if tmpl.IsSystem {
return errs.ErrSystemTemplateCannotDelete
}
return repository.DeleteTemplateRecord(ctx, &tmpl)
}
@@ -1,13 +1,14 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
package service
import (
"Wavelet/pkg/buildinfo"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
"archive/tar"
"archive/zip"
"compress/gzip"
@@ -25,7 +26,6 @@ import (
"sync"
"time"
"github.com/gin-gonic/gin"
"golang.org/x/mod/semver"
)
@@ -57,90 +57,22 @@ type githubRelease struct {
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 {
// UpdaterManager manages application binary updates from GitHub releases.
type UpdaterManager struct {
client releaseClient
mu sync.Mutex
upgrading bool
}
var defaultUpdaterManager = &updaterManager{
// DefaultUpdaterManager is the default singleton update manager.
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" {
@@ -158,7 +90,7 @@ func normalizeVersion(version string) string {
func parseRepository(raw string) (string, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return "", errors.New(errInvalidRepository)
return "", errors.New(errs.ErrInvalidRepository)
}
if !strings.Contains(raw, "://") {
@@ -166,16 +98,16 @@ func parseRepository(raw string) (string, error) {
if len(strings.Split(repo, "/")) == repositoryParts {
return repo, nil
}
return "", errors.New(errInvalidRepository)
return "", errors.New(errs.ErrInvalidRepository)
}
parsed, err := url.Parse(raw)
if err != nil || !strings.EqualFold(parsed.Hostname(), "github.com") {
return "", errors.New(errInvalidRepository)
return "", errors.New(errs.ErrInvalidRepository)
}
repo := strings.TrimSuffix(strings.Trim(parsed.Path, "/"), ".git")
if len(strings.Split(repo, "/")) != repositoryParts {
return "", errors.New(errInvalidRepository)
return "", errors.New(errs.ErrInvalidRepository)
}
return repo, nil
}
@@ -188,9 +120,9 @@ func expectedAssetName(tag string) string {
return fmt.Sprintf("wavelet_%s_%s_%s.%s", tag, runtime.GOOS, runtime.GOARCH, extension)
}
func expectedAssetNames(repository, tag string) []string {
func expectedAssetNames(repo, tag string) []string {
names := []string{expectedAssetName(tag)}
if parts := strings.Split(repository, "/"); len(parts) == repositoryParts {
if parts := strings.Split(repo, "/"); len(parts) == repositoryParts {
repoName := parts[1]
if repoName != "wavelet" {
extension := "tar.gz"
@@ -203,7 +135,7 @@ func expectedAssetNames(repository, tag string) []string {
return names
}
func selectLatestRelease(repository string, releases []githubRelease) (githubRelease, releaseAsset, error) {
func selectLatestRelease(repo string, releases []githubRelease) (githubRelease, releaseAsset, error) {
var selected githubRelease
var selectedAsset releaseAsset
selectedVersion := ""
@@ -213,7 +145,7 @@ func selectLatestRelease(repository string, releases []githubRelease) (githubRel
if release.Draft || version == "" {
continue
}
expectedNames := expectedAssetNames(repository, release.TagName)
expectedNames := expectedAssetNames(repo, release.TagName)
for _, asset := range release.Assets {
matched := false
for _, name := range expectedNames {
@@ -234,20 +166,20 @@ func selectLatestRelease(repository string, releases []githubRelease) (githubRel
}
if selectedVersion == "" {
return githubRelease{}, releaseAsset{}, errors.New(errNoCompatibleRelease)
return githubRelease{}, releaseAsset{}, errors.New(errs.ErrNoCompatibleRelease)
}
return selected, selectedAsset, nil
}
func (m *updaterManager) fetchRelease(ctx context.Context, repository string) (githubRelease, releaseAsset, error) {
func (m *UpdaterManager) fetchRelease(ctx context.Context, repo string) (githubRelease, releaseAsset, error) {
req, err := http.NewRequestWithContext(
ctx,
http.MethodGet,
fmt.Sprintf("%s/repos/%s/releases?per_page=30", githubAPIBaseURL, repository),
fmt.Sprintf("%s/repos/%s/releases?per_page=30", githubAPIBaseURL, repo),
nil,
)
if err != nil {
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseRequestFailed, err)
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errs.ErrReleaseRequestFailed, err)
}
req.Header.Set("Accept", "application/vnd.github+json")
req.Header.Set("User-Agent", "Wavelet-Updater")
@@ -255,22 +187,22 @@ func (m *updaterManager) fetchRelease(ctx context.Context, repository string) (g
resp, err := m.client.Do(req)
if err != nil {
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseRequestFailed, err)
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errs.ErrReleaseRequestFailed, err)
}
defer func() {
_ = resp.Body.Close()
}()
if resp.StatusCode != http.StatusOK {
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: HTTP %d", errReleaseRequestFailed, resp.StatusCode)
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: HTTP %d", errs.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)
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errs.ErrReleaseResponseInvalid, err)
}
release, asset, err := selectLatestRelease(repository, releases)
release, asset, err := selectLatestRelease(repo, releases)
if err != nil {
return githubRelease{}, releaseAsset{}, err
}
@@ -279,21 +211,22 @@ func (m *updaterManager) fetchRelease(ctx context.Context, repository string) (g
}
func loadRepository(ctx context.Context) (string, error) {
config, err := GetSystemConfigByKey(ctx, ConfigKeyUpdateUpstreamRepository)
cfg, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUpdateUpstreamRepository)
if err != nil {
return "", fmt.Errorf("%s: %w", errInvalidRepository, err)
return "", fmt.Errorf("%s: %w", errs.ErrInvalidRepository, err)
}
return parseRepository(config.Value)
return parseRepository(cfg.Value)
}
func (m *updaterManager) status(ctx context.Context) (UpdaterStatus, releaseAsset, error) {
// status returns current version and update status.
func (m *UpdaterManager) status(ctx context.Context) (model.UpdaterStatus, releaseAsset, error) {
upstreamRepo, err := loadRepository(ctx)
if err != nil {
return UpdaterStatus{}, releaseAsset{}, err
return model.UpdaterStatus{}, releaseAsset{}, err
}
release, asset, err := m.fetchRelease(ctx, upstreamRepo)
if err != nil {
return UpdaterStatus{}, releaseAsset{}, err
return model.UpdaterStatus{}, releaseAsset{}, err
}
currentVersion := normalizeVersion(buildinfo.Version)
@@ -302,7 +235,7 @@ func (m *updaterManager) status(ctx context.Context) (UpdaterStatus, releaseAsse
logger.InfoF(ctx, "[Updater] Check update complete. current: %s, latest: %s, update_available: %t", buildinfo.Version, release.TagName, updateAvailable)
return UpdaterStatus{
return model.UpdaterStatus{
CurrentVersion: buildinfo.Version,
BuildTime: buildinfo.BuildTime,
LatestVersion: release.TagName,
@@ -319,44 +252,50 @@ func (m *updaterManager) status(ctx context.Context) (UpdaterStatus, releaseAsse
}, asset, nil
}
// GetUpdateStatus returns current updater status.
func GetUpdateStatus(ctx context.Context) (model.UpdaterStatus, error) {
status, _, err := DefaultUpdaterManager.status(ctx)
return status, err
}
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)
return fmt.Errorf(errs.ErrReleaseAssetSizeInvalid, 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)
return fmt.Errorf(errs.ErrCreateUpgradeRequestFailed, err)
}
req.Header.Set("User-Agent", "Wavelet-Updater")
resp, err := client.Do(req)
if err != nil {
return fmt.Errorf("下载升级资产失败: %w", err)
return fmt.Errorf(errs.ErrDownloadUpgradeAssetFailed, err)
}
defer func() {
_ = resp.Body.Close()
}()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("下载升级资产失败: HTTP %d", resp.StatusCode)
return fmt.Errorf(errs.ErrUpgradeAssetHTTPFailed, 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)
return fmt.Errorf(errs.ErrCreateUpgradeArchiveFailed, err)
}
written, err := io.Copy(file, io.LimitReader(resp.Body, maxArchiveSize+1))
if err != nil {
_ = file.Close()
return fmt.Errorf("写入升级归档失败: %w", err)
return fmt.Errorf(errs.ErrWriteUpgradeArchiveFailed, err)
}
if err := file.Close(); err != nil {
return fmt.Errorf("关闭升级归档失败: %w", err)
return fmt.Errorf(errs.ErrCloseUpgradeArchiveFailed, err)
}
if written > maxArchiveSize || written != asset.Size {
return fmt.Errorf("升级归档大小不匹配: got %d, want %d", written, asset.Size)
return fmt.Errorf(errs.ErrUpgradeArchiveSizeMismatch, written, asset.Size)
}
logger.InfoF(ctx, "[Updater] Successfully downloaded release asset to %s", destination)
return nil
@@ -365,12 +304,12 @@ func downloadArchive(ctx context.Context, client releaseClient, asset releaseAss
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)
return "", fmt.Errorf(errs.ErrArchiveContainsIllegalPath, 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 "", fmt.Errorf(errs.ErrArchivePathOutOfDestination, name)
}
return target, nil
}
@@ -390,7 +329,7 @@ func matchBinaryName(name string, candidates []string) bool {
return false
}
func getCandidateBinaryNames(executable, repository string) []string {
func getCandidateBinaryNames(executable, repo string) []string {
execName := filepath.Base(executable)
names := []string{execName}
@@ -407,7 +346,7 @@ func getCandidateBinaryNames(executable, repository string) []string {
names = append(names, name)
}
if parts := strings.Split(repository, "/"); len(parts) == repositoryParts {
if parts := strings.Split(repo, "/"); len(parts) == repositoryParts {
addName(parts[1])
}
addName("wavelet")
@@ -480,7 +419,7 @@ func findBinaryInTarGz(archivePath string, candidates []string) (string, error)
}
}
return "", errors.New(errNoCompatibleAsset)
return "", errors.New(errs.ErrNoCompatibleAsset)
}
func findBinaryInZip(archivePath string, candidates []string) (string, error) {
@@ -509,7 +448,7 @@ func findBinaryInZip(archivePath string, candidates []string) (string, error) {
}
}
return "", errors.New(errNoCompatibleAsset)
return "", errors.New(errs.ErrNoCompatibleAsset)
}
func extractTarGz(ctx context.Context, archivePath, destination, targetName string, candidates []string) (string, error) {
@@ -565,12 +504,12 @@ func extractTarGz(ctx context.Context, archivePath, destination, targetName stri
return "", closeErr
}
if written > maxArchiveSize {
return "", errors.New("解压后的程序文件超过大小限制")
return "", errors.New(errs.ErrExtractedBinaryTooLarge)
}
logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target)
return target, nil
}
return "", errors.New(errNoCompatibleAsset)
return "", errors.New(errs.ErrNoCompatibleAsset)
}
func extractZip(ctx context.Context, archivePath, destination, targetName string, candidates []string) (string, error) {
@@ -617,26 +556,27 @@ func extractZip(ctx context.Context, archivePath, destination, targetName string
return "", outputCloseErr
}
if written > maxArchiveSize {
return "", errors.New("解压后的程序文件超过大小限制")
return "", errors.New(errs.ErrExtractedBinaryTooLarge)
}
logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target)
return target, nil
}
return "", errors.New(errNoCompatibleAsset)
return "", errors.New(errs.ErrNoCompatibleAsset)
}
func (m *updaterManager) prepareUpgrade(ctx context.Context) (string, string, error) {
// PrepareUpgrade validates preconditions and downloads the newest binary.
func (m *UpdaterManager) PrepareUpgrade(ctx context.Context) (string, string, error) {
if runtime.GOOS == windowsOS {
return "", "", errors.New(errAutomaticUpgradeBlocked)
return "", "", errors.New(errs.ErrAutomaticUpgradeBlocked)
}
if normalizeVersion(buildinfo.Version) == "" {
return "", "", errors.New(errDevelopmentBuild)
return "", "", errors.New(errs.ErrDevelopmentBuild)
}
m.mu.Lock()
defer m.mu.Unlock()
if m.upgrading {
return "", "", errors.New(errUpgradeAlreadyRunning)
return "", "", errors.New(errs.ErrUpgradeAlreadyRunning)
}
status, asset, err := m.status(ctx)
@@ -644,23 +584,23 @@ func (m *updaterManager) prepareUpgrade(ctx context.Context) (string, string, er
return "", "", err
}
if !status.UpdateAvailable {
return "", "", errors.New(errAlreadyUpToDate)
return "", "", errors.New(errs.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)
return "", "", fmt.Errorf(errs.ErrLocateExecutableFailed, err)
}
executable, err = filepath.EvalSymlinks(executable)
if err != nil {
return "", "", fmt.Errorf("解析当前程序路径失败: %w", err)
return "", "", fmt.Errorf(errs.ErrResolveExecutablePathFailed, err)
}
tempDir, err := os.MkdirTemp(filepath.Dir(executable), ".wavelet-update-*")
if err != nil {
return "", "", fmt.Errorf("创建升级目录失败: %w", err)
return "", "", fmt.Errorf(errs.ErrCreateUpgradeDirFailed, err)
}
archivePath := filepath.Join(tempDir, asset.Name)
@@ -680,14 +620,15 @@ func (m *updaterManager) prepareUpgrade(ctx context.Context) (string, string, er
}
if err != nil {
_ = os.RemoveAll(tempDir)
return "", "", fmt.Errorf("解压升级资产失败: %w", err)
return "", "", fmt.Errorf(errs.ErrExtractUpgradeAssetFailed, err)
}
logger.InfoF(ctx, "[Updater] Staged binary successfully prepared: %s", stagedBinary)
m.upgrading = true
return executable, stagedBinary, nil
}
func (m *updaterManager) finishUpgrade() {
// FinishUpgrade resets the upgrading flag.
func (m *UpdaterManager) FinishUpgrade() {
m.mu.Lock()
defer m.mu.Unlock()
m.upgrading = false
@@ -0,0 +1,120 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"context"
"errors"
)
// ToUserResponse projects the user contract DTO onto the console response shape.
func ToUserResponse(u *contracts.UserDTO) model.UserResponse {
if u == nil {
return model.UserResponse{}
}
return model.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,
}
}
// AdminListUsers pages users through the user contract service.
func AdminListUsers(
ctx context.Context,
filter contracts.AdminListUsersFilter,
) (int64, []*contracts.UserDTO, error) {
userSvc, err := requireUserService(ctx)
if err != nil {
return 0, nil, err
}
total, dtos, err := userSvc.AdminListUsers(ctx, filter)
if err != nil {
logger.ErrorF(ctx, "List admin users failed: %v", err)
return 0, nil, errors.New(errs.ListAdminUsersFailed)
}
return total, dtos, nil
}
// AdminGetUser loads a single user profile.
func AdminGetUser(ctx context.Context, id uint64) (*contracts.UserDTO, error) {
userSvc, err := requireUserService(ctx)
if err != nil {
return nil, err
}
targetUser, err := userSvc.AdminGetUser(ctx, id)
if err != nil {
return nil, translateNotFound(err, errs.ErrUserNotFound)
}
return targetUser, nil
}
// AdminUpdateUserStatus enables or disables a user account.
func AdminUpdateUserStatus(ctx context.Context, id uint64, isActive bool) error {
userSvc, err := requireUserService(ctx)
if err != nil {
return err
}
err = userSvc.AdminUpdateUserStatus(ctx, id, isActive)
return translateNotFound(err, errs.ErrUserNotFound)
}
// AdminDeleteUser removes a user on behalf of the acting administrator.
func AdminDeleteUser(ctx context.Context, operatorID, id uint64) error {
userSvc, err := requireUserService(ctx)
if err != nil {
return err
}
err = userSvc.AdminDeleteUser(ctx, operatorID, id)
return translateNotFound(err, errs.ErrUserNotFound)
}
// AdminCreateUser registers a local-password user.
func AdminCreateUser(ctx context.Context, req contracts.AdminCreateUserRequest) (*contracts.UserDTO, error) {
userSvc, err := requireUserService(ctx)
if err != nil {
return nil, err
}
newUser, err := userSvc.AdminCreateUser(ctx, req)
if err != nil {
return nil, translateNotFound(err, errs.ErrUserNotFound)
}
return newUser, nil
}
// AdminUpdateUser rewrites a user profile and optionally resets its password.
func AdminUpdateUser(
ctx context.Context,
operatorID uint64,
req contracts.AdminUpdateUserRequest,
) error {
userSvc, err := requireUserService(ctx)
if err != nil {
return err
}
err = userSvc.AdminUpdateUser(ctx, operatorID, req)
return translateNotFound(err, errs.ErrUserNotFound)
}
@@ -16,8 +16,8 @@ import (
)
func isOIDCLoginEnabled(ctx context.Context) bool {
var val string
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "oidc_login_enabled").Pluck("value", &val).Error; err != nil || val == "" {
val, err := GetSystemConfigValue(ctx, "oidc_login_enabled")
if err != nil || val == "" {
return true
}
b, err := strconv.ParseBool(val)
@@ -75,8 +75,8 @@ func activeLoginSources(ctx context.Context) []AuthSourceView {
}
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
var val string
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || strings.TrimSpace(val) == "" {
val, err := GetSystemConfigValue(ctx, "server_address")
if err != nil || strings.TrimSpace(val) == "" {
return "", errors.New(errServerAddressMissing)
}
return strings.TrimRight(val, "/") + "/login", nil
+16
View File
@@ -36,3 +36,19 @@ const (
errBannedAccount = "账号已被封禁"
errUnAuthorized = "未登录"
)
// Service 层与鉴权中间件内部错误文案(保持与重构前逐字一致)
const (
errUserNotInContext = "auth: user not found in context"
errEmptyToken = "auth: empty token" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errSystemUserTokenNotAllowed = "auth: system user token not allowed" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errUnauthorizedInternal = "unauthorized"
errSystemUserLoginNotAllowed = "system user is not allowed to login"
)
// OAuth 回调会话校验错误文案(保持与重构前逐字一致)
const (
errInvalidSessionContext = "invalid session context"
errSessionMismatchForOAuth = "session mismatch for oauth state"
errUserContextMismatch = "user context mismatch for oauth binding"
)
+21 -22
View File
@@ -9,7 +9,6 @@ import (
"Wavelet/pkg/idgen"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"context"
"errors"
"fmt"
@@ -235,17 +234,17 @@ func Callback(c *gin.Context) {
token, ok := session.Get(SessionTokenKey).(string)
if !ok || token == "" {
response.AbortBadRequest(c, "invalid session context")
response.AbortBadRequest(c, errInvalidSessionContext)
return
}
if hashSessionToken(token) != payload.SessionHash {
response.AbortBadRequest(c, "session mismatch for oauth state")
response.AbortBadRequest(c, errSessionMismatchForOAuth)
return
}
if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID {
response.AbortBadRequest(c, "user context mismatch for oauth binding")
response.AbortBadRequest(c, errUserContextMismatch)
return
}
@@ -298,8 +297,8 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource,
response.AbortUnauthorized(c, errUnAuthorized)
return
}
var user contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
user, err := GetUserByID(ctx, userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
@@ -314,41 +313,43 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource,
return
}
user.LastLoginAt = time.Now()
_ = getDB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound")))
_ = TouchUserLastLogin(ctx, user.ID, user.LastLoginAt)
c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "bound")))
}
func handleCallbackLogin(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
var user contracts.UserDTO
var user *contracts.UserDTO
account, err := FindExternalAccount(ctx, source.ID, userInfo.Sub)
switch {
case err == nil:
if loadErr := getDB(ctx).Table("w_users").Where("id = ?", account.UserID).First(&user).Error; loadErr != nil {
loaded, loadErr := 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
user = &newUser
default:
response.AbortInternal(c, err.Error())
return
}
user.LastLoginAt = time.Now()
_ = getDB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
if err := SetLoginSession(ctx, c, &user); err != nil {
_ = TouchUserLastLogin(ctx, user.ID, user.LastLoginAt)
if err := SetLoginSession(ctx, c, user); err != nil {
response.AbortInternal(c, err.Error())
return
}
SetCachedUser(ctx, user.ID, &user)
SetCachedUser(ctx, user.ID, user)
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in")))
c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "logged_in")))
}
func uniqueUsername(ctx context.Context, base string) (string, error) {
@@ -357,10 +358,8 @@ func uniqueUsername(ctx context.Context, base string) (string, error) {
base = "user"
}
var existingUsernames []string
if err := getDB(ctx).Table("w_users").
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
Pluck("username", &existingUsernames).Error; err != nil {
existingUsernames, err := ListSimilarUsernames(ctx, base)
if err != nil {
return "", err
}
@@ -385,8 +384,8 @@ func uniqueUsername(ctx context.Context, base string) (string, error) {
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) (contracts.UserDTO, bool) {
registrationEnabled := true
var val string
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "registration_enabled").Pluck("value", &val).Error; err == nil && val != "" {
val, cfgErr := GetSystemConfigValue(ctx, "registration_enabled")
if cfgErr == nil && val != "" {
if b, err := strconv.ParseBool(val); err == nil {
registrationEnabled = b
}
@@ -417,7 +416,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSou
UpdatedAt: now,
}
if err := getDB(ctx).Table("w_users").Create(&user).Error; err != nil {
if err := InsertUser(ctx, &user); err != nil {
response.AbortInternal(c, err.Error())
return contracts.UserDTO{}, false
}
+21 -20
View File
@@ -22,33 +22,35 @@ func hashToken(token string) string {
return hex.EncodeToString(h.Sum(nil))
}
// currentUserIDFromRequestContext 是接入层向 Service 层暴露的登录态桥接。
//
// Session 读取必须依赖 *gin.Context,而 Service 层禁止 import gin,
// 因此该类型断言收敛在本(接入层)文件中。ok 为 false 表示 ctx 不是 *gin.Context。
func currentUserIDFromRequestContext(ctx context.Context) (uint64, bool) {
ginCtx, ok := ctx.(*gin.Context)
if !ok {
return 0, false
}
return GetUserIDFromContext(ginCtx), true
}
func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *CachedToken, error) {
tokenHash := hashToken(tokenStr)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil || tokenRecord == nil {
var tokenRow struct {
ID uint64
UserID uint64
IsAdmin bool
}
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
tokenRecord, err = GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, nil, err
}
tokenRecord = &CachedToken{
ID: tokenRow.ID,
UserID: tokenRow.UserID,
IsAdmin: tokenRow.IsAdmin,
}
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
var userRow contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&userRow).Error; err != nil {
user, err = GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, nil, err
}
user = &userRow
SetCachedUser(ctx, tokenRecord.UserID, user)
}
@@ -74,7 +76,7 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
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")
return nil, errors.New(errSystemUserLoginNotAllowed)
}
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, true)
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin)
@@ -85,16 +87,15 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
// 降级使用 Session 鉴权
userID := GetUserIDFromContext(c)
if userID <= 0 {
return nil, errors.New("unauthorized")
return nil, errors.New(errUnauthorizedInternal)
}
user, err := GetCachedUser(ctx, userID)
if err != nil || user == nil || !user.IsActive {
var dbUser contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&dbUser).Error; err != nil {
user, err = GetActiveUserByID(ctx, userID)
if err != nil {
return nil, err
}
user = &dbUser
SetCachedUser(ctx, userID, user)
}
@@ -102,7 +103,7 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, false)
if user.Username == "system" {
return nil, errors.New("system user is not allowed to login")
return nil, errors.New(errSystemUserLoginNotAllowed)
}
return user, nil
+91
View File
@@ -6,8 +6,10 @@ package auth
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/util"
"context"
"sync"
"time"
"gorm.io/gorm"
)
@@ -58,6 +60,80 @@ func getCache(ctx context.Context) contracts.CacheService {
return s
}
// GetAccessTokenByHash 按令牌哈希读取访问令牌记录(仅取鉴权所需字段)
func GetAccessTokenByHash(ctx context.Context, tokenHash string) (*CachedToken, error) {
var row struct {
ID uint64
UserID uint64
IsAdmin bool
}
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&row).Error; err != nil {
return nil, err
}
return &CachedToken{
ID: row.ID,
UserID: row.UserID,
IsAdmin: row.IsAdmin,
}, nil
}
// GetActiveUserByID 读取仍处于启用状态的用户
func GetActiveUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
var user contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
// GetUserByID 按 ID 读取用户(不限制启用状态)
func GetUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
var user contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
// InsertUser 新建用户记录
func InsertUser(ctx context.Context, user *contracts.UserDTO) error {
return getDB(ctx).Table("w_users").Create(user).Error
}
// TouchUserLastLogin 刷新用户最后登录时间
func TouchUserLastLogin(ctx context.Context, userID uint64, at time.Time) error {
return getDB(ctx).Table("w_users").Where("id = ?", userID).Update("last_login_at", at).Error
}
// ListSimilarUsernames 查询与基础用户名相同或带 `-序号` 后缀的用户名(用于用户名去重)
func ListSimilarUsernames(ctx context.Context, base string) ([]string, error) {
var existingUsernames []string
if err := getDB(ctx).Table("w_users").
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
Pluck("username", &existingUsernames).Error; err != nil {
return nil, err
}
return existingUsernames, nil
}
// GetSystemConfigValue 读取系统配置项原始值
func GetSystemConfigValue(ctx context.Context, key string) (string, error) {
var val string
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", key).Pluck("value", &val).Error; err != nil {
return "", err
}
return val, nil
}
// ListAllAuthSources 获取全部认证源(含未启用),按 ID 升序
func ListAllAuthSources(ctx context.Context) ([]AuthSource, error) {
var sources []AuthSource
if err := getDB(ctx).Order("id ASC").Find(&sources).Error; err != nil {
return nil, err
}
return sources, nil
}
// GetAuthSourceByID 根据 ID 获取认证源
func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
var src AuthSource
@@ -85,6 +161,21 @@ func ListActiveAuthSources(ctx context.Context) ([]AuthSource, error) {
return sources, nil
}
// CreateAuthSourceRecord 新建认证源记录
func CreateAuthSourceRecord(ctx context.Context, source *AuthSource) error {
return getDB(ctx).Create(source).Error
}
// SaveAuthSourceRecord 全量保存认证源记录
func SaveAuthSourceRecord(ctx context.Context, source *AuthSource) error {
return getDB(ctx).Save(source).Error
}
// DeleteAuthSourceRecord 删除认证源记录
func DeleteAuthSourceRecord(ctx context.Context, source *AuthSource) error {
return getDB(ctx).Delete(source).Error
}
// GetActiveAuthSourcesCached 获取所有启用的认证源(带缓存或直接查询)
func GetActiveAuthSourcesCached(ctx context.Context) ([]AuthSource, error) {
return ListActiveAuthSources(ctx)
+35 -43
View File
@@ -5,12 +5,9 @@ package auth
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"context"
"errors"
"sync"
"github.com/gin-gonic/gin"
)
type authServiceImpl struct{}
@@ -27,58 +24,48 @@ func (s *authServiceImpl) RequireAdminMiddleware() any {
return AdminRequired()
}
// GetCurrentUser 从 context 中读取登录用户。
//
// 中间件通过 gin 的 c.Set(contracts.AuthUserObjKey, user) 写入登录态;
// *gin.Context 自身实现了 context.Context,且其 Value(key) 对 string 类型 key
// 等价于 c.Get(key)(未命中时再回落到 Request.Context().Value),
// 因此这里无需感知 gin 即可读取同一份登录态。
func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if ginCtx, ok := ctx.(*gin.Context); ok {
if u, ok := ginutil.GetFromContext[*contracts.UserDTO](ginCtx, contracts.AuthUserObjKey); ok && u != nil {
return u, nil
}
}
if v := ctx.Value(contracts.AuthUserObjKey); v != nil {
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
return u, nil
}
}
return nil, errors.New("auth: user not found in context")
return nil, errors.New(errUserNotInContext)
}
func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contracts.UserDTO, error) {
if token == "" {
return nil, errors.New("auth: empty token")
return nil, errors.New(errEmptyToken)
}
tokenHash := hashToken(token)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
var tokenRow struct {
ID uint64
UserID uint64
IsAdmin bool
}
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
tokenRecord, err = GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, err
}
tokenRecord = &CachedToken{
ID: tokenRow.ID,
UserID: tokenRow.UserID,
IsAdmin: tokenRow.IsAdmin,
}
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
var dbUser contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil {
user, err = GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, err
}
user = &dbUser
SetCachedUser(ctx, tokenRecord.UserID, user)
}
if user.Username == SystemUsername {
return nil, errors.New("auth: system user token not allowed")
return nil, errors.New(errSystemUserTokenNotAllowed)
}
return user, nil
@@ -93,11 +80,16 @@ func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64)
return nil
}
// GetCurrentUserID 从请求登录态中读取用户 ID。
//
// Session 读取依赖 gin,属于接入层职责,因此这里通过接入层桥接函数
// currentUserIDFromRequestContext(见 middleware.go)取值,Service 层本身不感知 gin。
func (s *authServiceImpl) GetCurrentUserID(ctx context.Context) (uint64, error) {
if ginCtx, ok := ctx.(*gin.Context); ok {
return GetUserIDFromContext(ginCtx), nil
userID, ok := currentUserIDFromRequestContext(ctx)
if !ok {
return 0, errors.New(errUserNotInContext)
}
return 0, errors.New("auth: user not found in context")
return userID, nil
}
func (s *authServiceImpl) RevokeToken(ctx context.Context, tokenHash string) error {
@@ -118,8 +110,8 @@ func (s *authServiceImpl) InvalidateCachedToken(ctx context.Context, tokenHash s
}
func (s *authServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
var sources []AuthSource
if err := getDB(ctx).Order("id ASC").Find(&sources).Error; err != nil {
sources, err := ListAllAuthSources(ctx)
if err != nil {
return nil, err
}
@@ -156,7 +148,7 @@ func (s *authServiceImpl) CreateAuthSource(ctx context.Context, source contracts
return nil, err
}
if err := getDB(ctx).Create(&model).Error; err != nil {
if err := CreateAuthSourceRecord(ctx, &model); err != nil {
return nil, err
}
@@ -165,8 +157,8 @@ func (s *authServiceImpl) CreateAuthSource(ctx context.Context, source contracts
}
func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
var existing AuthSource
if err := getDB(ctx).First(&existing, id).Error; err != nil {
existing, err := GetAuthSourceByID(ctx, id)
if err != nil {
return nil, err
}
@@ -183,36 +175,36 @@ func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, sourc
return nil, err
}
if err := getDB(ctx).Save(&existing).Error; err != nil {
if err := SaveAuthSourceRecord(ctx, existing); err != nil {
return nil, err
}
existing.Sanitize()
return toAuthSourceDTO(&existing), nil
return toAuthSourceDTO(existing), nil
}
func (s *authServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error {
var existing AuthSource
if err := getDB(ctx).First(&existing, id).Error; err != nil {
existing, err := GetAuthSourceByID(ctx, id)
if err != nil {
return err
}
return getDB(ctx).Delete(&existing).Error
return DeleteAuthSourceRecord(ctx, existing)
}
func (s *authServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
var existing AuthSource
if err := getDB(ctx).First(&existing, id).Error; err != nil {
existing, err := GetAuthSourceByID(ctx, id)
if err != nil {
return nil, err
}
existing.IsActive = !existing.IsActive
if err := getDB(ctx).Save(&existing).Error; err != nil {
if err := SaveAuthSourceRecord(ctx, existing); err != nil {
return nil, err
}
existing.Sanitize()
return toAuthSourceDTO(&existing), nil
return toAuthSourceDTO(existing), nil
}
func toAuthSourceDTO(s *AuthSource) *contracts.AuthSourceDTO {
+367
View File
@@ -0,0 +1,367 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth_test
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/auth"
"context"
"errors"
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
const (
testSessionCookieName = "auth-test-session"
errUserNotInContext = "auth: user not found in context"
)
// newTestAuthService 装配仅注册 auth 插件的 core.Context,并返回其对外契约实现。
func newTestAuthService(t *testing.T, db *gorm.DB) contracts.AuthService {
t.Helper()
ctx := core.NewContext(context.Background())
if db != nil {
core.Provide[contracts.DBService](ctx, &mockDBService{db: db})
core.Provide[contracts.CacheService](ctx, newMockCacheService())
}
require.NoError(t, auth.New().Apply(ctx))
svc, err := core.Inject[contracts.AuthService](ctx)
require.NoError(t, err)
auth.ResetAuthRAMCacheForTest()
return svc
}
// newSessionEngine 构造一个带 Session 中间件的 gin 引擎,用于走通真实登录态链路。
//
// response.Abort* 只把错误挂载到 gin 错误链,状态码由全局错误中间件渲染,
// 因此这里必须同时装配 response.ErrorHandlerMiddleware()。
func newSessionEngine() *gin.Engine {
engine := gin.New()
engine.Use(response.ErrorHandlerMiddleware())
engine.Use(sessions.Sessions(testSessionCookieName, cookie.NewStore([]byte("test-secret"))))
return engine
}
func TestGetCurrentUserFromGinContext(t *testing.T) {
gin.SetMode(gin.TestMode)
svc := newTestAuthService(t, nil)
user := &contracts.UserDTO{ID: 4242, Username: "ctx_user", IsActive: true}
t.Run("gin 上下文已由中间件写入用户时返回该用户", func(t *testing.T) {
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/user-info", nil)
c.Set(contracts.AuthUserObjKey, user)
got, err := svc.GetCurrentUser(c)
require.NoError(t, err)
assert.Same(t, user, got)
})
t.Run("开启 ContextWithFallback 时可从请求 context 回落读取", func(t *testing.T) {
// 说明:本项目引擎默认不开启 ContextWithFallback,此时 (*gin.Context).Value
// 等价于 c.Get,与改造前 ginutil.GetFromContext 的读取路径完全一致;
// 开启回落后还能额外读到写入 Request.Context() 的登录态。
reqCtx := context.WithValue(context.Background(), contracts.AuthUserObjKey, user) //nolint:staticcheck // 模拟写入请求 context 的登录态
engine := gin.New()
engine.ContextWithFallback = true
var (
gotUser *contracts.UserDTO
gotErr error
)
engine.GET("/probe", func(c *gin.Context) {
gotUser, gotErr = svc.GetCurrentUser(c)
c.Status(http.StatusNoContent)
})
engine.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/probe", nil).WithContext(reqCtx))
require.NoError(t, gotErr)
assert.Same(t, user, gotUser)
})
t.Run("未登录时报错且文案不变", func(t *testing.T) {
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/user-info", nil)
got, err := svc.GetCurrentUser(c)
require.Error(t, err)
assert.Nil(t, got)
assert.Equal(t, errUserNotInContext, err.Error())
})
t.Run("非 gin 的普通 context 仍按 Value 取值", func(t *testing.T) {
got, err := svc.GetCurrentUser(context.WithValue(context.Background(), contracts.AuthUserObjKey, user)) //nolint:staticcheck // 与中间件写入的 key 语义一致
require.NoError(t, err)
assert.Same(t, user, got)
_, err = svc.GetCurrentUser(context.Background())
require.Error(t, err)
assert.Equal(t, errUserNotInContext, err.Error())
})
}
func TestGetCurrentUserID(t *testing.T) {
gin.SetMode(gin.TestMode)
svc := newTestAuthService(t, nil)
t.Run("gin Session 中的用户 ID 可正常读取", func(t *testing.T) {
engine := newSessionEngine()
var (
gotUID uint64
gotErr error
)
engine.GET("/probe", func(c *gin.Context) {
session := sessions.Default(c)
session.Set(auth.UserIDKey, uint64(777))
require.NoError(t, session.Save())
gotUID, gotErr = svc.GetCurrentUserID(c)
c.Status(http.StatusNoContent)
})
engine.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/probe", nil))
require.NoError(t, gotErr)
assert.Equal(t, uint64(777), gotUID)
})
t.Run("gin 上下文存在但 Session 无用户时返回 0 且不报错", func(t *testing.T) {
engine := newSessionEngine()
var (
gotUID uint64
gotErr error
)
engine.GET("/probe", func(c *gin.Context) {
gotUID, gotErr = svc.GetCurrentUserID(c)
c.Status(http.StatusNoContent)
})
engine.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/probe", nil))
require.NoError(t, gotErr)
assert.Equal(t, uint64(0), gotUID)
})
t.Run("非 gin context 报错且文案不变", func(t *testing.T) {
// 即使普通 context 中已写入用户对象,该方法的 Session 语义也保持不变。
uid, err := svc.GetCurrentUserID(
context.WithValue(context.Background(), contracts.AuthUserObjKey, &contracts.UserDTO{ID: 1}), //nolint:staticcheck // 同上
)
require.Error(t, err)
assert.Equal(t, uint64(0), uid)
assert.Equal(t, errUserNotInContext, err.Error())
})
}
func TestLoginRequiredMiddlewarePopulatesServiceContext(t *testing.T) {
gin.SetMode(gin.TestMode)
db := setupTestDB(t)
require.NoError(t, db.Create(&testUser{ID: 9001, Username: "session_user", IsActive: true}).Error)
require.NoError(t, db.Create(&testUser{ID: 9002, Username: "token_user", IsActive: true}).Error)
tokenStr := "integration-secret-token"
require.NoError(t, db.Create(&testAccessToken{
ID: 9101,
UserID: 9002,
TokenHash: hashToken(tokenStr),
Name: "integration",
IsAdmin: false,
}).Error)
svc := newTestAuthService(t, db)
t.Run("Session 鉴权链路上 GetCurrentUser 与 GetCurrentUserID 一致", func(t *testing.T) {
engine := newSessionEngine()
engine.Use(func(c *gin.Context) {
session := sessions.Default(c)
session.Set(auth.UserIDKey, uint64(9001))
require.NoError(t, session.Save())
c.Next()
})
var (
gotUser *contracts.UserDTO
userErr error
gotUID uint64
uidErr error
)
engine.GET("/protected", auth.LoginRequired(), func(c *gin.Context) {
gotUser, userErr = svc.GetCurrentUser(c)
gotUID, uidErr = svc.GetCurrentUserID(c)
c.Status(http.StatusNoContent)
})
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/protected", nil))
require.Equal(t, http.StatusNoContent, recorder.Code)
require.NoError(t, userErr)
require.NotNil(t, gotUser)
assert.Equal(t, uint64(9001), gotUser.ID)
assert.Equal(t, "session_user", gotUser.Username)
require.NoError(t, uidErr)
assert.Equal(t, uint64(9001), gotUID)
})
t.Run("Access Token 鉴权链路上 GetCurrentUser 可用", func(t *testing.T) {
engine := newSessionEngine()
var (
gotUser *contracts.UserDTO
userErr error
)
engine.GET("/protected", auth.LoginRequired(), func(c *gin.Context) {
gotUser, userErr = svc.GetCurrentUser(c)
c.Status(http.StatusNoContent)
})
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
req.Header.Set("Authorization", "Bearer "+tokenStr)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
require.Equal(t, http.StatusNoContent, recorder.Code)
require.NoError(t, userErr)
require.NotNil(t, gotUser)
assert.Equal(t, uint64(9002), gotUser.ID)
assert.Equal(t, "token_user", gotUser.Username)
})
t.Run("未登录请求被中间件拒绝", func(t *testing.T) {
engine := newSessionEngine()
engine.GET("/protected", auth.LoginRequired(), func(c *gin.Context) {
c.Status(http.StatusNoContent)
})
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/protected", nil))
assert.Equal(t, http.StatusUnauthorized, recorder.Code)
})
}
// legacyGetCurrentUser 逐字复刻改造前 Service 层的取值实现
// (*gin.Context 类型断言 + ginutil.GetFromContext + ctx.Value 回落),
// 用于与新实现做 differential 等价性校验。
func legacyGetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if ginCtx, ok := ctx.(*gin.Context); ok {
if u, ok := ginutil.GetFromContext[*contracts.UserDTO](ginCtx, contracts.AuthUserObjKey); ok && u != nil {
return u, nil
}
}
if v := ctx.Value(contracts.AuthUserObjKey); v != nil {
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
return u, nil
}
}
return nil, errors.New(errUserNotInContext)
}
// legacyGetCurrentUserID 逐字复刻改造前 Service 层基于 gin Session 的实现。
func legacyGetCurrentUserID(ctx context.Context) (uint64, error) {
if ginCtx, ok := ctx.(*gin.Context); ok {
return auth.GetUserIDFromContext(ginCtx), nil
}
return 0, errors.New(errUserNotInContext)
}
func errText(err error) string {
if err == nil {
return ""
}
return err.Error()
}
// assertLoginStateParity 断言新实现与改造前实现在同一 ctx 上返回完全一致的结果与错误文案。
func assertLoginStateParity(t *testing.T, svc contracts.AuthService, ctx context.Context) {
t.Helper()
wantUser, wantUserErr := legacyGetCurrentUser(ctx)
gotUser, gotUserErr := svc.GetCurrentUser(ctx)
if (wantUser == nil) != (gotUser == nil) {
t.Fatalf("GetCurrentUser nil-ness mismatch: want %v, got %v", wantUser, gotUser)
}
if wantUser != nil {
assert.Same(t, wantUser, gotUser)
}
assert.Equal(t, errText(wantUserErr), errText(gotUserErr))
wantUID, wantUIDErr := legacyGetCurrentUserID(ctx)
gotUID, gotUIDErr := svc.GetCurrentUserID(ctx)
assert.Equal(t, wantUID, gotUID)
assert.Equal(t, errText(wantUIDErr), errText(gotUIDErr))
}
func TestLoginStateContextParityWithLegacyImplementation(t *testing.T) {
gin.SetMode(gin.TestMode)
svc := newTestAuthService(t, nil)
user := &contracts.UserDTO{ID: 5150, Username: "parity_user", IsActive: true}
t.Run("gin 上下文各分支", func(t *testing.T) {
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/user-info", nil)
assertLoginStateParity(t, svc, c)
c.Set(contracts.AuthUserObjKey, user)
assertLoginStateParity(t, svc, c)
c.Set(contracts.AuthUserObjKey, "not-a-user-dto")
assertLoginStateParity(t, svc, c)
var typedNil *contracts.UserDTO
c.Set(contracts.AuthUserObjKey, typedNil)
assertLoginStateParity(t, svc, c)
})
t.Run("普通 context 各分支", func(t *testing.T) {
assertLoginStateParity(t, svc, context.Background())
assertLoginStateParity(t, svc, context.WithValue(context.Background(), contracts.AuthUserObjKey, user))
assertLoginStateParity(t, svc, context.WithValue(context.Background(), contracts.AuthUserObjKey, "nope"))
})
t.Run("Session 登录态各分支", func(t *testing.T) {
cases := []struct {
name string
userID any
}{
{name: "无用户", userID: nil},
{name: "uint64 用户 ID", userID: uint64(3301)},
{name: "float64 用户 ID", userID: float64(3302)},
{name: "string 用户 ID", userID: "3303"},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
engine := newSessionEngine()
engine.GET("/probe", func(c *gin.Context) {
if tc.userID != nil {
session := sessions.Default(c)
session.Set(auth.UserIDKey, tc.userID)
require.NoError(t, session.Save())
}
assertLoginStateParity(t, svc, c)
})
engine.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/probe", nil))
})
}
})
}
+2 -2
View File
@@ -116,8 +116,8 @@ func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDT
maxAge := config.Config.App.SessionAge
isSessionCookie := false
var val string
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "login_session_ttl_hours").Pluck("value", &val).Error; err == nil && val != "" {
val, err := GetSystemConfigValue(ctx, "login_session_ttl_hours")
if err == nil && val != "" {
if ttlHours, err := strconv.Atoi(val); err == nil {
switch {
case ttlHours == -1:
+14
View File
@@ -4,7 +4,21 @@
// Package cap 提供人机验证中间件
package cap
// HTTP 响应错误文案
const (
errCapTokenMissing = "验证码验证失败,缺少验证码凭证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errCapTokenInvalidOrExpired = "验证码校验失败或已过期,请重试" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errCapNotConfigured = "captcha is not configured"
errChallengeGenerateFailed = "生成验证难题失败,请稍后再试"
errInvalidRequestParams = "无效的参数"
errSolutionVerifyFailed = "校验验证解答失败,请稍后再试"
)
// Redeem 结果码,属于 redeem 响应 JSON 的对外契约取值,禁止改写取值
const (
redeemErrInvalidToken = "invalid_token"
redeemErrNonceStoreFailed = "nonce_store_error"
redeemErrAlreadyRedeemed = "already_redeemed"
redeemErrSettingsLoad = "settings_load_error"
redeemErrTokenStoreFailed = "token_store_error" //nolint:gosec // error code, not hardcoded credentials
)
+5 -19
View File
@@ -6,25 +6,11 @@ package cap
import (
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/cap/pow"
"net/http"
"github.com/gin-gonic/gin"
)
// ChallengeResponse is a local type alias for the pow.ChallengeResponse struct
type ChallengeResponse = pow.ChallengeResponse
type challengeRequest struct {
Scope string `json:"scope" form:"scope"`
}
type redeemRequest struct {
Token string `json:"token" binding:"required"`
Solutions []int `json:"solutions" binding:"required"`
Scope string `json:"scope" form:"scope"`
}
// Challenge 生成 PoW 人机验证难题
// @Summary 生成人机验证难题
// @Description 客户端获取 PoW 难题和签名的 JWT Token,并在后台计算。
@@ -45,13 +31,13 @@ func Challenge(c *gin.Context) {
mgr := GetDefaultManager()
if mgr == nil {
response.AbortInternal(c, "captcha is not configured")
response.AbortInternal(c, errCapNotConfigured)
return
}
resp, err := mgr.Generate(c.Request.Context(), req.Scope)
if err != nil {
logger.ErrorF(c.Request.Context(), "Generate cap challenge failed: %v", err)
response.AbortInternal(c, "生成验证难题失败,请稍后再试")
response.AbortInternal(c, errChallengeGenerateFailed)
return
}
@@ -72,7 +58,7 @@ func Challenge(c *gin.Context) {
func Redeem(c *gin.Context) {
var req redeemRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, "无效的参数")
response.AbortBadRequest(c, errInvalidRequestParams)
return
}
@@ -82,13 +68,13 @@ func Redeem(c *gin.Context) {
mgr := GetDefaultManager()
if mgr == nil {
response.AbortInternal(c, "captcha is not configured")
response.AbortInternal(c, errCapNotConfigured)
return
}
resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
if err != nil {
logger.ErrorF(c.Request.Context(), "Redeem cap solutions failed: %v", err)
response.AbortInternal(c, "校验验证解答失败,请稍后再试")
response.AbortInternal(c, errSolutionVerifyFailed)
return
}
+37
View File
@@ -0,0 +1,37 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"Wavelet/plugins/domain/cap/pow"
)
// ChallengeResponse is a local type alias for the pow.ChallengeResponse struct
type ChallengeResponse = pow.ChallengeResponse
// challengeRequest is the CAPTCHA challenge request payload.
type challengeRequest struct {
Scope string `json:"scope" form:"scope"`
}
// redeemRequest is the CAPTCHA redeem request payload.
type redeemRequest struct {
Token string `json:"token" binding:"required"`
Solutions []int `json:"solutions" binding:"required"`
Scope string `json:"scope" form:"scope"`
}
// RedeemResponse is returned to the client on redeem.
type RedeemResponse struct {
Success bool `json:"success"`
Token string `json:"token,omitempty"`
Expires int64 `json:"expires,omitempty"`
Error string `json:"error,omitempty"`
}
// configRecord maps the columns selected from the system config table.
type configRecord struct {
Key string `gorm:"column:key"`
Value string `gorm:"column:value"`
}
+58
View File
@@ -0,0 +1,58 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"Wavelet/core"
"Wavelet/core/contracts"
"context"
"sync"
"gorm.io/gorm"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
)
// setDBService caches the DBService contract used by the persistence layer.
func setDBService(s contracts.DBService) {
dbMu.Lock()
defer dbMu.Unlock()
dbSvc = s
}
// getDB resolves a GORM handle, preferring the *core.Context when supplied by callers.
func getDB(ctx context.Context) *gorm.DB {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
return s.DB(ctx)
}
}
dbMu.RLock()
s := dbSvc
dbMu.RUnlock()
if s != nil {
return s.DB(ctx)
}
return nil
}
// loadRuntimeSettings reads the CAPTCHA owned rows from the system config table.
func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
var records []configRecord
db := getDB(ctx)
if db == nil {
return parseRuntimeSettings(nil), nil
}
if err := db.Table("w_system_configs").Where("key IN ?", runtimeConfigKeys).Find(&records).Error; err != nil {
return RuntimeSettings{}, err
}
configs := make(map[string]string, len(records))
for _, r := range records {
configs[r.Key] = r.Value
}
return parseRuntimeSettings(configs), nil
}
@@ -4,45 +4,15 @@
package cap
import (
"Wavelet/core"
"Wavelet/core/contracts"
"context"
"errors"
"strconv"
"sync"
"sync/atomic"
"time"
"golang.org/x/sync/singleflight"
"gorm.io/gorm"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
)
func setDBService(s contracts.DBService) {
dbMu.Lock()
defer dbMu.Unlock()
dbSvc = s
}
func getDB(ctx context.Context) *gorm.DB {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
return s.DB(ctx)
}
}
dbMu.RLock()
s := dbSvc
dbMu.RUnlock()
if s != nil {
return s.DB(ctx)
}
return nil
}
const (
defaultChallengeCount = 1
defaultChallengeSize = 32
@@ -165,26 +135,6 @@ func (s *runtimeSettingsStore) current(ctx context.Context) (RuntimeSettings, er
return settings, nil
}
func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
type configRecord struct {
Key string `gorm:"column:key"`
Value string `gorm:"column:value"`
}
var records []configRecord
db := getDB(ctx)
if db == nil {
return parseRuntimeSettings(nil), nil
}
if err := db.Table("w_system_configs").Where("key IN ?", runtimeConfigKeys).Find(&records).Error; err != nil {
return RuntimeSettings{}, err
}
configs := make(map[string]string, len(records))
for _, r := range records {
configs[r.Key] = r.Value
}
return parseRuntimeSettings(configs), nil
}
func parseRuntimeSettings(configs map[string]string) RuntimeSettings {
settings := RuntimeSettings{
ChallengeCount: defaultChallengeCount,
@@ -53,19 +53,11 @@ func (m *Manager) Generate(ctx context.Context, scope string) (*pow.ChallengeRes
return pow.GenerateChallenge(m.secret, challengeConfig, scope)
}
// RedeemResponse is returned to the client on redeem.
type RedeemResponse struct {
Success bool `json:"success"`
Token string `json:"token,omitempty"`
Expires int64 `json:"expires,omitempty"`
Error string `json:"error,omitempty"`
}
// Redeem verifies PoW solutions and returns a one-time redeem token.
func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*RedeemResponse, error) {
sigHex := pow.JwtSigHex(token)
if sigHex == "" {
return &RedeemResponse{Success: false, Error: "invalid_token"}, nil
return &RedeemResponse{Success: false, Error: redeemErrInvalidToken}, nil
}
nonceKey := "cap:nonce:" + sigHex
@@ -83,15 +75,15 @@ func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, sco
set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL)
if err != nil {
return &RedeemResponse{Success: false, Error: "nonce_store_error"}, err
return &RedeemResponse{Success: false, Error: redeemErrNonceStoreFailed}, err
}
if !set {
return &RedeemResponse{Success: false, Error: "already_redeemed"}, nil
return &RedeemResponse{Success: false, Error: redeemErrAlreadyRedeemed}, nil
}
settings, err := CurrentSettings(ctx)
if err != nil {
return &RedeemResponse{Success: false, Error: "settings_load_error"}, err
return &RedeemResponse{Success: false, Error: redeemErrSettingsLoad}, err
}
id := pow.RandomHex(redeemTokenIDLength)
@@ -104,7 +96,7 @@ func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, sco
storeVal := strconv.FormatInt(tokenExpires.UnixNano(), 10) + "|" + scope
if err := m.store.Set(ctx, tokenKey, storeVal, settings.TokenTTL); err != nil {
return &RedeemResponse{Success: false, Error: "token_store_error"}, err
return &RedeemResponse{Success: false, Error: redeemErrTokenStoreFailed}, err
}
return &RedeemResponse{
@@ -1,338 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strings"
"time"
"github.com/tencent-connect/botgo/token"
"gorm.io/gorm"
)
const defaultTelegramAPI = "https://api.telegram.org"
// Field is one admin form field.
type Field struct {
Key string `json:"key"`
Type string `json:"type"`
Required bool `json:"required"`
}
// Definition describes a channel type form.
type Definition struct {
Type string `json:"type"`
Fields []Field `json:"fields"`
}
// ChannelDTO represents a channel for admin consumption.
type ChannelDTO struct {
ID uint64 `json:"id,string"`
Name string `json:"name"`
Type string `json:"type"`
OwnerScope string `json:"owner_scope"`
OwnerID *uint64 `json:"owner_id,string,omitempty"`
Enabled bool `json:"enabled"`
Credentials map[string]string `json:"credentials"`
Extra map[string]string `json:"extra"`
}
// CreateChannelRequest is admin create payload.
type CreateChannelRequest struct {
Name string `json:"name"`
Type string `json:"type"`
Enabled *bool `json:"enabled"`
Credentials map[string]string `json:"credentials"`
Extra map[string]string `json:"extra"`
}
// UpdateChannelRequest is admin update payload.
type UpdateChannelRequest struct {
Name string `json:"name"`
Enabled *bool `json:"enabled"`
Credentials map[string]string `json:"credentials"`
Extra map[string]string `json:"extra"`
}
func listDefinitions() []Definition {
return []Definition{
{
Type: MessageChannelTypeTelegram,
Fields: []Field{
{Key: "token", Type: "password", Required: true},
{Key: "api_base", Type: "text", Required: false},
},
},
{
Type: MessageChannelTypeQQ,
Fields: []Field{
{Key: "app_id", Type: "text", Required: true},
{Key: "client_secret", Type: "password", Required: true},
},
},
}
}
func createChannel(ctx context.Context, req CreateChannelRequest) (ChannelDTO, error) {
name := strings.TrimSpace(req.Name)
if name == "" {
return ChannelDTO{}, errors.New(errNameRequired)
}
channelType := strings.TrimSpace(req.Type)
if channelType != MessageChannelTypeTelegram && channelType != MessageChannelTypeQQ {
return ChannelDTO{}, errors.New(errTypeInvalid)
}
creds := req.Credentials
if creds == nil {
creds = map[string]string{}
}
if err := validateCredentials(channelType, creds, false); err != nil {
return ChannelDTO{}, err
}
cipher, err := EncryptCredentials(creds)
if err != nil {
return ChannelDTO{}, err
}
extra := req.Extra
if extra == nil {
extra = map[string]string{}
}
enabled := true
if req.Enabled != nil {
enabled = *req.Enabled
}
row := &MessageChannel{
Name: name,
Type: channelType,
OwnerScope: MessageOwnerScopeSystem,
Enabled: enabled,
Credentials: cipher,
Extra: EncodeExtra(extra),
}
if err := CreateMessageChannel(ctx, row); err != nil {
return ChannelDTO{}, err
}
return toDTO(row, creds, extra), nil
}
func updateChannel(ctx context.Context, id uint64, req UpdateChannelRequest) (ChannelDTO, error) {
row, err := GetMessageChannel(ctx, id)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return ChannelDTO{}, errors.New(errChannelNotFound)
}
return ChannelDTO{}, err
}
creds, err := DecryptCredentials(row.Credentials)
if err != nil {
return ChannelDTO{}, err
}
extra := ParseExtra(row.Extra)
if name := strings.TrimSpace(req.Name); name != "" {
row.Name = name
}
if req.Enabled != nil {
row.Enabled = *req.Enabled
}
if req.Extra != nil {
extra = req.Extra
}
if len(req.Credentials) > 0 {
merged := make(map[string]string, len(creds))
for k, v := range creds {
merged[k] = v
}
for k, v := range req.Credentials {
if strings.TrimSpace(v) == "" {
continue
}
merged[k] = v
}
if err := validateCredentials(row.Type, merged, true); err != nil {
return ChannelDTO{}, err
}
creds = merged
}
cipher, err := EncryptCredentials(creds)
if err != nil {
return ChannelDTO{}, err
}
row.Credentials = cipher
row.Extra = EncodeExtra(extra)
if err := UpdateMessageChannel(ctx, row); err != nil {
return ChannelDTO{}, err
}
return toDTO(row, creds, extra), nil
}
func listChannels(ctx context.Context) ([]ChannelDTO, error) {
rows, err := ListMessageChannels(ctx)
if err != nil {
return nil, err
}
out := make([]ChannelDTO, 0, len(rows))
for i := range rows {
creds, _ := DecryptCredentials(rows[i].Credentials)
extra := ParseExtra(rows[i].Extra)
out = append(out, toDTO(&rows[i], creds, extra))
}
return out, nil
}
func deleteChannel(ctx context.Context, id uint64) error {
if _, err := GetMessageChannel(ctx, id); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errChannelNotFound)
}
return err
}
return DeleteMessageChannel(ctx, id)
}
func probeChannel(ctx context.Context, id uint64) error {
row, err := GetMessageChannel(ctx, id)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errChannelNotFound)
}
return err
}
creds, err := DecryptCredentials(row.Credentials)
if err != nil {
return err
}
switch row.Type {
case MessageChannelTypeTelegram:
return probeTelegram(ctx, creds)
case MessageChannelTypeQQ:
return probeQQ(ctx, creds)
default:
return errors.New(errTypeInvalid)
}
}
func probeTelegram(ctx context.Context, creds map[string]string) error {
tok := creds["token"]
if strings.TrimSpace(tok) == "" {
return errors.New("missing telegram bot token")
}
base := creds["api_base"]
base = strings.TrimRight(strings.TrimSpace(base), "/")
if base == "" {
base = defaultTelegramAPI
}
url := fmt.Sprintf("%s/bot%s/getMe", base, tok)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return err
}
client := &http.Client{Timeout: 10 * time.Second}
resp, err := client.Do(req)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
body, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("telegram getMe failed (%d): %s", resp.StatusCode, string(body))
}
var res struct {
OK bool `json:"ok"`
}
if err := json.Unmarshal(body, &res); err != nil {
return err
}
if !res.OK {
return fmt.Errorf("telegram returned ok=false: %s", string(body))
}
return nil
}
func probeQQ(_ context.Context, creds map[string]string) error {
appID := strings.TrimSpace(creds["app_id"])
secret := strings.TrimSpace(creds["app_secret"])
if appID == "" || secret == "" {
return errors.New("missing qq app_id or app_secret")
}
credentials := &token.QQBotCredentials{
AppID: appID,
AppSecret: secret,
}
tokSrc := token.NewQQBotTokenSource(credentials)
tok, err := tokSrc.Token()
if err != nil {
return fmt.Errorf("qq token fetch failed: %w", err)
}
if tok == nil || tok.AccessToken == "" {
return errors.New("qq returned empty access token")
}
return nil
}
func validateCredentials(t string, creds map[string]string, isUpdate bool) error {
switch t {
case MessageChannelTypeTelegram:
tok := creds["token"]
if strings.TrimSpace(tok) == "" && !isUpdate {
return errors.New(errTelegramTokenRequired)
}
if base, ok := creds["api_base"]; ok && strings.TrimSpace(base) != "" {
if !strings.HasPrefix(base, "http://") && !strings.HasPrefix(base, "https://") {
return errors.New("api_base must start with http:// or https://")
}
}
case MessageChannelTypeQQ:
appID := creds["app_id"]
secret := creds["client_secret"]
if (strings.TrimSpace(appID) == "" || strings.TrimSpace(secret) == "") && !isUpdate {
return errors.New(errQQCredentialsRequired)
}
default:
return errors.New(errTypeInvalid)
}
return nil
}
func toDTO(row *MessageChannel, creds, extra map[string]string) ChannelDTO {
return ChannelDTO{
ID: row.ID,
Name: row.Name,
Type: row.Type,
OwnerScope: row.OwnerScope,
OwnerID: row.OwnerID,
Enabled: row.Enabled,
Credentials: maskCredentials(row.Type, creds),
Extra: extra,
}
}
func maskCredentials(_ string, in map[string]string) map[string]string {
out := make(map[string]string, len(in))
for k, v := range in {
if k == "token" || k == "client_secret" {
out[k] = maskSecret(v)
} else {
out[k] = v
}
}
return out
}
const minMaskSecretLength = 8
func maskSecret(s string) string {
s = strings.TrimSpace(s)
if len(s) <= minMaskSecretLength {
return "******"
}
return s[:4] + "..." + s[len(s)-4:]
}
@@ -1,21 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import "context"
// Handler processes one inbound message.
type Handler func(ctx context.Context, msg InboundMessage) error
// Factory constructs a Channel from decrypted config.
type Factory func(cfg ChannelConfig, onInbound Handler) (Channel, error)
// Channel is one connected messaging adapter.
type Channel interface {
Type() string
Connect(ctx context.Context) error
Disconnect(ctx context.Context) error
Send(ctx context.Context, to Recipient, msg OutboundMessage) error
Capabilities() Capability
}
@@ -7,7 +7,8 @@ package qq
import (
"Wavelet/pkg/logger"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/message_gateway"
"Wavelet/plugins/domain/message_gateway/model"
"Wavelet/plugins/domain/message_gateway/service"
"context"
"fmt"
"strings"
@@ -32,8 +33,8 @@ type qqEvent struct {
// Adapter is an official QQ Bot C2C channel.
type Adapter struct {
cfg message_gateway.ChannelConfig
onInbound message_gateway.Handler
cfg model.ChannelConfig
onInbound service.Handler
api openapi.OpenAPI
tokenSrc oauth2.TokenSource
cancel context.CancelFunc
@@ -42,7 +43,7 @@ type Adapter struct {
}
// New constructs a QQ adapter.
func New(cfg message_gateway.ChannelConfig, onInbound message_gateway.Handler) (message_gateway.Channel, error) {
func New(cfg model.ChannelConfig, onInbound service.Handler) (service.Channel, error) {
if strings.TrimSpace(cfg.Credentials["app_id"]) == "" || strings.TrimSpace(cfg.Credentials["app_secret"]) == "" {
return nil, fmt.Errorf("qq: app_id and app_secret are required")
}
@@ -50,11 +51,11 @@ func New(cfg message_gateway.ChannelConfig, onInbound message_gateway.Handler) (
}
// Type returns qq.
func (a *Adapter) Type() string { return message_gateway.ChannelTypeQQ }
func (a *Adapter) Type() string { return model.ChannelTypeQQ }
// Capabilities reports C2C text/media support.
func (a *Adapter) Capabilities() message_gateway.Capability {
return message_gateway.Capability{Text: true, Image: true, File: true, Reply: true}
func (a *Adapter) Capabilities() model.Capability {
return model.Capability{Text: true, Image: true, File: true, Reply: true}
}
// Connect starts the official WebSocket session (C2C intent).
@@ -127,7 +128,7 @@ func (a *Adapter) Disconnect(_ context.Context) error {
}
// Send posts a C2C text reply.
func (a *Adapter) Send(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error {
func (a *Adapter) Send(ctx context.Context, to model.Recipient, msg model.OutboundMessage) error {
a.mu.Lock()
api := a.api
a.mu.Unlock()
@@ -151,9 +152,9 @@ func (a *Adapter) handleEvent(ctx context.Context, ev qqEvent) {
if disconnected || a.onInbound == nil {
return
}
_ = a.onInbound(ctx, message_gateway.InboundMessage{
_ = a.onInbound(ctx, model.InboundMessage{
ChannelID: a.cfg.ID,
ChannelType: message_gateway.ChannelTypeQQ,
ChannelType: model.ChannelTypeQQ,
PlatformUserID: ev.UserID,
ChatID: ev.UserID,
MessageID: ev.MessageID,
@@ -4,14 +4,14 @@
package qq
import (
"Wavelet/plugins/domain/message_gateway"
"Wavelet/plugins/domain/message_gateway/model"
"context"
"testing"
)
func TestHandleEvent_DropsNonC2C(t *testing.T) {
var got int
a := &Adapter{onInbound: func(ctx context.Context, msg message_gateway.InboundMessage) error {
a := &Adapter{onInbound: func(ctx context.Context, msg model.InboundMessage) error {
got++
return nil
}}
@@ -22,8 +22,8 @@ func TestHandleEvent_DropsNonC2C(t *testing.T) {
}
func TestHandleEvent_C2CText(t *testing.T) {
var got message_gateway.InboundMessage
a := &Adapter{cfg: message_gateway.ChannelConfig{ID: 3}, onInbound: func(ctx context.Context, msg message_gateway.InboundMessage) error {
var got model.InboundMessage
a := &Adapter{cfg: model.ChannelConfig{ID: 3}, onInbound: func(ctx context.Context, msg model.InboundMessage) error {
got = msg
return nil
}}
@@ -34,7 +34,7 @@ func TestHandleEvent_C2CText(t *testing.T) {
}
func TestNew_RequiresCreds(t *testing.T) {
_, err := New(message_gateway.ChannelConfig{}, nil)
_, err := New(model.ChannelConfig{}, nil)
if err == nil {
t.Fatal("expected error")
}
@@ -6,7 +6,8 @@ package telegram
import (
"Wavelet/pkg/util"
"Wavelet/plugins/domain/message_gateway"
"Wavelet/plugins/domain/message_gateway/model"
"Wavelet/plugins/domain/message_gateway/service"
"context"
"fmt"
"os"
@@ -19,13 +20,13 @@ import (
// Adapter is a Telegram private-chat channel.
type Adapter struct {
cfg message_gateway.ChannelConfig
onInbound message_gateway.Handler
cfg model.ChannelConfig
onInbound service.Handler
bot *tele.Bot
}
// New constructs a Telegram adapter. Call message_gateway.Register from the runner.
func New(cfg message_gateway.ChannelConfig, onInbound message_gateway.Handler) (message_gateway.Channel, error) {
// New constructs a Telegram adapter. Call service.Register from the runner.
func New(cfg model.ChannelConfig, onInbound service.Handler) (service.Channel, error) {
if strings.TrimSpace(cfg.Credentials["bot_token"]) == "" {
return nil, fmt.Errorf("telegram: bot_token is required")
}
@@ -33,11 +34,11 @@ func New(cfg message_gateway.ChannelConfig, onInbound message_gateway.Handler) (
}
// Type returns telegram.
func (a *Adapter) Type() string { return message_gateway.ChannelTypeTelegram }
func (a *Adapter) Type() string { return model.ChannelTypeTelegram }
// Capabilities reports private-chat media support.
func (a *Adapter) Capabilities() message_gateway.Capability {
return message_gateway.Capability{Text: true, Image: true, File: true, Reply: true}
func (a *Adapter) Capabilities() model.Capability {
return model.Capability{Text: true, Image: true, File: true, Reply: true}
}
// Connect starts long polling.
@@ -85,7 +86,7 @@ func (a *Adapter) Disconnect(_ context.Context) error {
}
// Send replies to a private chat.
func (a *Adapter) Send(_ context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error {
func (a *Adapter) Send(_ context.Context, to model.Recipient, msg model.OutboundMessage) error {
if a.bot == nil {
return fmt.Errorf("telegram: not connected")
}
@@ -104,9 +105,9 @@ func (a *Adapter) handleTeleMessage(ctx context.Context, m *tele.Message) {
if a.onInbound == nil {
return
}
msg := message_gateway.InboundMessage{
msg := model.InboundMessage{
ChannelID: a.cfg.ID,
ChannelType: message_gateway.ChannelTypeTelegram,
ChannelType: model.ChannelTypeTelegram,
PlatformUserID: strconv.FormatInt(m.Sender.ID, 10),
ChatID: strconv.FormatInt(m.Chat.ID, 10),
MessageID: strconv.Itoa(m.ID),
@@ -121,7 +122,7 @@ func (a *Adapter) handleTeleMessage(ctx context.Context, m *tele.Message) {
_ = a.onInbound(ctx, msg)
}
func (a *Adapter) downloadMedia(m *tele.Message) []message_gateway.Attachment {
func (a *Adapter) downloadMedia(m *tele.Message) []model.Attachment {
var files []*tele.File
var names []string
if m.Photo != nil {
@@ -141,16 +142,16 @@ func (a *Adapter) downloadMedia(m *tele.Message) []message_gateway.Attachment {
}
dir, err := os.MkdirTemp("", "wg-tg-*")
if err != nil {
return []message_gateway.Attachment{{Error: err.Error()}}
return []model.Attachment{{Error: err.Error()}}
}
out := make([]message_gateway.Attachment, 0, len(files))
out := make([]model.Attachment, 0, len(files))
for i, f := range files {
path := filepath.Join(dir, names[i])
if err := a.bot.Download(f, path); err != nil {
out = append(out, message_gateway.Attachment{FileName: names[i], Error: err.Error()})
out = append(out, model.Attachment{FileName: names[i], Error: err.Error()})
continue
}
out = append(out, message_gateway.Attachment{Path: path, FileName: names[i]})
out = append(out, model.Attachment{Path: path, FileName: names[i]})
}
return out
}
@@ -4,7 +4,7 @@
package telegram
import (
"Wavelet/plugins/domain/message_gateway"
"Wavelet/plugins/domain/message_gateway/model"
"context"
"testing"
@@ -13,7 +13,7 @@ import (
func TestHandleUpdate_DropsGroups(t *testing.T) {
var got int
a := &Adapter{onInbound: func(ctx context.Context, msg message_gateway.InboundMessage) error {
a := &Adapter{onInbound: func(ctx context.Context, msg model.InboundMessage) error {
got++
return nil
}}
@@ -29,10 +29,10 @@ func TestHandleUpdate_DropsGroups(t *testing.T) {
}
func TestHandleUpdate_PrivateText(t *testing.T) {
var got message_gateway.InboundMessage
var got model.InboundMessage
a := &Adapter{
cfg: message_gateway.ChannelConfig{ID: 7, Type: "telegram"},
onInbound: func(ctx context.Context, msg message_gateway.InboundMessage) error {
cfg: model.ChannelConfig{ID: 7, Type: "telegram"},
onInbound: func(ctx context.Context, msg model.InboundMessage) error {
got = msg
return nil
},
@@ -49,7 +49,7 @@ func TestHandleUpdate_PrivateText(t *testing.T) {
}
func TestNew_RequiresToken(t *testing.T) {
_, err := New(message_gateway.ChannelConfig{}, nil)
_, err := New(model.ChannelConfig{}, nil)
if err == nil {
t.Fatal("expected error")
}
@@ -1,41 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"Wavelet/core/contracts"
"context"
"time"
)
// AdminLogin is the metadata definition for the admin login event.
var AdminLogin = EventMetadata{
Key: "admin_login",
Name: "管理员登录",
DefaultTemplate: NotificationMessage{
Title: "管理员登录提醒",
Content: "管理员 {{user.username}} 于 {{time}} 从 IP {{ip}} 登录系统。",
Level: "INFO",
},
Description: "当管理员成功登录系统时触发此通知",
}
// HandleAdminLoggedIn 处理管理员登录事件并触发通知
func HandleAdminLoggedIn(ctx context.Context, event contracts.AdminLoggedIn) {
if event.User == nil {
return
}
body := map[string]any{
"user": event.User,
"ip": event.IP,
"time": time.Now().Format("2006-01-02 15:04:05"),
}
DefaultTrigger.Trigger(ctx, AdminLogin, body)
}
// RegisterCustomEvents registers default domain push notification events.
func RegisterCustomEvents() {
RegisterBuiltInEvent(AdminLogin)
}
@@ -1,97 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"Wavelet/core"
"Wavelet/core/contracts"
"context"
"sync"
"gorm.io/gorm"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
cacheMu sync.RWMutex
cacheSvc contracts.CacheService
taskMu sync.RWMutex
taskSvc contracts.TaskService
userMu sync.RWMutex
userSvc contracts.UserService
)
// SetDBServiceForTest injects a DBService for tests. Production wiring must use Apply.
func SetDBServiceForTest(s contracts.DBService) {
setDBService(s)
}
func setDBService(s contracts.DBService) {
dbMu.Lock()
defer dbMu.Unlock()
dbSvc = s
}
func setCacheService(s contracts.CacheService) {
cacheMu.Lock()
defer cacheMu.Unlock()
cacheSvc = s
}
func setTaskService(s contracts.TaskService) {
taskMu.Lock()
defer taskMu.Unlock()
taskSvc = s
}
func setUserService(s contracts.UserService) {
userMu.Lock()
defer userMu.Unlock()
userSvc = s
}
func getDB(ctx context.Context) *gorm.DB {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
return s.DB(ctx)
}
}
dbMu.RLock()
s := dbSvc
dbMu.RUnlock()
if s != nil {
return s.DB(ctx)
}
return nil
}
func getCache(ctx context.Context) contracts.CacheService {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
return s
}
}
cacheMu.RLock()
s := cacheSvc
cacheMu.RUnlock()
return s
}
func getTaskService() contracts.TaskService {
taskMu.RLock()
defer taskMu.RUnlock()
return taskSvc
}
func getUserService(ctx context.Context) contracts.UserService {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.UserService](c); err == nil && s != nil {
return s
}
}
userMu.RLock()
defer userMu.RUnlock()
return userSvc
}
@@ -1,26 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import "errors"
var (
errCodeInvalid = errors.New("invalid or expired pairing code")
errChannelMismatch = errors.New("pairing code does not match channel")
errPlatformAlreadyBound = errors.New("this platform account is already bound")
errBindingNotFound = errors.New("binding not found")
errBindingForbidden = errors.New("cannot unbind another user's binding")
errChannelIDRequired = errors.New("channel_id is required")
errChannelDisabled = errors.New("channel is not enabled")
)
const (
errNameRequired = "name is required"
errTypeInvalid = "type must be telegram or qq"
errTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text
errQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text
errChannelNotFound = "channel not found"
errChannelProbeFailed = "channel probe failed"
maskedSecret = "********"
)
@@ -0,0 +1,68 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package errs defines error sentinels and user-facing error message constants
// for the message_gateway plugin.
package errs
import "errors"
// Sentinel 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")
// ErrRecordNotFound maps GORM's missing-row sentinel at the repository boundary so
// upper layers never import gorm. Its text matches gorm.ErrRecordNotFound verbatim.
ErrRecordNotFound = errors.New("record not found")
)
// User-facing validation and error message constants.
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 = "********"
ErrLoginRequired = "login required"
ErrInvalidBindingID = "invalid binding id"
ErrInvalidChannelID = "invalid channel id"
ErrInvalidEventID = "invalid event id"
ErrEventNotFound = "notification event not found"
ErrValidationFailed = "validation failed"
ErrMissingTelegramToken = "missing telegram bot token"
ErrMissingQQCredentials = "missing qq app_id or app_secret" //nolint:gosec // user-facing validation text
ErrQQTokenFetchFailed = "qq token fetch failed" //nolint:gosec // user-facing validation text
ErrQQEmptyToken = "qq returned empty access token" //nolint:gosec // user-facing validation text
ErrTelegramGetMeFailed = "telegram getMe failed"
ErrTelegramNotOK = "telegram returned ok=false"
ErrAPIBaseInvalid = "api_base must start with http:// or https://"
ErrChannelNameExists = "channel name already exists"
ErrChannelNameRequired = "channel name is required"
ErrChannelTypeRequired = "channel type is required"
ErrEventKeyRequired = "event_key is required"
ErrEventAlreadyConfigured = "this notification event is already configured"
ErrTemplateInvalidJSON = "custom template is not a valid JSON format"
ErrEnableWithoutChannels = "cannot enable event without any push channels configured"
ErrEventKeyOrTaskType = "either event_key or task_type must be provided"
ErrUnsupportedEventKey = "unsupported built-in event key"
ErrTaskServiceUnavailable = "task service not available"
ErrUserNotFound = "user not found"
ErrNoAdminUser = "no admin user found"
ErrPayloadRequired = "payload is required"
ErrInvalidJSONFormat = "invalid json format"
ErrParsePayloadFailed = "parse payload failed"
ErrGetPusherFailed = "get pusher failed"
ErrPusherSendFailed = "pusher.Send failed"
)
@@ -1,62 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package message_gateway defines channel adapters, pairing codes, and inbound types.
package message_gateway
// ChannelTypeTelegram is the Telegram private-chat adapter type.
const ChannelTypeTelegram = "telegram"
// ChannelTypeQQ is the official QQ Bot C2C adapter type.
const ChannelTypeQQ = "qq"
// Capability describes what an adapter can send and receive.
type Capability struct {
Text bool
Image bool
File bool
Reply bool
Group bool
}
// ChannelConfig is the decrypted runtime config passed to a factory.
type ChannelConfig struct {
ID uint64
Type string
Name string
Credentials map[string]string
Extra map[string]string
}
// Recipient is the outbound destination on a platform.
type Recipient struct {
ChatID string
PlatformUserID string
}
// Attachment is a downloaded inbound file sitting on local disk.
type Attachment struct {
Path string
FileName string
MIME string
Error string
}
// InboundMessage is a normalized private-chat message.
type InboundMessage struct {
ChannelID uint64
ChannelType string
PlatformUserID string
ChatID string
MessageID string
Text string
Attachments []Attachment
BindingUserID *uint64
}
// OutboundMessage is a reply or probe send.
type OutboundMessage struct {
Text string
ReplyToID string
Attachments []Attachment
}
@@ -1,10 +1,12 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
package handler
import (
"Wavelet/pkg/response"
"Wavelet/plugins/domain/message_gateway/errs"
"Wavelet/plugins/domain/message_gateway/service"
"net/http"
"strconv"
@@ -17,10 +19,10 @@ import (
// @Tags admin-message-gateway
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]Definition}
// @Success 200 {object} response.Any{data=[]model.Definition}
// @Router /api/v1/admin/message-gateway/channels/definitions [get]
func ListAdminChannelDefinitions(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(listDefinitions()))
c.JSON(http.StatusOK, response.OK(service.ListDefinitions()))
}
// ListAdminChannels lists configured messaging channels with secrets masked.
@@ -29,10 +31,10 @@ func ListAdminChannelDefinitions(c *gin.Context) {
// @Tags admin-message-gateway
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]ChannelDTO}
// @Success 200 {object} response.Any{data=[]model.ChannelDTO}
// @Router /api/v1/admin/message-gateway/channels [get]
func ListAdminChannels(c *gin.Context) {
rows, err := listChannels(c.Request.Context())
rows, err := service.ListChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -43,14 +45,14 @@ func ListAdminChannels(c *gin.Context) {
func parseAdminChannelID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
response.AbortBadRequest(c, errs.ErrInvalidChannelID)
return 0, false
}
return id, true
}
func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
if err.Error() == errChannelNotFound {
if err.Error() == errs.ErrChannelNotFound {
response.AbortNotFound(c, err.Error())
return
}
@@ -64,12 +66,12 @@ func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Con
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body CreateChannelRequest true "create body"
// @Success 200 {object} response.Any{data=ChannelDTO}
// @Param request body model.CreateChannelRequest true "create body"
// @Success 200 {object} response.Any{data=model.ChannelDTO}
// @Failure 400 {object} response.Any
// @Router /api/v1/admin/message-gateway/channels [post]
func CreateAdminChannel(c *gin.Context) {
handleJSONRequest(c, createChannel)
handleJSONRequest(c, service.CreateChannel)
}
// UpdateAdminChannel patches a messaging channel. Empty secrets keep the previous values.
@@ -80,13 +82,13 @@ func CreateAdminChannel(c *gin.Context) {
// @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}
// @Param request body model.UpdateChannelRequest true "update body"
// @Success 200 {object} response.Any{data=model.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) {
handleEntityUpdate(c, parseAdminChannelID, updateChannel, func(c *gin.Context, err error) {
handleEntityUpdate(c, parseAdminChannelID, service.UpdateChannel, func(c *gin.Context, err error) {
handleAdminChannelError(c, err, response.AbortBadRequest)
})
}
@@ -106,7 +108,7 @@ func DeleteAdminChannel(c *gin.Context) {
if !ok {
return
}
if err := deleteChannel(c.Request.Context(), id); err != nil {
if err := service.DeleteChannel(c.Request.Context(), id); err != nil {
handleAdminChannelError(c, err, response.AbortInternal)
return
}
@@ -129,22 +131,9 @@ func TestAdminChannel(c *gin.Context) {
if !ok {
return
}
if err := probeChannel(c.Request.Context(), id); err != nil {
if err := service.ProbeChannel(c.Request.Context(), id); err != nil {
handleAdminChannelError(c, err, response.AbortBadRequest)
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)
}
}
@@ -1,12 +1,17 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
// Package handler provides HTTP endpoints for message_gateway.
package handler
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/message_gateway/errs"
"Wavelet/plugins/domain/message_gateway/model"
"Wavelet/plugins/domain/message_gateway/service"
"context"
"errors"
"net/http"
"strconv"
@@ -18,21 +23,62 @@ func currentUser(c *gin.Context) (*contracts.UserDTO, bool) {
return ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
}
// handleJSONRequest binds a JSON body, runs the service use case and writes the
// standard success envelope; any service error surfaces as a bad request.
func handleJSONRequest[Req any, Res any](c *gin.Context, handler func(ctx context.Context, req Req) (Res, error)) {
var req Req
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
res, err := handler(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(res))
}
// handleEntityUpdate resolves a path identifier plus JSON body, runs the updater
// use case and writes the success envelope; error classification is delegated to onErr.
func handleEntityUpdate[Req any, Res any](
c *gin.Context,
parseID func(*gin.Context) (uint64, bool),
updater func(ctx context.Context, id uint64, req Req) (Res, error),
onErr func(*gin.Context, error),
) {
id, ok := parseID(c)
if !ok {
return
}
var req Req
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := updater(c.Request.Context(), id, req)
if err != nil {
onErr(c, err)
return
}
c.JSON(http.StatusOK, response.OK(dto))
}
// 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}
// @Success 200 {object} response.Any{data=[]model.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")
response.AbortUnauthorized(c, errs.ErrLoginRequired)
return
}
rows, err := listEnabledPublicChannels(c.Request.Context())
rows, err := service.ListEnabledPublicChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -46,16 +92,16 @@ func ListChannels(c *gin.Context) {
// @Tags message-gateway
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]BindingDTO}
// @Success 200 {object} response.Any{data=[]model.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")
response.AbortUnauthorized(c, errs.ErrLoginRequired)
return
}
rows, err := listUserBindings(c.Request.Context(), user.ID)
rows, err := service.ListUserBindings(c.Request.Context(), user.ID)
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -70,25 +116,25 @@ func ListBindings(c *gin.Context) {
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body BindRequest true "bind body"
// @Success 200 {object} response.Any{data=BindingDTO}
// @Param request body model.BindRequest true "bind body"
// @Success 200 {object} response.Any{data=model.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")
response.AbortUnauthorized(c, errs.ErrLoginRequired)
return
}
var req BindRequest
var req model.BindRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := bindChannel(c.Request.Context(), user.ID, req)
dto, err := service.BindChannel(c.Request.Context(), user.ID, req)
if err != nil {
if errors.Is(err, errPlatformAlreadyBound) {
if errors.Is(err, errs.ErrPlatformAlreadyBound) {
response.AbortConflict(c, err.Error())
return
}
@@ -112,20 +158,20 @@ func BindBinding(c *gin.Context) {
func UnbindBinding(c *gin.Context) {
user, ok := currentUser(c)
if !ok || user == nil {
response.AbortUnauthorized(c, "login required")
response.AbortUnauthorized(c, errs.ErrLoginRequired)
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid binding id")
response.AbortBadRequest(c, errs.ErrInvalidBindingID)
return
}
if err := unbindChannel(c.Request.Context(), user.ID, id); err != nil {
if errors.Is(err, errBindingNotFound) {
if err := service.UnbindChannel(c.Request.Context(), user.ID, id); err != nil {
if errors.Is(err, errs.ErrBindingNotFound) {
response.AbortNotFound(c, err.Error())
return
}
if errors.Is(err, errBindingForbidden) {
if errors.Is(err, errs.ErrBindingForbidden) {
response.AbortForbidden(c, err.Error())
return
}
@@ -134,14 +180,3 @@ func UnbindBinding(c *gin.Context) {
}
c.JSON(http.StatusOK, response.OKNil())
}
// RegisterUserRoutes mounts user-facing message gateway endpoints.
func RegisterUserRoutes(r *gin.RouterGroup, loginMW gin.HandlerFunc) {
mg := r.Group("/message-gateway", loginMW)
{
mg.GET("/channels", ListChannels)
mg.GET("/bindings", ListBindings)
mg.POST("/bindings", BindBinding)
mg.DELETE("/bindings/:id", UnbindBinding)
}
}
@@ -0,0 +1,96 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"Wavelet/pkg/response"
"Wavelet/plugins/domain/message_gateway/errs"
"Wavelet/plugins/domain/message_gateway/model"
"Wavelet/plugins/domain/message_gateway/service"
"errors"
"net/http"
"strconv"
"github.com/gin-gonic/gin"
)
// ListPushChannelDefinitions returns channel definitions.
func ListPushChannelDefinitions(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(model.ListPushDefinitions()))
}
// ListPushChannels lists configured push channels.
func ListPushChannels(c *gin.Context) {
channels, err := service.ListPushChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(channels))
}
// parsePushChannelID reads the path identifier of a push channel.
func parsePushChannelID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, errs.ErrInvalidChannelID)
return 0, false
}
return id, true
}
// handlePushChannelNotFoundError maps a missing channel row to 404, others to fallback.
func handlePushChannelNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
if errors.Is(err, errs.ErrRecordNotFound) {
response.AbortNotFound(c, errs.ErrChannelNotFound)
return
}
fallback(c, err.Error())
}
// CreatePushChannel creates a push channel.
func CreatePushChannel(c *gin.Context) {
handleJSONRequest(c, service.CreatePushChannel)
}
// UpdatePushChannel updates a push channel.
func UpdatePushChannel(c *gin.Context) {
handleEntityUpdate(c, parsePushChannelID, service.UpdatePushChannel, func(c *gin.Context, err error) {
handlePushChannelNotFoundError(c, err, response.AbortInternal)
})
}
// DeletePushChannel deletes a push channel.
func DeletePushChannel(c *gin.Context) {
id, ok := parsePushChannelID(c)
if !ok {
return
}
if err := service.DeletePushChannel(c.Request.Context(), id); err != nil {
handlePushChannelNotFoundError(c, err, response.AbortInternal)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// TestPushChannel tests connectivity of a push channel.
func TestPushChannel(c *gin.Context) {
var req model.TestPushChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
payload, err := service.PreparePushChannelTest(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := service.EnqueuePushTask(c.Request.Context(), payload); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -1,49 +1,24 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
package handler
import (
"Wavelet/pkg/response"
"Wavelet/plugins/domain/message_gateway/errs"
"Wavelet/plugins/domain/message_gateway/model"
"Wavelet/plugins/domain/message_gateway/service"
"errors"
"fmt"
"net/http"
"strconv"
pkgpush "Wavelet/plugins/domain/message_gateway/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)
events, err := service.ListPushEvents(ctx)
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -53,21 +28,23 @@ func ListPushEvents(c *gin.Context) {
// ListBuiltInPushEvents lists system built-in push event definitions.
func ListBuiltInPushEvents(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(GetBuiltInEvents()))
c.JSON(http.StatusOK, response.OK(service.GetBuiltInEvents()))
}
// parsePushEventID reads the path identifier of a push event.
func parsePushEventID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
response.AbortBadRequest(c, errs.ErrInvalidEventID)
return 0, false
}
return id, true
}
// handlePushEventNotFoundError maps a missing event row to 404, others to fallback.
func handlePushEventNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
if errors.Is(err, errs.ErrRecordNotFound) {
response.AbortNotFound(c, errs.ErrEventNotFound)
return
}
fallback(c, err.Error())
@@ -75,7 +52,7 @@ func handlePushEventNotFoundError(c *gin.Context, err error, fallback func(c *gi
// CreatePushEvent creates a new push event configuration.
func CreatePushEvent(c *gin.Context) {
handleJSONRequest(c, createPushEvent)
handleJSONRequest(c, service.CreatePushEvent)
}
// DeletePushEvent deletes a push event configuration by ID.
@@ -85,7 +62,7 @@ func DeletePushEvent(c *gin.Context) {
return
}
if err := deletePushEvent(c.Request.Context(), id); err != nil {
if err := service.DeletePushEvent(c.Request.Context(), id); err != nil {
handlePushEventNotFoundError(c, err, response.AbortInternal)
return
}
@@ -99,13 +76,13 @@ func UpdatePushEvent(c *gin.Context) {
return
}
var req UpdatePushEventRequest
var req model.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 err := service.UpdatePushEvent(c.Request.Context(), id, req); err != nil {
handlePushEventNotFoundError(c, err, response.AbortBadRequest)
return
}
@@ -119,7 +96,7 @@ func TogglePushEvent(c *gin.Context) {
return
}
enabled, err := togglePushEvent(c.Request.Context(), id)
enabled, err := service.TogglePushEvent(c.Request.Context(), id)
if err != nil {
handlePushEventNotFoundError(c, err, response.AbortBadRequest)
return
@@ -138,7 +115,7 @@ func ListPushHistories(c *gin.Context) {
pageSize = 20
}
total, results, err := listPushHistories(c.Request.Context(), PushHistoryListFilter{
total, results, err := service.ListPushHistories(c.Request.Context(), model.PushHistoryListFilter{
EventKey: c.Query("event_key"),
Status: c.Query("status"),
Page: page,
@@ -157,30 +134,13 @@ func ListPushHistories(c *gin.Context) {
// TestPush executes a synchronous push test using the specified config.
func TestPush(c *gin.Context) {
var req TestPushRequest
var req model.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 {
if err := service.RunPushTest(c.Request.Context(), req.Config, req.Target); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
@@ -0,0 +1,63 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"Wavelet/core/extpoints"
"github.com/gin-gonic/gin"
)
// RegisterUserRoutes mounts user-facing message gateway endpoints.
func RegisterUserRoutes(r extpoints.RouterExtension, loginMW gin.HandlerFunc) {
mg := r.Group("/message-gateway", loginMW)
{
mg.GET("/channels", ListChannels)
mg.GET("/bindings", ListBindings)
mg.POST("/bindings", BindBinding)
mg.DELETE("/bindings/:id", UnbindBinding)
}
}
// RegisterAdminRoutes mounts admin message-gateway APIs under /admin.
func RegisterAdminRoutes(adminRouter extpoints.RouterExtension, loginMW, adminMW gin.HandlerFunc) {
g := adminRouter.Group("/message-gateway", loginMW, adminMW)
{
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)
}
}
// RegisterAdminPushRoutes mounts admin push notification APIs under /admin.
func RegisterAdminPushRoutes(adminRouter extpoints.RouterExtension, loginMW, adminMW gin.HandlerFunc) {
adminPushGroup := adminRouter.Group("/push", loginMW, adminMW)
{
events := adminPushGroup.Group("/events")
{
events.GET("", ListPushEvents)
events.GET("/builtin", ListBuiltInPushEvents)
events.POST("", CreatePushEvent)
events.PUT("/:id", UpdatePushEvent)
events.DELETE("/:id", DeletePushEvent)
events.POST("/:id/toggle", TogglePushEvent)
}
adminPushGroup.GET("/histories", ListPushHistories)
adminPushGroup.POST("/test", TestPush)
channels := adminPushGroup.Group("/channels")
{
channels.GET("/definitions", ListPushChannelDefinitions)
channels.GET("", ListPushChannels)
channels.POST("", CreatePushChannel)
channels.PUT("/:id", UpdatePushChannel)
channels.DELETE("/:id", DeletePushChannel)
channels.POST("/test", TestPushChannel)
}
}
}
@@ -1,49 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"Wavelet/pkg/response"
"context"
"net/http"
"github.com/gin-gonic/gin"
)
func handleJSONRequest[Req any, Res any](c *gin.Context, handler func(ctx context.Context, req Req) (Res, error)) {
var req Req
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
res, err := handler(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(res))
}
func handleEntityUpdate[Req any, Res any](
c *gin.Context,
parseID func(*gin.Context) (uint64, bool),
updater func(ctx context.Context, id uint64, req Req) (Res, error),
onErr func(*gin.Context, error),
) {
id, ok := parseID(c)
if !ok {
return
}
var req Req
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := updater(c.Request.Context(), id, req)
if err != nil {
onErr(c, err)
return
}
c.JSON(http.StatusOK, response.OK(dto))
}
@@ -1,155 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"errors"
"strconv"
"strings"
"time"
"gorm.io/gorm"
)
// BindRequest is the user bind body.
type BindRequest struct {
ChannelID string `json:"channel_id"`
Code string `json:"code"`
}
// BindingDTO is a user-facing binding row.
type BindingDTO struct {
ID uint64 `json:"id,string"`
UserID uint64 `json:"user_id,string"`
ChannelID uint64 `json:"channel_id,string"`
ChannelName string `json:"channel_name"`
ChannelType string `json:"channel_type"`
PlatformUserID string `json:"platform_user_id"`
CreatedAt time.Time `json:"created_at"`
}
func bindChannel(ctx context.Context, userID uint64, req BindRequest) (BindingDTO, error) {
channelID, err := strconv.ParseUint(strings.TrimSpace(req.ChannelID), 10, 64)
if err != nil || channelID == 0 {
return BindingDTO{}, errChannelIDRequired
}
code := NormalizeCode(req.Code)
if code == "" {
return BindingDTO{}, errCodeInvalid
}
pairing, err := GetPairingCode(ctx, code)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return BindingDTO{}, errCodeInvalid
}
return BindingDTO{}, err
}
if !pairing.ExpiresAt.After(time.Now()) {
return BindingDTO{}, errCodeInvalid
}
if pairing.ChannelID != channelID {
return BindingDTO{}, errChannelMismatch
}
ch, err := GetMessageChannel(ctx, channelID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return BindingDTO{}, errCodeInvalid
}
return BindingDTO{}, err
}
if !ch.Enabled {
return BindingDTO{}, errChannelDisabled
}
existing, err := GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID)
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return BindingDTO{}, err
}
if err == nil && existing != nil {
if existing.UserID != userID {
return BindingDTO{}, errPlatformAlreadyBound
}
_ = DeletePairingCode(ctx, pairing.Code)
return toBindingDTO(existing, ch), nil
}
row := &MessageBinding{
UserID: userID,
ChannelID: channelID,
PlatformUserID: pairing.PlatformUserID,
}
if err := CreateMessageBinding(ctx, row); err != nil {
return BindingDTO{}, err
}
if err := DeletePairingCode(ctx, pairing.Code); err != nil {
return BindingDTO{}, err
}
return toBindingDTO(row, ch), nil
}
// PublicChannelDTO is an enabled channel a user can bind to.
type PublicChannelDTO struct {
ID uint64 `json:"id,string"`
Name string `json:"name"`
Type string `json:"type"`
}
func listEnabledPublicChannels(ctx context.Context) ([]PublicChannelDTO, error) {
rows, err := ListEnabledMessageChannels(ctx)
if err != nil {
return nil, err
}
out := make([]PublicChannelDTO, 0, len(rows))
for _, row := range rows {
out = append(out, PublicChannelDTO{ID: row.ID, Name: row.Name, Type: row.Type})
}
return out, nil
}
func listUserBindings(ctx context.Context, userID uint64) ([]BindingDTO, error) {
rows, err := ListBindingsByUser(ctx, userID)
if err != nil {
return nil, err
}
out := make([]BindingDTO, 0, len(rows))
for i := range rows {
ch, err := GetMessageChannel(ctx, rows[i].ChannelID)
if err != nil {
out = append(out, toBindingDTO(&rows[i], nil))
continue
}
out = append(out, toBindingDTO(&rows[i], ch))
}
return out, nil
}
func unbindChannel(ctx context.Context, userID, bindingID uint64) error {
row, err := GetMessageBinding(ctx, bindingID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errBindingNotFound
}
return err
}
if row.UserID != userID {
return errBindingForbidden
}
return DeleteMessageBinding(ctx, bindingID)
}
func toBindingDTO(row *MessageBinding, ch *MessageChannel) BindingDTO {
dto := BindingDTO{
ID: row.ID,
UserID: row.UserID,
ChannelID: row.ChannelID,
PlatformUserID: row.PlatformUserID,
CreatedAt: row.CreatedAt,
}
if ch != nil {
dto.ChannelName = ch.Name
dto.ChannelType = ch.Type
}
return dto
}
@@ -0,0 +1,46 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
// Field is one admin form field.
type Field struct {
Key string `json:"key"`
Type string `json:"type"`
Required bool `json:"required"`
}
// Definition describes a channel type form.
type Definition struct {
Type string `json:"type"`
Fields []Field `json:"fields"`
}
// ChannelDTO represents a channel for admin consumption.
type ChannelDTO struct {
ID uint64 `json:"id,string"`
Name string `json:"name"`
Type string `json:"type"`
OwnerScope string `json:"owner_scope"`
OwnerID *uint64 `json:"owner_id,string,omitempty"`
Enabled bool `json:"enabled"`
Credentials map[string]string `json:"credentials"`
Extra map[string]string `json:"extra"`
}
// CreateChannelRequest is admin create payload.
type CreateChannelRequest struct {
Name string `json:"name"`
Type string `json:"type"`
Enabled *bool `json:"enabled"`
Credentials map[string]string `json:"credentials"`
Extra map[string]string `json:"extra"`
}
// UpdateChannelRequest is admin update payload.
type UpdateChannelRequest struct {
Name string `json:"name"`
Enabled *bool `json:"enabled"`
Credentials map[string]string `json:"credentials"`
Extra map[string]string `json:"extra"`
}
@@ -1,16 +1,20 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
// Package model defines the domain entities, DTOs, and schemas for message_gateway.
package model
import (
"Wavelet/plugins/domain/message_gateway/errs"
"errors"
"strings"
"time"
)
// Message channel and push channel constants.
// Channel type and scope constants.
const (
ChannelTypeTelegram = "telegram"
ChannelTypeQQ = "qq"
MessageChannelTypeTelegram = "telegram"
MessageChannelTypeQQ = "qq"
MessageOwnerScopeSystem = "system"
@@ -20,6 +24,57 @@ const (
TypeTelegram = "telegram"
)
// Capability describes what an adapter can send and receive.
type Capability struct {
Text bool
Image bool
File bool
Reply bool
Group bool
}
// ChannelConfig is the decrypted runtime config passed to a factory.
type ChannelConfig struct {
ID uint64
Type string
Name string
Credentials map[string]string
Extra map[string]string
}
// Recipient is the outbound destination on a platform.
type Recipient struct {
ChatID string
PlatformUserID string
}
// Attachment is a downloaded inbound file sitting on local disk.
type Attachment struct {
Path string
FileName string
MIME string
Error string
}
// InboundMessage is a normalized private-chat message.
type InboundMessage struct {
ChannelID uint64
ChannelType string
PlatformUserID string
ChatID string
MessageID string
Text string
Attachments []Attachment
BindingUserID *uint64
}
// OutboundMessage is a reply or probe send.
type OutboundMessage struct {
Text string
ReplyToID string
Attachments []Attachment
}
// MessageChannel is an admin-configured messaging adapter.
type MessageChannel struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
@@ -92,11 +147,11 @@ func (PushChannel) TableName() string {
func (c *PushChannel) Validate() error {
c.Name = strings.TrimSpace(c.Name)
if c.Name == "" {
return errors.New("channel name is required")
return errors.New(errs.ErrChannelNameRequired)
}
c.Type = strings.TrimSpace(c.Type)
if c.Type == "" {
return errors.New("channel type is required")
return errors.New(errs.ErrChannelTypeRequired)
}
return nil
}
@@ -124,11 +179,11 @@ func (PushEvent) TableName() string {
func (e *PushEvent) Validate() error {
e.EventKey = strings.TrimSpace(e.EventKey)
if e.EventKey == "" {
return errors.New("event_key is required")
return errors.New(errs.ErrEventKeyRequired)
}
e.Name = strings.TrimSpace(e.Name)
if e.Name == "" {
return errors.New("name is required")
return errors.New(errs.ErrNameRequired)
}
return nil
}
@@ -153,13 +208,35 @@ func (PushHistory) TableName() string {
return "w_push_histories"
}
// PushHistoryListFilter filters push history pagination queries.
type PushHistoryListFilter struct {
EventKey string
Channel string
Status string
StartTime *time.Time
EndTime *time.Time
Page int
PageSize int
// 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"`
}
// 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"`
}
// PushNotificationEvent defines the payload for eventbus notification trigger.
type PushNotificationEvent struct {
UserID uint64 `json:"user_id"`
Channel string `json:"channel"`
Title string `json:"title"`
Content string `json:"content"`
Metadata map[string]any `json:"metadata,omitempty"`
}
@@ -1,39 +1,49 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
package model
import (
"Wavelet/pkg/response"
"encoding/json"
"errors"
"net/http"
"strconv"
"strings"
"sync"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"time"
pkgpush "Wavelet/plugins/domain/message_gateway/push"
)
// Push channel and payload constants.
const (
// KeyURL represents the URL field key
ChannelCustom = "custom"
ChannelEmail = "email"
ChannelLark = "lark"
ChannelTelegram = "telegram"
DefaultLevelInfo = "INFO"
KeyTitle = "title"
KeyContent = "content"
KeyLevel = "level"
// KeyURL represents the URL field key.
KeyURL = "url"
// KeyToken represents the Token field key
// KeyToken represents the Token field key.
KeyToken = "token"
// KeyOther represents the Other field key
// KeyOther represents the Other field key.
KeyOther = "other"
// TypeText represents standard text input type
// TypeText represents standard text input type.
TypeText = "text"
// TypePassword represents password input type
// TypePassword represents password input type.
TypePassword = "password"
// TypeTextarea represents textarea input type
// TypeTextarea represents textarea input type.
TypeTextarea = "textarea"
)
// SMTPConfig mirrors the system SMTP settings consumed by the push service.
type SMTPConfig struct {
Host string
Port string
Username string
Password string
}
// PushField represents a form field configuration for a channel.
type PushField struct {
Key string `json:"key"`
@@ -52,6 +62,120 @@ type PushDefinition struct {
Fields []PushField `json:"fields"`
}
// 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"`
}
// 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"`
}
// 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"`
}
// 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"`
}
// 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"`
}
// 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"`
}
// TestPushRequest is the request body for testing push config.
type TestPushRequest struct {
Config pkgpush.Config `json:"config" binding:"required"`
Target string `json:"target"`
}
// 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 is the async push dispatch载荷 consumed by the notification 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"`
}
// PushHistoryListFilter filters push history pagination queries.
type PushHistoryListFilter struct {
EventKey string
Channel string
Status string
StartTime *time.Time
EndTime *time.Time
Page int
PageSize int
}
var (
pushDefMu sync.RWMutex
pushDefinitions = make(map[string]PushDefinition)
@@ -69,7 +193,7 @@ func ListPushDefinitions() []PushDefinition {
pushDefMu.RLock()
defer pushDefMu.RUnlock()
order := []string{channelCustom, channelLark, channelTelegram, channelEmail}
order := []string{ChannelCustom, ChannelLark, ChannelTelegram, ChannelEmail}
res := make([]PushDefinition, 0, len(pushDefinitions))
for _, t := range order {
if d, ok := pushDefinitions[t]; ok {
@@ -93,7 +217,7 @@ func ListPushDefinitions() []PushDefinition {
func init() {
RegisterPushChannelDefinition(PushDefinition{
Type: channelCustom,
Type: ChannelCustom,
Name: "自定义消息通道",
Description: "使用自定义 HTTP POST 请求向外部 Webhook 发送数据。",
Fields: []PushField{
@@ -117,7 +241,7 @@ func init() {
})
RegisterPushChannelDefinition(PushDefinition{
Type: channelLark,
Type: ChannelLark,
Name: "飞书群机器人",
Description: "配置飞书群自定义机器人的 Webhook 接口投递。",
Fields: []PushField{
@@ -149,7 +273,7 @@ func init() {
})
RegisterPushChannelDefinition(PushDefinition{
Type: channelTelegram,
Type: ChannelTelegram,
Name: "Telegram 机器人",
Description: "配置 Telegram 机器人推送消息。",
Fields: []PushField{
@@ -181,200 +305,9 @@ func init() {
})
RegisterPushChannelDefinition(PushDefinition{
Type: channelEmail,
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"`
}
func parsePushChannelID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return 0, false
}
return id, true
}
func handlePushChannelNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "channel not found")
return
}
fallback(c, err.Error())
}
// CreatePushChannel creates a push channel.
func CreatePushChannel(c *gin.Context) {
handleJSONRequest(c, createPushChannel)
}
// 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) {
handleEntityUpdate(c, parsePushChannelID, updatePushChannel, func(c *gin.Context, err error) {
handlePushChannelNotFoundError(c, err, response.AbortInternal)
})
}
// DeletePushChannel deletes a push channel.
func DeletePushChannel(c *gin.Context) {
id, ok := parsePushChannelID(c)
if !ok {
return
}
if err := deletePushChannel(c.Request.Context(), id); err != nil {
handlePushChannelNotFoundError(c, err, response.AbortInternal)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// TestPushChannelRequest is the test channel request payload.
type TestPushChannelRequest struct {
Name string `json:"name"`
Type string `json:"type"`
Token string `json:"token"`
URL string `json:"url"`
Other string `json:"other"`
Target string `json:"target"`
}
// TestPushChannel tests connectivity of a push channel.
func TestPushChannel(c *gin.Context) {
var req TestPushChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
ctx := c.Request.Context()
url, token, other, channelType, err := loadChannelForTest(ctx, req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if channelType == channelEmail {
url, token, other = resolveSMTPConfig(ctx, url, token, other)
}
tempChannel := PushChannel{
Name: "test_temp",
URL: url,
Token: token,
Other: other,
Type: channelType,
Enabled: true,
}
if err := tempChannel.Validate(); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
url = tempChannel.URL
var config pkgpush.Config
var renderedJSON string
switch channelType {
case channelLark:
config = pkgpush.Config{Channel: channelLark, URL: url, Secret: token}
renderedJSON = other
case channelEmail:
config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other}
case channelTelegram:
config = pkgpush.Config{Channel: channelTelegram, URL: url, Secret: token, Key: other}
default:
config = pkgpush.Config{Channel: channelCustom, URL: url}
customPushReq := CustomPushRequest{
Title: "通道测试通知",
Content: "这是一条来自系统的消息通道连通性测试消息。",
Description: "系统通道测试",
URL: "https://example.com",
To: req.Target,
}
renderedJSON = renderCustomPayload(other, customPushReq)
}
payload := SendPayload{
EventKey: "test_channel",
Config: config,
Target: req.Target,
Body: NotificationMessage{
Title: "通道测试通知",
Content: "这是一条来自系统的消息通道连通性测试消息。",
Level: defaultLevelInfo,
},
Template: renderedJSON,
}
if err := enqueuePushTask(ctx, payload); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// CustomPushRequest contains custom webhook parameters.
type CustomPushRequest struct {
Title string `json:"title" form:"title"`
Description string `json:"description" form:"description"`
Content string `json:"content" form:"content"`
URL string `json:"url" form:"url"`
To string `json:"to" form:"to"`
Token string `json:"token" form:"token"`
}
func escapeJSONString(s string) string {
b, _ := json.Marshal(s)
const minJSONLen = 2
if len(b) >= minJSONLen {
return string(b[1 : len(b)-1])
}
return s
}
func renderCustomPayload(template string, req CustomPushRequest) string {
result := template
result = strings.ReplaceAll(result, "$title", escapeJSONString(req.Title))
result = strings.ReplaceAll(result, "$description", escapeJSONString(req.Description))
result = strings.ReplaceAll(result, "$content", escapeJSONString(req.Content))
result = strings.ReplaceAll(result, "$url", escapeJSONString(req.URL))
result = strings.ReplaceAll(result, "$to", escapeJSONString(req.To))
return result
}
@@ -1,50 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"crypto/rand"
"strings"
"unicode"
)
// CodeAlphabet excludes easily confused runes 0/O/1/I.
const CodeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
// CodeLength is the raw pairing code size.
const CodeLength = 8
// GenerateCode returns an 8-character pairing code.
func GenerateCode() (string, error) {
buf := make([]byte, CodeLength)
if _, err := rand.Read(buf); err != nil {
return "", err
}
out := make([]byte, CodeLength)
for i, b := range buf {
out[i] = CodeAlphabet[int(b)%len(CodeAlphabet)]
}
return string(out), nil
}
// NormalizeCode strips separators and uppercases.
func NormalizeCode(s string) string {
var b strings.Builder
for _, r := range s {
if r == '-' || unicode.IsSpace(r) {
continue
}
b.WriteRune(unicode.ToUpper(r))
}
return b.String()
}
// FormatCode renders ABCD-EFGH.
func FormatCode(s string) string {
s = NormalizeCode(s)
if len(s) != CodeLength {
return s
}
return s[:4] + "-" + s[4:]
}
+148 -74
View File
@@ -9,6 +9,10 @@ import (
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/message_gateway/handler"
"Wavelet/plugins/domain/message_gateway/model"
"Wavelet/plugins/domain/message_gateway/repository"
"Wavelet/plugins/domain/message_gateway/service"
"context"
"embed"
"reflect"
@@ -68,51 +72,45 @@ func (p *Plugin) Manifest() core.Manifest {
}
}
// PushNotificationEvent defines the payload for eventbus notification trigger.
type PushNotificationEvent struct {
UserID uint64 `json:"user_id"`
Channel string `json:"channel"`
Title string `json:"title"`
Content string `json:"content"`
Metadata map[string]any `json:"metadata,omitempty"`
}
// Apply registers message_gateway migrations, routes, tasks, schedules, events, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
// 0. Bind DBService, CacheService, TaskService, UserService
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
setDBService(db)
repository.SetDBService(db)
} else {
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
setDBService(db)
repository.SetDBService(db)
})
}
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
setCacheService(cache)
repository.SetCacheService(cache)
service.SetCacheService(cache)
} else {
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
setCacheService(cache)
repository.SetCacheService(cache)
service.SetCacheService(cache)
})
}
if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil {
setTaskService(taskSvc)
service.SetTaskService(taskSvc)
} else {
core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) {
setTaskService(taskSvc)
service.SetTaskService(taskSvc)
})
}
if uSvc, err := core.Inject[contracts.UserService](ctx); err == nil && uSvc != nil {
setUserService(uSvc)
service.SetUserService(uSvc)
} else {
core.When[contracts.UserService](ctx, func(uSvc contracts.UserService) {
setUserService(uSvc)
service.SetUserService(uSvc)
})
}
ctx.OnDispose(func() error {
setDBService(nil)
setCacheService(nil)
setTaskService(nil)
setUserService(nil)
repository.SetDBService(nil)
repository.SetCacheService(nil)
service.SetCacheService(nil)
service.SetTaskService(nil)
service.SetUserService(nil)
return nil
})
@@ -132,61 +130,23 @@ func (p *Plugin) Apply(ctx *core.Context) error {
ctx.Migrations().Register("message_gateway", mgMigrations)
// 2. Register User HTTP Routes
mgGroup := ctx.Router().Group("/api/v1/message-gateway", loginMW)
{
mgGroup.GET("/channels", ListChannels)
mgGroup.GET("/bindings", ListBindings)
mgGroup.POST("/bindings", BindBinding)
mgGroup.DELETE("/bindings/:id", UnbindBinding)
}
handler.RegisterUserRoutes(ctx.Router().Group("/api/v1"), loginMW)
// 3. Register Admin Message Gateway HTTP Routes
adminMgGroup := ctx.Router().Group("/api/v1/admin/message-gateway", loginMW, adminMW)
{
adminMgGroup.GET("/channels/definitions", ListAdminChannelDefinitions)
adminMgGroup.GET("/channels", ListAdminChannels)
adminMgGroup.POST("/channels", CreateAdminChannel)
adminMgGroup.PATCH("/channels/:id", UpdateAdminChannel)
adminMgGroup.DELETE("/channels/:id", DeleteAdminChannel)
adminMgGroup.POST("/channels/:id/test", TestAdminChannel)
}
handler.RegisterAdminRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW)
// 4. Register Admin Push HTTP Routes
adminPushGroup := ctx.Router().Group("/api/v1/admin/push", loginMW, adminMW)
{
events := adminPushGroup.Group("/events")
{
events.GET("", ListPushEvents)
events.GET("/builtin", ListBuiltInPushEvents)
events.POST("", CreatePushEvent)
events.PUT("/:id", UpdatePushEvent)
events.DELETE("/:id", DeletePushEvent)
events.POST("/:id/toggle", TogglePushEvent)
}
adminPushGroup.GET("/histories", ListPushHistories)
adminPushGroup.POST("/test", TestPush)
channels := adminPushGroup.Group("/channels")
{
channels.GET("/definitions", ListPushChannelDefinitions)
channels.GET("", ListPushChannels)
channels.POST("", CreatePushChannel)
channels.PUT("/:id", UpdatePushChannel)
channels.DELETE("/:id", DeletePushChannel)
channels.POST("/test", TestPushChannel)
}
}
handler.RegisterAdminPushRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW)
const defaultTaskRetry = 3
pushHandler := &PushHandler{}
pushHandler := &service.PushHandler{}
// 5. Register background tasks
ctx.Task().Register("message_gateway:push_notification", func(c context.Context, payload []byte) error {
return pushHandler.Execute(c, payload)
}, extpoints.WithTaskRetry(defaultTaskRetry))
ctx.Task().Register(SendNotificationTask, func(c context.Context, payload []byte) error {
ctx.Task().Register(service.SendNotificationTask, func(c context.Context, payload []byte) error {
return pushHandler.Execute(c, payload)
}, extpoints.WithTaskRetry(defaultTaskRetry))
@@ -198,19 +158,19 @@ func (p *Plugin) Apply(ctx *core.Context) error {
ctx.Schedule().RegisterCron("*/10 * * * *", "message_gateway:cleanup_pairing_codes", map[string]any{"action": "cleanup"})
// 7. Register EventBus listeners for decoupled push triggers
ctx.Events().On("notification:push", func(c context.Context, e PushNotificationEvent) error {
meta := EventMetadata{
ctx.Events().On("notification:push", func(c context.Context, e model.PushNotificationEvent) error {
meta := model.EventMetadata{
Key: "eventbus:" + e.Channel,
Name: e.Title,
DefaultTemplate: NotificationMessage{
DefaultTemplate: model.NotificationMessage{
Title: e.Title,
Content: e.Content,
Level: defaultLevelInfo,
Level: model.DefaultLevelInfo,
Ext: e.Metadata,
},
Description: "EventBus triggered notification",
}
DefaultTrigger.Trigger(c, meta, map[string]any{
service.DefaultTrigger.Trigger(c, meta, map[string]any{
"user.id": e.UserID,
"title": e.Title,
"content": e.Content,
@@ -220,14 +180,14 @@ func (p *Plugin) Apply(ctx *core.Context) error {
// 8. Register task completed event listener
ctx.Events().On(contracts.EventTopicTaskCompleted, func(c context.Context, e contracts.TaskCompletedEvent) error {
handleTaskCompleted(c, e)
service.HandleTaskCompleted(c, e)
return nil
})
// 9. Register built-in domain events
RegisterCustomEvents()
service.RegisterCustomEvents()
// 9. Register Settings Schemas
// 10. Register Settings Schemas
ctx.Settings().Register(extpoints.SettingSchema{
Key: "message_gateway.pairing_code_expiry_minutes",
Default: 15,
@@ -243,12 +203,12 @@ func (p *Plugin) Apply(ctx *core.Context) error {
Category: "messaging",
})
// 10. Optional runner start & lifecycle
// 11. Optional runner start & lifecycle
if p.autoStartRunner {
runnerCtx, cancel := context.WithCancel(ctx.GoContext())
p.cancelRunner = cancel
util.Go(func() {
_ = Start(runnerCtx)
_ = service.Start(runnerCtx)
})
}
@@ -261,3 +221,117 @@ func (p *Plugin) Apply(ctx *core.Context) error {
return nil
}
// Re-exported constants.
const (
CodeAlphabet = service.CodeAlphabet
CodeLength = service.CodeLength
)
// MessageChannel is an alias for model.MessageChannel.
type MessageChannel = model.MessageChannel
// MessageBinding is an alias for model.MessageBinding.
type MessageBinding = model.MessageBinding
// MessagePairingCode is an alias for model.MessagePairingCode.
type MessagePairingCode = model.MessagePairingCode
// PushChannel is an alias for model.PushChannel.
type PushChannel = model.PushChannel
// PushEvent is an alias for model.PushEvent.
type PushEvent = model.PushEvent
// PushHistory is an alias for model.PushHistory.
type PushHistory = model.PushHistory
// PushNotificationEvent is an alias for model.PushNotificationEvent.
type PushNotificationEvent = model.PushNotificationEvent
// ChannelConfig is an alias for model.ChannelConfig.
type ChannelConfig = model.ChannelConfig
// Capability is an alias for model.Capability.
type Capability = model.Capability
// Recipient is an alias for model.Recipient.
type Recipient = model.Recipient
// Attachment is an alias for model.Attachment.
type Attachment = model.Attachment
// InboundMessage is an alias for model.InboundMessage.
type InboundMessage = model.InboundMessage
// OutboundMessage is an alias for model.OutboundMessage.
type OutboundMessage = model.OutboundMessage
// BindingDTO is an alias for model.BindingDTO.
type BindingDTO = model.BindingDTO
// PublicChannelDTO is an alias for model.PublicChannelDTO.
type PublicChannelDTO = model.PublicChannelDTO
// Definition is an alias for model.Definition.
type Definition = model.Definition
// ChannelDTO is an alias for model.ChannelDTO.
type ChannelDTO = model.ChannelDTO
// CreateChannelRequest is an alias for model.CreateChannelRequest.
type CreateChannelRequest = model.CreateChannelRequest
// UpdateChannelRequest is an alias for model.UpdateChannelRequest.
type UpdateChannelRequest = model.UpdateChannelRequest
// PushDefinition is an alias for model.PushDefinition.
type PushDefinition = model.PushDefinition
// PushField is an alias for model.PushField.
type PushField = model.PushField
// NotificationMessage is an alias for model.NotificationMessage.
type NotificationMessage = model.NotificationMessage
// EventMetadata is an alias for model.EventMetadata.
type EventMetadata = model.EventMetadata
// SendPayload is an alias for model.SendPayload.
type SendPayload = model.SendPayload
// Handler is an alias for service.Handler.
type Handler = service.Handler
// Factory is an alias for service.Factory.
type Factory = service.Factory
// Channel is an alias for service.Channel.
type Channel = service.Channel
// Runner is an alias for service.Runner.
type Runner = service.Runner
// EventTrigger is an alias for service.EventTrigger.
type EventTrigger = service.EventTrigger
// PushHandler is an alias for service.PushHandler.
type PushHandler = service.PushHandler
// Re-exported variables and functions.
var (
SetDBServiceForTest = repository.SetDBServiceForTest
UpsertPairingCode = repository.UpsertPairingCode
Register = service.Register
Lookup = service.Lookup
GenerateCode = service.GenerateCode
NormalizeCode = service.NormalizeCode
FormatCode = service.FormatCode
Start = service.Start
Stop = service.Stop
GlobalRunner = service.GlobalRunner
DefaultTrigger = service.DefaultTrigger
SyncEvents = service.SyncEvents
AdminLogin = service.AdminLogin
HandleAdminLoggedIn = service.HandleAdminLoggedIn
)
@@ -1,15 +0,0 @@
// 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"
)
@@ -1,273 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"Wavelet/pkg/logger"
"Wavelet/pkg/util"
"context"
"encoding/json"
"errors"
"sync"
pkgpush "Wavelet/plugins/domain/message_gateway/push"
"gorm.io/gorm"
)
// NotificationMessage represents the structured notification message payload.
type NotificationMessage struct {
Title string `json:"title"`
Content string `json:"content"`
Level string `json:"level"`
Ext map[string]any `json:"ext,omitempty"`
}
// Flatten converts the structured NotificationMessage back to a flat map (original json structure).
func (m NotificationMessage) Flatten() map[string]any {
res := map[string]any{
keyTitle: m.Title,
keyContent: m.Content,
keyLevel: m.Level,
}
for k, v := range m.Ext {
res[k] = v
}
return res
}
// EventMetadata represents the metadata of a push notification event.
type EventMetadata struct {
Key string `json:"key"`
Name string `json:"name"`
DefaultTemplate NotificationMessage `json:"default_template"`
Description string `json:"description"`
}
// SendPayload 异步投递推送载荷 (供 task/Worker 使用)
type SendPayload struct {
EventKey string `json:"event_key"`
Config pkgpush.Config `json:"config"`
Target string `json:"target"`
Body NotificationMessage `json:"body"`
Template string `json:"template"`
}
var (
builtInEventsMu sync.RWMutex
// BuiltInEvents lists all built-in events defined in custom_events.
BuiltInEvents []EventMetadata
)
// RegisterBuiltInEvent registers a built-in event definition.
func RegisterBuiltInEvent(meta EventMetadata) {
builtInEventsMu.Lock()
defer builtInEventsMu.Unlock()
for i, e := range BuiltInEvents {
if e.Key == meta.Key {
BuiltInEvents[i] = meta
return
}
}
BuiltInEvents = append(BuiltInEvents, meta)
}
// GetBuiltInEvents returns a copy of registered built-in events.
func GetBuiltInEvents() []EventMetadata {
builtInEventsMu.RLock()
defer builtInEventsMu.RUnlock()
out := make([]EventMetadata, len(BuiltInEvents))
copy(out, BuiltInEvents)
return out
}
// EventTrigger represents the unified event trigger class.
type EventTrigger struct{}
// DefaultTrigger is the singleton instance of EventTrigger.
var DefaultTrigger = &EventTrigger{}
// Trigger receives event metadata and processes the event notification dispatch asynchronously.
//
//nolint:contextcheck
func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map[string]any) {
asyncCtx := context.WithoutCancel(ctx)
util.Go(func() {
if body == nil {
body = make(map[string]any)
}
if _, hasUser := body["user"]; !hasUser || body["user"] == nil {
body["user"] = getSystemUser(asyncCtx)
}
eventPtr, err := GetActivePushEventByKey(asyncCtx, meta.Key)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return
}
logger.ErrorF(asyncCtx, "push_event_trigger: failed to get active event %s: %v", meta.Key, err)
return
}
event := *eventPtr
if len(event.Channels) == 0 {
return
}
flatBody := getFlatBody(body)
msg, _ := t.buildMessage(&event, meta, flatBody, body)
t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody)
})
}
func (t *EventTrigger) buildMessage(event *PushEvent, meta EventMetadata, flatBody, body map[string]any) (NotificationMessage, string) {
var msg NotificationMessage
renderedTemplate := ""
templateSource := event.Template
if templateSource != "" {
var err error
msg, renderedTemplate, err = t.parseCustomTemplate(event, templateSource, flatBody)
if err != nil {
msg.Title = event.Name
msg.Content = renderedTemplate
msg.Level = defaultLevelInfo
}
} else {
msg = t.parseDefaultTemplate(meta, flatBody)
}
if msg.Ext == nil {
msg.Ext = make(map[string]any)
}
for k, v := range body {
if k == keyTitle || k == keyContent || k == keyLevel {
continue
}
if _, exists := msg.Ext[k]; !exists {
msg.Ext[k] = v
}
}
return msg, renderedTemplate
}
func (t *EventTrigger) parseCustomTemplate(event *PushEvent, templateSource string, flatBody map[string]any) (NotificationMessage, string, error) {
var msg NotificationMessage
renderedTemplate := pkgpush.ParseTemplate(templateSource, flatBody)
var tMap map[string]any
if err := json.Unmarshal([]byte(renderedTemplate), &tMap); err != nil {
return msg, renderedTemplate, err
}
if title, ok := tMap[keyTitle].(string); ok && title != "" {
msg.Title = title
} else {
msg.Title = event.Name
}
delete(tMap, keyTitle)
if content, ok := tMap[keyContent].(string); ok && content != "" {
msg.Content = content
} else {
msg.Content = renderedTemplate
}
delete(tMap, keyContent)
if level, ok := tMap[keyLevel].(string); ok && level != "" {
msg.Level = level
} else {
msg.Level = defaultLevelInfo
}
delete(tMap, keyLevel)
msg.Ext = tMap
return msg, renderedTemplate, nil
}
func (t *EventTrigger) parseDefaultTemplate(meta EventMetadata, flatBody map[string]any) NotificationMessage {
var msg NotificationMessage
msg.Title = pkgpush.ParseTemplate(meta.DefaultTemplate.Title, flatBody)
msg.Content = pkgpush.ParseTemplate(meta.DefaultTemplate.Content, flatBody)
msg.Level = pkgpush.ParseTemplate(meta.DefaultTemplate.Level, flatBody)
if meta.DefaultTemplate.Ext != nil {
msg.Ext = make(map[string]any)
for k, v := range meta.DefaultTemplate.Ext {
if strVal, ok := v.(string); ok {
msg.Ext[k] = pkgpush.ParseTemplate(strVal, flatBody)
} else {
msg.Ext[k] = v
}
}
}
return msg
}
func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta EventMetadata, event *PushEvent, msg NotificationMessage, flatBody map[string]any) {
for _, channelName := range event.Channels {
customChannel, err := GetActivePushChannelByName(ctx, channelName)
if err == nil {
t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody)
continue
}
logger.WarnF(ctx, "push_event_trigger: channel %q not found in DB or disabled: %v", channelName, err)
}
}
func (t *EventTrigger) enqueueCustomPushChannelTasks(ctx context.Context, meta EventMetadata, event *PushEvent, channel *PushChannel, msg NotificationMessage, flatBody map[string]any) {
if len(event.Targets) == 0 {
t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, "", msg)
return
}
for _, target := range event.Targets {
resolvedTarget := resolveTarget(ctx, target, flatBody, channel.Name)
t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, resolvedTarget, msg)
}
}
func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, meta EventMetadata, channel *PushChannel, target string, msg NotificationMessage) {
var config pkgpush.Config
var renderedTemplate string
switch channel.Type {
case channelLark:
config = pkgpush.Config{Channel: channelLark, URL: channel.URL, Secret: channel.Token}
renderedTemplate = channel.Other
case channelEmail:
url, token, other := resolveSMTPConfig(ctx, channel.URL, channel.Token, channel.Other)
config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other}
case channelTelegram:
config = pkgpush.Config{Channel: channelTelegram, URL: channel.URL, Secret: channel.Token, Key: channel.Other}
default:
config = pkgpush.Config{Channel: channelCustom, URL: channel.URL}
customPushReq := CustomPushRequest{
Title: msg.Title,
Content: msg.Content,
Description: meta.Description,
To: target,
}
if urlVal, ok := msg.Ext["url"].(string); ok {
customPushReq.URL = urlVal
}
renderedTemplate = renderCustomPayload(channel.Other, customPushReq)
}
payload := SendPayload{
EventKey: meta.Key,
Config: config,
Target: target,
Body: msg,
Template: renderedTemplate,
}
if err := enqueuePushTask(ctx, payload); err != nil {
logger.ErrorF(ctx, "push_event_trigger: enqueuePushTask failed for %s channel %s -> %s: %v", channel.Type, channel.Name, target, err)
}
}
// SyncEvents automatically registers/updates built-in events in the database.
func SyncEvents(ctx context.Context) error {
return syncBuiltInEvents(ctx)
}
@@ -1,556 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"Wavelet/core/contracts"
"context"
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
"gorm.io/gorm"
pkgpush "Wavelet/plugins/domain/message_gateway/push"
)
type smtpConfig struct {
Host string
Port string
Username string
Password string
}
func loadSMTPConfig(ctx context.Context) smtpConfig {
var cfg smtpConfig
var host, port, user, pass string
_ = getDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_host").Pluck("value", &host).Error
_ = getDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_port").Pluck("value", &port).Error
_ = getDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_username").Pluck("value", &user).Error
_ = getDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_password").Pluck("value", &pass).Error
cfg.Host = host
cfg.Port = port
cfg.Username = user
cfg.Password = pass
return cfg
}
func syncBuiltInEvents(ctx context.Context) error {
for _, meta := range GetBuiltInEvents() {
_, err := GetPushEventByKeyRecord(ctx, meta.Key)
if errors.Is(err, gorm.ErrRecordNotFound) {
var defaultTemplateStr string
if defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate); err == nil {
defaultTemplateStr = string(defaultTemplateBytes)
}
event := PushEvent{
EventKey: meta.Key,
Name: meta.Name,
Channels: []string{},
Targets: []string{},
Template: defaultTemplateStr,
Enabled: false,
}
if err := CreatePushEventRecord(ctx, &event); err != nil {
return err
}
} else if err != nil {
return err
}
}
return nil
}
func listPushEvents(ctx context.Context) ([]PushEvent, error) {
return ListPushEventsRecord(ctx)
}
func createPushEvent(ctx context.Context, req CreatePushEventRequest) (PushEvent, error) {
eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req)
if err != nil {
return PushEvent{}, err
}
count, err := CountPushEventsByKeyRecord(ctx, eventKey)
if err != nil {
return PushEvent{}, err
}
if count > 0 {
return PushEvent{}, errors.New("this notification event is already configured")
}
templateStr := strings.TrimSpace(req.Template)
if templateStr == "" {
templateStr = string(defaultTemplateBytes)
} else {
var tempMap map[string]any
if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil {
return PushEvent{}, errors.New("custom template is not a valid JSON format")
}
}
channels := req.Channels
if channels == nil {
channels = []string{}
}
targets := req.Targets
if targets == nil {
targets = []string{}
}
event := PushEvent{
EventKey: eventKey,
Name: eventName,
TaskType: req.TaskType,
Channels: channels,
Targets: targets,
Template: templateStr,
Enabled: req.Enabled,
}
if err := event.Validate(); err != nil {
return PushEvent{}, err
}
if err := CreatePushEventRecord(ctx, &event); err != nil {
return PushEvent{}, err
}
return event, nil
}
func deletePushEvent(ctx context.Context, id uint64) error {
event, err := GetPushEventByIDRecord(ctx, id)
if err != nil {
return err
}
return DeletePushEventRecord(ctx, &event)
}
func updatePushEvent(ctx context.Context, id uint64, req UpdatePushEventRequest) error {
event, err := GetPushEventByIDRecord(ctx, id)
if err != nil {
return err
}
event.Channels = req.Channels
event.Targets = req.Targets
event.Template = req.Template
event.Enabled = req.Enabled
if err := event.Validate(); err != nil {
return err
}
return SavePushEventRecord(ctx, &event)
}
func togglePushEvent(ctx context.Context, id uint64) (bool, error) {
event, err := GetPushEventByIDRecord(ctx, id)
if err != nil {
return false, err
}
enabled := !event.Enabled
if enabled && len(event.Channels) == 0 {
return false, errors.New("cannot enable event without any push channels configured")
}
if err := UpdatePushEventEnabledRecord(ctx, &event, enabled); err != nil {
return false, err
}
return enabled, nil
}
func listPushHistories(ctx context.Context, filter PushHistoryListFilter) (int64, []PushHistory, error) {
return ListPushHistoriesRecord(ctx, filter)
}
func applySMTPFallbackToPushConfig(ctx context.Context, cfg *pkgpush.Config) {
if cfg.Channel != channelEmail || (cfg.URL != "" && cfg.Key != "") {
return
}
smtp := loadSMTPConfig(ctx)
if smtp.Host == "" || smtp.Username == "" {
return
}
port := smtp.Port
if port == "" {
port = "587"
}
cfg.URL = smtp.Host + ":" + port
cfg.Key = smtp.Username
cfg.Secret = smtp.Password
}
func listPushChannels(ctx context.Context) ([]PushChannel, error) {
return ListPushChannelsRecord(ctx)
}
func createPushChannel(ctx context.Context, req CreatePushChannelRequest) (PushChannel, error) {
count, err := CountPushChannelsByNameRecord(ctx, req.Name)
if err != nil {
return PushChannel{}, err
}
if count > 0 {
return PushChannel{}, errors.New("channel name already exists")
}
channel := PushChannel{
Name: req.Name,
Description: req.Description,
Type: req.Type,
Token: req.Token,
URL: req.URL,
Other: req.Other,
Enabled: req.Enabled,
}
if err := channel.Validate(); err != nil {
return PushChannel{}, err
}
if err := CreatePushChannelRecord(ctx, &channel); err != nil {
return PushChannel{}, err
}
return channel, nil
}
func updatePushChannel(ctx context.Context, id uint64, req UpdatePushChannelRequest) (PushChannel, error) {
channel, err := GetPushChannelByIDRecord(ctx, id)
if err != nil {
return PushChannel{}, err
}
channel.Description = req.Description
channel.Type = req.Type
channel.Token = req.Token
channel.URL = req.URL
channel.Other = req.Other
channel.Enabled = req.Enabled
if err := channel.Validate(); err != nil {
return PushChannel{}, err
}
if err := SavePushChannelRecord(ctx, &channel); err != nil {
return PushChannel{}, err
}
return channel, nil
}
func deletePushChannel(ctx context.Context, id uint64) error {
channel, err := GetPushChannelByIDRecord(ctx, id)
if err != nil {
return err
}
return DeletePushChannelRecord(ctx, &channel)
}
func loadChannelForTest(ctx context.Context, req TestPushChannelRequest) (string, string, string, string, error) {
if req.Name != "" {
channel, err := GetPushChannelByNameRecord(ctx, req.Name)
if err != nil {
return "", "", "", "", errors.New("channel not found")
}
return channel.URL, channel.Token, channel.Other, channel.Type, nil
}
return req.URL, req.Token, req.Other, req.Type, nil
}
func listActivePushEventsByTaskType(ctx context.Context, taskType string) ([]PushEvent, error) {
return ListActivePushEventsByTaskTypeRecord(ctx, taskType)
}
func queryUser(ctx context.Context, fromService func(contracts.UserService) (*contracts.UserDTO, error), dbField string, dbVal any) (*contracts.UserDTO, error) {
if userSvc := getUserService(ctx); userSvc != nil {
return fromService(userSvc)
}
if db := getDB(ctx); db != nil {
var user contracts.UserDTO
if err := db.Table("w_users").Where(dbField+" = ?", dbVal).First(&user).Error; err == nil {
return &user, nil
}
}
return nil, errors.New("user not found")
}
func findUserByID(ctx context.Context, id uint64) (*contracts.UserDTO, error) {
return queryUser(ctx, func(s contracts.UserService) (*contracts.UserDTO, error) {
return s.GetUserByID(ctx, id)
}, "id", id)
}
func findUserByUsername(ctx context.Context, username string) (*contracts.UserDTO, error) {
return queryUser(ctx, func(s contracts.UserService) (*contracts.UserDTO, error) {
return s.GetUserByUsername(ctx, username)
}, "username", username)
}
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 := findUserByID(ctx, userID); err == nil && user != nil {
return user
}
}
if username := extractUsername(data); username != "" {
if user, err := findUserByUsername(ctx, username); err == nil && user != nil {
return user
}
}
return nil
}
func recordPushHistory(ctx context.Context, req SendPayload, status, errMsg string) error {
title := req.Body.Title
content := req.Body.Content
level := req.Body.Level
if title == "" {
title = "系统通知"
}
if level == "" {
level = defaultLevelInfo
}
target := req.Target
if target == "" {
if req.Config.URL != "" {
target = req.Config.URL
const maxTargetLen = 50
const truncatedLen = 47
if len(target) > maxTargetLen {
target = target[:truncatedLen] + "..."
}
} else {
target = "default"
}
}
history := PushHistory{
EventKey: req.EventKey,
Channel: req.Config.Channel,
Target: target,
Title: title,
Content: content,
Level: level,
Status: status,
ErrorMsg: errMsg,
}
return CreatePushHistoryRecord(ctx, &history)
}
func resolveTarget(ctx context.Context, target string, flatBody map[string]any, channel string) string {
target = strings.TrimSpace(target)
if target == "" {
return ""
}
resolved := resolveDynamicKeyword(target, flatBody)
if strings.Contains(resolved, "@") {
return resolved
}
if val, matched := resolveSystemTarget(ctx, resolved, channel); matched {
return val
}
user, found := resolveTargetUser(ctx, resolved, channel)
if !found {
return resolved
}
if channel == channelEmail && user.Email != "" {
return user.Email
}
if channel != channelEmail && user.Username != "" {
return user.Username
}
return resolved
}
func resolveDynamicKeyword(target string, flatBody map[string]any) string {
switch target {
case "user.id", "id":
if val, ok := flatBody["user.id"]; ok {
return fmt.Sprintf("%v", val)
}
if val, ok := flatBody["id"]; ok {
return fmt.Sprintf("%v", val)
}
case "user.username", "username":
if val, ok := flatBody["user.username"]; ok {
return fmt.Sprintf("%v", val)
}
if val, ok := flatBody["username"]; ok {
return fmt.Sprintf("%v", val)
}
case "user.email", channelEmail:
if val, ok := flatBody["user.email"]; ok {
return fmt.Sprintf("%v", val)
}
if val, ok := flatBody["email"]; ok {
return fmt.Sprintf("%v", val)
}
}
return target
}
func resolveTargetUser(ctx context.Context, resolved, _ string) (contracts.UserDTO, bool) {
if id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
if u, err := findUserByID(ctx, id); err == nil && u != nil {
return *u, true
}
}
if u, err := findUserByUsername(ctx, resolved); err == nil && u != nil {
return *u, true
}
return contracts.UserDTO{}, false
}
func getFirstAdminUser(ctx context.Context) (*contracts.UserDTO, error) {
if userSvc := getUserService(ctx); userSvc != nil {
return userSvc.GetFirstAdminUser(ctx)
}
if db := getDB(ctx); db != nil {
var adminUser contracts.UserDTO
if err := db.Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err == nil {
return &adminUser, nil
}
}
return nil, errors.New("no admin user found")
}
func resolveSystemTarget(ctx context.Context, resolved, channel string) (string, bool) {
if resolved != "系统" && resolved != "system" && resolved != "0" {
return "", false
}
adminUser, err := getFirstAdminUser(ctx)
if err != nil || adminUser == nil {
return resolved, true
}
if channel == channelEmail && adminUser.Email != "" {
return adminUser.Email, true
}
if channel != channelEmail && adminUser.Username != "" {
return adminUser.Username, true
}
return resolved, true
}
func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, string, string) {
if url != "" && token != "" {
return url, token, other
}
smtp := loadSMTPConfig(ctx)
if smtp.Host == "" || smtp.Username == "" {
return url, token, other
}
port := smtp.Port
if port == "" {
port = "587"
}
if url == "" {
url = smtp.Host + ":" + port
}
if token == "" {
token = smtp.Username
}
if other == "" {
other = smtp.Password
}
return url, token, other
}
func getSystemUser(ctx context.Context) *contracts.UserDTO {
if adminUser, err := getFirstAdminUser(ctx); err == nil && adminUser != nil {
return adminUser
}
return &contracts.UserDTO{
Username: "system",
Nickname: "系统管理员",
}
}
func findBuiltInEvent(key string) (EventMetadata, bool) {
for _, meta := range GetBuiltInEvents() {
if meta.Key == key {
return meta, true
}
}
return EventMetadata{}, false
}
func getEventInfo(req CreatePushEventRequest) (string, string, []byte, error) {
if req.TaskType != "" {
taskName := req.TaskType
if taskSvc := getTaskService(); taskSvc != nil {
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
taskName = meta.DisplayName
}
}
eventKey := "task_completed:" + req.TaskType
eventName := "任务完成: " + taskName
defaultTemplate := NotificationMessage{
Title: "任务完成: " + taskName,
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
}
if taskSvc := getTaskService(); taskSvc != nil {
_, err = taskSvc.Dispatch(ctx, "send_notification", payloadBytes, "system")
return err
}
return errors.New("task service not available")
}
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, 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
}
}
}
@@ -1,109 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"context"
"encoding/json"
"strconv"
"time"
)
func handleTaskCompleted(ctx context.Context, e contracts.TaskCompletedEvent) {
events, err := listActivePushEventsByTaskType(ctx, e.TaskType)
if err != nil {
logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", e.TaskType, err)
return
}
if len(events) == 0 {
return
}
body := map[string]any{
"task_id": e.TaskID,
"task_name": e.TaskName,
"task_type": e.TaskType,
"task_status": e.Status,
"task_duration": e.Duration,
"time": time.Now().Format("2006-01-02 15:04:05"),
"task_error": e.ErrorMsg,
"task_result": e.ResultMsg,
}
var payloadMap map[string]any
if e.Payload != "" {
if err := json.Unmarshal([]byte(e.Payload), &payloadMap); err == nil {
body["payload"] = payloadMap
extractUserFromMap(ctx, payloadMap, body)
}
}
if e.Detail != "" {
var detailMap map[string]any
if err := json.Unmarshal([]byte(e.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, 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 ""
}
@@ -1,107 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/message_gateway/push"
"context"
"encoding/json"
"errors"
"fmt"
)
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 = contracts.TaskMetaDTO{
Name: TaskTypeSendNotification,
DisplayName: "推送通知",
Description: "异步执行系统通知的多渠道派发与推送",
MaxRetry: 3,
Queue: "default",
Params: []contracts.TaskParamDTO{
{
Name: "event_key",
Type: "string",
Description: "事件标识 (如 admin_login)",
Required: true,
},
{
Name: "target",
Type: "string",
Description: "目标接收者",
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) error {
var req SendPayload
if err := json.Unmarshal(payload, &req); err != nil {
logger.ErrorF(ctx, "[Push] 解析推送参数失败: %v", err)
return fmt.Errorf("parse payload failed: %w", err)
}
logger.InfoF(ctx, "[Push] 开始推送通知: 事件 = %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)
logger.ErrorF(ctx, "[Push] 推送失败: %v", errWrap)
h.recordHistory(ctx, req, "failed", errWrap.Error())
return 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 {
logger.ErrorF(ctx, "[Push] 消息推送失败 (标题: %s): %v, 上游返回: %s", title, err, upstreamResp)
h.recordHistory(ctx, req, "failed", err.Error())
return fmt.Errorf("pusher.Send failed: %w", err)
}
logger.InfoF(ctx, "[Push] 消息推送成功 (标题: %s, 内容摘要: %s), 上游返回: %s", title, content, upstreamResp)
h.recordHistory(ctx, req, "success", "")
return nil
}
func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status, errMsg string) {
if dbErr := recordPushHistory(ctx, req, status, errMsg); dbErr != nil {
logger.ErrorF(ctx, "[Push] 写入推送历史审计记录失败: %v", dbErr)
}
}
@@ -1,26 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import "sync"
var (
factoriesMu sync.RWMutex
factories = map[string]Factory{}
)
// Register stores a channel factory under typ.
func Register(typ string, fn Factory) {
factoriesMu.Lock()
defer factoriesMu.Unlock()
factories[typ] = fn
}
// Lookup returns a previously registered factory.
func Lookup(typ string) (Factory, bool) {
factoriesMu.RLock()
defer factoriesMu.RUnlock()
fn, ok := factories[typ]
return fn, ok
}
@@ -1,382 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"Wavelet/pkg/idgen"
"context"
"errors"
"time"
"gorm.io/gorm"
)
const (
activePushChannelCacheTTL = 24 * time.Hour
activePushEventCacheTTL = 24 * time.Hour
)
// CreateMessageChannel inserts a channel row.
func CreateMessageChannel(ctx context.Context, ch *MessageChannel) error {
if ch.ID == 0 {
ch.ID = idgen.NextUint64ID()
}
return getDB(ctx).Create(ch).Error
}
// UpdateMessageChannel saves a channel row.
func UpdateMessageChannel(ctx context.Context, ch *MessageChannel) error {
return getDB(ctx).Save(ch).Error
}
// GetMessageChannel loads a channel by id.
func GetMessageChannel(ctx context.Context, id uint64) (*MessageChannel, error) {
var ch MessageChannel
if err := getDB(ctx).Where("id = ?", id).First(&ch).Error; err != nil {
return nil, err
}
return &ch, nil
}
// ListMessageChannels returns all channels newest first.
func ListMessageChannels(ctx context.Context) ([]MessageChannel, error) {
var rows []MessageChannel
if err := getDB(ctx).Order("id DESC").Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// DeleteMessageChannel removes pairings, bindings, then the channel.
func DeleteMessageChannel(ctx context.Context, id uint64) error {
return getDB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("channel_id = ?", id).Delete(&MessagePairingCode{}).Error; err != nil {
return err
}
if err := tx.Where("channel_id = ?", id).Delete(&MessageBinding{}).Error; err != nil {
return err
}
return tx.Delete(&MessageChannel{}, id).Error
})
}
// CreateMessageBinding inserts a binding.
func CreateMessageBinding(ctx context.Context, b *MessageBinding) error {
if b.ID == 0 {
b.ID = idgen.NextUint64ID()
}
return getDB(ctx).Create(b).Error
}
// GetBindingByChannelPlatform finds a binding for a platform user on a channel.
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*MessageBinding, error) {
var b MessageBinding
err := getDB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error
if err != nil {
return nil, err
}
return &b, nil
}
// ListBindingsByUser lists bindings for a Wavelet user.
func ListBindingsByUser(ctx context.Context, userID uint64) ([]MessageBinding, error) {
var rows []MessageBinding
if err := getDB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// GetMessageBinding loads a binding by id.
func GetMessageBinding(ctx context.Context, id uint64) (*MessageBinding, error) {
var b MessageBinding
if err := getDB(ctx).Where("id = ?", id).First(&b).Error; err != nil {
return nil, err
}
return &b, nil
}
// DeleteMessageBinding deletes a binding by id.
func DeleteMessageBinding(ctx context.Context, id uint64) error {
return getDB(ctx).Delete(&MessageBinding{}, id).Error
}
// UpsertPairingCode reuses an unexpired code for the same channel+platform user.
func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*MessagePairingCode, error) {
var existing MessagePairingCode
err := getDB(ctx).
Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()).
First(&existing).Error
if err == nil {
return &existing, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
row := &MessagePairingCode{
Code: code,
ChannelID: channelID,
PlatformUserID: platformUserID,
ExpiresAt: expiresAt,
}
if err := getDB(ctx).Create(row).Error; err != nil {
return nil, err
}
return row, nil
}
// GetPairingCode loads a pairing code by normalized code string.
func GetPairingCode(ctx context.Context, code string) (*MessagePairingCode, error) {
var row MessagePairingCode
if err := getDB(ctx).Where("code = ?", code).First(&row).Error; err != nil {
return nil, err
}
return &row, nil
}
// DeletePairingCode removes a pairing code.
func DeletePairingCode(ctx context.Context, code string) error {
return getDB(ctx).Where("code = ?", code).Delete(&MessagePairingCode{}).Error
}
// DeleteExpiredPairingCodes removes expired pairing rows.
func DeleteExpiredPairingCodes(ctx context.Context) error {
return getDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&MessagePairingCode{}).Error
}
// ListEnabledMessageChannels returns enabled channels.
func ListEnabledMessageChannels(ctx context.Context) ([]MessageChannel, error) {
var rows []MessageChannel
if err := getDB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// ListPushChannelsRecord returns all push channels ordered by creation time descending.
func ListPushChannelsRecord(ctx context.Context) ([]PushChannel, error) {
var channels []PushChannel
if err := getDB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
return nil, err
}
return channels, nil
}
// GetPushChannelByIDRecord loads a push channel by primary key.
func GetPushChannelByIDRecord(ctx context.Context, id uint64) (PushChannel, error) {
var channel PushChannel
if err := getDB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
return PushChannel{}, err
}
return channel, nil
}
// GetPushChannelByNameRecord 根据名称获取消息通道。
func GetPushChannelByNameRecord(ctx context.Context, name string) (*PushChannel, error) {
var channel PushChannel
if err := getDB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
return nil, err
}
return &channel, nil
}
// CountPushChannelsByNameRecord returns how many channels share the given name.
func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, error) {
var count int64
if err := getDB(ctx).Model(&PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CreatePushChannelRecord persists a new channel and invalidates cache.
func CreatePushChannelRecord(ctx context.Context, channel *PushChannel) error {
if err := getDB(ctx).Create(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
// SavePushChannelRecord updates a channel and invalidates cache.
func SavePushChannelRecord(ctx context.Context, channel *PushChannel) error {
if err := getDB(ctx).Save(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
// DeletePushChannelRecord removes a channel and invalidates cache.
func DeletePushChannelRecord(ctx context.Context, channel *PushChannel) error {
if err := getDB(ctx).Delete(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
func getCachedOrQuery[T any](ctx context.Context, cacheKey string, ttl time.Duration, query func(db *gorm.DB, dest *T) error) (*T, error) {
var val T
if cache := getCache(ctx); cache != nil {
if err := cache.Get(ctx, cacheKey, &val); err == nil {
return &val, nil
}
}
db := getDB(ctx)
if err := query(db, &val); err != nil {
return nil, err
}
if cache := getCache(ctx); cache != nil {
_ = cache.Set(ctx, cacheKey, val, ttl)
}
return &val, nil
}
// GetActivePushChannelByName 根据名称获取启用的消息通道 (优先从 Redis 缓存获取)。
func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel, error) {
return getCachedOrQuery(ctx, "push:channel:active:"+name, activePushChannelCacheTTL, func(db *gorm.DB, dest *PushChannel) error {
return db.Where("name = ? AND enabled = ?", name, true).First(dest).Error
})
}
// DeleteActivePushChannelCache 清理启用消息通道的缓存。
func DeleteActivePushChannelCache(ctx context.Context, name string) {
if cache := getCache(ctx); cache != nil {
_ = cache.Delete(ctx, "push:channel:active:"+name)
}
}
// ListPushEventsRecord returns all push events ordered by creation time descending.
func ListPushEventsRecord(ctx context.Context) ([]PushEvent, error) {
var events []PushEvent
if err := getDB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
return nil, err
}
return events, nil
}
// GetPushEventByIDRecord loads a push event by primary key.
func GetPushEventByIDRecord(ctx context.Context, id uint64) (PushEvent, error) {
var event PushEvent
if err := getDB(ctx).First(&event, id).Error; err != nil {
return PushEvent{}, err
}
return event, nil
}
// GetPushEventByKeyRecord loads a push event by event key.
func GetPushEventByKeyRecord(ctx context.Context, key string) (PushEvent, error) {
var event PushEvent
if err := getDB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
return PushEvent{}, err
}
return event, nil
}
// CountPushEventsByKeyRecord returns how many events use the given event key.
func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error) {
var count int64
if err := getDB(ctx).Model(&PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CreatePushEventRecord persists a new push event and invalidates cache.
func CreatePushEventRecord(ctx context.Context, event *PushEvent) error {
if err := getDB(ctx).Create(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// SavePushEventRecord updates a push event and invalidates cache.
func SavePushEventRecord(ctx context.Context, event *PushEvent) error {
if err := getDB(ctx).Save(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// UpdatePushEventEnabledRecord toggles the enabled flag for a push event.
func UpdatePushEventEnabledRecord(ctx context.Context, event *PushEvent, enabled bool) error {
event.Enabled = enabled
if err := getDB(ctx).Model(event).Update("enabled", enabled).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// DeletePushEventRecord removes a push event and invalidates cache.
func DeletePushEventRecord(ctx context.Context, event *PushEvent) error {
if err := getDB(ctx).Delete(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// ListActivePushEventsByTaskTypeRecord returns enabled events bound to a task type.
func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]PushEvent, error) {
var events []PushEvent
if err := getDB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil {
return nil, err
}
return events, nil
}
// GetActivePushEventByKey 获取启用的通知事件 (优先从 Redis 缓存获取)。
func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error) {
return getCachedOrQuery(ctx, "push:event:active:"+key, activePushEventCacheTTL, func(db *gorm.DB, dest *PushEvent) error {
return db.Where("event_key = ? AND enabled = ?", key, true).First(dest).Error
})
}
// DeleteActivePushEventCache 清理启用通知事件的缓存。
func DeleteActivePushEventCache(ctx context.Context, key string) {
if cache := getCache(ctx); cache != nil {
_ = cache.Delete(ctx, "push:event:active:"+key)
}
}
// ListPushHistoriesRecord returns paginated push history records.
func ListPushHistoriesRecord(ctx context.Context, filter PushHistoryListFilter) (int64, []PushHistory, error) {
query := getDB(ctx).Model(&PushHistory{}).Order("created_at DESC")
if filter.EventKey != "" {
query = query.Where("event_key = ?", filter.EventKey)
}
if filter.Status != "" {
query = query.Where("status = ?", filter.Status)
}
var total int64
if err := query.Count(&total).Error; err != nil {
return 0, nil, err
}
var results []PushHistory
offset := (filter.Page - 1) * filter.PageSize
if err := query.Offset(offset).Limit(filter.PageSize).Find(&results).Error; err != nil {
return 0, nil, err
}
return total, results, nil
}
// CreatePushHistoryRecord persists a push history audit record.
func CreatePushHistoryRecord(ctx context.Context, history *PushHistory) error {
return getDB(ctx).Create(history).Error
}
// PushHistoryQuery returns a scoped query builder for push histories.
func PushHistoryQuery(ctx context.Context) *gorm.DB {
return getDB(ctx).Model(&PushHistory{})
}
@@ -0,0 +1,324 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/message_gateway/errs"
"Wavelet/plugins/domain/message_gateway/model"
"context"
"sync"
"time"
"gorm.io/gorm"
)
const (
activePushChannelCacheTTL = 24 * time.Hour
activePushEventCacheTTL = 24 * time.Hour
)
var (
cacheMu sync.RWMutex
cacheSvc contracts.CacheService
)
// SetCacheService sets the cache service singleton.
func SetCacheService(s contracts.CacheService) {
cacheMu.Lock()
defer cacheMu.Unlock()
cacheSvc = s
}
// GetCache resolves the cache service for the current call.
func GetCache(ctx context.Context) contracts.CacheService {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
return s
}
}
cacheMu.RLock()
s := cacheSvc
cacheMu.RUnlock()
return s
}
// ListPushChannelsRecord returns all push channels ordered by creation time descending.
func ListPushChannelsRecord(ctx context.Context) ([]model.PushChannel, error) {
var channels []model.PushChannel
if err := GetDB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
return nil, err
}
return channels, nil
}
// GetPushChannelByIDRecord loads a push channel by primary key.
func GetPushChannelByIDRecord(ctx context.Context, id uint64) (model.PushChannel, error) {
var channel model.PushChannel
if err := GetDB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
return model.PushChannel{}, mapNotFound(err)
}
return channel, nil
}
// GetPushChannelByNameRecord loads a push channel by its unique name.
func GetPushChannelByNameRecord(ctx context.Context, name string) (*model.PushChannel, error) {
var channel model.PushChannel
if err := GetDB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
return nil, mapNotFound(err)
}
return &channel, nil
}
// CountPushChannelsByNameRecord returns how many channels share the given name.
func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, error) {
var count int64
if err := GetDB(ctx).Model(&model.PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CreatePushChannelRecord persists a new channel and invalidates cache.
func CreatePushChannelRecord(ctx context.Context, channel *model.PushChannel) error {
if err := GetDB(ctx).Create(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
// SavePushChannelRecord updates a channel and invalidates cache.
func SavePushChannelRecord(ctx context.Context, channel *model.PushChannel) error {
if err := GetDB(ctx).Save(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
// DeletePushChannelRecord removes a channel and invalidates cache.
func DeletePushChannelRecord(ctx context.Context, channel *model.PushChannel) error {
if err := GetDB(ctx).Delete(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
func getCachedOrQuery[T any](ctx context.Context, cacheKey string, ttl time.Duration, query func(db *gorm.DB, dest *T) error) (*T, error) {
var val T
if cache := GetCache(ctx); cache != nil {
if err := cache.Get(ctx, cacheKey, &val); err == nil {
return &val, nil
}
}
db := GetDB(ctx)
if err := query(db, &val); err != nil {
return nil, err
}
if cache := GetCache(ctx); cache != nil {
_ = cache.Set(ctx, cacheKey, val, ttl)
}
return &val, nil
}
// GetActivePushChannelByName loads an enabled push channel, preferring the cache layer.
func GetActivePushChannelByName(ctx context.Context, name string) (*model.PushChannel, error) {
channel, err := getCachedOrQuery(ctx, "push:channel:active:"+name, activePushChannelCacheTTL, func(db *gorm.DB, dest *model.PushChannel) error {
return db.Where("name = ? AND enabled = ?", name, true).First(dest).Error
})
if err != nil {
return nil, mapNotFound(err)
}
return channel, nil
}
// DeleteActivePushChannelCache drops the cached enabled-channel entry.
func DeleteActivePushChannelCache(ctx context.Context, name string) {
if cache := GetCache(ctx); cache != nil {
_ = cache.Delete(ctx, "push:channel:active:"+name)
}
}
// ListPushEventsRecord returns all push events ordered by creation time descending.
func ListPushEventsRecord(ctx context.Context) ([]model.PushEvent, error) {
var events []model.PushEvent
if err := GetDB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
return nil, err
}
return events, nil
}
// GetPushEventByIDRecord loads a push event by primary key.
func GetPushEventByIDRecord(ctx context.Context, id uint64) (model.PushEvent, error) {
var event model.PushEvent
if err := GetDB(ctx).First(&event, id).Error; err != nil {
return model.PushEvent{}, mapNotFound(err)
}
return event, nil
}
// GetPushEventByKeyRecord loads a push event by event key.
func GetPushEventByKeyRecord(ctx context.Context, key string) (model.PushEvent, error) {
var event model.PushEvent
if err := GetDB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
return model.PushEvent{}, mapNotFound(err)
}
return event, nil
}
// CountPushEventsByKeyRecord returns how many events use the given event key.
func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error) {
var count int64
if err := GetDB(ctx).Model(&model.PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CreatePushEventRecord persists a new push event and invalidates cache.
func CreatePushEventRecord(ctx context.Context, event *model.PushEvent) error {
if err := GetDB(ctx).Create(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// SavePushEventRecord updates a push event and invalidates cache.
func SavePushEventRecord(ctx context.Context, event *model.PushEvent) error {
if err := GetDB(ctx).Save(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// UpdatePushEventEnabledRecord toggles the enabled flag for a push event.
func UpdatePushEventEnabledRecord(ctx context.Context, event *model.PushEvent, enabled bool) error {
event.Enabled = enabled
if err := GetDB(ctx).Model(event).Update("enabled", enabled).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// DeletePushEventRecord removes a push event and invalidates cache.
func DeletePushEventRecord(ctx context.Context, event *model.PushEvent) error {
if err := GetDB(ctx).Delete(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// ListActivePushEventsByTaskTypeRecord returns enabled events bound to a task type.
func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]model.PushEvent, error) {
var events []model.PushEvent
if err := GetDB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil {
return nil, err
}
return events, nil
}
// GetActivePushEventByKey loads an enabled push event, preferring the cache layer.
func GetActivePushEventByKey(ctx context.Context, key string) (*model.PushEvent, error) {
event, err := getCachedOrQuery(ctx, "push:event:active:"+key, activePushEventCacheTTL, func(db *gorm.DB, dest *model.PushEvent) error {
return db.Where("event_key = ? AND enabled = ?", key, true).First(dest).Error
})
if err != nil {
return nil, mapNotFound(err)
}
return event, nil
}
// DeleteActivePushEventCache drops the cached enabled-event entry.
func DeleteActivePushEventCache(ctx context.Context, key string) {
if cache := GetCache(ctx); cache != nil {
_ = cache.Delete(ctx, "push:event:active:"+key)
}
}
// ListPushHistoriesRecord returns paginated push history records.
func ListPushHistoriesRecord(ctx context.Context, filter model.PushHistoryListFilter) (int64, []model.PushHistory, error) {
query := GetDB(ctx).Model(&model.PushHistory{}).Order("created_at DESC")
if filter.EventKey != "" {
query = query.Where("event_key = ?", filter.EventKey)
}
if filter.Status != "" {
query = query.Where("status = ?", filter.Status)
}
var total int64
if err := query.Count(&total).Error; err != nil {
return 0, nil, err
}
var results []model.PushHistory
offset := (filter.Page - 1) * filter.PageSize
if err := query.Offset(offset).Limit(filter.PageSize).Find(&results).Error; err != nil {
return 0, nil, err
}
return total, results, nil
}
// CreatePushHistoryRecord persists a push history audit record.
func CreatePushHistoryRecord(ctx context.Context, history *model.PushHistory) error {
return GetDB(ctx).Create(history).Error
}
// PushHistoryQuery returns a scoped query builder for push histories.
func PushHistoryQuery(ctx context.Context) *gorm.DB {
return GetDB(ctx).Model(&model.PushHistory{})
}
// LoadSMTPConfigRecord reads the SMTP settings owned by the system config table.
func LoadSMTPConfigRecord(ctx context.Context) model.SMTPConfig {
var cfg model.SMTPConfig
var host, port, user, pass string
_ = GetDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_host").Pluck("value", &host).Error
_ = GetDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_port").Pluck("value", &port).Error
_ = GetDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_username").Pluck("value", &user).Error
_ = GetDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_password").Pluck("value", &pass).Error
cfg.Host = host
cfg.Port = port
cfg.Username = user
cfg.Password = pass
return cfg
}
// FindUserByFieldRecord is the user lookup fallback for when the UserService
// contract is not wired yet. field comes from call sites, never from user input.
func FindUserByFieldRecord(ctx context.Context, field string, value any) (*contracts.UserDTO, error) {
db := GetDB(ctx)
if db == nil {
return nil, errs.ErrRecordNotFound
}
var user contracts.UserDTO
if err := db.Table("w_users").Where(field+" = ?", value).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
// FindFirstAdminUserRecord is the admin lookup fallback for when the UserService
// contract is not wired yet.
func FindFirstAdminUserRecord(ctx context.Context) (*contracts.UserDTO, error) {
db := GetDB(ctx)
if db == nil {
return nil, errs.ErrRecordNotFound
}
var adminUser contracts.UserDTO
if err := db.Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
return nil, err
}
return &adminUser, nil
}
@@ -0,0 +1,199 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package repository provides data persistence for the message_gateway plugin.
package repository
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/idgen"
"Wavelet/plugins/domain/message_gateway/errs"
"Wavelet/plugins/domain/message_gateway/model"
"context"
"errors"
"sync"
"time"
"gorm.io/gorm"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
)
// SetDBServiceForTest injects a DBService for tests. Production wiring must use Apply.
func SetDBServiceForTest(s contracts.DBService) {
SetDBService(s)
}
// SetDBService sets the database service singleton.
func SetDBService(s contracts.DBService) {
dbMu.Lock()
defer dbMu.Unlock()
dbSvc = s
}
// GetDB resolves the persistence handle for the current call, preferring an
// explicitly injected *core.Context before falling back to the plugin singleton.
func GetDB(ctx context.Context) *gorm.DB {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
return s.DB(ctx)
}
}
dbMu.RLock()
s := dbSvc
dbMu.RUnlock()
if s != nil {
return s.DB(ctx)
}
return nil
}
// mapNotFound translates GORM's missing-row sentinel into the plugin-level
// errs.ErrRecordNotFound so the service and handler layers stay free of gorm imports.
func mapNotFound(err error) error {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errs.ErrRecordNotFound
}
return err
}
// CreateMessageChannel inserts a channel row.
func CreateMessageChannel(ctx context.Context, ch *model.MessageChannel) error {
if ch.ID == 0 {
ch.ID = idgen.NextUint64ID()
}
return GetDB(ctx).Create(ch).Error
}
// UpdateMessageChannel saves a channel row.
func UpdateMessageChannel(ctx context.Context, ch *model.MessageChannel) error {
return GetDB(ctx).Save(ch).Error
}
// GetMessageChannel loads a channel by id.
func GetMessageChannel(ctx context.Context, id uint64) (*model.MessageChannel, error) {
var ch model.MessageChannel
if err := GetDB(ctx).Where("id = ?", id).First(&ch).Error; err != nil {
return nil, mapNotFound(err)
}
return &ch, nil
}
// ListMessageChannels returns all channels newest first.
func ListMessageChannels(ctx context.Context) ([]model.MessageChannel, error) {
var rows []model.MessageChannel
if err := GetDB(ctx).Order("id DESC").Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// DeleteMessageChannel removes pairings, bindings, then the channel.
func DeleteMessageChannel(ctx context.Context, id uint64) error {
return GetDB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("channel_id = ?", id).Delete(&model.MessagePairingCode{}).Error; err != nil {
return err
}
if err := tx.Where("channel_id = ?", id).Delete(&model.MessageBinding{}).Error; err != nil {
return err
}
return tx.Delete(&model.MessageChannel{}, id).Error
})
}
// CreateMessageBinding inserts a binding.
func CreateMessageBinding(ctx context.Context, b *model.MessageBinding) error {
if b.ID == 0 {
b.ID = idgen.NextUint64ID()
}
return GetDB(ctx).Create(b).Error
}
// GetBindingByChannelPlatform finds a binding for a platform user on a channel.
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) {
var b model.MessageBinding
err := GetDB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error
if err != nil {
return nil, mapNotFound(err)
}
return &b, nil
}
// ListBindingsByUser lists bindings for a Wavelet user.
func ListBindingsByUser(ctx context.Context, userID uint64) ([]model.MessageBinding, error) {
var rows []model.MessageBinding
if err := GetDB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// GetMessageBinding loads a binding by id.
func GetMessageBinding(ctx context.Context, id uint64) (*model.MessageBinding, error) {
var b model.MessageBinding
if err := GetDB(ctx).Where("id = ?", id).First(&b).Error; err != nil {
return nil, mapNotFound(err)
}
return &b, nil
}
// DeleteMessageBinding deletes a binding by id.
func DeleteMessageBinding(ctx context.Context, id uint64) error {
return GetDB(ctx).Delete(&model.MessageBinding{}, id).Error
}
// UpsertPairingCode reuses an unexpired code for the same channel+platform user.
func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*model.MessagePairingCode, error) {
var existing model.MessagePairingCode
err := GetDB(ctx).
Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()).
First(&existing).Error
if err == nil {
return &existing, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
row := &model.MessagePairingCode{
Code: code,
ChannelID: channelID,
PlatformUserID: platformUserID,
ExpiresAt: expiresAt,
}
if err := GetDB(ctx).Create(row).Error; err != nil {
return nil, err
}
return row, nil
}
// GetPairingCode loads a pairing code by normalized code string.
func GetPairingCode(ctx context.Context, code string) (*model.MessagePairingCode, error) {
var row model.MessagePairingCode
if err := GetDB(ctx).Where("code = ?", code).First(&row).Error; err != nil {
return nil, mapNotFound(err)
}
return &row, nil
}
// DeletePairingCode removes a pairing code.
func DeletePairingCode(ctx context.Context, code string) error {
return GetDB(ctx).Where("code = ?", code).Delete(&model.MessagePairingCode{}).Error
}
// DeleteExpiredPairingCodes removes expired pairing rows.
func DeleteExpiredPairingCodes(ctx context.Context) error {
return GetDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&model.MessagePairingCode{}).Error
}
// ListEnabledMessageChannels returns enabled channels.
func ListEnabledMessageChannels(ctx context.Context) ([]model.MessageChannel, error) {
var rows []model.MessageChannel
if err := GetDB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
@@ -1,52 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"Wavelet/pkg/logger"
"context"
"sync"
)
// 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
}
@@ -1,77 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"Wavelet/pkg/config"
"Wavelet/pkg/util"
"crypto/sha256"
"encoding/hex"
"encoding/json"
)
// 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)
}

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