149 lines
3.6 KiB
Go
149 lines
3.6 KiB
Go
package runners
|
||
|
||
import (
|
||
"context"
|
||
"fmt"
|
||
"math/rand"
|
||
"remote-task-excutor-cli/pkg/models"
|
||
"remote-task-excutor-cli/pkg/service"
|
||
"time"
|
||
)
|
||
|
||
type taskRunner struct {
|
||
authService service.AuthService
|
||
preparationService service.PreparationService
|
||
taskRunService service.TaskService
|
||
}
|
||
|
||
func NewTaskRunner(authService service.AuthService, preparationService service.PreparationService, taskRunService service.TaskService) *taskRunner {
|
||
return &taskRunner{
|
||
authService: authService,
|
||
preparationService: preparationService,
|
||
taskRunService: taskRunService,
|
||
}
|
||
}
|
||
|
||
func (r *taskRunner) RunTask(ctx context.Context, config *models.RunConfig) error {
|
||
authData, err := r.authService.GetToken(ctx)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
fmt.Println("获取token成功")
|
||
clusterID, err := r.authService.GetClusterID(ctx, "openI")
|
||
if err != nil {
|
||
return err
|
||
}
|
||
fmt.Println("获取clusterID成功:", clusterID)
|
||
// 数据准备
|
||
bindResultSet, err := r.preparationService.PrepareAll(ctx, authData, config, clusterID)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
// 任务提交
|
||
jobSetID, err := r.taskRunService.SubmitTask(ctx, authData, config, clusterID, bindResultSet)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
fmt.Println("任务提交成功,开始查询任务状态")
|
||
time.Sleep(5 * time.Second)
|
||
go r.RecordLogFile(ctx, authData, jobSetID)
|
||
// 轮询任务状态
|
||
resp, err := r.PollTaskStatusWithBackoff(ctx, authData, jobSetID)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
|
||
if resp.Data.SubTaskInfos[0].Status != "Completed" {
|
||
return fmt.Errorf("远程任务执行失败:")
|
||
}
|
||
// 获取任务结果
|
||
err = r.taskRunService.GetTaskResult(ctx, authData, config.TaskOutput, jobSetID)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
func (r *taskRunner) RecordLogFile(ctx context.Context, authData *models.AuthData, jobSetID string) error {
|
||
for {
|
||
content, _ := r.taskRunService.GetTaskLogs(ctx, authData, jobSetID)
|
||
if len(content) == 0 {
|
||
time.Sleep(5 * time.Second)
|
||
}
|
||
fmt.Println("get log content is ", content)
|
||
}
|
||
}
|
||
|
||
func (r *taskRunner) PollTaskStatusWithBackoff(ctx context.Context, authData *models.AuthData, jobSetID string) (*models.TaskDetailResponse, error) {
|
||
const (
|
||
baseInterval = 5 * time.Second
|
||
maxInterval = 1 * time.Minute
|
||
maxAttempts = 5 // 最多尝试10次
|
||
)
|
||
|
||
var (
|
||
backoff = baseInterval
|
||
attemptCount = 0
|
||
)
|
||
|
||
for {
|
||
if attemptCount > maxAttempts {
|
||
return nil, fmt.Errorf("超出最大查询次数(%d)", maxAttempts)
|
||
}
|
||
|
||
// 检查上下文是否被取消
|
||
if ctx.Err() != nil {
|
||
return nil, ctx.Err()
|
||
}
|
||
|
||
// 获取任务状态
|
||
status, err := r.taskRunService.GetTaskStatus(ctx, authData, jobSetID)
|
||
if err != nil {
|
||
attemptCount++
|
||
// 指数回退
|
||
backoff = min(backoff*2, maxInterval)
|
||
sleepTime := backoff + time.Duration(rand.Int63n(int64(backoff/2))) // 随机抖动避免同步
|
||
fmt.Printf("任务查询失败(%s), 将在 %s 后重试\n", err, sleepTime)
|
||
select {
|
||
case <-ctx.Done():
|
||
return nil, ctx.Err()
|
||
case <-time.After(sleepTime):
|
||
continue
|
||
}
|
||
}
|
||
|
||
// 重置回退间隔
|
||
backoff = baseInterval
|
||
|
||
// 检查任务状态
|
||
switch {
|
||
case IsCompleted(status.Data.SubTaskInfos[0].Status):
|
||
fmt.Println("任务成功完成!, 完成状态:", status)
|
||
return status, nil
|
||
default:
|
||
fmt.Println("任务运行中: ", status)
|
||
}
|
||
|
||
// 等待下一次轮询
|
||
select {
|
||
case <-ctx.Done():
|
||
return nil, ctx.Err()
|
||
case <-time.After(baseInterval):
|
||
}
|
||
}
|
||
}
|
||
|
||
func IsCompleted(status string) bool {
|
||
if status == "Completed" || status == "Failed" || status == "Succeed" {
|
||
return true
|
||
}
|
||
return false
|
||
}
|
||
func min(a, b time.Duration) time.Duration {
|
||
if a < b {
|
||
return a
|
||
}
|
||
return b
|
||
}
|