235 lines
7.2 KiB
Go
235 lines
7.2 KiB
Go
package service
|
|
|
|
import (
|
|
"context"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/QuantumNous/new-api/common"
|
|
"github.com/QuantumNous/new-api/model"
|
|
|
|
"github.com/stretchr/testify/assert"
|
|
"github.com/stretchr/testify/require"
|
|
)
|
|
|
|
// withSystemTaskRegistry swaps the package registry for the given handlers for
|
|
// the duration of a test and restores the original registry afterward.
|
|
func withSystemTaskRegistry(t *testing.T, handlers ...SystemTaskHandler) {
|
|
t.Helper()
|
|
systemTaskHandlersMu.Lock()
|
|
saved := systemTaskHandlers
|
|
systemTaskHandlers = map[string]SystemTaskHandler{}
|
|
for _, h := range handlers {
|
|
systemTaskHandlers[h.Type()] = h
|
|
}
|
|
systemTaskHandlersMu.Unlock()
|
|
t.Cleanup(func() {
|
|
systemTaskHandlersMu.Lock()
|
|
systemTaskHandlers = saved
|
|
systemTaskHandlersMu.Unlock()
|
|
})
|
|
}
|
|
|
|
type stubScheduledHandler struct {
|
|
taskType string
|
|
enabled bool
|
|
interval time.Duration
|
|
onRun func(ctx context.Context, task *model.SystemTask, runnerID string)
|
|
}
|
|
|
|
type stubSystemTaskRunResult struct {
|
|
taskID string
|
|
taskType string
|
|
err error
|
|
}
|
|
|
|
func (h *stubScheduledHandler) Type() string { return h.taskType }
|
|
|
|
func (h *stubScheduledHandler) Run(ctx context.Context, task *model.SystemTask, runnerID string) {
|
|
if h.onRun != nil {
|
|
h.onRun(ctx, task, runnerID)
|
|
}
|
|
}
|
|
|
|
func (h *stubScheduledHandler) Enabled() bool { return h.enabled }
|
|
func (h *stubScheduledHandler) Interval() time.Duration { return h.interval }
|
|
func (h *stubScheduledHandler) NewPayload() any { return nil }
|
|
|
|
func countSystemTasks(t *testing.T, taskType string) int64 {
|
|
t.Helper()
|
|
var count int64
|
|
require.NoError(t, model.DB.Model(&model.SystemTask{}).Where("type = ?", taskType).Count(&count).Error)
|
|
return count
|
|
}
|
|
|
|
func TestSystemTaskSchedulerCreatesWhenDueAndDedups(t *testing.T) {
|
|
truncate(t)
|
|
|
|
handler := &stubScheduledHandler{taskType: "test_scheduled", enabled: true, interval: time.Minute}
|
|
withSystemTaskRegistry(t, handler)
|
|
|
|
runSystemTaskScheduler()
|
|
require.Equal(t, int64(1), countSystemTasks(t, handler.taskType))
|
|
|
|
// An active (pending) row already exists, so a second pass must not create
|
|
// another row.
|
|
runSystemTaskScheduler()
|
|
require.Equal(t, int64(1), countSystemTasks(t, handler.taskType))
|
|
|
|
// Finish the run; with a fresh updated_at the next run is not due yet.
|
|
latest, err := model.GetLatestSystemTask(handler.taskType)
|
|
require.NoError(t, err)
|
|
require.NotNil(t, latest)
|
|
_, claimed, err := model.ClaimSystemTask(latest.ID, handler.taskType, "runner-a", common.GetTimestamp()+60)
|
|
require.NoError(t, err)
|
|
require.True(t, claimed)
|
|
require.NoError(t, model.FinishSystemTask(latest.TaskID, "runner-a", model.SystemTaskStatusSucceeded, nil, ""))
|
|
|
|
runSystemTaskScheduler()
|
|
require.Equal(t, int64(1), countSystemTasks(t, handler.taskType))
|
|
|
|
// Backdate the finished row beyond the interval -> the job becomes due again.
|
|
require.NoError(t, model.DB.Model(&model.SystemTask{}).
|
|
Where("task_id = ?", latest.TaskID).
|
|
Update("updated_at", common.GetTimestamp()-120).Error)
|
|
|
|
runSystemTaskScheduler()
|
|
require.Equal(t, int64(2), countSystemTasks(t, handler.taskType))
|
|
}
|
|
|
|
func TestSystemTaskSchedulerSkipsDisabled(t *testing.T) {
|
|
truncate(t)
|
|
|
|
handler := &stubScheduledHandler{taskType: "test_disabled", enabled: false, interval: time.Minute}
|
|
withSystemTaskRegistry(t, handler)
|
|
|
|
runSystemTaskScheduler()
|
|
assert.Equal(t, int64(0), countSystemTasks(t, handler.taskType))
|
|
}
|
|
|
|
func TestSystemTaskClaimPassDispatchesByType(t *testing.T) {
|
|
truncate(t)
|
|
|
|
ran := make(chan stubSystemTaskRunResult, 1)
|
|
handler := &stubScheduledHandler{
|
|
taskType: "test_dispatch",
|
|
enabled: true,
|
|
interval: time.Minute,
|
|
onRun: func(_ context.Context, task *model.SystemTask, runnerID string) {
|
|
ran <- stubSystemTaskRunResult{
|
|
taskType: task.Type,
|
|
err: model.FinishSystemTask(task.TaskID, runnerID, model.SystemTaskStatusSucceeded, nil, ""),
|
|
}
|
|
},
|
|
}
|
|
withSystemTaskRegistry(t, handler)
|
|
|
|
_, err := model.CreateSystemTask(handler.taskType, nil, nil)
|
|
require.NoError(t, err)
|
|
|
|
runSystemTaskClaimPass("runner-dispatch")
|
|
|
|
select {
|
|
case got := <-ran:
|
|
require.NoError(t, got.err)
|
|
assert.Equal(t, handler.taskType, got.taskType)
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("claimed task was not dispatched to its handler")
|
|
}
|
|
|
|
require.Eventually(t, func() bool {
|
|
latest, err := model.GetLatestSystemTask(handler.taskType)
|
|
return err == nil && latest != nil && latest.Status == model.SystemTaskStatusSucceeded
|
|
}, 2*time.Second, 20*time.Millisecond)
|
|
}
|
|
|
|
func TestSystemTaskClaimPassDispatchesEarliestPendingByType(t *testing.T) {
|
|
truncate(t)
|
|
|
|
ran := make(chan stubSystemTaskRunResult, 2)
|
|
handlerA := &stubScheduledHandler{
|
|
taskType: "test_dispatch_a",
|
|
enabled: true,
|
|
interval: time.Minute,
|
|
onRun: func(_ context.Context, task *model.SystemTask, runnerID string) {
|
|
ran <- stubSystemTaskRunResult{
|
|
taskID: task.TaskID,
|
|
err: model.FinishSystemTask(task.TaskID, runnerID, model.SystemTaskStatusSucceeded, nil, ""),
|
|
}
|
|
},
|
|
}
|
|
handlerB := &stubScheduledHandler{
|
|
taskType: "test_dispatch_b",
|
|
enabled: true,
|
|
interval: time.Minute,
|
|
onRun: func(_ context.Context, task *model.SystemTask, runnerID string) {
|
|
ran <- stubSystemTaskRunResult{
|
|
taskID: task.TaskID,
|
|
err: model.FinishSystemTask(task.TaskID, runnerID, model.SystemTaskStatusSucceeded, nil, ""),
|
|
}
|
|
},
|
|
}
|
|
withSystemTaskRegistry(t, handlerA, handlerB)
|
|
|
|
firstA, err := model.CreateSystemTask(handlerA.taskType, nil, nil)
|
|
require.NoError(t, err)
|
|
secondTaskID, err := model.GenerateSystemTaskID()
|
|
require.NoError(t, err)
|
|
secondA := &model.SystemTask{
|
|
TaskID: secondTaskID,
|
|
Type: handlerA.taskType,
|
|
Status: model.SystemTaskStatusPending,
|
|
}
|
|
require.NoError(t, model.DB.Create(secondA).Error)
|
|
firstB, err := model.CreateSystemTask(handlerB.taskType, nil, nil)
|
|
require.NoError(t, err)
|
|
|
|
runSystemTaskClaimPass("runner-dispatch")
|
|
|
|
got := map[string]bool{}
|
|
for range 2 {
|
|
select {
|
|
case result := <-ran:
|
|
require.NoError(t, result.err)
|
|
got[result.taskID] = true
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("claimed tasks were not dispatched to their handlers")
|
|
}
|
|
}
|
|
|
|
assert.True(t, got[firstA.TaskID])
|
|
assert.True(t, got[firstB.TaskID])
|
|
assert.False(t, got[secondA.TaskID])
|
|
|
|
require.Eventually(t, func() bool {
|
|
reloaded, err := model.GetSystemTaskByTaskID(secondA.TaskID)
|
|
return err == nil && reloaded != nil && reloaded.Status == model.SystemTaskStatusPending
|
|
}, 2*time.Second, 20*time.Millisecond)
|
|
}
|
|
|
|
func TestEnqueueSystemTaskReportsCreatedAndExistingActive(t *testing.T) {
|
|
truncate(t)
|
|
|
|
first, created, err := EnqueueSystemTask("test_enqueue", map[string]bool{"manual": true})
|
|
require.NoError(t, err)
|
|
require.True(t, created)
|
|
require.NotNil(t, first)
|
|
|
|
existing, created, err := EnqueueSystemTask("test_enqueue", nil)
|
|
require.NoError(t, err)
|
|
require.False(t, created)
|
|
require.NotNil(t, existing)
|
|
assert.Equal(t, first.TaskID, existing.TaskID)
|
|
|
|
_, claimed, err := model.ClaimSystemTask(first.ID, first.Type, "runner-a", common.GetTimestamp()+60)
|
|
require.NoError(t, err)
|
|
require.True(t, claimed)
|
|
require.NoError(t, model.FinishSystemTask(first.TaskID, "runner-a", model.SystemTaskStatusSucceeded, nil, ""))
|
|
|
|
second, created, err := EnqueueSystemTask("test_enqueue", nil)
|
|
require.NoError(t, err)
|
|
require.True(t, created)
|
|
require.NotNil(t, second)
|
|
assert.NotEqual(t, first.TaskID, second.TaskID)
|
|
}
|