feat(cordis): add OpenFlare Cordis 架构改造设计

docs(changelog): 修正表述笔误

refactor(cordis): 磁盘缓存改用上上游能力并清理本地副本

按上游/下游归属规约:类型断言守卫已回流 Wavelet(f3d85d5,附回归用例),
本仓库删除 OpenFlare/plugins/server/pkg/cache 整包并改 import 到
Wavelet/pkg/cache/disk,同步后与上游零漂移。

验证:go build 通过;go test ./... exit 0(137 包 ok);256 条路由对拍与
232 条 swagger 操作均零差异;make build-all 四进制;前端零改动。

docs(cordis): 记录 T1 清理结果与五个复用阻塞点

refactor(cordis): server 复用上游 pkg 能力并删除等价本地副本

按上游/下游归属规约清理重复实现,删除 7 个与上游等价的本地包并改 import:
shared/response→pkg/response、pkg/{logger,mail,trace,httppool,cache/ram}→
上游同名包、infra/persistence/batchwriter→pkg/batchwriter。逐项核过差异:
httppool 逐字节相同;logger 的 Config 字段完全一致;response 的 7 个 Abort*
一致;cache/ram 换过去顺带把裸 go 变回带 panic 恢复的 util.Go。

两处非等价差异按语义处理:
- batchwriter.Stats 与 status DTO 原为类型别名,改为消费侧逐字段转换,
  避免 model 反向依赖基础设施类型;
- 上游 pkg/idgen 要求显式 Init(本地副本为懒加载自动初始化),本次保留本地
  副本,待与 infra 初始化一并迁移(已登记在清理计划)。

验证:go build 通过;go test ./... exit 0(138 包 ok);256 条路由对拍零差异;
make swagger 232 条操作零增减,且归一化后与旧文档深度相等——差异仅为
response.Any / logger.LogEntry 两个定义名随包路径改名,接口形状未变。

chore(cordis): 回流内核与 pkg/util 通用能力并清理 vendoring 污染

按新增的上游/下游归属规约:HandleRaw/BasePath 与版本比较、网络、格式化助手
属通用能力,已提交到 Wavelet 分支 feat/cordis-router-raw-routes,本仓库改为
纯同步获取(pkg/util 已零漂移),补丁登记保留至上游合并。

