334 lines
9.9 KiB
Go
334 lines
9.9 KiB
Go
package task
|
||
|
||
import (
|
||
"context"
|
||
"errors"
|
||
"reflect"
|
||
"sync"
|
||
"sync/atomic"
|
||
"testing"
|
||
"time"
|
||
)
|
||
|
||
func Test正常跑完N个商品(t *testing.T) {
|
||
var mu sync.Mutex
|
||
var doneValues []int
|
||
runner := newTestRunner(Options{
|
||
OnProgress: func(progress Progress) {
|
||
mu.Lock()
|
||
doneValues = append(doneValues, progress.Done)
|
||
mu.Unlock()
|
||
},
|
||
})
|
||
if err := runner.Start(context.Background(), []string{"1", "2", "3"}); err != nil {
|
||
t.Fatalf("启动任务失败:%v", err)
|
||
}
|
||
progress := waitForState(t, runner, StateCompleted)
|
||
if progress.Total != 3 || progress.Done != 3 {
|
||
t.Fatalf("完成进度错误:%+v", progress)
|
||
}
|
||
mu.Lock()
|
||
defer mu.Unlock()
|
||
if !reflect.DeepEqual(doneValues, []int{1, 2, 3}) {
|
||
t.Fatalf("进度没有从 0 依次推进到 N:%v", doneValues)
|
||
}
|
||
}
|
||
|
||
func Test中途停止后不再启动剩余商品(t *testing.T) {
|
||
started := make(chan string, 3)
|
||
release := make(chan struct{})
|
||
runner := newTestRunner(Options{
|
||
Prepare: func(_ context.Context, product Product) (Work, error) {
|
||
started <- product.ID
|
||
if product.ID == "1" {
|
||
<-release
|
||
}
|
||
return Work{}, nil
|
||
},
|
||
})
|
||
if err := runner.Start(context.Background(), []string{"1", "2", "3"}); err != nil {
|
||
t.Fatalf("启动任务失败:%v", err)
|
||
}
|
||
if got := <-started; got != "1" {
|
||
t.Fatalf("首个商品错误:%s", got)
|
||
}
|
||
runner.Stop()
|
||
close(release)
|
||
progress := waitForState(t, runner, StateStopped)
|
||
if progress.Done != 1 {
|
||
t.Fatalf("当前商品应处理完再停止:%+v", progress)
|
||
}
|
||
select {
|
||
case id := <-started:
|
||
t.Fatalf("停止后仍启动了商品:%s", id)
|
||
default:
|
||
}
|
||
}
|
||
|
||
func Test登录失效立即触发全局停止门(t *testing.T) {
|
||
var processed []string
|
||
var mu sync.Mutex
|
||
runner := newTestRunner(Options{
|
||
Prepare: func(_ context.Context, product Product) (Work, error) {
|
||
mu.Lock()
|
||
processed = append(processed, product.ID)
|
||
mu.Unlock()
|
||
if product.ID == "2" {
|
||
return Work{}, errors.Join(ErrLoginRequired, errors.New("深度检查不通过"))
|
||
}
|
||
return Work{}, nil
|
||
},
|
||
})
|
||
if err := runner.Start(context.Background(), []string{"1", "2", "3", "4"}); err != nil {
|
||
t.Fatalf("启动任务失败:%v", err)
|
||
}
|
||
progress := waitForState(t, runner, StateLoginRequired)
|
||
mu.Lock()
|
||
defer mu.Unlock()
|
||
if !reflect.DeepEqual(processed, []string{"1", "2"}) {
|
||
t.Fatalf("登录失效后仍处理了后续商品:%v", processed)
|
||
}
|
||
if progress.Done != 1 || progress.Failed != 0 {
|
||
t.Fatalf("登录失效商品不能记为单商品失败:%+v", progress)
|
||
}
|
||
}
|
||
|
||
func Test疑似风控触发停止门且后续商品不处理(t *testing.T) {
|
||
var processed []string
|
||
runner := newTestRunner(Options{Prepare: func(_ context.Context, product Product) (Work, error) {
|
||
processed = append(processed, product.ID)
|
||
if product.ID == "2" {
|
||
return Work{}, errors.Join(ErrLoginRequired, ErrRiskSuspected, RiskSuspectedError{Consecutive: 8})
|
||
}
|
||
return Work{}, nil
|
||
}})
|
||
if err := runner.Start(context.Background(), []string{"1", "2", "3"}); err != nil {
|
||
t.Fatalf("启动任务失败:%v", err)
|
||
}
|
||
progress := waitForState(t, runner, StateLoginRequired)
|
||
if !reflect.DeepEqual(processed, []string{"1", "2"}) {
|
||
t.Fatalf("疑似风控后仍处理后续商品:%v", processed)
|
||
}
|
||
if progress.StopReason != StopReasonRiskSuspected {
|
||
t.Fatalf("停止原因应为疑似风控,实际:%+v", progress)
|
||
}
|
||
if progress.RiskEmptyCount != 8 {
|
||
t.Fatalf("连续空结果数应传给前端,实际:%+v", progress)
|
||
}
|
||
}
|
||
|
||
func Test空结果中途发现视频会清零(t *testing.T) {
|
||
guard := NewEmptyRiskGuard(3)
|
||
if guard.Record(0) || guard.Record(0) {
|
||
t.Fatal("阈值前不应触发")
|
||
}
|
||
if guard.Record(1) || guard.Consecutive() != 0 {
|
||
t.Fatalf("发现视频应清零,实际连续数 %d", guard.Consecutive())
|
||
}
|
||
if guard.Record(0) || guard.Record(0) || !guard.Record(0) {
|
||
t.Fatal("清零后应重新累计到阈值")
|
||
}
|
||
}
|
||
|
||
func Test明确访问受限立即停批且不计失败(t *testing.T) {
|
||
processed := []string{}
|
||
runner := newTestRunner(Options{Prepare: func(_ context.Context, product Product) (Work, error) {
|
||
processed = append(processed, product.ID)
|
||
if product.ID == "2" {
|
||
return Work{}, ErrRiskBlocked
|
||
}
|
||
return Work{}, nil
|
||
}})
|
||
if err := runner.Start(context.Background(), []string{"1", "2", "3"}); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
progress := waitForState(t, runner, StateLoginRequired)
|
||
if !reflect.DeepEqual(processed, []string{"1", "2"}) || progress.Done != 1 || progress.Failed != 0 || progress.StopReason != StopReasonRiskBlocked {
|
||
t.Fatalf("incorrect risk stop: %+v %v", progress, processed)
|
||
}
|
||
}
|
||
|
||
func Test停止取消准备等待但不把当前商品记失败(t *testing.T) {
|
||
started := make(chan struct{})
|
||
runner := newTestRunner(Options{Prepare: func(ctx context.Context, _ Product) (Work, error) {
|
||
close(started)
|
||
<-ctx.Done()
|
||
return Work{}, ctx.Err()
|
||
}})
|
||
if err := runner.Start(context.Background(), []string{"1", "2"}); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
<-started
|
||
runner.Stop()
|
||
progress := waitForState(t, runner, StateStopped)
|
||
if progress.Done != 0 || progress.Failed != 0 {
|
||
t.Fatalf("cancelled preparation counted as failed: %+v", progress)
|
||
}
|
||
}
|
||
|
||
func Test单商品失败后继续处理后续商品(t *testing.T) {
|
||
var processed []string
|
||
var mu sync.Mutex
|
||
runner := newTestRunner(Options{
|
||
Prepare: func(_ context.Context, product Product) (Work, error) {
|
||
mu.Lock()
|
||
processed = append(processed, product.ID)
|
||
mu.Unlock()
|
||
if product.ID == "2" {
|
||
return Work{}, errors.New("图搜失败")
|
||
}
|
||
return Work{}, nil
|
||
},
|
||
})
|
||
if err := runner.Start(context.Background(), []string{"1", "2", "3"}); err != nil {
|
||
t.Fatalf("启动任务失败:%v", err)
|
||
}
|
||
progress := waitForState(t, runner, StateCompleted)
|
||
mu.Lock()
|
||
defer mu.Unlock()
|
||
if !reflect.DeepEqual(processed, []string{"1", "2", "3"}) {
|
||
t.Fatalf("单商品失败后没有继续:%v", processed)
|
||
}
|
||
if progress.Failed != 1 || progress.Done != 3 {
|
||
t.Fatalf("失败计数错误:%+v", progress)
|
||
}
|
||
}
|
||
|
||
func Test已有任务运行时拒绝再次启动(t *testing.T) {
|
||
release := make(chan struct{})
|
||
started := make(chan struct{})
|
||
runner := newTestRunner(Options{
|
||
Prepare: func(context.Context, Product) (Work, error) {
|
||
close(started)
|
||
<-release
|
||
return Work{}, nil
|
||
},
|
||
})
|
||
if err := runner.Start(context.Background(), []string{"1"}); err != nil {
|
||
t.Fatalf("首次启动失败:%v", err)
|
||
}
|
||
<-started
|
||
if err := runner.Start(context.Background(), []string{"2"}); err == nil || err.Error() != "已有任务在运行" {
|
||
t.Fatalf("重复启动应返回指定中文错误,实际:%v", err)
|
||
}
|
||
close(release)
|
||
waitForState(t, runner, StateCompleted)
|
||
}
|
||
|
||
func Test已完成和无同款商品不进入队列(t *testing.T) {
|
||
statuses := map[string]Product{
|
||
"done": {ID: "done", DownloadStatus: "done"},
|
||
"none": {ID: "none", VideoStatus: "none"},
|
||
"todo": {ID: "todo"},
|
||
}
|
||
var processed []string
|
||
var mu sync.Mutex
|
||
runner := newTestRunner(Options{
|
||
Load: func(_ context.Context, id string) (Product, error) { return statuses[id], nil },
|
||
Prepare: func(_ context.Context, product Product) (Work, error) {
|
||
mu.Lock()
|
||
processed = append(processed, product.ID)
|
||
mu.Unlock()
|
||
return Work{}, nil
|
||
},
|
||
})
|
||
if err := runner.Start(context.Background(), []string{"done", "none", "todo"}); err != nil {
|
||
t.Fatalf("启动任务失败:%v", err)
|
||
}
|
||
progress := waitForState(t, runner, StateCompleted)
|
||
mu.Lock()
|
||
defer mu.Unlock()
|
||
if !reflect.DeepEqual(processed, []string{"todo"}) {
|
||
t.Fatalf("断点跳过逻辑错误:%v", processed)
|
||
}
|
||
if progress.Total != 1 || progress.Done != 1 {
|
||
t.Fatalf("跳过后的总数错误:%+v", progress)
|
||
}
|
||
}
|
||
|
||
func Test下载并发数不超过配置值(t *testing.T) {
|
||
var active int32
|
||
var maximum int32
|
||
jobs := make([]DownloadFunc, 8)
|
||
for index := range jobs {
|
||
jobs[index] = func(context.Context) error {
|
||
current := atomic.AddInt32(&active, 1)
|
||
for {
|
||
old := atomic.LoadInt32(&maximum)
|
||
if current <= old || atomic.CompareAndSwapInt32(&maximum, old, current) {
|
||
break
|
||
}
|
||
}
|
||
time.Sleep(15 * time.Millisecond)
|
||
atomic.AddInt32(&active, -1)
|
||
return nil
|
||
}
|
||
}
|
||
runner := newTestRunner(Options{
|
||
Concurrency: 2,
|
||
Prepare: func(context.Context, Product) (Work, error) {
|
||
return Work{Downloads: jobs}, nil
|
||
},
|
||
})
|
||
if err := runner.Start(context.Background(), []string{"1"}); err != nil {
|
||
t.Fatalf("启动任务失败:%v", err)
|
||
}
|
||
progress := waitForState(t, runner, StateCompleted)
|
||
if maximum != 2 {
|
||
t.Fatalf("下载同时运行数应为 2,实际为 %d", maximum)
|
||
}
|
||
if progress.Downloaded != len(jobs) {
|
||
t.Fatalf("下载计数错误:%+v", progress)
|
||
}
|
||
}
|
||
|
||
func Test每处理完一个商品调用一次进度回调(t *testing.T) {
|
||
var calls int32
|
||
runner := newTestRunner(Options{
|
||
OnProgress: func(Progress) { atomic.AddInt32(&calls, 1) },
|
||
})
|
||
if err := runner.Start(context.Background(), []string{"1", "2", "3", "4"}); err != nil {
|
||
t.Fatalf("启动任务失败:%v", err)
|
||
}
|
||
waitForState(t, runner, StateCompleted)
|
||
if calls != 4 {
|
||
t.Fatalf("进度回调次数应等于处理商品数,实际为 %d", calls)
|
||
}
|
||
}
|
||
|
||
func newTestRunner(override Options) *Runner {
|
||
options := Options{
|
||
Concurrency: 1,
|
||
Load: func(_ context.Context, id string) (Product, error) {
|
||
return Product{ID: id, Name: "商品" + id}, nil
|
||
},
|
||
Prepare: func(context.Context, Product) (Work, error) { return Work{}, nil },
|
||
}
|
||
if override.Load != nil {
|
||
options.Load = override.Load
|
||
}
|
||
if override.Prepare != nil {
|
||
options.Prepare = override.Prepare
|
||
}
|
||
if override.Concurrency > 0 {
|
||
options.Concurrency = override.Concurrency
|
||
}
|
||
options.OnProgress = override.OnProgress
|
||
options.OnFinished = override.OnFinished
|
||
return NewRunner(options)
|
||
}
|
||
|
||
func waitForState(t *testing.T, runner *Runner, state string) Progress {
|
||
t.Helper()
|
||
deadline := time.Now().Add(2 * time.Second)
|
||
for time.Now().Before(deadline) {
|
||
progress := runner.Snapshot()
|
||
if progress.State == state {
|
||
return progress
|
||
}
|
||
time.Sleep(time.Millisecond)
|
||
}
|
||
t.Fatalf("等待任务状态 %s 超时,当前:%+v", state, runner.Snapshot())
|
||
return Progress{}
|
||
}
|