diff --git a/.agent/skills/new-api/SKILL.md b/.agent/skills/new-api/SKILL.md
index 9752dec3..6d829fba 100644
--- a/.agent/skills/new-api/SKILL.md
+++ b/.agent/skills/new-api/SKILL.md
@@ -119,8 +119,9 @@ internal/
### 步骤 2:在模块内实现业务逻辑 (`logics.go` / `service.go`)
业务逻辑逻辑应当实现于 `internal/apps/custom/` 目录下:
-- **优先使用纯函数(`logics.go`)**:定义接收 `context.Context` 且不依赖 `*gin.Context` 的函数,易于单元测试。
+- **优先使用纯函数(`logics.go`)**:定义接收 `context.Context` 且不依赖 `*gin.Context` 的函数,易于单元测试与 Worker 复用。参考 `internal/apps/user/logics.go`。
- **有状态服务(`service.go`)**:若需注入依赖(如 DB 连接、外部客户端等),可定义 Service 结构体和构造函数。
+- **跨模块副作用(推送、任务监听等)**:核心业务代码通过 `internal/listener` 发射域事件,禁止直接 `import` push 模块;装配在 `internal/bootstrap` 完成(参见 `push-notification` skill)。
### 步骤 3:编写 HTTP Handler (`routers.go`)
在 `internal/apps/custom/routers.go` 中编写 Handler:
diff --git a/.agent/skills/new-async-task/SKILL.md b/.agent/skills/new-async-task/SKILL.md
index 92774fe4..5305e14d 100644
--- a/.agent/skills/new-async-task/SKILL.md
+++ b/.agent/skills/new-async-task/SKILL.md
@@ -13,8 +13,9 @@ description: "Wavelet 项目专用:新增或修改 Asynq 异步任务、后台
- `internal/task/handler.go`:`TaskHandler`、`TaskResult`、`PayloadValidator`
- `internal/task/meta.go`:`TaskMeta`、`TaskParam`
-- `internal/task/executor.go`:下发、执行、日志、重试
-- `internal/task/handlers/register.go`:Handler 和元数据注册
+- `internal/task/executor.go`:下发、执行、日志、重试、`OnTaskCompleted` 订阅
+- `internal/task/handlers/register.go`:Handler 和元数据注册(由 bootstrap 调用)
+- `internal/bootstrap/bootstrap.go`:任务注册与进程级装配入口
- `internal/task/worker/worker.go`:Worker 路由和队列
- `internal/task/scheduler/scheduler.go`:定时调度
- `internal/apps/admin/task/routers.go`:Admin 任务 API
@@ -48,6 +49,24 @@ description: "Wavelet 项目专用:新增或修改 Asynq 异步任务、后台
- 在 `internal/task/handlers/register.go` 同时注册 Handler 和 `TaskMeta`。
- 不要在其他位置单独注册任务。
+- **禁止**在业务包 `routers.go` 或 `init()` 中调用 `task.RegisterHandler`;统一由 `bootstrap.RegisterTasks()` → `taskhandlers.Register()` 在进程启动时装配。
+- 任务完成钩子(如 push 通知)通过 `task.OnTaskCompleted` 注册,在 `bootstrap.RegisterTaskListeners()` 中装配(Worker/`all` 进程)。
+
+### 进程装配分工
+
+| 进程 | 注册入口 |
+| :--- | :--- |
+| `api` | `cmd/api.go` → `bootstrap.RegisterAPI()`(含 `RegisterTasks`) |
+| `worker` | `worker.StartWorker()` → `bootstrap.RegisterWorker()`(含 `RegisterTasks` + `RegisterTaskListeners`) |
+| `scheduler` | `scheduler.StartScheduler()` → `bootstrap.RegisterScheduler()` |
+| `all` | `cmd/all.go` → `bootstrap.RegisterAll()` |
+
+所有 `Register*` 使用 `sync.Once`,重复调用安全。
+
+### 测试
+
+- 依赖已注册任务类型或 Handler 的测试(如 `internal/apps/admin/task/routers_test.go`),必须在 setup 中显式调用 `bootstrap.RegisterTasks()`。
+- 不得依赖 `init()` 副作用或 import 链触发注册。
## 日志要求
diff --git a/.agent/skills/push-notification/SKILL.md b/.agent/skills/push-notification/SKILL.md
index 4e4f5e52..4e4e0f12 100644
--- a/.agent/skills/push-notification/SKILL.md
+++ b/.agent/skills/push-notification/SKILL.md
@@ -17,7 +17,9 @@ Wavelet 的消息推送机制采用了**元数据驱动 + 统一触发器 + 异
| :--- | :--- | :--- |
| **`pkg/push/`** | 推送基础设施层 | 静态定义、不依赖系统数据库和任何框架。定义了统一接口 `Pusher`、单例 `PusherPool` 和多实现(Lark, Webhook, Email 等),提供配置验证及发送功能。 |
| **`internal/apps/admin/push/`** | 通知服务与后台任务层 | 包含以下核心文件:
1. [events.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/events.go):定义通知事件的结构模型(`NotificationMessage`, `EventMetadata`)、内置事件的动态注册中心(`BuiltInEvents` 及 `RegisterBuiltInEvent` 函数)以及统一触发器类 `EventTrigger`(包括其底层的派发引擎逻辑)。
2. [tasks.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/tasks.go):定义 Asynq 后台异步发送任务、处理器 `PushHandler` 及其校验逻辑,并记录推送历史审计。
3. [routers.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/routers.go):管理端接口,负责获取事件配置列表和更新配置。 |
-| **`internal/apps/admin/push/custom_events/`** | 自定义通知事件包 | 自定义的事件定义和触发函数单独放到此包下,**一个 Go 文件代表一个事件**。例如:
[admin_login.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/custom_events/admin_login.go) 代表管理员登录事件。 |
+| **`internal/apps/admin/push/custom_events/`** | 自定义通知事件包 | 事件元数据定义与 push 侧处理逻辑;**一个 Go 文件代表一个事件**。在 [register.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/custom_events/register.go) 统一装配,禁止 `init()` 副作用。 |
+| **`internal/listener/`** | 域事件分发层 | 核心域发射事件(如 `EmitAdminLoggedIn`),push 在 bootstrap 阶段通过 `OnAdminLoggedIn` 订阅,避免 auth/user 直接依赖 push。 |
+| **`internal/bootstrap/`** | 应用装配根 | `RegisterPushDomainEvents()` 调用 `custom_events.Register()`;`Init` 中执行 `SyncEvents` 将内置事件元数据同步到数据库。 |
| **数据库审计表** | 状态与历史审计 | `w_push_events` 存放每个通知事件的启用状态、启用渠道、发送目标和自定义渲染模板。
`w_push_histories` 存放消息发送记录用于审计。 |
---
@@ -26,8 +28,8 @@ Wavelet 的消息推送机制采用了**元数据驱动 + 统一触发器 + 异
如果某个新业务(如“新用户注册”或“订单创建”)需要带有消息推送功能,请严格按照以下步骤开发:
-### 步骤 1:在 `custom_events/` 中以一个文件声明事件元数据
-在 `internal/apps/admin/push/custom_events/` 下新建一个 Go 文件(如 `user_registered.go`),声明其事件元数据并利用 `init()` 动态注册。
+### 步骤 1:在 `custom_events/` 中声明事件元数据与处理函数
+在 `internal/apps/admin/push/custom_events/` 下新建一个 Go 文件(如 `user_registered.go`),声明 `EventMetadata` 和 push 侧处理函数(组装 body 并调用 `DefaultTrigger.Trigger`)。
```go
package custom_events
@@ -37,10 +39,9 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/apps/admin/push"
- "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/Rain-kl/Wavelet/internal/listener"
)
-// NewUserRegistered is the metadata definition for the user registered event.
var NewUserRegistered = push.EventMetadata{
Key: "user_registered",
Name: "新用户注册提醒",
@@ -52,58 +53,58 @@ var NewUserRegistered = push.EventMetadata{
Description: "当系统有新用户注册成功时,向管理员或指定目标发送通知",
}
-func init() {
- push.RegisterBuiltInEvent(NewUserRegistered)
-}
-```
-
-### 步骤 2:在该事件文件中编写触发封装函数 (Wrapper)
-为了使业务层调用方便且类型安全,在同一个事件 Go 文件中为该事件定义一个封装函数。
-
-> [!NOTE]
-> - 底层框架 `EventTrigger.Trigger` 本身已经内置了**异步 Goroutine 执行**以及 **`context.WithoutCancel(ctx)` 衍生上下文转换**逻辑。
-> - 开发者只需在封装函数中组装数据,并直接调用 `DefaultTrigger.Trigger` 即可,无需在外部手动写 `go func()` 也不需要处理上下文防取消问题,从而通过框架底层强制约束了异步投递行为。
-
-```go
-// TriggerNewUserRegisteredEvent triggers the user registration notification event.
-func TriggerNewUserRegisteredEvent(ctx context.Context, user *model.User) {
- if user == nil {
+func handleUserRegistered(ctx context.Context, event listener.UserRegistered) {
+ if event.User == nil {
return
}
body := map[string]any{
- "user": user,
+ "user": event.User,
"time": time.Now().Format("2006-01-02 15:04:05"),
}
push.DefaultTrigger.Trigger(ctx, NewUserRegistered, body)
}
```
-### 步骤 3:在业务代码中调用触发函数
-在业务逻辑完成处(例如 `internal/apps/user/routers.go` 的注册 Handler 中)导入 `custom_events` 并调用该函数。
+> `EventTrigger.Trigger` 已内置异步 Goroutine 与 `context.WithoutCancel`;处理函数内直接调用即可,无需外层 `go func()`。
+
+### 步骤 2:在 `listener/` 定义域事件并在 `register.go` 装配
+1. 在 `internal/listener/` 新增域事件类型、`Emit*` 与 `On*` 注册函数(参考 [admin_login.go](file:///Users/ryan/DEV/Go/Wavelet/internal/listener/admin_login.go))。
+2. 在 [register.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/custom_events/register.go) 中注册元数据并订阅域事件:
```go
-import (
- "github.com/Rain-kl/Wavelet/internal/apps/admin/push/custom_events"
-)
-
-func Register(c *gin.Context) {
- // ... 注册成功逻辑 ...
-
- // 异步触发通知推送事件
- custom_events.TriggerNewUserRegisteredEvent(ctx, user)
+func Register() {
+ push.RegisterBuiltInEvent(NewUserRegistered)
+ listener.OnUserRegistered(handleUserRegistered)
}
```
-### 步骤 4:在主路由器或初始化模块进行匿名导入以确保注册
-由于事件是在 `custom_events` 的 `init()` 中注册到 `push` 包的,所以应用程序的执行路径(例如 [router.go](file:///Users/ryan/DEV/Go/Wavelet/internal/router/router.go))必须匿名导入 `custom_events` 包,以确保其在程序启动时被加载和初始化。
+**禁止**在 `custom_events` 或 `router` 中使用 `init()` 注册;**禁止**在 `router.go` 空白导入 `custom_events`。
+
+### 步骤 3:在业务代码中发射域事件(不 import push)
+在业务逻辑完成处(如 `internal/apps/user/routers.go`)仅 import `internal/listener` 并发射事件:
```go
-import (
- _ "github.com/Rain-kl/Wavelet/internal/apps/admin/push/custom_events"
-)
+import "github.com/Rain-kl/Wavelet/internal/listener"
+
+func Register(c *gin.Context) {
+ // ... 注册成功逻辑 ...
+ listener.EmitUserRegistered(ctx, user)
+}
```
-当程序启动后,系统在初始化阶段的 `SyncEvents` 流程中会自动将新声明的 `user_registered` 事件元数据插入数据库表 `w_push_events` 中。此后,管理员即可直接在管理端前端界面上为该事件配置推送渠道。
+### 步骤 4:在 bootstrap / cmd 入口显式装配
+新增事件后,确保 `custom_events.Register()` 已被 `bootstrap.RegisterPushDomainEvents()` 调用,且 API/`all` 进程在 `bootstrap.Init` 之前完成注册:
+
+| 进程 | cmd 入口调用 |
+| :--- | :--- |
+| `api` | `bootstrap.RegisterAPI()` → `bootstrap.Init(ctx, Options{API: true})` |
+| `all` | `bootstrap.RegisterAll()` → `bootstrap.Init(ctx, Options{API: true})` |
+| `worker` / `scheduler` | 不注册 push 域事件;仅 `bootstrap.Init` + 各自 `RegisterWorker`/`RegisterScheduler` |
+
+`Init` 中的 `SyncEvents` 会将 `user_registered` 元数据同步到 `w_push_events`,管理员即可在前端配置推送渠道。
+
+### 步骤 5:编写集成测试
+在 `custom_events/` 或 `listener/` 包内添加测试,验证 `Emit*` → handler → `DefaultTrigger.Trigger` 全链路。测试 setup 须显式调用 `custom_events.Register()`(或 `bootstrap.RegisterPushDomainEvents()`)和 `push.SyncEvents`,参考 [admin_login_test.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/custom_events/admin_login_test.go)。
---
@@ -159,3 +160,12 @@ import (
### 1. 禁止绕过统一触发器 (Always Use EventTrigger)
- 所有推送请求必须经过 `EventTrigger.Trigger`,以确保进行“事件是否启用”、“目标渠道过滤”、“全局推送配置读取”及“发送日志审计”等流程。
+
+### 2. 禁止业务模块直接依赖 push (Decouple via listener)
+- `oauth`、`user` 等核心域 **不得** `import` `internal/apps/admin/push` 或 `custom_events`。
+- 跨模块通知必须通过 `internal/listener` 发射域事件;push 在 `custom_events.Register()` 中订阅。
+
+### 3. 禁止 init() 与 router 副作用注册 (Explicit Bootstrap)
+- 不得在 `init()` 中调用 `RegisterBuiltInEvent` 或订阅 listener。
+- 不得在 `router.go` 空白导入 `custom_events` 触发注册。
+- 统一在 `internal/bootstrap` + `internal/cmd` 入口显式装配。
diff --git a/.grok/hooks/dmux-hooks.json b/.grok/hooks/dmux-hooks.json
deleted file mode 100644
index ebc20596..00000000
--- a/.grok/hooks/dmux-hooks.json
+++ /dev/null
@@ -1,35 +0,0 @@
-{
- "description": "dmux pane status hooks for Grok Build",
- "hooks": {
- "Stop": [
- {
- "hooks": [
- {
- "type": "command",
- "command": "node '/Users/ryan/DEV/Go/Wavelet/.dmux/worktrees/handler-model-logics-user/.grok/hooks/dmux-status-hook.cjs'",
- "timeout": 5,
- "env": {
- "DMUX_PANE_ID": "dmux-1781749411085",
- "DMUX_TMUX_PANE_ID": "%2"
- }
- }
- ]
- }
- ],
- "Notification": [
- {
- "hooks": [
- {
- "type": "command",
- "command": "node '/Users/ryan/DEV/Go/Wavelet/.dmux/worktrees/handler-model-logics-user/.grok/hooks/dmux-status-hook.cjs'",
- "timeout": 5,
- "env": {
- "DMUX_PANE_ID": "dmux-1781749411085",
- "DMUX_TMUX_PANE_ID": "%2"
- }
- }
- ]
- }
- ]
- }
-}
\ No newline at end of file
diff --git a/.grok/hooks/dmux-status-hook.cjs b/.grok/hooks/dmux-status-hook.cjs
deleted file mode 100755
index 4e68ef63..00000000
--- a/.grok/hooks/dmux-status-hook.cjs
+++ /dev/null
@@ -1,97 +0,0 @@
-#!/usr/bin/env node
-const fs = require('fs');
-
-function normalizeHookEventName(value) {
- const raw = String(value || '');
- const normalized = raw.trim().toLowerCase().replace(/-/g, '_');
- switch (normalized) {
- case 'stop':
- return 'Stop';
- case 'notification':
- return 'Notification';
- case 'user_prompt_submit':
- case 'userpromptsubmit':
- return 'UserPromptSubmit';
- case 'pre_tool_use':
- case 'pretooluse':
- return 'PreToolUse';
- case 'post_tool_use':
- case 'posttooluse':
- return 'PostToolUse';
- case 'post_tool_use_failure':
- case 'posttoolusefailure':
- return 'PostToolUseFailure';
- case 'session_start':
- case 'sessionstart':
- return 'SessionStart';
- case 'session_end':
- case 'sessionend':
- return 'SessionEnd';
- default:
- return raw;
- }
-}
-
-function stringValue(...values) {
- for (const value of values) {
- if (typeof value === 'string' && value.trim()) {
- return value;
- }
- }
- return '';
-}
-
-let input = '';
-process.stdin.setEncoding('utf8');
-process.stdin.on('data', (chunk) => {
- input += chunk;
-});
-process.stdin.on('end', () => {
- let payload = {};
- try {
- payload = input.trim() ? JSON.parse(input) : {};
- } catch (error) {
- payload = { parse_error: String(error), raw: input };
- }
-
- const hookEventName = normalizeHookEventName(
- payload.hookEventName || payload.hook_event_name || process.env.GROK_HOOK_EVENT
- );
- const sessionId = stringValue(
- payload.sessionId,
- payload.session_id,
- process.env.GROK_SESSION_ID
- );
- const message = stringValue(
- payload.lastAssistantMessage,
- payload.last_assistant_message,
- payload.message,
- payload.notificationMessage,
- payload.notification_message
- );
-
- const event = {
- source: 'grok-status-hook',
- dmuxPaneId: process.env.DMUX_PANE_ID || '',
- tmuxPaneId: process.env.DMUX_TMUX_PANE_ID || '',
- expectedDmuxPaneId: 'dmux-1781749411085',
- expectedTmuxPaneId: '%2',
- hookEventName,
- sessionId,
- turnId: stringValue(payload.turnId, payload.turn_id, sessionId),
- lastAssistantMessage: message || null,
- transcriptPath: stringValue(payload.transcriptPath, payload.transcript_path) || null,
- cwd: stringValue(payload.cwd, payload.workspaceRoot, process.env.GROK_WORKSPACE_ROOT) || process.cwd(),
- timestamp: Date.now()
- };
-
- if (event.dmuxPaneId !== event.expectedDmuxPaneId) {
- process.exit(0);
- }
-
- try {
- fs.writeFileSync('/Users/ryan/DEV/Go/Wavelet/.dmux/worktrees/handler-model-logics-user/.grok/dmux/dmux-1781749411085.json', JSON.stringify(event, null, 2));
- } catch (error) {
- process.exit(0);
- }
-});
diff --git a/AGENTS.md b/AGENTS.md
index 9cd3053b..ce55e5b2 100644
--- a/AGENTS.md
+++ b/AGENTS.md
@@ -39,6 +39,10 @@
- 当 API Handler 发生变化时,更新 Swagger 文档(运行 `make swagger`)。
- 在完成代码开发后必须运行 `make code-check`, 并修复报错。
- 需要缓存或文件管理能力时,必须复用现有平台实现,禁止在业务包中自行创建缓存目录、直接管理缓存文件或重复封装存储后端。
+- 禁止在 `init()` 中注册跨模块集成(任务 Handler、推送内置事件、域事件监听器、任务完成钩子)。统一通过 `internal/bootstrap` 在 `internal/cmd` 入口显式装配。
+- `internal/router/router.go` 的 `Serve()` 仅负责 HTTP 路由与中间件,禁止在其中执行 `SyncEvents`、`InitLogWriter` 等进程级运行时初始化。
+- 核心业务模块(如 `oauth`、`user`)禁止直接 `import` `internal/apps/admin/push` 或 `custom_events` 触发通知;应通过 `internal/listener` 发射域事件,由 push 模块在 bootstrap 阶段订阅。
+- 编写依赖任务注册或推送事件同步的测试时,必须在测试 setup 中显式调用 `bootstrap.RegisterTasks()`、`bootstrap.RegisterPushDomainEvents()` 等,不得依赖 `init()` 副作用。
## 项目介绍
@@ -67,7 +71,8 @@
后端目录:
-- `internal/cmd/`:用于 API、worker、scheduler、root init 的 Cobra 命令。
+- `internal/cmd/`:用于 API、worker、scheduler、root init 的 Cobra 命令。进程启动时在此调用 `bootstrap.Register*` 与 `bootstrap.Init`,再启动 router / worker / scheduler。
+- `internal/bootstrap/`:应用装配根(composition root)。集中注册任务 Handler、推送域事件订阅、任务完成监听器,并执行 `SyncEvents`、ClickHouse 访问日志写入等进程级初始化;所有注册函数使用 `sync.Once` 保证幂等。
- `internal/config/`:Viper 加载和配置结构体。运行时代码应使用 `config.Config..`。
- `internal/router/`:唯一的 HTTP 路由注册点。
- `internal/apps/`:按功能(Feature-based)组织的 HTTP Handler、中间件、内部服务与模块逻辑。移除全局 service 层,模块内部业务逻辑(如验证码业务逻辑管理器 `internal/apps/cap/manager.go`)均收敛于各自模块中;管理端模块位于 `internal/apps/admin/`。
@@ -79,7 +84,7 @@
- `internal/task/`:Asynq 任务框架;参见 `new-async-task` 了解变更。
- `internal/common/`:共享的通用模型及响应(如 `internal/common/response`)、绑定(bind)、常量以及通用错误。
- `internal/util/`:纯底层工具包,无任何 HTTP/数据库框架依赖。
-- `internal/listener/`:事件监听器和消息/Webhook 消费者。
+- `internal/listener/`:域事件分发层。核心域(auth、user 等)在此定义并发射事件(如 `EmitAdminLoggedIn`);运维模块(push、webhook 等)在 bootstrap 阶段订阅,实现跨模块解耦。
- `internal/otel_trace/`:链路追踪(tracing)助手。
- `internal/testhelper/`:后端测试共享辅助能力。
- `internal/buildinfo/`:暴露在发布/构建工作流中注入的元数据(如版本号、编译时间等)。
@@ -139,7 +144,13 @@ Handler 规范:
路由与模块:
- 仅在 `internal/router/router.go` 中作为统一高层入口进行路由分发委派,不允许在 `router.go` 中直接挂载业务 Handler。
-- 关于所有的路由归属划分、接口开发隔离防线以及详细的注册和开发步骤,请直接阅读并严格遵循 [new-api](file:///Users/ryan/DEV/Go/Wavelet/.agent/skills/new-api/SKILL.md) 技能。
+- 关于所有的路由归属划分、接口开发隔离防线以及详细的注册和开发步骤,请直接阅读并严格遵循 [new-api](file:///Users/ryan/DEV/Go/Wavelet/.claude/skills/new-api/SKILL.md) 技能。
+
+应用装配与跨模块集成:
+
+- 新增跨模块副作用(任务注册、推送订阅、后台监听器)时,在 `internal/bootstrap/bootstrap.go` 增加 `Register*` 函数,并在对应 `internal/cmd/*.go` 入口调用;参考现有 `RegisterAPI` / `RegisterWorker` / `RegisterAll` 分工。
+- `bootstrap.Init` 必须在 `RegisterPushDomainEvents()` 之后调用(API/`all` 模式),以确保 `SyncEvents` 能同步内置推送事件元数据。
+- Handler 与业务逻辑分离:HTTP Handler 负责绑定与响应;可复用逻辑放入 `logics.go`(接受 `context.Context`,不依赖 `*gin.Context`),便于 Worker 与单元测试复用。参考 `internal/apps/user/logics.go`。
中间件:
diff --git a/internal/apps/admin/auth_source/routers.go b/internal/apps/admin/auth_source/routers.go
index b5b08159..d18e2c8e 100644
--- a/internal/apps/admin/auth_source/routers.go
+++ b/internal/apps/admin/auth_source/routers.go
@@ -50,7 +50,7 @@ type ToggleAuthSourceRequest struct {
func ListAuthSources(c *gin.Context) {
sources, err := model.GetAuthSources(c.Request.Context())
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(sources))
@@ -72,7 +72,7 @@ func ListAuthSources(c *gin.Context) {
func CreateAuthSource(c *gin.Context) {
var req AuthSourceRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -88,7 +88,7 @@ func CreateAuthSource(c *gin.Context) {
IconURL: req.IconURL,
}
if err := model.CreateAuthSource(c.Request.Context(), &source); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
source.Sanitize()
@@ -113,13 +113,13 @@ func CreateAuthSource(c *gin.Context) {
func UpdateAuthSource(c *gin.Context) {
id, err := parseSourceID(c)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
var req AuthSourceRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -140,7 +140,7 @@ func UpdateAuthSource(c *gin.Context) {
}
keepSecret := source.ClientSecret == ""
if err := model.UpdateAuthSource(c.Request.Context(), &source, keepSecret); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -153,7 +153,7 @@ func UpdateAuthSource(c *gin.Context) {
updated, err := model.GetAuthSourceByID(c.Request.Context(), id)
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
updated.Sanitize()
@@ -177,18 +177,18 @@ func UpdateAuthSource(c *gin.Context) {
func ToggleAuthSource(c *gin.Context) {
id, err := parseSourceID(c)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
var req ToggleAuthSourceRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if err := model.ToggleAuthSource(c.Request.Context(), id, req.IsActive); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
@@ -209,11 +209,11 @@ func ToggleAuthSource(c *gin.Context) {
func DeleteAuthSource(c *gin.Context) {
id, err := parseSourceID(c)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if err := model.DeleteAuthSource(c.Request.Context(), id); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
diff --git a/internal/apps/admin/cache/routers.go b/internal/apps/admin/cache/routers.go
index 97ca546a..3b98a37b 100644
--- a/internal/apps/admin/cache/routers.go
+++ b/internal/apps/admin/cache/routers.go
@@ -56,7 +56,7 @@ func GetCacheStatus(c *gin.Context) {
func UpdateCacheConfig(c *gin.Context) {
var req updateCacheConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -64,19 +64,19 @@ func UpdateCacheConfig(c *gin.Context) {
// Update Max Size
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
// Update Default TTL
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
// Update LRU Enabled
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -99,7 +99,7 @@ func UpdateCacheConfig(c *gin.Context) {
// @Router /api/v1/admin/cache/clear [post]
func ClearCache(c *gin.Context) {
if err := diskcache.GetGlobalCache().Clear(); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
diff --git a/internal/apps/admin/db_manage/routers.go b/internal/apps/admin/db_manage/routers.go
index 897fe9e3..5d2b8f78 100644
--- a/internal/apps/admin/db_manage/routers.go
+++ b/internal/apps/admin/db_manage/routers.go
@@ -213,7 +213,7 @@ func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
func GetDBOverview(c *gin.Context) {
gormDB := db.DB(c.Request.Context())
if gormDB == nil {
- c.JSON(http.StatusInternalServerError, response.Err("数据库未初始化"))
+ response.AbortInternal(c, "数据库未初始化")
return
}
@@ -227,7 +227,7 @@ func GetDBOverview(c *gin.Context) {
}
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -248,7 +248,7 @@ func GetDBOverview(c *gin.Context) {
func ListDBTables(c *gin.Context) {
gormDB := db.DB(c.Request.Context())
if gormDB == nil {
- c.JSON(http.StatusInternalServerError, response.Err("数据库未初始化"))
+ response.AbortInternal(c, "数据库未初始化")
return
}
@@ -262,7 +262,7 @@ func ListDBTables(c *gin.Context) {
}
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -273,13 +273,13 @@ func ListDBTables(c *gin.Context) {
func GetDBTableData(c *gin.Context) {
var req GetTableDataRequest
if err := c.ShouldBindQuery(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
gormDB := db.DB(c.Request.Context())
if gormDB == nil {
- c.JSON(http.StatusInternalServerError, response.Err("数据库未初始化"))
+ response.AbortInternal(c, "数据库未初始化")
return
}
@@ -288,7 +288,7 @@ func GetDBTableData(c *gin.Context) {
var total int64
if err := gormDB.Raw("SELECT count(*) FROM " + quotedTable).Scan(&total).Error; err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -303,7 +303,7 @@ func GetDBTableData(c *gin.Context) {
rows, err := gormDB.Raw("SELECT * FROM "+quotedTable+" LIMIT ? OFFSET ?", limit, offset).Rows()
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
defer func() {
@@ -312,13 +312,13 @@ func GetDBTableData(c *gin.Context) {
cols, err := rows.Columns()
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
results, err := scanTableRows(rows, cols)
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -449,19 +449,19 @@ func executeSQLMutation(gormDB *gorm.DB, sqlStr string, startTime time.Time) (Ex
func ExecuteSQL(c *gin.Context) {
var req ExecuteSQLRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
gormDB := db.DB(c.Request.Context())
if gormDB == nil {
- c.JSON(http.StatusInternalServerError, response.Err("数据库未初始化"))
+ response.AbortInternal(c, "数据库未初始化")
return
}
trimmedSQL := strings.TrimSpace(req.SQL)
if trimmedSQL == "" {
- c.JSON(http.StatusBadRequest, response.Err("SQL 语句不能为空"))
+ response.AbortBadRequest(c, "SQL 语句不能为空")
return
}
@@ -488,7 +488,7 @@ func ExecuteSQL(c *gin.Context) {
}
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
diff --git a/internal/apps/admin/middlewares.go b/internal/apps/admin/middlewares.go
index 06856557..3bd1c25f 100644
--- a/internal/apps/admin/middlewares.go
+++ b/internal/apps/admin/middlewares.go
@@ -5,8 +5,7 @@
package admin
import (
- "net/http"
-
+ "github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
@@ -29,13 +28,13 @@ func LoginAdminRequired() gin.HandlerFunc {
if tokenAuth, _ := util.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth {
tokenAdmin, _ := util.GetFromContext[bool](c, oauth.TokenAdminKey)
if !tokenAdmin {
- c.AbortWithStatusJSON(http.StatusNotFound, gin.H{"error_msg": TokenAdminRequired, "data": nil})
+ response.AbortNotFound(c, TokenAdminRequired)
return
}
}
if !user.IsAdmin {
- c.AbortWithStatusJSON(http.StatusNotFound, gin.H{"error_msg": AdminRequired, "data": nil})
+ response.AbortNotFound(c, AdminRequired)
return
}
diff --git a/internal/apps/admin/push/channels.go b/internal/apps/admin/push/channels.go
index 34c32a7e..36ab1489 100644
--- a/internal/apps/admin/push/channels.go
+++ b/internal/apps/admin/push/channels.go
@@ -41,7 +41,7 @@ func ListChannels(c *gin.Context) {
ctx := c.Request.Context()
var channels []model.PushChannel
if err := db.DB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(channels))
@@ -71,7 +71,7 @@ type CreateChannelRequest struct {
func CreateChannel(c *gin.Context) {
var req CreateChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -79,11 +79,11 @@ func CreateChannel(c *gin.Context) {
var count int64
if err := db.DB(ctx).Model(&model.PushChannel{}).Where("name = ?", req.Name).Count(&count).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
if count > 0 {
- c.JSON(http.StatusBadRequest, response.Err("channel name already exists"))
+ response.AbortBadRequest(c, "channel name already exists")
return
}
@@ -98,12 +98,12 @@ func CreateChannel(c *gin.Context) {
}
if err := channel.Validate(); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if err := db.DB(ctx).Create(&channel).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -138,13 +138,13 @@ func UpdateChannel(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err("invalid channel id"))
+ response.AbortBadRequest(c, "invalid channel id")
return
}
var req UpdateChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -153,10 +153,10 @@ func UpdateChannel(c *gin.Context) {
var channel model.PushChannel
if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusNotFound, response.Err("channel not found"))
+ response.AbortNotFound(c, "channel not found")
return
}
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -168,12 +168,12 @@ func UpdateChannel(c *gin.Context) {
channel.Enabled = req.Enabled
if err := channel.Validate(); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if err := db.DB(ctx).Save(&channel).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -196,7 +196,7 @@ func DeleteChannel(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err("invalid channel id"))
+ response.AbortBadRequest(c, "invalid channel id")
return
}
@@ -204,15 +204,15 @@ func DeleteChannel(c *gin.Context) {
var channel model.PushChannel
if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusNotFound, response.Err("channel not found"))
+ response.AbortNotFound(c, "channel not found")
return
}
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
if err := db.DB(ctx).Delete(&channel).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -245,7 +245,7 @@ type TestChannelRequest struct {
func TestChannel(c *gin.Context) {
var req TestChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -255,7 +255,7 @@ func TestChannel(c *gin.Context) {
if req.Name != "" {
var channel model.PushChannel
if err := db.DB(ctx).Where("name = ?", req.Name).First(&channel).Error; err != nil {
- c.JSON(http.StatusBadRequest, response.Err("channel not found"))
+ response.AbortBadRequest(c, "channel not found")
return
}
url = channel.URL
@@ -284,7 +284,7 @@ func TestChannel(c *gin.Context) {
}
if err := tempChannel.Validate(); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
url = tempChannel.URL
@@ -342,7 +342,7 @@ func TestChannel(c *gin.Context) {
}
if err := enqueuePushTask(ctx, payload); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
diff --git a/internal/apps/admin/push/custom_events/admin_login_test.go b/internal/apps/admin/push/custom_events/admin_login_test.go
new file mode 100644
index 00000000..b9c54f81
--- /dev/null
+++ b/internal/apps/admin/push/custom_events/admin_login_test.go
@@ -0,0 +1,175 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package custom_events
+
+import (
+ "context"
+ "encoding/json"
+ "sync"
+ "testing"
+ "time"
+
+ "github.com/Rain-kl/Wavelet/internal/apps/admin/push"
+ "github.com/Rain-kl/Wavelet/internal/listener"
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/Rain-kl/Wavelet/internal/task"
+ "github.com/Rain-kl/Wavelet/internal/testhelper"
+ "github.com/hibiken/asynq"
+ "github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
+ "gorm.io/gorm"
+)
+
+var registerOnce sync.Once
+
+func ensureRegistered() {
+ registerOnce.Do(Register)
+}
+
+func setupAdminLoginIntegrationTest(t *testing.T) (*gorm.DB, func()) {
+ t.Helper()
+
+ dbConn, mr, cleanup := testhelper.SetupTestEnvironment(t)
+
+ err := dbConn.AutoMigrate(
+ &model.PushEvent{},
+ &model.PushHistory{},
+ &model.PushChannel{},
+ )
+ require.NoError(t, err)
+
+ sysUser := &model.User{
+ ID: 999,
+ Username: "system",
+ Nickname: "系统",
+ Password: "*",
+ IsActive: true,
+ }
+ require.NoError(t, dbConn.Create(sysUser).Error)
+
+ task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{Addr: mr.Addr()})
+ task.RegisterHandler(push.SendNotificationTask, &push.PushHandler{})
+ task.RegisterTaskMeta(push.SendNotificationMeta)
+
+ ensureRegistered()
+
+ require.NoError(t, push.SyncEvents(context.Background()))
+
+ return dbConn, func() {
+ cleanup()
+ if task.AsynqClient != nil {
+ task.AsynqClient.Close()
+ task.AsynqClient = nil
+ }
+ }
+}
+
+func seedMockPushChannel(t *testing.T, dbConn *gorm.DB) *model.PushChannel {
+ t.Helper()
+
+ channel := &model.PushChannel{
+ Name: "mock_channel",
+ Type: "custom",
+ URL: "https://webhook.site/admin-login",
+ Other: `{"text": "$content"}`,
+ Enabled: true,
+ }
+ require.NoError(t, dbConn.Create(channel).Error)
+ return channel
+}
+
+func enableAdminLoginEvent(t *testing.T, dbConn *gorm.DB, channelName string, targets []string) {
+ t.Helper()
+
+ var event model.PushEvent
+ require.NoError(t, dbConn.Where("event_key = ?", AdminLogin.Key).First(&event).Error)
+
+ event.Enabled = true
+ event.Channels = []string{channelName}
+ event.Targets = targets
+ require.NoError(t, dbConn.Save(&event).Error)
+}
+
+func waitForAsyncTrigger(t *testing.T) {
+ t.Helper()
+ time.Sleep(100 * time.Millisecond)
+}
+
+func countPushTasks(t *testing.T, dbConn *gorm.DB) int64 {
+ t.Helper()
+
+ var count int64
+ require.NoError(t, dbConn.Model(&model.TaskExecution{}).
+ Where("task_type = ?", push.SendNotificationTask).
+ Count(&count).Error)
+ return count
+}
+
+func TestAdminLoginPushIntegration(t *testing.T) {
+ dbConn, cleanup := setupAdminLoginIntegrationTest(t)
+ defer cleanup()
+
+ channel := seedMockPushChannel(t, dbConn)
+ defer dbConn.Delete(channel)
+
+ enableAdminLoginEvent(t, dbConn, channel.Name, []string{"ops_team"})
+
+ adminUser := &model.User{
+ ID: 1001,
+ Username: "super_admin",
+ IsAdmin: true,
+ IsActive: true,
+ }
+ require.NoError(t, dbConn.Create(adminUser).Error)
+
+ t.Run("admin login emits push task with user and ip", func(t *testing.T) {
+ dbConn.Where("task_type = ?", push.SendNotificationTask).Delete(&model.TaskExecution{})
+
+ listener.EmitAdminLoggedIn(context.Background(), adminUser, "203.0.113.42")
+ waitForAsyncTrigger(t)
+
+ var execution model.TaskExecution
+ require.NoError(t, dbConn.Where("task_type = ?", push.SendNotificationTask).First(&execution).Error)
+
+ var payload push.SendPayload
+ require.NoError(t, json.Unmarshal([]byte(execution.Payload), &payload))
+
+ assert.Equal(t, AdminLogin.Key, payload.EventKey)
+ assert.Equal(t, "ops_team", payload.Target)
+ assert.Equal(t, "管理员登录提醒", payload.Body.Title)
+ assert.Contains(t, payload.Body.Content, "super_admin")
+ assert.Contains(t, payload.Body.Content, "203.0.113.42")
+ })
+
+ t.Run("non-admin login does not trigger push", func(t *testing.T) {
+ dbConn.Where("task_type = ?", push.SendNotificationTask).Delete(&model.TaskExecution{})
+
+ nonAdmin := &model.User{
+ ID: 2002,
+ Username: "regular_user",
+ IsAdmin: false,
+ IsActive: true,
+ }
+ require.NoError(t, dbConn.Create(nonAdmin).Error)
+
+ listener.EmitAdminLoggedIn(context.Background(), nonAdmin, "198.51.100.1")
+ waitForAsyncTrigger(t)
+
+ assert.Equal(t, int64(0), countPushTasks(t, dbConn))
+ })
+
+ t.Run("disabled admin login event does not enqueue push", func(t *testing.T) {
+ dbConn.Where("task_type = ?", push.SendNotificationTask).Delete(&model.TaskExecution{})
+
+ var event model.PushEvent
+ require.NoError(t, dbConn.Where("event_key = ?", AdminLogin.Key).First(&event).Error)
+ event.Enabled = false
+ require.NoError(t, dbConn.Save(&event).Error)
+
+ listener.EmitAdminLoggedIn(context.Background(), adminUser, "10.0.0.1")
+ waitForAsyncTrigger(t)
+
+ assert.Equal(t, int64(0), countPushTasks(t, dbConn))
+ })
+}
\ No newline at end of file
diff --git a/internal/apps/admin/push/routers.go b/internal/apps/admin/push/routers.go
index 8858c5f7..1cc50d0c 100644
--- a/internal/apps/admin/push/routers.go
+++ b/internal/apps/admin/push/routers.go
@@ -76,7 +76,7 @@ func ListEvents(c *gin.Context) {
var events []model.PushEvent
if err := db.DB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(events))
@@ -166,7 +166,7 @@ func getEventInfo(req CreateEventRequest) (string, string, []byte, error) {
func CreateEvent(c *gin.Context) {
var req CreateEventRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -174,18 +174,18 @@ func CreateEvent(c *gin.Context) {
eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
// 2. 检查是否已经创建过该事件的配置
var count int64
if err := db.DB(ctx).Model(&model.PushEvent{}).Where("event_key = ?", eventKey).Count(&count).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
if count > 0 {
- c.JSON(http.StatusBadRequest, response.Err("this notification event is already configured"))
+ response.AbortBadRequest(c, "this notification event is already configured")
return
}
@@ -196,7 +196,7 @@ func CreateEvent(c *gin.Context) {
} else {
var tempMap map[string]any
if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil {
- c.JSON(http.StatusBadRequest, response.Err("custom template is not a valid JSON format"))
+ response.AbortBadRequest(c, "custom template is not a valid JSON format")
return
}
}
@@ -222,12 +222,12 @@ func CreateEvent(c *gin.Context) {
}
if err := event.Validate(); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if err := db.DB(ctx).Create(&event).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -250,7 +250,7 @@ func DeleteEvent(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err("invalid event id"))
+ response.AbortBadRequest(c, "invalid event id")
return
}
@@ -258,15 +258,15 @@ func DeleteEvent(c *gin.Context) {
var event model.PushEvent
if err := db.DB(ctx).First(&event, id).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusNotFound, response.Err("notification event not found"))
+ response.AbortNotFound(c, "notification event not found")
} else {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
}
return
}
if err := db.DB(ctx).Delete(&event).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -291,22 +291,22 @@ func UpdateEvent(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err("invalid event id"))
+ response.AbortBadRequest(c, "invalid event id")
return
}
var req UpdateEventRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
var event model.PushEvent
if err := db.DB(c.Request.Context()).First(&event, id).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusNotFound, response.Err("notification event not found"))
+ response.AbortNotFound(c, "notification event not found")
} else {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
}
return
}
@@ -317,12 +317,12 @@ func UpdateEvent(c *gin.Context) {
event.Enabled = req.Enabled
if err := event.Validate(); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if err := db.DB(c.Request.Context()).Save(&event).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -345,27 +345,27 @@ func ToggleEvent(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err("invalid event id"))
+ response.AbortBadRequest(c, "invalid event id")
return
}
var event model.PushEvent
if err := db.DB(c.Request.Context()).First(&event, id).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusNotFound, response.Err("notification event not found"))
+ response.AbortNotFound(c, "notification event not found")
} else {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
}
return
}
event.Enabled = !event.Enabled
if event.Enabled && len(event.Channels) == 0 {
- c.JSON(http.StatusBadRequest, response.Err("cannot enable event without any push channels configured"))
+ response.AbortBadRequest(c, "cannot enable event without any push channels configured")
return
}
if err := db.DB(c.Request.Context()).Model(&event).Update("enabled", event.Enabled).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -420,14 +420,14 @@ func ListHistories(c *gin.Context) {
var total int64
if err := query.Count(&total).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
var results []model.PushHistory
offset := (page - 1) * pageSize
if err := query.Offset(offset).Limit(pageSize).Find(&results).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -450,19 +450,19 @@ func ListHistories(c *gin.Context) {
func TestPush(c *gin.Context) {
var req TestPushRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
pusher, err := push.GetPusher(req.Config.Channel)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
// 校验配置
if err := pusher.ValidateConfig(req.Config); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(fmt.Sprintf("validation failed: %v", err)))
+ response.AbortBadRequest(c, fmt.Sprintf("validation failed: %v", err))
return
}
@@ -494,7 +494,7 @@ func TestPush(c *gin.Context) {
err = pusher.Send(c.Request.Context(), req.Config, req.Target, testBody, "", nil)
if err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
diff --git a/internal/apps/admin/status/routers.go b/internal/apps/admin/status/routers.go
index 2890ada8..513d701b 100644
--- a/internal/apps/admin/status/routers.go
+++ b/internal/apps/admin/status/routers.go
@@ -290,7 +290,7 @@ func exportSQLite(c *gin.Context) {
f, err := os.Open(path) //nolint:gosec // path is loaded from server startup configuration, not user input
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err("无法打开数据库文件: "+err.Error()))
+ response.AbortInternal(c, "无法打开数据库文件: "+err.Error())
return
}
defer func() {
@@ -301,7 +301,7 @@ func exportSQLite(c *gin.Context) {
fi, err := f.Stat()
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err("无法读取数据库文件信息: "+err.Error()))
+ response.AbortInternal(c, "无法读取数据库文件信息: "+err.Error())
return
}
@@ -319,7 +319,7 @@ func exportPostgres(c *gin.Context) {
// 检查 pg_dump 是否可用
pgDumpPath, err := exec.LookPath("pg_dump")
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err("pg_dump 不可用,请确保服务器已安装 PostgreSQL 客户端工具"))
+ response.AbortInternal(c, "pg_dump 不可用,请确保服务器已安装 PostgreSQL 客户端工具")
return
}
diff --git a/internal/apps/admin/system_config/routers.go b/internal/apps/admin/system_config/routers.go
index 1efc0c6c..89f15096 100644
--- a/internal/apps/admin/system_config/routers.go
+++ b/internal/apps/admin/system_config/routers.go
@@ -60,17 +60,17 @@ type UpdateSystemConfigRequest struct {
func CreateSystemConfig(c *gin.Context) {
var req CreateSystemConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
// 检查配置键是否已存在
var existing model.SystemConfig
if err := db.DB(c.Request.Context()).Where("key = ?", req.Key).First(&existing).Error; err == nil {
- c.JSON(http.StatusBadRequest, response.Err(ConfigKeyExists))
+ response.AbortBadRequest(c, ConfigKeyExists)
return
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -90,7 +90,7 @@ func CreateSystemConfig(c *gin.Context) {
return nil
}); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -124,7 +124,7 @@ func ListSystemConfigs(c *gin.Context) {
var configs []model.SystemConfig
if err := query.Find(&configs).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -152,9 +152,9 @@ func GetSystemConfig(c *gin.Context) {
var config model.SystemConfig
if err := db.DB(c.Request.Context()).Where("key = ?", c.Param("key")).First(&config).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusNotFound, response.Err(SystemConfigNotFound))
+ response.AbortNotFound(c, SystemConfigNotFound)
} else {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
}
return
}
@@ -183,7 +183,7 @@ func GetSystemConfig(c *gin.Context) {
func UpdateSystemConfig(c *gin.Context) {
var req UpdateSystemConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -193,9 +193,9 @@ func UpdateSystemConfig(c *gin.Context) {
var config model.SystemConfig
if err := db.DB(c.Request.Context()).Where("key = ?", key).First(&config).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusNotFound, response.Err(SystemConfigNotFound))
+ response.AbortNotFound(c, SystemConfigNotFound)
} else {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
}
return
}
@@ -209,7 +209,7 @@ func UpdateSystemConfig(c *gin.Context) {
validatedVal, err := validateAndMergeStorageConfig(c.Request.Context(), req.Value, config.Value)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
req.Value = validatedVal
@@ -242,7 +242,7 @@ func UpdateSystemConfig(c *gin.Context) {
return nil
}); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -339,7 +339,7 @@ type TestSMTPResponse struct {
func TestSMTP(c *gin.Context) {
var req TestSMTPRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
diff --git a/internal/apps/admin/task/routers.go b/internal/apps/admin/task/routers.go
index 974fd762..54965e33 100644
--- a/internal/apps/admin/task/routers.go
+++ b/internal/apps/admin/task/routers.go
@@ -60,13 +60,13 @@ type DispatchTaskRequest struct {
func DispatchTask(c *gin.Context) {
var req DispatchTaskRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
meta := task.GetTaskMeta(req.TaskType)
if meta == nil {
- c.JSON(http.StatusBadRequest, response.Err(InvalidTaskType))
+ response.AbortBadRequest(c, InvalidTaskType)
return
}
@@ -77,13 +77,13 @@ func DispatchTask(c *gin.Context) {
validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
taskID, err := task.DispatchTask(c.Request.Context(), req.TaskType, validated, "manual")
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(fmt.Sprintf("%s: %v", TaskDispatchFailed, err)))
+ response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err))
return
}
@@ -107,7 +107,7 @@ func DispatchTask(c *gin.Context) {
func ListTaskExecutions(c *gin.Context) {
var req model.ListTaskExecutionsRequest
if err := c.ShouldBindQuery(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -119,7 +119,7 @@ func ListTaskExecutions(c *gin.Context) {
executions, total, err := model.ListTaskExecutions(c.Request.Context(), req)
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -147,13 +147,13 @@ func ListTaskExecutions(c *gin.Context) {
func GetTaskExecution(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(admin.InvalidTaskExecutionID))
+ response.AbortBadRequest(c, admin.InvalidTaskExecutionID)
return
}
execution, err := model.GetTaskExecutionByID(c.Request.Context(), id)
if err != nil {
- c.JSON(http.StatusNotFound, response.Err(TaskNotFound))
+ response.AbortNotFound(c, TaskNotFound)
return
}
@@ -177,7 +177,7 @@ func GetTaskExecution(c *gin.Context) {
func RetryTask(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(admin.InvalidTaskExecutionID))
+ response.AbortBadRequest(c, admin.InvalidTaskExecutionID)
return
}
@@ -186,11 +186,11 @@ func RetryTask(c *gin.Context) {
errMsg := err.Error()
switch {
case strings.Contains(errMsg, "不存在"):
- c.JSON(http.StatusNotFound, response.Err(errMsg))
+ response.AbortNotFound(c, errMsg)
case strings.Contains(errMsg, "只有失败的任务") || strings.Contains(errMsg, "不支持重试") || strings.Contains(errMsg, "已达到最大重试"):
- c.JSON(http.StatusBadRequest, response.Err(errMsg))
+ response.AbortBadRequest(c, errMsg)
default:
- c.JSON(http.StatusInternalServerError, response.Err(fmt.Sprintf("%s: %v", TaskRetryFailed, err)))
+ response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskRetryFailed, err))
}
return
}
@@ -211,7 +211,7 @@ func RetryTask(c *gin.Context) {
func ListSchedules(c *gin.Context) {
schedules, err := model.ListSchedules(c.Request.Context())
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(schedules))
@@ -243,20 +243,20 @@ type CreateScheduleRequest struct {
func CreateSchedule(c *gin.Context) {
var req CreateScheduleRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
// 校验 Cron 表达式
if _, err := cron.ParseStandard(req.Cron); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(InvalidCronExpression))
+ response.AbortBadRequest(c, InvalidCronExpression)
return
}
// 校验关联的异步任务类型
meta := task.GetTaskMeta(req.TaskType)
if meta == nil {
- c.JSON(http.StatusBadRequest, response.Err(InvalidTaskType))
+ response.AbortBadRequest(c, InvalidTaskType)
return
}
@@ -267,7 +267,7 @@ func CreateSchedule(c *gin.Context) {
}
validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -280,7 +280,7 @@ func CreateSchedule(c *gin.Context) {
}
if err := model.CreateSchedule(c.Request.Context(), schedule); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(fmt.Sprintf("%s: %v", ScheduleSaveFailed, err)))
+ response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
@@ -320,33 +320,33 @@ type UpdateScheduleRequest struct {
func UpdateSchedule(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err("无效的定时任务ID"))
+ response.AbortBadRequest(c, "无效的定时任务ID")
return
}
var req UpdateScheduleRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
// 检查定时任务是否存在
schedule, err := model.GetScheduleByID(c.Request.Context(), id)
if err != nil {
- c.JSON(http.StatusNotFound, response.Err(ScheduleNotFound))
+ response.AbortNotFound(c, ScheduleNotFound)
return
}
// 校验 Cron 表达式
if _, err := cron.ParseStandard(req.Cron); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(InvalidCronExpression))
+ response.AbortBadRequest(c, InvalidCronExpression)
return
}
// 校验关联的异步任务类型
meta := task.GetTaskMeta(req.TaskType)
if meta == nil {
- c.JSON(http.StatusBadRequest, response.Err(InvalidTaskType))
+ response.AbortBadRequest(c, InvalidTaskType)
return
}
@@ -357,7 +357,7 @@ func UpdateSchedule(c *gin.Context) {
}
validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -368,7 +368,7 @@ func UpdateSchedule(c *gin.Context) {
schedule.IsActive = *req.IsActive
if err := model.UpdateSchedule(c.Request.Context(), schedule); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(fmt.Sprintf("%s: %v", ScheduleSaveFailed, err)))
+ response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
@@ -396,12 +396,12 @@ func UpdateSchedule(c *gin.Context) {
func DeleteSchedule(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
- c.JSON(http.StatusBadRequest, response.Err("无效的定时任务ID"))
+ response.AbortBadRequest(c, "无效的定时任务ID")
return
}
if err := model.DeleteSchedule(c.Request.Context(), id); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err)))
+ response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))
return
}
diff --git a/internal/apps/admin/task/routers_test.go b/internal/apps/admin/task/routers_test.go
index 2102c022..5c118ff4 100644
--- a/internal/apps/admin/task/routers_test.go
+++ b/internal/apps/admin/task/routers_test.go
@@ -16,6 +16,7 @@ import ("bytes"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/apps/user"
+ "github.com/Rain-kl/Wavelet/internal/bootstrap"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/testhelper"
@@ -29,6 +30,7 @@ import ("bytes"
func setupTaskTestEnvironment(t *testing.T) func() {
_, mr, cleanup := testhelper.SetupTestEnvironment(t)
+ bootstrap.RegisterTasks()
task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{
Addr: mr.Addr(),
})
diff --git a/internal/apps/admin/template/routers.go b/internal/apps/admin/template/routers.go
index 7580aa59..09d0dc0d 100644
--- a/internal/apps/admin/template/routers.go
+++ b/internal/apps/admin/template/routers.go
@@ -49,17 +49,17 @@ type UpdateTemplateRequest struct {
func CreateTemplate(c *gin.Context) {
var req CreateTemplateRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
// 检查模板 Key 是否已存在
var existing model.Template
if err := db.DB(c.Request.Context()).Where("key = ?", req.Key).First(&existing).Error; err == nil {
- c.JSON(http.StatusBadRequest, response.Err(TemplateKeyExists))
+ response.AbortBadRequest(c, TemplateKeyExists)
return
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -74,12 +74,12 @@ func CreateTemplate(c *gin.Context) {
}
if err := tmpl.Validate(); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if err := db.DB(c.Request.Context()).Create(&tmpl).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -100,7 +100,7 @@ func CreateTemplate(c *gin.Context) {
func ListTemplates(c *gin.Context) {
var templates []model.Template
if err := db.DB(c.Request.Context()).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -124,9 +124,9 @@ func GetTemplate(c *gin.Context) {
var tmpl model.Template
if err := db.DB(c.Request.Context()).Where("key = ?", c.Param("key")).First(&tmpl).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusNotFound, response.Err(TemplateNotFound))
+ response.AbortNotFound(c, TemplateNotFound)
} else {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
}
return
}
@@ -153,7 +153,7 @@ func GetTemplate(c *gin.Context) {
func UpdateTemplate(c *gin.Context) {
var req UpdateTemplateRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -163,9 +163,9 @@ func UpdateTemplate(c *gin.Context) {
var tmpl model.Template
if err := db.DB(c.Request.Context()).Where("key = ?", key).First(&tmpl).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusNotFound, response.Err(TemplateNotFound))
+ response.AbortNotFound(c, TemplateNotFound)
} else {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
}
return
}
@@ -177,12 +177,12 @@ func UpdateTemplate(c *gin.Context) {
tmpl.Description = req.Description
if err := tmpl.Validate(); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if err := db.DB(c.Request.Context()).Save(&tmpl).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -210,21 +210,21 @@ func DeleteTemplate(c *gin.Context) {
var tmpl model.Template
if err := db.DB(c.Request.Context()).Where("key = ?", key).First(&tmpl).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusNotFound, response.Err(TemplateNotFound))
+ response.AbortNotFound(c, TemplateNotFound)
} else {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
}
return
}
// 限制系统模板删除
if tmpl.IsSystem {
- c.JSON(http.StatusBadRequest, response.Err(SystemTemplateCannotDelete))
+ response.AbortBadRequest(c, SystemTemplateCannotDelete)
return
}
if err := db.DB(c.Request.Context()).Delete(&tmpl).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
diff --git a/internal/apps/admin/updater/routers.go b/internal/apps/admin/updater/routers.go
index 3e5ba0da..d96c447e 100644
--- a/internal/apps/admin/updater/routers.go
+++ b/internal/apps/admin/updater/routers.go
@@ -27,7 +27,7 @@ func GetUpdateStatus(c *gin.Context) {
status, _, err := defaultManager.status(c.Request.Context())
if err != nil {
logger.ErrorF(c.Request.Context(), "[Updater] check release failed: %v", err)
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(status))
@@ -49,7 +49,7 @@ func ApplyUpdate(c *gin.Context) {
executable, stagedBinary, err := defaultManager.prepareUpgrade(c.Request.Context())
if err != nil {
logger.ErrorF(c.Request.Context(), "[Updater] prepare upgrade failed: %v", err)
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
diff --git a/internal/apps/admin/user/routers.go b/internal/apps/admin/user/routers.go
index e6ef0081..509882d9 100644
--- a/internal/apps/admin/user/routers.go
+++ b/internal/apps/admin/user/routers.go
@@ -59,7 +59,7 @@ type listUsersResponse struct {
func parseUserID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil || id == 0 {
- c.JSON(http.StatusBadRequest, response.Err(userNotFound))
+ response.AbortBadRequest(c, userNotFound)
return 0, false
}
return id, true
@@ -102,7 +102,7 @@ func toUser(u model.User) user {
func ListUsers(c *gin.Context) {
var req listUsersRequest
if err := c.ShouldBindQuery(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -122,7 +122,7 @@ func ListUsers(c *gin.Context) {
}
if err := query.Count(&total).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -134,7 +134,7 @@ func ListUsers(c *gin.Context) {
Offset(offset).
Limit(req.PageSize).
Find(&modelUsers).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -176,10 +176,10 @@ func GetUser(c *gin.Context) {
Where("id = ?", id).
First(&targetUser).Error; err != nil {
if err == gorm.ErrRecordNotFound {
- c.JSON(http.StatusNotFound, response.Err(userNotFound))
+ response.AbortNotFound(c, userNotFound)
return
}
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
@@ -210,7 +210,7 @@ type updateUserStatusRequest struct {
func UpdateUserStatus(c *gin.Context) {
var req updateUserStatusRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -229,15 +229,15 @@ func UpdateUserStatus(c *gin.Context) {
Where("id = ?", id).
First(&targetUser).Error; err != nil {
if err == gorm.ErrRecordNotFound {
- c.JSON(http.StatusNotFound, response.Err(userNotFound))
+ response.AbortNotFound(c, userNotFound)
return
}
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
if !req.IsActive && targetUser.IsAdmin {
- c.JSON(http.StatusForbidden, response.Err(cannotDisable))
+ response.AbortForbidden(c, cannotDisable)
return
}
@@ -245,7 +245,7 @@ func UpdateUserStatus(c *gin.Context) {
Model(&model.User{}).
Where("id = ?", id).
Update("is_active", req.IsActive).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(updateUserFailed))
+ response.AbortInternal(c, updateUserFailed)
return
}
@@ -274,7 +274,7 @@ func DeleteUser(c *gin.Context) {
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
if currUser != nil && currUser.ID == id {
- c.JSON(http.StatusForbidden, response.Err(cannotDeleteSelf))
+ response.AbortForbidden(c, cannotDeleteSelf)
return
}
@@ -288,15 +288,15 @@ func DeleteUser(c *gin.Context) {
Where("id = ?", id).
First(&targetUser).Error; err != nil {
if err == gorm.ErrRecordNotFound {
- c.JSON(http.StatusNotFound, response.Err(userNotFound))
+ response.AbortNotFound(c, userNotFound)
return
}
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
if targetUser.IsAdmin {
- c.JSON(http.StatusForbidden, response.Err(cannotDelete))
+ response.AbortForbidden(c, cannotDelete)
return
}
@@ -309,7 +309,7 @@ func DeleteUser(c *gin.Context) {
}
return tx.Where("id = ?", id).Delete(&model.User{}).Error
}); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(deleteUserFailed))
+ response.AbortInternal(c, deleteUserFailed)
return
}
@@ -343,7 +343,7 @@ type createUserRequest struct {
func CreateUser(c *gin.Context) {
var req createUserRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -353,36 +353,36 @@ func CreateUser(c *gin.Context) {
req.Email = strings.TrimSpace(req.Email)
if req.Username == "" {
- c.JSON(http.StatusBadRequest, response.Err(usernameRequired))
+ response.AbortBadRequest(c, usernameRequired)
return
}
if req.Email == "" {
- c.JSON(http.StatusBadRequest, response.Err(emailRequired))
+ response.AbortBadRequest(c, emailRequired)
return
}
if len(req.Password) < minPasswordLength {
- c.JSON(http.StatusBadRequest, response.Err(passwordTooShort))
+ response.AbortBadRequest(c, passwordTooShort)
return
}
ctx := c.Request.Context()
var count int64
if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", req.Username).Count(&count).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
if count > 0 {
- c.JSON(http.StatusBadRequest, response.Err(usernameExists))
+ response.AbortBadRequest(c, usernameExists)
return
}
var emailCount int64
if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", req.Email).Count(&emailCount).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
if emailCount > 0 {
- c.JSON(http.StatusBadRequest, response.Err(emailExists))
+ response.AbortBadRequest(c, emailExists)
return
}
@@ -400,12 +400,12 @@ func CreateUser(c *gin.Context) {
}
if err := newUser.SetEncryptedPassword(req.Password); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
if err := db.DB(ctx).Create(&newUser).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
diff --git a/internal/apps/cap/middleware.go b/internal/apps/cap/middleware.go
index dfe97f78..e8039bf4 100644
--- a/internal/apps/cap/middleware.go
+++ b/internal/apps/cap/middleware.go
@@ -4,8 +4,6 @@
package cap
import (
- "net/http"
-
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
@@ -21,13 +19,13 @@ func VerifyMiddleware(mgr *Manager, scope string) gin.HandlerFunc {
token := c.GetHeader("X-Cap-Token")
if token == "" {
- c.AbortWithStatusJSON(http.StatusUnauthorized, response.Err(errCapTokenMissing))
+ response.AbortUnauthorized(c, errCapTokenMissing)
return
}
valid, err := mgr.VerifyToken(c.Request.Context(), token, scope)
if err != nil || !valid {
- c.AbortWithStatusJSON(http.StatusUnauthorized, response.Err(errCapTokenInvalidOrExpired))
+ response.AbortUnauthorized(c, errCapTokenInvalidOrExpired)
return
}
diff --git a/internal/apps/config/routers.go b/internal/apps/config/routers.go
index 8f693cb0..7226e7dc 100644
--- a/internal/apps/config/routers.go
+++ b/internal/apps/config/routers.go
@@ -24,7 +24,7 @@ func GetPublicConfig(c *gin.Context) {
ctx := c.Request.Context()
configs, err := model.ListVisibleSystemConfigs(ctx)
if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
diff --git a/internal/apps/oauth/auth_source_resolver.go b/internal/apps/oauth/auth_source_resolver.go
new file mode 100644
index 00000000..58dc7218
--- /dev/null
+++ b/internal/apps/oauth/auth_source_resolver.go
@@ -0,0 +1,117 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package oauth
+
+import (
+ "context"
+ "errors"
+ "strings"
+
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/coreos/go-oidc/v3/oidc"
+ "golang.org/x/oauth2"
+)
+
+func isOIDCLoginEnabled(ctx context.Context) bool {
+ enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
+ if err != nil {
+ return true
+ }
+ return enabled
+}
+
+func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) {
+ name := strings.TrimSpace(strings.ToLower(sourceName))
+ if name == "" {
+ sources, err := model.GetActiveAuthSources(ctx)
+ if err != nil {
+ return nil, err
+ }
+ if len(sources) == 0 {
+ return nil, errors.New(errNoActiveAuthSource)
+ }
+ return &sources[0], nil
+ }
+ return model.GetAuthSourceByName(ctx, name)
+}
+
+func activeLoginSources(ctx context.Context) []AuthSourceView {
+ enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
+ if err == nil && !enabled {
+ return nil
+ }
+
+ dbSources, err := model.GetActiveAuthSources(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) {
+ var sc model.SystemConfig
+ if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || strings.TrimSpace(sc.Value) == "" {
+ return "", errors.New(errServerAddressMissing)
+ }
+ return strings.TrimRight(sc.Value, "/") + "/login", nil
+}
+
+func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
+ if source == nil {
+ return nil, nil, errors.New(errAuthSourceRequired)
+ }
+
+ if source.OpenIDDiscoveryURL == "" {
+ return nil, nil, errors.New(errDiscoveryURLRequired)
+ }
+
+ // Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake)
+ 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,避免每次调用都向 issuer 发起
+ // /.well-known/openid-configuration HTTP 请求。
+ 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
+}
diff --git a/internal/apps/oauth/handler_authorize.go b/internal/apps/oauth/handler_authorize.go
new file mode 100644
index 00000000..42ac36b4
--- /dev/null
+++ b/internal/apps/oauth/handler_authorize.go
@@ -0,0 +1,173 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package oauth
+
+import (
+ "context"
+ "fmt"
+ "net/http"
+ "strings"
+
+ "github.com/Rain-kl/Wavelet/internal/common"
+ "github.com/Rain-kl/Wavelet/internal/common/response"
+ "github.com/Rain-kl/Wavelet/internal/db"
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/coreos/go-oidc/v3/oidc"
+ "github.com/gin-contrib/sessions"
+ "github.com/gin-gonic/gin"
+ "github.com/google/uuid"
+)
+
+// GetLoginURL 获取登录授权地址
+// @Summary 获取登录授权地址
+// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。
+// @Tags oauth
+// @Produce json
+// @Param source query string false "认证源名称,为空使用第一个启用的认证源"
+// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL"
+// @Failure 400 {object} response.Any "认证源不存在或未配置"
+// @Failure 500 {object} response.Any "Redis 异常 or 构造 URL 失败"
+// @Router /api/v1/oauth/login [get]
+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)
+
+ state := uuid.NewString()
+ payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
+ SourceName: source.Name,
+ Purpose: OAuthPurposeLogin,
+ UserID: userID,
+ SessionHash: sessionHash,
+ })
+ if err != nil {
+ response.AbortInternal(c, err.Error())
+ return
+ }
+ if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
+ response.AbortInternal(c, err.Error())
+ return
+ }
+
+ authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
+ if err != nil {
+ response.AbortBadRequest(c, err.Error())
+ return
+ }
+ c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
+}
+
+func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state string) (string, error) {
+ redirectURL, err := getFrontendLoginRedirectURL(ctx)
+ if err != nil {
+ return "", err
+ }
+ authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
+ if err != nil {
+ return "", err
+ }
+ if verifier != nil {
+ return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil
+ }
+ return authConfig.AuthCodeURL(state), nil
+}
+
+// Authorize 发起指定认证源授权
+// @Summary 发起指定认证源授权
+// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。
+// @Tags oauth
+// @Produce json
+// @Param source path string true "认证源名称"
+// @Param purpose query string false "授权目的:login(登录)或 bind(绑定账号),默认 login"
+// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL"
+// @Failure 400 {object} response.Any "认证源不存在或未启用"
+// @Failure 500 {object} response.Any "Redis 异常或构造 URL 失败"
+// @Router /api/v1/oauth/{source}/authorize [get]
+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, common.UnAuthorized)
+ return
+ }
+
+ token, isNew := ensureSessionToken(session)
+ if isNew {
+ if err := session.Save(); err != nil {
+ response.AbortInternal(c, err.Error())
+ return
+ }
+ }
+
+ sessionHash := hashSessionToken(token)
+
+ state := uuid.NewString()
+ payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
+ SourceName: source.Name,
+ Purpose: purpose,
+ UserID: userID,
+ SessionHash: sessionHash,
+ })
+ if err != nil {
+ response.AbortInternal(c, err.Error())
+ return
+ }
+ if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
+ response.AbortInternal(c, err.Error())
+ return
+ }
+
+ authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
+ if err != nil {
+ response.AbortBadRequest(c, err.Error())
+ return
+ }
+ c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
+}
\ No newline at end of file
diff --git a/internal/apps/oauth/handler_callback.go b/internal/apps/oauth/handler_callback.go
new file mode 100644
index 00000000..fe8584c6
--- /dev/null
+++ b/internal/apps/oauth/handler_callback.go
@@ -0,0 +1,226 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package oauth
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "net/http"
+ "time"
+
+ "github.com/Rain-kl/Wavelet/internal/common"
+ "github.com/Rain-kl/Wavelet/internal/common/response"
+ "github.com/Rain-kl/Wavelet/internal/db"
+ "github.com/Rain-kl/Wavelet/internal/listener"
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/Rain-kl/Wavelet/pkg/logger"
+ "github.com/gin-contrib/sessions"
+ "github.com/gin-gonic/gin"
+ "gorm.io/gorm"
+)
+
+// Callback OAuth 回调处理
+// @Summary OAuth 回调处理
+// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立会话。支持登录(login)和账号绑定(bind)两种场景。
+// @Tags oauth
+// @Accept json
+// @Produce json
+// @Param request body oauth.CallbackRequest true "回调请求参数"
+// @Success 200 {object} response.Any{data=oauth.OAuthCallbackResult} "登录或绑定成功"
+// @Failure 400 {object} response.Any "state 无效、参数错误或认证源错误"
+// @Failure 401 {object} response.Any "绑定场景未登录"
+// @Failure 500 {object} response.Any "OAuth 认证失败或内部错误"
+// @Router /api/v1/oauth/callback [post]
+func Callback(c *gin.Context) {
+ var req CallbackRequest
+ if err := c.ShouldBindJSON(&req); err != nil {
+ response.AbortBadRequest(c, err.Error())
+ return
+ }
+
+ ctx := c.Request.Context()
+ stateKey := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
+ payloadRaw, err := db.Redis.Get(ctx, stateKey).Result()
+ if err != nil {
+ response.AbortBadRequest(c, errInvalidState)
+ return
+ }
+ _ = db.Redis.Del(ctx, stateKey)
+
+ payload, err := decodeOAuthStatePayload(payloadRaw)
+ if err != nil {
+ response.AbortBadRequest(c, err.Error())
+ return
+ }
+
+ session := sessions.Default(c)
+ currentUserID := GetUserIDFromSession(session)
+
+ if payload.Purpose == OAuthPurposeBind && currentUserID == 0 {
+ response.AbortUnauthorized(c, common.UnAuthorized)
+ return
+ }
+
+ token, ok := session.Get(SessionTokenKey).(string)
+ if !ok || token == "" {
+ response.AbortBadRequest(c, "invalid session context")
+ return
+ }
+
+ if hashSessionToken(token) != payload.SessionHash {
+ response.AbortBadRequest(c, "session mismatch for oauth state")
+ return
+ }
+
+ if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID {
+ response.AbortBadRequest(c, "user context mismatch for oauth binding")
+ return
+ }
+
+ if !isOIDCLoginEnabled(ctx) {
+ response.AbortBadRequest(c, errAuthSourceDisabled)
+ return
+ }
+
+ source, err := resolveAuthSource(ctx, payload.SourceName)
+ if err != nil {
+ response.AbortBadRequest(c, err.Error())
+ return
+ }
+
+ if !source.IsActive {
+ response.AbortBadRequest(c, errAuthSourceDisabled)
+ return
+ }
+
+ redirectURL, err := getFrontendLoginRedirectURL(ctx)
+ if err != nil {
+ response.AbortBadRequest(c, err.Error())
+ return
+ }
+
+ userInfo, err := buildOAuthUserInfo(ctx, source, req.Code, req.State, redirectURL)
+ if err != nil {
+ response.AbortInternal(c, err.Error())
+ return
+ }
+ if err := normalizeOAuthUserInfo(userInfo); err != nil {
+ response.AbortBadRequest(c, err.Error())
+ return
+ }
+ if userInfo.Sub == "" {
+ userInfo.Sub = userInfo.Username
+ }
+
+ if payload.Purpose == OAuthPurposeBind {
+ handleCallbackBind(ctx, c, source, userInfo)
+ return
+ }
+
+ handleCallbackLogin(ctx, c, source, userInfo)
+}
+
+// handleCallbackBind 处理 OAuth 回调中的帐号绑定流程
+func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
+ userID := GetUserIDFromContext(c)
+ if userID == 0 {
+ response.AbortUnauthorized(c, common.UnAuthorized)
+ return
+ }
+ var user model.User
+ if err := db.DB(ctx).First(&user, "id = ?", userID).Error; err != nil {
+ response.AbortInternal(c, err.Error())
+ return
+ }
+ if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
+ AuthSourceID: source.ID,
+ UserID: user.ID,
+ ExternalID: userInfo.Sub,
+ ExternalUsername: userInfo.Username,
+ Email: userInfo.Email,
+ }); err != nil {
+ response.AbortBadRequest(c, err.Error())
+ return
+ }
+ user.LastLoginAt = time.Now()
+ _ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error
+ c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound")))
+}
+
+// handleCallbackLogin 处理 OAuth 回调中的登录流程(查找已有帐号或自动注册)
+func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
+ var user model.User
+
+ account, err := model.FindExternalAccount(ctx, source.ID, userInfo.Sub)
+ switch {
+ case err == nil:
+ if err := db.DB(ctx).First(&user, "id = ?", account.UserID).Error; err != nil {
+ response.AbortInternal(c, err.Error())
+ return
+ }
+ 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()
+ _ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error
+ if err := setLoginSession(ctx, c, &user); err != nil {
+ response.AbortInternal(c, err.Error())
+ return
+ }
+
+ logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
+
+ listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP())
+
+ c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in")))
+}
+
+// handleCallbackRegister 处理 OAuth 回调中的自动注册流程
+// 若注册被禁用则保存 pending 信息并返回 false;若注册成功则返回新用户;若出错则返回 false
+func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) {
+ registrationEnabled, regErr := model.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
+ if regErr != nil {
+ registrationEnabled = true
+ }
+
+ if !registrationEnabled {
+ c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind")))
+ return model.User{}, false
+ }
+
+ username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
+ if uniqueErr != nil {
+ response.AbortInternal(c, uniqueErr.Error())
+ return model.User{}, false
+ }
+ userInfo.Username = username
+
+ var user model.User
+ if err := user.CreateUser(ctx, db.DB(ctx), userInfo); err != nil {
+ response.AbortInternal(c, err.Error())
+ return model.User{}, false
+ }
+ if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
+ AuthSourceID: source.ID,
+ UserID: user.ID,
+ ExternalID: userInfo.Sub,
+ ExternalUsername: userInfo.Username,
+ Email: userInfo.Email,
+ }); err != nil {
+ response.AbortBadRequest(c, err.Error())
+ return model.User{}, false
+ }
+ logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
+
+ return user, true
+}
\ No newline at end of file
diff --git a/internal/apps/oauth/handler_external_accounts.go b/internal/apps/oauth/handler_external_accounts.go
new file mode 100644
index 00000000..7f766154
--- /dev/null
+++ b/internal/apps/oauth/handler_external_accounts.go
@@ -0,0 +1,65 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package oauth
+
+import (
+ "net/http"
+ "strconv"
+ "strings"
+
+ "github.com/Rain-kl/Wavelet/internal/common"
+ "github.com/Rain-kl/Wavelet/internal/common/response"
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/gin-gonic/gin"
+)
+
+// ListExternalAccounts 获取当前用户的外部帐号绑定列表
+// @Summary 获取外部帐号列表
+// @Description 返回当前登录用户已绑定的所有外部 OAuth 帐号信息,需要登录
+// @Tags oauth
+// @Produce json
+// @Security SessionCookie
+// @Success 200 {object} response.Any{data=[]model.ExternalAccountView} "外部帐号列表"
+// @Failure 401 {object} response.Any "未登录"
+// @Failure 500 {object} response.Any "内部错误"
+// @Router /api/v1/oauth/external-accounts [get]
+func ListExternalAccounts(c *gin.Context) {
+ userID := GetUserIDFromContext(c)
+ accounts, err := model.ListExternalAccountsByUserID(c.Request.Context(), userID)
+ if err != nil {
+ response.AbortInternal(c, err.Error())
+ return
+ }
+ c.JSON(http.StatusOK, response.OK(accounts))
+}
+
+// DeleteExternalAccount 解除外部帐号绑定
+// @Summary 解除外部帐号绑定
+// @Description 解除当前登录用户与指定外部帐号的绑定关系,需要登录
+// @Tags oauth
+// @Produce json
+// @Security SessionCookie
+// @Param id path uint64 true "外部帐号绑定记录 ID"
+// @Success 200 {object} response.Any{data=string} "解除绑定成功"
+// @Failure 400 {object} response.Any "ID 无效或解除失败"
+// @Failure 401 {object} response.Any "未登录"
+// @Router /api/v1/oauth/external-accounts/{id}/delete [post]
+func DeleteExternalAccount(c *gin.Context) {
+ userID := GetUserIDFromContext(c)
+ if userID == 0 {
+ response.AbortUnauthorized(c, common.UnAuthorized)
+ return
+ }
+ rawID := strings.TrimSpace(c.Param("id"))
+ id, err := strconv.ParseUint(rawID, 10, 64)
+ if err != nil || id == 0 {
+ response.AbortBadRequest(c, errInvalidExternalAccountBindingID)
+ return
+ }
+ if err := model.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil {
+ response.AbortBadRequest(c, err.Error())
+ return
+ }
+ c.JSON(http.StatusOK, response.OKNil())
+}
\ No newline at end of file
diff --git a/internal/apps/oauth/handler_sources.go b/internal/apps/oauth/handler_sources.go
new file mode 100644
index 00000000..fa1c628a
--- /dev/null
+++ b/internal/apps/oauth/handler_sources.go
@@ -0,0 +1,22 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package oauth
+
+import (
+ "net/http"
+
+ "github.com/Rain-kl/Wavelet/internal/common/response"
+ "github.com/gin-gonic/gin"
+)
+
+// GetLoginSources 获取可用登录源列表
+// @Summary 获取可用登录源
+// @Description 返回当前系统已启用的所有 OAuth 登录源,前端展示登录按钮列表时调用
+// @Tags oauth
+// @Produce json
+// @Success 200 {object} response.Any{data=[]oauth.AuthSourceView} "登录源列表"
+// @Router /api/v1/oauth/sources [get]
+func GetLoginSources(c *gin.Context) {
+ c.JSON(http.StatusOK, response.OK(activeLoginSources(c.Request.Context())))
+}
diff --git a/internal/apps/oauth/middlewares.go b/internal/apps/oauth/middlewares.go
index 8f4a4f68..5eacbbda 100644
--- a/internal/apps/oauth/middlewares.go
+++ b/internal/apps/oauth/middlewares.go
@@ -7,9 +7,9 @@ package oauth
import (
"context"
"errors"
- "net/http"
"github.com/Rain-kl/Wavelet/internal/common"
+ "github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
@@ -111,7 +111,7 @@ func LoginRequired() gin.HandlerFunc {
user, err := GetUserFromRequest(c)
if err != nil {
- c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil})
+ response.AbortUnauthorized(c, common.UnAuthorized)
return
}
@@ -130,7 +130,7 @@ func LoginRequired() gin.HandlerFunc {
func DisallowTokenAuth() gin.HandlerFunc {
return func(c *gin.Context) {
if tokenAuth, _ := util.GetFromContext[bool](c, TokenAuthKey); tokenAuth {
- c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error_msg": ErrTokenAuthNotAllowed, "data": nil})
+ response.AbortForbidden(c, ErrTokenAuthNotAllowed)
return
}
c.Next()
diff --git a/internal/apps/oauth/oauth_test.go b/internal/apps/oauth/oauth_test.go
index 5e2a9468..6d58fbf1 100644
--- a/internal/apps/oauth/oauth_test.go
+++ b/internal/apps/oauth/oauth_test.go
@@ -35,6 +35,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/internal/util"
)
@@ -127,6 +128,12 @@ func (m *mockRedisClient) HGet(ctx context.Context, key string, field string) *r
return cmd
}
+func (m *mockRedisClient) Subscribe(ctx context.Context, channels ...string) *redis.PubSub {
+ return redis.NewClient(&redis.Options{
+ Addr: "127.0.0.1:0",
+ }).Subscribe(ctx, channels...)
+}
+
type mockRoundTripper struct {
roundTripFunc func(req *http.Request) (*http.Response, error)
}
@@ -332,9 +339,7 @@ func mockContextMiddleware(mockClient *http.Client) gin.HandlerFunc {
}
func setupTestRouter(dbConn *gorm.DB, mockRedis *mockRedisClient, mockClient *http.Client) *gin.Engine {
- gin.SetMode(gin.TestMode)
- r := gin.New()
- r.Use(gin.Recovery())
+ r := testhelper.NewTestGinEngine(gin.Recovery())
// Inject context mock middleware
r.Use(mockContextMiddleware(mockClient))
@@ -1199,7 +1204,7 @@ func TestSystemUserBlockedByMiddleware(t *testing.T) {
// 3. 设置全局测试数据库连接并构建测试路由组
db.SetDB(dbConn)
- rProtected := gin.New()
+ rProtected := testhelper.NewTestGinEngine()
store := cookie.NewStore([]byte("secret"))
rProtected.Use(sessions.Sessions("mysession", store))
rProtected.Use(LoginRequired())
diff --git a/internal/apps/oauth/oauth_types.go b/internal/apps/oauth/oauth_types.go
new file mode 100644
index 00000000..17ab267c
--- /dev/null
+++ b/internal/apps/oauth/oauth_types.go
@@ -0,0 +1,36 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package oauth
+
+// AuthSourceView 登录源展示信息
+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 响应
+//
+//nolint:revive // OAuth 前缀保持包内语义清晰
+type OAuthAuthorizeResponse struct {
+ AuthorizeURL string `json:"authorize_url"`
+}
+
+// OAuthCallbackResult 回调处理结果
+//
+//nolint:revive // OAuth 前缀保持包内语义清晰
+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"`
+}
diff --git a/internal/apps/oauth/oauth_userinfo.go b/internal/apps/oauth/oauth_userinfo.go
new file mode 100644
index 00000000..84ae7124
--- /dev/null
+++ b/internal/apps/oauth/oauth_userinfo.go
@@ -0,0 +1,141 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package oauth
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "strings"
+
+ "github.com/Rain-kl/Wavelet/internal/db"
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/coreos/go-oidc/v3/oidc"
+ "golang.org/x/oauth2"
+)
+
+func uniqueUsername(ctx context.Context, base string) (string, error) {
+ base = strings.TrimSpace(base)
+ if base == "" {
+ base = "user"
+ }
+
+ var existingUsernames []string
+ if err := db.DB(ctx).Model(&model.User{}).
+ Where("username = ? OR username LIKE ?", base, base+"-%").
+ Pluck("username", &existingUsernames).Error; err != nil {
+ return "", err
+ }
+
+ // 将现有的用户名放入 map 中,以便 O(1) 查找
+ exists := make(map[string]bool, len(existingUsernames))
+ for _, u := range existingUsernames {
+ exists[strings.ToLower(u)] = true
+ }
+
+ // 检查 base 是否被占用
+ if !exists[strings.ToLower(base)] {
+ return base, nil
+ }
+
+ // 顺序查找第一个可用的带后缀用户名
+ for i := 1; i <= 1000; i++ {
+ candidate := fmt.Sprintf("%s-%d", base, i)
+ if !exists[strings.ToLower(candidate)] {
+ return candidate, nil
+ }
+ }
+
+ return "", errors.New(errUsernameGenerateFailed)
+}
+
+func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, error) {
+ authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
+ if err != nil {
+ return nil, err
+ }
+
+ token, err := authConfig.Exchange(ctx, code)
+ if err != nil {
+ return nil, err
+ }
+
+ userInfo := &model.OAuthUserInfo{Active: true}
+ if verifier != nil {
+ if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil {
+ return nil, verifyErr
+ }
+ }
+
+ if userInfo.Username == "" && userInfo.PreferredUsername != "" {
+ userInfo.Username = userInfo.PreferredUsername
+ }
+ if userInfo.Username == "" && userInfo.Email != "" {
+ userInfo.Username = strings.Split(userInfo.Email, "@")[0]
+ }
+ if userInfo.Username == "" && userInfo.Sub != "" {
+ userInfo.Username = userInfo.Sub
+ }
+ if userInfo.Name == "" {
+ userInfo.Name = userInfo.Username
+ }
+
+ return userInfo, nil
+}
+
+// verifyIDToken 验证 OIDC ID Token 并将 Claims 解析到 userInfo
+func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *model.OAuthUserInfo) error {
+ rawIDToken, ok := token.Extra("id_token").(string)
+ if !ok {
+ return nil
+ }
+ idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
+ if verifyErr != nil {
+ return fmt.Errorf(errIDTokenVerifyFailedFormat, errIDTokenVerifyFailed, verifyErr)
+ }
+ if nonce != "" && idToken.Nonce != nonce {
+ return errors.New(errNonceMismatch)
+ }
+ if claimsErr := idToken.Claims(userInfo); claimsErr != nil {
+ return claimsErr
+ }
+ return nil
+}
+
+func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
+ userInfo.Username = strings.TrimSpace(userInfo.Username)
+ userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername)
+ userInfo.Email = strings.TrimSpace(userInfo.Email)
+ userInfo.Name = strings.TrimSpace(userInfo.Name)
+ userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL)
+
+ if userInfo.Username == "" && userInfo.PreferredUsername != "" {
+ userInfo.Username = userInfo.PreferredUsername
+ }
+ if userInfo.Username == "" && userInfo.Email != "" {
+ userInfo.Username = strings.Split(userInfo.Email, "@")[0]
+ }
+ if userInfo.Username == "" && userInfo.Sub != "" {
+ userInfo.Username = userInfo.Sub
+ }
+ if userInfo.Username == "" {
+ return errors.New(errUsernameFromSourceFailed)
+ }
+ if userInfo.Name == "" {
+ userInfo.Name = userInfo.Username
+ }
+ if !userInfo.Active {
+ userInfo.Active = true
+ }
+ return nil
+}
+
+func buildCallbackResult(user *model.User, status string) OAuthCallbackResult {
+ result := OAuthCallbackResult{Status: status}
+ if user != nil {
+ info := BuildBasicUserInfo(user, false)
+ result.User = &info
+ }
+ return result
+}
diff --git a/internal/apps/oauth/routers.go b/internal/apps/oauth/routers.go
index e5be215a..6eeba2f3 100644
--- a/internal/apps/oauth/routers.go
+++ b/internal/apps/oauth/routers.go
@@ -98,7 +98,7 @@ func Logout(c *gin.Context) {
session.Options(GetSessionOptions(-1))
session.Clear()
if err := session.Save(); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
diff --git a/internal/apps/oauth/session_context.go b/internal/apps/oauth/session_context.go
new file mode 100644
index 00000000..42be0c55
--- /dev/null
+++ b/internal/apps/oauth/session_context.go
@@ -0,0 +1,82 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package oauth
+
+import (
+ "context"
+ "crypto/sha256"
+ "encoding/hex"
+
+ "github.com/Rain-kl/Wavelet/internal/config"
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/gin-contrib/sessions"
+ "github.com/gin-gonic/gin"
+ "github.com/google/uuid"
+)
+
+// GetUserIDFromSession 从 Session 中提取用户 ID
+func GetUserIDFromSession(s sessions.Session) uint64 {
+ userID, ok := s.Get(UserIDKey).(uint64)
+ if !ok {
+ return 0
+ }
+ return userID
+}
+
+// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID
+func GetUserIDFromContext(c *gin.Context) uint64 {
+ session := sessions.Default(c)
+ return GetUserIDFromSession(session)
+}
+
+func ensureSessionToken(s sessions.Session) (string, bool) {
+ token, ok := s.Get(SessionTokenKey).(string)
+ if !ok || token == "" {
+ token = uuid.NewString()
+ s.Set(SessionTokenKey, token)
+ return token, true
+ }
+ return token, false
+}
+
+func hashSessionToken(token string) string {
+ h := sha256.New()
+ h.Write([]byte(token))
+ return hex.EncodeToString(h.Sum(nil))
+}
+
+func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error {
+ session := sessions.Default(c)
+ session.Set(UserIDKey, user.ID)
+ session.Set(UserNameKey, user.Username)
+ session.Set(PasswordHashKey, user.Password)
+
+ // 根据系统配置动态设置 Session 过期时间
+ maxAge := config.Config.App.SessionAge
+ isSessionCookie := false
+
+ ttlHours, err := model.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
+ if err == nil {
+ switch {
+ case ttlHours == -1:
+ // 永不过期,设置为 10 年
+ maxAge = 10 * 365 * 24 * 3600
+ case ttlHours > 0:
+ maxAge = ttlHours * 3600
+ case ttlHours == 0:
+ isSessionCookie = true
+ }
+ }
+ session.Options(GetSessionOptions(maxAge))
+
+ if err := session.Save(); err != nil {
+ return err
+ }
+
+ if isSessionCookie {
+ StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName)
+ }
+
+ return nil
+}
diff --git a/internal/apps/oauth/sources.go b/internal/apps/oauth/sources.go
deleted file mode 100644
index d1caa0bb..00000000
--- a/internal/apps/oauth/sources.go
+++ /dev/null
@@ -1,774 +0,0 @@
-// Copyright 2026 Arctel.net
-// SPDX-License-Identifier: Apache-2.0
-
-package oauth
-
-import (
- "context"
- "crypto/sha256"
- "encoding/hex"
- "errors"
- "fmt"
- "net/http"
- "strconv"
- "strings"
- "time"
-
- "github.com/Rain-kl/Wavelet/internal/common"
- "github.com/Rain-kl/Wavelet/internal/common/response"
- "github.com/Rain-kl/Wavelet/internal/config"
- "github.com/Rain-kl/Wavelet/internal/db"
- "github.com/Rain-kl/Wavelet/internal/listener"
- "github.com/Rain-kl/Wavelet/internal/model"
- "github.com/Rain-kl/Wavelet/pkg/logger"
- "github.com/coreos/go-oidc/v3/oidc"
- "github.com/gin-contrib/sessions"
- "github.com/gin-gonic/gin"
- "github.com/google/uuid"
- "golang.org/x/oauth2"
- "gorm.io/gorm"
-)
-
-// AuthSourceView 登录源展示信息
-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 响应
-//
-//nolint:revive // OAuth 前缀保持包内语义清晰
-type OAuthAuthorizeResponse struct {
- AuthorizeURL string `json:"authorize_url"`
-}
-
-// OAuthCallbackResult 回调处理结果
-//
-//nolint:revive // OAuth 前缀保持包内语义清晰
-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"`
-}
-
-// GetUserIDFromSession 从 Session 中提取用户 ID
-func GetUserIDFromSession(s sessions.Session) uint64 {
- userID, ok := s.Get(UserIDKey).(uint64)
- if !ok {
- return 0
- }
- return userID
-}
-
-// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID
-func GetUserIDFromContext(c *gin.Context) uint64 {
- session := sessions.Default(c)
- return GetUserIDFromSession(session)
-}
-
-func ensureSessionToken(s sessions.Session) (string, bool) {
- token, ok := s.Get(SessionTokenKey).(string)
- if !ok || token == "" {
- token = uuid.NewString()
- s.Set(SessionTokenKey, token)
- return token, true
- }
- return token, false
-}
-
-func hashSessionToken(token string) string {
- h := sha256.New()
- h.Write([]byte(token))
- return hex.EncodeToString(h.Sum(nil))
-}
-
-func isOIDCLoginEnabled(ctx context.Context) bool {
- enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
- if err != nil {
- return true
- }
- return enabled
-}
-
-func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) {
- name := strings.TrimSpace(strings.ToLower(sourceName))
- if name == "" {
- sources, err := model.GetActiveAuthSources(ctx)
- if err != nil {
- return nil, err
- }
- if len(sources) == 0 {
- return nil, errors.New(errNoActiveAuthSource)
- }
- return &sources[0], nil
- }
- return model.GetAuthSourceByName(ctx, name)
-}
-
-func activeLoginSources(ctx context.Context) []AuthSourceView {
- enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
- if err == nil && !enabled {
- return nil
- }
-
- dbSources, err := model.GetActiveAuthSources(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) {
- var sc model.SystemConfig
- if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || strings.TrimSpace(sc.Value) == "" {
- return "", errors.New(errServerAddressMissing)
- }
- return strings.TrimRight(sc.Value, "/") + "/login", nil
-}
-
-func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
- if source == nil {
- return nil, nil, errors.New(errAuthSourceRequired)
- }
-
- if source.OpenIDDiscoveryURL == "" {
- return nil, nil, errors.New(errDiscoveryURLRequired)
- }
-
- // Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake)
- 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,避免每次调用都向 issuer 发起
- // /.well-known/openid-configuration HTTP 请求。
- 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 setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error {
- session := sessions.Default(c)
- session.Set(UserIDKey, user.ID)
- session.Set(UserNameKey, user.Username)
- session.Set(PasswordHashKey, user.Password)
-
- // 根据系统配置动态设置 Session 过期时间
- maxAge := config.Config.App.SessionAge
- isSessionCookie := false
-
- ttlHours, err := model.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
- if err == nil {
- switch {
- case ttlHours == -1:
- // 永不过期,设置为 10 年
- maxAge = 10 * 365 * 24 * 3600
- case ttlHours > 0:
- maxAge = ttlHours * 3600
- case ttlHours == 0:
- isSessionCookie = true
- }
- }
- session.Options(GetSessionOptions(maxAge))
-
- if err := session.Save(); err != nil {
- return err
- }
-
- if isSessionCookie {
- StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName)
- }
-
- return nil
-}
-
-func uniqueUsername(ctx context.Context, base string) (string, error) {
- base = strings.TrimSpace(base)
- if base == "" {
- base = "user"
- }
-
- var existingUsernames []string
- if err := db.DB(ctx).Model(&model.User{}).
- Where("username = ? OR username LIKE ?", base, base+"-%").
- Pluck("username", &existingUsernames).Error; err != nil {
- return "", err
- }
-
- // 将现有的用户名放入 map 中,以便 O(1) 查找
- exists := make(map[string]bool, len(existingUsernames))
- for _, u := range existingUsernames {
- exists[strings.ToLower(u)] = true
- }
-
- // 检查 base 是否被占用
- if !exists[strings.ToLower(base)] {
- return base, nil
- }
-
- // 顺序查找第一个可用的带后缀用户名
- for i := 1; i <= 1000; i++ {
- candidate := fmt.Sprintf("%s-%d", base, i)
- if !exists[strings.ToLower(candidate)] {
- return candidate, nil
- }
- }
-
- return "", errors.New(errUsernameGenerateFailed)
-}
-
-func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, error) {
- authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
- if err != nil {
- return nil, err
- }
-
- token, err := authConfig.Exchange(ctx, code)
- if err != nil {
- return nil, err
- }
-
- userInfo := &model.OAuthUserInfo{Active: true}
- if verifier != nil {
- if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil {
- return nil, verifyErr
- }
- }
-
- if userInfo.Username == "" && userInfo.PreferredUsername != "" {
- userInfo.Username = userInfo.PreferredUsername
- }
- if userInfo.Username == "" && userInfo.Email != "" {
- userInfo.Username = strings.Split(userInfo.Email, "@")[0]
- }
- if userInfo.Username == "" && userInfo.Sub != "" {
- userInfo.Username = userInfo.Sub
- }
- if userInfo.Name == "" {
- userInfo.Name = userInfo.Username
- }
-
- return userInfo, nil
-}
-
-// verifyIDToken 验证 OIDC ID Token 并将 Claims 解析到 userInfo
-func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *model.OAuthUserInfo) error {
- rawIDToken, ok := token.Extra("id_token").(string)
- if !ok {
- return nil
- }
- idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
- if verifyErr != nil {
- return fmt.Errorf(errIDTokenVerifyFailedFormat, errIDTokenVerifyFailed, verifyErr)
- }
- if nonce != "" && idToken.Nonce != nonce {
- return errors.New(errNonceMismatch)
- }
- if claimsErr := idToken.Claims(userInfo); claimsErr != nil {
- return claimsErr
- }
- return nil
-}
-
-func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
- userInfo.Username = strings.TrimSpace(userInfo.Username)
- userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername)
- userInfo.Email = strings.TrimSpace(userInfo.Email)
- userInfo.Name = strings.TrimSpace(userInfo.Name)
- userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL)
-
- if userInfo.Username == "" && userInfo.PreferredUsername != "" {
- userInfo.Username = userInfo.PreferredUsername
- }
- if userInfo.Username == "" && userInfo.Email != "" {
- userInfo.Username = strings.Split(userInfo.Email, "@")[0]
- }
- if userInfo.Username == "" && userInfo.Sub != "" {
- userInfo.Username = userInfo.Sub
- }
- if userInfo.Username == "" {
- return errors.New(errUsernameFromSourceFailed)
- }
- if userInfo.Name == "" {
- userInfo.Name = userInfo.Username
- }
- if !userInfo.Active {
- userInfo.Active = true
- }
- return nil
-}
-
-func buildCallbackResult(user *model.User, status string) OAuthCallbackResult {
- result := OAuthCallbackResult{Status: status}
- if user != nil {
- info := BuildBasicUserInfo(user, false)
- result.User = &info
- }
- return result
-}
-
-// GetLoginSources 获取可用登录源列表
-// @Summary 获取可用登录源
-// @Description 返回当前系统已启用的所有 OAuth 登录源,前端展示登录按钮列表时调用
-// @Tags oauth
-// @Produce json
-// @Success 200 {object} response.Any{data=[]oauth.AuthSourceView} "登录源列表"
-// @Router /api/v1/oauth/sources [get]
-func GetLoginSources(c *gin.Context) {
- c.JSON(http.StatusOK, response.OK(activeLoginSources(c.Request.Context())))
-}
-
-// GetLoginURL 获取登录授权地址
-// @Summary 获取登录授权地址
-// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。
-// @Tags oauth
-// @Produce json
-// @Param source query string false "认证源名称,为空使用第一个启用的认证源"
-// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL"
-// @Failure 400 {object} response.Any "认证源不存在或未配置"
-// @Failure 500 {object} response.Any "Redis 异常 or 构造 URL 失败"
-// @Router /api/v1/oauth/login [get]
-func GetLoginURL(c *gin.Context) {
- ctx := c.Request.Context()
- if !isOIDCLoginEnabled(ctx) {
- c.JSON(http.StatusBadRequest, response.Err(errAuthSourceDisabled))
- return
- }
-
- source, err := resolveAuthSource(ctx, c.Query("source"))
- if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
- return
- }
-
- if !source.IsActive {
- c.JSON(http.StatusBadRequest, response.Err(errAuthSourceDisabled))
- return
- }
-
- session := sessions.Default(c)
- token, isNew := ensureSessionToken(session)
- if isNew {
- if err := session.Save(); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return
- }
- }
-
- userID := GetUserIDFromSession(session)
- sessionHash := hashSessionToken(token)
-
- state := uuid.NewString()
- payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
- SourceName: source.Name,
- Purpose: OAuthPurposeLogin,
- UserID: userID,
- SessionHash: sessionHash,
- })
- if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return
- }
- if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return
- }
-
- authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
- if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
- return
- }
- c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
-}
-
-func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state string) (string, error) {
- redirectURL, err := getFrontendLoginRedirectURL(ctx)
- if err != nil {
- return "", err
- }
- authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
- if err != nil {
- return "", err
- }
- if verifier != nil {
- return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil
- }
- return authConfig.AuthCodeURL(state), nil
-}
-
-// Authorize 发起指定认证源授权
-// @Summary 发起指定认证源授权
-// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。
-// @Tags oauth
-// @Produce json
-// @Param source path string true "认证源名称"
-// @Param purpose query string false "授权目的:login(登录)或 bind(绑定账号),默认 login"
-// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL"
-// @Failure 400 {object} response.Any "认证源不存在或未启用"
-// @Failure 500 {object} response.Any "Redis 异常或构造 URL 失败"
-// @Router /api/v1/oauth/{source}/authorize [get]
-func Authorize(c *gin.Context) {
- ctx := c.Request.Context()
- if !isOIDCLoginEnabled(ctx) {
- c.JSON(http.StatusBadRequest, response.Err(errAuthSourceDisabled))
- return
- }
-
- source, err := resolveAuthSource(ctx, c.Param("source"))
- if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
- return
- }
-
- if !source.IsActive {
- c.JSON(http.StatusBadRequest, response.Err(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 {
- c.JSON(http.StatusUnauthorized, response.Err(common.UnAuthorized))
- return
- }
-
- token, isNew := ensureSessionToken(session)
- if isNew {
- if err := session.Save(); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return
- }
- }
-
- sessionHash := hashSessionToken(token)
-
- state := uuid.NewString()
- payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
- SourceName: source.Name,
- Purpose: purpose,
- UserID: userID,
- SessionHash: sessionHash,
- })
- if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return
- }
- if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return
- }
-
- authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
- if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
- return
- }
- c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
-}
-
-// Callback OAuth 回调处理
-// @Summary OAuth 回调处理
-// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立会话。支持登录(login)和账号绑定(bind)两种场景。
-// @Tags oauth
-// @Accept json
-// @Produce json
-// @Param request body oauth.CallbackRequest true "回调请求参数"
-// @Success 200 {object} response.Any{data=oauth.OAuthCallbackResult} "登录或绑定成功"
-// @Failure 400 {object} response.Any "state 无效、参数错误或认证源错误"
-// @Failure 401 {object} response.Any "绑定场景未登录"
-// @Failure 500 {object} response.Any "OAuth 认证失败或内部错误"
-// @Router /api/v1/oauth/callback [post]
-func Callback(c *gin.Context) {
- var req CallbackRequest
- if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
- return
- }
-
- ctx := c.Request.Context()
- stateKey := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
- payloadRaw, err := db.Redis.Get(ctx, stateKey).Result()
- if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(errInvalidState))
- return
- }
- _ = db.Redis.Del(ctx, stateKey)
-
- payload, err := decodeOAuthStatePayload(payloadRaw)
- if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
- return
- }
-
- session := sessions.Default(c)
- currentUserID := GetUserIDFromSession(session)
-
- if payload.Purpose == OAuthPurposeBind && currentUserID == 0 {
- c.JSON(http.StatusUnauthorized, response.Err(common.UnAuthorized))
- return
- }
-
- token, ok := session.Get(SessionTokenKey).(string)
- if !ok || token == "" {
- c.JSON(http.StatusBadRequest, response.Err("invalid session context"))
- return
- }
-
- if hashSessionToken(token) != payload.SessionHash {
- c.JSON(http.StatusBadRequest, response.Err("session mismatch for oauth state"))
- return
- }
-
- if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID {
- c.JSON(http.StatusBadRequest, response.Err("user context mismatch for oauth binding"))
- return
- }
-
- if !isOIDCLoginEnabled(ctx) {
- c.JSON(http.StatusBadRequest, response.Err(errAuthSourceDisabled))
- return
- }
-
- source, err := resolveAuthSource(ctx, payload.SourceName)
- if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
- return
- }
-
- if !source.IsActive {
- c.JSON(http.StatusBadRequest, response.Err(errAuthSourceDisabled))
- return
- }
-
- redirectURL, err := getFrontendLoginRedirectURL(ctx)
- if err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
- return
- }
-
- userInfo, err := buildOAuthUserInfo(ctx, source, req.Code, req.State, redirectURL)
- if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return
- }
- if err := normalizeOAuthUserInfo(userInfo); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(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)
-}
-
-// handleCallbackBind 处理 OAuth 回调中的帐号绑定流程
-func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
- userID := GetUserIDFromContext(c)
- if userID == 0 {
- c.JSON(http.StatusUnauthorized, response.Err(common.UnAuthorized))
- return
- }
- var user model.User
- if err := db.DB(ctx).First(&user, "id = ?", userID).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return
- }
- if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
- AuthSourceID: source.ID,
- UserID: user.ID,
- ExternalID: userInfo.Sub,
- ExternalUsername: userInfo.Username,
- Email: userInfo.Email,
- }); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
- return
- }
- user.LastLoginAt = time.Now()
- _ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error
- c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound")))
-}
-
-// handleCallbackLogin 处理 OAuth 回调中的登录流程(查找已有帐号或自动注册)
-func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
- var user model.User
-
- account, err := model.FindExternalAccount(ctx, source.ID, userInfo.Sub)
- switch {
- case err == nil:
- if err := db.DB(ctx).First(&user, "id = ?", account.UserID).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return
- }
- case errors.Is(err, gorm.ErrRecordNotFound):
- newUser, ok := handleCallbackRegister(ctx, c, source, userInfo)
- if !ok {
- return
- }
- user = newUser
- default:
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return
- }
-
- user.LastLoginAt = time.Now()
- _ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error
- if err := setLoginSession(ctx, c, &user); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return
- }
-
- logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
-
- listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP())
-
- c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in")))
-}
-
-// handleCallbackRegister 处理 OAuth 回调中的自动注册流程
-// 若注册被禁用则保存 pending 信息并返回 false;若注册成功则返回新用户;若出错则返回 false
-func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) {
- registrationEnabled, regErr := model.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
- if regErr != nil {
- registrationEnabled = true
- }
-
- if !registrationEnabled {
- c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind")))
- return model.User{}, false
- }
-
- username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
- if uniqueErr != nil {
- c.JSON(http.StatusInternalServerError, response.Err(uniqueErr.Error()))
- return model.User{}, false
- }
- userInfo.Username = username
-
- var user model.User
- if err := user.CreateUser(ctx, db.DB(ctx), userInfo); err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return model.User{}, false
- }
- if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
- AuthSourceID: source.ID,
- UserID: user.ID,
- ExternalID: userInfo.Sub,
- ExternalUsername: userInfo.Username,
- Email: userInfo.Email,
- }); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
- return model.User{}, false
- }
- logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
-
- return user, true
-}
-
-// ListExternalAccounts 获取当前用户的外部帐号绑定列表
-// @Summary 获取外部帐号列表
-// @Description 返回当前登录用户已绑定的所有外部 OAuth 帐号信息,需要登录
-// @Tags oauth
-// @Produce json
-// @Security SessionCookie
-// @Success 200 {object} response.Any{data=[]model.ExternalAccountView} "外部帐号列表"
-// @Failure 401 {object} response.Any "未登录"
-// @Failure 500 {object} response.Any "内部错误"
-// @Router /api/v1/oauth/external-accounts [get]
-func ListExternalAccounts(c *gin.Context) {
- userID := GetUserIDFromContext(c)
- accounts, err := model.ListExternalAccountsByUserID(c.Request.Context(), userID)
- if err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
- return
- }
- c.JSON(http.StatusOK, response.OK(accounts))
-}
-
-// DeleteExternalAccount 解除外部帐号绑定
-// @Summary 解除外部帐号绑定
-// @Description 解除当前登录用户与指定外部帐号的绑定关系,需要登录
-// @Tags oauth
-// @Produce json
-// @Security SessionCookie
-// @Param id path uint64 true "外部帐号绑定记录 ID"
-// @Success 200 {object} response.Any{data=string} "解除绑定成功"
-// @Failure 400 {object} response.Any "ID 无效或解除失败"
-// @Failure 401 {object} response.Any "未登录"
-// @Router /api/v1/oauth/external-accounts/{id}/delete [post]
-func DeleteExternalAccount(c *gin.Context) {
- userID := GetUserIDFromContext(c)
- if userID == 0 {
- c.JSON(http.StatusUnauthorized, response.Err(common.UnAuthorized))
- return
- }
- rawID := strings.TrimSpace(c.Param("id"))
- id, err := strconv.ParseUint(rawID, 10, 64)
- if err != nil || id == 0 {
- c.JSON(http.StatusBadRequest, response.Err(errInvalidExternalAccountBindingID))
- return
- }
- if err := model.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
- return
- }
- c.JSON(http.StatusOK, response.OKNil())
-}
diff --git a/internal/apps/risk_control/middleware.go b/internal/apps/risk_control/middleware.go
index 8b88afa5..ee5daf06 100644
--- a/internal/apps/risk_control/middleware.go
+++ b/internal/apps/risk_control/middleware.go
@@ -29,7 +29,7 @@ func RiskControlMiddleware() gin.HandlerFunc {
// 1. 限流背压检测(检测本地缓冲队列是否已满)
if IsBufferFull() {
- c.AbortWithStatusJSON(http.StatusTooManyRequests, response.Err("系统繁忙,请稍后再试"))
+ response.AbortTooManyRequests(c, "系统繁忙,请稍后再试")
return
}
diff --git a/internal/apps/risk_control/middleware_test.go b/internal/apps/risk_control/middleware_test.go
index b4f5652a..58d1227f 100644
--- a/internal/apps/risk_control/middleware_test.go
+++ b/internal/apps/risk_control/middleware_test.go
@@ -13,6 +13,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
@@ -25,8 +26,7 @@ func TestRiskControlMiddleware(t *testing.T) {
config.Config.ClickHouse.Enabled = false
defer func() { config.Config.ClickHouse.Enabled = false }()
- r := gin.New()
- r.Use(RiskControlMiddleware())
+ r := testhelper.NewTestGinEngine(RiskControlMiddleware())
r.GET("/test", func(c *gin.Context) {
c.String(http.StatusOK, "ok")
})
@@ -91,8 +91,7 @@ func TestRiskControlMiddleware(t *testing.T) {
logChan = nil
}()
- r := gin.New()
- r.Use(RiskControlMiddleware())
+ r := testhelper.NewTestGinEngine(RiskControlMiddleware())
r.GET("/test", func(c *gin.Context) {
c.String(http.StatusOK, "ok")
})
@@ -126,8 +125,7 @@ func TestRiskControlMiddleware(t *testing.T) {
logChan <- &UserAccessLog{}
}
- r := gin.New()
- r.Use(RiskControlMiddleware())
+ r := testhelper.NewTestGinEngine(RiskControlMiddleware())
r.GET("/test", func(c *gin.Context) {
c.String(http.StatusOK, "ok")
})
diff --git a/internal/apps/upload/access_cache.go b/internal/apps/upload/cache/access_cache.go
similarity index 61%
rename from internal/apps/upload/access_cache.go
rename to internal/apps/upload/cache/access_cache.go
index efddb0ee..eff0681d 100644
--- a/internal/apps/upload/access_cache.go
+++ b/internal/apps/upload/cache/access_cache.go
@@ -1,7 +1,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+// Package cache provides in-process upload access-control caches.
+package cache
import (
"context"
@@ -10,32 +11,19 @@ import (
"sync"
"time"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
+ uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
)
-const accessCacheTTL = 5 * time.Second
-
const fileAccessInvalidationChannel = "upload:file_access_invalidation"
-type migrationAccessState struct {
- readOnly bool
- target storage.Config
- hasTarget bool
- targetErr error
- loadErr error
-}
-
var (
accessCacheOnce sync.Once
- migrationAccessMu sync.RWMutex
- migrationAccessCached migrationAccessState
- migrationAccessValid bool
- migrationAccessCheckedAt time.Time
-
- fileAccessWhitelistMu sync.RWMutex
+ fileAccessWhitelistMu sync.RWMutex
fileAccessWhitelistTypes map[string]struct{}
fileAccessWhitelistValid bool
fileAccessWhitelistCheckedAt time.Time
@@ -43,9 +31,7 @@ var (
// ResetAccessCaches clears in-process upload access caches.
func ResetAccessCaches() {
- migrationAccessMu.Lock()
- migrationAccessValid = false
- migrationAccessMu.Unlock()
+ uploadstorage.ResetMigrationAccessCache()
fileAccessWhitelistMu.Lock()
fileAccessWhitelistValid = false
@@ -85,62 +71,18 @@ func startAccessCacheInvalidationListener() {
}()
}
-func loadMigrationAccessState(ctx context.Context) migrationAccessState {
- ensureAccessCacheListener()
-
- migrationAccessMu.RLock()
- if migrationAccessValid && time.Since(migrationAccessCheckedAt) < accessCacheTTL {
- state := migrationAccessCached
- migrationAccessMu.RUnlock()
- return state
- }
- migrationAccessMu.RUnlock()
-
- migrationAccessMu.Lock()
- defer migrationAccessMu.Unlock()
-
- if migrationAccessValid && time.Since(migrationAccessCheckedAt) < accessCacheTTL {
- return migrationAccessCached
- }
-
- migrationAccessCached = buildMigrationAccessState(ctx)
- migrationAccessValid = true
- migrationAccessCheckedAt = time.Now()
- return migrationAccessCached
-}
-
-func buildMigrationAccessState(ctx context.Context) migrationAccessState {
- execution, ok, err := latestStorageMigrationExecution(ctx)
- if err != nil {
- return migrationAccessState{loadErr: err, readOnly: true}
- }
- if !ok {
- return migrationAccessState{}
- }
-
- state := migrationAccessState{
- readOnly: execution.Status != model.TaskExecutionStatusSucceeded,
- }
- if execution.Status == model.TaskExecutionStatusSucceeded {
- return state
- }
-
- target, err := parseMigrationTargetConfig(ctx, []byte(execution.Payload))
- if err != nil {
- state.targetErr = err
- return state
- }
-
- state.target = target
- state.hasTarget = true
- return state
+// IsFilePublic reports whether uploadType is in the public access whitelist.
+func IsFilePublic(ctx context.Context, uploadType string) bool {
+ whitelist := loadFileAccessWhitelist(ctx)
+ _, ok := whitelist[strings.ToLower(uploadType)]
+ return ok
}
func loadFileAccessWhitelist(ctx context.Context) map[string]struct{} {
ensureAccessCacheListener()
fileAccessWhitelistMu.RLock()
- if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < accessCacheTTL {
+ if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
types := fileAccessWhitelistTypes
fileAccessWhitelistMu.RUnlock()
return types
@@ -150,7 +92,7 @@ func loadFileAccessWhitelist(ctx context.Context) map[string]struct{} {
fileAccessWhitelistMu.Lock()
defer fileAccessWhitelistMu.Unlock()
- if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < accessCacheTTL {
+ if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
return fileAccessWhitelistTypes
}
@@ -172,7 +114,7 @@ func fetchFileAccessWhitelist(ctx context.Context) map[string]struct{} {
func parseFileAccessWhitelist(ctx context.Context) []string {
var sc model.SystemConfig
if err := sc.GetByKey(ctx, model.ConfigKeyFileAccessWhitelist); err != nil || sc.Value == "" {
- return []string{defaultPublicUploadType}
+ return []string{shared.DefaultPublicUploadType}
}
var whitelist []string
@@ -182,7 +124,7 @@ func parseFileAccessWhitelist(ctx context.Context) []string {
whitelist = parseCommaSeparatedWhitelist(sc.Value)
if len(whitelist) == 0 {
- return []string{defaultPublicUploadType}
+ return []string{shared.DefaultPublicUploadType}
}
return whitelist
}
diff --git a/internal/apps/upload/access_cache_test.go b/internal/apps/upload/cache/access_cache_test.go
similarity index 69%
rename from internal/apps/upload/access_cache_test.go
rename to internal/apps/upload/cache/access_cache_test.go
index 155023b5..1e7fb3d5 100644
--- a/internal/apps/upload/access_cache_test.go
+++ b/internal/apps/upload/cache/access_cache_test.go
@@ -1,13 +1,15 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+package cache
import (
"context"
"testing"
"time"
+ uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
@@ -19,14 +21,14 @@ func TestLoadMigrationAccessStateCachesResult(t *testing.T) {
ResetAccessCaches()
ctx := context.Background()
- first := loadMigrationAccessState(ctx)
- second := loadMigrationAccessState(ctx)
+ first := uploadstorage.LoadMigrationAccessState(ctx)
+ second := uploadstorage.LoadMigrationAccessState(ctx)
- if first.readOnly != second.readOnly {
- t.Fatalf("readOnly mismatch: first=%v second=%v", first.readOnly, second.readOnly)
+ if first.ReadOnly != second.ReadOnly {
+ t.Fatalf("readOnly mismatch: first=%v second=%v", first.ReadOnly, second.ReadOnly)
}
- if first.hasTarget != second.hasTarget {
- t.Fatalf("hasTarget mismatch: first=%v second=%v", first.hasTarget, second.hasTarget)
+ if first.HasTarget != second.HasTarget {
+ t.Fatalf("hasTarget mismatch: first=%v second=%v", first.HasTarget, second.HasTarget)
}
}
@@ -36,13 +38,13 @@ func TestIsFilePublicUsesCachedWhitelist(t *testing.T) {
ResetAccessCaches()
ctx := context.Background()
- if !isFilePublic(ctx, "avatar") {
+ if !IsFilePublic(ctx, "avatar") {
t.Fatal("expected avatar to be public by default")
}
- if isFilePublic(ctx, "attachment") {
+ if IsFilePublic(ctx, "attachment") {
t.Fatal("expected attachment to be private by default")
}
- if !isFilePublic(ctx, "AVATAR") {
+ if !IsFilePublic(ctx, "AVATAR") {
t.Fatal("expected whitelist lookup to be case-insensitive")
}
}
@@ -53,7 +55,7 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
ResetAccessCaches()
ctx := context.Background()
- if !isFilePublic(ctx, "avatar") {
+ if !IsFilePublic(ctx, "avatar") {
t.Fatal("expected seeded avatar whitelist before reset")
}
@@ -68,12 +70,13 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
if err := db.HSetJSON(ctx, model.SystemConfigRedisHashKey, model.ConfigKeyFileAccessWhitelist, &sc); err != nil {
t.Fatalf("refresh whitelist redis cache: %v", err)
}
+ model.ResetSystemConfigRAMCacheForTest()
ResetAccessCaches()
- if !isFilePublic(ctx, "attachment") {
+ if !IsFilePublic(ctx, "attachment") {
t.Fatal("expected attachment to be public after whitelist refresh")
}
- if isFilePublic(ctx, "avatar") {
+ if IsFilePublic(ctx, "avatar") {
t.Fatal("expected avatar to be private after whitelist refresh")
}
}
@@ -87,11 +90,11 @@ func TestAccessCacheTTLExpires(t *testing.T) {
_ = loadFileAccessWhitelist(ctx)
fileAccessWhitelistMu.Lock()
- fileAccessWhitelistCheckedAt = time.Now().Add(-accessCacheTTL - time.Second)
+ fileAccessWhitelistCheckedAt = time.Now().Add(-time.Duration(shared.AccessCacheTTL)*time.Second - time.Second)
fileAccessWhitelistMu.Unlock()
// Should still work after TTL by reloading from config.
- if !isFilePublic(ctx, "avatar") {
+ if !IsFilePublic(ctx, "avatar") {
t.Fatal("expected whitelist reload after TTL expiration")
}
}
\ No newline at end of file
diff --git a/internal/apps/upload/constants.go b/internal/apps/upload/constants.go
deleted file mode 100644
index 22bff1d6..00000000
--- a/internal/apps/upload/constants.go
+++ /dev/null
@@ -1,20 +0,0 @@
-// Copyright 2026 Arctel.net
-// SPDX-License-Identifier: Apache-2.0
-
-package upload
-
-import "github.com/Rain-kl/Wavelet/internal/storage"
-
-const (
- maxUploadSize = 32 * 1024 * 1024 // 32MB
- detectContentBytes = 512 // http.DetectContentType 需要的最小字节数
- uploadDirPerm = 0755 // 上传目录权限
- uploadFilePerm = 0644 // 上传文件权限
- imageQualityLow = "low"
- imageQualityMedium = "medium"
- imageQualityHigh = "high"
- imageQualityOrigin = "origin"
- storageDriverLocal = string(storage.DriverLocal)
- defaultPublicUploadType = "avatar"
- fileStatsTrendDays = 7
-)
diff --git a/internal/apps/upload/errs.go b/internal/apps/upload/errs.go
index 5b2cafd6..e146bc4a 100644
--- a/internal/apps/upload/errs.go
+++ b/internal/apps/upload/errs.go
@@ -5,37 +5,34 @@
// Package upload 提供文件上传与下载功能
package upload
+import "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
+
// 文件管理常量
const (
- ErrNoFileSelected = "请选择要上传的文件"
- ErrUnsupportedFormat = "只支持 JPG、PNG、WEBP 格式的图片"
- ErrProcessFileFailed = "处理文件失败"
- ErrSaveFileFailed = "保存文件失败"
- ErrOpenFileFailed = "打开文件失败"
- ErrSaveUploadRecordFailed = "保存上传记录失败"
- ErrGenericFileTooLarge = "文件大小不能超过 32MB"
- ErrFileContentExtensionMismatch = "文件内容与扩展名不匹配,可能包含安全风险"
- ErrFileValidationFailed = "文件校验失败"
- ErrInvalidMetadataJSON = "元数据 JSON 格式不合法"
- ErrInvalidFileID = "无效的文件 ID"
- ErrQueryUploadRecordFailed = "查询文件记录失败"
- ErrInvalidBatchDownloadRequest = "参数绑定失败,请传入有效的文件 ID 数组"
- ErrInvalidIDValueFormat = "无效的 ID 值: %s"
- ErrRetrieveUploadRecordsFailed = "检索文件记录失败"
- ErrNoValidFilesForArchive = "没有找到任何有效的文件记录进行打包"
- ErrInvalidParams = "参数错误"
- ErrQueryFileCountFailed = "查询文件数量失败"
- ErrQueryFileListFailed = "查询文件列表失败"
- ErrDeleteFileFailed = "删除文件失败"
- ErrStorageReadOnly = "存储迁移维护中,当前仅允许读取文件"
- ErrS3KeyRequired = "s3 key must not be empty"
- ErrS3KeyTooLongFormat = "s3 key exceeds maximum length of %d"
- ErrS3KeyStartsWithSlash = "s3 key must not start with /"
- ErrS3KeyContainsNullBytes = "s3 key must not contain null bytes"
- ErrQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w"
- errImageCacheWarmupPayloadRequired = "图片缓存预热参数不能为空"
- errInvalidImageCacheWarmupPayload = "图片缓存预热参数格式无效: %w"
- errInvalidImageCacheWarmupQuality = "图片质量仅支持 low、medium、high"
- errParseImageCacheWarmupPayload = "解析图片缓存预热参数失败: %w"
- errQueryImagesForCacheWarmup = "查询待预热图片失败: %w"
-)
+ ErrNoFileSelected = shared.ErrNoFileSelected
+ ErrUnsupportedFormat = shared.ErrUnsupportedFormat
+ ErrProcessFileFailed = shared.ErrProcessFileFailed
+ ErrSaveFileFailed = shared.ErrSaveFileFailed
+ ErrOpenFileFailed = shared.ErrOpenFileFailed
+ ErrSaveUploadRecordFailed = shared.ErrSaveUploadRecordFailed
+ ErrGenericFileTooLarge = shared.ErrGenericFileTooLarge
+ ErrFileContentExtensionMismatch = shared.ErrFileContentExtensionMismatch
+ ErrFileValidationFailed = shared.ErrFileValidationFailed
+ ErrInvalidMetadataJSON = shared.ErrInvalidMetadataJSON
+ ErrInvalidFileID = shared.ErrInvalidFileID
+ ErrQueryUploadRecordFailed = shared.ErrQueryUploadRecordFailed
+ ErrInvalidBatchDownloadRequest = shared.ErrInvalidBatchDownloadRequest
+ ErrInvalidIDValueFormat = shared.ErrInvalidIDValueFormat
+ ErrRetrieveUploadRecordsFailed = shared.ErrRetrieveUploadRecordsFailed
+ ErrNoValidFilesForArchive = shared.ErrNoValidFilesForArchive
+ ErrInvalidParams = shared.ErrInvalidParams
+ ErrQueryFileCountFailed = shared.ErrQueryFileCountFailed
+ ErrQueryFileListFailed = shared.ErrQueryFileListFailed
+ ErrDeleteFileFailed = shared.ErrDeleteFileFailed
+ ErrStorageReadOnly = shared.ErrStorageReadOnly
+ ErrS3KeyRequired = shared.ErrS3KeyRequired
+ ErrS3KeyTooLongFormat = shared.ErrS3KeyTooLongFormat
+ ErrS3KeyStartsWithSlash = shared.ErrS3KeyStartsWithSlash
+ ErrS3KeyContainsNullBytes = shared.ErrS3KeyContainsNullBytes
+ ErrQueryUnusedUploadsFailed = shared.ErrQueryUnusedUploadsFailed
+)
\ No newline at end of file
diff --git a/internal/apps/upload/exports.go b/internal/apps/upload/exports.go
new file mode 100644
index 00000000..0ea8a489
--- /dev/null
+++ b/internal/apps/upload/exports.go
@@ -0,0 +1,86 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package upload
+
+import (
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/handler"
+ uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
+ uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/util"
+ "github.com/Rain-kl/Wavelet/internal/task"
+)
+
+// HTTP handlers
+var (
+ UploadFile = handler.UploadFile
+ DownloadFile = handler.DownloadFile
+ BatchDownloadFiles = handler.BatchDownloadFiles
+ ListFiles = handler.ListFiles
+ DeleteFile = handler.DeleteFile
+ GetDistinctUploadTypes = handler.GetDistinctUploadTypes
+ ListMyFiles = handler.ListMyFiles
+ DeleteMyFile = handler.DeleteMyFile
+ UpdateMyFile = handler.UpdateMyFile
+ GetFileStats = handler.GetFileStats
+ ServeFileByID = filesrv.ServeFileByID
+)
+
+// Cache management
+var (
+ ResetAccessCaches = cache.ResetAccessCaches
+ PublishAccessCacheInvalidation = cache.PublishAccessCacheInvalidation
+)
+
+// Stats
+var (
+ ApplyUploadStatsAdd = uploadstats.ApplyUploadStatsAdd
+ ApplyUploadStatsRemove = uploadstats.ApplyUploadStatsRemove
+ RebuildUploadStats = uploadstats.RebuildUploadStats
+)
+
+// Utilities
+var (
+ CompressImageToWebP = util.CompressImageToWebP
+ ValidateS3Key = util.ValidateS3Key
+)
+
+// Task identifiers and metadata
+const (
+ StorageMigrationTask = uploadtask.StorageMigrationTask
+ SystemCleanupTask = uploadtask.SystemCleanupTask
+ WarmImageCacheTask = uploadtask.WarmImageCacheTask
+)
+
+var (
+ // StorageMigrationMeta describes the storage migration async task.
+ StorageMigrationMeta = uploadtask.StorageMigrationMeta
+ // SystemCleanupMeta describes the orphaned upload cleanup task.
+ SystemCleanupMeta = uploadtask.SystemCleanupMeta
+ // WarmImageCacheMeta describes the image compression cache warmup task.
+ WarmImageCacheMeta = uploadtask.WarmImageCacheMeta
+)
+
+// MigrationHandler executes storage migration tasks.
+type MigrationHandler = uploadtask.MigrationHandler
+
+// SystemCleanupHandler removes orphaned upload files.
+type SystemCleanupHandler = uploadtask.SystemCleanupHandler
+
+// WarmImageCacheHandler pre-warms compressed image caches.
+type WarmImageCacheHandler = uploadtask.WarmImageCacheHandler
+
+// WarmImageCachePayload is the payload for image cache warmup tasks.
+type WarmImageCachePayload = uploadtask.WarmImageCachePayload
+
+// Ensure task handler types implement required interfaces.
+var (
+ _ task.TaskHandler = (*MigrationHandler)(nil)
+ _ task.TaskHandler = (*SystemCleanupHandler)(nil)
+ _ interface {
+ task.TaskHandler
+ ValidatePayload([]byte) ([]byte, error)
+ } = (*WarmImageCacheHandler)(nil)
+)
\ No newline at end of file
diff --git a/internal/apps/upload/file_server.go b/internal/apps/upload/filesrv/file_server.go
similarity index 70%
rename from internal/apps/upload/file_server.go
rename to internal/apps/upload/filesrv/file_server.go
index 4fd247ca..d4a3156a 100644
--- a/internal/apps/upload/file_server.go
+++ b/internal/apps/upload/filesrv/file_server.go
@@ -2,7 +2,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+// Package filesrv serves uploaded files with access control and image compression.
+package filesrv
import (
"bytes"
@@ -15,11 +16,16 @@ import (
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
+ uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/util"
"github.com/Rain-kl/Wavelet/internal/common"
+ "github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
- "github.com/Rain-kl/Wavelet/internal/util"
+ apputil "github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"golang.org/x/sync/singleflight"
@@ -34,6 +40,15 @@ type compressedImageCacheResult struct {
err error
}
+type fileTypeCategory string
+
+const (
+ fileTypeImage fileTypeCategory = "image"
+ fileTypeVideo fileTypeCategory = "video"
+ fileTypeAudio fileTypeCategory = "audio"
+ fileTypeOther fileTypeCategory = "other"
+)
+
// ServeFileByID 根据 ID 获取并提供已上传的文件
// @Summary 获取已上传文件
// @Description 根据文件 ID 获取并提供已上传的临时或正式文件,若配置了缓存则优先走本地缓存,否则从 S3 等后端存储读取并流式返回
@@ -48,7 +63,7 @@ type compressedImageCacheResult struct {
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /f/{id} [get]
func ServeFileByID(c *gin.Context) {
- upload, err := getUploadRecordByID(c)
+ upload, err := GetUploadRecordByID(c)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.AbortWithStatus(http.StatusNotFound)
@@ -62,18 +77,16 @@ func ServeFileByID(c *gin.Context) {
return
}
- // 校验业务白名单与访问权限
- if err := checkFileAccessPermission(c, upload); err != nil {
- c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil})
+ if err := CheckFileAccessPermission(c, upload); err != nil {
+ response.AbortUnauthorized(c, common.UnAuthorized)
return
}
ServeUpload(c, upload)
}
-// getUploadRecordByID 从请求路径参数中解析文件 ID 并从数据库中检索处于 Pending 或 Used 状态的上传记录。
-// 同时会自动设置通用的安全响应头。
-func getUploadRecordByID(c *gin.Context) (*model.Upload, error) {
+// GetUploadRecordByID 从请求路径参数中解析文件 ID 并从数据库中检索处于 Pending 或 Used 状态的上传记录。
+func GetUploadRecordByID(c *gin.Context) (*model.Upload, error) {
c.Header("X-Content-Type-Options", "nosniff")
c.Header("Content-Security-Policy", "sandbox")
@@ -93,22 +106,11 @@ func getUploadRecordByID(c *gin.Context) (*model.Upload, error) {
return &upload, nil
}
-// fileTypeCategory 定义文件的大类,便于未来扩展不同的处理方式
-type fileTypeCategory string
-
-const (
- fileTypeImage fileTypeCategory = "image"
- fileTypeVideo fileTypeCategory = "video"
- fileTypeAudio fileTypeCategory = "audio"
- fileTypeOther fileTypeCategory = "other"
-)
-
-// getFileTypeCategory 判断并返回文件的大类
func getFileTypeCategory(upload *model.Upload) fileTypeCategory {
mime := strings.ToLower(upload.MimeType)
ext := strings.ToLower(upload.Extension)
- if strings.HasPrefix(mime, "image/") || isImageExtension(ext) {
+ if strings.HasPrefix(mime, "image/") || util.IsImageExtension(ext) {
return fileTypeImage
}
if strings.HasPrefix(mime, "video/") {
@@ -120,32 +122,27 @@ func getFileTypeCategory(upload *model.Upload) fileTypeCategory {
return fileTypeOther
}
-// ServeUpload 将已存在的文件内容读取并流式响应给客户端,支持本地和 S3/CDN 驱动,并可选支持 WebP 图片压缩与本地缓存。
+// ServeUpload 将已存在的文件内容读取并流式响应给客户端。
func ServeUpload(c *gin.Context, upload *model.Upload) {
- // 设置通用的缓存控制响应头
setCacheHeaders(c, upload)
category := getFileTypeCategory(upload)
- quality := normalizeImageQuality(c.Query("quality"))
+ quality := util.NormalizeImageQuality(c.Query("quality"))
switch category {
case fileTypeImage:
- // 如果是图片且不是原图质量,则提供压缩优化后的图片预览
- if quality != imageQualityOrigin {
+ if quality != shared.ImageQualityOrigin {
serveCompressedImage(c, upload, quality)
return
}
- // 请求原图质量时,退化到默认提供原文件
fallthrough
-
default:
- // 默认提供原文件,并执行协商缓存校验
serveOriginalWithConditionalCheck(c, upload)
}
}
func setCacheHeaders(c *gin.Context, upload *model.Upload) {
- if isFilePublic(c.Request.Context(), upload.Type) {
+ if cache.IsFilePublic(c.Request.Context(), upload.Type) {
c.Header("Cache-Control", "public, max-age=31536000")
} else {
c.Header("Cache-Control", "private, no-cache")
@@ -173,7 +170,7 @@ func serveCompressedImage(c *gin.Context, upload *model.Upload, quality string)
return
}
- webpBytes, _, err := ensureCompressedImageCache(c.Request.Context(), upload, quality)
+ webpBytes, _, err := EnsureCompressedImageCache(c.Request.Context(), upload, quality)
if err != nil {
if len(webpBytes) > 0 {
logger.WarnF(c.Request.Context(), "failed to cache compressed image: %v", err)
@@ -188,14 +185,15 @@ func serveCompressedImage(c *gin.Context, upload *model.Upload, quality string)
c.Data(http.StatusOK, "image/webp", webpBytes)
}
-func ensureCompressedImageCache(
+// EnsureCompressedImageCache returns cached or freshly generated WebP bytes for an upload.
+func EnsureCompressedImageCache(
ctx context.Context,
upload *model.Upload,
quality string,
) ([]byte, bool, error) {
- cache := diskcache.GetGlobalCache()
- cacheKey := imageCompressionCacheKey(upload, quality)
- webpBytes, err := cache.Get(cacheKey)
+ cacheStore := diskcache.GetGlobalCache()
+ cacheKey := ImageCompressionCacheKey(upload, quality)
+ webpBytes, err := cacheStore.Get(cacheKey)
if err == nil {
return webpBytes, true, nil
}
@@ -220,9 +218,9 @@ func generateCompressedImageCache(
quality string,
cacheKey string,
) (compressedImageCacheResult, error) {
- cache := diskcache.GetGlobalCache()
+ cacheStore := diskcache.GetGlobalCache()
- webpBytes, err := cache.Get(cacheKey)
+ webpBytes, err := cacheStore.Get(cacheKey)
if err == nil {
return compressedImageCacheResult{bytes: webpBytes, cached: true}, nil
}
@@ -235,12 +233,12 @@ func generateCompressedImageCache(
return compressedImageCacheResult{}, fmt.Errorf("read original image: %w", err)
}
- webpBytes, err = CompressImageToWebP(bytes.NewReader(origBytes), quality)
+ webpBytes, err = util.CompressImageToWebP(bytes.NewReader(origBytes), quality)
if err != nil {
return compressedImageCacheResult{}, fmt.Errorf("compress image to WebP: %w", err)
}
- if err := cache.Set(cacheKey, webpBytes, diskcache.NoExpiration); err != nil {
+ if err := cacheStore.Set(cacheKey, webpBytes, diskcache.NoExpiration); err != nil {
return compressedImageCacheResult{
bytes: webpBytes,
err: fmt.Errorf("write compressed image cache: %w", err),
@@ -250,7 +248,8 @@ func generateCompressedImageCache(
return compressedImageCacheResult{bytes: webpBytes}, nil
}
-func imageCompressionCacheKey(upload *model.Upload, quality string) string {
+// ImageCompressionCacheKey returns the disk cache key for a compressed upload image.
+func ImageCompressionCacheKey(upload *model.Upload, quality string) string {
return fmt.Sprintf(
"upload_webp_v1_%d_%d_%d_%s_%s",
upload.ID,
@@ -261,18 +260,8 @@ func imageCompressionCacheKey(upload *model.Upload, quality string) string {
)
}
-func normalizeImageQuality(quality string) string {
- switch strings.ToLower(quality) {
- case imageQualityLow, imageQualityMedium, imageQualityHigh:
- return strings.ToLower(quality)
- default:
- return imageQualityOrigin
- }
-}
-
-// serveOriginal 原始文件的流式响应逻辑
func serveOriginal(c *gin.Context, upload *model.Upload) {
- obj, err := openStoredObject(c.Request.Context(), upload)
+ obj, err := uploadstorage.OpenStoredObject(c.Request.Context(), upload)
if err != nil {
c.AbortWithStatus(http.StatusNotFound)
return
@@ -281,9 +270,8 @@ func serveOriginal(c *gin.Context, upload *model.Upload) {
c.DataFromReader(http.StatusOK, obj.ContentLength, obj.ContentType, obj.Body, nil)
}
-// getOriginalFileBytes 获取原始文件所有字节
func getOriginalFileBytes(ctx context.Context, upload *model.Upload) ([]byte, error) {
- obj, err := openStoredObject(ctx, upload)
+ obj, err := uploadstorage.OpenStoredObject(ctx, upload)
if err != nil {
return nil, err
}
@@ -291,17 +279,10 @@ func getOriginalFileBytes(ctx context.Context, upload *model.Upload) ([]byte, er
return io.ReadAll(obj.Body)
}
-// isFilePublic 校验文件类型是否在公开访问白名单中
-func isFilePublic(ctx context.Context, uploadType string) bool {
- whitelist := loadFileAccessWhitelist(ctx)
- _, ok := whitelist[strings.ToLower(uploadType)]
- return ok
-}
-
func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
var currUser *model.User
var err error
- if u, ok := util.GetFromContext[*model.User](c, oauth.UserObjKey); ok && u != nil {
+ if u, ok := apputil.GetFromContext[*model.User](c, oauth.UserObjKey); ok && u != nil {
currUser = u
} else {
currUser, err = oauth.GetUserFromRequest(c)
@@ -318,21 +299,18 @@ func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
return nil
}
-// checkFileAccessPermission 校验文件是否可以被当前请求访问
-func checkFileAccessPermission(c *gin.Context, upload *model.Upload) error {
- // 1. 私有文件校验(优先级高于当前白名单逻辑)
+// CheckFileAccessPermission 校验文件是否可以被当前请求访问
+func CheckFileAccessPermission(c *gin.Context, upload *model.Upload) error {
if upload.AccessMode == 0 {
return checkPrivateFileOwner(c, upload.UserID)
}
- // 2. 如果类型为公开的则再进行校验白名单
- if !isFilePublic(c.Request.Context(), upload.Type) {
- // 必须进行鉴权
- if _, ok := util.GetFromContext[*model.User](c, oauth.UserObjKey); !ok {
+ if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
+ if _, ok := apputil.GetFromContext[*model.User](c, oauth.UserObjKey); !ok {
if _, err := oauth.GetUserFromRequest(c); err != nil {
return err
}
}
}
return nil
-}
+}
\ No newline at end of file
diff --git a/internal/apps/upload/file_server_test.go b/internal/apps/upload/filesrv/file_server_test.go
similarity index 83%
rename from internal/apps/upload/file_server_test.go
rename to internal/apps/upload/filesrv/file_server_test.go
index 779fcb33..5dc426e2 100644
--- a/internal/apps/upload/file_server_test.go
+++ b/internal/apps/upload/filesrv/file_server_test.go
@@ -2,7 +2,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+package filesrv
import (
"bytes"
@@ -15,7 +15,11 @@ import (
"os"
"testing"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/util"
"github.com/Rain-kl/Wavelet/internal/common"
+ "github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
@@ -27,6 +31,7 @@ import (
func TestServeFileByIDAccessControl(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
+ cache.ResetAccessCaches()
// Ensure uploads dir is cleaned up
defer func() { _ = os.RemoveAll("uploads") }()
@@ -91,6 +96,7 @@ func TestServeFileByIDAccessControl(t *testing.T) {
// Set up router
gin.SetMode(gin.TestMode)
r := gin.New()
+ r.Use(response.ErrorHandlerMiddleware())
store := cookie.NewStore([]byte("secret"))
r.Use(sessions.Sessions("test_session", store))
r.GET("/f/:id", ServeFileByID)
@@ -151,58 +157,6 @@ func TestServeFileByIDAccessControl(t *testing.T) {
})
}
-func TestGetDistinctUploadTypes(t *testing.T) {
- dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
- defer cleanup()
-
- // Seed some uploads with new custom types
- user := model.User{ID: 2222, Username: "test_user_2"}
- dbConn.Create(&user)
-
- customUpload := model.Upload{
- ID: 9001,
- UserID: user.ID,
- FileName: "custom.txt",
- FilePath: "uploads/custom.txt",
- FileSize: 10,
- MimeType: "text/plain",
- Extension: "txt",
- StorageDriver: "local",
- Type: "custom_type_xyz",
- Status: model.UploadStatusUsed,
- }
- dbConn.Create(&customUpload)
-
- gin.SetMode(gin.TestMode)
- r := gin.New()
- r.GET("/api/v1/admin/uploads/types", GetDistinctUploadTypes)
-
- req, _ := http.NewRequest("GET", "/api/v1/admin/uploads/types", nil)
- w := httptest.NewRecorder()
- r.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("expected 200, got %d", w.Code)
- }
-
- var resp struct {
- ErrorMsg string `json:"error_msg"`
- Data []string `json:"data"`
- }
- if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
- t.Fatalf("failed to parse JSON: %v", err)
- }
-
- if resp.ErrorMsg != "" {
- t.Fatalf("unexpected error: %s", resp.ErrorMsg)
- }
-
- // Verify that only custom_type_xyz is present
- if len(resp.Data) != 1 || resp.Data[0] != "custom_type_xyz" {
- t.Errorf("expected only custom_type_xyz in types list, got: %v", resp.Data)
- }
-}
-
func TestImageCompression(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
@@ -295,7 +249,7 @@ func TestImageCompression(t *testing.T) {
t.Errorf("expected Content-Type image/webp, got %s", w.Header().Get("Content-Type"))
}
- cacheKey := imageCompressionCacheKey(&uploadRecord, imageQualityMedium)
+ cacheKey := ImageCompressionCacheKey(&uploadRecord, shared.ImageQualityMedium)
cachedBytes, err := cache.Get(cacheKey)
if err != nil {
t.Fatalf("disk cache Get(%q) returned error: %v", cacheKey, err)
@@ -376,19 +330,19 @@ func TestNormalizeImageQuality(t *testing.T) {
quality string
want string
}{
- {name: imageQualityLow, quality: imageQualityLow, want: imageQualityLow},
- {name: imageQualityMedium, quality: imageQualityMedium, want: imageQualityMedium},
- {name: imageQualityHigh, quality: imageQualityHigh, want: imageQualityHigh},
+ {name: shared.ImageQualityLow, quality: shared.ImageQualityLow, want: shared.ImageQualityLow},
+ {name: shared.ImageQualityMedium, quality: shared.ImageQualityMedium, want: shared.ImageQualityMedium},
+ {name: shared.ImageQualityHigh, quality: shared.ImageQualityHigh, want: shared.ImageQualityHigh},
{name: "origin", quality: "origin", want: "origin"},
- {name: "uppercase", quality: "LOW", want: imageQualityLow},
+ {name: "uppercase", quality: "LOW", want: shared.ImageQualityLow},
{name: "empty", quality: "", want: "origin"},
{name: "invalid", quality: "maximum", want: "origin"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- if got := normalizeImageQuality(tt.quality); got != tt.want {
- t.Errorf("normalizeImageQuality(%q) = %q, want %q", tt.quality, got, tt.want)
+ if got := util.NormalizeImageQuality(tt.quality); got != tt.want {
+ t.Errorf("NormalizeImageQuality(%q) = %q, want %q", tt.quality, got, tt.want)
}
})
}
diff --git a/internal/apps/upload/file_management.go b/internal/apps/upload/handler/file_management.go
similarity index 84%
rename from internal/apps/upload/file_management.go
rename to internal/apps/upload/handler/file_management.go
index af14f8cf..2f005c2d 100644
--- a/internal/apps/upload/file_management.go
+++ b/internal/apps/upload/handler/file_management.go
@@ -1,7 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+package handler
import (
"errors"
@@ -11,10 +11,13 @@ import (
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
+ uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
+ uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
- "github.com/Rain-kl/Wavelet/internal/util"
+ apputil "github.com/Rain-kl/Wavelet/internal/util"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
@@ -56,7 +59,7 @@ func ListFiles(c *gin.Context) {
var req listFilesRequest
if err := c.ShouldBindQuery(&req); err != nil {
- c.JSON(http.StatusOK, response.Err(ErrInvalidParams))
+ response.AbortBadRequest(c, shared.ErrInvalidParams)
return
}
if req.Page <= 0 {
@@ -84,14 +87,14 @@ func ListFiles(c *gin.Context) {
var total int64
if err := query.Count(&total).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(ErrQueryFileCountFailed))
+ response.AbortBadRequest(c, shared.ErrQueryFileCountFailed)
return
}
var items []model.Upload
offset := (req.Page - 1) * req.PageSize
if err := query.Order("created_at DESC").Offset(offset).Limit(req.PageSize).Find(&items).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(ErrQueryFileListFailed))
+ response.AbortBadRequest(c, shared.ErrQueryFileListFailed)
return
}
@@ -116,14 +119,14 @@ func ListFiles(c *gin.Context) {
// @Router /api/v1/admin/uploads/{id} [delete]
func DeleteFile(c *gin.Context) {
ctx := c.Request.Context()
- if StorageReadOnly(ctx) {
- c.JSON(http.StatusConflict, response.Err(ErrStorageReadOnly))
+ if uploadstorage.ReadOnly(ctx) {
+ response.AbortConflict(c, shared.ErrStorageReadOnly)
return
}
uploadID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
- c.JSON(http.StatusOK, response.Err(ErrInvalidFileID))
+ response.AbortBadRequest(c, shared.ErrInvalidFileID)
return
}
@@ -133,14 +136,14 @@ func DeleteFile(c *gin.Context) {
c.AbortWithStatus(http.StatusNotFound)
return
}
- c.JSON(http.StatusOK, response.Err(ErrQueryUploadRecordFailed))
+ response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed)
return
}
if err := db.DB(ctx).Model(&upload).Update("status", model.UploadStatusDeleted).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(ErrDeleteFileFailed))
+ response.AbortBadRequest(c, shared.ErrDeleteFileFailed)
return
}
- recordUploadStatsRemove(ctx, &upload)
+ uploadstats.RecordUploadStatsRemove(ctx, &upload)
c.JSON(http.StatusOK, response.OKNil())
}
@@ -161,7 +164,7 @@ func GetDistinctUploadTypes(c *gin.Context) {
Where("type IS NOT NULL AND type != ''").
Distinct().
Pluck("type", &dbTypes).Error; err != nil {
- c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
+ response.AbortInternal(c, err.Error())
return
}
sort.Strings(dbTypes)
@@ -198,12 +201,12 @@ type listMyFilesResponse struct {
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/upload/my [get]
func ListMyFiles(c *gin.Context) {
- currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
+ currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
var req listMyFilesRequest
if err := c.ShouldBindQuery(&req); err != nil {
- c.JSON(http.StatusOK, response.Err(ErrInvalidParams))
+ response.AbortBadRequest(c, shared.ErrInvalidParams)
return
}
if req.Page <= 0 {
@@ -228,14 +231,14 @@ func ListMyFiles(c *gin.Context) {
var total int64
if err := query.Count(&total).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(ErrQueryFileCountFailed))
+ response.AbortBadRequest(c, shared.ErrQueryFileCountFailed)
return
}
var items []model.Upload
offset := (req.Page - 1) * req.PageSize
if err := query.Order("created_at DESC").Offset(offset).Limit(req.PageSize).Find(&items).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(ErrQueryFileListFailed))
+ response.AbortBadRequest(c, shared.ErrQueryFileListFailed)
return
}
@@ -259,16 +262,16 @@ func ListMyFiles(c *gin.Context) {
// @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/upload/{id} [delete]
func DeleteMyFile(c *gin.Context) {
- currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
+ currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
- if StorageReadOnly(ctx) {
- c.JSON(http.StatusConflict, response.Err(ErrStorageReadOnly))
+ if uploadstorage.ReadOnly(ctx) {
+ response.AbortConflict(c, shared.ErrStorageReadOnly)
return
}
uploadID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
- c.JSON(http.StatusOK, response.Err(ErrInvalidFileID))
+ response.AbortBadRequest(c, shared.ErrInvalidFileID)
return
}
@@ -278,7 +281,7 @@ func DeleteMyFile(c *gin.Context) {
c.AbortWithStatus(http.StatusNotFound)
return
}
- c.JSON(http.StatusOK, response.Err(ErrQueryUploadRecordFailed))
+ response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed)
return
}
@@ -288,10 +291,10 @@ func DeleteMyFile(c *gin.Context) {
}
if err := db.DB(ctx).Model(&upload).Update("status", model.UploadStatusDeleted).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(ErrDeleteFileFailed))
+ response.AbortBadRequest(c, shared.ErrDeleteFileFailed)
return
}
- recordUploadStatsRemove(ctx, &upload)
+ uploadstats.RecordUploadStatsRemove(ctx, &upload)
c.JSON(http.StatusOK, response.OKNil())
}
@@ -314,22 +317,22 @@ type updateMyFileRequest struct {
// @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/upload/{id} [put]
func UpdateMyFile(c *gin.Context) {
- currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
+ currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
- if StorageReadOnly(ctx) {
- c.JSON(http.StatusConflict, response.Err(ErrStorageReadOnly))
+ if uploadstorage.ReadOnly(ctx) {
+ response.AbortConflict(c, shared.ErrStorageReadOnly)
return
}
uploadID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
- c.JSON(http.StatusOK, response.Err(ErrInvalidFileID))
+ response.AbortBadRequest(c, shared.ErrInvalidFileID)
return
}
var req updateMyFileRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusOK, response.Err(ErrInvalidParams))
+ response.AbortBadRequest(c, shared.ErrInvalidParams)
return
}
@@ -339,7 +342,7 @@ func UpdateMyFile(c *gin.Context) {
c.AbortWithStatus(http.StatusNotFound)
return
}
- c.JSON(http.StatusOK, response.Err(ErrQueryUploadRecordFailed))
+ response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed)
return
}
@@ -358,10 +361,10 @@ func UpdateMyFile(c *gin.Context) {
if len(updates) > 0 {
if err := db.DB(ctx).Model(&upload).Updates(updates).Error; err != nil {
- c.JSON(http.StatusOK, response.Err("更新文件记录失败"))
+ response.AbortBadRequest(c, "更新文件记录失败")
return
}
}
c.JSON(http.StatusOK, response.OK(upload))
-}
+}
\ No newline at end of file
diff --git a/internal/apps/upload/handler/file_management_test.go b/internal/apps/upload/handler/file_management_test.go
new file mode 100644
index 00000000..7607b526
--- /dev/null
+++ b/internal/apps/upload/handler/file_management_test.go
@@ -0,0 +1,65 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package handler
+
+import (
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/Rain-kl/Wavelet/internal/testhelper"
+ "github.com/gin-gonic/gin"
+)
+
+func TestGetDistinctUploadTypes(t *testing.T) {
+ dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
+ defer cleanup()
+
+ user := model.User{ID: 2222, Username: "test_user_2"}
+ dbConn.Create(&user)
+
+ customUpload := model.Upload{
+ ID: 9001,
+ UserID: user.ID,
+ FileName: "custom.txt",
+ FilePath: "uploads/custom.txt",
+ FileSize: 10,
+ MimeType: "text/plain",
+ Extension: "txt",
+ StorageDriver: "local",
+ Type: "custom_type_xyz",
+ Status: model.UploadStatusUsed,
+ }
+ dbConn.Create(&customUpload)
+
+ gin.SetMode(gin.TestMode)
+ r := gin.New()
+ r.GET("/api/v1/admin/uploads/types", GetDistinctUploadTypes)
+
+ req, _ := http.NewRequest("GET", "/api/v1/admin/uploads/types", nil)
+ w := httptest.NewRecorder()
+ r.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("expected 200, got %d", w.Code)
+ }
+
+ var resp struct {
+ ErrorMsg string `json:"error_msg"`
+ Data []string `json:"data"`
+ }
+ if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
+ t.Fatalf("failed to parse JSON: %v", err)
+ }
+
+ if resp.ErrorMsg != "" {
+ t.Fatalf("unexpected error: %s", resp.ErrorMsg)
+ }
+
+ if len(resp.Data) != 1 || resp.Data[0] != "custom_type_xyz" {
+ t.Errorf("expected only custom_type_xyz in types list, got: %v", resp.Data)
+ }
+}
\ No newline at end of file
diff --git a/internal/apps/upload/routers.go b/internal/apps/upload/handler/routers.go
similarity index 71%
rename from internal/apps/upload/routers.go
rename to internal/apps/upload/handler/routers.go
index 15d71e37..7f0953e7 100644
--- a/internal/apps/upload/routers.go
+++ b/internal/apps/upload/handler/routers.go
@@ -2,9 +2,11 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+// Package handler provides upload HTTP API handlers.
+package handler
-import ("archive/zip"
+import (
+ "archive/zip"
"bytes"
"context"
"crypto/sha256"
@@ -22,16 +24,21 @@ import ("archive/zip"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
+ uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
+ uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/util"
"github.com/Rain-kl/Wavelet/internal/common"
+ "github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
- "github.com/Rain-kl/Wavelet/internal/util"
+ apputil "github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
- "github.com/Rain-kl/Wavelet/internal/common/response"
)
type batchDownloadRequest struct {
@@ -59,59 +66,53 @@ func UploadFile(c *gin.Context) {
c.Header("X-Content-Type-Options", "nosniff")
c.Header("Content-Security-Policy", "sandbox")
- // 限制请求体大小以防止 DoS
- c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, maxUploadSize)
+ c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize)
- currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
+ currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
header, err := c.FormFile("file")
if err != nil {
- c.JSON(http.StatusOK, response.Err(ErrNoFileSelected))
+ response.AbortBadRequest(c, shared.ErrNoFileSelected)
return
}
file, err := header.Open()
if err != nil {
- c.JSON(http.StatusOK, response.Err(ErrOpenFileFailed))
+ response.AbortBadRequest(c, shared.ErrOpenFileFailed)
return
}
defer func() { _ = file.Close() }()
- // 校验大小
- if header.Size > maxUploadSize {
- c.JSON(http.StatusOK, response.Err(ErrGenericFileTooLarge))
+ if header.Size > shared.MaxUploadSize {
+ response.AbortBadRequest(c, shared.ErrGenericFileTooLarge)
return
}
- // 2. 提取文件基本元数据
origName := header.Filename
ext := strings.ToLower(strings.TrimPrefix(filepath.Ext(origName), "."))
if ext == "" {
ext = "bin"
}
- // 3. 校验文件后缀是否在允许的系统配置列表中
if errMsg := validateUploadExtension(ctx, ext); errMsg != "" {
- c.JSON(http.StatusOK, response.Err(errMsg))
+ response.AbortBadRequest(c, errMsg)
return
}
- // 4. 读取文件并计算 Hash
hashWriter := sha256.New()
var buf bytes.Buffer
size, err := io.Copy(&buf, io.TeeReader(file, hashWriter))
if err != nil {
- c.JSON(http.StatusOK, response.Err(ErrProcessFileFailed))
+ response.AbortBadRequest(c, shared.ErrProcessFileFailed)
return
}
fileHash := hex.EncodeToString(hashWriter.Sum(nil))
mimeType := detectMimeType(&buf, header, size)
- // 校验真实 MIME Type 是否与常见图片扩展名匹配,防止 Polyglot / HTML 注入攻击
- if isImageExtension(ext) && !strings.HasPrefix(mimeType, "image/") {
- c.JSON(http.StatusOK, response.Err(ErrFileContentExtensionMismatch))
+ if util.IsImageExtension(ext) && !strings.HasPrefix(mimeType, "image/") {
+ response.AbortBadRequest(c, shared.ErrFileContentExtensionMismatch)
return
}
@@ -119,38 +120,34 @@ func UploadFile(c *gin.Context) {
accessMode, errMsg := resolveUploadAccessMode(c, uploadType)
if errMsg != "" {
- c.JSON(http.StatusOK, response.Err(errMsg))
+ response.AbortBadRequest(c, errMsg)
return
}
- // 6. 秒传匹配校验:校验数据库中是否存在相同 Hash 且大小一致的可用文件
handled, lookupErr := tryInstantUpload(ctx, c, currUser, fileHash, size, mimeType, ext, origName, accessMode)
if handled {
return
}
if lookupErr != nil && !errors.Is(lookupErr, gorm.ErrRecordNotFound) {
- c.JSON(http.StatusOK, response.Err(ErrFileValidationFailed))
+ response.AbortBadRequest(c, shared.ErrFileValidationFailed)
return
}
- // 7. 解析可选元数据字段
meta, errMsg := parseUploadMetadata(c, mimeType)
if errMsg != "" {
- c.JSON(http.StatusOK, response.Err(errMsg))
+ response.AbortBadRequest(c, errMsg)
return
}
id := idgen.NextUint64ID()
subPath := fmt.Sprintf("uploads/%s/%d.%s", time.Now().Format("2006/01/02"), id, ext)
- // 8. 写入当前活动存储驱动。
storageDriver, subPath, errMsg := storeUploadFile(ctx, subPath, size, mimeType, &buf, &meta)
if errMsg != "" {
- c.JSON(http.StatusOK, response.Err(errMsg))
+ response.AbortBadRequest(c, errMsg)
return
}
- // 9. 保存文件记录至数据库
newUpload := model.Upload{
ID: id,
UserID: currUser.ID,
@@ -168,7 +165,7 @@ func UploadFile(c *gin.Context) {
}
if err := saveUploadRecord(ctx, &newUpload, storageDriver, subPath); err != "" {
- c.JSON(http.StatusOK, response.Err(err))
+ response.AbortBadRequest(c, err)
return
}
@@ -189,31 +186,30 @@ func UploadFile(c *gin.Context) {
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/uploads/download/{id} [get]
func DownloadFile(c *gin.Context) {
- upload, err := getUploadRecordByID(c)
+ upload, err := filesrv.GetUploadRecordByID(c)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.AbortWithStatus(http.StatusNotFound)
return
}
if _, ok := err.(*strconv.NumError); ok {
- c.JSON(http.StatusOK, response.Err(ErrInvalidFileID))
+ response.AbortBadRequest(c, shared.ErrInvalidFileID)
return
}
- c.JSON(http.StatusOK, response.Err(ErrQueryUploadRecordFailed))
+ response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed)
return
}
- // 校验文件访问权限
- if err := checkFileAccessPermission(c, upload); err != nil {
- c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil})
+ if err := filesrv.CheckFileAccessPermission(c, upload); err != nil {
+ response.AbortUnauthorized(c, common.UnAuthorized)
return
}
fileName := upload.FileName
- quality := normalizeImageQuality(c.Query("quality"))
- isImage := strings.HasPrefix(strings.ToLower(upload.MimeType), "image/") || isImageExtension(strings.ToLower(upload.Extension))
+ quality := util.NormalizeImageQuality(c.Query("quality"))
+ isImage := strings.HasPrefix(strings.ToLower(upload.MimeType), "image/") || util.IsImageExtension(strings.ToLower(upload.Extension))
- if quality != imageQualityOrigin && isImage {
+ if quality != shared.ImageQualityOrigin && isImage {
ext := filepath.Ext(fileName)
if ext != "" {
fileName = strings.TrimSuffix(fileName, ext) + ".webp"
@@ -222,9 +218,8 @@ func DownloadFile(c *gin.Context) {
}
}
- // 设置下载 Attachment 响应头 (支持 UTF-8 中文文件名转义)
c.Header("Content-Disposition", fmt.Sprintf("attachment; filename*=UTF-8''%s", url.PathEscape(fileName)))
- ServeUpload(c, upload)
+ filesrv.ServeUpload(c, upload)
}
// BatchDownloadFiles 批量打包 ZIP 下载接口
@@ -233,7 +228,7 @@ func DownloadFile(c *gin.Context) {
// @Tags admin
// @Accept json
// @Produce octet-stream
-// @Param request body upload.batchDownloadRequest true "包含文件 ID 数组 of string 的请求体"
+// @Param request body handler.batchDownloadRequest true "包含文件 ID 数组 of string 的请求体"
// @Security SessionCookie
// @Success 200 {file} file "成功下载打包后的 ZIP"
// @Failure 400 {object} response.Any "参数错误"
@@ -244,52 +239,45 @@ func BatchDownloadFiles(c *gin.Context) {
var req batchDownloadRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusOK, response.Err(ErrInvalidBatchDownloadRequest))
+ response.AbortBadRequest(c, shared.ErrInvalidBatchDownloadRequest)
return
}
- // 转换 ID 列表
var ids []uint64
for _, idStr := range req.IDs {
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
- c.JSON(http.StatusOK, response.Err(fmt.Sprintf(ErrInvalidIDValueFormat, idStr)))
+ response.AbortBadRequest(c, fmt.Sprintf(shared.ErrInvalidIDValueFormat, idStr))
return
}
ids = append(ids, id)
}
- // 查库获取所有匹配且正常的文件记录
var uploads []model.Upload
if err := db.DB(ctx).Where("id IN ? AND status IN (?, ?)", ids, model.UploadStatusPending, model.UploadStatusUsed).Find(&uploads).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(ErrRetrieveUploadRecordsFailed))
+ response.AbortBadRequest(c, shared.ErrRetrieveUploadRecordsFailed)
return
}
if len(uploads) == 0 {
- c.JSON(http.StatusOK, response.Err(ErrNoValidFilesForArchive))
+ response.AbortBadRequest(c, shared.ErrNoValidFilesForArchive)
return
}
- // 设置 ZIP 格式流的响应头
c.Header("Content-Type", "application/zip")
c.Header("Content-Disposition", "attachment; filename=\"batch_download.zip\"")
- // 开启实时 ZIP 压缩器并直接输出给 Response Writer
zipWriter := zip.NewWriter(c.Writer)
defer func() { _ = zipWriter.Close() }()
- // 用于解决 ZIP 内部文件名称发生碰撞冲突的问题
usedNames := make(map[string]int)
for _, upload := range uploads {
- // 校验文件访问权限
- if err := checkFileAccessPermission(c, &upload); err != nil {
+ if err := filesrv.CheckFileAccessPermission(c, &upload); err != nil {
logger.WarnF(ctx, "Batch download: skip file %d due to permission denied: %v", upload.ID, err)
continue
}
- // 校验防冲突重命名逻辑
fileName := upload.FileName
if count, exists := usedNames[fileName]; exists {
usedNames[fileName] = count + 1
@@ -300,23 +288,19 @@ func BatchDownloadFiles(c *gin.Context) {
usedNames[fileName] = 1
}
- // 在 ZIP 包内建新条目
zipFileEntry, err := zipWriter.Create(fileName)
if err != nil {
logger.ErrorF(ctx, "ZIP 添加条目失败 [%s]: %v", fileName, err)
continue
}
- // 打开底层文件数据源
- var rc io.ReadCloser
- obj, err := openStoredObject(ctx, &upload)
+ obj, err := uploadstorage.OpenStoredObject(ctx, &upload)
if err != nil {
logger.ErrorF(ctx, "打包时读取文件失败: %v", err)
continue
}
- rc = obj.Body
+ rc := obj.Body
- // 流式拷贝到 ZIP entry
_, err = io.Copy(zipFileEntry, rc)
_ = rc.Close()
if err != nil {
@@ -328,7 +312,7 @@ func BatchDownloadFiles(c *gin.Context) {
func resolveUploadAccessMode(c *gin.Context, uploadType string) (int, string) {
accessModeStr := c.PostForm("access_mode")
if accessModeStr == "" {
- if uploadType == defaultPublicUploadType {
+ if uploadType == shared.DefaultPublicUploadType {
return 1, ""
}
return 0, ""
@@ -341,7 +325,6 @@ func resolveUploadAccessMode(c *gin.Context, uploadType string) (int, string) {
return accessMode, ""
}
-// validateUploadExtension 校验文件后缀是否在系统允许的上传扩展名列表中
func validateUploadExtension(ctx context.Context, ext string) string {
var sc model.SystemConfig
if err := sc.GetByKey(ctx, model.ConfigKeyUploadAllowedExtensions); err == nil && sc.Value != "" {
@@ -354,21 +337,20 @@ func validateUploadExtension(ctx context.Context, ext string) string {
}
}
if !allowed {
- return ErrUnsupportedFormat
+ return shared.ErrUnsupportedFormat
}
}
return ""
}
-// tryInstantUpload 尝试秒传:若数据库已存在相同 Hash 且大小一致的可用文件,直接生成新记录
func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User, fileHash string, size int64, mimeType, ext, origName string, accessMode int) (bool, error) {
var existing model.Upload
err := db.DB(ctx).Where("hash = ? AND file_size = ? AND status IN (?, ?)", fileHash, size, model.UploadStatusPending, model.UploadStatusUsed).First(&existing).Error
if err != nil {
return false, err
}
- if StorageReadOnly(ctx) {
- c.JSON(http.StatusConflict, response.Err(ErrStorageReadOnly))
+ if uploadstorage.ReadOnly(ctx) {
+ response.AbortConflict(c, shared.ErrStorageReadOnly)
return true, nil
}
@@ -390,52 +372,40 @@ func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User,
}
if err := db.DB(ctx).Create(&newUpload).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(ErrSaveUploadRecordFailed))
+ response.AbortBadRequest(c, shared.ErrSaveUploadRecordFailed)
return true, err
}
- recordUploadStatsAdd(ctx, &newUpload)
+ uploadstats.RecordUploadStatsAdd(ctx, &newUpload)
logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", id, existing.FilePath)
c.JSON(http.StatusOK, response.OK(newUpload))
return true, nil
}
-// storeUploadFile 将文件写入当前活动存储驱动。
func storeUploadFile(ctx context.Context, subPath string, size int64, mimeType string, buf *bytes.Buffer, meta *model.UploadMetadata) (string, string, string) {
- if StorageReadOnly(ctx) {
- return "", "", ErrStorageReadOnly
+ if uploadstorage.ReadOnly(ctx) {
+ return "", "", shared.ErrStorageReadOnly
}
driver, backend, err := storage.Active(ctx)
if err != nil {
logger.ErrorF(ctx, "初始化活动存储失败: %v", err)
- return "", "", ErrSaveFileFailed
+ return "", "", shared.ErrSaveFileFailed
}
result, err := backend.Put(ctx, subPath, bytes.NewReader(buf.Bytes()), size, mimeType)
if err != nil {
logger.ErrorF(ctx, "写入 %s 存储失败: %v", driver, err)
- return "", "", ErrSaveFileFailed
+ return "", "", shared.ErrSaveFileFailed
}
meta.Bucket = result.Bucket
return string(driver), result.Key, ""
}
-// isImageExtension 判断文件扩展名是否属于常见图片格式
-func isImageExtension(ext string) bool {
- for _, imgExt := range []string{"jpg", "jpeg", "png", "webp", "gif"} {
- if ext == imgExt {
- return true
- }
- }
- return false
-}
-
-// parseUploadMetadata 解析上传元数据字段
func parseUploadMetadata(c *gin.Context, mimeType string) (model.UploadMetadata, string) {
var meta model.UploadMetadata
metadataStr := c.DefaultPostForm("metadata", "")
if metadataStr != "" {
if err := json.Unmarshal([]byte(metadataStr), &meta); err != nil {
- return meta, ErrInvalidMetadataJSON
+ return meta, shared.ErrInvalidMetadataJSON
}
}
meta.OriginalMime = mimeType
@@ -444,16 +414,14 @@ func parseUploadMetadata(c *gin.Context, mimeType string) (model.UploadMetadata,
return meta, ""
}
-// detectMimeType 检测文件的 MIME 类型,优先使用 Content-Type 头部信息
func detectMimeType(buf *bytes.Buffer, header *multipart.FileHeader, size int64) string {
- mimeType := http.DetectContentType(buf.Bytes()[:min(detectContentBytes, int(size))])
+ mimeType := http.DetectContentType(buf.Bytes()[:min(shared.DetectContentBytes, int(size))])
if mimeType == "application/octet-stream" && header.Header.Get("Content-Type") != "" {
mimeType = header.Header.Get("Content-Type")
}
return mimeType
}
-// saveUploadRecord 保存上传记录到数据库,失败时清理本地垃圾文件
func saveUploadRecord(ctx context.Context, upload *model.Upload, storageDriver, filePath string) string {
if err := db.DB(ctx).Create(upload).Error; err != nil {
backend, backendErr := storage.ForDriver(ctx, storage.Driver(storageDriver))
@@ -462,8 +430,8 @@ func saveUploadRecord(ctx context.Context, upload *model.Upload, storageDriver,
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr)
}
}
- return ErrSaveUploadRecordFailed
+ return shared.ErrSaveUploadRecordFailed
}
- recordUploadStatsAdd(ctx, upload)
+ uploadstats.RecordUploadStatsAdd(ctx, upload)
return ""
-}
+}
\ No newline at end of file
diff --git a/internal/apps/upload/routers_test.go b/internal/apps/upload/handler/routers_test.go
similarity index 98%
rename from internal/apps/upload/routers_test.go
rename to internal/apps/upload/handler/routers_test.go
index 98a6aeff..1dbf9d4e 100644
--- a/internal/apps/upload/routers_test.go
+++ b/internal/apps/upload/handler/routers_test.go
@@ -2,7 +2,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+package handler
import (
"archive/zip"
@@ -20,6 +20,9 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
+ "github.com/Rain-kl/Wavelet/internal/common/response"
+ uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
@@ -36,6 +39,7 @@ type testResponse struct {
func setupTestRouter(authUser *model.User) *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
+ r.Use(response.ErrorHandlerMiddleware())
authMiddleware := func(c *gin.Context) {
if authUser != nil {
@@ -211,13 +215,13 @@ func TestUploadFile(t *testing.T) {
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
- if w.Code != http.StatusOK {
- t.Fatalf("expected status 200, got %d. Body: %s", w.Code, w.Body.String())
+ if w.Code != http.StatusBadRequest {
+ t.Fatalf("expected status 400, got %d. Body: %s", w.Code, w.Body.String())
}
var resp testResponse
_ = json.Unmarshal(w.Body.Bytes(), &resp)
- if resp.ErrorMsg == "" || !strings.Contains(resp.ErrorMsg, ErrUnsupportedFormat) {
+ if resp.ErrorMsg == "" || !strings.Contains(resp.ErrorMsg, shared.ErrUnsupportedFormat) {
t.Errorf("expected unsupported format error, got: %v", resp)
}
})
@@ -294,6 +298,7 @@ func TestUploadFile(t *testing.T) {
sc.Value = "jpg,png,webp,txt"
dbConn.Save(&sc)
_ = db.HSetJSON(context.Background(), model.SystemConfigRedisHashKey, sc.Key, &sc)
+ model.ResetSystemConfigRAMCacheForTest()
contentType, body := createMultipartRequest(t, "file", "doc.txt", []byte("hello world generic document file"), map[string]string{
"type": "document",
@@ -806,7 +811,7 @@ func TestGetFileStats(t *testing.T) {
t.Fatalf("failed to create upload: %v", err)
}
}
- if err := RebuildUploadStats(context.Background()); err != nil {
+ if err := uploadstats.RebuildUploadStats(context.Background()); err != nil {
t.Fatalf("failed to rebuild upload stats: %v", err)
}
diff --git a/internal/apps/upload/stats.go b/internal/apps/upload/handler/stats.go
similarity index 66%
rename from internal/apps/upload/stats.go
rename to internal/apps/upload/handler/stats.go
index 46527730..88a0a753 100644
--- a/internal/apps/upload/stats.go
+++ b/internal/apps/upload/handler/stats.go
@@ -1,27 +1,17 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+package handler
import (
"net/http"
- "strings"
"time"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
+ "github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
-
- "github.com/Rain-kl/Wavelet/internal/common/response"
-)
-
-const (
- catImage = "图片"
- catVideo = "视频"
- catAudio = "音频"
- catDocument = "文档"
- catArchive = "压缩包"
- catOther = "其他"
)
type trendItem struct {
@@ -60,15 +50,15 @@ func GetFileStats(c *gin.Context) {
var stats []model.UploadStat
if err := db.DB(ctx).Find(&stats).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
now := time.Now()
- trendDates := make([]string, 0, fileStatsTrendDays)
- trendCountMap := make(map[string]int64, fileStatsTrendDays)
- trendSizeMap := make(map[string]int64, fileStatsTrendDays)
- for i := fileStatsTrendDays - 1; i >= 0; i-- {
+ trendDates := make([]string, 0, shared.FileStatsTrendDays)
+ trendCountMap := make(map[string]int64, shared.FileStatsTrendDays)
+ trendSizeMap := make(map[string]int64, shared.FileStatsTrendDays)
+ for i := shared.FileStatsTrendDays - 1; i >= 0; i-- {
date := now.AddDate(0, 0, -i).Format("2006-01-02")
trendDates = append(trendDates, date)
trendCountMap[date] = 0
@@ -82,7 +72,7 @@ func GetFileStats(c *gin.Context) {
categories []distributionItem
)
- categoriesList := []string{catImage, catVideo, catAudio, catDocument, catArchive, catOther}
+ categoriesList := []string{"图片", "视频", "音频", "文档", "压缩包", "其他"}
categoryMap := make(map[string]distributionItem, len(categoriesList))
for _, cat := range categoriesList {
categoryMap[cat] = distributionItem{Name: cat}
@@ -134,44 +124,4 @@ func GetFileStats(c *gin.Context) {
Categories: categories,
Types: types,
}))
-}
-
-func getFileCategory(mimeType, ext string) string {
- mimeType = strings.ToLower(mimeType)
- ext = strings.ToLower(ext)
-
- if strings.HasPrefix(mimeType, "image/") || isImageExtension(ext) {
- return catImage
- }
- if strings.HasPrefix(mimeType, "video/") {
- return catVideo
- }
- if strings.HasPrefix(mimeType, "audio/") {
- return catAudio
- }
- if isArchiveExtension(ext) || strings.Contains(mimeType, "zip") || strings.Contains(mimeType, "tar") || strings.Contains(mimeType, "gzip") {
- return catArchive
- }
- if isDocumentExtension(ext) || strings.HasPrefix(mimeType, "text/") || mimeType == "application/pdf" {
- return catDocument
- }
- return catOther
-}
-
-func isArchiveExtension(ext string) bool {
- for _, e := range []string{"zip", "rar", "7z", "tar", "gz", "tgz", "bz2", "xz"} {
- if ext == e {
- return true
- }
- }
- return false
-}
-
-func isDocumentExtension(ext string) bool {
- for _, e := range []string{"pdf", "doc", "docx", "xls", "xlsx", "ppt", "pptx", "txt", "md", "csv", "json", "yaml", "yml", "xml"} {
- if ext == e {
- return true
- }
- }
- return false
}
\ No newline at end of file
diff --git a/internal/apps/upload/shared/constants.go b/internal/apps/upload/shared/constants.go
new file mode 100644
index 00000000..405bf08e
--- /dev/null
+++ b/internal/apps/upload/shared/constants.go
@@ -0,0 +1,23 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package shared
+
+import "github.com/Rain-kl/Wavelet/internal/storage"
+
+// Upload size, path, media quality, and cache constants shared across subpackages.
+const (
+ MaxUploadSize = 32 * 1024 * 1024 // 32MB
+ DetectContentBytes = 512 // http.DetectContentType 需要的最小字节数
+ UploadDirPerm = 0755 // 上传目录权限
+ UploadFilePerm = 0644 // 上传文件权限
+ ImageQualityLow = "low"
+ ImageQualityMedium = "medium"
+ ImageQualityHigh = "high"
+ ImageQualityOrigin = "origin"
+ StorageDriverLocal = string(storage.DriverLocal)
+ DefaultPublicUploadType = "avatar"
+ FileStatsTrendDays = 7
+ MaxS3KeyLength = 1024
+ AccessCacheTTL = 5 // seconds; multiplied by time.Second at use site
+)
\ No newline at end of file
diff --git a/internal/apps/upload/shared/errs.go b/internal/apps/upload/shared/errs.go
new file mode 100644
index 00000000..e9a466e6
--- /dev/null
+++ b/internal/apps/upload/shared/errs.go
@@ -0,0 +1,41 @@
+// Copyright 2025 linux.do
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+// Package shared holds upload error and configuration constants shared across subpackages.
+package shared
+
+// 文件管理常量
+const (
+ ErrNoFileSelected = "请选择要上传的文件"
+ ErrUnsupportedFormat = "只支持 JPG、PNG、WEBP 格式的图片"
+ ErrProcessFileFailed = "处理文件失败"
+ ErrSaveFileFailed = "保存文件失败"
+ ErrOpenFileFailed = "打开文件失败"
+ ErrSaveUploadRecordFailed = "保存上传记录失败"
+ ErrGenericFileTooLarge = "文件大小不能超过 32MB"
+ ErrFileContentExtensionMismatch = "文件内容与扩展名不匹配,可能包含安全风险"
+ ErrFileValidationFailed = "文件校验失败"
+ ErrInvalidMetadataJSON = "元数据 JSON 格式不合法"
+ ErrInvalidFileID = "无效的文件 ID"
+ ErrQueryUploadRecordFailed = "查询文件记录失败"
+ ErrInvalidBatchDownloadRequest = "参数绑定失败,请传入有效的文件 ID 数组"
+ ErrInvalidIDValueFormat = "无效的 ID 值: %s"
+ ErrRetrieveUploadRecordsFailed = "检索文件记录失败"
+ ErrNoValidFilesForArchive = "没有找到任何有效的文件记录进行打包"
+ ErrInvalidParams = "参数错误"
+ ErrQueryFileCountFailed = "查询文件数量失败"
+ ErrQueryFileListFailed = "查询文件列表失败"
+ ErrDeleteFileFailed = "删除文件失败"
+ ErrStorageReadOnly = "存储迁移维护中,当前仅允许读取文件"
+ ErrS3KeyRequired = "s3 key must not be empty"
+ ErrS3KeyTooLongFormat = "s3 key exceeds maximum length of %d"
+ ErrS3KeyStartsWithSlash = "s3 key must not start with /"
+ ErrS3KeyContainsNullBytes = "s3 key must not contain null bytes"
+ ErrQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w"
+ ErrImageCacheWarmupPayloadRequired = "图片缓存预热参数不能为空"
+ ErrInvalidImageCacheWarmupPayload = "图片缓存预热参数格式无效: %w"
+ ErrInvalidImageCacheWarmupQuality = "图片质量仅支持 low、medium、high"
+ ErrParseImageCacheWarmupPayload = "解析图片缓存预热参数失败: %w"
+ ErrQueryImagesForCacheWarmup = "查询待预热图片失败: %w"
+)
\ No newline at end of file
diff --git a/internal/apps/upload/stats/category.go b/internal/apps/upload/stats/category.go
new file mode 100644
index 00000000..5ed059e6
--- /dev/null
+++ b/internal/apps/upload/stats/category.go
@@ -0,0 +1,43 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+// Package stats maintains incremental upload statistics and aggregations.
+package stats
+
+import (
+ "strings"
+
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/util"
+)
+
+const (
+ catImage = "图片"
+ catVideo = "视频"
+ catAudio = "音频"
+ catDocument = "文档"
+ catArchive = "压缩包"
+ catOther = "其他"
+)
+
+// GetFileCategory classifies a file by mime type and extension.
+func GetFileCategory(mimeType, ext string) string {
+ mimeType = strings.ToLower(mimeType)
+ ext = strings.ToLower(ext)
+
+ if strings.HasPrefix(mimeType, "image/") || util.IsImageExtension(ext) {
+ return catImage
+ }
+ if strings.HasPrefix(mimeType, "video/") {
+ return catVideo
+ }
+ if strings.HasPrefix(mimeType, "audio/") {
+ return catAudio
+ }
+ if util.IsArchiveExtension(ext) || strings.Contains(mimeType, "zip") || strings.Contains(mimeType, "tar") || strings.Contains(mimeType, "gzip") {
+ return catArchive
+ }
+ if util.IsDocumentExtension(ext) || strings.HasPrefix(mimeType, "text/") || mimeType == "application/pdf" {
+ return catDocument
+ }
+ return catOther
+}
\ No newline at end of file
diff --git a/internal/apps/upload/stats_counter.go b/internal/apps/upload/stats/stats_counter.go
similarity index 91%
rename from internal/apps/upload/stats_counter.go
rename to internal/apps/upload/stats/stats_counter.go
index 3f5c28d0..eb2170d2 100644
--- a/internal/apps/upload/stats_counter.go
+++ b/internal/apps/upload/stats/stats_counter.go
@@ -1,7 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+package stats
import (
"context"
@@ -72,7 +72,7 @@ func applyUploadStatsDeltaTx(tx *gorm.DB, upload *model.Upload, sign int64) erro
}{
{model.UploadStatDimensionTotal, ""},
{model.UploadStatDimensionType, typeKey},
- {model.UploadStatDimensionCategory, getFileCategory(upload.MimeType, upload.Extension)},
+ {model.UploadStatDimensionCategory, GetFileCategory(upload.MimeType, upload.Extension)},
{model.UploadStatDimensionTrend, upload.CreatedAt.Format("2006-01-02")},
}
@@ -111,13 +111,15 @@ func upsertUploadStatDelta(tx *gorm.DB, dimension, key string, countDelta, sizeD
}).Error
}
-func recordUploadStatsAdd(ctx context.Context, upload *model.Upload) {
+// RecordUploadStatsAdd logs and applies upload stats increment.
+func RecordUploadStatsAdd(ctx context.Context, upload *model.Upload) {
if err := ApplyUploadStatsAdd(ctx, upload); err != nil {
logger.WarnF(ctx, "increment upload stats failed: %v", err)
}
}
-func recordUploadStatsRemove(ctx context.Context, upload *model.Upload) {
+// RecordUploadStatsRemove logs and applies upload stats decrement.
+func RecordUploadStatsRemove(ctx context.Context, upload *model.Upload) {
if err := ApplyUploadStatsRemove(ctx, upload); err != nil {
logger.WarnF(ctx, "decrement upload stats failed: %v", err)
}
diff --git a/internal/apps/upload/stats_counter_test.go b/internal/apps/upload/stats/stats_counter_test.go
similarity index 99%
rename from internal/apps/upload/stats_counter_test.go
rename to internal/apps/upload/stats/stats_counter_test.go
index dad330a6..597c929a 100644
--- a/internal/apps/upload/stats_counter_test.go
+++ b/internal/apps/upload/stats/stats_counter_test.go
@@ -1,7 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+package stats
import (
"context"
diff --git a/internal/apps/upload/storage/access_state.go b/internal/apps/upload/storage/access_state.go
new file mode 100644
index 00000000..2f8054c0
--- /dev/null
+++ b/internal/apps/upload/storage/access_state.go
@@ -0,0 +1,88 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+// Package storage provides upload storage backend operations and migration state.
+package storage
+
+import (
+ "context"
+ "sync"
+ "time"
+
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/Rain-kl/Wavelet/internal/storage"
+)
+
+// MigrationAccessState captures cached migration maintenance state.
+type MigrationAccessState struct {
+ ReadOnly bool
+ Target storage.Config
+ HasTarget bool
+ TargetErr error
+ LoadErr error
+}
+
+var (
+ migrationAccessMu sync.RWMutex
+ migrationAccessCached MigrationAccessState
+ migrationAccessValid bool
+ migrationAccessCheckedAt time.Time
+)
+
+// ResetMigrationAccessCache clears the in-process migration access cache.
+func ResetMigrationAccessCache() {
+ migrationAccessMu.Lock()
+ migrationAccessValid = false
+ migrationAccessMu.Unlock()
+}
+
+// LoadMigrationAccessState returns cached migration maintenance state.
+func LoadMigrationAccessState(ctx context.Context) MigrationAccessState {
+ migrationAccessMu.RLock()
+ if migrationAccessValid && time.Since(migrationAccessCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
+ state := migrationAccessCached
+ migrationAccessMu.RUnlock()
+ return state
+ }
+ migrationAccessMu.RUnlock()
+
+ migrationAccessMu.Lock()
+ defer migrationAccessMu.Unlock()
+
+ if migrationAccessValid && time.Since(migrationAccessCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
+ return migrationAccessCached
+ }
+
+ migrationAccessCached = buildMigrationAccessState(ctx)
+ migrationAccessValid = true
+ migrationAccessCheckedAt = time.Now()
+ return migrationAccessCached
+}
+
+func buildMigrationAccessState(ctx context.Context) MigrationAccessState {
+ execution, ok, err := LatestMigrationExecution(ctx)
+ if err != nil {
+ return MigrationAccessState{LoadErr: err, ReadOnly: true}
+ }
+ if !ok {
+ return MigrationAccessState{}
+ }
+
+ state := MigrationAccessState{
+ ReadOnly: execution.Status != model.TaskExecutionStatusSucceeded,
+ }
+ if execution.Status == model.TaskExecutionStatusSucceeded {
+ return state
+ }
+
+ target, err := ParseMigrationTargetConfig(ctx, []byte(execution.Payload))
+ if err != nil {
+ state.TargetErr = err
+ return state
+ }
+
+ state.Target = target
+ state.HasTarget = true
+ return state
+}
\ No newline at end of file
diff --git a/internal/apps/upload/storage/migration.go b/internal/apps/upload/storage/migration.go
new file mode 100644
index 00000000..1f8536e0
--- /dev/null
+++ b/internal/apps/upload/storage/migration.go
@@ -0,0 +1,93 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package storage
+
+import (
+ "context"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "strings"
+
+ "github.com/Rain-kl/Wavelet/internal/db"
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/Rain-kl/Wavelet/internal/storage"
+ "gorm.io/gorm"
+)
+
+// StorageMigrationTask is the Asynq task name for storage migration.
+const StorageMigrationTask = "storage:migrate"
+
+// LatestMigrationExecution returns the most recent storage migration task execution.
+func LatestMigrationExecution(ctx context.Context) (*model.TaskExecution, bool, error) {
+ var execution model.TaskExecution
+ err := db.DB(ctx).
+ Where("task_type = ?", StorageMigrationTask).
+ Order("id DESC").
+ First(&execution).Error
+ if err == nil {
+ return &execution, true, nil
+ }
+ if !errors.Is(err, gorm.ErrRecordNotFound) {
+ return nil, false, err
+ }
+ return nil, false, nil
+}
+
+// ParseMigrationTargetConfig parses and validates a storage migration target payload.
+func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (storage.Config, error) {
+ if strings.TrimSpace(string(payload)) == "" {
+ return storage.Config{}, errors.New("storage migration target payload is required")
+ }
+
+ var raw struct {
+ Target json.RawMessage `json:"target"`
+ }
+ if err := json.Unmarshal(payload, &raw); err != nil {
+ return storage.Config{}, fmt.Errorf("parse storage migration payload envelope: %w", err)
+ }
+
+ if len(raw.Target) == 0 {
+ return storage.Config{}, errors.New("storage migration target payload is required")
+ }
+
+ var targetBytes []byte
+ var targetStr string
+ if err := json.Unmarshal(raw.Target, &targetStr); err == nil {
+ targetBytes = []byte(targetStr)
+ } else {
+ targetBytes = raw.Target
+ }
+
+ var target storage.Config
+ if err := json.Unmarshal(targetBytes, &target); err != nil {
+ return storage.Config{}, fmt.Errorf("parse target storage config: %w", err)
+ }
+
+ current, err := storage.LoadConfig(ctx)
+ if err != nil {
+ return storage.Config{}, fmt.Errorf("load active storage config: %w", err)
+ }
+ target = storage.MergeMaskedSecrets(target, current)
+ if err := storage.ValidateConfig(target); err != nil {
+ return storage.Config{}, fmt.Errorf("validate target storage config: %w", err)
+ }
+ return target, nil
+}
+
+// NormalizeMigrationPayload validates and normalizes a storage migration payload.
+func NormalizeMigrationPayload(ctx context.Context, payload []byte) ([]byte, storage.Config, error) {
+ target, err := ParseMigrationTargetConfig(ctx, payload)
+ if err != nil {
+ return nil, storage.Config{}, err
+ }
+ type storageMigrationPayload struct {
+ Target storage.Config `json:"target"`
+ }
+ normalized, err := json.Marshal(storageMigrationPayload{Target: target})
+ if err != nil {
+ return nil, storage.Config{}, fmt.Errorf("marshal storage migration payload: %w", err)
+ }
+ return normalized, target, nil
+}
\ No newline at end of file
diff --git a/internal/apps/upload/storage_ops.go b/internal/apps/upload/storage/storage_ops.go
similarity index 55%
rename from internal/apps/upload/storage_ops.go
rename to internal/apps/upload/storage/storage_ops.go
index 369f9f43..8e4e36e7 100644
--- a/internal/apps/upload/storage_ops.go
+++ b/internal/apps/upload/storage/storage_ops.go
@@ -1,7 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+package storage
import (
"context"
@@ -12,17 +12,18 @@ import (
"github.com/Rain-kl/Wavelet/pkg/logger"
)
-// StorageReadOnly checks if the storage system is in read-only maintenance mode.
-func StorageReadOnly(ctx context.Context) bool {
- state := loadMigrationAccessState(ctx)
- if state.loadErr != nil {
- logger.ErrorF(ctx, "读取存储维护状态失败: %v", state.loadErr)
+// ReadOnly checks if the storage system is in read-only maintenance mode.
+func ReadOnly(ctx context.Context) bool {
+ state := LoadMigrationAccessState(ctx)
+ if state.LoadErr != nil {
+ logger.ErrorF(ctx, "读取存储维护状态失败: %v", state.LoadErr)
return true
}
- return state.readOnly
+ return state.ReadOnly
}
-func openStoredObject(ctx context.Context, upload *model.Upload) (*storage.Object, error) {
+// OpenStoredObject opens a stored upload object from its configured backend.
+func OpenStoredObject(ctx context.Context, upload *model.Upload) (*storage.Object, error) {
driver := storage.Driver(upload.StorageDriver)
if driver == "" {
driver = storage.DriverLocal
@@ -40,7 +41,7 @@ func backendForStoredDriver(ctx context.Context, driver storage.Driver) (storage
return backend, nil
}
- target, ok, targetErr := currentMigrationTargetConfig(ctx)
+ target, ok, targetErr := CurrentMigrationTargetConfig(ctx)
if targetErr != nil {
return nil, targetErr
}
@@ -50,16 +51,17 @@ func backendForStoredDriver(ctx context.Context, driver storage.Driver) (storage
return nil, fmt.Errorf("storage configuration for driver %q is unavailable", driver)
}
-func currentMigrationTargetConfig(ctx context.Context) (storage.Config, bool, error) {
- state := loadMigrationAccessState(ctx)
- if state.loadErr != nil {
- return storage.Config{}, false, state.loadErr
+// CurrentMigrationTargetConfig returns the pending migration target config when available.
+func CurrentMigrationTargetConfig(ctx context.Context) (storage.Config, bool, error) {
+ state := LoadMigrationAccessState(ctx)
+ if state.LoadErr != nil {
+ return storage.Config{}, false, state.LoadErr
}
- if state.targetErr != nil {
- return storage.Config{}, false, state.targetErr
+ if state.TargetErr != nil {
+ return storage.Config{}, false, state.TargetErr
}
- if !state.hasTarget {
+ if !state.HasTarget {
return storage.Config{}, false, nil
}
- return state.target, true, nil
-}
+ return state.Target, true, nil
+}
\ No newline at end of file
diff --git a/internal/apps/upload/cleanup.go b/internal/apps/upload/task/cleanup.go
similarity index 78%
rename from internal/apps/upload/cleanup.go
rename to internal/apps/upload/task/cleanup.go
index adfa80c1..82ac2fc1 100644
--- a/internal/apps/upload/cleanup.go
+++ b/internal/apps/upload/task/cleanup.go
@@ -1,8 +1,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-// Package upload implements upload tasks and file cleanup services.
-package upload
+// Package task provides upload-related async background task handlers.
+package task
import (
"context"
@@ -10,6 +10,9 @@ import (
"fmt"
"time"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
+ uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
+ uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
@@ -18,16 +21,11 @@ import (
"gorm.io/gorm"
)
-// 异步任务名称与管理类型定义
const (
// SystemCleanupTask 系统定期垃圾清理任务标识
SystemCleanupTask = "system:cleanup"
// TaskTypeSystemCleanup 系统定期垃圾清理管理类型
TaskTypeSystemCleanup = "system_cleanup"
-
- // 错误描述常量
- errStorageReadOnly = "存储迁移维护中,当前仅允许读取文件"
- errQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w"
)
// SystemCleanupMeta represents the task metadata.
@@ -47,21 +45,19 @@ type SystemCleanupHandler struct{}
// Execute 执行系统清理(包含文件清理、历史推送日志和任务执行日志清理)
func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
- if storageReadOnly(ctx) {
- return nil, errors.New(errStorageReadOnly)
+ if uploadstorage.ReadOnly(ctx) {
+ return nil, errors.New(shared.ErrStorageReadOnly)
}
- const batchSize = 100 // 每批处理100个文件
+ const batchSize = 100
var lastID uint64
var totalProcessed int
var totalDeleted int
- // 计算1小时前的时间
oneHourAgo := time.Now().Add(-1 * time.Hour)
task.AppendLog(ctx, "开始扫描未使用上传文件,阈值: %s", oneHourAgo.Format(time.RFC3339))
for {
- // 使用游标分页查询未使用且超过1小时的上传记录
var unusedUploads []model.Upload
if err := db.DB(ctx).
Where("id > ? AND status = ? AND created_at < ?", lastID, model.UploadStatusPending, oneHourAgo).
@@ -69,22 +65,19 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
Limit(batchSize).
Find(&unusedUploads).Error; err != nil {
task.AppendLog(ctx, "查询未使用的上传文件失败: %v", err)
- return nil, fmt.Errorf(errQueryUnusedUploadsFailed, err)
+ return nil, fmt.Errorf(shared.ErrQueryUnusedUploadsFailed, err)
}
- // 没有更多数据,退出循环
if len(unusedUploads) == 0 {
break
}
task.AppendLog(ctx, "本批次找到 %d 个需要清理的上传文件", len(unusedUploads))
- // 处理每个未使用的上传文件
for _, u := range unusedUploads {
totalProcessed++
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
- // 更新上传记录状态
if err := tx.Model(&model.Upload{}).
Where("id = ? AND status = ?", u.ID, model.UploadStatusPending).
Update("status", model.UploadStatusDeleted).Error; err != nil {
@@ -110,13 +103,12 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
continue
}
- recordUploadStatsRemove(ctx, &u)
+ uploadstats.RecordUploadStatsRemove(ctx, &u)
totalDeleted++
lastID = u.ID
}
}
- // 2. 清理超过7天的历史推送日志
task.AppendLog(ctx, "开始清理历史推送审计日志,只保留最近7天数据...")
cutoff := time.Now().AddDate(0, 0, -7)
var pushHistoryCount int64
@@ -132,7 +124,6 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
task.AppendLog(ctx, "没有需要清理的历史推送记录 (截止时间: %s)", cutoff.Format("2006-01-02 15:04:05"))
}
- // 3. 清理任务执行日志:高频任务保留3天,低频任务保留30天。
task.AppendLog(ctx, "开始清理任务执行日志:高频任务保留最近3天,低频任务保留最近30天...")
taskLogStats, err := model.CleanupTaskExecutionLogs(ctx, time.Now())
if err != nil {
@@ -154,17 +145,4 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
)
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
-}
-
-func storageReadOnly(ctx context.Context) bool {
- var execution model.TaskExecution
- err := db.DB(ctx).Where("task_type = ?", "storage:migrate").Order("id DESC").First(&execution).Error
- if err != nil {
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return false
- }
- logger.ErrorF(ctx, "读取存储维护状态失败: %v", err)
- return true
- }
- return execution.Status != model.TaskExecutionStatusSucceeded
-}
+}
\ No newline at end of file
diff --git a/internal/apps/upload/storage_migration_task.go b/internal/apps/upload/task/storage_migration.go
similarity index 79%
rename from internal/apps/upload/storage_migration_task.go
rename to internal/apps/upload/task/storage_migration.go
index 5b089d74..1a54f8c3 100644
--- a/internal/apps/upload/storage_migration_task.go
+++ b/internal/apps/upload/task/storage_migration.go
@@ -1,13 +1,12 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+package task
import (
"context"
"crypto/sha256"
"encoding/hex"
- "encoding/json"
"errors"
"fmt"
"io"
@@ -16,17 +15,18 @@ import (
"sync/atomic"
"time"
+ uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
+ uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/task"
"golang.org/x/sync/errgroup"
- "gorm.io/gorm"
)
const (
// StorageMigrationTask is the Asynq task name for storage migration.
- StorageMigrationTask = "storage:migrate"
+ StorageMigrationTask = uploadstorage.StorageMigrationTask
// TaskTypeStorageMigration is the task metadata type for storage migration.
TaskTypeStorageMigration = "storage_migration"
@@ -58,13 +58,9 @@ var StorageMigrationMeta = task.TaskMeta{
// MigrationHandler copies stored objects and activates the target backend.
type MigrationHandler struct{}
-type storageMigrationPayload struct {
- Target storage.Config `json:"target"`
-}
-
// ValidatePayload rejects duplicate active migrations through the task framework.
func (h *MigrationHandler) ValidatePayload(payload []byte) ([]byte, error) {
- normalized, _, err := normalizeStorageMigrationPayload(context.Background(), payload)
+ normalized, _, err := uploadstorage.NormalizeMigrationPayload(context.Background(), payload)
if err != nil {
return payload, err
}
@@ -95,7 +91,6 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
return nil, errors.New("另一个存储迁移任务正在运行中")
}
- // 任务结束时清理锁,使用 Background context 避免受任务 context 取消的影响
stopRenewal := make(chan struct{})
//nolint:contextcheck
defer func() {
@@ -105,7 +100,6 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
_ = db.Redis.Del(cleanupCtx, lockKey)
}()
- // 启动看门狗续租协程,每 10 分钟将锁的 TTL 自动延长为 1 小时
//nolint:contextcheck,gosec
go func() {
ticker := time.NewTicker(renewalInterval)
@@ -129,7 +123,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
if err != nil {
return nil, fmt.Errorf("load active storage config: %w", err)
}
- target, err := parseMigrationTargetConfig(ctx, payload)
+ target, err := uploadstorage.ParseMigrationTargetConfig(ctx, payload)
if err != nil {
return nil, err
}
@@ -178,62 +172,6 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
return &task.TaskResult{Message: message}, nil
}
-func normalizeStorageMigrationPayload(ctx context.Context, payload []byte) ([]byte, storage.Config, error) {
- target, err := parseMigrationTargetConfig(ctx, payload)
- if err != nil {
- return nil, storage.Config{}, err
- }
- normalized, err := json.Marshal(storageMigrationPayload{Target: target})
- if err != nil {
- return nil, storage.Config{}, fmt.Errorf("marshal storage migration payload: %w", err)
- }
- return normalized, target, nil
-}
-
-func parseMigrationTargetConfig(ctx context.Context, payload []byte) (storage.Config, error) {
- if strings.TrimSpace(string(payload)) == "" {
- return storage.Config{}, errors.New("storage migration target payload is required")
- }
-
- // Try to parse using raw JSON message to handle both struct and string payload formats
- var raw struct {
- Target json.RawMessage `json:"target"`
- }
- if err := json.Unmarshal(payload, &raw); err != nil {
- return storage.Config{}, fmt.Errorf("parse storage migration payload envelope: %w", err)
- }
-
- if len(raw.Target) == 0 {
- return storage.Config{}, errors.New("storage migration target payload is required")
- }
-
- var targetBytes []byte
- var targetStr string
- // Check if Target is a JSON string
- if err := json.Unmarshal(raw.Target, &targetStr); err == nil {
- // It is a string (e.g. from dynamic form input), parse its content as JSON
- targetBytes = []byte(targetStr)
- } else {
- // It is a JSON object, use directly
- targetBytes = raw.Target
- }
-
- var target storage.Config
- if err := json.Unmarshal(targetBytes, &target); err != nil {
- return storage.Config{}, fmt.Errorf("parse target storage config: %w", err)
- }
-
- current, err := storage.LoadConfig(ctx)
- if err != nil {
- return storage.Config{}, fmt.Errorf("load active storage config: %w", err)
- }
- target = storage.MergeMaskedSecrets(target, current)
- if err := storage.ValidateConfig(target); err != nil {
- return storage.Config{}, fmt.Errorf("validate target storage config: %w", err)
- }
- return target, nil
-}
-
func countStorageObjects(ctx context.Context, driver storage.Driver) (int64, error) {
var count int64
err := db.DB(ctx).Model(&model.Upload{}).
@@ -244,28 +182,13 @@ func countStorageObjects(ctx context.Context, driver storage.Driver) (int64, err
}
func hasUnresolvedMigrationTask(ctx context.Context) (bool, error) {
- execution, ok, err := latestStorageMigrationExecution(ctx)
+ execution, ok, err := uploadstorage.LatestMigrationExecution(ctx)
if err != nil || !ok {
return false, err
}
return execution.Status == model.TaskExecutionStatusPending || execution.Status == model.TaskExecutionStatusRunning, nil
}
-func latestStorageMigrationExecution(ctx context.Context) (*model.TaskExecution, bool, error) {
- var execution model.TaskExecution
- err := db.DB(ctx).
- Where("task_type = ?", StorageMigrationTask).
- Order("id DESC").
- First(&execution).Error
- if err == nil {
- return &execution, true, nil
- }
- if !errors.Is(err, gorm.ErrRecordNotFound) {
- return nil, false, err
- }
- return nil, false, nil
-}
-
func migrateObjects(
ctx context.Context,
sourceBackend storage.Backend,
@@ -311,7 +234,7 @@ func migrateObjects(
g.SetLimit(migrationConcurrency)
for _, object := range objects {
- obj := object // Capture range variable
+ obj := object
g.Go(func() error {
if err := migrateSingleObject(ctx, sourceBackend, targetBackend, sourceDriver, targetDriver, obj, sha256HexLength); err != nil {
return err
@@ -344,7 +267,6 @@ func migrateSingleObject(
},
sha256HexLength int,
) error {
- // Check if the file already exists in target storage and has matching size
if shouldSkipMigration(ctx, targetBackend, obj) {
task.AppendLog(ctx, "[跳过迁移] 目标存储已存在相同文件且校验一致: %s", obj.FilePath)
if err := db.DB(ctx).Model(&model.Upload{}).
@@ -375,7 +297,6 @@ func migrateSingleObject(
return fmt.Errorf("close source object %q: %w", obj.FilePath, closeErr)
}
- // Data integrity check (SHA-256 hash verification)
if len(obj.Hash) == sha256HexLength {
task.AppendLog(ctx, "[校验中] 正在对目标文件进行数据一致性校验 (SHA-256): %s", targetResult.Key)
targetObj, getErr := targetBackend.Get(ctx, targetResult.Key)
@@ -450,13 +371,13 @@ func markMissingMigrationObjectDeleted(
if err := db.DB(ctx).Model(&model.Upload{}).
Where("storage_driver = ? AND file_path = ?", sourceDriver, filePath).
Updates(map[string]any{
- "status": model.UploadStatusDeleted,
- colStorageDriver: targetDriver,
+ "status": model.UploadStatusDeleted,
+ colStorageDriver: targetDriver,
}).Error; err != nil {
return fmt.Errorf("update missing object %q: %w", filePath, err)
}
for i := range affectedUploads {
- recordUploadStatsRemove(ctx, &affectedUploads[i])
+ uploadstats.RecordUploadStatsRemove(ctx, &affectedUploads[i])
}
return nil
}
@@ -475,4 +396,4 @@ func isNotFoundError(err error) bool {
}
}
return false
-}
+}
\ No newline at end of file
diff --git a/internal/apps/upload/storage_migration_task_test.go b/internal/apps/upload/task/storage_migration_task_test.go
similarity index 96%
rename from internal/apps/upload/storage_migration_task_test.go
rename to internal/apps/upload/task/storage_migration_task_test.go
index 86d4177f..f7216971 100644
--- a/internal/apps/upload/storage_migration_task_test.go
+++ b/internal/apps/upload/task/storage_migration_task_test.go
@@ -1,7 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+package task
import (
"bytes"
@@ -52,7 +52,9 @@ func TestMigrationHandlerExecute(t *testing.T) {
AccessKeyID: "key",
SecretAccessKey: "secret",
}
- payload, err := json.Marshal(storageMigrationPayload{Target: target})
+ payload, err := json.Marshal(struct {
+ Target storage.Config `json:"target"`
+ }{Target: target})
if err != nil {
t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err)
}
@@ -150,7 +152,9 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) {
AccessKeyID: "key",
SecretAccessKey: "secret",
}
- payload, err := json.Marshal(storageMigrationPayload{Target: target})
+ payload, err := json.Marshal(struct {
+ Target storage.Config `json:"target"`
+ }{Target: target})
if err != nil {
t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err)
}
@@ -259,7 +263,9 @@ func TestMigrationHandlerExecuteWithRedisLock(t *testing.T) {
t.Fatalf("SaveActiveConfig() returned error: %v", err)
}
- payload, err := json.Marshal(storageMigrationPayload{Target: active})
+ payload, err := json.Marshal(struct {
+ Target storage.Config `json:"target"`
+ }{Target: active})
if err != nil {
t.Fatalf("Marshal payload failed: %v", err)
}
diff --git a/internal/apps/upload/tasks.go b/internal/apps/upload/task/tasks.go
similarity index 85%
rename from internal/apps/upload/tasks.go
rename to internal/apps/upload/task/tasks.go
index f203a7cd..bcfec085 100644
--- a/internal/apps/upload/tasks.go
+++ b/internal/apps/upload/task/tasks.go
@@ -2,7 +2,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+package task
import (
"context"
@@ -12,12 +12,13 @@ import (
"strings"
"sync"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
)
-// 异步任务名称与管理类型定义
const (
// WarmImageCacheTask 图片压缩缓存预热任务标识
WarmImageCacheTask = "upload:warm_image_cache"
@@ -54,29 +55,25 @@ type WarmImageCachePayload struct {
Quality string `json:"quality"`
}
-// SystemCleanupHandler 系统定期垃圾清理异步任务处理器
-
// WarmImageCacheHandler serially warms compressed image cache entries.
type WarmImageCacheHandler struct{}
-// Execute 执行系统清理(包含文件清理和历史消息推送日志清理)
-
// ValidatePayload validates and normalizes image cache warmup parameters.
func (h *WarmImageCacheHandler) ValidatePayload(payload []byte) ([]byte, error) {
if len(payload) == 0 {
- return nil, errors.New(errImageCacheWarmupPayloadRequired)
+ return nil, errors.New(shared.ErrImageCacheWarmupPayloadRequired)
}
var req WarmImageCachePayload
if err := json.Unmarshal(payload, &req); err != nil {
- return nil, fmt.Errorf(errInvalidImageCacheWarmupPayload, err)
+ return nil, fmt.Errorf(shared.ErrInvalidImageCacheWarmupPayload, err)
}
req.Quality = strings.ToLower(strings.TrimSpace(req.Quality))
- if req.Quality != imageQualityLow &&
- req.Quality != imageQualityMedium &&
- req.Quality != imageQualityHigh {
- return nil, errors.New(errInvalidImageCacheWarmupQuality)
+ if req.Quality != shared.ImageQualityLow &&
+ req.Quality != shared.ImageQualityMedium &&
+ req.Quality != shared.ImageQualityHigh {
+ return nil, errors.New(shared.ErrInvalidImageCacheWarmupQuality)
}
return json.Marshal(req)
@@ -92,7 +89,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
var req WarmImageCachePayload
if err := json.Unmarshal(normalizedPayload, &req); err != nil {
- return nil, fmt.Errorf(errParseImageCacheWarmupPayload, err)
+ return nil, fmt.Errorf(shared.ErrParseImageCacheWarmupPayload, err)
}
task.AppendLog(ctx, "等待获取图片缓存预热执行锁,质量: %s", req.Quality)
@@ -128,7 +125,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
Limit(batchSize).
Find(&uploads).Error; err != nil {
task.AppendLog(ctx, "查询图片上传记录失败: %v", err)
- return nil, fmt.Errorf(errQueryImagesForCacheWarmup, err)
+ return nil, fmt.Errorf(shared.ErrQueryImagesForCacheWarmup, err)
}
if len(uploads) == 0 {
@@ -147,7 +144,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
totalProcessed++
lastID = upload.ID
- _, cacheHit, err := ensureCompressedImageCache(ctx, upload, req.Quality)
+ _, cacheHit, err := filesrv.EnsureCompressedImageCache(ctx, upload, req.Quality)
if err != nil {
totalFailed++
batchFailed++
@@ -184,4 +181,4 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
)
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
-}
+}
\ No newline at end of file
diff --git a/internal/apps/upload/tasks_test.go b/internal/apps/upload/task/tasks_test.go
similarity index 96%
rename from internal/apps/upload/tasks_test.go
rename to internal/apps/upload/task/tasks_test.go
index 27bb5439..79b3293b 100644
--- a/internal/apps/upload/tasks_test.go
+++ b/internal/apps/upload/task/tasks_test.go
@@ -2,7 +2,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+package task
import (
"bytes"
@@ -17,6 +17,8 @@ import (
"testing"
"time"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
@@ -201,7 +203,7 @@ func TestWarmImageCacheHandlerValidatePayload(t *testing.T) {
{
name: "normalizes quality",
payload: []byte(`{"quality":" HIGH "}`),
- wantQuality: imageQualityHigh,
+ wantQuality: shared.ImageQualityHigh,
},
{
name: "empty payload",
@@ -281,7 +283,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
FilePath: firstPath,
MimeType: "image/png",
Extension: "png",
- StorageDriver: storageDriverLocal,
+ StorageDriver: shared.StorageDriverLocal,
Status: model.UploadStatusUsed,
},
{
@@ -291,7 +293,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
FilePath: secondPath,
MimeType: "application/octet-stream",
Extension: "jpg",
- StorageDriver: storageDriverLocal,
+ StorageDriver: shared.StorageDriverLocal,
Status: model.UploadStatusPending,
},
{
@@ -301,7 +303,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
FilePath: filepath.Join(testDir, "notes.txt"),
MimeType: "text/plain",
Extension: "txt",
- StorageDriver: storageDriverLocal,
+ StorageDriver: shared.StorageDriverLocal,
Status: model.UploadStatusUsed,
},
{
@@ -311,7 +313,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
FilePath: firstPath,
MimeType: "image/png",
Extension: "png",
- StorageDriver: storageDriverLocal,
+ StorageDriver: shared.StorageDriverLocal,
Status: model.UploadStatusDeleted,
},
}
@@ -339,7 +341,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
}
for i := range records[:2] {
- key := imageCompressionCacheKey(&records[i], imageQualityLow)
+ key := filesrv.ImageCompressionCacheKey(&records[i], shared.ImageQualityLow)
got, err := cache.Get(key)
if err != nil {
t.Errorf("cache.Get(%q) returned error: %v", key, err)
diff --git a/internal/apps/upload/util/media.go b/internal/apps/upload/util/media.go
new file mode 100644
index 00000000..0550422c
--- /dev/null
+++ b/internal/apps/upload/util/media.go
@@ -0,0 +1,50 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package util
+
+import (
+ "strings"
+
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
+)
+
+// IsImageExtension reports whether ext is a common image format.
+func IsImageExtension(ext string) bool {
+ for _, imgExt := range []string{"jpg", "jpeg", "png", "webp", "gif"} {
+ if ext == imgExt {
+ return true
+ }
+ }
+ return false
+}
+
+// IsArchiveExtension reports whether ext is a common archive format.
+func IsArchiveExtension(ext string) bool {
+ for _, e := range []string{"zip", "rar", "7z", "tar", "gz", "tgz", "bz2", "xz"} {
+ if ext == e {
+ return true
+ }
+ }
+ return false
+}
+
+// IsDocumentExtension reports whether ext is a common document format.
+func IsDocumentExtension(ext string) bool {
+ for _, e := range []string{"pdf", "doc", "docx", "xls", "xlsx", "ppt", "pptx", "txt", "md", "csv", "json", "yaml", "yml", "xml"} {
+ if ext == e {
+ return true
+ }
+ }
+ return false
+}
+
+// NormalizeImageQuality normalizes the requested image quality query parameter.
+func NormalizeImageQuality(quality string) string {
+ switch strings.ToLower(quality) {
+ case shared.ImageQualityLow, shared.ImageQualityMedium, shared.ImageQualityHigh:
+ return strings.ToLower(quality)
+ default:
+ return shared.ImageQualityOrigin
+ }
+}
\ No newline at end of file
diff --git a/internal/apps/upload/utils.go b/internal/apps/upload/util/utils.go
similarity index 73%
rename from internal/apps/upload/utils.go
rename to internal/apps/upload/util/utils.go
index 5f65aa4c..177f7365 100644
--- a/internal/apps/upload/utils.go
+++ b/internal/apps/upload/util/utils.go
@@ -2,7 +2,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
-package upload
+// Package util provides upload media helpers and image utilities.
+package util
import (
"bytes"
@@ -15,28 +16,27 @@ import (
"io"
"strings"
+ "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/deepteams/webp"
_ "golang.org/x/image/webp" // Register WebP decoder for image.Decode
)
-const maxS3KeyLength = 1024
-
// ValidateS3Key validates an S3 object key for safety.
func ValidateS3Key(key string) error {
if key == "" {
- return errors.New(ErrS3KeyRequired)
+ return errors.New(shared.ErrS3KeyRequired)
}
- if len(key) > maxS3KeyLength {
- return fmt.Errorf(ErrS3KeyTooLongFormat, maxS3KeyLength)
+ if len(key) > shared.MaxS3KeyLength {
+ return fmt.Errorf(shared.ErrS3KeyTooLongFormat, shared.MaxS3KeyLength)
}
if strings.HasPrefix(key, "/") {
- return errors.New(ErrS3KeyStartsWithSlash)
+ return errors.New(shared.ErrS3KeyStartsWithSlash)
}
if strings.Contains(key, "\x00") {
- return errors.New(ErrS3KeyContainsNullBytes)
+ return errors.New(shared.ErrS3KeyContainsNullBytes)
}
return nil
@@ -45,34 +45,31 @@ func ValidateS3Key(key string) error {
// CompressImageToWebP decodes an image from srcReader and encodes it into WebP format
// using the specified quality (low -> 60, medium -> 75, high -> 85).
func CompressImageToWebP(srcReader io.Reader, quality string) ([]byte, error) {
- // Decode the image
img, format, err := image.Decode(srcReader)
if err != nil {
return nil, fmt.Errorf("failed to decode image (format: %s): %w", format, err)
}
- // Determine quality
var qualityScore float32
switch strings.ToLower(quality) {
- case imageQualityLow:
+ case shared.ImageQualityLow:
qualityScore = 60
- case imageQualityMedium:
+ case shared.ImageQualityMedium:
qualityScore = 75
- case imageQualityHigh, "":
+ case shared.ImageQualityHigh, "":
qualityScore = 85
default:
qualityScore = 85
}
- // Encode to WebP
var buf bytes.Buffer
err = webp.Encode(&buf, img, &webp.EncoderOptions{
Quality: qualityScore,
- Method: 4, // Default method
+ Method: 4,
})
if err != nil {
return nil, fmt.Errorf("failed to encode WebP: %w", err)
}
return buf.Bytes(), nil
-}
+}
\ No newline at end of file
diff --git a/internal/apps/user/access_tokens.go b/internal/apps/user/access_tokens.go
index e128a981..28e98644 100644
--- a/internal/apps/user/access_tokens.go
+++ b/internal/apps/user/access_tokens.go
@@ -43,7 +43,7 @@ func ListAccessTokens(c *gin.Context) {
var tokens []model.AccessToken
if err := db.DB(ctx).Where("user_id = ?", currUser.ID).Order("created_at desc").Find(&tokens).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -67,19 +67,19 @@ func CreateAccessToken(c *gin.Context) {
var req createTokenRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusOK, response.Err(errBindParamsFailed))
+ response.AbortBadRequest(c, errBindParamsFailed)
return
}
req.Name = strings.TrimSpace(req.Name)
if req.Name == "" {
- c.JSON(http.StatusOK, response.Err(errTokenNameRequired))
+ response.AbortBadRequest(c, errTokenNameRequired)
return
}
// 只有管理员才能创建具有管理员权限的令牌
if req.IsAdmin && !currUser.IsAdmin {
- c.JSON(http.StatusOK, response.Err(errAdminTokenRequiresAdmin))
+ response.AbortBadRequest(c, errAdminTokenRequiresAdmin)
return
}
@@ -91,19 +91,19 @@ func CreateAccessToken(c *gin.Context) {
var count int64
if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", currUser.ID).Count(&count).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if int(count) >= maxLimit {
- c.JSON(http.StatusOK, response.Err(errAccessTokenLimitReached))
+ response.AbortBadRequest(c, errAccessTokenLimitReached)
return
}
// 生成 Token
tokenStr, err := model.GenerateTokenString()
if err != nil {
- c.JSON(http.StatusOK, response.Err(errGenerateTokenFailed))
+ response.AbortBadRequest(c, errGenerateTokenFailed)
return
}
@@ -119,7 +119,7 @@ func CreateAccessToken(c *gin.Context) {
}
if err := db.DB(ctx).Create(&tokenRecord).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -146,18 +146,18 @@ func DeleteAccessToken(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
- c.JSON(http.StatusOK, response.Err(errInvalidTokenID))
+ response.AbortBadRequest(c, errInvalidTokenID)
return
}
tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).Delete(&model.AccessToken{})
if tx.Error != nil {
- c.JSON(http.StatusOK, response.Err(tx.Error.Error()))
+ response.AbortBadRequest(c, tx.Error.Error())
return
}
if tx.RowsAffected == 0 {
- c.JSON(http.StatusOK, response.Err(errTokenNotFoundOrForbidden))
+ response.AbortBadRequest(c, errTokenNotFoundOrForbidden)
return
}
@@ -181,20 +181,20 @@ func RotateAccessToken(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
- c.JSON(http.StatusOK, response.Err(errInvalidTokenID))
+ response.AbortBadRequest(c, errInvalidTokenID)
return
}
var tokenRecord model.AccessToken
if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).First(&tokenRecord).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(errTokenNotFoundOrForbidden))
+ response.AbortBadRequest(c, errTokenNotFoundOrForbidden)
return
}
// 生成新的 Token
newTokenStr, err := model.GenerateTokenString()
if err != nil {
- c.JSON(http.StatusOK, response.Err(errGenerateTokenFailed))
+ response.AbortBadRequest(c, errGenerateTokenFailed)
return
}
@@ -205,7 +205,7 @@ func RotateAccessToken(c *gin.Context) {
tokenRecord.MaskedToken = newMaskedToken
if err := db.DB(ctx).Save(&tokenRecord).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go
index 1607d36b..43092dbf 100644
--- a/internal/apps/user/routers.go
+++ b/internal/apps/user/routers.go
@@ -125,17 +125,17 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro
// @Router /api/v1/user/login [post]
func Login(c *gin.Context) {
if !isPasswordLoginEnabled() {
- c.JSON(http.StatusOK, response.Err(errPasswordLoginDisabled))
+ response.AbortBadRequest(c, errPasswordLoginDisabled)
return
}
var req loginRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
req.Username = strings.TrimSpace(req.Username)
if req.Username == "" || req.Password == "" {
- c.JSON(http.StatusOK, response.Err(errInvalidParams))
+ response.AbortBadRequest(c, errInvalidParams)
return
}
@@ -143,12 +143,12 @@ func Login(c *gin.Context) {
ctx := c.Request.Context()
if err := db.DB(ctx).Where("username = ? OR email = ?", req.Username, req.Username).First(&user).Error; err != nil {
logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP())
- c.JSON(http.StatusOK, response.Err(errUsernameOrPasswordWrong))
+ response.AbortBadRequest(c, errUsernameOrPasswordWrong)
return
}
if !user.IsActive {
logger.WarnF(ctx, "[LoginAudit] banned user login attempt for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
- c.JSON(http.StatusOK, response.Err(common.BannedAccount))
+ response.AbortBadRequest(c, common.BannedAccount)
return
}
@@ -157,18 +157,18 @@ func Login(c *gin.Context) {
if !user.CheckPassword(req.Password) {
logger.WarnF(ctx, "[LoginAudit] failed login attempt (incorrect password) for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
- c.JSON(http.StatusOK, response.Err(errUsernameOrPasswordWrong))
+ response.AbortBadRequest(c, errUsernameOrPasswordWrong)
return
}
if isEmailLoginVerificationEnabled(ctx) {
result, err := processLoginEmailVerification(ctx, req.Code, &user)
if err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if result.Status != LoginEmailVerificationPassed {
- c.JSON(http.StatusOK, response.Err(result.Message))
+ response.AbortBadRequest(c, result.Message)
return
}
}
@@ -184,11 +184,11 @@ func Login(c *gin.Context) {
user.LastLoginAt = time.Now()
if err := db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if err := setLoginSession(ctx, c, &user); err != nil {
- c.JSON(http.StatusOK, response.Err(errSaveSessionFailed))
+ response.AbortBadRequest(c, errSaveSessionFailed)
return
}
@@ -212,13 +212,13 @@ func Login(c *gin.Context) {
// @Router /api/v1/user/register [post]
func Register(c *gin.Context) {
if !isRegistrationEnabled() || !isPasswordRegisterEnabled() {
- c.JSON(http.StatusOK, response.Err(errRegistrationDisabled))
+ response.AbortBadRequest(c, errRegistrationDisabled)
return
}
var req registerRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -230,15 +230,15 @@ func Register(c *gin.Context) {
req.Code = strings.TrimSpace(req.Code)
if req.Username == "" || req.Password == "" {
- c.JSON(http.StatusOK, response.Err(errInvalidParams))
+ response.AbortBadRequest(c, errInvalidParams)
return
}
if req.Email == "" {
- c.JSON(http.StatusOK, response.Err(errEmailRequired))
+ response.AbortBadRequest(c, errEmailRequired)
return
}
if len(req.Password) < minPasswordLength {
- c.JSON(http.StatusOK, response.Err(errPasswordTooShort))
+ response.AbortBadRequest(c, errPasswordTooShort)
return
}
@@ -246,7 +246,7 @@ func Register(c *gin.Context) {
// 邮箱注册验证校验
if err := validateRegisterEmailVerification(ctx, req.Email, req.Code); err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -267,17 +267,17 @@ func Register(c *gin.Context) {
user.Nickname = req.Username
}
if err := user.SetEncryptedPassword(req.Password); err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if err := user.RegisterUser(ctx, db.DB(ctx)); err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
if err := setLoginSession(ctx, c, &user); err != nil {
- c.JSON(http.StatusOK, response.Err(errSaveSessionFailed))
+ response.AbortBadRequest(c, errSaveSessionFailed)
return
}
@@ -303,7 +303,7 @@ func Logout(c *gin.Context) {
session.Options(oauth.GetSessionOptions(-1))
session.Clear()
if err := session.Save(); err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(""))
@@ -328,7 +328,7 @@ type changePasswordRequest struct {
func ChangePassword(c *gin.Context) {
var req changePasswordRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -336,47 +336,47 @@ func ChangePassword(c *gin.Context) {
req.NewPassword = strings.TrimSpace(req.NewPassword)
if req.OldPassword == "" || req.NewPassword == "" {
- c.JSON(http.StatusOK, response.Err(errInvalidParams))
+ response.AbortBadRequest(c, errInvalidParams)
return
}
if len(req.NewPassword) < minPasswordLength {
- c.JSON(http.StatusOK, response.Err(errNewPasswordTooShort))
+ response.AbortBadRequest(c, errNewPasswordTooShort)
return
}
userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
if userObj == nil {
- c.JSON(http.StatusUnauthorized, response.Err(errLoginRequired))
+ response.AbortUnauthorized(c, errLoginRequired)
return
}
ctx := c.Request.Context()
var dbUser model.User
if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(errUserNotFound))
+ response.AbortBadRequest(c, errUserNotFound)
return
}
// 校验旧密码
if !dbUser.CheckPassword(req.OldPassword) {
- c.JSON(http.StatusOK, response.Err(errOldPasswordIncorrect))
+ response.AbortBadRequest(c, errOldPasswordIncorrect)
return
}
// 加密并更新为新密码
if err := dbUser.SetEncryptedPassword(req.NewPassword); err != nil {
- c.JSON(http.StatusOK, response.Err(errPasswordEncryptFailed))
+ response.AbortBadRequest(c, errPasswordEncryptFailed)
return
}
if err := db.DB(ctx).Model(&dbUser).Update("password", dbUser.Password).Error; err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
// 吊销该用户所有的 Access Token
if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Delete(&model.AccessToken{}).Error; err != nil {
- c.JSON(http.StatusOK, response.Err("吊销 Access Token 失败: "+err.Error()))
+ response.AbortBadRequest(c, "吊销 Access Token 失败: "+err.Error())
return
}
@@ -401,24 +401,24 @@ func ChangePassword(c *gin.Context) {
func SendEmailCode(c *gin.Context) {
var req sendEmailCodeRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
req.Email = strings.TrimSpace(req.Email)
if req.Email == "" {
- c.JSON(http.StatusOK, response.Err(errEmailRequired))
+ response.AbortBadRequest(c, errEmailRequired)
return
}
if req.Scene != "register" {
- c.JSON(http.StatusOK, response.Err(errUnsupportedEmailScene))
+ response.AbortBadRequest(c, errUnsupportedEmailScene)
return
}
ctx := c.Request.Context()
if err := sendRegisterEmailCode(ctx, req.Email); err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
@@ -439,20 +439,20 @@ func SendEmailCode(c *gin.Context) {
func UpdateProfile(c *gin.Context) {
var req updateProfileRequest
if err := c.ShouldBindJSON(&req); err != nil {
- c.JSON(http.StatusBadRequest, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
if userObj == nil {
- c.JSON(http.StatusUnauthorized, response.Err(errLoginRequired))
+ response.AbortUnauthorized(c, errLoginRequired)
return
}
ctx := c.Request.Context()
dbUser, err := updateUserProfile(ctx, userObj.ID, updateProfileInput(req))
if err != nil {
- c.JSON(http.StatusOK, response.Err(err.Error()))
+ response.AbortBadRequest(c, err.Error())
return
}
diff --git a/internal/bootstrap/bootstrap_test.go b/internal/bootstrap/bootstrap_test.go
index ef678c03..7823d14b 100644
--- a/internal/bootstrap/bootstrap_test.go
+++ b/internal/bootstrap/bootstrap_test.go
@@ -20,15 +20,30 @@ func TestInitSyncsPushEventsOnce(t *testing.T) {
t.Fatalf("auto migrate push events failed: %v", err)
}
+ RegisterPushDomainEvents()
+
+ wantCount := len(admin_push.BuiltInEvents)
+ if wantCount < 1 {
+ t.Fatalf("built-in push events = %d, want at least 1", wantCount)
+ }
+
ctx := context.Background()
Init(ctx, Options{})
- Init(ctx, Options{API: true})
+ Init(ctx, Options{API: true}) // second Init must not duplicate events (initRuntimeOnce)
var count int64
if err := dbConn.Model(&model.PushEvent{}).Count(&count).Error; err != nil {
t.Fatalf("count push events failed: %v", err)
}
- if count != int64(len(admin_push.BuiltInEvents)) {
- t.Fatalf("push event count = %d, want %d", count, len(admin_push.BuiltInEvents))
+ if count != int64(wantCount) {
+ t.Fatalf("push event count = %d, want %d", count, wantCount)
+ }
+
+ var adminLogin model.PushEvent
+ if err := dbConn.Where("event_key = ?", "admin_login").First(&adminLogin).Error; err != nil {
+ t.Fatalf("admin_login event not found after Init: %v", err)
+ }
+ if adminLogin.Name != "管理员登录" {
+ t.Fatalf("admin_login name = %q, want %q", adminLogin.Name, "管理员登录")
}
}
\ No newline at end of file
diff --git a/internal/common/response/abort.go b/internal/common/response/abort.go
new file mode 100644
index 00000000..581629ab
--- /dev/null
+++ b/internal/common/response/abort.go
@@ -0,0 +1,45 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package response
+
+import (
+ "net/http"
+
+ "github.com/gin-gonic/gin"
+)
+
+// AbortBadRequest 以 400 中断请求并将错误挂载到 Gin Error 链,供全局中间件统一记录 Trace 并响应。
+func AbortBadRequest(c *gin.Context, msg string) {
+ AbortWithError(c, http.StatusBadRequest, msg)
+}
+
+// AbortUnauthorized 以 401 中断请求并将错误挂载到 Gin Error 链。
+func AbortUnauthorized(c *gin.Context, msg string) {
+ AbortWithError(c, http.StatusUnauthorized, msg)
+}
+
+// AbortForbidden 以 403 中断请求并将错误挂载到 Gin Error 链。
+func AbortForbidden(c *gin.Context, msg string) {
+ AbortWithError(c, http.StatusForbidden, msg)
+}
+
+// AbortNotFound 以 404 中断请求并将错误挂载到 Gin Error 链。
+func AbortNotFound(c *gin.Context, msg string) {
+ AbortWithError(c, http.StatusNotFound, msg)
+}
+
+// AbortInternal 以 500 中断请求并将错误挂载到 Gin Error 链。
+func AbortInternal(c *gin.Context, msg string) {
+ AbortWithError(c, http.StatusInternalServerError, msg)
+}
+
+// AbortTooManyRequests 以 429 中断请求并将错误挂载到 Gin Error 链。
+func AbortTooManyRequests(c *gin.Context, msg string) {
+ AbortWithError(c, http.StatusTooManyRequests, msg)
+}
+
+// AbortConflict 以 409 中断请求并将错误挂载到 Gin Error 链。
+func AbortConflict(c *gin.Context, msg string) {
+ AbortWithError(c, http.StatusConflict, msg)
+}
\ No newline at end of file
diff --git a/internal/common/response/middleware.go b/internal/common/response/middleware.go
new file mode 100644
index 00000000..5f895085
--- /dev/null
+++ b/internal/common/response/middleware.go
@@ -0,0 +1,40 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package response
+
+import (
+ "errors"
+ "net/http"
+
+ "github.com/gin-gonic/gin"
+ "go.opentelemetry.io/otel/codes"
+ "go.opentelemetry.io/otel/trace"
+)
+
+// ErrorHandlerMiddleware 捕获 c.Errors 并统一格式化为 JSON 返回给客户端,同时将其记录到 Span 异常中。
+// 与 AbortWithError / AbortBadRequest 等配合使用,是全局 OTel 友好错误响应的唯一出口。
+func ErrorHandlerMiddleware() gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.Next()
+
+ if len(c.Errors) == 0 || c.Writer.Written() {
+ return
+ }
+
+ err := c.Errors.Last().Err
+ span := trace.SpanFromContext(c.Request.Context())
+ if span.IsRecording() {
+ span.RecordError(err)
+ span.SetStatus(codes.Error, err.Error())
+ }
+
+ var apiErr *APIError
+ if errors.As(err, &apiErr) {
+ c.JSON(apiErr.Code, Err(apiErr.Msg))
+ return
+ }
+
+ c.JSON(http.StatusInternalServerError, Err("内部系统错误"))
+ }
+}
\ No newline at end of file
diff --git a/internal/router/middlewares.go b/internal/router/middlewares.go
index 77f05633..73fb89b3 100644
--- a/internal/router/middlewares.go
+++ b/internal/router/middlewares.go
@@ -7,7 +7,6 @@ package router
import (
"context"
- "errors"
"net/http"
"strconv"
"strings"
@@ -106,29 +105,7 @@ func corsMiddleware() gin.HandlerFunc {
}
}
-// errorHandlerMiddleware 捕获 c.Errors 并统一格式化为 JSON 返回给客户端,同时将其记录到 Span 异常中
+// errorHandlerMiddleware 委托给 response.ErrorHandlerMiddleware,保持路由层单一入口。
func errorHandlerMiddleware() gin.HandlerFunc {
- return func(c *gin.Context) {
- c.Next()
-
- if len(c.Errors) > 0 {
- err := c.Errors.Last().Err
- span := trace.SpanFromContext(c.Request.Context())
-
- // 1. 如果有活跃的 Span,将错误信息记录到 Trace 中,并把 Span 状态置为 Error
- if span.IsRecording() {
- span.RecordError(err)
- span.SetStatus(codes.Error, err.Error())
- }
-
- // 2. 将错误转化为统一的 JSON 格式响应给客户端
- var apiErr *response.APIError
- if errors.As(err, &apiErr) {
- c.JSON(apiErr.Code, response.Err(apiErr.Msg))
- } else {
- // 兜底策略:未知的系统级错误
- c.JSON(http.StatusInternalServerError, response.Err("内部系统错误"))
- }
- }
- }
+ return response.ErrorHandlerMiddleware()
}
diff --git a/internal/testhelper/gin.go b/internal/testhelper/gin.go
new file mode 100644
index 00000000..11b881b4
--- /dev/null
+++ b/internal/testhelper/gin.go
@@ -0,0 +1,20 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package testhelper
+
+import (
+ "github.com/Rain-kl/Wavelet/internal/common/response"
+ "github.com/gin-gonic/gin"
+)
+
+// NewTestGinEngine 创建带 ErrorHandlerMiddleware 的 Gin 引擎,与生产环境错误响应行为一致。
+func NewTestGinEngine(middlewares ...gin.HandlerFunc) *gin.Engine {
+ gin.SetMode(gin.TestMode)
+ r := gin.New()
+ r.Use(response.ErrorHandlerMiddleware())
+ for _, middleware := range middlewares {
+ r.Use(middleware)
+ }
+ return r
+}
\ No newline at end of file