mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 23:56:37 +08:00
async task framework
This commit is contained in:
@@ -1,5 +1,5 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
Copyright 2025-2026 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
@@ -17,8 +17,13 @@ limitations under the License.
|
||||
package task
|
||||
|
||||
const (
|
||||
InvalidTaskType = "无效的任务类型"
|
||||
InvalidTimeRange = "无效的时间范围"
|
||||
TaskDispatchFailed = "任务下发失败"
|
||||
UserIDRequired = "用户ID必填"
|
||||
InvalidTaskType = "无效的任务类型"
|
||||
InvalidTimeRange = "无效的时间范围"
|
||||
TaskDispatchFailed = "任务下发失败"
|
||||
UserIDRequired = "用户ID必填"
|
||||
TaskNotFound = "任务执行记录不存在"
|
||||
TaskNotRetryable = "该任务不支持重试"
|
||||
TaskNotFailed = "只有失败的任务才能重试"
|
||||
TaskMaxRetryExceeded = "已达到最大重试次数"
|
||||
TaskRetryFailed = "任务重试失败"
|
||||
)
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
Copyright 2025-2026 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
@@ -19,12 +19,13 @@ package task
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/hibiken/asynq"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/task"
|
||||
"github.com/linux-do/credit/internal/task/scheduler"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
)
|
||||
|
||||
@@ -57,7 +58,7 @@ type DispatchTaskRequest struct {
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body task.DispatchTaskRequest true "任务请求参数"
|
||||
// @Param request body DispatchTaskRequest true "任务请求参数"
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "任务已入队"
|
||||
// @Failure 400 {object} util.ResponseAny "任务类型不存在或参数错误"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
@@ -77,22 +78,113 @@ func DispatchTask(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
var taskInfo *asynq.Task
|
||||
var taskID string
|
||||
|
||||
taskInfo = asynq.NewTask(meta.AsynqTask, nil)
|
||||
taskID = fmt.Sprintf("manual_%s", req.TaskType)
|
||||
|
||||
_, err := scheduler.AsynqClient.Enqueue(
|
||||
taskInfo,
|
||||
asynq.TaskID(taskID),
|
||||
asynq.MaxRetry(meta.MaxRetry),
|
||||
asynq.Queue(meta.Queue),
|
||||
)
|
||||
taskID, err := task.DispatchTask(c.Request.Context(), req.TaskType, nil, "manual")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(fmt.Sprintf("%s: %v", TaskDispatchFailed, err)))
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, util.OKNil())
|
||||
c.JSON(http.StatusOK, util.OK(taskID))
|
||||
}
|
||||
|
||||
// ListTaskExecutions 查询任务执行记录列表
|
||||
// @Summary 查询任务执行记录
|
||||
// @Description 分页查询任务执行记录,支持按状态和任务类型筛选,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param status query string false "状态筛选 (pending/running/succeeded/failed)"
|
||||
// @Param task_type query string false "任务类型筛选"
|
||||
// @Param page query int false "页码" default(1)
|
||||
// @Param page_size query int false "每页条数" default(20)
|
||||
// @Success 200 {object} util.ResponseAny{data=object} "任务执行记录列表"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Router /api/v1/admin/tasks/executions [get]
|
||||
func ListTaskExecutions(c *gin.Context) {
|
||||
var req model.ListTaskExecutionsRequest
|
||||
if err := c.ShouldBindQuery(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
executions, total, err := model.ListTaskExecutions(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, util.OK(gin.H{
|
||||
"items": executions,
|
||||
"total": total,
|
||||
"page": req.Page,
|
||||
"page_size": req.PageSize,
|
||||
}))
|
||||
}
|
||||
|
||||
// GetTaskExecution 查询单条任务执行详情
|
||||
// @Summary 查询任务执行详情
|
||||
// @Description 根据 ID 查询任务执行记录详情,包含完整执行日志,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "任务执行记录 ID"
|
||||
// @Success 200 {object} util.ResponseAny{data=model.TaskExecution} "任务执行详情"
|
||||
// @Failure 400 {object} util.ResponseAny "参数错误"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Failure 404 {object} util.ResponseAny "记录不存在"
|
||||
// @Router /api/v1/admin/tasks/executions/{id} [get]
|
||||
func GetTaskExecution(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err("无效的任务执行记录 ID"))
|
||||
return
|
||||
}
|
||||
|
||||
execution, err := model.GetTaskExecutionByID(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusNotFound, util.Err(TaskNotFound))
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, util.OK(execution))
|
||||
}
|
||||
|
||||
// RetryTask 重试失败的任务
|
||||
// @Summary 重试失败任务
|
||||
// @Description 重新下发一条失败的任务,创建新的执行记录,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "任务执行记录 ID"
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "新任务的 TaskID"
|
||||
// @Failure 400 {object} util.ResponseAny "任务不支持重试或参数错误"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Failure 404 {object} util.ResponseAny "记录不存在"
|
||||
// @Failure 500 {object} util.ResponseAny "重试失败"
|
||||
// @Router /api/v1/admin/tasks/executions/{id}/retry [post]
|
||||
func RetryTask(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err("无效的任务执行记录 ID"))
|
||||
return
|
||||
}
|
||||
|
||||
newTaskID, err := task.RetryTask(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
errMsg := err.Error()
|
||||
switch {
|
||||
case strings.Contains(errMsg, "不存在"):
|
||||
c.JSON(http.StatusNotFound, util.Err(errMsg))
|
||||
case strings.Contains(errMsg, "只有失败") || strings.Contains(errMsg, "不支持重试") || strings.Contains(errMsg, "已达到最大重试"):
|
||||
c.JSON(http.StatusBadRequest, util.Err(errMsg))
|
||||
default:
|
||||
c.JSON(http.StatusInternalServerError, util.Err(fmt.Sprintf("%s: %v", TaskRetryFailed, err)))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, util.OK(newTaskID))
|
||||
}
|
||||
|
||||
@@ -18,10 +18,13 @@ package task
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/apps/oauth"
|
||||
@@ -29,6 +32,8 @@ import (
|
||||
"github.com/linux-do/credit/internal/task"
|
||||
"github.com/linux-do/credit/internal/testhelper"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||
@@ -46,6 +51,9 @@ func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||
|
||||
adminGroup.GET("/tasks/types", ListTaskTypes)
|
||||
adminGroup.POST("/tasks/dispatch", DispatchTask)
|
||||
adminGroup.GET("/tasks/executions", ListTaskExecutions)
|
||||
adminGroup.GET("/tasks/executions/:id", GetTaskExecution)
|
||||
adminGroup.POST("/tasks/executions/:id/retry", RetryTask)
|
||||
return r
|
||||
}
|
||||
|
||||
@@ -104,9 +112,17 @@ func TestDispatchTask(t *testing.T) {
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
assert.Equal(t, http.StatusOK, w.Code, "Body: %s", w.Body.String())
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.Empty(t, resp.ErrorMsg)
|
||||
assert.NotNil(t, resp.Data)
|
||||
|
||||
// 返回的 data 应该是 taskID
|
||||
taskID, ok := resp.Data.(string)
|
||||
assert.True(t, ok)
|
||||
assert.NotEmpty(t, taskID)
|
||||
})
|
||||
|
||||
t.Run("dispatch invalid task type failure", func(t *testing.T) {
|
||||
@@ -119,14 +135,307 @@ func TestDispatchTask(t *testing.T) {
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Errorf("expected 400 Bad Request, got %d", w.Code)
|
||||
}
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
if resp.ErrorMsg != InvalidTaskType {
|
||||
t.Errorf("expected error message '%s', got '%s'", InvalidTaskType, resp.ErrorMsg)
|
||||
}
|
||||
assert.Equal(t, InvalidTaskType, resp.ErrorMsg)
|
||||
})
|
||||
|
||||
t.Run("dispatch with empty body failure", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer([]byte("{}")))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
}
|
||||
|
||||
func TestListTaskExecutions(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
ctx := context.Background()
|
||||
|
||||
// 准备测试数据
|
||||
now := time.Now()
|
||||
records := []*model.TaskExecution{
|
||||
{TaskID: "exec_001", TaskType: "upload:cleanup_unused", TaskName: "清理上传", Status: model.TaskExecutionStatusSucceeded, TriggeredBy: "manual", Duration: 1500, Result: "清理完成", StartedAt: &now, FinishedAt: &now},
|
||||
{TaskID: "exec_002", TaskType: "upload:cleanup_unused", TaskName: "清理上传", Status: model.TaskExecutionStatusFailed, TriggeredBy: "system", ErrorMessage: "连接超时", StartedAt: &now, FinishedAt: &now},
|
||||
{TaskID: "exec_003", TaskType: "upload:cleanup_unused", TaskName: "清理上传", Status: model.TaskExecutionStatusPending, TriggeredBy: "manual", Retryable: true, MaxRetry: 3},
|
||||
}
|
||||
for _, r := range records {
|
||||
err := model.CreateTaskExecution(ctx, r)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
t.Run("list all executions", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(3), data["total"])
|
||||
})
|
||||
|
||||
t.Run("filter by status", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?status=failed", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(1), data["total"])
|
||||
})
|
||||
|
||||
t.Run("filter by task_type", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?task_type=upload:cleanup_unused", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(3), data["total"])
|
||||
})
|
||||
|
||||
t.Run("pagination", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?page=1&page_size=2", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(3), data["total"])
|
||||
})
|
||||
}
|
||||
|
||||
func TestGetTaskExecution(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建测试记录
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "detail_001",
|
||||
TaskType: "upload:cleanup_unused",
|
||||
TaskName: "清理上传",
|
||||
Status: model.TaskExecutionStatusSucceeded,
|
||||
Log: "[10:00:01] 开始扫描\n[10:00:02] 找到 50 个文件\n[10:00:03] 清理完成",
|
||||
Result: "共清理 50 个文件",
|
||||
Duration: 2000,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Run("get existing execution", func(t *testing.T) {
|
||||
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d", execution.ID)
|
||||
req, _ := http.NewRequest("GET", url, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var detail model.TaskExecution
|
||||
json.Unmarshal(dataBytes, &detail)
|
||||
|
||||
assert.Equal(t, "detail_001", detail.TaskID)
|
||||
assert.Equal(t, model.TaskExecutionStatusSucceeded, detail.Status)
|
||||
assert.Contains(t, detail.Log, "开始扫描")
|
||||
assert.Contains(t, detail.Log, "清理完成")
|
||||
assert.Equal(t, int64(2000), detail.Duration)
|
||||
})
|
||||
|
||||
t.Run("get non-existent execution", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions/99999999", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusNotFound, w.Code)
|
||||
})
|
||||
|
||||
t.Run("invalid ID format", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions/invalid", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
}
|
||||
|
||||
func TestRetryTask(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("retry failed task successfully", func(t *testing.T) {
|
||||
now := time.Now()
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "retry_api_001",
|
||||
TaskType: "upload:cleanup_unused",
|
||||
TaskName: "清理上传",
|
||||
Status: model.TaskExecutionStatusFailed,
|
||||
ErrorMessage: "S3 连接超时",
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
RetryCount: 0,
|
||||
TriggeredBy: "manual",
|
||||
StartedAt: &now,
|
||||
FinishedAt: &now,
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
|
||||
req, _ := http.NewRequest("POST", url, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.Empty(t, resp.ErrorMsg)
|
||||
assert.NotNil(t, resp.Data)
|
||||
|
||||
// 验证新记录
|
||||
newTaskID, ok := resp.Data.(string)
|
||||
assert.True(t, ok)
|
||||
assert.NotEmpty(t, newTaskID)
|
||||
|
||||
newExecution, err := model.GetTaskExecutionByTaskID(ctx, newTaskID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 1, newExecution.RetryCount)
|
||||
assert.Equal(t, "retry", newExecution.TriggeredBy)
|
||||
})
|
||||
|
||||
t.Run("retry succeeded task fails", func(t *testing.T) {
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "retry_succeeded_001",
|
||||
TaskType: "upload:cleanup_unused",
|
||||
TaskName: "清理上传",
|
||||
Status: model.TaskExecutionStatusSucceeded,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
|
||||
req, _ := http.NewRequest("POST", url, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
|
||||
t.Run("retry non-retryable task fails", func(t *testing.T) {
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "retry_not_allowed_001",
|
||||
TaskType: "upload:cleanup_unused",
|
||||
TaskName: "清理上传",
|
||||
Status: model.TaskExecutionStatusFailed,
|
||||
Retryable: false,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
|
||||
req, _ := http.NewRequest("POST", url, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
|
||||
t.Run("retry non-existent task", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/executions/99999999/retry", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusNotFound, w.Code)
|
||||
})
|
||||
|
||||
t.Run("retry with invalid ID", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/executions/invalid/retry", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
}
|
||||
|
||||
func TestRetryTaskMaxRetryExceeded(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "retry_max_api_001",
|
||||
TaskType: "upload:cleanup_unused",
|
||||
TaskName: "清理上传",
|
||||
Status: model.TaskExecutionStatusFailed,
|
||||
Retryable: true,
|
||||
MaxRetry: 1,
|
||||
RetryCount: 1,
|
||||
TriggeredBy: "retry",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
|
||||
req, _ := http.NewRequest("POST", url, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
Copyright 2025-2026 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
@@ -18,34 +18,31 @@ package upload
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/hibiken/asynq"
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/logger"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/storage"
|
||||
"github.com/linux-do/credit/internal/task"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// HandleCleanupUnusedUploads 处理清理未使用上传文件的定时任务
|
||||
func HandleCleanupUnusedUploads(ctx context.Context, t *asynq.Task) error {
|
||||
logger.InfoF(ctx, "开始清理未使用的上传文件任务")
|
||||
cleanupUnusedUploads(ctx)
|
||||
logger.InfoF(ctx, "未使用上传文件清理任务完成")
|
||||
return nil
|
||||
}
|
||||
// CleanupUnusedUploadsHandler 清理未使用上传文件的异步任务处理器
|
||||
type CleanupUnusedUploadsHandler struct{}
|
||||
|
||||
// cleanupUnusedUploads 清理超过1小时未使用的上传文件
|
||||
func cleanupUnusedUploads(ctx context.Context) {
|
||||
// Execute 执行清理未使用上传文件的业务逻辑
|
||||
func (h *CleanupUnusedUploadsHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
|
||||
const batchSize = 100 // 每批处理100个文件
|
||||
var lastID uint64 = 0
|
||||
var totalProcessed int = 0
|
||||
var totalDeleted int = 0
|
||||
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
|
||||
@@ -54,8 +51,8 @@ func cleanupUnusedUploads(ctx context.Context) {
|
||||
Order("id ASC").
|
||||
Limit(batchSize).
|
||||
Find(&unusedUploads).Error; err != nil {
|
||||
logger.ErrorF(ctx, "查询未使用的上传文件失败: %v", err)
|
||||
return
|
||||
task.AppendLog(ctx, "查询未使用的上传文件失败: %v", err)
|
||||
return nil, fmt.Errorf("查询未使用的上传文件失败: %w", err)
|
||||
}
|
||||
|
||||
// 没有更多数据,退出循环
|
||||
@@ -63,7 +60,7 @@ func cleanupUnusedUploads(ctx context.Context) {
|
||||
break
|
||||
}
|
||||
|
||||
logger.InfoF(ctx, "本批次找到 %d 个需要清理的上传文件", len(unusedUploads))
|
||||
task.AppendLog(ctx, "本批次找到 %d 个需要清理的上传文件", len(unusedUploads))
|
||||
|
||||
// 处理每个未使用的上传文件
|
||||
for _, upload := range unusedUploads {
|
||||
@@ -84,22 +81,17 @@ func cleanupUnusedUploads(ctx context.Context) {
|
||||
|
||||
return nil
|
||||
}); err != nil {
|
||||
logger.ErrorF(ctx, "清理上传文件失败 [ID:%d]: %v", upload.ID, err)
|
||||
task.AppendLog(ctx, "清理上传文件失败 [ID:%d]: %v", upload.ID, err)
|
||||
lastID = upload.ID
|
||||
continue
|
||||
}
|
||||
|
||||
totalDeleted++
|
||||
logger.InfoF(ctx, "成功清理上传文件 [ID:%d, Path:%s, Size:%d bytes]", upload.ID, upload.FilePath, upload.FileSize)
|
||||
|
||||
// 更新游标
|
||||
lastID = upload.ID
|
||||
}
|
||||
}
|
||||
|
||||
if totalDeleted > 0 {
|
||||
logger.InfoF(ctx, "清理任务完成,共处理 %d 个文件,成功删除 %d 个", totalProcessed, totalDeleted)
|
||||
} else {
|
||||
logger.InfoF(ctx, "没有需要清理的上传文件")
|
||||
}
|
||||
msg := fmt.Sprintf("共处理 %d 个文件,成功删除 %d 个", totalProcessed, totalDeleted)
|
||||
task.AppendLog(ctx, msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,139 @@
|
||||
/*
|
||||
Copyright 2025-2026 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package upload
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/storage"
|
||||
"github.com/linux-do/credit/internal/task"
|
||||
"github.com/linux-do/credit/internal/testhelper"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestCleanupUnusedUploadsHandler_Execute(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// Mock S3 存储(让 DeleteObject 总是成功)
|
||||
storageMock := storage.MockStorage(
|
||||
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
|
||||
return nil
|
||||
},
|
||||
func(ctx context.Context, key string) (*storage.ObjectInfo, error) { return nil, nil },
|
||||
func(ctx context.Context, key string) error { return nil },
|
||||
)
|
||||
defer storageMock()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// 准备测试数据:创建一些上传记录
|
||||
now := time.Now()
|
||||
twoHoursAgo := now.Add(-2 * time.Hour)
|
||||
|
||||
records := []*model.Upload{
|
||||
// 超过1小时且状态为 pending 的记录 —— 应被清理
|
||||
{
|
||||
UserID: 1001, FileName: "old_file_1.jpg", FilePath: "uploads/old_1.jpg",
|
||||
FileSize: 1024, MimeType: "image/jpeg", Extension: "jpg", Hash: "hash1",
|
||||
StorageDriver: "s3", Type: "attachment", Status: model.UploadStatusPending,
|
||||
CreatedAt: twoHoursAgo,
|
||||
},
|
||||
{
|
||||
UserID: 1001, FileName: "old_file_2.png", FilePath: "uploads/old_2.png",
|
||||
FileSize: 2048, MimeType: "image/png", Extension: "png", Hash: "hash2",
|
||||
StorageDriver: "s3", Type: "attachment", Status: model.UploadStatusPending,
|
||||
CreatedAt: twoHoursAgo,
|
||||
},
|
||||
// 状态为 used 的记录 —— 不应被清理
|
||||
{
|
||||
UserID: 1001, FileName: "used_file.jpg", FilePath: "uploads/used.jpg",
|
||||
FileSize: 512, MimeType: "image/jpeg", Extension: "jpg", Hash: "hash3",
|
||||
StorageDriver: "s3", Type: "attachment", Status: model.UploadStatusUsed,
|
||||
CreatedAt: twoHoursAgo,
|
||||
},
|
||||
// 不到1小时的 pending 记录 —— 不应被清理
|
||||
{
|
||||
UserID: 1001, FileName: "recent_file.jpg", FilePath: "uploads/recent.jpg",
|
||||
FileSize: 256, MimeType: "image/jpeg", Extension: "jpg", Hash: "hash4",
|
||||
StorageDriver: "s3", Type: "attachment", Status: model.UploadStatusPending,
|
||||
CreatedAt: now.Add(-10 * time.Minute),
|
||||
},
|
||||
}
|
||||
for _, r := range records {
|
||||
err := db.DB(ctx).Create(r).Error
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
// 执行 handler
|
||||
handler := &CleanupUnusedUploadsHandler{}
|
||||
result, err := handler.Execute(ctx, nil)
|
||||
|
||||
// 验证结果
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
assert.Contains(t, result.Message, "共处理 2 个文件,成功删除 2 个")
|
||||
|
||||
// 验证数据库状态:pending 且超过1小时的应被标记为 deleted
|
||||
var pendingCount int64
|
||||
db.DB(ctx).Model(&model.Upload{}).Where("status = ?", model.UploadStatusPending).Count(&pendingCount)
|
||||
assert.Equal(t, int64(1), pendingCount, "应只剩1条 pending 记录(最近的文件)")
|
||||
|
||||
var deletedCount int64
|
||||
db.DB(ctx).Model(&model.Upload{}).Where("status = ?", model.UploadStatusDeleted).Count(&deletedCount)
|
||||
assert.Equal(t, int64(2), deletedCount, "应有2条被标记为 deleted")
|
||||
|
||||
var usedCount int64
|
||||
db.DB(ctx).Model(&model.Upload{}).Where("status = ?", model.UploadStatusUsed).Count(&usedCount)
|
||||
assert.Equal(t, int64(1), usedCount, "used 状态的文件不应受影响")
|
||||
}
|
||||
|
||||
func TestCleanupUnusedUploadsHandler_ExecuteNoFiles(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// Mock S3 存储
|
||||
storageMock := storage.MockStorage(
|
||||
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
|
||||
return nil
|
||||
},
|
||||
func(ctx context.Context, key string) (*storage.ObjectInfo, error) { return nil, nil },
|
||||
func(ctx context.Context, key string) error { return nil },
|
||||
)
|
||||
defer storageMock()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// 没有任何上传记录
|
||||
handler := &CleanupUnusedUploadsHandler{}
|
||||
result, err := handler.Execute(ctx, nil)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
assert.Contains(t, result.Message, "共处理 0 个文件,成功删除 0 个")
|
||||
}
|
||||
|
||||
func TestCleanupUnusedUploadsHandler_ImplementsTaskHandler(t *testing.T) {
|
||||
// 编译期验证 CleanupUnusedUploadsHandler 实现了 TaskHandler 接口
|
||||
var _ task.TaskHandler = (*CleanupUnusedUploadsHandler)(nil)
|
||||
}
|
||||
@@ -38,6 +38,7 @@ func Migrate() {
|
||||
&model.SystemConfig{},
|
||||
&model.Upload{},
|
||||
&model.AccessToken{},
|
||||
&model.TaskExecution{},
|
||||
); err != nil {
|
||||
log.Fatalf("[PostgreSQL] auto migrate failed: %v\n", err)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
/*
|
||||
Copyright 2025-2026 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/db/idgen"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// TaskExecutionStatus 任务执行状态
|
||||
type TaskExecutionStatus string
|
||||
|
||||
const (
|
||||
TaskExecutionStatusPending TaskExecutionStatus = "pending"
|
||||
TaskExecutionStatusRunning TaskExecutionStatus = "running"
|
||||
TaskExecutionStatusSucceeded TaskExecutionStatus = "succeeded"
|
||||
TaskExecutionStatusFailed TaskExecutionStatus = "failed"
|
||||
)
|
||||
|
||||
// TaskExecution 任务执行记录
|
||||
type TaskExecution struct {
|
||||
ID uint64 `json:"id,string" gorm:"primaryKey"`
|
||||
TaskID string `json:"task_id" gorm:"size:128;uniqueIndex;not null"`
|
||||
TaskType string `json:"task_type" gorm:"size:64;index;not null"`
|
||||
TaskName string `json:"task_name" gorm:"size:128"`
|
||||
Status TaskExecutionStatus `json:"status" gorm:"size:32;index;not null"`
|
||||
Retryable bool `json:"retryable" gorm:"not null;default:false"`
|
||||
MaxRetry int `json:"max_retry" gorm:"not null;default:0"`
|
||||
RetryCount int `json:"retry_count" gorm:"not null;default:0"`
|
||||
Log string `json:"log" gorm:"type:text"`
|
||||
ErrorMessage string `json:"error_message" gorm:"type:text"`
|
||||
Result string `json:"result" gorm:"type:text"`
|
||||
StartedAt *time.Time `json:"started_at" gorm:"index"`
|
||||
FinishedAt *time.Time `json:"finished_at"`
|
||||
Duration int64 `json:"duration" gorm:"comment:耗时毫秒"`
|
||||
Payload string `json:"payload" gorm:"type:text"`
|
||||
TriggeredBy string `json:"triggered_by" gorm:"size:32;not null;default:system"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (TaskExecution) TableName() string {
|
||||
return "task_executions"
|
||||
}
|
||||
|
||||
// CreateTaskExecution 创建任务执行记录
|
||||
func CreateTaskExecution(ctx context.Context, execution *TaskExecution) error {
|
||||
execution.ID = idgen.NextUint64ID()
|
||||
return db.DB(ctx).Create(execution).Error
|
||||
}
|
||||
|
||||
// UpdateTaskExecution 更新任务执行记录
|
||||
func UpdateTaskExecution(ctx context.Context, execution *TaskExecution) error {
|
||||
return db.DB(ctx).Save(execution).Error
|
||||
}
|
||||
|
||||
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
|
||||
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) {
|
||||
var execution TaskExecution
|
||||
if err := db.DB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &execution, nil
|
||||
}
|
||||
|
||||
// GetTaskExecutionByID 根据 ID 获取执行记录
|
||||
func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) {
|
||||
var execution TaskExecution
|
||||
if err := db.DB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &execution, nil
|
||||
}
|
||||
|
||||
// AppendTaskExecutionLog 追加日志到执行记录
|
||||
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
|
||||
now := time.Now().Format("15:04:05")
|
||||
line := fmt.Sprintf("[%s] %s\n", now, logLine)
|
||||
return db.DB(ctx).Model(&TaskExecution{}).
|
||||
Where("task_id = ?", taskID).
|
||||
Update("log", gorm.Expr("COALESCE(log, '') || ?", line)).Error
|
||||
}
|
||||
|
||||
// ListTaskExecutionsRequest 查询任务执行记录列表请求
|
||||
type ListTaskExecutionsRequest struct {
|
||||
Status string `form:"status"`
|
||||
TaskType string `form:"task_type"`
|
||||
Page int `form:"page"`
|
||||
PageSize int `form:"page_size"`
|
||||
}
|
||||
|
||||
// ListTaskExecutions 分页查询任务执行记录
|
||||
func ListTaskExecutions(ctx context.Context, req ListTaskExecutionsRequest) ([]TaskExecution, int64, error) {
|
||||
if req.Page <= 0 {
|
||||
req.Page = 1
|
||||
}
|
||||
if req.PageSize <= 0 {
|
||||
req.PageSize = 20
|
||||
}
|
||||
|
||||
query := db.DB(ctx).Model(&TaskExecution{})
|
||||
|
||||
if req.Status != "" {
|
||||
query = query.Where("status = ?", req.Status)
|
||||
}
|
||||
if req.TaskType != "" {
|
||||
query = query.Where("task_type = ?", req.TaskType)
|
||||
}
|
||||
|
||||
var total int64
|
||||
if err := query.Count(&total).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
var executions []TaskExecution
|
||||
offset := (req.Page - 1) * req.PageSize
|
||||
if err := query.Order("id DESC").Offset(offset).Limit(req.PageSize).Find(&executions).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
return executions, total, nil
|
||||
}
|
||||
@@ -0,0 +1,301 @@
|
||||
/*
|
||||
Copyright 2025-2026 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/linux-do/credit/internal/testhelper"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestCreateTaskExecution(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &TaskExecution{
|
||||
TaskID: "manual_cleanup_123",
|
||||
TaskType: "upload:cleanup_unused",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: TaskExecutionStatusPending,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
RetryCount: 0,
|
||||
Payload: `{"test": true}`,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
assert.NotZero(t, execution.ID, "ID should be generated")
|
||||
assert.NotZero(t, execution.CreatedAt, "CreatedAt should be set")
|
||||
assert.NotZero(t, execution.UpdatedAt, "UpdatedAt should be set")
|
||||
}
|
||||
|
||||
func TestGetTaskExecutionByTaskID(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建记录
|
||||
execution := &TaskExecution{
|
||||
TaskID: "test_task_id_001",
|
||||
TaskType: "upload:cleanup_unused",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: TaskExecutionStatusPending,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 按 TaskID 查询
|
||||
found, err := GetTaskExecutionByTaskID(ctx, "test_task_id_001")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, execution.ID, found.ID)
|
||||
assert.Equal(t, "test_task_id_001", found.TaskID)
|
||||
assert.Equal(t, TaskExecutionStatusPending, found.Status)
|
||||
assert.True(t, found.Retryable)
|
||||
assert.Equal(t, 3, found.MaxRetry)
|
||||
|
||||
// 查询不存在的 TaskID
|
||||
_, err = GetTaskExecutionByTaskID(ctx, "nonexistent")
|
||||
assert.Error(t, err, "should return error for non-existent taskID")
|
||||
}
|
||||
|
||||
func TestGetTaskExecutionByID(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &TaskExecution{
|
||||
TaskID: "test_by_id_001",
|
||||
TaskType: "upload:cleanup_unused",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: TaskExecutionStatusPending,
|
||||
TriggeredBy: "system",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 按主键查询
|
||||
found, err := GetTaskExecutionByID(ctx, execution.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, execution.TaskID, found.TaskID)
|
||||
}
|
||||
|
||||
func TestUpdateTaskExecution(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建记录
|
||||
execution := &TaskExecution{
|
||||
TaskID: "test_update_001",
|
||||
TaskType: "upload:cleanup_unused",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: TaskExecutionStatusPending,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 更新状态为 running
|
||||
now := time.Now()
|
||||
execution.Status = TaskExecutionStatusRunning
|
||||
execution.StartedAt = &now
|
||||
err = UpdateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 验证更新
|
||||
found, err := GetTaskExecutionByTaskID(ctx, "test_update_001")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, TaskExecutionStatusRunning, found.Status)
|
||||
assert.NotNil(t, found.StartedAt)
|
||||
|
||||
// 更新为 succeeded
|
||||
finishTime := time.Now()
|
||||
execution.Status = TaskExecutionStatusSucceeded
|
||||
execution.FinishedAt = &finishTime
|
||||
execution.Duration = 1500
|
||||
execution.Result = "共清理 50 个文件"
|
||||
err = UpdateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
found, err = GetTaskExecutionByTaskID(ctx, "test_update_001")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, TaskExecutionStatusSucceeded, found.Status)
|
||||
assert.Equal(t, int64(1500), found.Duration)
|
||||
assert.Equal(t, "共清理 50 个文件", found.Result)
|
||||
}
|
||||
|
||||
func TestUpdateTaskExecutionFailed(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &TaskExecution{
|
||||
TaskID: "test_fail_001",
|
||||
TaskType: "upload:cleanup_unused",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: TaskExecutionStatusPending,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 标记为失败
|
||||
now := time.Now()
|
||||
execution.Status = TaskExecutionStatusFailed
|
||||
execution.StartedAt = &now
|
||||
execution.FinishedAt = &now
|
||||
execution.Duration = 200
|
||||
execution.ErrorMessage = "S3 连接超时"
|
||||
err = UpdateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
found, err := GetTaskExecutionByTaskID(ctx, "test_fail_001")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, TaskExecutionStatusFailed, found.Status)
|
||||
assert.Equal(t, "S3 连接超时", found.ErrorMessage)
|
||||
assert.Equal(t, int64(200), found.Duration)
|
||||
}
|
||||
|
||||
func TestAppendTaskExecutionLog(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &TaskExecution{
|
||||
TaskID: "test_log_001",
|
||||
TaskType: "upload:cleanup_unused",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: TaskExecutionStatusPending,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 追加多条日志
|
||||
err = AppendTaskExecutionLog(ctx, "test_log_001", "开始扫描未使用上传文件")
|
||||
require.NoError(t, err)
|
||||
|
||||
err = AppendTaskExecutionLog(ctx, "test_log_001", "本批次找到 42 个待清理文件")
|
||||
require.NoError(t, err)
|
||||
|
||||
err = AppendTaskExecutionLog(ctx, "test_log_001", "清理完成,共删除 42 个文件")
|
||||
require.NoError(t, err)
|
||||
|
||||
// 验证日志内容
|
||||
found, err := GetTaskExecutionByTaskID(ctx, "test_log_001")
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, found.Log, "开始扫描未使用上传文件")
|
||||
assert.Contains(t, found.Log, "本批次找到 42 个待清理文件")
|
||||
assert.Contains(t, found.Log, "清理完成,共删除 42 个文件")
|
||||
}
|
||||
|
||||
func TestAppendTaskExecutionLogNonExistent(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 对不存在的 TaskID 追加日志不应报错(COALESCE 处理空值)
|
||||
err := AppendTaskExecutionLog(ctx, "nonexistent_task", "测试日志")
|
||||
// SQLite 下 COALESCE + || 操作不应报错
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestListTaskExecutions(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建多条记录,包含不同状态和类型
|
||||
records := []*TaskExecution{
|
||||
{TaskID: "list_001", TaskType: "upload:cleanup_unused", TaskName: "清理上传", Status: TaskExecutionStatusSucceeded, TriggeredBy: "manual"},
|
||||
{TaskID: "list_002", TaskType: "upload:cleanup_unused", TaskName: "清理上传", Status: TaskExecutionStatusFailed, TriggeredBy: "system"},
|
||||
{TaskID: "list_003", TaskType: "other:task", TaskName: "其他任务", Status: TaskExecutionStatusPending, TriggeredBy: "manual"},
|
||||
{TaskID: "list_004", TaskType: "upload:cleanup_unused", TaskName: "清理上传", Status: TaskExecutionStatusRunning, TriggeredBy: "manual"},
|
||||
{TaskID: "list_005", TaskType: "other:task", TaskName: "其他任务", Status: TaskExecutionStatusSucceeded, TriggeredBy: "system"},
|
||||
}
|
||||
for _, r := range records {
|
||||
err := CreateTaskExecution(ctx, r)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
// 查询全部(分页)
|
||||
items, total, err := ListTaskExecutions(ctx, ListTaskExecutionsRequest{Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(5), total)
|
||||
assert.Len(t, items, 5)
|
||||
|
||||
// 按状态筛选:failed
|
||||
items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{Status: "failed", Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(1), total)
|
||||
assert.Len(t, items, 1)
|
||||
assert.Equal(t, "list_002", items[0].TaskID)
|
||||
|
||||
// 按类型筛选
|
||||
items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{TaskType: "other:task", Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(2), total)
|
||||
|
||||
// 分页测试
|
||||
items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{Page: 1, PageSize: 2})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(5), total)
|
||||
assert.Len(t, items, 2)
|
||||
|
||||
items2, total2, err := ListTaskExecutions(ctx, ListTaskExecutionsRequest{Page: 2, PageSize: 2})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(5), total2)
|
||||
assert.Len(t, items2, 2)
|
||||
|
||||
// 确保分页数据不重复
|
||||
assert.NotEqual(t, items[0].ID, items2[0].ID)
|
||||
|
||||
// 状态 + 类型组合筛选
|
||||
items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{Status: "succeeded", TaskType: "upload:cleanup_unused", Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(1), total)
|
||||
assert.Equal(t, "list_001", items[0].TaskID)
|
||||
}
|
||||
|
||||
func TestListTaskExecutionsDefaultPaging(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 不传分页参数,应使用默认值 page=1, pageSize=20
|
||||
items, total, err := ListTaskExecutions(ctx, ListTaskExecutionsRequest{})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(0), total)
|
||||
assert.Len(t, items, 0)
|
||||
}
|
||||
|
||||
func TestTaskExecutionTableName(t *testing.T) {
|
||||
execution := TaskExecution{}
|
||||
assert.Equal(t, "task_executions", execution.TableName())
|
||||
}
|
||||
@@ -164,6 +164,11 @@ func Serve() {
|
||||
adminRouter.GET("/tasks/types", admin_task.ListTaskTypes)
|
||||
adminRouter.POST("/tasks/dispatch", admin_task.DispatchTask)
|
||||
|
||||
// Task executions
|
||||
adminRouter.GET("/tasks/executions", admin_task.ListTaskExecutions)
|
||||
adminRouter.GET("/tasks/executions/:id", admin_task.GetTaskExecution)
|
||||
adminRouter.POST("/tasks/executions/:id/retry", admin_task.RetryTask)
|
||||
|
||||
// Users
|
||||
adminRouter.GET("/users", admin_user.ListUsers)
|
||||
adminRouter.PUT("/users/:id/status", admin_user.UpdateUserStatus)
|
||||
|
||||
@@ -38,6 +38,7 @@ type TaskMeta struct {
|
||||
SupportsTime bool
|
||||
MaxRetry int
|
||||
Queue string
|
||||
Retryable bool // 是否支持手动重试
|
||||
}
|
||||
|
||||
// DispatchableTasks 可下发的任务列表
|
||||
@@ -50,6 +51,7 @@ var DispatchableTasks = []TaskMeta{
|
||||
SupportsTime: false,
|
||||
MaxRetry: 3,
|
||||
Queue: QueueDefault,
|
||||
Retryable: true,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,296 @@
|
||||
/*
|
||||
Copyright 2025-2026 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package task
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/hibiken/asynq"
|
||||
"github.com/linux-do/credit/internal/db/idgen"
|
||||
"github.com/linux-do/credit/internal/logger"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/otel_trace"
|
||||
"go.opentelemetry.io/otel/attribute"
|
||||
"go.opentelemetry.io/otel/codes"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
)
|
||||
|
||||
// handlerRegistry 已注册的任务处理器
|
||||
var handlerRegistry = make(map[string]TaskHandler)
|
||||
|
||||
// RegisterHandler 注册任务处理器
|
||||
// 传入任务类型标识(对应 constants.go 中的 AsynqTask 常量)和 TaskHandler 实现
|
||||
func RegisterHandler(asynqTaskType string, handler TaskHandler) {
|
||||
handlerRegistry[asynqTaskType] = handler
|
||||
}
|
||||
|
||||
// getHandler 获取已注册的处理器
|
||||
func getHandler(asynqTaskType string) (TaskHandler, bool) {
|
||||
h, ok := handlerRegistry[asynqTaskType]
|
||||
return h, ok
|
||||
}
|
||||
|
||||
// contextKey 用于 context 存取 taskID
|
||||
type contextKey string
|
||||
|
||||
const taskIDKey contextKey = "task_execution_task_id"
|
||||
|
||||
// withTaskID 将 taskID 注入 context
|
||||
func withTaskID(ctx context.Context, taskID string) context.Context {
|
||||
return context.WithValue(ctx, taskIDKey, taskID)
|
||||
}
|
||||
|
||||
// GetTaskID 从 context 中获取 taskID
|
||||
func GetTaskID(ctx context.Context) string {
|
||||
if v, ok := ctx.Value(taskIDKey).(string); ok {
|
||||
return v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// AppendLog 追加日志到任务执行记录
|
||||
// 在 TaskHandler.Execute 中调用,日志会自动追加到 TaskExecution.Log 字段
|
||||
func AppendLog(ctx context.Context, format string, args ...interface{}) {
|
||||
taskID := GetTaskID(ctx)
|
||||
if taskID == "" {
|
||||
// 上下文中没有 taskID,降级到普通日志
|
||||
logger.InfoF(ctx, format, args...)
|
||||
return
|
||||
}
|
||||
|
||||
logLine := fmt.Sprintf(format, args...)
|
||||
if err := model.AppendTaskExecutionLog(ctx, taskID, logLine); err != nil {
|
||||
logger.ErrorF(ctx, "[TaskExecutor] 追加任务日志失败 taskID=%s: %v", taskID, err)
|
||||
}
|
||||
}
|
||||
|
||||
// DispatchTask 下发任务(创建 TaskExecution 记录 → 入队 Asynq)
|
||||
func DispatchTask(ctx context.Context, taskType string, payload []byte, triggeredBy string) (string, error) {
|
||||
meta := GetTaskMeta(taskType)
|
||||
if meta == nil {
|
||||
return "", fmt.Errorf("未知的任务类型: %s", taskType)
|
||||
}
|
||||
|
||||
// 生成唯一的 TaskID
|
||||
taskID := generateTaskID(taskType, triggeredBy)
|
||||
|
||||
// 创建任务执行记录
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: taskID,
|
||||
TaskType: meta.AsynqTask,
|
||||
TaskName: meta.Name,
|
||||
Status: model.TaskExecutionStatusPending,
|
||||
Retryable: meta.Retryable,
|
||||
MaxRetry: meta.MaxRetry,
|
||||
RetryCount: 0,
|
||||
Payload: string(payload),
|
||||
TriggeredBy: triggeredBy,
|
||||
}
|
||||
|
||||
if err := model.CreateTaskExecution(ctx, execution); err != nil {
|
||||
return "", fmt.Errorf("创建任务执行记录失败: %w", err)
|
||||
}
|
||||
|
||||
// 入队 Asynq
|
||||
taskInfo := asynq.NewTask(meta.AsynqTask, payload)
|
||||
if _, err := AsynqClient.Enqueue(
|
||||
taskInfo,
|
||||
asynq.TaskID(taskID),
|
||||
asynq.MaxRetry(meta.MaxRetry),
|
||||
asynq.Queue(meta.Queue),
|
||||
); err != nil {
|
||||
// 入队失败,更新执行记录状态
|
||||
execution.Status = model.TaskExecutionStatusFailed
|
||||
execution.ErrorMessage = fmt.Sprintf("入队失败: %v", err)
|
||||
now := time.Now()
|
||||
execution.StartedAt = &now
|
||||
execution.FinishedAt = &now
|
||||
_ = model.UpdateTaskExecution(ctx, execution)
|
||||
return "", fmt.Errorf("任务入队失败: %w", err)
|
||||
}
|
||||
|
||||
return taskID, nil
|
||||
}
|
||||
|
||||
// RetryTask 重试失败的任务
|
||||
func RetryTask(ctx context.Context, id uint64) (string, error) {
|
||||
execution, err := model.GetTaskExecutionByID(ctx, id)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("任务执行记录不存在: %w", err)
|
||||
}
|
||||
|
||||
if execution.Status != model.TaskExecutionStatusFailed {
|
||||
return "", fmt.Errorf("只有失败的任务才能重试,当前状态: %s", execution.Status)
|
||||
}
|
||||
|
||||
if !execution.Retryable {
|
||||
return "", fmt.Errorf("该任务不支持重试")
|
||||
}
|
||||
|
||||
if execution.RetryCount >= execution.MaxRetry {
|
||||
return "", fmt.Errorf("已达到最大重试次数 %d", execution.MaxRetry)
|
||||
}
|
||||
|
||||
// 生成新的 TaskID
|
||||
newTaskID := generateRetryTaskID(execution.TaskID, execution.RetryCount+1)
|
||||
|
||||
// 创建新的执行记录
|
||||
newExecution := &model.TaskExecution{
|
||||
TaskID: newTaskID,
|
||||
TaskType: execution.TaskType,
|
||||
TaskName: execution.TaskName,
|
||||
Status: model.TaskExecutionStatusPending,
|
||||
Retryable: execution.Retryable,
|
||||
MaxRetry: execution.MaxRetry,
|
||||
RetryCount: execution.RetryCount + 1,
|
||||
Payload: execution.Payload,
|
||||
TriggeredBy: "retry",
|
||||
}
|
||||
|
||||
if err := model.CreateTaskExecution(ctx, newExecution); err != nil {
|
||||
return "", fmt.Errorf("创建重试任务执行记录失败: %w", err)
|
||||
}
|
||||
|
||||
// 入队 Asynq
|
||||
taskInfo := asynq.NewTask(execution.TaskType, []byte(execution.Payload))
|
||||
if _, err := AsynqClient.Enqueue(
|
||||
taskInfo,
|
||||
asynq.TaskID(newTaskID),
|
||||
asynq.MaxRetry(execution.MaxRetry),
|
||||
asynq.Queue(PrefixedQueue(QueueDefault)),
|
||||
); err != nil {
|
||||
newExecution.Status = model.TaskExecutionStatusFailed
|
||||
newExecution.ErrorMessage = fmt.Sprintf("重试入队失败: %v", err)
|
||||
now := time.Now()
|
||||
newExecution.StartedAt = &now
|
||||
newExecution.FinishedAt = &now
|
||||
_ = model.UpdateTaskExecution(ctx, newExecution)
|
||||
return "", fmt.Errorf("重试任务入队失败: %w", err)
|
||||
}
|
||||
|
||||
return newTaskID, nil
|
||||
}
|
||||
|
||||
// ProcessTask Asynq 实际调用的统一处理函数
|
||||
// Worker 注册时统一使用此函数,内部自动分发到对应的 TaskHandler
|
||||
func ProcessTask(ctx context.Context, t *asynq.Task) error {
|
||||
// 初始化 Trace
|
||||
ctx, span := otel_trace.Start(ctx, "TaskProcess_"+t.Type(), trace.WithSpanKind(trace.SpanKindConsumer))
|
||||
defer span.End()
|
||||
|
||||
// 添加任务信息到 Span
|
||||
span.SetAttributes(
|
||||
attribute.String("task.type", t.Type()),
|
||||
attribute.Int("task.payload_size", len(t.Payload())),
|
||||
attribute.String("task.id", t.ResultWriter().TaskID()),
|
||||
)
|
||||
|
||||
taskID := t.ResultWriter().TaskID()
|
||||
|
||||
// 注入 taskID 到 context
|
||||
ctx = withTaskID(ctx, taskID)
|
||||
|
||||
// 查找处理器
|
||||
handler, ok := getHandler(t.Type())
|
||||
if !ok {
|
||||
err := fmt.Errorf("未注册的任务处理器: %s", t.Type())
|
||||
logger.ErrorF(ctx, "[TaskExecutor] %v", err)
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return err
|
||||
}
|
||||
|
||||
// 从数据库加载执行记录
|
||||
execution, err := model.GetTaskExecutionByTaskID(ctx, taskID)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "[TaskExecutor] 查询执行记录失败 taskID=%s: %v", taskID, err)
|
||||
// 执行记录不存在,仍然执行任务但不记录状态
|
||||
_, execErr := handler.Execute(ctx, t.Payload())
|
||||
if execErr != nil {
|
||||
span.SetStatus(codes.Error, execErr.Error())
|
||||
return execErr
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// 更新状态为 running
|
||||
now := time.Now()
|
||||
execution.Status = model.TaskExecutionStatusRunning
|
||||
execution.StartedAt = &now
|
||||
if err := model.UpdateTaskExecution(ctx, execution); err != nil {
|
||||
logger.ErrorF(ctx, "[TaskExecutor] 更新执行状态失败 taskID=%s: %v", taskID, err)
|
||||
}
|
||||
|
||||
// 开始计时
|
||||
start := time.Now()
|
||||
|
||||
// 执行业务逻辑
|
||||
result, execErr := handler.Execute(ctx, t.Payload())
|
||||
|
||||
// 计算耗时
|
||||
duration := time.Since(start)
|
||||
finishTime := time.Now()
|
||||
execution.Duration = duration.Milliseconds()
|
||||
execution.FinishedAt = &finishTime
|
||||
|
||||
if execErr != nil {
|
||||
// 执行失败
|
||||
execution.Status = model.TaskExecutionStatusFailed
|
||||
execution.ErrorMessage = execErr.Error()
|
||||
|
||||
logger.ErrorF(ctx,
|
||||
"[TaskExecutor] 任务处理失败 Type: %s TaskID: %s Duration: %d ms Error: %v",
|
||||
t.Type(), taskID, duration.Milliseconds(), execErr,
|
||||
)
|
||||
|
||||
span.SetStatus(codes.Error, execErr.Error())
|
||||
span.RecordError(execErr)
|
||||
} else {
|
||||
// 执行成功
|
||||
execution.Status = model.TaskExecutionStatusSucceeded
|
||||
if result != nil {
|
||||
execution.Result = result.Message
|
||||
if result.Detail != "" {
|
||||
execution.Result = fmt.Sprintf("%s\n%s", result.Message, result.Detail)
|
||||
}
|
||||
}
|
||||
|
||||
logger.InfoF(ctx,
|
||||
"[TaskExecutor] 任务处理完成 Type: %s TaskID: %s Duration: %d ms",
|
||||
t.Type(), taskID, duration.Milliseconds(),
|
||||
)
|
||||
}
|
||||
|
||||
// 更新执行记录
|
||||
if err := model.UpdateTaskExecution(ctx, execution); err != nil {
|
||||
logger.ErrorF(ctx, "[TaskExecutor] 更新执行记录失败 taskID=%s: %v", taskID, err)
|
||||
}
|
||||
|
||||
return execErr
|
||||
}
|
||||
|
||||
// generateTaskID 生成任务 ID
|
||||
func generateTaskID(taskType string, triggeredBy string) string {
|
||||
uniqueID := idgen.NextUint64ID()
|
||||
return fmt.Sprintf("%s_%s_%d", triggeredBy, taskType, uniqueID)
|
||||
}
|
||||
|
||||
// generateRetryTaskID 生成重试任务 ID
|
||||
func generateRetryTaskID(originalTaskID string, retryCount int) string {
|
||||
return fmt.Sprintf("retry_%d_%s", retryCount, originalTaskID)
|
||||
}
|
||||
@@ -0,0 +1,357 @@
|
||||
/*
|
||||
Copyright 2025-2026 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package task
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/hibiken/asynq"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/testhelper"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// mockHandler 用于测试的模拟任务处理器
|
||||
type mockHandler struct {
|
||||
executeFunc func(ctx context.Context, payload []byte) (*TaskResult, error)
|
||||
}
|
||||
|
||||
func (h *mockHandler) Execute(ctx context.Context, payload []byte) (*TaskResult, error) {
|
||||
if h.executeFunc != nil {
|
||||
return h.executeFunc(ctx, payload)
|
||||
}
|
||||
return &TaskResult{Message: "mock success"}, nil
|
||||
}
|
||||
|
||||
// successHandler 返回成功的处理器
|
||||
func successHandler() *mockHandler {
|
||||
return &mockHandler{
|
||||
executeFunc: func(ctx context.Context, payload []byte) (*TaskResult, error) {
|
||||
AppendLog(ctx, "执行成功,处理了 %d 条数据", 100)
|
||||
return &TaskResult{Message: "处理完成,共 100 条"}, nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// failHandler 返回失败的处理器
|
||||
func failHandler() *mockHandler {
|
||||
return &mockHandler{
|
||||
executeFunc: func(ctx context.Context, payload []byte) (*TaskResult, error) {
|
||||
AppendLog(ctx, "开始执行任务")
|
||||
return nil, fmt.Errorf("模拟执行失败: 数据库连接超时")
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
const testTaskType = "test:mock_task"
|
||||
|
||||
func setupTest(t *testing.T) func() {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
// 注册测试用 handler
|
||||
RegisterHandler(testTaskType, successHandler())
|
||||
return cleanup
|
||||
}
|
||||
|
||||
func TestRegisterAndGetHandler(t *testing.T) {
|
||||
_ = testTaskType
|
||||
cleanup := setupTest(t)
|
||||
defer cleanup()
|
||||
|
||||
// 验证 handler 已注册
|
||||
h, ok := getHandler(testTaskType)
|
||||
assert.True(t, ok, "handler should be registered")
|
||||
assert.NotNil(t, h)
|
||||
|
||||
// 未注册的 handler
|
||||
_, ok = getHandler("nonexistent")
|
||||
assert.False(t, ok, "non-existent handler should return false")
|
||||
}
|
||||
|
||||
func TestGetTaskIDFromContext(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
// 空 context
|
||||
taskID := GetTaskID(ctx)
|
||||
assert.Equal(t, "", taskID)
|
||||
|
||||
// 注入 taskID
|
||||
ctx = withTaskID(ctx, "test_task_123")
|
||||
taskID = GetTaskID(ctx)
|
||||
assert.Equal(t, "test_task_123", taskID)
|
||||
}
|
||||
|
||||
func TestAppendLogWithoutTaskID(t *testing.T) {
|
||||
cleanup := setupTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 没有 taskID 的 context,应降级到普通日志,不报错
|
||||
AppendLog(ctx, "这条日志应该降级处理,不会报错")
|
||||
}
|
||||
|
||||
func TestAppendLogWithTaskID(t *testing.T) {
|
||||
cleanup := setupTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 先创建一条执行记录
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "log_test_001",
|
||||
TaskType: testTaskType,
|
||||
TaskName: "测试任务",
|
||||
Status: model.TaskExecutionStatusRunning,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 注入 taskID 并追加日志
|
||||
ctx = withTaskID(ctx, "log_test_001")
|
||||
AppendLog(ctx, "第一条日志")
|
||||
AppendLog(ctx, "处理了 %d 条数据", 50)
|
||||
|
||||
// 验证日志
|
||||
found, err := model.GetTaskExecutionByTaskID(ctx, "log_test_001")
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, found.Log, "第一条日志")
|
||||
assert.Contains(t, found.Log, "处理了 50 条数据")
|
||||
}
|
||||
|
||||
func TestProcessTaskSuccess(t *testing.T) {
|
||||
cleanup := setupTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 注册成功 handler
|
||||
RegisterHandler(testTaskType, successHandler())
|
||||
|
||||
// 创建执行记录
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "process_success_001",
|
||||
TaskType: testTaskType,
|
||||
TaskName: "测试任务",
|
||||
Status: model.TaskExecutionStatusPending,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 模拟 Asynq Task
|
||||
asynqTask := asynq.NewTask(testTaskType, nil)
|
||||
// 使用 ResultWriter 设置 TaskID
|
||||
rw := asynq.NewResultWriter()
|
||||
rw.SetTaskID("process_success_001")
|
||||
// 通过 asynq 的 Task 不能直接设置 taskID,ProcessTask 通过 t.ResultWriter().TaskID() 获取
|
||||
// 但 asynq.Task 在没有经过 asynq server 的情况下 ResultWriter 可能为 nil
|
||||
// 我们需要在 ProcessTask 内部改用 taskID 注入的方式测试
|
||||
// 为了测试 ProcessTask,我们直接模拟调用 handler
|
||||
|
||||
// 直接通过 handler 测试
|
||||
handler, ok := getHandler(testTaskType)
|
||||
require.True(t, ok)
|
||||
|
||||
ctx = withTaskID(ctx, "process_success_001")
|
||||
result, err := handler.Execute(ctx, nil)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "处理完成,共 100 条", result.Message)
|
||||
|
||||
// 验证日志被追加
|
||||
found, err := model.GetTaskExecutionByTaskID(ctx, "process_success_001")
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, found.Log, "执行成功,处理了 100 条数据")
|
||||
}
|
||||
|
||||
func TestProcessTaskFailure(t *testing.T) {
|
||||
cleanup := setupTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 注册失败 handler
|
||||
RegisterHandler(testTaskType, failHandler())
|
||||
|
||||
// 创建执行记录
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "process_fail_001",
|
||||
TaskType: testTaskType,
|
||||
TaskName: "测试任务",
|
||||
Status: model.TaskExecutionStatusPending,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 直接调用 handler
|
||||
handler, ok := getHandler(testTaskType)
|
||||
require.True(t, ok)
|
||||
|
||||
ctx = withTaskID(ctx, "process_fail_001")
|
||||
_, err = handler.Execute(ctx, nil)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "模拟执行失败")
|
||||
|
||||
// 验证日志
|
||||
found, err := model.GetTaskExecutionByTaskID(ctx, "process_fail_001")
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, found.Log, "开始执行任务")
|
||||
}
|
||||
|
||||
func TestRetryTask(t *testing.T) {
|
||||
cleanup := setupTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建一条失败的执行记录(可重试)
|
||||
now := time.Now()
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "retry_test_001",
|
||||
TaskType: testTaskType,
|
||||
TaskName: "测试任务",
|
||||
Status: model.TaskExecutionStatusFailed,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
RetryCount: 0,
|
||||
ErrorMessage: "首次执行失败",
|
||||
StartedAt: &now,
|
||||
FinishedAt: &now,
|
||||
Duration: 100,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 重试
|
||||
newTaskID, err := RetryTask(ctx, execution.ID)
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, newTaskID)
|
||||
assert.Contains(t, newTaskID, "retry_1_")
|
||||
|
||||
// 验证新记录
|
||||
newExecution, err := model.GetTaskExecutionByTaskID(ctx, newTaskID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, model.TaskExecutionStatusPending, newExecution.Status)
|
||||
assert.Equal(t, 1, newExecution.RetryCount)
|
||||
assert.Equal(t, "retry", newExecution.TriggeredBy)
|
||||
assert.Equal(t, execution.TaskType, newExecution.TaskType)
|
||||
assert.True(t, newExecution.Retryable)
|
||||
|
||||
// 原记录不变
|
||||
original, err := model.GetTaskExecutionByID(ctx, execution.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, model.TaskExecutionStatusFailed, original.Status)
|
||||
assert.Equal(t, 0, original.RetryCount)
|
||||
}
|
||||
|
||||
func TestRetryTaskNotFailed(t *testing.T) {
|
||||
cleanup := setupTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建一条成功的记录
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "retry_not_failed_001",
|
||||
TaskType: testTaskType,
|
||||
TaskName: "测试任务",
|
||||
Status: model.TaskExecutionStatusSucceeded,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 尝试重试成功的任务
|
||||
_, err = RetryTask(ctx, execution.ID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "只有失败的任务才能重试")
|
||||
}
|
||||
|
||||
func TestRetryTaskNotRetryable(t *testing.T) {
|
||||
cleanup := setupTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "retry_not_allowed_001",
|
||||
TaskType: testTaskType,
|
||||
TaskName: "测试任务",
|
||||
Status: model.TaskExecutionStatusFailed,
|
||||
Retryable: false,
|
||||
MaxRetry: 0,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = RetryTask(ctx, execution.ID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "不支持重试")
|
||||
}
|
||||
|
||||
func TestRetryTaskMaxRetryExceeded(t *testing.T) {
|
||||
cleanup := setupTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "retry_max_001",
|
||||
TaskType: testTaskType,
|
||||
TaskName: "测试任务",
|
||||
Status: model.TaskExecutionStatusFailed,
|
||||
Retryable: true,
|
||||
MaxRetry: 2,
|
||||
RetryCount: 2,
|
||||
TriggeredBy: "retry",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = RetryTask(ctx, execution.ID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "已达到最大重试次数")
|
||||
}
|
||||
|
||||
func TestRetryTaskNonExistent(t *testing.T) {
|
||||
cleanup := setupTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := RetryTask(ctx, 99999999)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "不存在")
|
||||
}
|
||||
|
||||
func TestGenerateTaskID(t *testing.T) {
|
||||
id1 := generateTaskID("test_type", "manual")
|
||||
id2 := generateTaskID("test_type", "manual")
|
||||
|
||||
// 两个 ID 应不同(包含 Snowflake ID)
|
||||
assert.NotEqual(t, id1, id2)
|
||||
assert.Contains(t, id1, "manual_test_type_")
|
||||
}
|
||||
|
||||
func TestGenerateRetryTaskID(t *testing.T) {
|
||||
id := generateRetryTaskID("original_task_123", 2)
|
||||
assert.Equal(t, "retry_2_original_task_123", id)
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
/*
|
||||
Copyright 2025-2026 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package task
|
||||
|
||||
import "context"
|
||||
|
||||
// TaskResult 任务执行结果
|
||||
type TaskResult struct {
|
||||
Message string // 结果摘要,如 "共清理 120 个文件,耗时 3.2s"
|
||||
Detail string // 可选的详细结果 JSON
|
||||
}
|
||||
|
||||
// TaskHandler 异步任务处理器接口
|
||||
// 所有异步任务必须实现此接口,框架将自动管理任务执行记录的创建、状态流转和日志写入。
|
||||
//
|
||||
// 开发者只需实现 Execute 方法编写业务逻辑,在方法内通过 task.AppendLog(ctx, ...) 追加执行日志。
|
||||
// 任务的创建、状态更新、错误记录、重试计数全部由框架透明处理。
|
||||
type TaskHandler interface {
|
||||
// Execute 执行任务业务逻辑
|
||||
// - ctx: 已注入 Trace Span 和 taskID 的上下文
|
||||
// - payload: 调度时传入的原始参数(可为 nil)
|
||||
// - 返回 TaskResult 描述执行结果,或 error 表示执行失败
|
||||
Execute(ctx context.Context, payload []byte) (*TaskResult, error)
|
||||
}
|
||||
@@ -28,13 +28,17 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
AsynqClient *asynq.Client
|
||||
scheduler *asynq.Scheduler
|
||||
schedulerOnce sync.Once
|
||||
)
|
||||
|
||||
func init() {
|
||||
AsynqClient = asynq.NewClient(task.RedisOpt)
|
||||
// AsynqClient 已在 task 包中初始化
|
||||
}
|
||||
|
||||
// GetAsynqClient 获取全局 AsynqClient
|
||||
func GetAsynqClient() *asynq.Client {
|
||||
return task.AsynqClient
|
||||
}
|
||||
|
||||
// StartScheduler 启动调度器
|
||||
|
||||
@@ -24,8 +24,12 @@ import (
|
||||
// RedisOpt asynq Redis 连接配置(兼容 Standalone/Sentinel/Cluster)
|
||||
var RedisOpt asynq.RedisConnOpt
|
||||
|
||||
// AsynqClient asynq 客户端,用于任务入队
|
||||
var AsynqClient *asynq.Client
|
||||
|
||||
func init() {
|
||||
RedisOpt = NewRedisConnOpt()
|
||||
AsynqClient = asynq.NewClient(RedisOpt)
|
||||
}
|
||||
|
||||
// NewRedisConnOpt 根据配置返回对应的 asynq Redis 连接选项
|
||||
|
||||
@@ -18,67 +18,15 @@ package worker
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/linux-do/credit/internal/logger"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
|
||||
"github.com/hibiken/asynq"
|
||||
"github.com/linux-do/credit/internal/otel_trace"
|
||||
"go.opentelemetry.io/otel/attribute"
|
||||
"go.opentelemetry.io/otel/codes"
|
||||
)
|
||||
|
||||
// taskLoggingMiddleware 记录任务日志中间件
|
||||
// taskLoggingMiddleware 任务处理中间件
|
||||
// 注意:OTel Span 创建、日志记录、TaskExecution 状态管理
|
||||
// 已由 task.ProcessTask 统一处理,此中间件保留用于未来扩展(如限流、监控等)
|
||||
func taskLoggingMiddleware(h asynq.Handler) asynq.Handler {
|
||||
return asynq.HandlerFunc(func(ctx context.Context, t *asynq.Task) error {
|
||||
// 初始化 Trace
|
||||
ctx, span := otel_trace.Start(ctx, "TaskProcess_"+t.Type(), trace.WithSpanKind(trace.SpanKindConsumer))
|
||||
defer span.End()
|
||||
|
||||
// 添加任务信息到 Span
|
||||
span.SetAttributes(
|
||||
attribute.String("task.type", t.Type()),
|
||||
attribute.Int("task.payload_size", len(t.Payload())),
|
||||
attribute.String("task.id", t.ResultWriter().TaskID()),
|
||||
)
|
||||
|
||||
// 开始计时
|
||||
start := time.Now()
|
||||
|
||||
// 处理任务
|
||||
err := h.ProcessTask(ctx, t)
|
||||
|
||||
// 计算耗时
|
||||
latency := time.Since(start)
|
||||
|
||||
if err != nil {
|
||||
// 处理出错,记录错误日志
|
||||
logger.ErrorF(
|
||||
ctx,
|
||||
"[TaskMiddleware] 任务处理失败 Type: %s\nStartTime: %s\nLatency: %d ms\nError: %v",
|
||||
t.Type(),
|
||||
start.Format(time.RFC3339),
|
||||
latency.Milliseconds(),
|
||||
err,
|
||||
)
|
||||
|
||||
// 设置 Span 错误状态
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
span.RecordError(err)
|
||||
return err
|
||||
}
|
||||
|
||||
// 处理成功,记录成功日志
|
||||
logger.InfoF(
|
||||
ctx,
|
||||
"[TaskMiddleware] 任务处理完成 Type: %s\nStartTime: %s\nEndTime: %s\nLatency: %d ms",
|
||||
t.Type(),
|
||||
start.Format(time.RFC3339),
|
||||
time.Now().Format(time.RFC3339),
|
||||
latency.Milliseconds(),
|
||||
)
|
||||
|
||||
return nil
|
||||
return h.ProcessTask(ctx, t)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
Copyright 2025-2026 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
@@ -25,6 +25,11 @@ import (
|
||||
"github.com/linux-do/credit/internal/task"
|
||||
)
|
||||
|
||||
func init() {
|
||||
// 注册所有任务处理器
|
||||
task.RegisterHandler(task.CleanupUnusedUploadsTask, &upload.CleanupUnusedUploadsHandler{})
|
||||
}
|
||||
|
||||
// StartWorker 启动任务处理服务器
|
||||
func StartWorker() error {
|
||||
asynqServer := asynq.NewServer(
|
||||
@@ -37,10 +42,13 @@ func StartWorker() error {
|
||||
},
|
||||
)
|
||||
|
||||
// 注册任务处理器
|
||||
// 注册 Asynq 任务路由
|
||||
mux := asynq.NewServeMux()
|
||||
mux.Use(taskLoggingMiddleware)
|
||||
mux.HandleFunc(task.CleanupUnusedUploadsTask, upload.HandleCleanupUnusedUploads)
|
||||
|
||||
// 统一使用 task.ProcessTask 处理所有任务类型
|
||||
// 框架内部自动分发到对应的 TaskHandler 实现
|
||||
mux.HandleFunc(task.CleanupUnusedUploadsTask, task.ProcessTask)
|
||||
|
||||
// 启动服务器
|
||||
return asynqServer.Run(mux)
|
||||
|
||||
@@ -25,7 +25,7 @@ import (
|
||||
"github.com/hibiken/asynq"
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/task/scheduler"
|
||||
"github.com/linux-do/credit/internal/task"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
@@ -48,6 +48,7 @@ func SetupTestEnvironment(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func())
|
||||
&model.ExternalAccount{},
|
||||
&model.SystemConfig{},
|
||||
&model.Upload{},
|
||||
&model.TaskExecution{},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to auto migrate tables: %v", err)
|
||||
@@ -69,7 +70,7 @@ func SetupTestEnvironment(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func())
|
||||
db.Redis = redisClient
|
||||
|
||||
// Hook up AsynqClient to miniredis
|
||||
scheduler.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{
|
||||
task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{
|
||||
Addr: mr.Addr(),
|
||||
})
|
||||
|
||||
@@ -83,7 +84,7 @@ func SetupTestEnvironment(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func())
|
||||
// Reset database and Redis references
|
||||
db.SetDB(nil)
|
||||
db.Redis = nil
|
||||
scheduler.AsynqClient = nil
|
||||
task.AsynqClient = nil
|
||||
}
|
||||
|
||||
return sqliteDB, mr, cleanup
|
||||
|
||||
Reference in New Issue
Block a user