同时修掉我此前 git add -A 造成的污染:首次 vendoring 把上游工作区里被
gitignore 的运行期产物一起提交进来(upload 的 diskcache 缓存块 650 个与
driver_http/dist 前端构建物 380 个,共 12872 行/1030 文件)。sync-upstream.sh
现显式排除 uploads/dist/data/*.db,.gitignore 补上对应兜底规则。

AGENTS.md 增加上游/下游改动归属规约,并把仍指向前 Cordis 布局的硬性约束
(internal/router + Serve、internal/repository/logstore、internal/platform/bootstrap、
internal/cmd)改到当前插件路径。

验证:go build 通过;go test ./... exit 0(144 包 ok);make swagger 232 条
操作与基线逐条一致;make build-all 四进制;gofmt 干净。

feat(cordis): server 插件化并改由内核挂载控制面路由

新增 plugins/server/plugin.go:Apply 以 ctx.Router().Group(app.api_prefix)
声明根级与 /v1 全部路由;33 个注册函数由 *gin.RouterGroup 改为
core.RouterExtension,RegisterCollection 改用内核新增的 HandleRaw 保留
尾部斜杠变体,AdminMiddlewares 返回 []any(Go 不允许把 []T 展开为 ...any)。
删除 router.Serve 与 registerRoutes,装配根改为 core.App +
driver_http.New(WithEngine(router.BuildEngine())),监听、信号与优雅退出归内核;
前端 SPA 的 NoRoute 兜底因内核暂无贡献点而保留在引擎层。

路由保真证据:plugin_parity_test 对拍 baseline/routes-engine.txt 的 256 条
(方法 路径) 零差异;go test ./... exit 0(144 包 ok,含真实 handler 的
openflare/integration 用例走同一条挂载路径);make swagger 232 条操作与基线
逐条一致;golangci-lint 0 issues;make build-all 四进制;embed_frontend
标签编译通过;前端零改动。

已知待补:带 Redis 的实机 HTTP 冒烟(本机 6379 未启动,session store 与
改造前一样在建店阶段即 fatal),以及 bootstrap 的任务/设置/迁移注册迁入 Apply。

feat(core): RouterExtension 增加 HandleRaw 与 BasePath 以保真尾部斜杠路由

server 插件化的前置:Handle 经 cleanPath 会剥掉尾部斜杠,无法表达
/resource 与 /resource/ 两条不同路由,而 OpenFlare 有 20 个历史 list
端点两者都注册且部署关闭了 RedirectTrailingSlash,缺失即 404。新增
HandleRaw 与 BasePath(作用域包装器同样登记反注册),补 extpoints 用例;
并把 router.Serve 拆出 BuildEngine 以便交给 driver_http.WithEngine 复用,
新增路由表导出 harness,固化 256 条 (方法 路径) 基线供插件化对拍。
上游补丁登记于 backend/OpenFlare/upstream-patches.md,同步脚本改为按目录
前缀输出差异并在同步后提醒确认补丁是否仍在。

验证:go build 通过;go test ./... exit 0(143 包 ok);gofmt 干净。

docs(cordis): 记录 server 插件接入内核的可行路径与内核能力缺口

feat(cordis): agent/relay/flared 落地为内核驱动插件

三个边缘守护进程各新增 plugin.go,实现 core.Plugin + core.Driver
(自定义 DriverType 与同名 profile),装配与生命周期从 main 迁入
Apply/Start/Stop:Apply 负责 JSON 配置加载、运行环境与用户确保、
openresty/frps/frpc 管理器与各服务装配;Start 以 util.Go 拉起阻塞式
runner 与 GeoIP 周期更新;Stop 收敛主循环结果并在超时时报错而非静默。

入口改为 core.NewApp(core.WithProfile(...)) + Prepare/Run,保持
-config 旗标、默认路径、退出码与启动/停止日志不变。

验证:go build 通过;go test ./... exit 0(143 包 ok,含 3 个插件身份
与配置失败路径测试);make build-all 四进制产出;三进制实跑缺失配置
均 exit 1 且错误链保留 load {agent,relay,flared} config 原因;gofmt 干净。

refactor(cordis): 按功能职责拆分为 4 个插件与 share 共享层

backend/OpenFlare 不再平铺遗留分层,改为 plugins/{server,agent,relay,flared}
加 share/:控制面业务(openflare/admin/oauth/user/upload/cap/config/health 与
repository/model/infra/router 等支撑层)归 server;三个边缘守护进程各自成插件;
被两个以上插件消费的 protocol/geoip/wsclient/render/pagesarchive/edge 归 share。
同时把 pkg/util 与 buildinfo 合并回上游 pkg(上游已覆盖全部符号,仅 8 个函数与
2 个类型为 OpenFlare 独有,已一并迁入),装配根统一到 backend/cmd(含三个 daemon
入口),Dockerfile 与 release 工作流的构建路径和 -X 注入路径同步更新。

验证:go build 通过;go test ./... exit 0(141 包 ok);make swagger exit 0 且
232 条 API 操作与基线逐条一致;make build-all 产出 4 进制;-X 注入经二进制
strings 实测生效;日志后端直连门禁改写为按 server 插件业务域扫描并在扫描数为 0
时报错(防门禁静默失效);前端零改动。

feat(cordis): 落地 backend/share 共享层与上游同步脚本

跨插件共享资源(控制消息协议、GeoIP+iputil、边缘守护进程日志)从下游包
移入 backend/share,并声明其只能依赖 core/pkg 与标准/第三方库,禁止反向
引用下游业务与具体插件实现;新增 scripts/sync-upstream.sh 只覆盖
backend/{core,pkg,plugins},同步后 --check 报告零差异,证明与上游逐字一致。

go build 通过,go test ./... exit 0(142 包 ok),前端零改动。

refactor(cordis): 采用与 Wavelet 同构的单模块布局并引入上游内核

按上游结构落位:backend/{core,pkg,plugins} 为 Wavelet 上游拷贝,OpenFlare
全部业务收拢到上游 downstream 所对应的位置 backend/OpenFlare/,模块名保持
Wavelet 以保证上游 import 路径逐字一致、同步零改写;三个 daemon 入口移至
backend/OpenFlare/cmd,backend/cmd 与 main.go 作为控制面装配根。

行为不变:go build 通过,142 个测试包全绿(含上游插件测试),232 条 API
操作与改造前逐条一致,四进制产物正常,前端零改动。swagger 暂只扫描下游代码,
待 P4 挂载上游路由后再纳入 plugins/。

style: 修正模块路径改写导致的 import 分组排序漂移

refactor(layout): Go 代码迁入 backend/ 并将模块名简化为 OpenFlare

对齐上游 Wavelet 的仓库布局,为以第二 module 形态 vendoring Cordis 内核与
平台插件做准备:模块路径整体改写为 OpenFlare,Go 目标加 cd backend,
swaggo 产物移至 backend/docs 并把 json/yaml 复制回 docs/ 供站点消费,
Dockerfile 与 release 工作流的构建目录、ldflags 模块路径同步更新。

行为保持不变:232 条路由与改造前逐条一致,95 个测试包全绿,
四进制产物正常,前端零改动。

chore(cordis): 落地改造计划与 schema/路由基线

新增 legacy_dump_test 迁移快照 harness:在临时 sqlite 库上按生产顺序
(goose.UpTo → zone 导入 → goose.Up)跑完 76 个历史迁移并导出 schema 与
版本序列,作为改造前后一致性门禁的唯一事实来源。同时记录 232 条路由清单
与 foundation 实施计划。

docs(cordis): add OpenFlare Cordis 架构改造设计

明确上游以第二 module 形态 vendoring 进 backend/Wavelet、4 个插件
(server/agent/relay/flared) 全部装载内核,并规定保留 76 个历史 goose
迁移 + 一次性版本 stamp 桥接的迁移方案,配套三方 schema 一致性门禁,
确保已部署库不重跑历史、不丢数据。
This commit is contained in:
ryan
2026-08-29 19:28:39 +08:00
parent 9f79fb9969
commit dbaa3bf140
1327 changed files with 91634 additions and 4157 deletions
+12
View File
@@ -0,0 +1,12 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import "Wavelet/plugins/domain/admin/model"
// DatabaseConfig aliases model.DatabaseConfig.
type DatabaseConfig = model.DatabaseConfig
// ClickHouseConfig aliases model.ClickHouseConfig.
type ClickHouseConfig = model.ClickHouseConfig
+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
}
@@ -0,0 +1,76 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"Wavelet/pkg/response"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/service"
"net/http"
"github.com/gin-gonic/gin"
)
// GetCacheStatus 获取磁盘缓存状态与当前统计数据
// @Summary 获取缓存状态
// @Description 获取当前系统磁盘缓存的使用情况(已占用字节、Key 数量等)与策略配置
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=disk.Status} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/cache/status [get]
func GetCacheStatus(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(service.DiskCacheStatus()))
}
// UpdateCacheConfig 更新磁盘缓存策略配置
// @Summary 更新缓存配置
// @Description 更改磁盘缓存最大容量限制、文件生存时间(TTL)以及是否启用 LRU 淘汰淘汰算法,并进行热更新
// @Tags admin
// @Accept json
// @Produce json
// @Param request body model.UpdateCacheConfigRequest true "缓存配置请求体"
// @Security SessionCookie
// @Success 200 {object} response.Any "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/cache/config [post]
func UpdateCacheConfig(c *gin.Context) {
var req model.UpdateCacheConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := service.UpdateDiskCachePolicy(c.Request.Context(), req); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// ClearCache 一键清空所有磁盘缓存数据
// @Summary 清空缓存
// @Description 清除系统磁盘缓存目录中的所有临时文件,并重置缓存容量和 Key 追踪数据
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any "清理成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/cache/clear [post]
func ClearCache(c *gin.Context) {
if err := service.ClearDiskCache(); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -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/logger"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/service"
"errors"
"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 errors.Is(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 !service.GetDBConfig().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,205 @@
// 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"
"errors"
"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)
},
}
}
// errNegativeParam 表示查询参数解析出了负数。
var errNegativeParam = errors.New("parameter must not be negative")
// parsePositiveInt 解析非负整数查询参数;返回错误时 result 保持调用前的值。
func parsePositiveInt(s string, result *int) error {
if s == "" {
*result = 0
return nil
}
n, err := strconv.Atoi(s)
if err != nil {
return err
}
if n < 0 {
return errNegativeParam
}
*result = n
return nil
}
@@ -0,0 +1,41 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"testing"
)
func TestParsePositiveInt(t *testing.T) {
const untouched = 77
tests := []struct {
name string
input string
want int
wantErr bool
}{
{name: "empty means zero", input: "", want: 0},
{name: "zero accepted", input: "0", want: 0},
{name: "positive accepted", input: "42", want: 42},
{name: "negative rejected", input: "-5", want: untouched, wantErr: true},
{name: "oversized rejected", input: "99999999999999999999", want: untouched, wantErr: true},
{name: "non numeric rejected", input: "abc", want: untouched, wantErr: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
got := untouched
err := parsePositiveInt(tt.input, &got)
if (err != nil) != tt.wantErr {
t.Fatalf("parsePositiveInt(%q) error = %v, wantErr %v", tt.input, err, tt.wantErr)
}
if got != tt.want {
t.Errorf("parsePositiveInt(%q) left result %d, want %d", tt.input, got, tt.want)
}
})
}
}
@@ -0,0 +1,46 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/trace"
"Wavelet/plugins/domain/admin/errs"
"github.com/gin-gonic/gin"
)
// LoginAdminRequired 返回管理员权限校验中间件
func LoginAdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
ctx, span := trace.Start(c.Request.Context(), "LoginAdminRequired")
defer span.End()
user, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if user == nil {
response.AbortNotFound(c, errs.AdminRequired)
return
}
// 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限
if tokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
tokenAdmin, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
if !tokenAdmin {
response.AbortNotFound(c, errs.TokenAdminRequired)
return
}
}
if !user.IsAdmin {
response.AbortNotFound(c, errs.AdminRequired)
return
}
logger.InfoF(ctx, "[LoginAdminRequired] %d %s", user.ID, user.Username)
c.Next()
}
}
@@ -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())))
}
@@ -0,0 +1,318 @@
// 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"
"strconv"
"github.com/gin-gonic/gin"
)
// 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 返回系统支持的所有可调度任务类型列表,包括任务名称、描述、是否支持时间范围等元数据,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]contracts.TaskMetaDTO} "任务类型列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/types [get]
func ListTaskTypes(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(service.ListTaskTypes()))
}
// DispatchTask 下发任务
// @Summary 下发异步任务
// @Description 手动触发指定类型的异步任务,支持指定时间范围和用户,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body model.DispatchTaskRequest true "任务请求参数"
// @Success 200 {object} response.Any{data=string} "任务已入队"
// @Failure 400 {object} response.Any "任务类型不存在或参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "任务入队失败"
// @Router /api/v1/admin/tasks/dispatch [post]
func DispatchTask(c *gin.Context) {
var req model.DispatchTaskRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
taskID, err := service.DispatchTask(c.Request.Context(), req)
if err != nil {
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
}
c.JSON(http.StatusOK, response.OK(taskID))
}
// ListTaskExecutions 查询任务执行记录列表
// @Summary 查询任务执行记录
// @Description 分页查询任务执行记录,支持按状态和任务类型筛选,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param status query string false "状态筛选 (pending/running/succeeded/failed)"
// @Param task_type query string false "任务类型筛选"
// @Param page query int false "页码" default(1)
// @Param page_size query int false "每页条数" default(20)
// @Success 200 {object} response.Any{data=object} "任务执行记录列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/executions [get]
func ListTaskExecutions(c *gin.Context) {
var req model.ListTaskExecutionsRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
executions, total, err := service.ListTaskExecutions(c.Request.Context(), req)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(gin.H{
"items": executions,
"total": total,
"page": req.Page,
"page_size": req.PageSize,
}))
}
// GetTaskExecution 查询单条任务执行详情
// @Summary 查询任务执行详情
// @Description 根据 ID 查询任务执行记录详情,包含完整执行日志,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "任务执行记录 ID"
// @Success 200 {object} response.Any{data=model.TaskExecution} "任务执行详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "记录不存在"
// @Router /api/v1/admin/tasks/executions/{id} [get]
func GetTaskExecution(c *gin.Context) {
id, err := parseUintParam(c, errs.InvalidTaskExecutionID)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
execution, err := service.TaskExecution(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, errs.TaskNotFound)
return
}
c.JSON(http.StatusOK, response.OK(execution))
}
// RetryTask 重试失败的任务
// @Summary 重试失败任务
// @Description 重新下发一条失败的任务,创建新的执行记录,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "任务执行记录 ID"
// @Success 200 {object} response.Any{data=string} "新任务的 TaskID"
// @Failure 400 {object} response.Any "任务不支持重试或参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "记录不存在"
// @Failure 500 {object} response.Any "重试失败"
// @Router /api/v1/admin/tasks/executions/{id}/retry [post]
func RetryTask(c *gin.Context) {
id, err := parseUintParam(c, errs.InvalidTaskExecutionID)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
newTaskID, err := service.RetryTask(c.Request.Context(), id)
if err != nil {
switch {
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, err.Error())
}
return
}
c.JSON(http.StatusOK, response.OK(newTaskID))
}
// ListSchedules 获取定时任务列表
// @Summary 获取定时任务列表
// @Description 返回系统所有的定时任务配置列表,包括名称、关联的异步任务类型、Cron 表达式和启用状态,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.Schedule} "定时任务列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/schedules [get]
func ListSchedules(c *gin.Context) {
schedules, err := service.ListSchedules(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(schedules))
}
// CreateSchedule 创建定时任务
// @Summary 创建定时任务
// @Description 新增一个动态定时任务配置,关联已有的异步任务,配置 Cron 表达式和执行参数,并触发调度器热加载,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @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 model.CreateScheduleRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
schedule, err := service.CreateSchedule(c.Request.Context(), req)
if abortTaskLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(schedule))
}
// UpdateSchedule 修改定时任务
// @Summary 修改定时任务
// @Description 修改一个定时任务的配置(名称、Cron 表达式、异步任务参数和是否启用等),并触发调度器热加载,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "定时任务 ID"
// @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 "无管理员权限"
// @Failure 404 {object} response.Any "定时任务不存在"
// @Failure 500 {object} response.Any "修改定时任务失败"
// @Router /api/v1/admin/tasks/schedules/{id} [put]
func UpdateSchedule(c *gin.Context) {
id, err := parseUintParam(c, errs.InvalidScheduleID)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
var req model.UpdateScheduleRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
schedule, err := service.UpdateSchedule(c.Request.Context(), id, req)
if abortTaskLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(schedule))
}
// DeleteSchedule 删除定时任务
// @Summary 删除定时任务
// @Description 删除指定的定时任务配置,并触发调度器热加载,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "定时任务 ID"
// @Success 200 {object} response.Any{data=string} "删除结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "删除定时任务失败"
// @Router /api/v1/admin/tasks/schedules/{id} [delete]
func DeleteSchedule(c *gin.Context) {
id, err := parseUintParam(c, errs.InvalidScheduleID)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := service.DeleteSchedule(c.Request.Context(), id); err != nil {
response.AbortInternal(c, err.Error())
return
}
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
}
@@ -0,0 +1,97 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler_test
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/plugins/domain/admin/handler"
"Wavelet/plugins/domain/admin/service"
"Wavelet/plugins/drivers/driver_asynq_worker"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type listTaskTypesResponse struct {
ErrorMsg string `json:"error_msg"`
Data []contracts.TaskMetaDTO `json:"data"`
}
func TestListTaskTypesHandler(t *testing.T) {
gin.SetMode(gin.TestMode)
ctx := core.NewContext(t.Context())
worker := driver_asynq_worker.New()
require.NoError(t, worker.Apply(ctx))
ctx.Task().Register("logs:db_switch", func(_ any) error { return nil },
extpoints.WithTaskType("logs_db_switch"),
extpoints.WithTaskName("切换日志数据库"),
extpoints.WithTaskDescription("复制迁移用户访问日志并在成功后切换日志主库"),
extpoints.WithTaskCategory("system"),
extpoints.WithTaskRetry(3),
extpoints.WithTaskQueue("default"),
extpoints.WithTaskRetryable(true),
extpoints.WithTaskParams(contracts.TaskParamDTO{
Name: "target",
Label: "目标日志库",
Type: "string",
Required: true,
Placeholder: "postgres|sqlite|clickhouse",
Description: "迁移目标",
}),
)
taskSvc, err := core.Inject[contracts.TaskService](ctx)
require.NoError(t, err)
service.SetTaskService(taskSvc)
r := gin.New()
r.GET("/api/v1/admin/tasks/types", handler.ListTaskTypes)
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/tasks/types", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp listTaskTypesResponse
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.Empty(t, resp.ErrorMsg)
require.NotEmpty(t, resp.Data)
var found bool
for _, task := range resp.Data {
if task.Type == "logs_db_switch" {
found = true
assert.Equal(t, "logs:db_switch", task.AsynqTask)
assert.Equal(t, "切换日志数据库", task.Name)
assert.Equal(t, "复制迁移用户访问日志并在成功后切换日志主库", task.Description)
assert.Equal(t, "system", task.Category)
assert.Equal(t, 3, task.MaxRetry)
assert.Equal(t, "default", task.Queue)
assert.True(t, task.Retryable)
require.Len(t, task.Params, 1)
assert.Equal(t, "target", task.Params[0].Name)
assert.Equal(t, "目标日志库", task.Params[0].Label)
assert.Equal(t, "string", task.Params[0].Type)
assert.True(t, task.Params[0].Required)
assert.Equal(t, "postgres|sqlite|clickhouse", task.Params[0].Placeholder)
}
}
assert.True(t, found, "expected logs_db_switch in task types")
// Ensure NO task in the list has an empty Type or Name, which breaks SelectItem key/value in frontend
for _, task := range resp.Data {
assert.NotEmpty(t, task.Type, "task.type must never be empty")
assert.NotEmpty(t, task.Name, "task.name must never be empty")
}
}
@@ -0,0 +1,159 @@
// 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"
)
// 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, errs.ErrTemplateNotFound) {
response.AbortNotFound(c, errs.TemplateNotFound)
return true
}
msg := err.Error()
switch msg {
case errs.TemplateKeyExists, errs.SystemTemplateCannotDelete:
response.AbortBadRequest(c, msg)
return true
}
response.AbortInternal(c, msg)
return true
}
// CreateTemplate 创建模板
// @Summary 创建模板
// @Description 创建一条新的自定义通知模板,模板标识符(Key)不可重复,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body model.CreateTemplateRequest true "创建请求参数"
// @Success 200 {object} response.Any{data=string} "创建成功"
// @Failure 400 {object} response.Any "参数错误或模板标识符已存在"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates [post]
func CreateTemplate(c *gin.Context) {
var req model.CreateTemplateRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
tmpl, err := service.CreateTemplate(c.Request.Context(), req)
if abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(tmpl))
}
// ListTemplates 获取模板列表
// @Summary 获取模板列表
// @Description 返回所有通知模板列表,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.Template} "模板列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates [get]
func ListTemplates(c *gin.Context) {
templates, err := service.ListTemplates(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(templates))
}
// GetTemplate 获取单个模板
// @Summary 获取单个模板
// @Description 根据模板标识符获取对应的模板详情,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Success 200 {object} response.Any{data=model.Template} "模板详情"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "模板不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates/{key} [get]
func GetTemplate(c *gin.Context) {
tmpl, err := service.GetTemplate(c.Request.Context(), c.Param("key"))
if abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(tmpl))
}
// UpdateTemplate 更新模板
// @Summary 更新模板
// @Description 根据模板标识符更新对应的模板内容,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Param request body 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 "无管理员权限"
// @Failure 404 {object} response.Any "模板不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates/{key} [put]
func UpdateTemplate(c *gin.Context) {
var req model.UpdateTemplateRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
tmpl, err := service.UpdateTemplate(c.Request.Context(), c.Param("key"), req)
if abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(tmpl))
}
// DeleteTemplate 删除模板
// @Summary 删除模板
// @Description 根据模板标识符删除对应模板,系统预置模板不可删除,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Failure 400 {object} response.Any "不可删除系统模板"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "模板不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates/{key} [delete]
func DeleteTemplate(c *gin.Context) {
if err := service.DeleteTemplate(c.Request.Context(), c.Param("key")); abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -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)
}
})
}
@@ -0,0 +1,294 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
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"
"github.com/gin-gonic/gin"
)
func parseUserID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil || id == 0 {
response.AbortBadRequest(c, errs.UserNotFound)
return 0, false
}
return id, true
}
// 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, errs.ErrUserServiceUnavailable) {
response.AbortInternal(c, err.Error())
return true
}
if errors.Is(err, errs.ErrUserNotFound) {
response.AbortNotFound(c, notFoundMsg)
return true
}
msg := err.Error()
for _, m := range badRequestMsgs {
if msg == m {
response.AbortBadRequest(c, msg)
return true
}
}
for _, m := range forbiddenMsgs {
if msg == m {
response.AbortForbidden(c, msg)
return true
}
}
logger.ErrorF(c.Request.Context(), "Admin user error: %v", err)
response.AbortInternal(c, errs.InternalServerError)
return true
}
// ListUsers 获取用户列表
// @Summary 获取用户列表
// @Description 分页返回用户列表,支持按用户 ID 和用户名筛选,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @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 model.ListUsersRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
total, dtos, err := service.AdminListUsers(c.Request.Context(), contracts.AdminListUsersFilter{
Page: req.Page,
PageSize: req.PageSize,
UserID: req.UserID,
Username: req.Username,
Email: req.Email,
})
if err != nil {
response.AbortInternal(c, err.Error())
return
}
users := make([]model.UserResponse, 0, len(dtos))
for _, dto := range dtos {
users = append(users, service.ToUserResponse(dto))
}
c.JSON(http.StatusOK, response.OK(model.ListUsersResponse{
Users: users,
Total: total,
}))
}
// GetUser 获取用户详情
// @Summary 获取用户详情
// @Description 返回指定用户的完整个人资料和系统状态,需要管理员权限,不返回密码等敏感字段
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Success 200 {object} response.Any{data=model.UserResponse} "用户详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "用户不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id} [get]
func GetUser(c *gin.Context) {
id, ok := parseUserID(c)
if !ok {
return
}
targetUser, err := service.AdminGetUser(c.Request.Context(), id)
if abortUserLogicError(c, err, errs.UserNotFound, nil, nil) {
return
}
c.JSON(http.StatusOK, response.OK(service.ToUserResponse(targetUser)))
}
// UpdateUserStatus 更新用户状态(启用/禁用)
// @Summary 更新用户状态
// @Description 启用或禁用指定用户,管理员账号无法被禁用,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Param request body model.UpdateUserStatusRequest true "状态参数"
// @Success 200 {object} response.Any{data=string} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限或尝试禁用管理员"
// @Failure 404 {object} response.Any "用户不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id}/status [put]
func UpdateUserStatus(c *gin.Context) {
var req model.UpdateUserStatusRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
id, ok := parseUserID(c)
if !ok {
return
}
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, errs.UpdateUserFailed)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// DeleteUser 删除用户
// @Summary 删除用户
// @Description 删除指定非管理员用户,需要管理员权限,不能删除当前登录用户
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限、尝试删除管理员或当前用户"
// @Failure 404 {object} response.Any "用户不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id} [delete]
func DeleteUser(c *gin.Context) {
id, ok := parseUserID(c)
if !ok {
return
}
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if currUser == nil {
response.AbortUnauthorized(c, errs.AdminRequired)
return
}
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, errs.DeleteUserFailed)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// CreateUser 创建用户
// @Summary 创建用户
// @Description 创建一个本地密码登录的新用户,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @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 model.CreateUserRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
newUser, err := service.AdminCreateUser(c.Request.Context(), contracts.AdminCreateUserRequest{
Username: req.Username,
Password: req.Password,
Nickname: req.Nickname,
Email: req.Email,
IsActive: req.IsActive,
IsAdmin: req.IsAdmin,
})
if abortUserLogicError(c, err, "", nil, []string{errs.UsernameRequired, errs.EmailRequired, errs.PasswordTooShort, errs.UsernameExists, errs.EmailExists}) {
return
}
c.JSON(http.StatusOK, response.OK(service.ToUserResponse(newUser)))
}
// UpdateUser 更新用户信息
// @Summary 更新用户信息
// @Description 更新指定用户的昵称、邮箱、管理员权限,并可选重置密码,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Param request body model.UpdateUserRequest true "更新参数"
// @Success 200 {object} response.Any{data=string} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限或尝试修改自身权限"
// @Failure 404 {object} response.Any "用户不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id} [put]
func UpdateUser(c *gin.Context) {
var req model.UpdateUserRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
id, ok := parseUserID(c)
if !ok {
return
}
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if currUser == nil {
response.AbortUnauthorized(c, errs.AdminRequired)
return
}
err := service.AdminUpdateUser(c.Request.Context(), currUser.ID, contracts.AdminUpdateUserRequest{
ID: id,
Nickname: req.Nickname,
Email: req.Email,
IsAdmin: req.IsAdmin,
Password: req.Password,
})
if err != nil {
if abortUserLogicError(c, err, errs.UserNotFound, []string{errs.CannotRevokeSelfAdmin}, []string{errs.EmailRequired, errs.EmailExists, errs.PasswordTooShort}) {
return
}
response.AbortInternal(c, errs.UpdateUserInfoFailed)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,131 @@
-- +goose Up
-- +goose StatementBegin
CREATE TABLE IF NOT EXISTS w_system_configs (
key VARCHAR(64) PRIMARY KEY,
value TEXT NOT NULL,
type VARCHAR(32) NOT NULL DEFAULT 'system',
visibility INTEGER NOT NULL DEFAULT 0,
description VARCHAR(255),
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE TABLE IF NOT EXISTS w_templates (
id BIGINT PRIMARY KEY,
key VARCHAR(80) NOT NULL UNIQUE,
name VARCHAR(100) NOT NULL,
type VARCHAR(20) NOT NULL DEFAULT 'email',
subject VARCHAR(255),
content TEXT NOT NULL,
description VARCHAR(255),
is_system BOOLEAN NOT NULL DEFAULT FALSE,
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_templates_is_system ON w_templates (is_system);
CREATE INDEX IF NOT EXISTS idx_w_templates_created_at ON w_templates (created_at);
CREATE INDEX IF NOT EXISTS idx_w_templates_updated_at ON w_templates (updated_at);
CREATE TABLE IF NOT EXISTS w_schedules (
id BIGINT PRIMARY KEY,
name VARCHAR(128) NOT NULL,
task_type VARCHAR(64) NOT NULL,
cron VARCHAR(64) NOT NULL,
payload TEXT,
is_active BOOLEAN NOT NULL DEFAULT TRUE,
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_schedules_is_active ON w_schedules (is_active);
-- Seed initial cleanup task
INSERT INTO w_schedules (id, name, task_type, cron, payload, is_active, created_at, updated_at)
VALUES (1, '系统定期垃圾清理', 'system_cleanup', '0 3 * * *', '{}', TRUE, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
ON CONFLICT (id) DO NOTHING;
CREATE TABLE IF NOT EXISTS w_task_executions (
id BIGINT PRIMARY KEY,
task_id VARCHAR(128) NOT NULL UNIQUE,
task_type VARCHAR(64) NOT NULL,
task_name VARCHAR(128),
status VARCHAR(32) NOT NULL,
retryable BOOLEAN NOT NULL DEFAULT FALSE,
max_retry INTEGER NOT NULL DEFAULT 0,
retry_count INTEGER NOT NULL DEFAULT 0,
log TEXT,
error_message TEXT,
result TEXT,
started_at TIMESTAMPTZ,
finished_at TIMESTAMPTZ,
duration BIGINT,
payload TEXT,
triggered_by VARCHAR(32) NOT NULL DEFAULT 'system',
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_task_executions_task_type ON w_task_executions (task_type);
CREATE INDEX IF NOT EXISTS idx_w_task_executions_status ON w_task_executions (status);
CREATE INDEX IF NOT EXISTS idx_w_task_executions_started_at ON w_task_executions (started_at);
CREATE INDEX IF NOT EXISTS idx_w_task_executions_created_at ON w_task_executions (created_at);
-- Seed system configs (all default platform configs)
INSERT INTO w_system_configs (key, value, type, visibility, description, created_at, updated_at) VALUES
('cap_login_enabled', 'false', 'system', 1, '是否启用登录人机验证(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_auto_solve', 'true', 'system', 1, '打开页面后是否自动开始计算,关闭则需用户手动点击触发', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_challenge_count', '1', 'system', 0, '客户端需求解的 PoW 难题总数,默认 1,推荐 1~5', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_challenge_size', '32', 'system', 0, '人机验证盐值长度', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_challenge_difficulty', '4', 'system', 0, '人机验证 PoW 难度(目标前缀长度)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_challenge_ttl_seconds', '600', 'system', 0, '人机验证难题有效时间(秒)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_token_ttl_seconds', '1200', 'system', 0, '人机验证兑换凭证有效时间(秒)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('server_address', '', 'system', 0, '服务器地址(用于跨域源控制,不设定则允许任意源)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('smtp_host', '', 'system', 0, 'SMTP 服务器地址(例如 smtp.example.com)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('smtp_port', '587', 'system', 0, 'SMTP 端口(例如 587 或 465)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('smtp_username', '', 'system', 0, 'SMTP 账户(如 sender@example.com)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('smtp_password', '', 'system', 0, 'SMTP 访问凭证(授权码/密码)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('upload_allowed_extensions', 'jpg,png,webp', 'system', 1, '允许上传的图片扩展名(逗号分隔)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('site_name', 'Wavelet', 'system', 1, '系统平台的展示名称', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('password_login_enabled', 'true', 'system', 1, '是否允许使用账号密码登录', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('registration_enabled', 'true', 'system', 1, '控制普通用户是否可以自主注册(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('password_register_enabled', 'true', 'system', 1, '是否允许通过密码创建本地账号', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('oidc_login_enabled', 'true', 'system', 1, '是否允许使用第三方 OIDC 认证源登录', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('max_api_keys_per_user', '5', 'business', 1, '限制每个普通用户可以创建的 API Key 最大数量', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('email_login_verification_enabled', 'false', 'system', 1, '是否开启邮箱登录验证(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('email_register_verification_enabled', 'false', 'system', 1, '是否开启邮箱注册验证(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('menu_display_config', '{}', 'system', 1, '目录显示配置(JSON 字符串,格式为 {url: enabled})', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('search_engine_indexing_enabled', 'false', 'system', 1, '是否允许搜索引擎爬取/检索该站点(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('update_upstream_repository', 'Rain-kl/Wavelet', 'system', 0, 'GitHub Actions Release 上游仓库(owner/repo 或 GitHub 仓库地址)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('storage_config', '{"driver":"local","local":{"root":"."},"s3":{"region":"us-east-1"},"r2":{"region":"auto"},"minio":{"region":"us-east-1","path_style":true},"oss":{},"webdav":{}}', 'system', 0, '文件存储驱动及连接配置(JSON)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('disk_cache_max_size_mb', '1024', 'system', 0, '磁盘缓存最大空间大小(MB)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('disk_cache_ttl_minutes', '1440', 'system', 0, '磁盘缓存默认有效期(分钟)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('disk_cache_lru_enabled', 'true', 'system', 0, '是否启用 LRU 淘汰机制', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('file_access_whitelist', '["avatar"]', 'system', 0, '免登录访问的文件业务类型白名单 (JSON 数组)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('login_session_ttl_hours', '168', 'system', 0, '登录会话过期时间(小时)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('log_database', '', 'system', 0, '当前日志主库(postgres/sqlite/clickhouse),由切换任务写入', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('log_db_migration', '', 'system', 0, '日志库迁移冻结标记(空或 migrating)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
ON CONFLICT (key) DO NOTHING;
INSERT INTO w_templates (id, key, name, type, subject, content, description, is_system, created_at, updated_at) VALUES
(1, 'login_email', '登录验证码邮件', 'email', 'Wavelet 登录验证码', '<h3>Wavelet 登录验证</h3><p>您的登录验证码为:<strong>{{.Code}}</strong>,5分钟内有效,请勿将验证码泄露给他人。</p>', '用户密码登录时发送的验证码邮件模板,支持变量:{{.Code}}', TRUE, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
(2, 'register_email', '注册验证码邮件', 'email', 'Wavelet 注册验证码', '<h3>Wavelet 注册验证</h3><p>您的注册验证码为:<strong>{{.Code}}</strong>,5分钟内有效,请勿泄露给他人。</p>', '用户注册时发送的验证码邮件模板,支持变量:{{.Code}}', TRUE, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
ON CONFLICT (key) DO NOTHING;
-- +goose StatementEnd
-- +goose Down
-- +goose StatementBegin
DELETE FROM w_templates WHERE key IN ('login_email', 'register_email');
DELETE FROM w_system_configs WHERE key IN (
'cap_login_enabled', 'cap_auto_solve', 'cap_challenge_count', 'cap_challenge_size',
'cap_challenge_difficulty', 'cap_challenge_ttl_seconds', 'cap_token_ttl_seconds',
'server_address', 'smtp_host', 'smtp_port', 'smtp_username', 'smtp_password',
'upload_allowed_extensions', 'site_name', 'password_login_enabled', 'registration_enabled',
'password_register_enabled', 'oidc_login_enabled', 'max_api_keys_per_user',
'email_login_verification_enabled', 'email_register_verification_enabled',
'menu_display_config', 'search_engine_indexing_enabled', 'update_upstream_repository',
'storage_config', 'disk_cache_max_size_mb', 'disk_cache_ttl_minutes', 'disk_cache_lru_enabled',
'file_access_whitelist', 'login_session_ttl_hours', 'log_database', 'log_db_migration'
);
DROP TABLE IF EXISTS w_task_executions;
DROP TABLE IF EXISTS w_schedules;
DROP TABLE IF EXISTS w_templates;
DROP TABLE IF EXISTS w_system_configs;
-- +goose StatementEnd
@@ -0,0 +1,131 @@
-- +goose Up
-- +goose StatementBegin
CREATE TABLE IF NOT EXISTS w_system_configs (
key VARCHAR(64) PRIMARY KEY,
value TEXT NOT NULL,
type VARCHAR(32) NOT NULL DEFAULT 'system',
visibility INTEGER NOT NULL DEFAULT 0,
description VARCHAR(255),
updated_at DATETIME DEFAULT CURRENT_TIMESTAMP,
created_at DATETIME DEFAULT CURRENT_TIMESTAMP
);
CREATE TABLE IF NOT EXISTS w_templates (
id BIGINT PRIMARY KEY,
key VARCHAR(80) NOT NULL UNIQUE,
name VARCHAR(100) NOT NULL,
type VARCHAR(20) NOT NULL DEFAULT 'email',
subject VARCHAR(255),
content TEXT NOT NULL,
description VARCHAR(255),
is_system BOOLEAN NOT NULL DEFAULT 0,
created_at DATETIME DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_templates_is_system ON w_templates (is_system);
CREATE INDEX IF NOT EXISTS idx_w_templates_created_at ON w_templates (created_at);
CREATE INDEX IF NOT EXISTS idx_w_templates_updated_at ON w_templates (updated_at);
CREATE TABLE IF NOT EXISTS w_schedules (
id BIGINT PRIMARY KEY,
name VARCHAR(128) NOT NULL,
task_type VARCHAR(64) NOT NULL,
cron VARCHAR(64) NOT NULL,
payload TEXT,
is_active BOOLEAN NOT NULL DEFAULT 1,
created_at DATETIME DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_schedules_is_active ON w_schedules (is_active);
-- Seed initial cleanup task
INSERT INTO w_schedules (id, name, task_type, cron, payload, is_active, created_at, updated_at)
VALUES (1, '系统定期垃圾清理', 'system_cleanup', '0 3 * * *', '{}', 1, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
ON CONFLICT (id) DO NOTHING;
CREATE TABLE IF NOT EXISTS w_task_executions (
id BIGINT PRIMARY KEY,
task_id VARCHAR(128) NOT NULL UNIQUE,
task_type VARCHAR(64) NOT NULL,
task_name VARCHAR(128),
status VARCHAR(32) NOT NULL,
retryable BOOLEAN NOT NULL DEFAULT 0,
max_retry INTEGER NOT NULL DEFAULT 0,
retry_count INTEGER NOT NULL DEFAULT 0,
log TEXT,
error_message TEXT,
result TEXT,
started_at DATETIME,
finished_at DATETIME,
duration BIGINT,
payload TEXT,
triggered_by VARCHAR(32) NOT NULL DEFAULT 'system',
created_at DATETIME DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_task_executions_task_type ON w_task_executions (task_type);
CREATE INDEX IF NOT EXISTS idx_w_task_executions_status ON w_task_executions (status);
CREATE INDEX IF NOT EXISTS idx_w_task_executions_started_at ON w_task_executions (started_at);
CREATE INDEX IF NOT EXISTS idx_w_task_executions_created_at ON w_task_executions (created_at);
-- Seed system configs (all default platform configs)
INSERT INTO w_system_configs (key, value, type, visibility, description, created_at, updated_at) VALUES
('cap_login_enabled', 'false', 'system', 1, '是否启用登录人机验证(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_auto_solve', 'true', 'system', 1, '打开页面后是否自动开始计算,关闭则需用户手动点击触发', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_challenge_count', '1', 'system', 0, '客户端需求解的 PoW 难题总数,默认 1,推荐 1~5', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_challenge_size', '32', 'system', 0, '人机验证盐值长度', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_challenge_difficulty', '4', 'system', 0, '人机验证 PoW 难度(目标前缀长度)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_challenge_ttl_seconds', '600', 'system', 0, '人机验证难题有效时间(秒)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('cap_token_ttl_seconds', '1200', 'system', 0, '人机验证兑换凭证有效时间(秒)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('server_address', '', 'system', 0, '服务器地址(用于跨域源控制,不设定则允许任意源)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('smtp_host', '', 'system', 0, 'SMTP 服务器地址(例如 smtp.example.com)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('smtp_port', '587', 'system', 0, 'SMTP 端口(例如 587 或 465)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('smtp_username', '', 'system', 0, 'SMTP 账户(如 sender@example.com)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('smtp_password', '', 'system', 0, 'SMTP 访问凭证(授权码/密码)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('upload_allowed_extensions', 'jpg,png,webp', 'system', 1, '允许上传的图片扩展名(逗号分隔)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('site_name', 'Wavelet', 'system', 1, '系统平台的展示名称', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('password_login_enabled', 'true', 'system', 1, '是否允许使用账号密码登录', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('registration_enabled', 'true', 'system', 1, '控制普通用户是否可以自主注册(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('password_register_enabled', 'true', 'system', 1, '是否允许通过密码创建本地账号', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('oidc_login_enabled', 'true', 'system', 1, '是否允许使用第三方 OIDC 认证源登录', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('max_api_keys_per_user', '5', 'business', 1, '限制每个普通用户可以创建的 API Key 最大数量', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('email_login_verification_enabled', 'false', 'system', 1, '是否开启邮箱登录验证(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('email_register_verification_enabled', 'false', 'system', 1, '是否开启邮箱注册验证(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('menu_display_config', '{}', 'system', 1, '目录显示配置(JSON 字符串,格式为 {url: enabled})', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('search_engine_indexing_enabled', 'false', 'system', 1, '是否允许搜索引擎爬取/检索该站点(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('update_upstream_repository', 'Rain-kl/Wavelet', 'system', 0, 'GitHub Actions Release 上游仓库(owner/repo 或 GitHub 仓库地址)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('storage_config', '{"driver":"local","local":{"root":"."},"s3":{"region":"us-east-1"},"r2":{"region":"auto"},"minio":{"region":"us-east-1","path_style":true},"oss":{},"webdav":{}}', 'system', 0, '文件存储驱动及连接配置(JSON)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('disk_cache_max_size_mb', '1024', 'system', 0, '磁盘缓存最大空间大小(MB)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('disk_cache_ttl_minutes', '1440', 'system', 0, '磁盘缓存默认有效期(分钟)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('disk_cache_lru_enabled', 'true', 'system', 0, '是否启用 LRU 淘汰机制', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('file_access_whitelist', '["avatar"]', 'system', 0, '免登录访问的文件业务类型白名单 (JSON 数组)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('login_session_ttl_hours', '168', 'system', 0, '登录会话过期时间(小时)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('log_database', '', 'system', 0, '当前日志主库(postgres/sqlite/clickhouse),由切换任务写入', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('log_db_migration', '', 'system', 0, '日志库迁移冻结标记(空或 migrating)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
ON CONFLICT (key) DO NOTHING;
INSERT INTO w_templates (id, key, name, type, subject, content, description, is_system, created_at, updated_at) VALUES
(1, 'login_email', '登录验证码邮件', 'email', 'Wavelet 登录验证码', '<h3>Wavelet 登录验证</h3><p>您的登录验证码为:<strong>{{.Code}}</strong>,5分钟内有效,请勿将验证码泄露给他人。</p>', '用户密码登录时发送的验证码邮件模板,支持变量:{{.Code}}', 1, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
(2, 'register_email', '注册验证码邮件', 'email', 'Wavelet 注册验证码', '<h3>Wavelet 注册验证</h3><p>您的注册验证码为:<strong>{{.Code}}</strong>,5分钟内有效,请勿泄露给他人。</p>', '用户注册时发送的验证码邮件模板,支持变量:{{.Code}}', 1, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
ON CONFLICT (key) DO NOTHING;
-- +goose StatementEnd
-- +goose Down
-- +goose StatementBegin
DELETE FROM w_templates WHERE key IN ('login_email', 'register_email');
DELETE FROM w_system_configs WHERE key IN (
'cap_login_enabled', 'cap_auto_solve', 'cap_challenge_count', 'cap_challenge_size',
'cap_challenge_difficulty', 'cap_challenge_ttl_seconds', 'cap_token_ttl_seconds',
'server_address', 'smtp_host', 'smtp_port', 'smtp_username', 'smtp_password',
'upload_allowed_extensions', 'site_name', 'password_login_enabled', 'registration_enabled',
'password_register_enabled', 'oidc_login_enabled', 'max_api_keys_per_user',
'email_login_verification_enabled', 'email_register_verification_enabled',
'menu_display_config', 'search_engine_indexing_enabled', 'update_upstream_repository',
'storage_config', 'disk_cache_max_size_mb', 'disk_cache_ttl_minutes', 'disk_cache_lru_enabled',
'file_access_whitelist', 'login_session_ttl_hours', 'log_database', 'log_db_migration'
);
DROP TABLE IF EXISTS w_task_executions;
DROP TABLE IF EXISTS w_schedules;
DROP TABLE IF EXISTS w_templates;
DROP TABLE IF EXISTS w_system_configs;
-- +goose StatementEnd
@@ -0,0 +1,20 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
// DatabaseConfig holds database configuration needed by the admin plugin.
type DatabaseConfig struct {
Enabled bool `config:"enabled" env:"DB_ENABLED" default:"false" autoEnable:"DB_HOST"`
Host string `config:"host" env:"DB_HOST"`
Port int `config:"port" env:"DB_PORT" default:"5432"`
Database string `config:"database" env:"DB_DATABASE"`
Username string `config:"username" env:"DB_USERNAME"`
Password string `config:"password" env:"DB_PASSWORD" secret:"true"`
SQLitePath string `config:"sqlite_path" env:"DB_SQLITE_PATH" default:"./data/wavelet.db"`
}
// ClickHouseConfig holds clickhouse enablement status needed by admin log queries/switching.
type ClickHouseConfig struct {
Enabled bool `config:"enabled" env:"CLICKHOUSE_ENABLED" default:"false" autoEnable:"CLICKHOUSE_HOST"`
}
+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"`
}
@@ -0,0 +1,214 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package model contains database entities and data transfer objects for the admin domain.
package model
import (
"Wavelet/plugins/domain/admin/errs"
"bytes"
"errors"
"strings"
"text/template"
"time"
)
// 配置键常量 - 所有系统配置的 key 定义
const (
ConfigKeyUploadAllowedExtensions = "upload_allowed_extensions" // 允许上传的文件扩展名,逗号分隔
ConfigKeySiteName = "site_name" // 站点名称
ConfigKeyPasswordLoginEnabled = "password_login_enabled" // 是否允许密码登录
ConfigKeyRegistrationEnabled = "registration_enabled" // 是否允许注册
ConfigKeyPasswordRegisterEnabled = "password_register_enabled" // 是否允许密码注册
ConfigKeyOIDCLoginEnabled = "oidc_login_enabled" // 是否允许 OIDC 登录
ConfigKeyMaxAPIKeysPerUser = "max_api_keys_per_user" //nolint:gosec // false positive: config key name, not credentials
ConfigKeyCapLoginEnabled = "cap_login_enabled" // 是否启用登录人机验证
ConfigKeyCapAutoSolve = "cap_auto_solve" // 打开页面后是否自动开始计算(false 则需用户手动点击)
ConfigKeyCapChallengeCount = "cap_challenge_count" // 客户端需求解的 PoW 难题总数,默认 1,推荐 1~5
ConfigKeyCapChallengeSize = "cap_challenge_size" // 人机验证盐值长度
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty" // 人机验证 PoW 难度(目标前缀长度)
ConfigKeyCapChallengeTTL = "cap_challenge_ttl_seconds" // 人机验证难题有效时间(秒)
ConfigKeyCapTokenTTL = "cap_token_ttl_seconds" //nolint:gosec // false positive: config key name, not credentials
ConfigKeyServerAddress = "server_address" // 服务器地址
ConfigKeySMTPHost = "smtp_host" // SMTP 服务器地址
ConfigKeySMTPPort = "smtp_port" // SMTP 端口
ConfigKeySMTPUsername = "smtp_username" // SMTP 账户
ConfigKeySMTPPassword = "smtp_password" // SMTP 访问凭证
ConfigKeyEmailLoginVerificationEnabled = "email_login_verification_enabled" // 是否启用邮箱登录验证
ConfigKeyEmailRegisterVerificationEnabled = "email_register_verification_enabled" // 是否启用邮箱注册验证
ConfigKeyMenuDisplayConfig = "menu_display_config" // 目录显示配置 (JSON 字符串)
ConfigKeySearchEngineIndexingEnabled = "search_engine_indexing_enabled" // 是否允许搜索引擎检索
ConfigKeyFileAccessWhitelist = "file_access_whitelist" // 免登录访问的文件业务类型白名单 (JSON 数组格式)
ConfigKeyDiskCacheMaxSizeMB = "disk_cache_max_size_mb" // 磁盘缓存最大空间大小 (MB)
ConfigKeyDiskCacheTTLMinutes = "disk_cache_ttl_minutes" // 磁盘缓存默认有效期 (分钟)
ConfigKeyDiskCacheLRUEnabled = "disk_cache_lru_enabled" // 是否启用 LRU 淘汰机制
ConfigKeyLoginSessionTTLHours = "login_session_ttl_hours" // 登录会话过期时间 (小时)
ConfigKeyUpdateUpstreamRepository = "update_upstream_repository" // GitHub Actions Release 上游仓库
ConfigKeyStorageConfig = "storage_config" // 文件存储配置 (JSON)
ConfigKeyLogDatabase = "log_database" // 当前日志主库(postgres/sqlite/clickhouse),受保护
ConfigKeyLogDBMigration = "log_db_migration" // 日志库迁移冻结标记(空/migrating),受保护
ConfigKeyLogRetentionDaysPostgres = "log_retention_days_postgres" // PostgreSQL 用户访问日志保留天数
ConfigKeyLogRetentionDaysSQLite = "log_retention_days_sqlite" // SQLite 用户访问日志保留天数
ConfigKeyLogRetentionDaysClickHouse = "log_retention_days_clickhouse" // ClickHouse 用户访问日志保留天数
)
const (
// ConfigVisibilityHidden 表示配置不通过公共配置接口暴露
ConfigVisibilityHidden = 0
// ConfigVisibilityVisible 表示配置通过公共配置接口暴露
ConfigVisibilityVisible = 1
)
// SystemConfig 系统配置实体
type SystemConfig struct {
Key string `json:"key" gorm:"primaryKey;size:64;not null"`
Value string `json:"value" gorm:"type:text;not null"`
Type string `json:"type" gorm:"size:32;not null;default:'system'"`
Visibility int `json:"visibility" gorm:"not null;default:0"`
Description string `json:"description" gorm:"size:255"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName 表名
func (SystemConfig) TableName() string {
return "w_system_configs"
}
// Template 邮件/消息模板实体
type Template struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
Key string `json:"key" gorm:"uniqueIndex;size:80;not null"`
Name string `json:"name" gorm:"size:100;not null"`
Type string `json:"type" gorm:"size:20;not null;default:'email'"`
Subject string `json:"subject" gorm:"size:255"`
Content string `json:"content" gorm:"type:text;not null"`
Description string `json:"description" gorm:"size:255"`
IsSystem bool `json:"is_system" gorm:"index;not null;default:false"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
}
// TableName 表名
func (Template) TableName() string {
return "w_templates"
}
// TemplateTypeEmail 邮件模板类型
const TemplateTypeEmail = "email"
// Normalize 规范化模板字段
func (t *Template) Normalize() {
t.Key = strings.TrimSpace(t.Key)
t.Name = strings.TrimSpace(t.Name)
t.Type = strings.ToLower(strings.TrimSpace(t.Type))
t.Subject = strings.TrimSpace(t.Subject)
t.Content = strings.TrimSpace(t.Content)
t.Description = strings.TrimSpace(t.Description)
if t.Type == "" {
t.Type = TemplateTypeEmail
}
}
// Validate 校验模板必填字段
func (t *Template) Validate() error {
t.Normalize()
if t.Key == "" {
return errors.New(errs.TemplateKeyRequired)
}
if t.Name == "" {
return errors.New(errs.TemplateNameRequired)
}
if t.Content == "" {
return errors.New(errs.TemplateContentRequired)
}
return nil
}
// Render 渲染模板的 Subject 和 Content
func (t *Template) Render(data any) (string, string, error) {
var subject string
if t.Subject != "" {
tmplSubject, err := template.New(t.Key + "_subject").Parse(t.Subject)
if err != nil {
return "", "", err
}
var subBuf bytes.Buffer
if err := tmplSubject.Execute(&subBuf, data); err != nil {
return "", "", err
}
subject = subBuf.String()
}
tmplContent, err := template.New(t.Key + "_content").Parse(t.Content)
if err != nil {
return "", "", err
}
var bodyBuf bytes.Buffer
if err := tmplContent.Execute(&bodyBuf, data); err != nil {
return "", "", err
}
return subject, bodyBuf.String(), nil
}
// Schedule 定时任务配置表
type Schedule struct {
ID uint64 `json:"id,string" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:128;not null"`
TaskType string `json:"task_type" gorm:"size:64;not null"`
Cron string `json:"cron" gorm:"size:64;not null"`
Payload string `json:"payload" gorm:"type:text"`
IsActive bool `json:"is_active" gorm:"not null;default:true"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (Schedule) TableName() string {
return "w_schedules"
}
// TaskExecutionStatus 任务执行状态
type TaskExecutionStatus string
// 任务执行状态
const (
TaskExecutionStatusPending TaskExecutionStatus = "pending"
TaskExecutionStatusRunning TaskExecutionStatus = "running"
TaskExecutionStatusSucceeded TaskExecutionStatus = "succeeded"
TaskExecutionStatusFailed TaskExecutionStatus = "failed"
)
// TaskExecution 任务执行记录
type TaskExecution struct {
ID uint64 `json:"id,string" gorm:"primaryKey"`
TaskID string `json:"task_id" gorm:"size:128;uniqueIndex;not null"`
TaskType string `json:"task_type" gorm:"size:64;index;not null"`
TaskName string `json:"task_name" gorm:"size:128"`
Status TaskExecutionStatus `json:"status" gorm:"size:32;index;not null"`
Retryable bool `json:"retryable" gorm:"not null;default:false"`
MaxRetry int `json:"max_retry" gorm:"not null;default:0"`
RetryCount int `json:"retry_count" gorm:"not null;default:0"`
Log string `json:"log" gorm:"type:text"`
ErrorMessage string `json:"error_message" gorm:"type:text"`
Result string `json:"result" gorm:"type:text"`
StartedAt *time.Time `json:"started_at" gorm:"index"`
FinishedAt *time.Time `json:"finished_at"`
Duration int64 `json:"duration" gorm:"comment:耗时毫秒"`
Payload string `json:"payload" gorm:"type:text"`
TriggeredBy string `json:"triggered_by" gorm:"size:32;not null;default:system"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (TaskExecution) TableName() string {
return "w_task_executions"
}
// TaskExecutionCleanupStats 任务日志清理结果统计
type TaskExecutionCleanupStats struct {
HighFrequencyDeleted int64 `json:"high_frequency_deleted"`
LowFrequencyDeleted int64 `json:"low_frequency_deleted"`
}
+205
View File
@@ -0,0 +1,205 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package admin provides the system management console, diagnostics, audit logging, and configuration hot-reloading domain plugin for Cordis.
package admin
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/pkg/ginutil"
"Wavelet/plugins/domain/admin/handler"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/service"
"context"
"embed"
"reflect"
"github.com/gin-gonic/gin"
)
// SystemConfig aliases model.SystemConfig for external compatibility.
type SystemConfig = model.SystemConfig
//go:embed migrations/*/*.sql
var adminMigrations embed.FS
// Option configures the admin plugin.
type Option func(*Plugin)
// Plugin implements core.Plugin to provide system administration and management APIs.
type Plugin struct{}
// New creates a new admin domain plugin.
func New(opts ...Option) *Plugin {
p := &Plugin{}
for _, opt := range opts {
if opt != nil {
opt(p)
}
}
return p
}
// Name returns the unique identifier for the admin domain plugin.
func (p *Plugin) Name() string {
return "admin"
}
// Inject declares required dependencies for the admin domain plugin.
func (p *Plugin) Inject() []reflect.Type {
return []reflect.Type{
reflect.TypeFor[contracts.DBService](),
reflect.TypeFor[contracts.CacheService](),
reflect.TypeFor[contracts.UserService](),
reflect.TypeFor[contracts.AuthService](),
}
}
// Manifest returns the plugin metadata.
func (p *Plugin) Manifest() core.Manifest {
return core.Manifest{
Name: "admin",
Version: "1.0.0",
Description: "System administration console, diagnostic monitoring, and configuration hot-reload plugin",
Author: "Wavelet Team",
}
}
// DeclareConfig declares configuration bindings consumed by the admin plugin.
func (p *Plugin) DeclareConfig() []core.ConfigBinding {
return []core.ConfigBinding{
{Prefix: "database", Target: &model.DatabaseConfig{}},
{Prefix: "clickhouse", Target: &model.ClickHouseConfig{}},
}
}
// Apply registers admin routes, tasks, schedules, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
var dbCfg model.DatabaseConfig
_ = ctx.Config().Bind("database", &dbCfg)
service.SetDBConfig(dbCfg)
var chCfg model.ClickHouseConfig
_ = ctx.Config().Bind("clickhouse", &chCfg)
service.SetClickHouseConfig(chCfg)
// 0. Bind Services reactively
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
service.SetDBService(db)
} else {
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
service.SetDBService(db)
})
}
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
service.SetCacheService(cache)
} else {
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
service.SetCacheService(cache)
})
}
if user, err := core.Inject[contracts.UserService](ctx); err == nil && user != nil {
service.SetUserService(user)
} else {
core.When[contracts.UserService](ctx, func(user contracts.UserService) {
service.SetUserService(user)
})
}
if auth, err := core.Inject[contracts.AuthService](ctx); err == nil && auth != nil {
service.SetAuthService(auth)
} else {
core.When[contracts.AuthService](ctx, func(auth contracts.AuthService) {
service.SetAuthService(auth)
})
}
if task, err := core.Inject[contracts.TaskService](ctx); err == nil && task != nil {
service.SetTaskService(task)
} else {
core.When[contracts.TaskService](ctx, func(task contracts.TaskService) {
service.SetTaskService(task)
})
}
if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil {
service.SetStorageService(storage)
} else {
core.When[contracts.StorageService](ctx, func(storage contracts.StorageService) {
service.SetStorageService(storage)
})
}
if rc, err := core.Inject[contracts.RiskControlService](ctx); err == nil && rc != nil {
service.SetRiskControlService(rc)
} else {
core.When[contracts.RiskControlService](ctx, func(rc contracts.RiskControlService) {
service.SetRiskControlService(rc)
})
}
service.SetEventEmitter(ctx.Events().Emit)
ctx.OnDispose(func() error {
service.ResetServices()
return nil
})
// 0a. Dynamic Auth Middlewares
denyAuth := ginutil.AuthUnavailable()
var loginMW gin.HandlerFunc = func(c *gin.Context) {
if authSvc := service.GetAuthService(c.Request.Context()); authSvc != nil {
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
mw(c)
return
}
}
denyAuth(c)
}
var adminMW gin.HandlerFunc = func(c *gin.Context) {
if authSvc := service.GetAuthService(c.Request.Context()); authSvc != nil {
if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok {
mw(c)
return
}
}
denyAuth(c)
}
// 0b. Register migrations
ctx.Migrations().Register("admin", adminMigrations)
// 1. Register Admin HTTP Routes
adminRouter := ctx.Router().Group("/api/v1/admin", loginMW, adminMW)
handler.RegisterRoutes(adminRouter)
// 2. Register Background Tasks
logSwitchHandler := &service.LogDBSwitchHandler{}
ctx.Task().Register(service.LogDBSwitchTask, func(c context.Context, payload []byte) error {
_, err := logSwitchHandler.Execute(c, payload)
return err
}, extpoints.WithTaskMeta(service.LogDBSwitchMeta))
ctx.Task().Register("admin:system_cleanup", func(_ context.Context, _ []byte) error {
return nil
},
extpoints.WithTaskType("system_cleanup"),
extpoints.WithTaskName("系统垃圾清理"),
extpoints.WithTaskDescription("定期清理未使用上传文件、历史推送记录和过期任务执行日志"),
extpoints.WithTaskCategory("maintenance"),
extpoints.WithTaskRetry(1),
extpoints.WithTaskQueue("default"),
extpoints.WithTaskRetryable(true),
)
// 3. Register Cron Schedules
ctx.Schedule().RegisterCron("0 4 * * *", "admin:system_cleanup", map[string]string{"type": "daily"})
// 4. Register Settings Schemas
ctx.Settings().Register(extpoints.SettingSchema{
Key: "admin.system_cleanup_cron",
Default: "0 4 * * *",
Description: "Cron expression for nightly system logs and expired tokens cleanup",
Type: "string",
Category: "maintenance",
})
return nil
}
@@ -0,0 +1,84 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin_test
import (
"Wavelet/core"
"Wavelet/plugins/domain/admin"
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestAdminPluginUnit(t *testing.T) {
ctx := core.NewContext(context.Background())
p := admin.New()
assert.Equal(t, "admin", p.Name())
assert.Equal(t, "1.0.0", p.Manifest().Version)
require.NoError(t, p.Apply(ctx))
// Verify routes
routes := ctx.Router().Routes()
assert.NotEmpty(t, routes)
// Verify tasks
_, ok := ctx.Tasks().Get("admin:system_cleanup")
require.True(t, ok)
// Verify schedules
sched, ok := ctx.Schedules().Get("admin:system_cleanup")
require.True(t, ok)
assert.Equal(t, "0 4 * * *", sched.Spec)
// Verify settings
setting, ok := ctx.Settings().Get("admin.system_cleanup_cron")
require.True(t, ok)
assert.Equal(t, "0 4 * * *", setting.Default)
}
func TestAdminMigrationsIncludeTaskExecutionsAndSchedules(t *testing.T) {
ctx := core.NewContext(context.Background())
p := admin.New()
require.NoError(t, p.Apply(ctx))
entry, ok := ctx.Migrations().Get("admin")
require.True(t, ok, "admin plugin must register migrations")
assert.Equal(t, "admin", entry.PluginID)
// Verify sqlite migration files include w_schedules and w_task_executions
sqliteDir, err := entry.FS.Open("migrations/sqlite/00001_initial.sql")
require.NoError(t, err)
defer sqliteDir.Close()
stat, err := sqliteDir.Stat()
require.NoError(t, err)
buf := make([]byte, stat.Size())
_, err = sqliteDir.Read(buf)
require.NoError(t, err)
content := string(buf)
assert.Contains(t, content, "CREATE TABLE IF NOT EXISTS w_task_executions")
assert.Contains(t, content, "CREATE TABLE IF NOT EXISTS w_schedules")
assert.Contains(t, content, "CREATE TABLE IF NOT EXISTS w_system_configs")
assert.Contains(t, content, "CREATE TABLE IF NOT EXISTS w_templates")
// Verify postgres migration files include w_schedules and w_task_executions
pgDir, err := entry.FS.Open("migrations/postgres/00001_initial.sql")
require.NoError(t, err)
defer pgDir.Close()
stat, err = pgDir.Stat()
require.NoError(t, err)
buf = make([]byte, stat.Size())
_, err = pgDir.Read(buf)
require.NoError(t, err)
pgContent := string(buf)
assert.Contains(t, pgContent, "CREATE TABLE IF NOT EXISTS w_task_executions")
assert.Contains(t, pgContent, "CREATE TABLE IF NOT EXISTS w_schedules")
assert.Contains(t, pgContent, "CREATE TABLE IF NOT EXISTS w_system_configs")
assert.Contains(t, pgContent, "CREATE TABLE IF NOT EXISTS w_templates")
}
@@ -0,0 +1,144 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"Wavelet/pkg/cache/ram"
"Wavelet/plugins/domain/admin/model"
"context"
"encoding/json"
"errors"
"time"
"gorm.io/gorm"
)
const (
// SystemConfigBroadcastChannel broadcasts system config cache updates across nodes.
SystemConfigBroadcastChannel = "system:config_broadcast"
// SystemConfigInvalidationChannel is kept as an alias for backward compatibility.
SystemConfigInvalidationChannel = SystemConfigBroadcastChannel
// SystemConfigRedisHashKey is kept for backward compatibility in tests.
SystemConfigRedisHashKey = "system:system_configs"
// SystemConfigVisibleListRedisKey is kept for backward compatibility in tests.
SystemConfigVisibleListRedisKey = "system:visible_configs"
// ConfigCacheType is the cache type for all system configs.
ConfigCacheType = "config"
)
// ConfigLoader loads configuration data from the database.
type ConfigLoader struct{}
// LoadAll loads all system configs from database as CacheItems.
func (ConfigLoader) LoadAll(ctx context.Context, configType string) ([]ram.CacheItem, error) {
configs, err := PreheatSystemConfigs(ctx)
if err != nil {
return nil, err
}
items := make([]ram.CacheItem, len(configs))
for i, cfg := range configs {
valBytes, err := json.Marshal(cfg)
if err != nil {
return nil, err
}
items[i] = ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(cfg.Key),
}
}
return items, nil
}
// LoadOne loads a single system config from database as CacheItem.
func (ConfigLoader) LoadOne(ctx context.Context, configType, key string) (ram.CacheItem, error) {
cfg, err := GetSystemConfigByKey(ctx, key)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return ram.CacheItem{}, ram.ErrNotFound
}
return ram.CacheItem{}, err
}
valBytes, err := json.Marshal(cfg)
if err != nil {
return ram.CacheItem{}, err
}
return ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(cfg.Key),
}, nil
}
// GetCachedSystemConfig retrieves a single system config with RAM L1 fallback to DB.
func GetCachedSystemConfig(ctx context.Context, key string) (*model.SystemConfig, error) {
if item, ok := ram.Get(ConfigCacheType, key); ok {
var cfg model.SystemConfig
if err := json.Unmarshal([]byte(item.Value), &cfg); err == nil {
return &cfg, nil
}
}
cfg, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return nil, err
}
valBytes, err := json.Marshal(cfg)
if err == nil {
ram.Set(ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(key),
})
}
return &cfg, nil
}
// StopSystemConfigCacheListener stops the cache invalidation listener (kept for backward compatibility).
func StopSystemConfigCacheListener() {
}
// StartSystemConfigCacheListener starts the cache listener (kept for backward compatibility).
func StartSystemConfigCacheListener() {
}
func ensureSystemConfigCacheListener() {
}
func determineTTL(_ string) time.Duration {
return -1
}
// InvalidateSystemConfigCache triggers a broadcast to refresh the cache for key.
func InvalidateSystemConfigCache(ctx context.Context, key string) error {
ram.Delete(ConfigCacheType, key)
if cacheSvc := GetCache(ctx); cacheSvc != nil {
_ = cacheSvc.Delete(ctx, "system:config:"+key)
_ = cacheSvc.Delete(ctx, SystemConfigVisibleListRedisKey)
}
return nil
}
// InvalidateAllSystemConfigCaches triggers a broadcast to refresh the entire config cache.
func InvalidateAllSystemConfigCaches(ctx context.Context) error {
ram.UpdateTypeItems(ConfigCacheType, nil)
if cacheSvc := GetCache(ctx); cacheSvc != nil {
_ = cacheSvc.Delete(ctx, SystemConfigRedisHashKey)
_ = cacheSvc.Delete(ctx, SystemConfigVisibleListRedisKey)
}
return nil
}
// ResetSystemConfigRAMCacheForTest clears only the process-local RAM cache.
func ResetSystemConfigRAMCacheForTest() {
ram.ResetForTest()
}
@@ -0,0 +1,394 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"context"
"database/sql"
"errors"
"fmt"
"os"
"os/exec"
"strings"
"sync"
"time"
)
const (
defaultSQLiteDBPath = "./data/wavelet.db"
logDBNameSQLite = "sqlite"
)
var (
dbConfigMu sync.RWMutex
dbConfig = model.DatabaseConfig{
SQLitePath: defaultSQLiteDBPath,
}
)
// SetDBConfig sets the database configuration.
func SetDBConfig(cfg model.DatabaseConfig) {
dbConfigMu.Lock()
defer dbConfigMu.Unlock()
dbConfig = cfg
}
// GetDBConfig gets the database configuration.
func GetDBConfig() model.DatabaseConfig {
dbConfigMu.RLock()
defer dbConfigMu.RUnlock()
return dbConfig
}
// sqliteDatabasePath resolves the effective SQLite file path from configuration.
func sqliteDatabasePath() string {
name := GetDBConfig().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 := GetDBConfig().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 !GetDBConfig().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 {
cfg := GetDBConfig()
info := model.DatabaseInfoResponse{
Type: logDBNameSQLite,
Name: cfg.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 {
cfg := GetDBConfig()
info := model.DatabaseInfoResponse{
Type: "postgres",
Name: cfg.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) {
// 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 := GetDBConfig()
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
}
@@ -0,0 +1,157 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository_test
import (
"context"
"errors"
"fmt"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
cacheplugin "Wavelet/plugins/infra/cache"
)
// stubDBService 用内存 SQLite 满足 DBService 契约,隔离外部依赖。
type stubDBService struct{ db *gorm.DB }
func (s stubDBService) GORM() *gorm.DB { return s.db }
func (s stubDBService) DB(context.Context) *gorm.DB { return s.db }
func (s stubDBService) Named(string) *gorm.DB { return s.db }
// newFlushLogTestCache 构建真实多层缓存服务并注入 admin 插件上下文。
func newFlushLogTestCache(t *testing.T) (contracts.CacheService, *miniredis.Miniredis, func()) {
t.Helper()
mr, err := miniredis.Run()
require.NoError(t, err)
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr(), MaintNotificationsConfig: &maintnotifications.Config{Mode: maintnotifications.ModeDisabled}})
p := cacheplugin.New(cacheplugin.WithRedis(rdb), cacheplugin.WithRAMCapacity(64))
ctx := core.NewContext(context.Background())
ctx.Config().SetSource(core.NewMapSource(map[string]any{
"redis": map[string]any{
"enabled": true,
"addrs": []string{mr.Addr()},
},
}))
require.NoError(t, ctx.Config().Resolve())
require.NoError(t, p.Apply(ctx))
svc, err := core.Inject[contracts.CacheService](ctx)
require.NoError(t, err)
repository.SetCacheService(svc)
cleanup := func() {
repository.SetCacheService(nil)
_ = rdb.Close()
mr.Close()
}
return svc, mr, cleanup
}
// TestFlushTaskExecutionLogPropagatesCacheError 回归:缓存读取失败(非未命中)时,
// FlushTaskExecutionLog 必须返回错误而不是静默吞掉日志并误报成功(nilerr 修复)。
func TestFlushTaskExecutionLogPropagatesCacheError(t *testing.T) {
_, mr, cleanup := newFlushLogTestCache(t)
defer cleanup()
ctx := context.Background()
const taskID = "flush-err-task"
// 先缓冲一行日志
require.NoError(t, repository.AppendTaskExecutionLog(ctx, taskID, "step-1 ok"))
// 关闭 miniredis 模拟缓存基础设施故障(读取出错而非未命中)
mr.Close()
err := repository.FlushTaskExecutionLog(ctx, taskID)
assert.Error(t, err, "缓存故障时必须返回错误,防止缓冲日志被静默丢弃")
}
// TestFlushTaskExecutionLogCacheMissIsNoop 回归:任务无缓冲日志(未命中)时应为空操作成功。
func TestFlushTaskExecutionLogCacheMissIsNoop(t *testing.T) {
_, _, cleanup := newFlushLogTestCache(t)
defer cleanup()
ctx := context.Background()
assert.NoError(t, repository.FlushTaskExecutionLog(ctx, "missing-task"))
}
// TestFlushTaskExecutionLogPersistsAndClears 验证正常路径:缓冲日志写入执行记录后清理缓存。
func TestFlushTaskExecutionLogPersistsAndClears(t *testing.T) {
svc, _, cleanup := newFlushLogTestCache(t)
defer cleanup()
ctx := context.Background()
const taskID = "flush-ok-task"
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(&model.TaskExecution{}))
repository.SetDBService(stubDBService{db: sqliteDB})
defer repository.SetDBService(nil)
gormDB := sqliteDB
exec := &model.TaskExecution{TaskID: taskID, TaskType: "upload:test", TaskName: "t", Status: model.TaskExecutionStatusSucceeded}
require.NoError(t, gormDB.Create(exec).Error)
require.NoError(t, repository.FlushTaskExecutionLog(ctx, taskID))
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, repository.TaskExecutionLogRedisKey(taskID), &buf)
assert.True(t, errors.Is(err, contracts.ErrCacheMiss), "flush 后缓存应清空, got %v", err)
}
// readFailCache 读取永远报错而写入成功,用于区分「未命中」与「缓存故障」两种语义。
type readFailCache struct {
writes []string
}
func (c *readFailCache) Get(context.Context, string, any) error {
return errors.New("cache unavailable")
}
func (c *readFailCache) Set(_ context.Context, key string, value any, _ time.Duration) error {
c.writes = append(c.writes, fmt.Sprintf("%s=%v", key, value))
return nil
}
func (c *readFailCache) Delete(context.Context, string) error { return nil }
func (c *readFailCache) GetOrSet(context.Context, string, any, time.Duration, func() (any, error)) error {
return errors.New("cache unavailable")
}
func (c *readFailCache) Invalidate(context.Context, string) error { return nil }
// TestAppendTaskExecutionLogKeepsBufferOnCacheReadError 回归:缓存读取失败(而非未命中)时
// 不得把「读不到」当成「没有缓冲」继续写入,否则整段任务日志会被最新一行覆盖丢失。
func TestAppendTaskExecutionLogKeepsBufferOnCacheReadError(t *testing.T) {
fake := &readFailCache{}
repository.SetCacheService(fake)
defer repository.SetCacheService(nil)
err := repository.AppendTaskExecutionLog(context.Background(), "append-err-task", "step-2")
assert.Error(t, err, "缓存故障必须上抛,而不是覆盖缓冲")
assert.Empty(t, fake.writes, "读取失败时不得写入,避免覆盖已缓冲日志")
}
@@ -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,330 @@
// 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
}
loadTaskExecutionLog(ctx, &execution)
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
}
loadTaskExecutionLog(ctx, &execution)
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 {
loadTaskExecutionLog(ctx, &execution)
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
if err := cacheSvc.Get(ctx, key, &existing); err != nil {
// 只有未命中才代表「尚无缓冲」;其余读取失败若被当作空缓冲继续写入,
// 会用这一行覆盖掉整段已缓冲的任务日志。
if !errors.Is(err, contracts.ErrCacheMiss) {
return fmt.Errorf("load buffered task execution log: %w", err)
}
}
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
}
loadTaskExecutionLogs(ctx, executions)
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
}
// loadTaskExecutionLog best-effort enriches an execution with its cached log;
// a cache miss or failure simply leaves the stored log column in place.
func loadTaskExecutionLog(ctx context.Context, execution *model.TaskExecution) {
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return
}
var logText string
if err := cacheSvc.Get(ctx, TaskExecutionLogRedisKey(execution.TaskID), &logText); err == nil && logText != "" {
execution.Log = logText
}
}
// loadTaskExecutionLogs best-effort enriches every execution with its cached log.
func loadTaskExecutionLogs(ctx context.Context, executions []model.TaskExecution) {
cacheSvc := GetCache(ctx)
if cacheSvc == nil || len(executions) == 0 {
return
}
for i := range executions {
var logText string
if err := cacheSvc.Get(ctx, TaskExecutionLogRedisKey(executions[i].TaskID), &logText); err == nil && logText != "" {
executions[i].Log = logText
}
}
}
@@ -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
}
@@ -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
}
+166
View File
@@ -0,0 +1,166 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
"context"
"os"
"os/exec"
"strings"
"sync"
"time"
)
var (
dbConfigMu sync.RWMutex
dbConfig model.DatabaseConfig
chConfig model.ClickHouseConfig
)
// SetDBConfig sets the database configuration in service and repository.
func SetDBConfig(cfg model.DatabaseConfig) {
dbConfigMu.Lock()
defer dbConfigMu.Unlock()
dbConfig = cfg
repository.SetDBConfig(cfg)
}
// GetDBConfig returns the database configuration.
func GetDBConfig() model.DatabaseConfig {
dbConfigMu.RLock()
defer dbConfigMu.RUnlock()
return dbConfig
}
// SetClickHouseConfig sets the clickhouse configuration.
func SetClickHouseConfig(cfg model.ClickHouseConfig) {
dbConfigMu.Lock()
defer dbConfigMu.Unlock()
chConfig = cfg
}
// GetClickHouseConfig returns the clickhouse configuration.
func GetClickHouseConfig() model.ClickHouseConfig {
dbConfigMu.RLock()
defer dbConfigMu.RUnlock()
return chConfig
}
// 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 !GetDBConfig().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 !GetDBConfig().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)
}
+240
View File
@@ -0,0 +1,240 @@
// 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 {
if users, err := userSvc.GetUsersByIDs(ctx, userIDs); err == nil {
for _, u := range users {
if u != nil {
userMap[u.ID] = 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,185 @@
// 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"
"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{
Type: TaskTypeLogDBSwitch,
AsynqTask: LogDBSwitchTask,
Name: "切换日志数据库",
DisplayName: "切换日志数据库",
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
Category: "system",
SupportsTime: false,
MaxRetry: 3,
Queue: "default",
Retryable: true,
Params: []contracts.TaskParamDTO{
{
Name: "target",
Label: "目标日志库",
Type: "string",
Required: true,
Placeholder: "postgres|sqlite|clickhouse",
Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse",
},
},
}
type logDBSwitchPayload struct {
Target string `json:"target"`
}
// LogDBSwitchHandler 切换日志数据库任务处理器。
type LogDBSwitchHandler struct{}
// ValidatePayload 校验并规范化参数。
func (h *LogDBSwitchHandler) ValidatePayload(payload []byte) ([]byte, error) {
var p logDBSwitchPayload
if err := json.Unmarshal(payload, &p); err != nil {
return nil, fmt.Errorf(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 !GetClickHouseConfig().Enabled {
return errors.New(errs.ErrClickHouseNotEnabled)
}
case targetPostgres:
if !GetDBConfig().Enabled {
return errors.New(errs.ErrPostgresNotEnabled)
}
case targetSQLite:
if GetDBConfig().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)
}
@@ -0,0 +1,50 @@
//go:build !windows
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/pkg/logger"
"context"
"fmt"
"os"
"path/filepath"
"syscall"
)
const installedBinaryMode = 0o755
// 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"
if err := os.Remove(backup); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("删除旧备份失败: %w", err)
}
if err := os.Rename(executable, backup); err != nil {
return fmt.Errorf("备份当前程序失败: %w", err)
}
if err := os.Rename(stagedBinary, executable); err != nil {
_ = os.Rename(backup, executable)
return fmt.Errorf("替换当前程序失败: %w", err)
}
if err := os.Chmod(executable, installedBinaryMode); err != nil {
_ = os.Remove(executable)
_ = os.Rename(backup, executable)
return fmt.Errorf("设置程序执行权限失败: %w", err)
}
stagingDir := filepath.Dir(stagedBinary)
_ = os.RemoveAll(stagingDir)
logger.InfoF(ctx, "[Updater] Executing syscall.Exec to restart service: %s %v", executable, os.Args)
//nolint:gosec // restart process via exec with same binary and args
return syscall.Exec(executable, os.Args, os.Environ())
}
@@ -0,0 +1,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)
}
@@ -0,0 +1,204 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// 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"
)
var (
servicesMu sync.RWMutex
dbService contracts.DBService
cacheService contracts.CacheService
userService contracts.UserService
authService contracts.AuthService
taskService contracts.TaskService
storageSvc contracts.StorageService
riskControlService contracts.RiskControlService
eventEmitter func(ctx context.Context, topic string, payload any) error
)
// SetDBService injects the DBService contract.
func SetDBService(s contracts.DBService) {
servicesMu.Lock()
defer servicesMu.Unlock()
dbService = s
repository.SetDBService(s)
}
// SetCacheService injects the CacheService contract.
func SetCacheService(s contracts.CacheService) {
servicesMu.Lock()
defer servicesMu.Unlock()
cacheService = s
repository.SetCacheService(s)
}
// SetUserService injects the UserService contract.
func SetUserService(s contracts.UserService) {
servicesMu.Lock()
defer servicesMu.Unlock()
userService = s
}
// SetAuthService injects the AuthService contract.
func SetAuthService(s contracts.AuthService) {
servicesMu.Lock()
defer servicesMu.Unlock()
authService = s
}
// SetTaskService injects the TaskService contract.
func SetTaskService(s contracts.TaskService) {
servicesMu.Lock()
defer servicesMu.Unlock()
taskService = s
}
// SetStorageService injects the StorageService contract.
func SetStorageService(s contracts.StorageService) {
servicesMu.Lock()
defer servicesMu.Unlock()
storageSvc = s
}
// SetRiskControlService injects the RiskControlService contract.
func SetRiskControlService(s contracts.RiskControlService) {
servicesMu.Lock()
defer servicesMu.Unlock()
riskControlService = s
}
// SetEventEmitter sets the event emission callback.
func SetEventEmitter(fn func(ctx context.Context, topic string, payload any) error) {
servicesMu.Lock()
defer servicesMu.Unlock()
eventEmitter = fn
}
// EmitEvent publishes a domain event if an emitter is registered.
func EmitEvent(ctx context.Context, topic string, payload any) error {
servicesMu.RLock()
defer servicesMu.RUnlock()
if eventEmitter == nil {
return nil
}
return eventEmitter(ctx, topic, payload)
}
// ResetServices clears all injected services (used on disposal and testing).
func ResetServices() {
servicesMu.Lock()
defer servicesMu.Unlock()
dbService = nil
cacheService = nil
userService = nil
authService = nil
taskService = nil
storageSvc = nil
riskControlService = nil
eventEmitter = nil
repository.ResetServices()
}
// GetDB returns the GORM DB instance bound to the context if available.
func GetDB(ctx context.Context) *gorm.DB {
servicesMu.RLock()
defer servicesMu.RUnlock()
if dbService == nil {
return nil
}
return dbService.DB(ctx)
}
// GetCache returns the unified CacheService instance.
func GetCache(_ context.Context) contracts.CacheService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return cacheService
}
// GetUserService returns the UserService instance.
func GetUserService(_ context.Context) contracts.UserService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return userService
}
// GetAuthService returns the AuthService instance.
func GetAuthService(_ context.Context) contracts.AuthService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return authService
}
// GetTaskService returns the TaskService instance.
func GetTaskService() contracts.TaskService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return taskService
}
// GetStorageService returns the StorageService instance.
func GetStorageService() contracts.StorageService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return storageSvc
}
// GetRiskControlService returns the RiskControlService instance.
func GetRiskControlService() contracts.RiskControlService {
servicesMu.RLock()
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,163 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"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 GetDBConfig().Enabled {
return []string{logDBNamePostgres}
}
return []string{logDBNameSQLite}
}
if GetClickHouseConfig().Enabled {
return []string{logDBNameClickHouse}
}
return []string{}
}
@@ -0,0 +1,159 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service_test
import (
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
"Wavelet/plugins/domain/admin/service"
"context"
"testing"
"time"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
type testDBService struct {
db *gorm.DB
}
func (s *testDBService) DB(ctx context.Context) *gorm.DB {
return s.db
}
func (s *testDBService) MasterDB(ctx context.Context) *gorm.DB {
return s.db
}
func (s *testDBService) GORM() *gorm.DB {
return s.db
}
func (s *testDBService) Named(_ string) *gorm.DB {
return s.db
}
func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
if err != nil {
t.Fatalf("gorm.Open(sqlite) error = %v", err)
}
if err := sqliteDB.AutoMigrate(&model.SystemConfig{}); err != nil {
t.Fatalf("AutoMigrate(SystemConfig) error = %v", err)
}
siteConfig := model.SystemConfig{
Key: model.ConfigKeySiteName,
Value: "Wavelet",
Type: "system",
Description: "系统平台的展示名称",
}
if err := sqliteDB.Create(&siteConfig).Error; err != nil {
t.Fatalf("Create(site_name) error = %v", err)
}
service.SetDBService(&testDBService{db: sqliteDB})
cleanup := func() {
repository.StopSystemConfigCacheListener()
repository.ResetSystemConfigRAMCacheForTest()
service.ResetServices()
}
return sqliteDB, cleanup
}
func TestListSystemConfigsByKeys_EmptyKeys(t *testing.T) {
result, err := repository.ListSystemConfigsByKeys(context.Background(), nil)
if err != nil {
t.Fatalf("ListSystemConfigsByKeys(nil) error = %v", err)
}
if len(result) != 0 {
t.Fatalf("ListSystemConfigsByKeys(nil) = %#v, want empty map", result)
}
}
func TestListSystemConfigsByKeys_LoadsFromRAMCache(t *testing.T) {
dbConn, cleanup := setupSystemConfigTest(t)
defer cleanup()
ctx := context.Background()
repository.ResetSystemConfigRAMCacheForTest()
// Initial load
warm, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey(site_name) warm error = %v", err)
}
if warm.Value != "Wavelet" {
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", warm.Value, "Wavelet")
}
// Update DB directly
if err := dbConn.Model(&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 := repository.ListSystemConfigsByKeys(ctx, []string{model.ConfigKeySiteName})
if err != nil {
t.Fatalf("ListSystemConfigsByKeys(site_name) error = %v", err)
}
sc, ok := configs[model.ConfigKeySiteName]
if !ok {
t.Fatal("ListSystemConfigsByKeys(site_name) missing site_name entry")
}
if sc.Value != "Wavelet" {
t.Fatalf("ListSystemConfigsByKeys(site_name).Value = %q, want cached value %q", sc.Value, "Wavelet")
}
}
func TestGetSystemConfigByGroupAndInvalidation(t *testing.T) {
dbConn, cleanup := setupSystemConfigTest(t)
defer cleanup()
ctx := context.Background()
repository.ResetSystemConfigRAMCacheForTest()
// Get via specific group/type
cfg, err := repository.GetSystemConfigByGroup(ctx, repository.ConfigCacheType, model.ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByGroup error = %v", err)
}
if cfg.Value != "Wavelet" {
t.Fatalf("value = %q, want %q", cfg.Value, "Wavelet")
}
// Direct DB update
if err := dbConn.Model(&model.SystemConfig{}).
Where("key = ?", model.ConfigKeySiteName).
Update("value", "new_site_name").Error; err != nil {
t.Fatalf("DB Update error = %v", err)
}
// Invalidate
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
t.Fatalf("InvalidateSystemConfigCache error = %v", err)
}
// Wait for broadcast execution
time.Sleep(100 * time.Millisecond)
// Fetch again
updated, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey error = %v", err)
}
if updated.Value != "new_site_name" {
t.Fatalf("value = %q, want %q", updated.Value, "new_site_name")
}
}
@@ -0,0 +1,217 @@
// 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
}
if _, ok := taskSvc.GetTaskMeta(req.TaskType); !ok {
return "", errs.ErrInvalidTaskType
}
validated, err := validateTaskPayload(taskSvc, req.TaskType, 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)
}
@@ -0,0 +1,635 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package service
import (
"Wavelet/pkg/buildinfo"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/admin/errs"
"Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/domain/admin/repository"
"archive/tar"
"archive/zip"
"compress/gzip"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"os"
"path/filepath"
"runtime"
"strings"
"sync"
"time"
"golang.org/x/mod/semver"
)
const (
githubAPIBaseURL = "https://api.github.com"
maxArchiveSize = int64(1024 * 1024 * 1024)
maxReleaseSize = int64(4 * 1024 * 1024)
repositoryParts = 2
windowsOS = "windows"
archiveFileMode = 0o600
stagedBinaryMode = 0o700
)
type releaseAsset struct {
Name string `json:"name"`
BrowserDownloadURL string `json:"browser_download_url"`
Size int64 `json:"size"`
State string `json:"state"`
}
type githubRelease struct {
TagName string `json:"tag_name"`
Name string `json:"name"`
Body string `json:"body"`
HTMLURL string `json:"html_url"`
Draft bool `json:"draft"`
Prerelease bool `json:"prerelease"`
Published time.Time `json:"published_at"`
Assets []releaseAsset `json:"assets"`
}
type releaseClient interface {
Do(req *http.Request) (*http.Response, error)
}
// UpdaterManager manages application binary updates from GitHub releases.
type UpdaterManager struct {
client releaseClient
mu sync.Mutex
upgrading bool
}
// DefaultUpdaterManager is the default singleton update manager.
var DefaultUpdaterManager = &UpdaterManager{
client: &http.Client{Timeout: 10 * time.Minute},
}
func normalizeVersion(version string) string {
version = strings.TrimSpace(version)
if version == "" || version == "dev" {
return ""
}
if !strings.HasPrefix(version, "v") {
version = "v" + version
}
if !semver.IsValid(version) {
return ""
}
return version
}
func parseRepository(raw string) (string, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return "", errors.New(errs.ErrInvalidRepository)
}
if !strings.Contains(raw, "://") {
repo := strings.TrimSuffix(strings.Trim(raw, "/"), ".git")
if len(strings.Split(repo, "/")) == repositoryParts {
return repo, nil
}
return "", errors.New(errs.ErrInvalidRepository)
}
parsed, err := url.Parse(raw)
if err != nil || !strings.EqualFold(parsed.Hostname(), "github.com") {
return "", errors.New(errs.ErrInvalidRepository)
}
repo := strings.TrimSuffix(strings.Trim(parsed.Path, "/"), ".git")
if len(strings.Split(repo, "/")) != repositoryParts {
return "", errors.New(errs.ErrInvalidRepository)
}
return repo, nil
}
func expectedAssetName(tag string) string {
extension := "tar.gz"
if runtime.GOOS == windowsOS {
extension = "zip"
}
return fmt.Sprintf("wavelet_%s_%s_%s.%s", tag, runtime.GOOS, runtime.GOARCH, extension)
}
func expectedAssetNames(repo, tag string) []string {
names := []string{expectedAssetName(tag)}
if parts := strings.Split(repo, "/"); len(parts) == repositoryParts {
repoName := parts[1]
if repoName != "wavelet" {
extension := "tar.gz"
if runtime.GOOS == windowsOS {
extension = "zip"
}
names = append(names, fmt.Sprintf("%s_%s_%s_%s.%s", repoName, tag, runtime.GOOS, runtime.GOARCH, extension))
}
}
return names
}
func selectLatestRelease(repo string, releases []githubRelease) (githubRelease, releaseAsset, error) {
var selected githubRelease
var selectedAsset releaseAsset
selectedVersion := ""
for _, release := range releases {
version := normalizeVersion(release.TagName)
if release.Draft || version == "" {
continue
}
expectedNames := expectedAssetNames(repo, release.TagName)
for _, asset := range release.Assets {
matched := false
for _, name := range expectedNames {
if asset.Name == name {
matched = true
break
}
}
if !matched || asset.BrowserDownloadURL == "" || asset.State != "uploaded" {
continue
}
if selectedVersion == "" || semver.Compare(version, selectedVersion) > 0 {
selected = release
selectedAsset = asset
selectedVersion = version
}
}
}
if selectedVersion == "" {
return githubRelease{}, releaseAsset{}, errors.New(errs.ErrNoCompatibleRelease)
}
return selected, selectedAsset, nil
}
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, repo),
nil,
)
if err != nil {
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")
req.Header.Set("X-GitHub-Api-Version", "2022-11-28")
resp, err := m.client.Do(req)
if err != nil {
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errs.ErrReleaseRequestFailed, err)
}
defer func() {
_ = resp.Body.Close()
}()
if resp.StatusCode != http.StatusOK {
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", errs.ErrReleaseResponseInvalid, err)
}
release, asset, err := selectLatestRelease(repo, releases)
if err != nil {
return githubRelease{}, releaseAsset{}, err
}
logger.InfoF(ctx, "[Updater] Selected latest compatible release: %s (Asset: %s)", release.TagName, asset.Name)
return release, asset, nil
}
func loadRepository(ctx context.Context) (string, error) {
cfg, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUpdateUpstreamRepository)
if err != nil {
return "", fmt.Errorf("%s: %w", errs.ErrInvalidRepository, err)
}
return parseRepository(cfg.Value)
}
// 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 model.UpdaterStatus{}, releaseAsset{}, err
}
release, asset, err := m.fetchRelease(ctx, upstreamRepo)
if err != nil {
return model.UpdaterStatus{}, releaseAsset{}, err
}
currentVersion := normalizeVersion(buildinfo.Version)
latestVersion := normalizeVersion(release.TagName)
updateAvailable := currentVersion != "" && semver.Compare(latestVersion, currentVersion) > 0
logger.InfoF(ctx, "[Updater] Check update complete. current: %s, latest: %s, update_available: %t", buildinfo.Version, release.TagName, updateAvailable)
return model.UpdaterStatus{
CurrentVersion: buildinfo.Version,
BuildTime: buildinfo.BuildTime,
LatestVersion: release.TagName,
UpdateAvailable: updateAvailable,
CanUpgrade: updateAvailable && runtime.GOOS != windowsOS,
Prerelease: release.Prerelease,
ReleaseName: release.Name,
ReleaseNotes: release.Body,
ReleaseURL: release.HTMLURL,
PublishedAt: release.Published.Format(time.RFC3339),
UpstreamRepository: upstreamRepo,
AssetName: asset.Name,
Platform: runtime.GOOS + "/" + runtime.GOARCH,
}, asset, nil
}
// 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(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(errs.ErrCreateUpgradeRequestFailed, err)
}
req.Header.Set("User-Agent", "Wavelet-Updater")
resp, err := client.Do(req)
if err != nil {
return fmt.Errorf(errs.ErrDownloadUpgradeAssetFailed, err)
}
defer func() {
_ = resp.Body.Close()
}()
if resp.StatusCode != http.StatusOK {
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(errs.ErrCreateUpgradeArchiveFailed, err)
}
written, err := io.Copy(file, io.LimitReader(resp.Body, maxArchiveSize+1))
if err != nil {
_ = file.Close()
return fmt.Errorf(errs.ErrWriteUpgradeArchiveFailed, err)
}
if err := file.Close(); err != nil {
return fmt.Errorf(errs.ErrCloseUpgradeArchiveFailed, err)
}
if written > maxArchiveSize || written != asset.Size {
return fmt.Errorf(errs.ErrUpgradeArchiveSizeMismatch, written, asset.Size)
}
logger.InfoF(ctx, "[Updater] Successfully downloaded release asset to %s", destination)
return nil
}
func safeArchivePath(destination, name string) (string, error) {
cleanName := filepath.Clean(name)
if filepath.IsAbs(cleanName) || cleanName == "." || strings.HasPrefix(cleanName, ".."+string(filepath.Separator)) {
return "", fmt.Errorf(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(errs.ErrArchivePathOutOfDestination, name)
}
return target, nil
}
func matchBinaryName(name string, candidates []string) bool {
for _, candidate := range candidates {
if runtime.GOOS == windowsOS {
if strings.EqualFold(name, candidate) {
return true
}
} else {
if name == candidate {
return true
}
}
}
return false
}
func getCandidateBinaryNames(executable, repo string) []string {
execName := filepath.Base(executable)
names := []string{execName}
addName := func(base string) {
name := base
if runtime.GOOS == windowsOS && !strings.HasSuffix(strings.ToLower(name), ".exe") {
name += ".exe"
}
for _, existing := range names {
if existing == name {
return
}
}
names = append(names, name)
}
if parts := strings.Split(repo, "/"); len(parts) == repositoryParts {
addName(parts[1])
}
addName("wavelet")
return names
}
func isLikelyBinary(name string, isDir bool, mode os.FileMode) bool {
if isDir {
return false
}
base := strings.ToLower(filepath.Base(name))
exclusions := []string{
"license", "licence", "copying", "notice", "readme", "changelog",
}
for _, excl := range exclusions {
if strings.HasPrefix(base, excl) {
return false
}
}
if runtime.GOOS == windowsOS {
return filepath.Ext(base) == ".exe"
}
return (mode.Perm()&0o111 != 0) || (filepath.Ext(base) == "")
}
func findBinaryInTarGz(archivePath string, candidates []string) (string, error) {
//nolint:gosec // updater archivePath is verified
file, err := os.Open(archivePath)
if err != nil {
return "", err
}
defer func() {
_ = file.Close()
}()
gzipReader, err := gzip.NewReader(file)
if err != nil {
return "", err
}
defer func() {
_ = gzipReader.Close()
}()
reader := tar.NewReader(gzipReader)
var binaries []string
for {
header, err := reader.Next()
if errors.Is(err, io.EOF) {
break
}
if err != nil {
return "", err
}
if header.Typeflag == tar.TypeReg && isLikelyBinary(header.Name, false, header.FileInfo().Mode()) {
binaries = append(binaries, header.Name)
}
}
if len(binaries) == 1 {
return binaries[0], nil
}
for _, name := range binaries {
if matchBinaryName(filepath.Base(name), candidates) {
return name, nil
}
}
return "", errors.New(errs.ErrNoCompatibleAsset)
}
func findBinaryInZip(archivePath string, candidates []string) (string, error) {
reader, err := zip.OpenReader(archivePath)
if err != nil {
return "", err
}
defer func() {
_ = reader.Close()
}()
var binaries []string
for _, file := range reader.File {
if !file.FileInfo().IsDir() && isLikelyBinary(file.Name, false, file.FileInfo().Mode()) {
binaries = append(binaries, file.Name)
}
}
if len(binaries) == 1 {
return binaries[0], nil
}
for _, name := range binaries {
if matchBinaryName(filepath.Base(name), candidates) {
return name, nil
}
}
return "", errors.New(errs.ErrNoCompatibleAsset)
}
func extractTarGz(ctx context.Context, archivePath, destination, targetName string, candidates []string) (string, error) {
binaryPathInArchive, err := findBinaryInTarGz(archivePath, candidates)
if err != nil {
return "", err
}
logger.InfoF(ctx, "[Updater] Extracting tar.gz archive: %s (extracting: %s)", archivePath, binaryPathInArchive)
//nolint:gosec // updater archivePath is verified
file, err := os.Open(archivePath)
if err != nil {
return "", err
}
defer func() {
_ = file.Close()
}()
gzipReader, err := gzip.NewReader(file)
if err != nil {
return "", err
}
defer func() {
_ = gzipReader.Close()
}()
reader := tar.NewReader(gzipReader)
for {
header, err := reader.Next()
if errors.Is(err, io.EOF) {
break
}
if err != nil {
return "", err
}
if header.Name != binaryPathInArchive {
continue
}
target, err := safeArchivePath(destination, targetName)
if err != nil {
return "", err
}
//nolint:gosec // updater destination is sanitized
output, err := os.OpenFile(target, os.O_CREATE|os.O_EXCL|os.O_WRONLY, stagedBinaryMode)
if err != nil {
return "", err
}
written, copyErr := io.Copy(output, io.LimitReader(reader, maxArchiveSize+1))
closeErr := output.Close()
if copyErr != nil {
return "", copyErr
}
if closeErr != nil {
return "", closeErr
}
if written > maxArchiveSize {
return "", errors.New(errs.ErrExtractedBinaryTooLarge)
}
logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target)
return target, nil
}
return "", errors.New(errs.ErrNoCompatibleAsset)
}
func extractZip(ctx context.Context, archivePath, destination, targetName string, candidates []string) (string, error) {
binaryPathInArchive, err := findBinaryInZip(archivePath, candidates)
if err != nil {
return "", err
}
logger.InfoF(ctx, "[Updater] Extracting zip archive: %s (extracting: %s)", archivePath, binaryPathInArchive)
reader, err := zip.OpenReader(archivePath)
if err != nil {
return "", err
}
defer func() {
_ = reader.Close()
}()
for _, file := range reader.File {
if file.Name != binaryPathInArchive {
continue
}
target, err := safeArchivePath(destination, targetName)
if err != nil {
return "", err
}
input, err := file.Open()
if err != nil {
return "", err
}
//nolint:gosec // updater extraction target is safe
output, err := os.OpenFile(target, os.O_CREATE|os.O_EXCL|os.O_WRONLY, stagedBinaryMode)
if err != nil {
return "", err
}
written, copyErr := io.Copy(output, io.LimitReader(input, maxArchiveSize+1))
inputCloseErr := input.Close()
outputCloseErr := output.Close()
if copyErr != nil {
return "", copyErr
}
if inputCloseErr != nil {
return "", inputCloseErr
}
if outputCloseErr != nil {
return "", outputCloseErr
}
if written > maxArchiveSize {
return "", errors.New(errs.ErrExtractedBinaryTooLarge)
}
logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target)
return target, nil
}
return "", errors.New(errs.ErrNoCompatibleAsset)
}
// 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(errs.ErrAutomaticUpgradeBlocked)
}
if normalizeVersion(buildinfo.Version) == "" {
return "", "", errors.New(errs.ErrDevelopmentBuild)
}
m.mu.Lock()
defer m.mu.Unlock()
if m.upgrading {
return "", "", errors.New(errs.ErrUpgradeAlreadyRunning)
}
status, asset, err := m.status(ctx)
if err != nil {
return "", "", err
}
if !status.UpdateAvailable {
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(errs.ErrLocateExecutableFailed, err)
}
executable, err = filepath.EvalSymlinks(executable)
if err != nil {
return "", "", fmt.Errorf(errs.ErrResolveExecutablePathFailed, err)
}
tempDir, err := os.MkdirTemp(filepath.Dir(executable), ".wavelet-update-*")
if err != nil {
return "", "", fmt.Errorf(errs.ErrCreateUpgradeDirFailed, err)
}
archivePath := filepath.Join(tempDir, asset.Name)
if err := downloadArchive(ctx, m.client, asset, archivePath); err != nil {
_ = os.RemoveAll(tempDir)
return "", "", err
}
targetName := filepath.Base(executable)
candidates := getCandidateBinaryNames(executable, status.UpstreamRepository)
var stagedBinary string
if strings.HasSuffix(asset.Name, ".zip") {
stagedBinary, err = extractZip(ctx, archivePath, tempDir, targetName, candidates)
} else {
stagedBinary, err = extractTarGz(ctx, archivePath, tempDir, targetName, candidates)
}
if err != nil {
_ = os.RemoveAll(tempDir)
return "", "", fmt.Errorf(errs.ErrExtractUpgradeAssetFailed, err)
}
logger.InfoF(ctx, "[Updater] Staged binary successfully prepared: %s", stagedBinary)
m.upgrading = true
return executable, stagedBinary, nil
}
// 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)
}
+37
View File
@@ -0,0 +1,37 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"context"
"encoding/json"
"github.com/gin-gonic/gin"
)
// LogForAudit 将登录鉴权审计日志写入 Logger
func LogForAudit(ctx context.Context, user *contracts.UserDTO, c *gin.Context) {
if user == nil || c == nil {
return
}
auditLog := loginRequiredAuditLog{
UserID: user.ID,
Username: user.Username,
ClientIP: c.ClientIP(),
Method: c.Request.Method,
Path: c.Request.URL.Path,
RequestURI: c.Request.RequestURI,
UserAgent: c.Request.UserAgent(),
Referer: c.Request.Referer(),
}
auditJSON, err := json.Marshal(auditLog)
if err != nil {
logger.ErrorF(ctx, "[LoginRequiredAudit] marshal failed: %v", err)
logger.DebugF(ctx, "[LoginRequiredAudit] %s %d %s", c.ClientIP(), user.ID, user.Username)
} else {
logger.DebugF(ctx, "[LoginRequiredAudit] %s", auditJSON)
}
}
@@ -0,0 +1,217 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"context"
"errors"
"fmt"
"strconv"
"strings"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
)
func isOIDCLoginEnabled(ctx context.Context) bool {
val, err := GetSystemConfigValue(ctx, "oidc_login_enabled")
if err != nil || val == "" {
return true
}
b, err := strconv.ParseBool(val)
if err != nil {
return true
}
return b
}
func resolveAuthSource(ctx context.Context, sourceName string) (*AuthSource, error) {
name := strings.TrimSpace(strings.ToLower(sourceName))
if name == "" {
sources, err := GetActiveAuthSourcesCached(ctx)
if err != nil {
return nil, err
}
if len(sources) == 0 {
return nil, errors.New(errNoActiveAuthSource)
}
src, err := GetAuthSourceByNameCached(ctx, sources[0].Name)
if err != nil {
return nil, err
}
return src, nil
}
src, err := GetAuthSourceByNameCached(ctx, name)
if err != nil {
return nil, err
}
return src, nil
}
func activeLoginSources(ctx context.Context) []AuthSourceView {
if !isOIDCLoginEnabled(ctx) {
return nil
}
dbSources, err := GetActiveAuthSourcesCached(ctx)
if err != nil {
return nil
}
sources := make([]AuthSourceView, 0, len(dbSources))
for _, source := range dbSources {
sources = append(sources, AuthSourceView{
ID: source.ID,
Name: source.Name,
Type: source.Type,
DisplayName: source.DisplayName,
IsActive: source.IsActive,
IconURL: source.IconURL,
ClientSecretConfigured: source.ClientSecretConfigured,
})
}
return sources
}
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
val, err := GetSystemConfigValue(ctx, "server_address")
if err != nil || strings.TrimSpace(val) == "" {
return "", errors.New(errServerAddressMissing)
}
return strings.TrimRight(val, "/") + "/login", nil
}
func buildOAuthConfig(ctx context.Context, source *AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
if source == nil {
return nil, nil, errors.New(errAuthSourceRequired)
}
if source.OpenIDDiscoveryURL == "" {
return nil, nil, errors.New(errDiscoveryURLRequired)
}
// Clean the issuer URL
issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/")
issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration")
issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server")
provider, err := globalOIDCProviderCache.get(ctx, issuer)
if err != nil {
return nil, nil, err
}
verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID})
scopes := strings.Fields(source.Scopes)
if len(scopes) == 0 {
scopes = []string{oidc.ScopeOpenID, "profile", "email"}
}
if !containsScope(scopes, oidc.ScopeOpenID) {
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
}
return &oauth2.Config{
ClientID: source.ClientID,
ClientSecret: source.ClientSecret,
RedirectURL: redirectURL,
Scopes: scopes,
Endpoint: provider.Endpoint(),
}, verifier, nil
}
func containsScope(scopes []string, scope string) bool {
for _, item := range scopes {
if item == scope {
return true
}
}
return false
}
func buildOAuthUserInfo(ctx context.Context, source *AuthSource, code, nonce, redirectURL string) (*contracts.OAuthUserInfoDTO, error) {
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
if err != nil {
return nil, err
}
token, err := authConfig.Exchange(ctx, code)
if err != nil {
return nil, err
}
userInfo := &contracts.OAuthUserInfoDTO{Active: true}
if verifier != nil {
if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil {
return nil, verifyErr
}
}
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
userInfo.Username = userInfo.PreferredUsername
}
if userInfo.Username == "" && userInfo.Email != "" {
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
}
if userInfo.Username == "" && userInfo.Sub != "" {
userInfo.Username = userInfo.Sub
}
if userInfo.Name == "" {
userInfo.Name = userInfo.Username
}
return userInfo, nil
}
func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *contracts.OAuthUserInfoDTO) error {
rawIDToken, ok := token.Extra("id_token").(string)
if !ok {
return nil
}
idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
if verifyErr != nil {
return fmt.Errorf(errIDTokenVerifyFailedFormat, errIDTokenVerifyFailed, verifyErr)
}
if nonce != "" && idToken.Nonce != nonce {
return errors.New(errNonceMismatch)
}
if claimsErr := idToken.Claims(userInfo); claimsErr != nil {
return claimsErr
}
return nil
}
func normalizeOAuthUserInfo(userInfo *contracts.OAuthUserInfoDTO) error {
userInfo.Username = strings.TrimSpace(userInfo.Username)
userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername)
userInfo.Email = strings.TrimSpace(userInfo.Email)
userInfo.Name = strings.TrimSpace(userInfo.Name)
userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL)
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
userInfo.Username = userInfo.PreferredUsername
}
if userInfo.Username == "" && userInfo.Email != "" {
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
}
if userInfo.Username == "" && userInfo.Sub != "" {
userInfo.Username = userInfo.Sub
}
if userInfo.Username == "" {
return errors.New(errUsernameFromSourceFailed)
}
if userInfo.Name == "" {
userInfo.Name = userInfo.Username
}
if !userInfo.Active {
userInfo.Active = true
}
return nil
}
func buildCallbackResult(user *contracts.UserDTO, status string) OAuthCallbackResult {
result := OAuthCallbackResult{Status: status}
if user != nil {
info := BuildBasicUserInfo(user, false)
result.User = &info
}
return result
}
+116
View File
@@ -0,0 +1,116 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"Wavelet/pkg/cache/ram"
"context"
"fmt"
"time"
)
const (
tokenCacheTTL = 5 * time.Minute
userCacheTTL = 5 * time.Minute
)
// CachedToken represents the minimal cached representation of an access token.
type CachedToken struct {
ID uint64 `json:"id"`
UserID uint64 `json:"user_id"`
IsAdmin bool `json:"is_admin"`
}
var (
tokenRAM = ram.MustNew[string, *CachedToken](ram.Options{MaximumSize: 2048})
userRAM = ram.MustNew[uint64, *contracts.UserDTO](ram.Options{MaximumSize: 2048})
)
func tokenCacheKey(tokenHash string) string {
return "oauth:token:" + tokenHash
}
func userCacheKey(userID uint64) string {
return fmt.Sprintf("oauth:user:%d", userID)
}
// GetCachedToken 获取缓存的 Token
func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) {
if val, ok := tokenRAM.GetIfPresent(tokenHash); ok {
return val, nil
}
if cache := getCache(ctx); cache != nil {
var token CachedToken
key := tokenCacheKey(tokenHash)
if err := cache.Get(ctx, key, &token); err == nil {
tokenRAM.Set(tokenHash, &token)
return &token, nil
}
}
return nil, fmt.Errorf("cache miss")
}
// SetCachedToken 设置 Token 缓存
func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) {
tokenRAM.Set(tokenHash, token)
if cache := getCache(ctx); cache != nil {
key := tokenCacheKey(tokenHash)
_ = cache.Set(ctx, key, token, tokenCacheTTL)
}
}
// InvalidateCachedToken 吊销/删除 token 缓存
func InvalidateCachedToken(ctx context.Context, tokenHash string) {
tokenRAM.Invalidate(tokenHash)
if cache := getCache(ctx); cache != nil {
key := tokenCacheKey(tokenHash)
_ = cache.Delete(ctx, key)
}
}
// GetCachedUser 获取缓存的 UserDTO
func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
if val, ok := userRAM.GetIfPresent(userID); ok {
return val, nil
}
if cache := getCache(ctx); cache != nil {
var u contracts.UserDTO
key := userCacheKey(userID)
if err := cache.Get(ctx, key, &u); err == nil {
userRAM.Set(userID, &u)
return &u, nil
}
}
return nil, fmt.Errorf("cache miss")
}
// SetCachedUser 设置 UserDTO 缓存
func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) {
userRAM.Set(userID, u)
if cache := getCache(ctx); cache != nil {
key := userCacheKey(userID)
_ = cache.Set(ctx, key, u, userCacheTTL)
}
}
// InvalidateCachedUser 吊销/失效 UserDTO 缓存
func InvalidateCachedUser(ctx context.Context, userID uint64) {
userRAM.Invalidate(userID)
if cache := getCache(ctx); cache != nil {
key := userCacheKey(userID)
_ = cache.Delete(ctx, key)
}
}
// StopAuthCacheListener compatibility stub for tests
func StopAuthCacheListener() {}
// ResetAuthRAMCacheForTest clears only the process-local RAM cache.
func ResetAuthRAMCacheForTest() {
tokenRAM.InvalidateAll()
userRAM.InvalidateAll()
}
+144
View File
@@ -0,0 +1,144 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth_test
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/auth"
"context"
"encoding/json"
"testing"
"time"
)
type mockCacheService struct {
items map[string][]byte
}
func newMockCacheService() *mockCacheService {
return &mockCacheService{items: make(map[string][]byte)}
}
func (m *mockCacheService) Get(ctx context.Context, key string, target any) error {
b, ok := m.items[key]
if !ok {
return contracts.ErrCacheMiss
}
return json.Unmarshal(b, target)
}
func (m *mockCacheService) Set(ctx context.Context, key string, value any, ttl time.Duration) error {
b, err := json.Marshal(value)
if err != nil {
return err
}
m.items[key] = b
return nil
}
func (m *mockCacheService) Delete(ctx context.Context, key string) error {
delete(m.items, key)
return nil
}
func (m *mockCacheService) Invalidate(ctx context.Context, key string) error {
return m.Delete(ctx, key)
}
func (m *mockCacheService) GetOrSet(ctx context.Context, key string, target any, ttl time.Duration, loader func() (any, error)) error {
err := m.Get(ctx, key, target)
if err == nil {
return nil
}
val, err := loader()
if err != nil {
return err
}
if err := m.Set(ctx, key, val, ttl); err != nil {
return err
}
b, _ := json.Marshal(val)
return json.Unmarshal(b, target)
}
func TestTokenCache_GetSetInvalidate(t *testing.T) {
ctx := core.NewContext(context.Background())
mockCache := newMockCacheService()
core.Provide[contracts.CacheService](ctx, mockCache)
tokenHash := "test-token-hash"
token := &auth.CachedToken{
ID: 123,
UserID: 456,
IsAdmin: true,
}
// 1. Get from empty cache -> miss
_, err := auth.GetCachedToken(ctx, tokenHash)
if err == nil {
t.Fatal("expected cache miss for un-cached token")
}
// 2. Set to cache
auth.SetCachedToken(ctx, tokenHash, token)
// 3. Get from cache -> hit
cached, err := auth.GetCachedToken(ctx, tokenHash)
if err != nil {
t.Fatalf("GetCachedToken() failed: %v", err)
}
if cached.ID != token.ID || cached.UserID != token.UserID || cached.IsAdmin != token.IsAdmin {
t.Fatalf("expected cached token %+v, got %+v", token, cached)
}
// 4. Invalidate cache
auth.InvalidateCachedToken(ctx, tokenHash)
// 5. Get from cache -> miss
_, err = auth.GetCachedToken(ctx, tokenHash)
if err == nil {
t.Fatal("expected cache miss after invalidation")
}
}
func TestUserCache_GetSetInvalidate(t *testing.T) {
ctx := core.NewContext(context.Background())
mockCache := newMockCacheService()
core.Provide[contracts.CacheService](ctx, mockCache)
userID := uint64(789)
user := &contracts.UserDTO{
ID: userID,
Username: "testuser",
Email: "test@example.com",
}
// 1. Get from empty cache -> miss
_, err := auth.GetCachedUser(ctx, userID)
if err == nil {
t.Fatal("expected cache miss for un-cached user")
}
// 2. Set to cache
auth.SetCachedUser(ctx, userID, user)
// 3. Get from cache -> hit
cached, err := auth.GetCachedUser(ctx, userID)
if err != nil {
t.Fatalf("GetCachedUser() failed: %v", err)
}
if cached.ID != user.ID || cached.Username != user.Username {
t.Fatalf("expected cached user %+v, got %+v", user, cached)
}
// 4. Invalidate cache
auth.InvalidateCachedUser(ctx, userID)
// 5. Get from cache -> miss
_, err = auth.GetCachedUser(ctx, userID)
if err == nil {
t.Fatal("expected cache miss after invalidation")
}
}
+14
View File
@@ -0,0 +1,14 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// SessionConfig defines the session configuration declared by the auth plugin.
type SessionConfig struct {
SessionCookieName string `config:"session_cookie_name" env:"APP_SESSION_COOKIE_NAME" default:"wavelet_session"`
SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"`
SessionDomain string `config:"session_domain" env:"APP_SESSION_DOMAIN"`
SessionAge int `config:"session_age" env:"APP_SESSION_AGE" default:"86400"`
SessionHTTPOnly bool `config:"session_http_only" env:"APP_SESSION_HTTP_ONLY" default:"true"`
SessionSecure bool `config:"session_secure" env:"APP_SESSION_SECURE"`
}
+39
View File
@@ -0,0 +1,39 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"time"
)
// Session and Context Keys
const (
UserNameKey = "username"
UserIDKey = "user_id"
UserObjKey = "user_obj"
TokenAuthKey = "token_auth" // 标记当前请求是否通过 Access Token 鉴权
TokenAdminKey = "token_admin" // Access Token 本身是否具有管理员权限
SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: this is a session key, not hardcoded credentials
PasswordHashKey = "password_hash"
SystemUsername = "system"
)
// OAuth State Cache Keys and Expirations
const (
OAuthStateCacheKeyFormat = "oauth:state:%s"
OAuthStateCacheKeyExpiration = 10 * time.Minute
oauthStateLimitKeyFormat = "oauth:state:limit:%s"
oauthStateLimitMax = 10
)
// OAuth Purpose Constants
const (
OAuthPurposeLogin = "login"
OAuthPurposeBind = "bind"
)
// Auth Source Types
const (
AuthSourceTypeOIDC = "oidc"
)
+54
View File
@@ -0,0 +1,54 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// OAuth and Auth error messages
const (
errInvalidState = "非法登录请求"
errIDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errIDTokenVerifyFailedFormat = "%s: %w"
errNonceMismatch = "nonce 不匹配,可能存在重放攻击"
errNoActiveAuthSource = "未配置可用认证源"
errServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试"
errAuthSourceRequired = "认证源不能为空"
errDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
errUsernameGenerateFailed = "无法生成可用用户名"
errUsernameFromSourceFailed = "无法从认证源获取用户名"
errAuthSourceDisabled = "认证源未启用"
errInvalidExternalAccountBindingID = "绑定记录 ID 无效"
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errOAuthStateRateLimited = "请求授权过于频繁,请稍后重试"
errAuthSourceNameRequired = "认证源名称不能为空"
errAuthSourceNameInvalid = "认证源名称格式不正确"
errAuthSourceTypeUnsupported = "不支持的认证源类型"
errAuthSourceDiscoveryURLRequired = "Discovery URL 不能为空"
//nolint:gosec // error message, not hardcoded credentials
errAuthSourceClientCredentialsRequired = "启用认证源时必须配置 Client ID 和 Client Secret"
errAuthSourceIDRequired = "认证源 ID 不能为空"
errUserIDRequired = "用户 ID 不能为空"
errExternalAccountBindingIncomplete = "外部帐号绑定信息不完整"
errExternalAccountAlreadyBoundToAnother = "该外部帐号已被其他用户绑定"
errExternalAccountBindingIDRequired = "外部帐号绑定记录 ID 不能为空"
errAdminRequired = "无权访问"
//nolint:gosec // error message, not hardcoded credentials
errTokenAdminRequired = "令牌无管理员权限"
errBannedAccount = "账号已被封禁"
errUnAuthorized = "未登录"
)
// 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"
)
+500
View File
@@ -0,0 +1,500 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/idgen"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"context"
"errors"
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/coreos/go-oidc/v3/oidc"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"gorm.io/gorm"
)
// GetLoginSources 获取可用登录源列表
func GetLoginSources(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(activeLoginSources(c.Request.Context())))
}
// GetLoginURL 获取登录授权地址
func GetLoginURL(c *gin.Context) {
ctx := c.Request.Context()
if !isOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
source, err := resolveAuthSource(ctx, c.Query("source"))
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
session := sessions.Default(c)
token, isNew := ensureSessionToken(session)
if isNew {
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
userID := GetUserIDFromSession(session)
sessionHash := hashSessionToken(token)
if err := reserveOAuthStateSlot(ctx, sessionHash); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
state := uuid.NewString()
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
SourceName: source.Name,
Purpose: OAuthPurposeLogin,
UserID: userID,
SessionHash: sessionHash,
})
if err != nil {
response.AbortInternal(c, err.Error())
return
}
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, state)
if cache := getCache(ctx); cache != nil {
if err := cache.Set(ctx, stateKey, payloadValue, OAuthStateCacheKeyExpiration); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
}
func buildAuthorizeURL(ctx context.Context, source *AuthSource, state string) (string, error) {
redirectURL, err := getFrontendLoginRedirectURL(ctx)
if err != nil {
return "", err
}
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
if err != nil {
return "", err
}
if verifier != nil {
return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil
}
return authConfig.AuthCodeURL(state), nil
}
func reserveOAuthStateSlot(ctx context.Context, sessionHash string) error {
if sessionHash == "" {
return nil
}
cache := getCache(ctx)
if cache == nil {
return nil
}
key := fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash)
var count int
_ = cache.Get(ctx, key, &count)
count++
_ = cache.Set(ctx, key, count, OAuthStateCacheKeyExpiration)
if count > oauthStateLimitMax {
return errors.New(errOAuthStateRateLimited)
}
return nil
}
// Authorize 发起指定认证源授权
func Authorize(c *gin.Context) {
ctx := c.Request.Context()
if !isOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
source, err := resolveAuthSource(ctx, c.Param("source"))
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose")))
if purpose != OAuthPurposeBind {
purpose = OAuthPurposeLogin
}
session := sessions.Default(c)
userID := GetUserIDFromSession(session)
if purpose == OAuthPurposeBind && userID == 0 {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
token, isNew := ensureSessionToken(session)
if isNew {
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
sessionHash := hashSessionToken(token)
if err := reserveOAuthStateSlot(ctx, sessionHash); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
state := uuid.NewString()
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
SourceName: source.Name,
Purpose: purpose,
UserID: userID,
SessionHash: sessionHash,
})
if err != nil {
response.AbortInternal(c, err.Error())
return
}
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, state)
if cache := getCache(ctx); cache != nil {
if err := cache.Set(ctx, stateKey, payloadValue, OAuthStateCacheKeyExpiration); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
}
// Callback OAuth 回调处理
func Callback(c *gin.Context) {
var req CallbackRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
ctx := c.Request.Context()
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, req.State)
var payloadRaw string
cache := getCache(ctx)
if cache == nil {
response.AbortBadRequest(c, errInvalidState)
return
}
if err := cache.Get(ctx, stateKey, &payloadRaw); err != nil {
response.AbortBadRequest(c, errInvalidState)
return
}
_ = cache.Delete(ctx, stateKey)
payload, err := decodeOAuthStatePayload(payloadRaw)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
session := sessions.Default(c)
currentUserID := GetUserIDFromSession(session)
if payload.Purpose == OAuthPurposeBind && currentUserID == 0 {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
token, ok := session.Get(SessionTokenKey).(string)
if !ok || token == "" {
response.AbortBadRequest(c, errInvalidSessionContext)
return
}
if hashSessionToken(token) != payload.SessionHash {
response.AbortBadRequest(c, errSessionMismatchForOAuth)
return
}
if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID {
response.AbortBadRequest(c, errUserContextMismatch)
return
}
if !isOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
source, err := resolveAuthSource(ctx, payload.SourceName)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
redirectURL, err := getFrontendLoginRedirectURL(ctx)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
userInfo, err := buildOAuthUserInfo(ctx, source, req.Code, req.State, redirectURL)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := normalizeOAuthUserInfo(userInfo); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if userInfo.Sub == "" {
userInfo.Sub = userInfo.Username
}
if payload.Purpose == OAuthPurposeBind {
handleCallbackBind(ctx, c, source, userInfo)
return
}
handleCallbackLogin(ctx, c, source, userInfo)
}
func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
user, err := GetUserByID(ctx, userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := BindExternalAccount(ctx, &ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
ExternalUsername: userInfo.Username,
Email: userInfo.Email,
}); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
user.LastLoginAt = time.Now()
_ = 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
account, err := FindExternalAccount(ctx, source.ID, userInfo.Sub)
switch {
case err == 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
default:
response.AbortInternal(c, err.Error())
return
}
user.LastLoginAt = time.Now()
_ = 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)
c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "logged_in")))
}
func uniqueUsername(ctx context.Context, base string) (string, error) {
base = strings.TrimSpace(base)
if base == "" {
base = "user"
}
existingUsernames, err := ListSimilarUsernames(ctx, base)
if err != nil {
return "", err
}
exists := make(map[string]bool, len(existingUsernames))
for _, u := range existingUsernames {
exists[strings.ToLower(u)] = true
}
if !exists[strings.ToLower(base)] {
return base, nil
}
for i := 1; i <= 1000; i++ {
candidate := fmt.Sprintf("%s-%d", base, i)
if !exists[strings.ToLower(candidate)] {
return candidate, nil
}
}
return "", errors.New(errUsernameGenerateFailed)
}
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) (contracts.UserDTO, bool) {
registrationEnabled := true
val, cfgErr := GetSystemConfigValue(ctx, "registration_enabled")
if cfgErr == nil && val != "" {
if b, err := strconv.ParseBool(val); err == nil {
registrationEnabled = b
}
}
if !registrationEnabled {
c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind")))
return contracts.UserDTO{}, false
}
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
if uniqueErr != nil {
response.AbortInternal(c, uniqueErr.Error())
return contracts.UserDTO{}, false
}
userInfo.Username = username
now := time.Now()
user := contracts.UserDTO{
ID: idgen.NextUint64ID(),
Username: userInfo.Username,
Nickname: userInfo.Name,
Email: userInfo.Email,
AvatarURL: userInfo.AvatarURL,
IsActive: userInfo.Active,
LastLoginAt: now,
CreatedAt: now,
UpdatedAt: now,
}
if err := InsertUser(ctx, &user); err != nil {
response.AbortInternal(c, err.Error())
return contracts.UserDTO{}, false
}
if err := BindExternalAccount(ctx, &ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
ExternalUsername: userInfo.Username,
Email: userInfo.Email,
}); err != nil {
response.AbortBadRequest(c, err.Error())
return contracts.UserDTO{}, false
}
logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
return user, true
}
// UserInfo 获取当前登录用户信息
func UserInfo(c *gin.Context) {
user, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
session := sessions.Default(c)
needChange := session.Get("need_change_password") == true || (user != nil && user.NeedChangePassword)
c.JSON(
http.StatusOK,
response.OK(BuildBasicUserInfo(user, needChange)),
)
}
// Logout 退出登录
func Logout(c *gin.Context) {
session := sessions.Default(c)
userID := session.Get(UserIDKey)
username := session.Get(UserNameKey)
if userID != nil {
logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP())
if id := ParseUserID(userID); id > 0 {
InvalidateCachedUser(c.Request.Context(), id)
}
}
session.Options(GetSessionOptions(-1))
session.Clear()
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// ListExternalAccounts 获取当前用户的外部帐号绑定列表
func ListExternalAccounts(c *gin.Context) {
userID := GetUserIDFromContext(c)
accounts, err := ListExternalAccountsByUserID(c.Request.Context(), userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(accounts))
}
// DeleteExternalAccount 解除外部帐号绑定
func DeleteExternalAccount(c *gin.Context) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
rawID := strings.TrimSpace(c.Param("id"))
id, err := strconv.ParseUint(rawID, 10, 64)
if err != nil || id == 0 {
response.AbortBadRequest(c, errInvalidExternalAccountBindingID)
return
}
if err := UnbindExternalAccount(c.Request.Context(), id, userID); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
+197
View File
@@ -0,0 +1,197 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/pkg/trace"
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"github.com/gin-gonic/gin"
)
// whitelist holds the no-auth route patterns. They are registered during Apply and
// matched on every request, so PathWhitelist parses them once up front.
var whitelist = extpoints.NewPathWhitelist()
// RegisterWhitelist registers route patterns that bypass mandatory authentication.
func RegisterWhitelist(patterns ...string) {
whitelist.Add(patterns...)
}
// IsWhitelisted checks if the specified path matches the auth whitelist.
func IsWhitelisted(path string) bool {
return whitelist.Match(path)
}
func hashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
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 {
tokenRecord, err = GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, nil, err
}
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
user, err = GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, nil, err
}
SetCachedUser(ctx, tokenRecord.UserID, user)
}
return user, tokenRecord, nil
}
// GetUserFromRequest 从请求中获取当前用户(优先 Access Token,其次 Session)
func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
ctx := c.Request.Context()
var tokenStr string
tokenFromQuery := c.Query("token")
if tokenFromQuery != "" {
tokenStr = tokenFromQuery
} else {
authHeader := c.GetHeader("Authorization")
if len(authHeader) > 7 && authHeader[:7] == "Bearer " {
tokenStr = authHeader[7:]
}
}
// 优先使用 Access Token 鉴权
if tokenStr != "" {
if user, tokenRecord, err := getUserByToken(ctx, tokenStr); err == nil {
if user.Username == SystemUsername {
return nil, errors.New(errSystemUserLoginNotAllowed)
}
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, true)
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin)
return user, nil
}
}
// 降级使用 Session 鉴权
userID := GetUserIDFromContext(c)
if userID <= 0 {
return nil, errors.New(errUnauthorizedInternal)
}
user, err := GetCachedUser(ctx, userID)
if err != nil || user == nil || !user.IsActive {
user, err = GetActiveUserByID(ctx, userID)
if err != nil {
return nil, err
}
SetCachedUser(ctx, userID, user)
}
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, false)
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, false)
if user.Username == "system" {
return nil, errors.New(errSystemUserLoginNotAllowed)
}
return user, nil
}
// LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session
func LoginRequired() gin.HandlerFunc {
return func(c *gin.Context) {
if IsWhitelisted(c.Request.URL.Path) {
c.Next()
return
}
_, span := trace.Start(c.Request.Context(), "LoginRequired")
defer span.End()
user, err := GetUserFromRequest(c)
if err != nil {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
LogForAudit(c.Request.Context(), user, c)
ginutil.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next()
}
}
// AdminRequired 校验管理员权限(支持 Session 和 Token 鉴权)
func AdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
_, span := trace.Start(c.Request.Context(), "AdminRequired")
defer span.End()
user, err := GetUserFromRequest(c)
if err != nil {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
isTokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey)
isTokenAdmin, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
// 如果是通过 Token 鉴权,要求该 Token 具备管理员权限或者用户本身为管理员
if isTokenAuth && !isTokenAdmin && !user.IsAdmin {
response.AbortNotFound(c, errTokenAdminRequired)
return
}
// 如果是通过 Session 鉴权,直接检查用户的 is_admin 属性
if !isTokenAuth && !user.IsAdmin {
response.AbortNotFound(c, errAdminRequired)
return
}
LogForAudit(c.Request.Context(), user, c)
ginutil.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next()
}
}
// LoginAdminRequired is an alias for AdminRequired.
func LoginAdminRequired() gin.HandlerFunc {
return AdminRequired()
}
// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点
func DisallowTokenAuth() gin.HandlerFunc {
return func(c *gin.Context) {
if tokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
response.AbortForbidden(c, ErrTokenAuthNotAllowed)
return
}
c.Next()
}
}
@@ -0,0 +1,53 @@
-- +goose Up
-- +goose StatementBegin
CREATE TABLE IF NOT EXISTS w_auth_sources (
id BIGINT PRIMARY KEY,
name VARCHAR(80) NOT NULL UNIQUE,
type VARCHAR(20) NOT NULL,
display_name VARCHAR(100),
is_active BOOLEAN NOT NULL DEFAULT FALSE,
client_id VARCHAR(255),
client_secret VARCHAR(1024),
openid_discovery_url VARCHAR(1024),
scopes VARCHAR(255),
icon_url VARCHAR(1024),
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_auth_sources_is_active ON w_auth_sources (is_active);
CREATE TABLE IF NOT EXISTS w_external_accounts (
id BIGINT PRIMARY KEY,
auth_source_id BIGINT,
user_id BIGINT NOT NULL,
external_id VARCHAR(255) NOT NULL,
external_username VARCHAR(255),
email VARCHAR(255),
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_external_accounts_auth_source_id ON w_external_accounts (auth_source_id);
CREATE INDEX IF NOT EXISTS idx_w_external_accounts_user_id ON w_external_accounts (user_id);
CREATE UNIQUE INDEX IF NOT EXISTS idx_w_external_accounts_source_external ON w_external_accounts (auth_source_id, external_id);
CREATE TABLE IF NOT EXISTS w_access_tokens (
id BIGINT PRIMARY KEY,
user_id BIGINT NOT NULL,
token_hash VARCHAR(64) NOT NULL UNIQUE,
name VARCHAR(128) NOT NULL,
masked_token VARCHAR(64) NOT NULL DEFAULT '',
description VARCHAR(255),
is_admin BOOLEAN DEFAULT FALSE,
expires_at TIMESTAMPTZ,
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_access_tokens_user_id ON w_access_tokens (user_id);
-- +goose StatementEnd
-- +goose Down
-- +goose StatementBegin
DROP TABLE IF EXISTS w_access_tokens;
DROP TABLE IF EXISTS w_external_accounts;
DROP TABLE IF EXISTS w_auth_sources;
-- +goose StatementEnd
@@ -0,0 +1,53 @@
-- +goose Up
-- +goose StatementBegin
CREATE TABLE IF NOT EXISTS w_auth_sources (
id BIGINT PRIMARY KEY,
name VARCHAR(80) NOT NULL UNIQUE,
type VARCHAR(20) NOT NULL,
display_name VARCHAR(100),
is_active BOOLEAN NOT NULL DEFAULT 0,
client_id VARCHAR(255),
client_secret VARCHAR(1024),
openid_discovery_url VARCHAR(1024),
scopes VARCHAR(255),
icon_url VARCHAR(1024),
created_at DATETIME DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_auth_sources_is_active ON w_auth_sources (is_active);
CREATE TABLE IF NOT EXISTS w_external_accounts (
id BIGINT PRIMARY KEY,
auth_source_id BIGINT,
user_id BIGINT NOT NULL,
external_id VARCHAR(255) NOT NULL,
external_username VARCHAR(255),
email VARCHAR(255),
created_at DATETIME DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_external_accounts_auth_source_id ON w_external_accounts (auth_source_id);
CREATE INDEX IF NOT EXISTS idx_w_external_accounts_user_id ON w_external_accounts (user_id);
CREATE UNIQUE INDEX IF NOT EXISTS idx_w_external_accounts_source_external ON w_external_accounts (auth_source_id, external_id);
CREATE TABLE IF NOT EXISTS w_access_tokens (
id BIGINT PRIMARY KEY,
user_id BIGINT NOT NULL,
token_hash VARCHAR(64) NOT NULL UNIQUE,
name VARCHAR(128) NOT NULL,
masked_token VARCHAR(64) NOT NULL DEFAULT '',
description VARCHAR(255),
is_admin BOOLEAN DEFAULT 0,
expires_at DATETIME,
created_at DATETIME DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_access_tokens_user_id ON w_access_tokens (user_id);
-- +goose StatementEnd
-- +goose Down
-- +goose StatementBegin
DROP TABLE IF EXISTS w_access_tokens;
DROP TABLE IF EXISTS w_external_accounts;
DROP TABLE IF EXISTS w_auth_sources;
-- +goose StatementEnd
+241
View File
@@ -0,0 +1,241 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"encoding/json"
"errors"
"regexp"
"strconv"
"strings"
"time"
)
var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`)
// AuthSource 认证源实体
//
//nolint:revive // auth.AuthSource is standard domain entity name
type AuthSource struct {
ID uint64 `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"uniqueIndex;size:80;not null"`
Type string `json:"type" gorm:"size:20;not null"`
DisplayName string `json:"display_name" gorm:"size:100"`
IsActive bool `json:"is_active" gorm:"index;not null;default:false"`
ClientID string `json:"client_id" gorm:"size:255"`
ClientSecret string `json:"-" gorm:"size:1024"`
OpenIDDiscoveryURL string `json:"openid_discovery_url" gorm:"column:openid_discovery_url;size:1024"`
Scopes string `json:"scopes" gorm:"size:255"`
IconURL string `json:"icon_url" gorm:"size:1024"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
ClientSecretConfigured bool `json:"client_secret_configured" gorm:"-"`
}
// TableName 表名
func (AuthSource) TableName() string {
return "w_auth_sources"
}
// Normalize 对认证源字段进行标准化处理
func (source *AuthSource) Normalize() {
source.Type = strings.ToLower(strings.TrimSpace(source.Type))
source.Name = strings.TrimSpace(source.Name)
source.DisplayName = strings.TrimSpace(source.DisplayName)
source.ClientID = strings.TrimSpace(source.ClientID)
source.ClientSecret = strings.TrimSpace(source.ClientSecret)
source.OpenIDDiscoveryURL = strings.TrimSpace(source.OpenIDDiscoveryURL)
source.Scopes = strings.TrimSpace(source.Scopes)
source.IconURL = strings.TrimSpace(source.IconURL)
if source.DisplayName == "" {
source.DisplayName = source.Name
}
if source.Type == AuthSourceTypeOIDC && source.Scopes == "" {
source.Scopes = "openid profile email"
}
}
// Validate 校验认证源字段合法性
func (source *AuthSource) Validate() error {
source.Normalize()
if source.Name == "" {
return errors.New(errAuthSourceNameRequired)
}
if !authSourceNamePattern.MatchString(source.Name) {
return errors.New(errAuthSourceNameInvalid)
}
if source.Type != AuthSourceTypeOIDC {
return errors.New(errAuthSourceTypeUnsupported)
}
if source.OpenIDDiscoveryURL == "" {
//nolint:staticcheck // descriptive error constant
return errors.New(errAuthSourceDiscoveryURLRequired)
}
if source.IsActive && (source.ClientID == "" || source.ClientSecret == "") {
return errors.New(errAuthSourceClientCredentialsRequired)
}
return nil
}
// Sanitize 脱敏处理,将 ClientSecret 清空并设置 ClientSecretConfigured 标志
func (source *AuthSource) Sanitize() {
source.ClientSecretConfigured = source.ClientSecret != ""
source.ClientSecret = ""
}
// ExternalAccount 外部账号绑定实体
type ExternalAccount struct {
ID uint64 `json:"id" gorm:"primaryKey"`
AuthSourceID uint64 `json:"auth_source_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:1;index"`
UserID uint64 `json:"user_id" gorm:"index;not null"`
ExternalID string `json:"external_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:2;size:255;not null"`
ExternalUsername string `json:"external_username" gorm:"size:255"`
Email string `json:"email" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// TableName 表名
func (ExternalAccount) TableName() string {
return "w_external_accounts"
}
// ExternalAccountView 外部帐号绑定视图(脱敏展示用)
type ExternalAccountView struct {
ID uint64 `json:"id"`
AuthSourceID uint64 `json:"auth_source_id"`
AuthSourceName string `json:"auth_source_name"`
AuthSourceType string `json:"auth_source_type"`
AuthSourceLabel string `json:"auth_source_label"`
ExternalUsername string `json:"external_username"`
Email string `json:"email"`
CreatedAt time.Time `json:"created_at"`
}
// AuthSourceView 登录源展示信息
//
//nolint:revive // auth.AuthSourceView is standard domain presentation struct
type AuthSourceView struct {
ID uint64 `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
DisplayName string `json:"display_name"`
IsActive bool `json:"is_active"`
IconURL string `json:"icon_url"`
ClientSecretConfigured bool `json:"client_secret_configured"`
}
// OAuthAuthorizeResponse 授权 URL 响应
type OAuthAuthorizeResponse struct {
AuthorizeURL string `json:"authorize_url"`
}
// OAuthCallbackResult 回调处理结果
type OAuthCallbackResult struct {
Status string `json:"status"`
User *BasicUserInfo `json:"user,omitempty"`
}
// CallbackRequest OAuth 回调请求参数
type CallbackRequest struct {
State string `json:"state" binding:"required"`
Code string `json:"code" binding:"required"`
}
// BasicUserInfo 用户基本信息结构体
type BasicUserInfo struct {
ID uint64 `json:"id"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Email string `json:"email"`
AvatarURL string `json:"avatar_url"`
IsAdmin bool `json:"is_admin"`
NeedChangePassword bool `json:"need_change_password"`
Bio string `json:"bio"`
Phone string `json:"phone"`
Gender string `json:"gender"`
Website string `json:"website"`
Location string `json:"location"`
}
// BuildBasicUserInfo 将 UserDTO 转换为 BasicUserInfo
func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo {
if user == nil {
return BasicUserInfo{}
}
return BasicUserInfo{
ID: user.ID,
Username: user.Username,
Nickname: user.Nickname,
Email: user.Email,
AvatarURL: user.AvatarURL,
IsAdmin: user.IsAdmin,
NeedChangePassword: needChange || user.NeedChangePassword,
Bio: user.Bio,
Phone: user.Phone,
Gender: user.Gender,
Website: user.Website,
Location: user.Location,
}
}
type oauthStatePayload struct {
SourceName string `json:"source_name"`
Purpose string `json:"purpose"`
UserID uint64 `json:"user_id,omitempty"`
SessionHash string `json:"session_hash"`
}
func encodeOAuthStatePayload(payload oauthStatePayload) (string, error) {
data, err := json.Marshal(payload)
if err != nil {
return "", err
}
return string(data), nil
}
func decodeOAuthStatePayload(value string) (oauthStatePayload, error) {
var payload oauthStatePayload
if err := json.Unmarshal([]byte(value), &payload); err != nil {
return oauthStatePayload{}, err
}
return payload, nil
}
type loginRequiredAuditLog struct {
UserID uint64 `json:"user_id"`
Username string `json:"username"`
ClientIP string `json:"client_ip"`
Method string `json:"method"`
Path string `json:"path"`
RequestURI string `json:"request_uri"`
UserAgent string `json:"user_agent"`
Referer string `json:"referer"`
}
// ParseUserID parses a string or float64 user ID representation.
func ParseUserID(v any) uint64 {
switch val := v.(type) {
case uint64:
return val
case int64:
if val > 0 {
return uint64(val)
}
case int:
if val > 0 {
return uint64(val)
}
case float64:
if val > 0 {
return uint64(val)
}
case string:
if id, err := strconv.ParseUint(val, 10, 64); err == nil {
return id
}
}
return 0
}
+185
View File
@@ -0,0 +1,185 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package auth provides the authentication, OAuth, session management, and access token domain plugin for Cordis.
package auth
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"context"
"embed"
"reflect"
)
//go:embed migrations/*/*.sql
var authMigrations embed.FS
// Option configures the auth plugin.
type Option func(*Plugin)
// WithAuthService sets a custom AuthService implementation.
func WithAuthService(svc contracts.AuthService) Option {
return func(p *Plugin) {
p.authSvc = svc
}
}
// WithAuthRegistry sets a custom AuthRegistry implementation.
func WithAuthRegistry(reg contracts.AuthRegistry) Option {
return func(p *Plugin) {
p.authRegistry = reg
}
}
// Plugin implements core.Plugin to provide authentication and OAuth domain services.
type Plugin struct {
authSvc contracts.AuthService
authRegistry contracts.AuthRegistry
}
// New creates a new auth domain plugin.
func New(opts ...Option) *Plugin {
p := &Plugin{}
for _, opt := range opts {
if opt != nil {
opt(p)
}
}
return p
}
// Name returns the unique identifier for the auth domain plugin.
func (p *Plugin) Name() string {
return "auth"
}
// Inject declares required dependencies for the auth domain plugin.
func (p *Plugin) Inject() []reflect.Type {
return []reflect.Type{
reflect.TypeFor[contracts.DBService](),
reflect.TypeFor[contracts.CacheService](),
}
}
// Manifest returns the plugin metadata.
func (p *Plugin) Manifest() core.Manifest {
return core.Manifest{
Name: "auth",
Version: "1.0.0",
Description: "Authentication, OAuth, Session and Passkey domain plugin",
Author: "Wavelet Team",
}
}
// DeclareConfig declares configuration bindings for the auth plugin.
func (p *Plugin) DeclareConfig() []core.ConfigBinding {
return []core.ConfigBinding{
{Prefix: "app", Target: &SessionConfig{}},
}
}
// Apply registers the auth migrations, services, routes, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
var cfg SessionConfig
if err := ctx.Config().Bind("app", &cfg); err == nil {
SetSessionConfig(cfg)
}
// 0. Bind DBService & CacheService from Context
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
setDBService(db)
} else {
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
setDBService(db)
})
}
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
setCacheService(cache)
} else {
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
setCacheService(cache)
})
}
ctx.OnDispose(func() error {
setDBService(nil)
setCacheService(nil)
return nil
})
// 1. Register migrations
ctx.Migrations().Register("auth", authMigrations)
// 2. Initialize and provide AuthService & AuthRegistry
if p.authSvc == nil {
p.authSvc = newAuthService()
}
if p.authRegistry == nil {
p.authRegistry = newAuthRegistry()
}
core.Provide[contracts.AuthService](ctx, p.authSvc)
core.Provide[contracts.AuthRegistry](ctx, p.authRegistry)
// 2.1 Register Public / Auth Whitelist Endpoints
publicEndpoints := []string{
"/api/v1/oauth/sources",
"/api/v1/oauth/login",
"/api/v1/oauth/*/authorize",
"/api/v1/oauth/:source/authorize",
"/api/v1/oauth/callback",
"/api/v1/user/login",
"/api/v1/user/register",
"/api/v1/user/send-email-code",
"/api/v1/cap/challenge",
"/api/v1/cap/redeem",
"/healthz",
"/metrics",
}
RegisterWhitelist(publicEndpoints...)
ctx.Router().RegisterWhitelist(publicEndpoints...)
// 3. Register HTTP Routes
oauthGroup := ctx.Router().Group("/api/v1/oauth")
{
oauthGroup.GET("/sources", GetLoginSources)
oauthGroup.GET("/login", GetLoginURL)
oauthGroup.GET("/:source/authorize", Authorize)
oauthGroup.GET("/logout", Logout)
oauthGroup.POST("/callback", Callback)
oauthGroup.GET("/user-info", LoginRequired(), UserInfo)
oauthGroup.GET("/external-accounts", LoginRequired(), ListExternalAccounts)
oauthGroup.POST("/external-accounts/:id/delete", LoginRequired(), DeleteExternalAccount)
}
ctx.Router().GET("/api/v1/user-info", LoginRequired(), UserInfo)
// 4. Register Settings Schemas
ctx.Settings().Register(extpoints.SettingSchema{
Key: "auth.session_age",
Default: 86400 * 7,
Description: "Default session lifetime in seconds",
Type: "integer",
Category: "security",
})
ctx.Settings().Register(extpoints.SettingSchema{
Key: "auth.login_rate_limit_max_attempts",
Default: 5,
Description: "Max login failure attempts before temporary IP lock",
Type: "integer",
Category: "security",
})
// 5. Register Event Listeners for domain events
ctx.Events().On(contracts.EventTopicUserStatusChanged, func(c context.Context, e contracts.UserStatusChangedEvent) error {
InvalidateCachedUser(c, e.UserID)
return nil
})
ctx.Events().On(contracts.EventTopicUserDeleted, func(c context.Context, e contracts.UserDeletedEvent) error {
InvalidateCachedUser(c, e.TargetUserID)
return nil
})
return nil
}
+155
View File
@@ -0,0 +1,155 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth_test
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/auth"
"context"
"crypto/sha256"
"encoding/hex"
"path/filepath"
"testing"
"time"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
type mockDBService struct {
db *gorm.DB
}
func (m *mockDBService) GORM() *gorm.DB {
return m.db
}
func (m *mockDBService) DB(ctx context.Context) *gorm.DB {
return m.db.WithContext(ctx)
}
func (m *mockDBService) Named(_ string) *gorm.DB {
return m.db
}
type testUser struct {
ID uint64 `gorm:"primaryKey"`
Username string
IsActive bool
LastLoginAt time.Time
}
func (testUser) TableName() string { return "w_users" }
type testAccessToken struct {
ID uint64 `gorm:"primaryKey"`
UserID uint64
TokenHash string
Name string
IsAdmin bool
}
func (testAccessToken) TableName() string { return "w_access_tokens" }
func hashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func setupTestDB(t *testing.T) *gorm.DB {
t.Helper()
dbPath := filepath.Join(t.TempDir(), "auth_test.db")
testDB, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, testDB.AutoMigrate(
&testUser{},
&testAccessToken{},
&auth.AuthSource{},
&auth.ExternalAccount{},
))
return testDB
}
type mockProvider struct{}
func (m *mockProvider) Name() string { return "custom" }
func (m *mockProvider) GetAuthURL(state string) string {
return "https://custom.com/auth?state=" + state
}
func (m *mockProvider) ExchangeCode(ctx context.Context, code string) (*contracts.OAuthUserInfoDTO, error) {
return &contracts.OAuthUserInfoDTO{
ID: 555,
Username: "custom_user",
Email: "custom@example.com",
}, nil
}
func TestAuthPluginUnit(t *testing.T) {
ctx := core.NewContext(context.Background())
testDB := setupTestDB(t)
core.Provide[contracts.DBService](ctx, &mockDBService{db: testDB})
p := auth.New()
assert.Equal(t, "auth", p.Name())
assert.Equal(t, "1.0.0", p.Manifest().Version)
require.NoError(t, p.Apply(ctx))
// Test AuthService injection
authSvc, err := core.Inject[contracts.AuthService](ctx)
require.NoError(t, err)
assert.NotNil(t, authSvc.RequireAuthMiddleware())
assert.NotNil(t, authSvc.RequireAdminMiddleware())
// Test AuthRegistry injection
authReg, err := core.Inject[contracts.AuthRegistry](ctx)
require.NoError(t, err)
authReg.RegisterOAuthProvider("custom", &mockProvider{})
prov, ok := authReg.GetOAuthProvider("custom")
require.True(t, ok)
assert.Equal(t, "custom", prov.Name())
// Test User Token Verification with dummy token
user := testUser{
ID: 101,
Username: "token_user",
IsActive: true,
}
require.NoError(t, testDB.Create(&user).Error)
tokenStr := "test-secret-token-123456"
tokenHash := hashToken(tokenStr)
tokenRecord := testAccessToken{
ID: 201,
UserID: user.ID,
TokenHash: tokenHash,
Name: "test-token",
IsAdmin: false,
}
require.NoError(t, testDB.Create(&tokenRecord).Error)
userDTO, err := authSvc.VerifyToken(context.Background(), tokenStr)
require.NoError(t, err)
assert.Equal(t, user.ID, userDTO.ID)
assert.Equal(t, "token_user", userDTO.Username)
// Empty token fails
_, err = authSvc.VerifyToken(context.Background(), "")
assert.Error(t, err)
// Revoke sessions
require.NoError(t, authSvc.RevokeUserSessions(context.Background(), user.ID))
// GetCurrentUser from context
userCtx := context.WithValue(context.Background(), contracts.AuthUserObjKey, userDTO)
current, err := authSvc.GetCurrentUser(userCtx)
require.NoError(t, err)
assert.Equal(t, user.ID, current.ID)
}
@@ -0,0 +1,81 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"net/http"
"sync"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
"golang.org/x/sync/singleflight"
)
// oidcProviderCache 进程级 OIDC provider 缓存。
type oidcProviderCache struct {
mu sync.RWMutex
entries map[string]*oidc.Provider // key: normalized issuer URL
sfGroup singleflight.Group
}
// globalOIDCProviderCache 是包级单例缓存,与进程同生命周期。
var globalOIDCProviderCache = &oidcProviderCache{
entries: make(map[string]*oidc.Provider),
}
// discoveryContext 从请求 ctx 提取 HTTP 客户端,并绑定到不可取消的 Background ctx。
func discoveryContext(ctx context.Context) context.Context {
bg := context.Background()
if client, ok := ctx.Value(oauth2.HTTPClient).(*http.Client); ok && client != nil {
bg = oidc.ClientContext(bg, client)
}
return bg
}
// get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。
func (c *oidcProviderCache) get(ctx context.Context, issuer string) (*oidc.Provider, error) {
c.mu.RLock()
if p, ok := c.entries[issuer]; ok {
c.mu.RUnlock()
return p, nil
}
c.mu.RUnlock()
discCtx := discoveryContext(ctx)
v, err, _ := c.sfGroup.Do(issuer, func() (any, error) {
c.mu.RLock()
if p, ok := c.entries[issuer]; ok {
c.mu.RUnlock()
return p, nil
}
c.mu.RUnlock()
p, err := oidc.NewProvider(discCtx, issuer)
if err != nil {
return nil, err
}
c.mu.Lock()
c.entries[issuer] = p
c.mu.Unlock()
return p, nil
})
if err != nil {
return nil, err
}
return v.(*oidc.Provider), nil //nolint:forcetypeassert
}
// invalidate 从缓存中移除指定 issuer 对应的 provider。
func (c *oidcProviderCache) invalidate(issuer string) {
c.mu.Lock()
delete(c.entries, issuer)
c.mu.Unlock()
}
// InvalidateOIDCProviderCache 从进程级缓存中清除指定 issuer 的 provider 条目。
func InvalidateOIDCProviderCache(issuer string) {
globalOIDCProviderCache.invalidate(issuer)
}
+215
View File
@@ -0,0 +1,215 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/util"
"context"
"sync"
"time"
"gorm.io/gorm"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
cacheMu sync.RWMutex
cacheSvc contracts.CacheService
)
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 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
}
// 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
if err := getDB(ctx).First(&src, id).Error; err != nil {
return nil, err
}
return &src, nil
}
// GetAuthSourceByName 根据名称获取认证源
func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) {
var src AuthSource
if err := getDB(ctx).Where("name = ?", name).First(&src).Error; err != nil {
return nil, err
}
return &src, nil
}
// ListActiveAuthSources 获取所有启用的认证源
func ListActiveAuthSources(ctx context.Context) ([]AuthSource, error) {
var sources []AuthSource
if err := getDB(ctx).Where("is_active = ?", true).Order("id ASC").Find(&sources).Error; err != nil {
return nil, err
}
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)
}
// GetAuthSourceByNameCached 根据名称获取认证源(带缓存或直接查询)
func GetAuthSourceByNameCached(ctx context.Context, name string) (*AuthSource, error) {
return GetAuthSourceByName(ctx, name)
}
// FindExternalAccount 查询指定认证源的外部账号绑定
func FindExternalAccount(ctx context.Context, authSourceID uint64, externalID string) (*ExternalAccount, error) {
var account ExternalAccount
if err := getDB(ctx).Where("auth_source_id = ? AND external_id = ?", authSourceID, externalID).First(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
// BindExternalAccount 绑定外部账号
func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
return getDB(ctx).Create(account).Error
}
// ListExternalAccountsByUserID 获取用户绑定的所有外部账号
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccount, error) {
var accounts []ExternalAccount
if err := getDB(ctx).Where("user_id = ?", userID).Find(&accounts).Error; err != nil {
return nil, err
}
return accounts, nil
}
// UnbindExternalAccount 解绑外部账号
func UnbindExternalAccount(ctx context.Context, id, userID uint64) error {
return getDB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
}
+262
View File
@@ -0,0 +1,262 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"context"
"errors"
"sync"
)
type authServiceImpl struct{}
func newAuthService() contracts.AuthService {
return &authServiceImpl{}
}
func (s *authServiceImpl) RequireAuthMiddleware() any {
return LoginRequired()
}
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 v := ctx.Value(contracts.AuthUserObjKey); v != nil {
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
return u, nil
}
}
return nil, errors.New(errUserNotInContext)
}
func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contracts.UserDTO, error) {
if token == "" {
return nil, errors.New(errEmptyToken)
}
tokenHash := hashToken(token)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
tokenRecord, err = GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, err
}
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
user, err = GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, err
}
SetCachedUser(ctx, tokenRecord.UserID, user)
}
if user.Username == SystemUsername {
return nil, errors.New(errSystemUserTokenNotAllowed)
}
return user, nil
}
func (s *authServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
return "", nil
}
func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error {
InvalidateCachedUser(ctx, userID)
return nil
}
// GetCurrentUserID 从请求登录态中读取用户 ID。
//
// Session 读取依赖 gin,属于接入层职责,因此这里通过接入层桥接函数
// currentUserIDFromRequestContext(见 middleware.go)取值,Service 层本身不感知 gin。
func (s *authServiceImpl) GetCurrentUserID(ctx context.Context) (uint64, error) {
userID, ok := currentUserIDFromRequestContext(ctx)
if !ok {
return 0, errors.New(errUserNotInContext)
}
return userID, nil
}
func (s *authServiceImpl) RevokeToken(ctx context.Context, tokenHash string) error {
InvalidateCachedToken(ctx, tokenHash)
return nil
}
func (s *authServiceImpl) DisallowTokenAuthMiddleware() any {
return DisallowTokenAuth()
}
func (s *authServiceImpl) InvalidateCachedUser(ctx context.Context, userID uint64) {
InvalidateCachedUser(ctx, userID)
}
func (s *authServiceImpl) InvalidateCachedToken(ctx context.Context, tokenHash string) {
InvalidateCachedToken(ctx, tokenHash)
}
func (s *authServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
sources, err := ListAllAuthSources(ctx)
if err != nil {
return nil, err
}
views := make([]contracts.AuthSourceViewDTO, len(sources))
for i := range sources {
views[i] = contracts.AuthSourceViewDTO{
ID: sources[i].ID,
Name: sources[i].Name,
Type: sources[i].Type,
DisplayName: sources[i].DisplayName,
IsActive: sources[i].IsActive,
IconURL: sources[i].IconURL,
ClientSecretConfigured: sources[i].ClientSecret != "",
}
}
return views, nil
}
func (s *authServiceImpl) CreateAuthSource(ctx context.Context, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
model := AuthSource{
ID: source.ID,
Name: source.Name,
Type: source.Type,
DisplayName: source.DisplayName,
ClientID: source.ClientID,
ClientSecret: source.ClientSecret,
OpenIDDiscoveryURL: source.OpenIDDiscoveryURL,
Scopes: source.Scopes,
IconURL: source.IconURL,
IsActive: source.IsActive,
}
if err := model.Validate(); err != nil {
return nil, err
}
if err := CreateAuthSourceRecord(ctx, &model); err != nil {
return nil, err
}
model.Sanitize()
return toAuthSourceDTO(&model), nil
}
func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
existing, err := GetAuthSourceByID(ctx, id)
if err != nil {
return nil, err
}
existing.DisplayName = source.DisplayName
existing.ClientID = source.ClientID
if source.ClientSecret != "" {
existing.ClientSecret = source.ClientSecret
}
existing.OpenIDDiscoveryURL = source.OpenIDDiscoveryURL
existing.Scopes = source.Scopes
existing.IconURL = source.IconURL
if err := existing.Validate(); err != nil {
return nil, err
}
if err := SaveAuthSourceRecord(ctx, existing); err != nil {
return nil, err
}
existing.Sanitize()
return toAuthSourceDTO(existing), nil
}
func (s *authServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error {
existing, err := GetAuthSourceByID(ctx, id)
if err != nil {
return err
}
return DeleteAuthSourceRecord(ctx, existing)
}
func (s *authServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
existing, err := GetAuthSourceByID(ctx, id)
if err != nil {
return nil, err
}
existing.IsActive = !existing.IsActive
if err := SaveAuthSourceRecord(ctx, existing); err != nil {
return nil, err
}
existing.Sanitize()
return toAuthSourceDTO(existing), nil
}
func toAuthSourceDTO(s *AuthSource) *contracts.AuthSourceDTO {
if s == nil {
return nil
}
return &contracts.AuthSourceDTO{
ID: s.ID,
Name: s.Name,
Type: s.Type,
DisplayName: s.DisplayName,
ClientID: s.ClientID,
ClientSecret: s.ClientSecret,
OpenIDDiscoveryURL: s.OpenIDDiscoveryURL,
Scopes: s.Scopes,
IconURL: s.IconURL,
IsActive: s.IsActive,
CreatedAt: s.CreatedAt,
UpdatedAt: s.UpdatedAt,
}
}
type authRegistryImpl struct {
mu sync.RWMutex
providers map[string]contracts.OAuthProvider
}
func newAuthRegistry() contracts.AuthRegistry {
return &authRegistryImpl{
providers: make(map[string]contracts.OAuthProvider),
}
}
func (r *authRegistryImpl) RegisterOAuthProvider(name string, provider contracts.OAuthProvider) {
r.mu.Lock()
defer r.mu.Unlock()
r.providers[name] = provider
}
func (r *authRegistryImpl) GetOAuthProvider(name string) (contracts.OAuthProvider, bool) {
r.mu.RLock()
defer r.mu.RUnlock()
p, ok := r.providers[name]
return p, ok
}
func (r *authRegistryImpl) ListOAuthProviders() []string {
r.mu.RLock()
defer r.mu.RUnlock()
res := make([]string, 0, len(r.providers))
for name := range r.providers {
res = append(res, name)
}
return res
}
+401
View File
@@ -0,0 +1,401 @@
// 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))
})
}
})
}
func TestAuthWhitelistMiddleware(t *testing.T) {
gin.SetMode(gin.TestMode)
ctx := core.NewContext(context.Background())
p := auth.New()
require.NoError(t, p.Apply(ctx))
svc, err := core.Inject[contracts.AuthService](ctx)
require.NoError(t, err)
mw, ok := svc.RequireAuthMiddleware().(gin.HandlerFunc)
require.True(t, ok)
engine := newSessionEngine()
engine.Use(mw)
engine.POST("/api/v1/user/login", func(c *gin.Context) {
c.JSON(http.StatusOK, response.OK("login-ok"))
})
engine.GET("/api/v1/secret-profile", func(c *gin.Context) {
c.JSON(http.StatusOK, response.OK("profile-ok"))
})
// 1. Whitelisted route /api/v1/user/login passes through without auth
w1 := httptest.NewRecorder()
req1, _ := http.NewRequest(http.MethodPost, "/api/v1/user/login", nil)
engine.ServeHTTP(w1, req1)
assert.Equal(t, http.StatusOK, w1.Code)
// 2. Non-whitelisted route /api/v1/secret-profile gets 401 Unauthorized
w2 := httptest.NewRecorder()
req2, _ := http.NewRequest(http.MethodGet, "/api/v1/secret-profile", nil)
engine.ServeHTTP(w2, req2)
assert.Equal(t, http.StatusUnauthorized, w2.Code)
}
+169
View File
@@ -0,0 +1,169 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"context"
"crypto/sha256"
"encoding/hex"
"net/http"
"strconv"
"strings"
"sync"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
gsessions "github.com/gorilla/sessions"
)
var (
sessConfigMu sync.RWMutex
sessConfig = SessionConfig{
SessionCookieName: "wavelet_session",
SessionAge: 86400,
SessionHTTPOnly: true,
}
)
// SetSessionConfig updates the active session configuration.
func SetSessionConfig(cfg SessionConfig) {
sessConfigMu.Lock()
defer sessConfigMu.Unlock()
sessConfig = cfg
}
// GetSessionConfig returns the active session configuration.
func GetSessionConfig() SessionConfig {
sessConfigMu.RLock()
defer sessConfigMu.RUnlock()
return sessConfig
}
// GetSessionOptions 根据配置构建 Session 选项
func GetSessionOptions(maxAge int) sessions.Options {
cfg := GetSessionConfig()
return sessions.Options{
Path: "/",
Domain: cfg.SessionDomain,
MaxAge: maxAge,
HttpOnly: cfg.SessionHTTPOnly,
Secure: cfg.SessionSecure,
SameSite: http.SameSiteLaxMode,
}
}
// StripCookieMaxAgeAndExpires 从 Set-Cookie 响应头中移除 Max-Age 和 Expires,从而使其成为浏览器会话 Cookie
func StripCookieMaxAgeAndExpires(header http.Header, cookieName string) {
headers := header["Set-Cookie"]
if len(headers) == 0 {
return
}
newHeaders := make([]string, 0, len(headers))
for _, h := range headers {
if strings.HasPrefix(h, cookieName+"=") {
parts := strings.Split(h, ";")
newParts := make([]string, 0, len(parts))
for _, p := range parts {
trimmed := strings.TrimSpace(p)
lower := strings.ToLower(trimmed)
if strings.HasPrefix(lower, "max-age=") || strings.HasPrefix(lower, "expires=") {
continue
}
newParts = append(newParts, p)
}
newHeaders = append(newHeaders, strings.Join(newParts, ";"))
} else {
newHeaders = append(newHeaders, h)
}
}
header["Set-Cookie"] = newHeaders
}
// GetUserIDFromSession 从 Session 中提取用户 ID
func GetUserIDFromSession(s sessions.Session) uint64 {
val := s.Get(UserIDKey)
return ParseUserID(val)
}
// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID
func GetUserIDFromContext(c *gin.Context) (uid uint64) {
defer func() {
_ = recover()
}()
session := sessions.Default(c)
return GetUserIDFromSession(session)
}
func ensureSessionToken(s sessions.Session) (string, bool) {
token, ok := s.Get(SessionTokenKey).(string)
if !ok || token == "" {
token = uuid.NewString()
s.Set(SessionTokenKey, token)
return token, true
}
return token, false
}
func hashSessionToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func rotateSessionID(s sessions.Session) {
if inner, ok := s.(interface{ Session() *gsessions.Session }); ok {
if sess := inner.Session(); sess != nil {
sess.ID = ""
}
}
}
// SetLoginSession writes the authenticated user into a freshly rotated session.
func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDTO, extras ...map[string]any) error {
session := sessions.Default(c)
session.Clear()
rotateSessionID(session)
session.Set(UserIDKey, user.ID)
session.Set(UserNameKey, user.Username)
if len(extras) > 0 {
for key, value := range extras[0] {
session.Set(key, value)
}
}
// 根据系统配置动态设置 Session 过期时间
cfg := GetSessionConfig()
maxAge := cfg.SessionAge
isSessionCookie := false
val, err := GetSystemConfigValue(ctx, "login_session_ttl_hours")
if err == nil && val != "" {
if ttlHours, err := strconv.Atoi(val); err == nil {
switch {
case ttlHours == -1:
// 永不过期,设置为 10 年
maxAge = 10 * 365 * 24 * 3600
case ttlHours > 0:
maxAge = ttlHours * 3600
case ttlHours == 0:
isSessionCookie = true
}
}
}
session.Options(GetSessionOptions(maxAge))
if err := session.Save(); err != nil {
return err
}
if isSessionCookie {
StripCookieMaxAgeAndExpires(c.Writer.Header(), cfg.SessionCookieName)
}
return nil
}
@@ -0,0 +1,219 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package domain_test
import (
"context"
"net/http"
"net/http/httptest"
"reflect"
"strings"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/plugins/domain/admin"
"Wavelet/plugins/domain/message_gateway"
"Wavelet/plugins/domain/user"
)
// The kernel only gates a plugin's Apply on the services it DECLARES in Inject,
// so a plugin that consumes contracts.AuthService without declaring it can be
// mounted before auth exists. Those plugins fall back to a pass-through
// middleware, which silently un-guards their routes.
func sentinelLogin(c *gin.Context) { c.Next() }
func sentinelAdmin(c *gin.Context) { c.Next() }
func sentinelNoToken(c *gin.Context) { c.Next() }
type stubAuthService struct{ contracts.AuthService }
func (stubAuthService) RequireAuthMiddleware() any { return gin.HandlerFunc(sentinelLogin) }
func (stubAuthService) RequireAdminMiddleware() any { return gin.HandlerFunc(sentinelAdmin) }
func (stubAuthService) DisallowTokenAuthMiddleware() any { return gin.HandlerFunc(sentinelNoToken) }
type stubDBService struct{ contracts.DBService }
// providerPlugin publishes a contract into the container at Apply time, the way
// the real infra and domain plugins do.
type providerPlugin struct {
name string
provide func(*core.Context) error
}
func (p providerPlugin) Name() string { return p.name }
func (p providerPlugin) Apply(ctx *core.Context) error { return p.provide(ctx) }
func dbProvider() core.Plugin {
return providerPlugin{name: "stub-database", provide: func(ctx *core.Context) error {
core.Provide[contracts.DBService](ctx, stubDBService{})
return nil
}}
}
func authProvider() core.Plugin {
return providerPlugin{name: "stub-auth", provide: func(ctx *core.Context) error {
core.Provide[contracts.AuthService](ctx, stubAuthService{})
return nil
}}
}
func findRoute(routes []extpoints.RouteDefinition, method, path string) (extpoints.RouteDefinition, bool) {
for _, rd := range routes {
if rd.Method == method && rd.Path == path {
return rd, true
}
}
return extpoints.RouteDefinition{}, false
}
// codePointer resolves the function code pointer of a registered handler so
// identity can be compared without depending on gin internals.
func codePointer(handler any) uintptr {
switch fn := handler.(type) {
case gin.HandlerFunc:
return reflect.ValueOf(fn).Pointer()
case func(*gin.Context):
return reflect.ValueOf(fn).Pointer()
default:
return 0
}
}
func assertsMiddleware(t *testing.T, routes []extpoints.RouteDefinition, method, path string, want func(*gin.Context)) {
t.Helper()
rd, ok := findRoute(routes, method, path)
if !ok {
t.Fatalf("route %s %s was never registered", method, path)
}
wantPtr := codePointer(gin.HandlerFunc(want))
for _, h := range rd.Handlers {
if codePointer(h) == wantPtr {
return
}
}
for _, m := range rd.Middlewares {
if codePointer(m) == wantPtr {
return
}
}
t.Errorf("route %s %s is not guarded by the auth middleware: handlers=%d middlewares=%d",
method, path, len(rd.Handlers), len(rd.Middlewares))
}
// TestRoutesMountedBeforeAuthServiceAreGuarded 回归:cmd/app.go 把 user/message_gateway
// 注册在 auth 之前,若插件未在 Inject 中声明 contracts.AuthService,reconcile 会先
// Apply 它们,导致鉴权中间件退化为透传闭包,路由完全不受保护。
func TestRoutesMountedBeforeAuthServiceAreGuarded(t *testing.T) {
gin.SetMode(gin.TestMode)
ctx := core.NewContext(context.Background())
// Registration order mirrors cmd/app.go: the auth provider comes last.
app := core.NewApp(
core.WithContext(ctx),
core.WithPlugins(
dbProvider(),
user.New(),
message_gateway.New(),
authProvider(),
),
)
require.NoError(t, app.ApplyPlugins())
routes := ctx.Router().Routes()
assertsMiddleware(t, routes, "POST", "/api/v1/user/change-password", sentinelLogin)
assertsMiddleware(t, routes, "PUT", "/api/v1/user/profile", sentinelLogin)
assertsMiddleware(t, routes, "GET", "/api/v1/user/access-tokens", sentinelLogin)
assertsMiddleware(t, routes, "GET", "/api/v1/user/access-tokens", sentinelNoToken)
assertsMiddleware(t, routes, "GET", "/api/v1/message-gateway/channels", sentinelLogin)
assertsMiddleware(t, routes, "GET", "/api/v1/admin/message-gateway/channels", sentinelAdmin)
}
// TestAuthConsumersDeclareAuthDependency 是同一缺陷的架构面:声明依赖是内核排序的唯一
// 依据,漏声明会让正确性取决于注册表顺序。
func TestAuthConsumersDeclareAuthDependency(t *testing.T) {
want := reflect.TypeFor[contracts.AuthService]()
for _, tc := range []struct {
name string
deps []reflect.Type
}{
{"user", user.New().Inject()},
{"message_gateway", message_gateway.New().Inject()},
} {
t.Run(tc.name, func(t *testing.T) {
assert.Contains(t, tc.deps, want,
"%s resolves contracts.AuthService in Apply, so it must declare it in Inject", tc.name)
})
}
}
// firstGuard returns the outermost middleware registered for the first route
// matching method and path prefix — the auth guard the plugin resolved at Apply.
func firstGuard(t *testing.T, routes []extpoints.RouteDefinition, method, prefix string) gin.HandlerFunc {
t.Helper()
for _, rd := range routes {
if rd.Method != method || !strings.HasPrefix(rd.Path, prefix) {
continue
}
for _, candidate := range rd.Middlewares {
if mw, ok := candidate.(gin.HandlerFunc); ok {
return mw
}
}
for _, candidate := range rd.Handlers {
if mw, ok := candidate.(gin.HandlerFunc); ok {
return mw
}
}
t.Fatalf("route %s %s has no inspectable guard", method, rd.Path)
}
t.Fatalf("no route registered for %s %s*", method, prefix)
return nil
}
// TestAuthGuardFailsClosed 回归:鉴权服务无法解析时兜底必须是拒绝。旧实现兜底为
// c.Next(),所以任何装配缺失——例如 admin 的 OnDispose 调用 service.ResetServices
// 把全局 authService 置 nil——都会让路由以“已登录”的姿态直达业务处理函数。
func TestAuthGuardFailsClosed(t *testing.T) {
gin.SetMode(gin.TestMode)
tests := []struct {
name string
apply func(*core.Context) error
method string
prefix string
}{
{"user", user.New().Apply, http.MethodPost, "/api/v1/user/change-password"},
{"message_gateway", message_gateway.New().Apply, http.MethodGet, "/api/v1/message-gateway"},
{"admin", admin.New().Apply, http.MethodGet, "/api/v1/admin"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
ctx := core.NewContext(context.Background())
require.NoError(t, tc.apply(ctx))
guard := firstGuard(t, ctx.Router().Routes(), tc.method, tc.prefix)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(tc.method, tc.prefix, nil)
guard(c)
assert.True(t, c.IsAborted(),
"%s guard must reject the request when contracts.AuthService is unavailable", tc.name)
assert.NotEmpty(t, c.Errors,
"%s guard must record why the request was rejected", tc.name)
})
}
}
+24
View File
@@ -0,0 +1,24 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// 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
)
+87
View File
@@ -0,0 +1,87 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"net/http"
"github.com/gin-gonic/gin"
)
// Challenge 生成 PoW 人机验证难题
// @Summary 生成人机验证难题
// @Description 客户端获取 PoW 难题和签名的 JWT Token,并在后台计算。
// @Tags cap
// @Accept json
// @Produce json
// @Param request body challengeRequest false "可选范围限制参数"
// @Success 200 {object} response.Any{data=cap.ChallengeResponse} "成功返回 PoW 难题"
// @Failure 500 {object} response.Any "内部服务错误"
// @Router /api/cap/challenge [post]
func Challenge(c *gin.Context) {
var req challengeRequest
_ = c.ShouldBind(&req) // 允许不传 body,默认使用 login scope
if req.Scope == "" {
req.Scope = "login"
}
mgr := GetDefaultManager()
if mgr == nil {
response.AbortInternal(c, 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, errChallengeGenerateFailed)
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
// Redeem 提交 PoW 解答并兑换一次性凭证 Token
// @Summary 校验人机验证解答
// @Description 提交 PoW 解答进行核销,成功后返回一次性 X-Cap-Token 凭证
// @Tags cap
// @Accept json
// @Produce json
// @Param request body redeemRequest true "难题 Token 与解答 solutions 数组"
// @Success 200 {object} response.Any{data=cap.RedeemResponse} "核销成功,返回 X-Cap-Token"
// @Failure 400 {object} response.Any "参数错误或核销失败"
// @Failure 500 {object} response.Any "内部服务错误"
// @Router /api/cap/redeem [post]
func Redeem(c *gin.Context) {
var req redeemRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, errInvalidRequestParams)
return
}
if req.Scope == "" {
req.Scope = "login"
}
mgr := GetDefaultManager()
if mgr == nil {
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, errSolutionVerifyFailed)
return
}
if !resp.Success {
response.AbortBadRequest(c, resp.Error)
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
+38
View File
@@ -0,0 +1,38 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
// VerifyMiddleware returns a Gin middleware that checks and consumes the X-Cap-Token header.
func VerifyMiddleware(mgr *Manager, scope string) gin.HandlerFunc {
return func(c *gin.Context) {
if !ProtectionEnabled(c.Request.Context()) {
c.Next()
return
}
if mgr == nil {
response.AbortUnauthorized(c, errCapTokenInvalidOrExpired)
return
}
token := c.GetHeader("X-Cap-Token")
if token == "" {
response.AbortUnauthorized(c, errCapTokenMissing)
return
}
valid, err := mgr.VerifyToken(c.Request.Context(), token, scope)
if err != nil || !valid {
response.AbortUnauthorized(c, errCapTokenInvalidOrExpired)
return
}
c.Next()
}
}
+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"`
}
+105
View File
@@ -0,0 +1,105 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package cap provides the proof-of-work (PoW) CAPTCHA verification domain plugin for Cordis.
package cap
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"reflect"
)
// Plugin implements core.Plugin to provide CAPTCHA generation, validation, and route protection.
type Plugin struct{}
// New creates a new cap domain plugin.
func New() *Plugin {
return &Plugin{}
}
// Name returns the unique identifier for the cap domain plugin.
func (p *Plugin) Name() string {
return "cap"
}
// Inject declares required dependencies for the cap domain plugin.
func (p *Plugin) Inject() []reflect.Type {
return []reflect.Type{
reflect.TypeFor[contracts.DBService](),
}
}
// Manifest returns the plugin metadata.
func (p *Plugin) Manifest() core.Manifest {
return core.Manifest{
Name: "cap",
Version: "1.0.0",
Description: "Proof-of-work CAPTCHA challenge and verification domain plugin",
Author: "Wavelet Team",
}
}
type capAppConfig struct {
SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"`
}
// DeclareConfig declares configuration bindings for the cap plugin.
func (p *Plugin) DeclareConfig() []core.ConfigBinding {
return []core.ConfigBinding{
{Prefix: "app", Target: &capAppConfig{}},
}
}
// Apply registers the cap routes and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
var cfg capAppConfig
if err := ctx.Config().Bind("app", &cfg); err == nil && cfg.SessionSecret != "" {
SetSecret([]byte(cfg.SessionSecret))
}
// 0. Bind DBService from Context
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
setDBService(db)
} else {
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
setDBService(db)
})
}
ctx.OnDispose(func() error {
setDBService(nil)
return nil
})
// Listen to system config changed events to invalidate cached settings
ctx.Events().On(contracts.EventTopicConfigChanged, func(_ any) {
InvalidateRuntimeSettings()
})
// Register HTTP Routes
capGroup := ctx.Router().Group("/api/v1/cap")
{
capGroup.GET("/challenge", Challenge)
capGroup.POST("/challenge", Challenge)
capGroup.POST("/redeem", Redeem)
}
// Register Settings Schemas
ctx.Settings().Register(extpoints.SettingSchema{
Key: "cap.login_enabled",
Default: false,
Description: "Whether to require CAPTCHA verification for user login",
Type: "boolean",
Category: "security",
})
ctx.Settings().Register(extpoints.SettingSchema{
Key: "cap.challenge_count",
Default: 1,
Description: "Number of PoW puzzle challenges to solve",
Type: "integer",
Category: "security",
})
return nil
}
+259
View File
@@ -0,0 +1,259 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package pow provides proof-of-work challenge generation and verification.
package pow
import (
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"encoding/json"
"errors"
"strconv"
"strings"
"time"
)
const (
jwtHeaderB64 = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9"
jwtPartsCount = 3 // JWT 三段结构
defaultChallengeCount = 50 // 默认 PoW 难题数
defaultChallengeSize = 32 // 默认盐值长度
defaultDifficulty = 4 // 默认难度
defaultNonceLength = 25 // 随机 Nonce 字节长度
defaultExpires = 10 * time.Minute // 默认过期时间
)
// ChallengeConfig holds parameters for the PoW challenge
type ChallengeConfig struct {
Count int // Number of puzzles (c)
Size int // Salt length (s)
Difficulty int // Difficulty prefix length (d)
Expires time.Duration // Challenge TTL
}
// ChallengeResponse is returned to the client
type ChallengeResponse struct {
Challenge struct {
C int `json:"c"`
S int `json:"s"`
D int `json:"d"`
} `json:"challenge"`
Token string `json:"token"`
Expires int64 `json:"expires"` // ms timestamp
}
// ChallengePayload represents the signed JWT payload
type ChallengePayload struct {
Nonce string `json:"n"`
Count int `json:"c"`
Size int `json:"s"`
Difficulty int `json:"d"`
Expires int64 `json:"exp"` // ms timestamp
IssuedAt int64 `json:"iat"` // ms timestamp
Scope string `json:"sk,omitempty"`
}
// RedeemRequest payload sent by client
type RedeemRequest struct {
Token string `json:"token"`
Solutions []int `json:"solutions"`
}
// RedeemResponse returned to client after verification
type RedeemResponse struct {
Success bool `json:"success"`
Token string `json:"token,omitempty"`
Expires int64 `json:"expires,omitempty"`
Error string `json:"error,omitempty"`
}
func b64urlEncode(data []byte) string {
return base64.RawURLEncoding.EncodeToString(data)
}
func b64urlDecode(str string) ([]byte, error) {
return base64.RawURLEncoding.DecodeString(str)
}
// RandomHex generates a cryptographically secure random hexadecimal string of the specified byte length.
func RandomHex(byteLen int) string {
bytes := make([]byte, byteLen)
if _, err := rand.Read(bytes); err != nil {
panic(err)
}
return hex.EncodeToString(bytes)
}
func jwtSign(payload, secret []byte) string {
body := b64urlEncode(payload)
sigInput := jwtHeaderB64 + "." + body
mac := hmac.New(sha256.New, secret)
mac.Write([]byte(sigInput))
sig := mac.Sum(nil)
return sigInput + "." + b64urlEncode(sig)
}
func jwtVerify(token string, secret []byte) ([]byte, error) {
parts := strings.Split(token, ".")
if len(parts) != jwtPartsCount {
return nil, errors.New(errInvalidTokenFormat)
}
if parts[0] != jwtHeaderB64 {
return nil, errors.New(errInvalidHeader)
}
sigInput := parts[0] + "." + parts[1]
mac := hmac.New(sha256.New, secret)
mac.Write([]byte(sigInput))
expectedSig := mac.Sum(nil)
actualSig, err := b64urlDecode(parts[2])
if err != nil {
return nil, err
}
if !hmac.Equal(expectedSig, actualSig) {
return nil, errors.New(errSignatureMismatch)
}
payload, err := b64urlDecode(parts[1])
if err != nil {
return nil, err
}
return payload, nil
}
// JwtSigHex extracts the signature part of a JWT token and returns it as a hexadecimal string.
func JwtSigHex(token string) string {
parts := strings.Split(token, ".")
if len(parts) != jwtPartsCount {
return ""
}
sigBytes, err := b64urlDecode(parts[2])
if err != nil {
return ""
}
return hex.EncodeToString(sigBytes)
}
// GenerateChallenge produces a new challenge and signed token
func GenerateChallenge(secret []byte, conf ChallengeConfig, scope string) (*ChallengeResponse, error) {
if conf.Count <= 0 {
conf.Count = defaultChallengeCount
}
if conf.Size <= 0 {
conf.Size = defaultChallengeSize
}
if conf.Difficulty <= 0 {
conf.Difficulty = defaultDifficulty
}
if conf.Expires <= 0 {
conf.Expires = defaultExpires
}
now := time.Now().UnixNano() / int64(time.Millisecond)
expires := now + int64(conf.Expires/time.Millisecond)
payload := ChallengePayload{
Nonce: RandomHex(defaultNonceLength),
Count: conf.Count,
Size: conf.Size,
Difficulty: conf.Difficulty,
Expires: expires,
IssuedAt: now,
Scope: scope,
}
payloadBytes, err := json.Marshal(payload)
if err != nil {
return nil, err
}
token := jwtSign(payloadBytes, secret)
resp := &ChallengeResponse{
Token: token,
Expires: expires,
}
resp.Challenge.C = conf.Count
resp.Challenge.S = conf.Size
resp.Challenge.D = conf.Difficulty
return resp, nil
}
// VerifyChallengeSolutions verifies client submitted solutions
func VerifyChallengeSolutions(token string, solutions []int, secret []byte, expectedScope string) (*ChallengePayload, error) {
payloadBytes, err := jwtVerify(token, secret)
if err != nil {
return nil, errors.New(errInvalidToken)
}
var payload ChallengePayload
if err := json.Unmarshal(payloadBytes, &payload); err != nil {
return nil, errors.New(errInvalidToken)
}
if expectedScope != "" && payload.Scope != expectedScope {
return nil, errors.New(errScopeMismatch)
}
now := time.Now().UnixNano() / int64(time.Millisecond)
if payload.Expires < now {
return nil, errors.New(errExpired)
}
if len(solutions) != payload.Count {
return nil, errors.New(errInvalidSolutions)
}
tokenFnv := fnv1a(token)
for i := 0; i < payload.Count; i++ {
idxStr := strconv.Itoa(i + 1)
saltSeed := fnv1aResume(tokenFnv, idxStr)
targetSeed := fnv1aResume(saltSeed, "d")
salt := prngFromHash(saltSeed, payload.Size)
target := prngFromHash(targetSeed, payload.Difficulty)
hashInput := salt + strconv.Itoa(solutions[i])
hashBytes := sha256.Sum256([]byte(hashInput))
hashHex := hex.EncodeToString(hashBytes[:])
if !strings.HasPrefix(hashHex, target) {
return nil, errors.New(errInvalidSolution)
}
}
return &payload, nil
}
// Solve is a utility function to solve a challenge (mainly used for tests and reference implementation)
func Solve(token string, count, size, difficulty int) []int {
solutions := make([]int, count)
tokenFnv := fnv1a(token)
for i := 0; i < count; i++ {
idxStr := strconv.Itoa(i + 1)
saltSeed := fnv1aResume(tokenFnv, idxStr)
targetSeed := fnv1aResume(saltSeed, "d")
salt := prngFromHash(saltSeed, size)
target := prngFromHash(targetSeed, difficulty)
for nonce := 0; nonce < 1000000; nonce++ {
hashInput := salt + strconv.Itoa(nonce)
hashBytes := sha256.Sum256([]byte(hashInput))
hashHex := hex.EncodeToString(hashBytes[:])
if strings.HasPrefix(hashHex, target) {
solutions[i] = nonce
break
}
}
}
return solutions
}
+15
View File
@@ -0,0 +1,15 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pow
const (
errInvalidTokenFormat = "invalid token format"
errInvalidHeader = "invalid header"
errSignatureMismatch = "signature mismatch"
errInvalidToken = "invalid_token"
errScopeMismatch = "scope_mismatch"
errExpired = "expired"
errInvalidSolutions = "invalid_solutions"
errInvalidSolution = "invalid_solution"
)
+111
View File
@@ -0,0 +1,111 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pow
import (
"context"
"testing"
"time"
)
func TestPowChallengeFlow(t *testing.T) {
secret := []byte("test-secret-key-1234567890123456")
conf := ChallengeConfig{
Count: 2,
Size: 16,
Difficulty: 1,
Expires: 1 * time.Minute,
}
scope := "login"
resp, err := GenerateChallenge(secret, conf, scope)
if err != nil {
t.Fatalf("GenerateChallenge failed: %v", err)
}
if resp.Token == "" {
t.Fatal("expected non-empty token")
}
if resp.Challenge.C != 2 {
t.Fatalf("expected count 2, got %d", resp.Challenge.C)
}
sigHex := JwtSigHex(resp.Token)
if sigHex == "" {
t.Fatal("expected non-empty sigHex")
}
solutions := Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
if len(solutions) != 2 {
t.Fatalf("expected 2 solutions, got %d", len(solutions))
}
payload, err := VerifyChallengeSolutions(resp.Token, solutions, secret, scope)
if err != nil {
t.Fatalf("VerifyChallengeSolutions failed: %v", err)
}
if payload.Scope != scope {
t.Fatalf("expected scope %s, got %s", scope, payload.Scope)
}
// Scope mismatch test
_, err = VerifyChallengeSolutions(resp.Token, solutions, secret, "other_scope")
if err == nil {
t.Fatal("expected scope mismatch error")
}
// Invalid solutions test
_, err = VerifyChallengeSolutions(resp.Token, []int{9999999, 9999999}, secret, scope)
if err == nil {
t.Fatal("expected invalid solution error")
}
// Invalid token test
_, err = VerifyChallengeSolutions("invalid.jwt.token", solutions, secret, scope)
if err == nil {
t.Fatal("expected invalid token error")
}
}
func TestMemoryStore(t *testing.T) {
ctx := context.Background()
store := NewMemoryStore(100 * time.Millisecond)
// Set and Get
err := store.Set(ctx, "k1", "v1", 200*time.Millisecond)
if err != nil {
t.Fatalf("Set failed: %v", err)
}
val, ok, err := store.Get(ctx, "k1")
if err != nil || !ok || val != "v1" {
t.Fatalf("Get failed: val=%s, ok=%v, err=%v", val, ok, err)
}
// SetNX
set, err := store.SetNX(ctx, "k1", "v2", 200*time.Millisecond)
if err != nil || set {
t.Fatalf("SetNX should have failed because key exists: set=%v, err=%v", set, err)
}
set, err = store.SetNX(ctx, "k2", "v2", 200*time.Millisecond)
if err != nil || !set {
t.Fatalf("SetNX should have succeeded: set=%v, err=%v", set, err)
}
// GetAndDelete
val, ok, err = store.GetAndDelete(ctx, "k2")
if err != nil || !ok || val != "v2" {
t.Fatalf("GetAndDelete failed: val=%s, ok=%v, err=%v", val, ok, err)
}
_, ok, _ = store.Get(ctx, "k2")
if ok {
t.Fatal("k2 should be deleted")
}
// Delete
_ = store.Delete(ctx, "k1")
_, ok, _ = store.Get(ctx, "k1")
if ok {
t.Fatal("k1 should be deleted")
}
}
+54
View File
@@ -0,0 +1,54 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pow
import (
"fmt"
"strings"
)
// fnv1a returns the 32-bit FNV-1a hash of a string
//
// FNV-1a 算法位移常量
func fnv1a(str string) uint32 {
var hash uint32 = 2166136261
for i := 0; i < len(str); i++ {
hash ^= uint32(str[i])
hash += (hash << 1) + (hash << 4) + (hash << 7) + (hash << 8) + (hash << 24)
}
return hash
}
// fnv1aResume resumes FNV-1a hashing from a given state
//
// FNV-1a 算法位移常量
func fnv1aResume(state uint32, str string) uint32 {
h := state
for i := 0; i < len(str); i++ {
h ^= uint32(str[i])
h += (hashShift(h))
}
return h
}
// hashShift computes FNV-1a mix additions
func hashShift(h uint32) uint32 {
return (h << 1) + (h << 4) + (h << 7) + (h << 8) + (h << 24)
}
// prngFromHash generates a hex string of specified length using an initial hash state
//
// xorshift 算法位移常量
func prngFromHash(initialHash uint32, length int) string {
state := initialHash
var result strings.Builder
for result.Len() < length {
state ^= state << 13
state ^= state >> 17
state ^= state << 5
hexStr := fmt.Sprintf("%08x", state)
result.WriteString(hexStr)
}
return result.String()[:length]
}
+188
View File
@@ -0,0 +1,188 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pow
import (
"Wavelet/pkg/util"
"context"
"errors"
"sync"
"time"
"github.com/redis/go-redis/v9"
)
// Store defines the storage interface for challenge nonces and verification tokens
type Store interface {
Get(ctx context.Context, key string) (string, bool, error)
Set(ctx context.Context, key, val string, ttl time.Duration) error
Delete(ctx context.Context, key string) error
// SetNX atomically sets key=val with the given TTL only when the key does not
// exist yet. It returns true when the key was actually written (i.e. this
// caller "won" the race), and false when the key already existed.
SetNX(ctx context.Context, key, val string, ttl time.Duration) (bool, error)
// GetAndDelete atomically retrieves the value of key and removes it in a
// single operation. Returns ("", false, nil) when the key does not exist.
GetAndDelete(ctx context.Context, key string) (string, bool, error)
}
type memoryItem struct {
value string
expiresAt time.Time
}
// MemoryStore is a thread-safe in-memory implementation of Store
type MemoryStore struct {
items map[string]memoryItem
mu sync.Mutex // unified write-lock; promotes to exclusive for all ops
}
// NewMemoryStore creates and initializes a new MemoryStore
func NewMemoryStore(cleanupInterval time.Duration) *MemoryStore {
store := &MemoryStore{
items: make(map[string]memoryItem),
}
if cleanupInterval > 0 {
util.Go(func() { store.startCleanupLoop(cleanupInterval) })
}
return store
}
// Get 从 MemoryStore 获取指定 key 的值
func (s *MemoryStore) Get(_ context.Context, key string) (string, bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
return s.getLocked(key)
}
// getLocked is the internal helper – caller must hold s.mu.
func (s *MemoryStore) getLocked(key string) (string, bool, error) {
item, found := s.items[key]
if !found {
return "", false, nil
}
if time.Now().After(item.expiresAt) {
delete(s.items, key)
return "", false, nil
}
return item.value, true, nil
}
// Set 向 MemoryStore 写入指定 key 的值
func (s *MemoryStore) Set(_ context.Context, key, val string, ttl time.Duration) error {
s.mu.Lock()
defer s.mu.Unlock()
s.items[key] = memoryItem{
value: val,
expiresAt: time.Now().Add(ttl),
}
return nil
}
// Delete 从 MemoryStore 删除指定 key
func (s *MemoryStore) Delete(_ context.Context, key string) error {
s.mu.Lock()
defer s.mu.Unlock()
delete(s.items, key)
return nil
}
// SetNX atomically sets key only when it is absent (or expired).
// Returns true if the key was written by this call.
func (s *MemoryStore) SetNX(_ context.Context, key, val string, ttl time.Duration) (bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
_, exists, _ := s.getLocked(key)
if exists {
return false, nil
}
s.items[key] = memoryItem{
value: val,
expiresAt: time.Now().Add(ttl),
}
return true, nil
}
// GetAndDelete atomically retrieves and removes key in one critical section.
func (s *MemoryStore) GetAndDelete(_ context.Context, key string) (string, bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
val, exists, err := s.getLocked(key)
if err != nil || !exists {
return "", false, err
}
delete(s.items, key)
return val, true, nil
}
func (s *MemoryStore) startCleanupLoop(interval time.Duration) {
ticker := time.NewTicker(interval)
for range ticker.C {
s.cleanupExpired()
}
}
func (s *MemoryStore) cleanupExpired() {
now := time.Now()
s.mu.Lock()
defer s.mu.Unlock()
for k, v := range s.items {
if now.After(v.expiresAt) {
delete(s.items, k)
}
}
}
// RedisStore is a GORM-compatible/standalone Redis-backed implementation of Store
type RedisStore struct {
client redis.UniversalClient
}
// NewRedisStore creates a new RedisStore wrapping a redis.UniversalClient
func NewRedisStore(client redis.UniversalClient) *RedisStore {
return &RedisStore{
client: client,
}
}
// Get 从 RedisStore 获取指定 key 的值
func (s *RedisStore) Get(ctx context.Context, key string) (string, bool, error) {
val, err := s.client.Get(ctx, key).Result()
if errors.Is(err, redis.Nil) {
return "", false, nil
}
if err != nil {
return "", false, err
}
return val, true, nil
}
// Set 向 RedisStore 写入指定 key 的值
func (s *RedisStore) Set(ctx context.Context, key, val string, ttl time.Duration) error {
return s.client.Set(ctx, key, val, ttl).Err()
}
// Delete 从 RedisStore 删除指定 key
func (s *RedisStore) Delete(ctx context.Context, key string) error {
return s.client.Del(ctx, key).Err()
}
// SetNX wraps Redis SET NX – returns true only when the key was newly created.
func (s *RedisStore) SetNX(ctx context.Context, key, val string, ttl time.Duration) (bool, error) {
return s.client.SetNX(ctx, key, val, ttl).Result()
}
// GetAndDelete wraps Redis GETDEL (available since Redis 6.2).
func (s *RedisStore) GetAndDelete(ctx context.Context, key string) (string, bool, error) {
val, err := s.client.GetDel(ctx, key).Result()
if errors.Is(err, redis.Nil) {
return "", false, nil
}
if err != nil {
return "", false, err
}
return val, true, nil
}
+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
}
@@ -0,0 +1,185 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"context"
"errors"
"strconv"
"sync/atomic"
"time"
"golang.org/x/sync/singleflight"
)
const (
defaultChallengeCount = 1
defaultChallengeSize = 32
defaultChallengeDifficulty = 4
defaultChallengeTTL = 10 * time.Minute
defaultTokenTTL = 20 * time.Minute
)
// RuntimeSettings is the parsed CAPTCHA runtime configuration loaded from system_configs.
type RuntimeSettings struct {
LoginEnabled bool
ChallengeCount int
ChallengeSize int
ChallengeDifficulty int
ChallengeTTL time.Duration
TokenTTL time.Duration
}
// CAP 动态配置键常量
const (
ConfigKeyCapLoginEnabled = "cap_login_enabled"
ConfigKeyCapChallengeCount = "cap_challenge_count"
ConfigKeyCapChallengeSize = "cap_challenge_size"
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty"
ConfigKeyCapChallengeTTL = "cap_challenge_ttl"
// ConfigKeyCapTokenTTL 验证码 Token 过期时间键
// #nosec G101
ConfigKeyCapTokenTTL = "cap_token_ttl"
)
var runtimeConfigKeys = []string{
ConfigKeyCapLoginEnabled,
ConfigKeyCapChallengeCount,
ConfigKeyCapChallengeSize,
ConfigKeyCapChallengeDifficulty,
ConfigKeyCapChallengeTTL,
ConfigKeyCapTokenTTL,
}
var runtimeConfigKeySet = func() map[string]struct{} {
set := make(map[string]struct{}, len(runtimeConfigKeys))
for _, key := range runtimeConfigKeys {
set[key] = struct{}{}
}
return set
}()
type runtimeSettingsStore struct {
snapshot atomic.Pointer[RuntimeSettings]
loadGroup singleflight.Group
}
var settingsStore = &runtimeSettingsStore{}
// IsRuntimeConfigKey reports whether a system config key affects CAPTCHA runtime settings.
func IsRuntimeConfigKey(key string) bool {
_, ok := runtimeConfigKeySet[key]
return ok
}
// CurrentSettings returns the cached CAPTCHA runtime settings snapshot.
func CurrentSettings(ctx context.Context) (RuntimeSettings, error) {
return settingsStore.current(ctx)
}
// ProtectionEnabled reports whether CAPTCHA verification is required for protected routes.
func ProtectionEnabled(ctx context.Context) bool {
settings, err := CurrentSettings(ctx)
if err != nil {
return false
}
return settings.LoginEnabled
}
// InvalidateRuntimeSettings drops the in-process CAPTCHA settings snapshot.
func InvalidateRuntimeSettings() {
settingsStore.snapshot.Store(nil)
}
// ResetRuntimeSettingsForTest clears the CAPTCHA runtime snapshot.
func ResetRuntimeSettingsForTest() {
InvalidateRuntimeSettings()
}
// InstallTestRuntimeSettings installs a fixed snapshot for unit tests.
func InstallTestRuntimeSettings(settings RuntimeSettings) func() {
snapshot := settings
settingsStore.snapshot.Store(&snapshot)
return InvalidateRuntimeSettings
}
func (s *runtimeSettingsStore) current(ctx context.Context) (RuntimeSettings, error) {
s.ensureInvalidationListener()
if snapshot := s.snapshot.Load(); snapshot != nil {
return *snapshot, nil
}
loaded, err, _ := s.loadGroup.Do("cap-runtime-settings", func() (any, error) {
if snapshot := s.snapshot.Load(); snapshot != nil {
return *snapshot, nil
}
settings, loadErr := loadRuntimeSettings(ctx)
if loadErr != nil {
return RuntimeSettings{}, loadErr
}
s.snapshot.Store(&settings)
return settings, nil
})
if err != nil {
return RuntimeSettings{}, err
}
settings, ok := loaded.(RuntimeSettings)
if !ok {
return RuntimeSettings{}, errors.New("cap runtime settings loader returned unexpected type")
}
return settings, nil
}
func parseRuntimeSettings(configs map[string]string) RuntimeSettings {
settings := RuntimeSettings{
ChallengeCount: defaultChallengeCount,
ChallengeSize: defaultChallengeSize,
ChallengeDifficulty: defaultChallengeDifficulty,
ChallengeTTL: defaultChallengeTTL,
TokenTTL: defaultTokenTTL,
}
if len(configs) == 0 {
return settings
}
if val, ok := configs[ConfigKeyCapLoginEnabled]; ok {
if enabled, err := strconv.ParseBool(val); err == nil {
settings.LoginEnabled = enabled
}
}
if val, ok := configs[ConfigKeyCapChallengeCount]; ok {
if count, err := strconv.Atoi(val); err == nil && count > 0 {
settings.ChallengeCount = count
}
}
if val, ok := configs[ConfigKeyCapChallengeSize]; ok {
if size, err := strconv.Atoi(val); err == nil && size > 0 {
settings.ChallengeSize = size
}
}
if val, ok := configs[ConfigKeyCapChallengeDifficulty]; ok {
if diff, err := strconv.Atoi(val); err == nil && diff > 0 {
settings.ChallengeDifficulty = diff
}
}
if val, ok := configs[ConfigKeyCapChallengeTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.ChallengeTTL = time.Duration(ttlSeconds) * time.Second
}
}
if val, ok := configs[ConfigKeyCapTokenTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.TokenTTL = time.Duration(ttlSeconds) * time.Second
}
}
return settings
}
func (s *runtimeSettingsStore) ensureInvalidationListener() {}
+182
View File
@@ -0,0 +1,182 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package cap provides CAPTCHA and proof-of-work (PoW) verification services.
package cap
import (
"Wavelet/plugins/domain/cap/pow"
"context"
"crypto/sha256"
"encoding/hex"
"strconv"
"strings"
"sync"
"time"
)
const (
redeemTokenIDLength = 8 // 兑换 Token ID 字节长度
redeemVerTokenLength = 15 // 兑换验证 Token 字节长度
tokenPartsCount = 2 // 兑换 Token 由两部分组成
valuePartsCount = 2 // 存储值由 scope 和过期时间组成
)
// Manager orchestrates challenge generation and solution validation.
type Manager struct {
secret []byte
store pow.Store
}
// NewManager creates a new CAPTCHA Manager.
func NewManager(secret []byte, store pow.Store) *Manager {
return &Manager{
secret: secret,
store: store,
}
}
// Generate creates a challenge response.
func (m *Manager) Generate(ctx context.Context, scope string) (*pow.ChallengeResponse, error) {
settings, err := CurrentSettings(ctx)
if err != nil {
return nil, err
}
challengeConfig := pow.ChallengeConfig{
Count: settings.ChallengeCount,
Size: settings.ChallengeSize,
Difficulty: settings.ChallengeDifficulty,
Expires: settings.ChallengeTTL,
}
return pow.GenerateChallenge(m.secret, challengeConfig, scope)
}
// 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: redeemErrInvalidToken}, nil
}
nonceKey := "cap:nonce:" + sigHex
payload, err := pow.VerifyChallengeSolutions(token, solutions, m.secret, scope)
if err != nil {
return &RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // validation errors are returned as response, not system errors
}
now := time.Now().UnixNano() / int64(time.Millisecond)
nonceTTL := time.Duration(payload.Expires-now) * time.Millisecond
if nonceTTL < time.Second {
nonceTTL = time.Second
}
set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL)
if err != nil {
return &RedeemResponse{Success: false, Error: redeemErrNonceStoreFailed}, err
}
if !set {
return &RedeemResponse{Success: false, Error: redeemErrAlreadyRedeemed}, nil
}
settings, err := CurrentSettings(ctx)
if err != nil {
return &RedeemResponse{Success: false, Error: redeemErrSettingsLoad}, err
}
id := pow.RandomHex(redeemTokenIDLength)
verToken := pow.RandomHex(redeemVerTokenLength)
verHashBytes := sha256.Sum256([]byte(verToken))
verHashHex := hex.EncodeToString(verHashBytes[:])
tokenKey := "cap:token:" + id + ":" + verHashHex
tokenExpires := time.Now().Add(settings.TokenTTL)
storeVal := strconv.FormatInt(tokenExpires.UnixNano(), 10) + "|" + scope
if err := m.store.Set(ctx, tokenKey, storeVal, settings.TokenTTL); err != nil {
return &RedeemResponse{Success: false, Error: redeemErrTokenStoreFailed}, err
}
return &RedeemResponse{
Success: true,
Token: id + ":" + verToken,
Expires: tokenExpires.UnixNano() / int64(time.Millisecond),
}, nil
}
// VerifyToken validates and consumes the redeem token (single-use).
func (m *Manager) VerifyToken(ctx context.Context, token, expectedScope string) (bool, error) {
if token == "" {
return false, nil
}
parts := strings.Split(token, ":")
if len(parts) != tokenPartsCount {
return false, nil
}
id := parts[0]
verToken := parts[1]
verHashBytes := sha256.Sum256([]byte(verToken))
verHashHex := hex.EncodeToString(verHashBytes[:])
tokenKey := "cap:token:" + id + ":" + verHashHex
val, exists, err := sGetAndDelete(ctx, m.store, tokenKey)
if err != nil {
return false, err
}
if !exists {
return false, nil
}
valParts := strings.Split(val, "|")
if len(valParts) != valuePartsCount {
return false, nil
}
expNano, err := strconv.ParseInt(valParts[0], 10, 64)
if err != nil {
return false, nil //nolint:nilerr // invalid format is treated as validation failure
}
tokenScope := valParts[1]
if expectedScope != "" && tokenScope != expectedScope {
return false, nil
}
if time.Now().UnixNano() > expNano {
return false, nil
}
return true, nil
}
func sGetAndDelete(ctx context.Context, store pow.Store, key string) (string, bool, error) {
if store == nil {
return "", false, nil
}
return store.GetAndDelete(ctx, key)
}
var (
defaultManagerMu sync.RWMutex
defaultManager *Manager
)
// SetSecret sets the shared secret used by the default manager.
func SetSecret(secret []byte) {
defaultManagerMu.Lock()
defer defaultManagerMu.Unlock()
if len(secret) > 0 {
store := pow.NewMemoryStore(1 * time.Minute)
defaultManager = NewManager(secret, store)
}
}
// GetDefaultManager yields the global singleton CAPTCHA manager.
func GetDefaultManager() *Manager {
defaultManagerMu.RLock()
defer defaultManagerMu.RUnlock()
return defaultManager
}
+474
View File
@@ -0,0 +1,474 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package domain_test
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/idgen"
"Wavelet/plugins/domain/admin"
"Wavelet/plugins/domain/auth"
"Wavelet/plugins/domain/message_gateway"
"Wavelet/plugins/domain/risk_control"
"Wavelet/plugins/domain/user"
"Wavelet/plugins/infra/cache"
"Wavelet/plugins/infra/logger"
"Wavelet/plugins/infra/storage"
"context"
"io/fs"
"path/filepath"
"testing"
"github.com/alicebob/miniredis/v2"
"github.com/glebarez/sqlite"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
db "Wavelet/plugins/infra/database"
)
func setupTestDB(t *testing.T) *gorm.DB {
t.Helper()
_ = idgen.Init(1)
dbPath := filepath.Join(t.TempDir(), "domain_test.db")
testDB, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, testDB.AutoMigrate(
&user.User{},
&user.AccessToken{},
&auth.AuthSource{},
&auth.ExternalAccount{},
&message_gateway.MessageChannel{},
&message_gateway.MessageBinding{},
&message_gateway.MessagePairingCode{},
&admin.SystemConfig{},
&message_gateway.PushChannel{},
&message_gateway.PushEvent{},
&message_gateway.PushHistory{},
))
db.SetDB(testDB)
return testDB
}
type mockOAuthProvider struct {
name string
}
func (m *mockOAuthProvider) Name() string {
return m.name
}
func (m *mockOAuthProvider) GetAuthURL(state string) string {
return "https://auth.example.com/auth?state=" + state
}
func (m *mockOAuthProvider) ExchangeCode(ctx context.Context, code string) (*contracts.OAuthUserInfoDTO, error) {
return &contracts.OAuthUserInfoDTO{
ID: 1001,
Username: "mock_user",
Email: "mock@example.com",
Active: true,
}, nil
}
func TestAuthPlugin(t *testing.T) {
ctx := core.NewContext(context.Background())
ctx.Config().SetSource(core.NewMapSource(nil))
require.NoError(t, ctx.Config().Resolve())
testDB := setupTestDB(t)
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
require.NoError(t, cache.New().Apply(ctx))
require.NoError(t, logger.New().Apply(ctx))
p := auth.New()
assert.Equal(t, "auth", p.Name())
assert.Equal(t, "auth", p.Manifest().Name)
require.NoError(t, p.Apply(ctx))
// 1. Verify migrations registered
entry, ok := ctx.Migrations().Get("auth")
require.True(t, ok)
assert.Equal(t, "auth", entry.PluginID)
entries, err := fs.ReadDir(entry.FS, entry.Dir)
require.NoError(t, err)
assert.NotEmpty(t, entries)
// 2. Verify AuthService
authSvc, err := core.Inject[contracts.AuthService](ctx)
require.NoError(t, err)
require.NotNil(t, authSvc)
assert.NotNil(t, authSvc.RequireAuthMiddleware())
assert.NotNil(t, authSvc.RequireAdminMiddleware())
// 3. Verify AuthRegistry
authReg, err := core.Inject[contracts.AuthRegistry](ctx)
require.NoError(t, err)
require.NotNil(t, authReg)
mockProv := &mockOAuthProvider{name: "github"}
authReg.RegisterOAuthProvider("github", mockProv)
retrieved, ok := authReg.GetOAuthProvider("github")
require.True(t, ok)
assert.Equal(t, "github", retrieved.Name())
assert.Contains(t, authReg.ListOAuthProviders(), "github")
// 4. Verify Routes
routes := ctx.Router().Routes()
var hasSources, hasLogin, hasUserInfo bool
for _, r := range routes {
if r.Path == "/api/v1/oauth/sources" {
hasSources = true
}
if r.Path == "/api/v1/oauth/login" {
hasLogin = true
}
if r.Path == "/api/v1/user-info" {
hasUserInfo = true
}
}
assert.True(t, hasSources)
assert.True(t, hasLogin)
assert.True(t, hasUserInfo)
// 5. Verify Settings
schema, ok := ctx.Settings().Get("auth.session_age")
require.True(t, ok)
assert.Equal(t, 86400*7, schema.Default)
}
func TestUserPlugin(t *testing.T) {
ctx := core.NewContext(context.Background())
ctx.Config().SetSource(core.NewMapSource(nil))
require.NoError(t, ctx.Config().Resolve())
testDB := setupTestDB(t)
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
require.NoError(t, cache.New().Apply(ctx))
require.NoError(t, logger.New().Apply(ctx))
require.NoError(t, auth.New().Apply(ctx))
p := user.New()
assert.Equal(t, "user", p.Name())
assert.Equal(t, "user", p.Manifest().Name)
require.NoError(t, p.Apply(ctx))
// 1. Verify migrations
entry, ok := ctx.Migrations().Get("user")
require.True(t, ok)
assert.Equal(t, "user", entry.PluginID)
// 2. Verify UserService
userSvc, err := core.Inject[contracts.UserService](ctx)
require.NoError(t, err)
require.NotNil(t, userSvc)
testCtx := context.Background()
// 3. Create user
created, err := userSvc.CreateUser(testCtx, contracts.CreateUserRequest{
Username: "bob",
Password: "SecurePassword123!",
Nickname: "Bob Builder",
Email: "bob@example.com",
IsAdmin: false,
})
require.NoError(t, err)
require.NotNil(t, created)
assert.Equal(t, "bob", created.Username)
assert.Equal(t, "Bob Builder", created.Nickname)
assert.Equal(t, "bob@example.com", created.Email)
assert.False(t, created.IsAdmin)
// 4. Query user
byID, err := userSvc.GetUserByID(testCtx, created.ID)
require.NoError(t, err)
assert.Equal(t, "bob", byID.Username)
byUsername, err := userSvc.GetUserByUsername(testCtx, "bob")
require.NoError(t, err)
assert.Equal(t, created.ID, byUsername.ID)
byEmail, err := userSvc.GetUserByEmail(testCtx, "bob@example.com")
require.NoError(t, err)
assert.Equal(t, created.ID, byEmail.ID)
// 5. Password verification and update
assert.True(t, userSvc.VerifyPassword(testCtx, created.ID, "SecurePassword123!"))
assert.False(t, userSvc.VerifyPassword(testCtx, created.ID, "WrongPass"))
require.NoError(t, userSvc.UpdatePassword(testCtx, created.ID, "SecurePassword123!", "NewSecurePassword456!"))
assert.True(t, userSvc.VerifyPassword(testCtx, created.ID, "NewSecurePassword456!"))
// 6. Update Profile
newBio := "I build things"
newPhone := "13800138000"
updated, err := userSvc.UpdateProfile(testCtx, created.ID, contracts.UpdateUserProfileRequest{
Bio: &newBio,
Phone: &newPhone,
})
require.NoError(t, err)
assert.Equal(t, newBio, updated.Bio)
assert.Equal(t, newPhone, updated.Phone)
// 7. Update Last Login
require.NoError(t, userSvc.UpdateLastLogin(testCtx, created.ID, "127.0.0.1"))
// 8. Admin operations: SetUserActive, SetUserAdmin, ListUsers
require.NoError(t, userSvc.SetUserAdmin(testCtx, created.ID, true))
reloaded, err := userSvc.GetUserByID(testCtx, created.ID)
require.NoError(t, err)
assert.True(t, reloaded.IsAdmin)
require.NoError(t, userSvc.SetUserActive(testCtx, created.ID, false))
reloadedBanned, err := userSvc.GetUserByID(testCtx, created.ID)
require.NoError(t, err)
assert.False(t, reloadedBanned.IsActive)
list, total, err := userSvc.ListUsers(testCtx, 1, 10, "bob")
require.NoError(t, err)
assert.Equal(t, int64(1), total)
assert.Len(t, list, 1)
assert.Equal(t, "bob", list[0].Username)
// 9. Tasks & Schedules
taskDef, ok := ctx.Tasks().Get("user:send_email_code")
require.True(t, ok)
assert.Equal(t, 3, taskDef.Retry)
// 10. Settings
sReg, ok := ctx.Settings().Get("user.registration_enabled")
require.True(t, ok)
assert.Equal(t, true, sReg.Default)
}
func TestMessageGatewayPlugin(t *testing.T) {
ctx := core.NewContext(context.Background())
ctx.Config().SetSource(core.NewMapSource(nil))
require.NoError(t, ctx.Config().Resolve())
testDB := setupTestDB(t)
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
require.NoError(t, cache.New().Apply(ctx))
require.NoError(t, logger.New().Apply(ctx))
p := message_gateway.New()
assert.Equal(t, "message_gateway", p.Name())
assert.Equal(t, "message_gateway", p.Manifest().Name)
require.NoError(t, p.Apply(ctx))
// 1. Migrations
entry, ok := ctx.Migrations().Get("message_gateway")
require.True(t, ok)
assert.Equal(t, "message_gateway", entry.PluginID)
// 2. Routes
routes := ctx.Router().Routes()
var hasChannels, hasBindings bool
for _, r := range routes {
if r.Path == "/api/v1/message-gateway/channels" {
hasChannels = true
}
if r.Path == "/api/v1/message-gateway/bindings" {
hasBindings = true
}
}
assert.True(t, hasChannels)
assert.True(t, hasBindings)
// 3. Tasks & Schedules
taskDef, ok := ctx.Tasks().Get("message_gateway:push_notification")
require.True(t, ok)
assert.Equal(t, 3, taskDef.Retry)
schedDef, ok := ctx.Schedules().Get("message_gateway:cleanup_pairing_codes")
require.True(t, ok)
assert.Equal(t, "*/10 * * * *", schedDef.Spec)
// 4. EventBus Trigger
var receivedEvent message_gateway.PushNotificationEvent
var eventFired bool
ctx.Events().On("notification:push", func(c context.Context, e message_gateway.PushNotificationEvent) error {
eventFired = true
receivedEvent = e
return nil
})
err := ctx.Events().Emit(context.Background(), "notification:push", message_gateway.PushNotificationEvent{
UserID: 99,
Channel: "telegram",
Title: "System Alert",
Content: "Disk 85% full",
})
require.NoError(t, err)
assert.True(t, eventFired)
assert.Equal(t, uint64(99), receivedEvent.UserID)
assert.Equal(t, "telegram", receivedEvent.Channel)
assert.Equal(t, "System Alert", receivedEvent.Title)
// 5. Settings
schema, ok := ctx.Settings().Get("message_gateway.pairing_code_expiry_minutes")
require.True(t, ok)
assert.Equal(t, 15, schema.Default)
}
func TestRiskControlPlugin(t *testing.T) {
ctx := core.NewContext(context.Background())
ctx.Config().SetSource(core.NewMapSource(nil))
require.NoError(t, ctx.Config().Resolve())
p := risk_control.New()
assert.Equal(t, "risk_control", p.Name())
assert.Equal(t, "risk_control", p.Manifest().Name)
require.NoError(t, p.Apply(ctx))
// 1. Middleware registered on Router
mws := ctx.Router().Middlewares()
assert.NotEmpty(t, mws)
// 2. Settings
schema, ok := ctx.Settings().Get("risk_control.ip_rate_limit_per_minute")
require.True(t, ok)
assert.Equal(t, 60, schema.Default)
// 3. Disposal cleanup
require.NoError(t, ctx.Dispose())
}
func TestAdminPlugin(t *testing.T) {
ctx := core.NewContext(context.Background())
ctx.Config().SetSource(core.NewMapSource(nil))
require.NoError(t, ctx.Config().Resolve())
testDB := setupTestDB(t)
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
require.NoError(t, cache.New().Apply(ctx))
require.NoError(t, logger.New().Apply(ctx))
p := admin.New()
assert.Equal(t, "admin", p.Name())
assert.Equal(t, "admin", p.Manifest().Name)
require.NoError(t, p.Apply(ctx))
// 1. Admin Routes
routes := ctx.Router().Routes()
var hasStatus, hasDBOverview, hasUsers, hasTasks, hasConfigs bool
for _, r := range routes {
if r.Path == "/api/v1/admin/status" {
hasStatus = true
}
if r.Path == "/api/v1/admin/db-manage/overview" {
hasDBOverview = true
}
if r.Path == "/api/v1/admin/users" {
hasUsers = true
}
if r.Path == "/api/v1/admin/tasks/types" {
hasTasks = true
}
if r.Path == "/api/v1/admin/system-configs" {
hasConfigs = true
}
}
assert.True(t, hasStatus)
assert.True(t, hasDBOverview)
assert.True(t, hasUsers)
assert.True(t, hasTasks)
assert.True(t, hasConfigs)
// 2. Task & Schedule
_, ok := ctx.Tasks().Get("admin:system_cleanup")
require.True(t, ok)
sched, ok := ctx.Schedules().Get("admin:system_cleanup")
require.True(t, ok)
assert.Equal(t, "0 4 * * *", sched.Spec)
// 3. Settings
schema, ok := ctx.Settings().Get("admin.system_cleanup_cron")
require.True(t, ok)
assert.Equal(t, "0 4 * * *", schema.Default)
}
func TestAllDomainPluginsCombined(t *testing.T) {
mr, err := miniredis.Run()
require.NoError(t, err)
defer mr.Close()
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
defer func() { _ = rdb.Close() }()
ctx := core.NewContext(context.Background())
ctx.Config().SetSource(core.NewMapSource(map[string]any{
"redis": map[string]any{
"enabled": true,
"addrs": []string{mr.Addr()},
},
}))
require.NoError(t, ctx.Config().Resolve())
testDB := setupTestDB(t)
// Apply Infra plugins
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
require.NoError(t, cache.New(cache.WithRedis(rdb)).Apply(ctx))
require.NoError(t, logger.New().Apply(ctx))
require.NoError(t, storage.New().Apply(ctx))
// Apply Domain plugins
require.NoError(t, auth.New().Apply(ctx))
require.NoError(t, user.New().Apply(ctx))
require.NoError(t, message_gateway.New().Apply(ctx))
require.NoError(t, risk_control.New().Apply(ctx))
require.NoError(t, admin.New().Apply(ctx))
// Verify cross-plugin service injection via Using3
var resolved bool
err = core.Using3(ctx, func(authSvc contracts.AuthService, userSvc contracts.UserService, authReg contracts.AuthRegistry) {
resolved = true
assert.NotNil(t, authSvc)
assert.NotNil(t, userSvc)
assert.NotNil(t, authReg)
})
require.NoError(t, err)
assert.True(t, resolved)
// Verify all migration entries
allMigrations := ctx.Migrations().Entries()
assert.GreaterOrEqual(t, len(allMigrations), 3)
// Verify total routes registered
allRoutes := ctx.Router().Routes()
assert.GreaterOrEqual(t, len(allRoutes), 20)
// Verify total tasks registered
allTasks := ctx.Tasks().Tasks()
assert.GreaterOrEqual(t, len(allTasks), 4)
for _, task := range allTasks {
dto := task.ToDTO()
assert.NotEmpty(t, dto.Type, "task %s should have type", task.Pattern)
assert.NotEmpty(t, dto.AsynqTask, "task %s should have asynq_task", task.Pattern)
assert.NotEmpty(t, dto.Name, "task %s should have name", task.Pattern)
}
// Verify total schedules registered
allSchedules := ctx.Schedules().Schedules()
assert.GreaterOrEqual(t, len(allSchedules), 2)
// 每个调度指向的任务类型都必须已注册 Handler,否则触发时会投递到无人处理的
// 任务类型,预期的清理逻辑静默失效。
for _, sched := range allSchedules {
_, ok := ctx.Tasks().Get(sched.TaskType)
assert.Truef(t, ok, "schedule %q dispatches to task %q, which is never registered",
sched.Spec, sched.TaskType)
}
// Verify total settings schemas registered
allSettings := ctx.Settings().Schemas()
assert.GreaterOrEqual(t, len(allSettings), 7)
// Clean shutdown
require.NoError(t, ctx.Dispose())
}
@@ -0,0 +1,163 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package qq implements the official QQ Bot C2C adapter.
package qq
import (
"Wavelet/pkg/logger"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/message_gateway/model"
"Wavelet/plugins/domain/message_gateway/service"
"context"
"fmt"
"strings"
"sync"
"time"
"github.com/tencent-connect/botgo"
"github.com/tencent-connect/botgo/dto"
"github.com/tencent-connect/botgo/event"
"github.com/tencent-connect/botgo/openapi"
"github.com/tencent-connect/botgo/token"
"golang.org/x/oauth2"
)
// qqEvent is a testable inbound envelope.
type qqEvent struct {
Kind string
UserID string
Text string
MessageID string
}
// Adapter is an official QQ Bot C2C channel.
type Adapter struct {
cfg model.ChannelConfig
onInbound service.Handler
api openapi.OpenAPI
tokenSrc oauth2.TokenSource
cancel context.CancelFunc
mu sync.Mutex
disconnected bool
}
// New constructs a QQ adapter.
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")
}
return &Adapter{cfg: cfg, onInbound: onInbound}, nil
}
// Type returns qq.
func (a *Adapter) Type() string { return model.ChannelTypeQQ }
// Capabilities reports C2C text/media support.
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).
func (a *Adapter) Connect(ctx context.Context) error {
credentials := &token.QQBotCredentials{
AppID: a.cfg.Credentials["app_id"],
AppSecret: a.cfg.Credentials["app_secret"],
}
tokSrc := token.NewQQBotTokenSource(credentials)
runCtx, cancel := context.WithCancel(ctx)
if err := token.StartRefreshAccessToken(runCtx, tokSrc); err != nil {
cancel()
return fmt.Errorf("qq: refresh token: %w", err)
}
var api openapi.OpenAPI
const apiTimeout = 5 * time.Second
if strings.EqualFold(strings.TrimSpace(a.cfg.Extra["sandbox"]), "true") {
api = botgo.NewSandboxOpenAPI(credentials.AppID, tokSrc).WithTimeout(apiTimeout)
} else {
api = botgo.NewOpenAPI(credentials.AppID, tokSrc).WithTimeout(apiTimeout)
}
wsAP, err := api.WS(ctx, nil, "")
if err != nil {
cancel()
return fmt.Errorf("qq: websocket ap: %w", err)
}
intent := event.RegisterHandlers(event.C2CMessageEventHandler(func(_ *dto.WSPayload, data *dto.WSC2CMessageData) error {
authorID := ""
if data != nil && data.Author != nil {
authorID = data.Author.ID
}
text := ""
id := ""
if data != nil {
text = data.Content
id = data.ID
}
a.handleEvent(runCtx, qqEvent{Kind: "c2c", UserID: authorID, Text: text, MessageID: id})
return nil
}))
a.mu.Lock()
a.api = api
a.tokenSrc = tokSrc
a.cancel = cancel
a.disconnected = false
a.mu.Unlock()
util.Go(func() {
if err := botgo.NewSessionManager().Start(wsAP, tokSrc, &intent); err != nil {
logger.ErrorF(runCtx, "qq session stopped: %v", err)
}
})
return nil
}
// Disconnect stops token refresh and drops further inbound events.
func (a *Adapter) Disconnect(_ context.Context) error {
a.mu.Lock()
defer a.mu.Unlock()
a.disconnected = true
if a.cancel != nil {
a.cancel()
a.cancel = nil
}
return nil
}
// Send posts a C2C text reply.
func (a *Adapter) Send(ctx context.Context, to model.Recipient, msg model.OutboundMessage) error {
a.mu.Lock()
api := a.api
a.mu.Unlock()
if api == nil {
return fmt.Errorf("qq: not connected")
}
_, err := api.PostC2CMessage(ctx, to.PlatformUserID, &dto.MessageToCreate{
Content: msg.Text,
MsgID: msg.ReplyToID,
})
return err
}
func (a *Adapter) handleEvent(ctx context.Context, ev qqEvent) {
if ev.Kind != "c2c" {
return
}
a.mu.Lock()
disconnected := a.disconnected
a.mu.Unlock()
if disconnected || a.onInbound == nil {
return
}
_ = a.onInbound(ctx, model.InboundMessage{
ChannelID: a.cfg.ID,
ChannelType: model.ChannelTypeQQ,
PlatformUserID: ev.UserID,
ChatID: ev.UserID,
MessageID: ev.MessageID,
Text: ev.Text,
})
}
@@ -0,0 +1,41 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package qq
import (
"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 model.InboundMessage) error {
got++
return nil
}}
a.handleEvent(context.Background(), qqEvent{Kind: "group", UserID: "u1", Text: "hi"})
if got != 0 {
t.Fatal("non-C2C must be ignored")
}
}
func TestHandleEvent_C2CText(t *testing.T) {
var got model.InboundMessage
a := &Adapter{cfg: model.ChannelConfig{ID: 3}, onInbound: func(ctx context.Context, msg model.InboundMessage) error {
got = msg
return nil
}}
a.handleEvent(context.Background(), qqEvent{Kind: "c2c", UserID: "openid-1", Text: "hello", MessageID: "m1"})
if got.Text != "hello" || got.PlatformUserID != "openid-1" || got.ChannelID != 3 {
t.Fatalf("%+v", got)
}
}
func TestNew_RequiresCreds(t *testing.T) {
_, err := New(model.ChannelConfig{}, nil)
if err == nil {
t.Fatal("expected error")
}
}
@@ -0,0 +1,181 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package telegram implements the Telegram private-chat adapter.
package telegram
import (
"Wavelet/pkg/logger"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/message_gateway/model"
"Wavelet/plugins/domain/message_gateway/service"
"context"
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
"time"
tele "gopkg.in/telebot.v4"
)
// Adapter is a Telegram private-chat channel.
type Adapter struct {
cfg model.ChannelConfig
onInbound service.Handler
bot *tele.Bot
}
// 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")
}
return &Adapter{cfg: cfg, onInbound: onInbound}, nil
}
// Type returns telegram.
func (a *Adapter) Type() string { return model.ChannelTypeTelegram }
// Capabilities reports private-chat media support.
func (a *Adapter) Capabilities() model.Capability {
return model.Capability{Text: true, Image: true, File: true, Reply: true}
}
// longPollWindow is how long Telegram may hold a getUpdates call open before
// returning empty. telebot converts it with int(timeout / time.Second), so a
// bare integer here would mean nanoseconds, send timeout=0 and turn the poller
// into a tight loop against the Bot API.
const longPollWindow = 10 * time.Second
// buildTeleSettings assembles the telebot settings.
func buildTeleSettings(cfg model.ChannelConfig) tele.Settings {
pref := tele.Settings{
Token: cfg.Credentials["bot_token"],
Poller: &tele.LongPoller{Timeout: longPollWindow},
}
if base := strings.TrimSpace(cfg.Extra["base_url"]); base != "" {
pref.URL = strings.TrimSuffix(base, "/")
}
return pref
}
// Connect starts long polling.
func (a *Adapter) Connect(ctx context.Context) error {
bot, err := tele.NewBot(buildTeleSettings(a.cfg))
if err != nil {
return fmt.Errorf("telegram: new bot: %w", err)
}
a.bot = bot
bot.Handle(tele.OnText, func(c tele.Context) error {
a.handleTeleMessage(ctx, c.Message())
return nil
})
bot.Handle(tele.OnPhoto, func(c tele.Context) error {
a.handleTeleMessage(ctx, c.Message())
return nil
})
bot.Handle(tele.OnDocument, func(c tele.Context) error {
a.handleTeleMessage(ctx, c.Message())
return nil
})
util.Go(func() {
bot.Start()
})
util.Go(func() {
<-ctx.Done()
bot.Stop()
})
return nil
}
// Disconnect stops the bot.
func (a *Adapter) Disconnect(_ context.Context) error {
if a.bot != nil {
a.bot.Stop()
}
return nil
}
// Send replies to a private chat.
func (a *Adapter) Send(_ context.Context, to model.Recipient, msg model.OutboundMessage) error {
if a.bot == nil {
return fmt.Errorf("telegram: not connected")
}
chatID, err := strconv.ParseInt(to.ChatID, 10, 64)
if err != nil {
return fmt.Errorf("telegram: chat id: %w", err)
}
_, err = a.bot.Send(tele.ChatID(chatID), msg.Text)
return err
}
func (a *Adapter) handleTeleMessage(ctx context.Context, m *tele.Message) {
if m == nil || m.Chat == nil || m.Chat.Type != tele.ChatPrivate {
return
}
if a.onInbound == nil {
return
}
msg := model.InboundMessage{
ChannelID: a.cfg.ID,
ChannelType: model.ChannelTypeTelegram,
PlatformUserID: strconv.FormatInt(m.Sender.ID, 10),
ChatID: strconv.FormatInt(m.Chat.ID, 10),
MessageID: strconv.Itoa(m.ID),
Text: m.Text,
}
if m.Caption != "" && msg.Text == "" {
msg.Text = m.Caption
}
if a.bot != nil {
dir, attachments := a.downloadMedia(m)
if dir != "" {
defer func() {
if err := os.RemoveAll(dir); err != nil {
logger.WarnF(ctx, "telegram: 清理临时媒体目录 %q 失败: %v", dir, err)
}
}()
}
msg.Attachments = attachments
}
_ = a.onInbound(ctx, msg)
}
// downloadMedia fetches message media into a scratch directory, returned so the
// caller can remove it once the inbound handler no longer needs the paths.
// An empty dir means nothing was downloaded.
func (a *Adapter) downloadMedia(m *tele.Message) (string, []model.Attachment) {
var files []*tele.File
var names []string
if m.Photo != nil {
files = append(files, m.Photo.MediaFile())
names = append(names, "photo.jpg")
}
if m.Document != nil {
files = append(files, &m.Document.File)
name := m.Document.FileName
if name == "" {
name = "file"
}
names = append(names, name)
}
if len(files) == 0 {
return "", nil
}
dir, err := os.MkdirTemp("", "wg-tg-*")
if err != nil {
return "", []model.Attachment{{Error: err.Error()}}
}
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, model.Attachment{FileName: names[i], Error: err.Error()})
continue
}
out = append(out, model.Attachment{Path: path, FileName: names[i]})
}
return dir, out
}
@@ -0,0 +1,78 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package telegram
import (
"Wavelet/plugins/domain/message_gateway/model"
"context"
"testing"
"time"
tele "gopkg.in/telebot.v4"
)
// TestBuildTeleSettingsLongPollWindow 回归:LongPoller.Timeout 是 time.Duration,
// telebot 以 int(timeout/time.Second) 下发给 getUpdates。写成裸整数会被解释为
// 纳秒,令 timeout=0,长轮询退化为对 Bot API 的空转轮询。
func TestBuildTeleSettingsLongPollWindow(t *testing.T) {
pref := buildTeleSettings(model.ChannelConfig{
Credentials: map[string]string{"bot_token": "token"},
Extra: map[string]string{"base_url": "https://tg.example.com/api/"},
})
poller, ok := pref.Poller.(*tele.LongPoller)
if !ok {
t.Fatalf("expected *tele.LongPoller, got %T", pref.Poller)
}
if got := int(poller.Timeout / time.Second); got != 10 {
t.Errorf("getUpdates would receive timeout=%d seconds, want 10", got)
}
if pref.URL != "https://tg.example.com/api" {
t.Errorf("base_url trailing slash should be trimmed, got %q", pref.URL)
}
}
func TestHandleUpdate_DropsGroups(t *testing.T) {
var got int
a := &Adapter{onInbound: func(ctx context.Context, msg model.InboundMessage) error {
got++
return nil
}}
a.handleTeleMessage(context.Background(), &tele.Message{
ID: 1,
Text: "hi",
Chat: &tele.Chat{ID: -100, Type: tele.ChatGroup},
Sender: &tele.User{ID: 1},
})
if got != 0 {
t.Fatalf("group must be ignored")
}
}
func TestHandleUpdate_PrivateText(t *testing.T) {
var got model.InboundMessage
a := &Adapter{
cfg: model.ChannelConfig{ID: 7, Type: "telegram"},
onInbound: func(ctx context.Context, msg model.InboundMessage) error {
got = msg
return nil
},
}
a.handleTeleMessage(context.Background(), &tele.Message{
ID: 9,
Text: "hi",
Chat: &tele.Chat{ID: 42, Type: tele.ChatPrivate},
Sender: &tele.User{ID: 42},
})
if got.Text != "hi" || got.PlatformUserID != "42" || got.ChannelID != 7 {
t.Fatalf("%+v", got)
}
}
func TestNew_RequiresToken(t *testing.T) {
_, err := New(model.ChannelConfig{}, nil)
if err == nil {
t.Fatal("expected error")
}
}
@@ -0,0 +1,72 @@
// 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")
// ErrUnsupportedUserLookupField rejects a column name that the repository is not
// allowed to interpolate into a WHERE clause.
ErrUnsupportedUserLookupField = errors.New("unsupported user lookup field")
)
// 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"
)
@@ -0,0 +1,139 @@
// 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/service"
"net/http"
"strconv"
"github.com/gin-gonic/gin"
)
// ListAdminChannelDefinitions returns form schemas for supported channel types.
// @Summary List message gateway channel definitions
// @Description Returns form field definitions for Telegram and QQ channels
// @Tags admin-message-gateway
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.Definition}
// @Router /api/v1/admin/message-gateway/channels/definitions [get]
func ListAdminChannelDefinitions(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(service.ListDefinitions()))
}
// ListAdminChannels lists configured messaging channels with secrets masked.
// @Summary List message gateway channels
// @Description Returns all messaging channels; secrets are masked
// @Tags admin-message-gateway
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.ChannelDTO}
// @Router /api/v1/admin/message-gateway/channels [get]
func ListAdminChannels(c *gin.Context) {
rows, err := service.ListChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(rows))
}
func parseAdminChannelID(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
}
func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
if err.Error() == errs.ErrChannelNotFound {
response.AbortNotFound(c, err.Error())
return
}
fallback(c, err.Error())
}
// CreateAdminChannel creates a messaging channel.
// @Summary Create message gateway channel
// @Description Creates a Telegram or QQ channel with encrypted credentials
// @Tags admin-message-gateway
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body 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, service.CreateChannel)
}
// UpdateAdminChannel patches a messaging channel. Empty secrets keep the previous values.
// @Summary Update message gateway channel
// @Description Updates a channel; empty secrets keep the current ciphertext
// @Tags admin-message-gateway
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "channel id"
// @Param request body 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, service.UpdateChannel, func(c *gin.Context, err error) {
handleAdminChannelError(c, err, response.AbortBadRequest)
})
}
// DeleteAdminChannel removes a channel and its bindings/pairing codes.
// @Summary Delete message gateway channel
// @Description Deletes a channel and cascaded bindings and pairing codes
// @Tags admin-message-gateway
// @Produce json
// @Security SessionCookie
// @Param id path int true "channel id"
// @Success 200 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/admin/message-gateway/channels/{id} [delete]
func DeleteAdminChannel(c *gin.Context) {
id, ok := parseAdminChannelID(c)
if !ok {
return
}
if err := service.DeleteChannel(c.Request.Context(), id); err != nil {
handleAdminChannelError(c, err, response.AbortInternal)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// TestAdminChannel probes stored credentials (Telegram getMe or QQ token).
// @Summary Test message gateway channel
// @Description Probes stored credentials without returning secrets
// @Tags admin-message-gateway
// @Produce json
// @Security SessionCookie
// @Param id path int true "channel id"
// @Success 200 {object} response.Any
// @Failure 400 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/admin/message-gateway/channels/{id}/test [post]
func TestAdminChannel(c *gin.Context) {
id, ok := parseAdminChannelID(c)
if !ok {
return
}
if err := service.ProbeChannel(c.Request.Context(), id); err != nil {
handleAdminChannelError(c, err, response.AbortBadRequest)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,182 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// 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"
"github.com/gin-gonic/gin"
)
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=[]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, errs.ErrLoginRequired)
return
}
rows, err := service.ListEnabledPublicChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(rows))
}
// ListBindings lists the current user's bot bindings.
// @Summary List message gateway bindings
// @Description Returns the current user's bound messaging channels
// @Tags message-gateway
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]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, errs.ErrLoginRequired)
return
}
rows, err := service.ListUserBindings(c.Request.Context(), user.ID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(rows))
}
// BindBinding consumes a pairing code and binds the platform identity.
// @Summary Bind a messaging channel
// @Description Binds the current user to a platform identity using a one-time pairing code
// @Tags message-gateway
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body 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, errs.ErrLoginRequired)
return
}
var req model.BindRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := service.BindChannel(c.Request.Context(), user.ID, req)
if err != nil {
if errors.Is(err, errs.ErrPlatformAlreadyBound) {
response.AbortConflict(c, err.Error())
return
}
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(dto))
}
// UnbindBinding removes the current user's binding.
// @Summary Unbind a messaging channel
// @Description Removes a binding owned by the current user
// @Tags message-gateway
// @Produce json
// @Security SessionCookie
// @Param id path int true "binding id"
// @Success 200 {object} response.Any
// @Failure 403 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/message-gateway/bindings/{id} [delete]
func UnbindBinding(c *gin.Context) {
user, ok := currentUser(c)
if !ok || user == nil {
response.AbortUnauthorized(c, errs.ErrLoginRequired)
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, errs.ErrInvalidBindingID)
return
}
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, errs.ErrBindingForbidden) {
response.AbortForbidden(c, err.Error())
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -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())
}
@@ -0,0 +1,148 @@
// 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"
)
// ListPushEvents lists configured push events.
func ListPushEvents(c *gin.Context) {
ctx := c.Request.Context()
events, err := service.ListPushEvents(ctx)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(events))
}
// ListBuiltInPushEvents lists system built-in push event definitions.
func ListBuiltInPushEvents(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(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, 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, errs.ErrRecordNotFound) {
response.AbortNotFound(c, errs.ErrEventNotFound)
return
}
fallback(c, err.Error())
}
// CreatePushEvent creates a new push event configuration.
func CreatePushEvent(c *gin.Context) {
handleJSONRequest(c, service.CreatePushEvent)
}
// DeletePushEvent deletes a push event configuration by ID.
func DeletePushEvent(c *gin.Context) {
id, ok := parsePushEventID(c)
if !ok {
return
}
if err := service.DeletePushEvent(c.Request.Context(), id); err != nil {
handlePushEventNotFoundError(c, err, response.AbortInternal)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// UpdatePushEvent updates an existing push event.
func UpdatePushEvent(c *gin.Context) {
id, ok := parsePushEventID(c)
if !ok {
return
}
var req model.UpdatePushEventRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := service.UpdatePushEvent(c.Request.Context(), id, req); err != nil {
handlePushEventNotFoundError(c, err, response.AbortBadRequest)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// TogglePushEvent toggles the enabled state of a push event.
func TogglePushEvent(c *gin.Context) {
id, ok := parsePushEventID(c)
if !ok {
return
}
enabled, err := service.TogglePushEvent(c.Request.Context(), id)
if err != nil {
handlePushEventNotFoundError(c, err, response.AbortBadRequest)
return
}
c.JSON(http.StatusOK, response.OK(enabled))
}
// ListPushHistories returns paginated push notification delivery histories.
func ListPushHistories(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
if page < 1 {
page = 1
}
if pageSize < 1 {
pageSize = 20
}
total, results, err := service.ListPushHistories(c.Request.Context(), model.PushHistoryListFilter{
EventKey: c.Query("event_key"),
Status: c.Query("status"),
Page: page,
PageSize: pageSize,
})
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(map[string]any{
"total": total,
"results": results,
}))
}
// TestPush executes a synchronous push test using the specified config.
func TestPush(c *gin.Context) {
var req model.TestPushRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := service.RunPushTest(c.Request.Context(), req.Config, req.Target); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -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)
}
}
}
@@ -0,0 +1,51 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway_test
import (
"Wavelet/pkg/testhelper"
"Wavelet/plugins/domain/message_gateway"
"context"
"testing"
"time"
"gorm.io/gorm"
)
type mockDBService struct {
db *gorm.DB
}
func (m *mockDBService) GORM() *gorm.DB {
return m.db
}
func (m *mockDBService) DB(ctx context.Context) *gorm.DB {
return m.db.WithContext(ctx)
}
func (m *mockDBService) Named(_ string) *gorm.DB {
return m.db
}
func TestUpsertPairingCode_ReusesUnexpired(t *testing.T) {
testDB, _, cleanup := testhelper.SetupTestEnvironment(t)
message_gateway.SetDBServiceForTest(&mockDBService{db: testDB})
defer func() {
message_gateway.SetDBServiceForTest(nil)
cleanup()
}()
ctx := context.Background()
first, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ABCD1234", time.Now().Add(15*time.Minute))
if err != nil {
t.Fatal(err)
}
second, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ZZZZ9999", time.Now().Add(15*time.Minute))
if err != nil {
t.Fatal(err)
}
if first.Code != second.Code || first.Code != "ABCD1234" {
t.Fatalf("reuse failed: %+v %+v", first, second)
}
}
@@ -0,0 +1,93 @@
-- +goose Up
-- +goose StatementBegin
CREATE TABLE IF NOT EXISTS w_message_channels (
id BIGINT PRIMARY KEY,
name VARCHAR(128) NOT NULL,
type VARCHAR(32) NOT NULL,
owner_scope VARCHAR(16) NOT NULL DEFAULT 'system',
owner_id BIGINT NULL,
enabled BOOLEAN NOT NULL DEFAULT TRUE,
credentials TEXT NOT NULL DEFAULT '',
extra TEXT NOT NULL DEFAULT '',
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_message_channels_type ON w_message_channels (type);
CREATE TABLE IF NOT EXISTS w_message_bindings (
id BIGINT PRIMARY KEY,
user_id BIGINT NOT NULL,
channel_id BIGINT NOT NULL,
platform_user_id VARCHAR(128) NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_message_bindings_channel_platform
ON w_message_bindings (channel_id, platform_user_id);
CREATE INDEX IF NOT EXISTS idx_w_message_bindings_user ON w_message_bindings (user_id);
CREATE TABLE IF NOT EXISTS w_message_pairing_codes (
code VARCHAR(16) PRIMARY KEY,
channel_id BIGINT NOT NULL,
platform_user_id VARCHAR(128) NOT NULL,
expires_at TIMESTAMPTZ NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_message_pairing_lookup
ON w_message_pairing_codes (channel_id, platform_user_id);
CREATE TABLE IF NOT EXISTS w_push_events (
id BIGINT PRIMARY KEY,
event_key VARCHAR(80) NOT NULL,
name VARCHAR(100) NOT NULL,
task_type VARCHAR(100) NOT NULL DEFAULT '',
channels TEXT NOT NULL DEFAULT '',
targets TEXT NOT NULL DEFAULT '',
template TEXT NOT NULL DEFAULT '',
enabled BOOLEAN NOT NULL DEFAULT FALSE,
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_push_events_key ON w_push_events(event_key);
CREATE INDEX IF NOT EXISTS idx_w_push_events_enabled ON w_push_events(enabled);
CREATE INDEX IF NOT EXISTS idx_w_push_events_task_type ON w_push_events(task_type);
CREATE TABLE IF NOT EXISTS w_push_channels (
id BIGINT PRIMARY KEY,
name VARCHAR(80) NOT NULL,
description VARCHAR(255) NOT NULL DEFAULT '',
type VARCHAR(50) NOT NULL DEFAULT 'custom',
token VARCHAR(100) NOT NULL DEFAULT '',
url TEXT NOT NULL DEFAULT '',
other TEXT NOT NULL DEFAULT '',
enabled BOOLEAN NOT NULL DEFAULT TRUE,
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_push_channels_name ON w_push_channels(name);
CREATE INDEX IF NOT EXISTS idx_w_push_channels_enabled ON w_push_channels(enabled);
CREATE TABLE IF NOT EXISTS w_push_histories (
id BIGINT PRIMARY KEY,
event_key VARCHAR(80) NOT NULL,
channel VARCHAR(50) NOT NULL,
target VARCHAR(255) NOT NULL,
title VARCHAR(255) NOT NULL,
content TEXT NOT NULL,
level VARCHAR(20) NOT NULL,
status VARCHAR(20) NOT NULL,
error_msg TEXT NOT NULL DEFAULT '',
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_push_histories_event ON w_push_histories(event_key);
CREATE INDEX IF NOT EXISTS idx_w_push_histories_created ON w_push_histories(created_at);
-- +goose StatementEnd
-- +goose Down
-- +goose StatementBegin
DROP TABLE IF EXISTS w_push_histories;
DROP TABLE IF EXISTS w_push_channels;
DROP TABLE IF EXISTS w_push_events;
DROP TABLE IF EXISTS w_message_pairing_codes;
DROP TABLE IF EXISTS w_message_bindings;
DROP TABLE IF EXISTS w_message_channels;
-- +goose StatementEnd
@@ -0,0 +1,93 @@
-- +goose Up
-- +goose StatementBegin
CREATE TABLE IF NOT EXISTS w_message_channels (
id BIGINT PRIMARY KEY,
name VARCHAR(128) NOT NULL,
type VARCHAR(32) NOT NULL,
owner_scope VARCHAR(16) NOT NULL DEFAULT 'system',
owner_id BIGINT NULL,
enabled BOOLEAN NOT NULL DEFAULT 1,
credentials TEXT NOT NULL DEFAULT '',
extra TEXT NOT NULL DEFAULT '',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_message_channels_type ON w_message_channels (type);
CREATE TABLE IF NOT EXISTS w_message_bindings (
id BIGINT PRIMARY KEY,
user_id BIGINT NOT NULL,
channel_id BIGINT NOT NULL,
platform_user_id VARCHAR(128) NOT NULL,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_message_bindings_channel_platform
ON w_message_bindings (channel_id, platform_user_id);
CREATE INDEX IF NOT EXISTS idx_w_message_bindings_user ON w_message_bindings (user_id);
CREATE TABLE IF NOT EXISTS w_message_pairing_codes (
code VARCHAR(16) PRIMARY KEY,
channel_id BIGINT NOT NULL,
platform_user_id VARCHAR(128) NOT NULL,
expires_at DATETIME NOT NULL,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_message_pairing_lookup
ON w_message_pairing_codes (channel_id, platform_user_id);
CREATE TABLE IF NOT EXISTS w_push_events (
id BIGINT PRIMARY KEY,
event_key VARCHAR(80) NOT NULL,
name VARCHAR(100) NOT NULL,
task_type VARCHAR(100) NOT NULL DEFAULT '',
channels TEXT NOT NULL DEFAULT '',
targets TEXT NOT NULL DEFAULT '',
template TEXT NOT NULL DEFAULT '',
enabled BOOLEAN NOT NULL DEFAULT 0,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_push_events_key ON w_push_events(event_key);
CREATE INDEX IF NOT EXISTS idx_w_push_events_enabled ON w_push_events(enabled);
CREATE INDEX IF NOT EXISTS idx_w_push_events_task_type ON w_push_events(task_type);
CREATE TABLE IF NOT EXISTS w_push_channels (
id BIGINT PRIMARY KEY,
name VARCHAR(80) NOT NULL,
description VARCHAR(255) NOT NULL DEFAULT '',
type VARCHAR(50) NOT NULL DEFAULT 'custom',
token VARCHAR(100) NOT NULL DEFAULT '',
url TEXT NOT NULL DEFAULT '',
other TEXT NOT NULL DEFAULT '',
enabled BOOLEAN NOT NULL DEFAULT 1,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_w_push_channels_name ON w_push_channels(name);
CREATE INDEX IF NOT EXISTS idx_w_push_channels_enabled ON w_push_channels(enabled);
CREATE TABLE IF NOT EXISTS w_push_histories (
id BIGINT PRIMARY KEY,
event_key VARCHAR(80) NOT NULL,
channel VARCHAR(50) NOT NULL,
target VARCHAR(255) NOT NULL,
title VARCHAR(255) NOT NULL,
content TEXT NOT NULL,
level VARCHAR(20) NOT NULL,
status VARCHAR(20) NOT NULL,
error_msg TEXT NOT NULL DEFAULT '',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_w_push_histories_event ON w_push_histories(event_key);
CREATE INDEX IF NOT EXISTS idx_w_push_histories_created ON w_push_histories(created_at);
-- +goose StatementEnd
-- +goose Down
-- +goose StatementBegin
DROP TABLE IF EXISTS w_push_histories;
DROP TABLE IF EXISTS w_push_channels;
DROP TABLE IF EXISTS w_push_events;
DROP TABLE IF EXISTS w_message_pairing_codes;
DROP TABLE IF EXISTS w_message_bindings;
DROP TABLE IF EXISTS w_message_channels;
-- +goose StatementEnd
@@ -0,0 +1,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"`
}
@@ -0,0 +1,242 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package model defines the domain entities, DTOs, and schemas for message_gateway.
package model
import (
"Wavelet/plugins/domain/message_gateway/errs"
"errors"
"strings"
"time"
)
// Channel type and scope constants.
const (
ChannelTypeTelegram = "telegram"
ChannelTypeQQ = "qq"
MessageChannelTypeTelegram = "telegram"
MessageChannelTypeQQ = "qq"
MessageOwnerScopeSystem = "system"
TypeCustom = "custom"
TypeEmail = "email"
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"`
Type string `json:"type" gorm:"size:32;not null"`
Name string `json:"name" gorm:"size:64;not null"`
OwnerScope string `json:"owner_scope" gorm:"size:32;not null;default:'system'"`
OwnerID *uint64 `json:"owner_id,omitempty"`
Credentials string `json:"credentials" gorm:"type:text;not null"`
Extra string `json:"extra" gorm:"type:text"`
Enabled bool `json:"enabled" gorm:"default:false;not null"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (MessageChannel) TableName() string {
return "w_message_channels"
}
// MessageBinding maps a platform user to a Wavelet user on one channel.
type MessageBinding struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
ChannelID uint64 `json:"channel_id" gorm:"not null;index"`
PlatformUserID string `json:"platform_user_id" gorm:"size:128;not null;index"`
UserID uint64 `json:"user_id" gorm:"not null;index"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName 表名
func (MessageBinding) TableName() string {
return "w_message_bindings"
}
// MessagePairingCode is a one-time bind code.
type MessagePairingCode struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
Code string `json:"code" gorm:"size:32;uniqueIndex;not null"`
ChannelID uint64 `json:"channel_id" gorm:"not null;index"`
PlatformUserID string `json:"platform_user_id" gorm:"size:128;not null;index"`
UserID uint64 `json:"user_id" gorm:"not null;index"`
ExpiresAt time.Time `json:"expires_at" gorm:"not null;index"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName 表名
func (MessagePairingCode) TableName() string {
return "w_message_pairing_codes"
}
// PushChannel 消息通道模型
type PushChannel struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"size:100;not null"`
Description string `json:"description" gorm:"size:255"`
Type string `json:"type" gorm:"size:50;not null;index"`
URL string `json:"url" gorm:"type:text"`
Token string `json:"token" gorm:"type:text"`
Other string `json:"other" gorm:"type:text"`
Enabled bool `json:"enabled" gorm:"index;not null;default:true"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
}
// TableName 指定 GORM 表名
func (PushChannel) TableName() string {
return "w_push_channels"
}
// Validate 验证与标准化字段
func (c *PushChannel) Validate() error {
c.Name = strings.TrimSpace(c.Name)
if c.Name == "" {
return errors.New(errs.ErrChannelNameRequired)
}
c.Type = strings.TrimSpace(c.Type)
if c.Type == "" {
return errors.New(errs.ErrChannelTypeRequired)
}
return nil
}
// PushEvent 系统通知事件模型
type PushEvent struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
EventKey string `json:"event_key" gorm:"uniqueIndex;size:80;not null"`
Name string `json:"name" gorm:"size:100;not null"`
TaskType string `json:"task_type" gorm:"size:100;index;not null;default:''"`
Channels []string `json:"channels" gorm:"type:text;serializer:json;not null"`
Targets []string `json:"targets" gorm:"type:text;serializer:json;not null"`
Template string `json:"template" gorm:"type:text;not null"`
Enabled bool `json:"enabled" gorm:"index;not null;default:false"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
}
// TableName 指定 GORM 表名
func (PushEvent) TableName() string {
return "w_push_events"
}
// Validate 验证 PushEvent 实体字段
func (e *PushEvent) Validate() error {
e.EventKey = strings.TrimSpace(e.EventKey)
if e.EventKey == "" {
return errors.New(errs.ErrEventKeyRequired)
}
e.Name = strings.TrimSpace(e.Name)
if e.Name == "" {
return errors.New(errs.ErrNameRequired)
}
return nil
}
// PushHistory 推送日志/历史实体
type PushHistory struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
EventKey string `json:"event_key" gorm:"size:80;not null;index"`
Channel string `json:"channel" gorm:"size:50;not null;index"`
Target string `json:"target" gorm:"size:255;not null"`
Title string `json:"title" gorm:"size:255;not null"`
Content string `json:"content" gorm:"type:text;not null"`
Level string `json:"level" gorm:"size:20;not null;default:'INFO'"`
Status string `json:"status" gorm:"size:20;not null;index"`
ErrorMsg string `json:"error_msg" gorm:"type:text"`
Payload string `json:"payload" gorm:"type:text"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
}
// TableName 指定 GORM 表名
func (PushHistory) TableName() string {
return "w_push_histories"
}
// 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"`
}
@@ -0,0 +1,313 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"sync"
"time"
pkgpush "Wavelet/plugins/domain/message_gateway/push"
)
// Push channel and payload constants.
const (
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 = "token"
// KeyOther represents the Other field key.
KeyOther = "other"
// TypeText represents standard text input type.
TypeText = "text"
// TypePassword represents password input type.
TypePassword = "password"
// TypeTextarea represents textarea input type.
TypeTextarea = "textarea"
)
// 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"`
Label string `json:"label"`
Type string `json:"type"`
Required bool `json:"required"`
Placeholder string `json:"placeholder"`
Description string `json:"description"`
}
// PushDefinition represents the metadata and form schema for a notification channel.
type PushDefinition struct {
Type string `json:"type"`
Name string `json:"name"`
Description string `json:"description"`
Fields []PushField `json:"fields"`
}
// 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)
)
// RegisterPushChannelDefinition registers a channel definition.
func RegisterPushChannelDefinition(def PushDefinition) {
pushDefMu.Lock()
defer pushDefMu.Unlock()
pushDefinitions[def.Type] = def
}
// ListPushDefinitions returns all registered channel definitions.
func ListPushDefinitions() []PushDefinition {
pushDefMu.RLock()
defer pushDefMu.RUnlock()
order := []string{ChannelCustom, ChannelLark, ChannelTelegram, ChannelEmail}
res := make([]PushDefinition, 0, len(pushDefinitions))
for _, t := range order {
if d, ok := pushDefinitions[t]; ok {
res = append(res, d)
}
}
for t, d := range pushDefinitions {
found := false
for _, o := range order {
if o == t {
found = true
break
}
}
if !found {
res = append(res, d)
}
}
return res
}
func init() {
RegisterPushChannelDefinition(PushDefinition{
Type: ChannelCustom,
Name: "自定义消息通道",
Description: "使用自定义 HTTP POST 请求向外部 Webhook 发送数据。",
Fields: []PushField{
{
Key: KeyURL,
Label: "请求地址",
Type: TypeText,
Required: true,
Placeholder: "在此填写完整的请求地址,必须使用 HTTPS 协议",
Description: "接口请求的完整 HTTPS URL,例如 https://api.example.com/webhook",
},
{
Key: KeyOther,
Label: "请求体 (JSON)",
Type: TypeTextarea,
Required: true,
Placeholder: "在此输入请求体,支持模板变量,必须为合法的 JSON 格式",
Description: "可使用的变量:$title, $description, $content, $url, $to。例如 {\"text\": \"$content\"}",
},
},
})
RegisterPushChannelDefinition(PushDefinition{
Type: ChannelLark,
Name: "飞书群机器人",
Description: "配置飞书群自定义机器人的 Webhook 接口投递。",
Fields: []PushField{
{
Key: KeyURL,
Label: "Webhook 地址",
Type: TypeText,
Required: true,
Placeholder: "https://open.feishu.cn/open-apis/bot/v2/hook/YOUR_TOKEN",
Description: "从飞书群机器人设置中复制的 Webhook URL",
},
{
Key: KeyToken,
Label: "签名校验密钥 (Secret) (可选)",
Type: TypeText,
Required: false,
Placeholder: "可选,若机器人启用了安全设置中的签名校验,请在此输入",
Description: "飞书群机器人安全设置中的签名校验 Key",
},
{
Key: KeyOther,
Label: "自定义卡片 JSON 模版 (可选)",
Type: TypeTextarea,
Required: false,
Placeholder: "可选,留空则默认使用系统内置的精美互动卡片",
Description: "若填写,必须是合法的飞书卡片 JSON 格式",
},
},
})
RegisterPushChannelDefinition(PushDefinition{
Type: ChannelTelegram,
Name: "Telegram 机器人",
Description: "配置 Telegram 机器人推送消息。",
Fields: []PushField{
{
Key: KeyURL,
Label: "API 基础地址 (可选)",
Type: TypeText,
Required: false,
Placeholder: "https://api.telegram.org",
Description: "接口请求的 HTTPS 基础地址,留空默认为 https://api.telegram.org",
},
{
Key: KeyToken,
Label: "机器人 Token (Bot Token)",
Type: TypePassword,
Required: true,
Placeholder: "在此输入 Telegram 机器人的 Bot Token",
Description: "通过 BotFather 申请到的机器人 Access Token",
},
{
Key: KeyOther,
Label: "默认会话 ID (Chat ID) (可选)",
Type: TypeText,
Required: false,
Placeholder: "例如 -100123456789 或 @channel_name",
Description: "默认的消息接收 Chat ID。如果通知事件中未配置 targets,将推送到此 ID",
},
},
})
RegisterPushChannelDefinition(PushDefinition{
Type: ChannelEmail,
Name: "邮件推送通道",
Description: "邮件推送通道直接使用系统全局 SMTP 设置进行发送,无需在此填写服务器配置。",
Fields: []PushField{},
})
}
@@ -0,0 +1,33 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"strings"
"testing"
)
func TestGenerateCode_AlphabetAndLength(t *testing.T) {
code, err := GenerateCode()
if err != nil {
t.Fatal(err)
}
if len(code) != 8 {
t.Fatalf("len=%d", len(code))
}
for _, r := range code {
if !strings.ContainsRune(CodeAlphabet, r) {
t.Fatalf("bad rune %q", r)
}
}
}
func TestNormalizeAndFormat(t *testing.T) {
if got := NormalizeCode("ab-cd-ef-gh"); got != "ABCDEFGH" {
t.Fatalf("got %q", got)
}
if got := FormatCode("ABCDEFGH"); got != "ABCD-EFGH" {
t.Fatalf("got %q", got)
}
}
@@ -0,0 +1,384 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package message_gateway provides the Bot gateway, multi-channel notification dispatching, and asynchronous push worker domain plugin for Cordis.
package message_gateway
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/pkg/ginutil"
"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"
"github.com/gin-gonic/gin"
)
//go:embed migrations/*/*.sql
var mgMigrations embed.FS
// Option configures the message_gateway plugin.
type Option func(*Plugin)
// WithAutoStartRunner enables automatic bot runner startup in the background.
func WithAutoStartRunner(enable bool) Option {
return func(p *Plugin) {
p.autoStartRunner = enable
}
}
// Plugin implements core.Plugin to provide Bot gateway and notification dispatch domain services.
type Plugin struct {
autoStartRunner bool
cancelRunner context.CancelFunc
}
// New creates a new message_gateway domain plugin.
func New(opts ...Option) *Plugin {
p := &Plugin{}
for _, opt := range opts {
if opt != nil {
opt(p)
}
}
return p
}
// Name returns the unique identifier for the message_gateway domain plugin.
func (p *Plugin) Name() string {
return "message_gateway"
}
// Inject declares required dependencies for the message_gateway domain plugin.
func (p *Plugin) Inject() []reflect.Type {
return []reflect.Type{
reflect.TypeFor[contracts.DBService](),
// AuthService is captured as a middleware value in Apply, so it cannot
// be late-bound with core.When like the other services below; the
// kernel must mount auth first or the routes get a pass-through guard.
reflect.TypeFor[contracts.AuthService](),
}
}
// Manifest returns the plugin metadata.
func (p *Plugin) Manifest() core.Manifest {
return core.Manifest{
Name: "message_gateway",
Version: "1.0.0",
Description: "Bot gateway, multi-channel notification push, and async worker dispatch plugin",
Author: "Wavelet Team",
}
}
type mgAppConfig struct {
SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"`
}
// DeclareConfig declares configuration bindings for the message_gateway plugin.
func (p *Plugin) DeclareConfig() []core.ConfigBinding {
return []core.ConfigBinding{
{Prefix: "app", Target: &mgAppConfig{}},
}
}
// Apply registers message_gateway migrations, routes, tasks, schedules, events, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
var cfg mgAppConfig
if err := ctx.Config().Bind("app", &cfg); err == nil && cfg.SessionSecret != "" {
service.SetCredentialSecret(cfg.SessionSecret)
}
// 0. Bind DBService, CacheService, TaskService, UserService
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
repository.SetDBService(db)
} else {
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
repository.SetDBService(db)
})
}
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
repository.SetCacheService(cache)
service.SetCacheService(cache)
} else {
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
repository.SetCacheService(cache)
service.SetCacheService(cache)
})
}
if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil {
service.SetTaskService(taskSvc)
} else {
core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) {
service.SetTaskService(taskSvc)
})
}
if uSvc, err := core.Inject[contracts.UserService](ctx); err == nil && uSvc != nil {
service.SetUserService(uSvc)
} else {
core.When[contracts.UserService](ctx, func(uSvc contracts.UserService) {
service.SetUserService(uSvc)
})
}
ctx.OnDispose(func() error {
repository.SetDBService(nil)
repository.SetCacheService(nil)
service.SetCacheService(nil)
service.SetTaskService(nil)
service.SetUserService(nil)
return nil
})
// 0. Resolve auth service for middleware (via IoC, not direct import)
denyAuth := ginutil.AuthUnavailable()
loginMW := denyAuth
adminMW := denyAuth
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
loginMW = mw
}
if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok {
adminMW = mw
}
}
// 1. Register migrations
ctx.Migrations().Register("message_gateway", mgMigrations)
// 2. Register User HTTP Routes
handler.RegisterUserRoutes(ctx.Router().Group("/api/v1"), loginMW)
// 3. Register Admin Message Gateway HTTP Routes
handler.RegisterAdminRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW)
// 4. Register Admin Push HTTP Routes
handler.RegisterAdminPushRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW)
const defaultTaskRetry = 3
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.WithTaskType("push_notification"),
extpoints.WithTaskName("消息网关推送通知"),
extpoints.WithTaskDescription("异步执行系统通知的多渠道派发与推送"),
extpoints.WithTaskCategory("push"),
extpoints.WithTaskRetry(defaultTaskRetry),
extpoints.WithTaskQueue("default"),
extpoints.WithTaskRetryable(true),
)
ctx.Task().Register(service.SendNotificationTask, func(c context.Context, payload []byte) error {
return pushHandler.Execute(c, payload)
}, extpoints.WithTaskMeta(service.SendNotificationMeta), extpoints.WithTaskRetry(defaultTaskRetry))
ctx.Task().Register("message_gateway:dispatch_bot_msg", func(_ context.Context, _ []byte) error {
return nil
},
extpoints.WithTaskType("dispatch_bot_msg"),
extpoints.WithTaskName("分发 Bot 消息"),
extpoints.WithTaskDescription("异步处理与分发 Bot 下行消息"),
extpoints.WithTaskCategory("messaging"),
extpoints.WithTaskQueue("default"),
)
ctx.Task().Register("message_gateway:cleanup_pairing_codes", func(c context.Context, _ []byte) error {
return repository.DeleteExpiredPairingCodes(c)
},
extpoints.WithTaskType("cleanup_pairing_codes"),
extpoints.WithTaskName("清理过期配对码"),
extpoints.WithTaskDescription("定时清理已过期的平台 Bot 配对码"),
extpoints.WithTaskCategory("messaging"),
extpoints.WithTaskRetry(defaultTaskRetry),
extpoints.WithTaskQueue("default"),
extpoints.WithTaskRetryable(true),
)
// 6. Register Cron Schedules
ctx.Schedule().RegisterCron("*/10 * * * *", "message_gateway:cleanup_pairing_codes", map[string]any{"action": "cleanup"})
// 7. Register EventBus listeners for decoupled push triggers
ctx.Events().On("notification:push", func(c context.Context, e model.PushNotificationEvent) error {
meta := model.EventMetadata{
Key: "eventbus:" + e.Channel,
Name: e.Title,
DefaultTemplate: model.NotificationMessage{
Title: e.Title,
Content: e.Content,
Level: model.DefaultLevelInfo,
Ext: e.Metadata,
},
Description: "EventBus triggered notification",
}
service.DefaultTrigger.Trigger(c, meta, map[string]any{
"user.id": e.UserID,
"title": e.Title,
"content": e.Content,
})
return nil
})
// 8. Register task completed event listener
ctx.Events().On(contracts.EventTopicTaskCompleted, func(c context.Context, e contracts.TaskCompletedEvent) error {
service.HandleTaskCompleted(c, e)
return nil
})
// 9. Register built-in domain events
service.RegisterCustomEvents()
// 10. Register Settings Schemas
ctx.Settings().Register(extpoints.SettingSchema{
Key: "message_gateway.pairing_code_expiry_minutes",
Default: 15,
Description: "Expiry duration for bot pairing codes in minutes",
Type: "integer",
Category: "messaging",
})
ctx.Settings().Register(extpoints.SettingSchema{
Key: "message_gateway.max_bindings_per_user",
Default: 5,
Description: "Maximum platform bot bindings per user",
Type: "integer",
Category: "messaging",
})
// 11. Optional runner start & lifecycle
if p.autoStartRunner {
runnerCtx, cancel := context.WithCancel(ctx.GoContext())
p.cancelRunner = cancel
util.Go(func() {
_ = service.Start(runnerCtx)
})
}
ctx.OnDispose(func() error {
if p.cancelRunner != nil {
p.cancelRunner()
}
return nil
})
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
)
@@ -0,0 +1,61 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway_test
import (
"Wavelet/core"
"Wavelet/plugins/domain/message_gateway"
"context"
"io/fs"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestMessageGatewayPluginUnit(t *testing.T) {
ctx := core.NewContext(context.Background())
p := message_gateway.New()
assert.Equal(t, "message_gateway", p.Name())
assert.Equal(t, "1.0.0", p.Manifest().Version)
require.NoError(t, p.Apply(ctx))
// Verify migrations
entry, ok := ctx.Migrations().Get("message_gateway")
require.True(t, ok)
entries, err := fs.ReadDir(entry.FS, entry.Dir)
require.NoError(t, err)
assert.NotEmpty(t, entries)
// Verify tasks
task, ok := ctx.Tasks().Get("message_gateway:push_notification")
require.True(t, ok)
assert.Equal(t, 3, task.Retry)
// Verify schedules
sched, ok := ctx.Schedules().Get("message_gateway:cleanup_pairing_codes")
require.True(t, ok)
assert.Equal(t, "*/10 * * * *", sched.Spec)
// Verify settings
setting, ok := ctx.Settings().Get("message_gateway.max_bindings_per_user")
require.True(t, ok)
assert.Equal(t, 5, setting.Default)
}
// TestEveryScheduleHasTaskHandler 回归:RegisterCron 仅登记调度;若同名任务从未
// Register,则每次触发都投递到无人处理的任务类型,清理逻辑静默失效。
func TestEveryScheduleHasTaskHandler(t *testing.T) {
ctx := core.NewContext(context.Background())
require.NoError(t, message_gateway.New().Apply(ctx))
schedules := ctx.Schedules().Schedules()
require.NotEmpty(t, schedules)
for _, sched := range schedules {
_, ok := ctx.Tasks().Get(sched.TaskType)
assert.Truef(t, ok, "schedule %q dispatches to task %q, which is never registered",
sched.Spec, sched.TaskType)
}
}
@@ -0,0 +1,97 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"Wavelet/pkg/httppool"
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strings"
)
func init() {
Register("custom", &CustomPusher{})
}
// maxCustomResponseBytes 限制读取 Webhook 响应体的最大字节数,防止无界读取。
const maxCustomResponseBytes = 4096
// CustomPusher 自定义 Webhook 发送实现
type CustomPusher struct{}
// Send 发送自定义 webhook
func (p *CustomPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, template string, _ map[string]any) (string, error) {
if cfg.URL == "" {
return "", errors.New("custom: URL is required")
}
var reqBody []byte
if template != "" {
// 替换模板中的 {{key}} 占位符
rendered := ParseTemplate(template, body)
reqBody = []byte(rendered)
} else {
// 兜底:直接把 body 转为 JSON 字符串发送
var err error
reqBody, err = json.Marshal(body)
if err != nil {
return "", fmt.Errorf("custom: marshal body failed: %w", err)
}
}
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.URL, bytes.NewReader(reqBody))
if err != nil {
return "", fmt.Errorf("custom: create http request failed: %w", err)
}
httpReq.Header.Set("Content-Type", "application/json")
// 如果配置了 Key 且格式为 "HeaderName:HeaderValue",我们可以附加测试用 Header
if cfg.Key != "" && strings.Contains(cfg.Key, ":") {
parts := strings.SplitN(cfg.Key, ":", 2) //nolint:mnd
httpReq.Header.Set(strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1]))
}
client := httppool.NewClient(defaultHTTPClientTimeout)
resp, err := client.Do(httpReq)
if err != nil {
return "", fmt.Errorf("custom: http request failed: %w", err)
}
defer func() { _ = resp.Body.Close() }()
bodyBytes, _ := io.ReadAll(io.LimitReader(resp.Body, maxCustomResponseBytes))
upstreamResp := strings.TrimSpace(string(bodyBytes))
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return upstreamResp, fmt.Errorf("custom: http status %s", resp.Status)
}
// 部分 Webhook(如企业微信、钉钉)即使业务失败也返回 HTTP 200,
// 仅当响应体包含非零 errcode 时才判定为发送失败,避免审计记录误报成功。
var apiResp struct {
ErrCode int `json:"errcode"`
ErrMsg string `json:"errmsg"`
}
if err := json.Unmarshal(bodyBytes, &apiResp); err == nil && apiResp.ErrCode != 0 {
return upstreamResp, fmt.Errorf("custom: webhook rejected: errcode=%d errmsg=%q", apiResp.ErrCode, apiResp.ErrMsg)
}
return upstreamResp, nil
}
// ValidateConfig 校验自定义配置
func (p *CustomPusher) ValidateConfig(cfg Config) error {
if cfg.URL == "" {
return errors.New("webhook URL is required")
}
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
return errors.New("webhook URL must start with http:// or https://")
}
return nil
}
@@ -0,0 +1,91 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestCustomPusherSend_ResponseBodyErrcode(t *testing.T) {
tests := []struct {
name string
statusCode int
body string
wantErr bool
wantErrMsg string
}{
{
name: "wechat business error returns HTTP 200 with non-zero errcode",
statusCode: http.StatusOK,
body: `{"errcode":93000,"errmsg":"invalid request data"}`,
wantErr: true,
wantErrMsg: "errcode=93000",
},
{
name: "wechat success returns errcode 0",
statusCode: http.StatusOK,
body: `{"errcode":0,"errmsg":"ok"}`,
wantErr: false,
},
{
name: "json response without errcode is tolerated",
statusCode: http.StatusOK,
body: `{"success":true}`,
wantErr: false,
},
{
name: "non-json response body is tolerated",
statusCode: http.StatusOK,
body: "ok",
wantErr: false,
},
{
name: "empty response body is tolerated",
statusCode: http.StatusNoContent,
body: "",
wantErr: false,
},
{
name: "http error status still fails",
statusCode: http.StatusInternalServerError,
body: `{"errcode":0,"errmsg":"ok"}`,
wantErr: true,
wantErrMsg: "http status",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(tt.statusCode)
_, _ = w.Write([]byte(tt.body))
}))
defer srv.Close()
pusher := &CustomPusher{}
upstreamResp, err := pusher.Send(context.Background(),
Config{Channel: "custom", URL: srv.URL},
"",
map[string]any{"title": "t", "content": "c"},
`{"title":"$title","content":"$content"}`,
nil,
)
if tt.wantErr {
require.Error(t, err)
assert.Contains(t, err.Error(), tt.wantErrMsg)
return
}
assert.NoError(t, err)
if tt.body != "" {
assert.Contains(t, upstreamResp, tt.body)
}
})
}
}

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