mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-30 06:16:37 +08:00
merge main into handler-model-logics-user
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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 链触发注册。
|
||||
|
||||
## 日志要求
|
||||
|
||||
|
||||
@@ -17,7 +17,9 @@ Wavelet 的消息推送机制采用了**元数据驱动 + 统一触发器 + 异
|
||||
| :--- | :--- | :--- |
|
||||
| **`pkg/push/`** | 推送基础设施层 | 静态定义、不依赖系统数据库和任何框架。定义了统一接口 `Pusher`、单例 `PusherPool` 和多实现(Lark, Webhook, Email 等),提供配置验证及发送功能。 |
|
||||
| **`internal/apps/admin/push/`** | 通知服务与后台任务层 | 包含以下核心文件:<br>1. [events.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/events.go):定义通知事件的结构模型(`NotificationMessage`, `EventMetadata`)、内置事件的动态注册中心(`BuiltInEvents` 及 `RegisterBuiltInEvent` 函数)以及统一触发器类 `EventTrigger`(包括其底层的派发引擎逻辑)。<br>2. [tasks.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/tasks.go):定义 Asynq 后台异步发送任务、处理器 `PushHandler` 及其校验逻辑,并记录推送历史审计。<br>3. [routers.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/routers.go):管理端接口,负责获取事件配置列表和更新配置。 |
|
||||
| **`internal/apps/admin/push/custom_events/`** | 自定义通知事件包 | 自定义的事件定义和触发函数单独放到此包下,**一个 Go 文件代表一个事件**。例如:<br>[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` 存放每个通知事件的启用状态、启用渠道、发送目标和自定义渲染模板。<br>`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` 入口显式装配。
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
});
|
||||
@@ -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.<Section>.<Field>`。
|
||||
- `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`。
|
||||
|
||||
中间件:
|
||||
|
||||
|
||||
@@ -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())
|
||||
|
||||
Vendored
+5
-5
@@ -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())
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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(),
|
||||
})
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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}))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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())))
|
||||
}
|
||||
@@ -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()
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -29,7 +29,7 @@ func RiskControlMiddleware() gin.HandlerFunc {
|
||||
|
||||
// 1. 限流背压检测(检测本地缓冲队列是否已满)
|
||||
if IsBufferFull() {
|
||||
c.AbortWithStatusJSON(http.StatusTooManyRequests, response.Err("系统繁忙,请稍后再试"))
|
||||
response.AbortTooManyRequests(c, "系统繁忙,请稍后再试")
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -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")
|
||||
})
|
||||
|
||||
+15
-73
@@ -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
|
||||
}
|
||||
Vendored
+18
-15
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
)
|
||||
@@ -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
|
||||
)
|
||||
@@ -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)
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
+14
-60
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
+34
-31
@@ -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))
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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 ""
|
||||
}
|
||||
}
|
||||
+10
-5
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
)
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package upload
|
||||
package stats
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
+12
-91
@@ -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
|
||||
}
|
||||
}
|
||||
+10
-4
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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, "管理员登录")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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("内部系统错误"))
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user