mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-02 14:56:38 +08:00
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:
@@ -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
|
||||
@@ -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)))
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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), ¤tCfg); 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), ¤tCfg); 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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
)
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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]
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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() {}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
Reference in New Issue
Block a user