remote-task-excutor-cli/pkg/runners/task_runner.go

286 lines
7.3 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package runners
import (
"context"
"fmt"
"math/rand"
"os"
"os/exec"
"path"
"path/filepath"
"remote-task-excutor-cli/pkg/models"
"remote-task-excutor-cli/pkg/service"
"strings"
"time"
"github.com/otiai10/copy"
)
type taskRunner struct {
authService service.AuthService
preparationService service.PreparationService
taskRunService service.TaskService
logService service.LogService
}
func NewTaskRunner(authService service.AuthService, preparationService service.PreparationService,
taskRunService service.TaskService, logService service.LogService) *taskRunner {
return &taskRunner{
authService: authService,
preparationService: preparationService,
taskRunService: taskRunService,
logService: logService,
}
}
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")
clusterID := config.Resource.ClusterID
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)
done := make(chan bool)
defer func() {
done <- true
close(done)
}()
go func() {
err := r.RecordLogFile(ctx, authData, jobSetID, done)
if err != nil {
fmt.Println("获取日志失败:", err)
}
}()
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
}
//if err := r.MergeAimRepo(ctx, config); err != nil {
// return err
//}
//
//if err := r.MergeTensorboard(ctx, config); err != nil {
// return err
//}
return nil
}
func (r *taskRunner) doMergeAim(repoPath string) error {
// 构建 Python 命令
pythonScript := "/app/merge_aim.py"
args := []string{"--repo_path=" + repoPath}
// 创建命令
cmd := exec.Command("python", append([]string{pythonScript}, args...)...)
// 设置命令输出和错误流
cmd.Stdout = os.Stdout
cmd.Stderr = os.Stderr
// 执行命令
fmt.Printf("执行命令: python %s %s\n", pythonScript, strings.Join(args, " "))
err := cmd.Run()
if err != nil {
return fmt.Errorf("执行 Python 脚本失败: %v", err)
}
return nil
}
func (r *taskRunner) MergeAimRepo(ctx context.Context, config *models.RunConfig) error {
// 1. 检查 config.TaskOuput
aimPath := path.Join(config.TaskOutput, "aim")
if !r.hasSubDirSimple(aimPath, "aim") {
fmt.Println("没有.aim目录不需要合并")
return nil
}
// 2.合并aim
if err := r.doMergeAim(aimPath); err != nil {
return err
}
// 3. 删除aim目录
return os.Remove(aimPath)
}
func (r *taskRunner) MergeTensorboard(ctx context.Context, config *models.RunConfig) error {
if !r.hasSubDirSimple(config.TaskOutput, "tensorboard-logs") {
fmt.Println("没有tensorboard-logs目录不需要合并")
return nil
}
tensorboardPath := path.Join(config.TaskOutput, "tensorboard")
if err := copy.Copy(tensorboardPath, " /tensorboard-logs"); err != nil {
return err
}
return os.Remove(tensorboardPath)
}
func (r *taskRunner) hasSubDirSimple(dirPath string, subDir string) bool {
// 检查目录是否存在且不为空
if !r.isDirExistAndNotEmpty(dirPath) {
return false
}
// 检查 .aim 子目录是否存在
aimPath := filepath.Join(dirPath, subDir)
info, err := os.Stat(aimPath)
return err == nil && info.IsDir()
}
// isDirExistAndNotEmpty 检查目录是否存在且不为空
func (r *taskRunner) isDirExistAndNotEmpty(dirPath string) bool {
// 检查目录是否存在
info, err := os.Stat(dirPath)
if err != nil || !info.IsDir() {
return false
}
// 检查目录是否为空
dir, err := os.Open(dirPath)
if err != nil {
return false
}
defer dir.Close()
// 读取第一个实际文件(排除 "." 和 ".."
files, err := dir.Readdir(3) // 读取最多3个文件
if err != nil {
return false
}
// 检查是否有实际文件(非 "." 和 ".."
for _, file := range files {
if file.Name() != "." && file.Name() != ".." {
return true
}
}
return false
}
func (r *taskRunner) RecordLogFile(ctx context.Context, authData *models.AuthData, jobSetID string, done chan bool) error {
return r.logService.RecordLogToFile(ctx, authData, jobSetID, done)
}
func (r *taskRunner) PollTaskStatusWithBackoff(ctx context.Context, authData *models.AuthData, jobSetID string) (*models.TaskDetailResponse, error) {
const (
baseInterval = 10 * time.Second
maxInterval = 1 * time.Minute
maxFailAttempts = 5 // 连续失败最多尝试5次
maxPollCount = 0 // 0表示不限制轮询次数可以设置一个最大值
)
var (
backoff = baseInterval
failAttemptCount = 0 // 连续失败次数
pollCount = 0 // 总轮询次数
)
for {
// 检查上下文是否被取消
if ctx.Err() != nil {
return nil, ctx.Err()
}
// 获取任务状态
status, err := r.taskRunService.GetTaskStatus(ctx, authData, jobSetID)
if err != nil {
failAttemptCount++
// 如果连续失败次数超过限制,返回错误
if failAttemptCount > maxFailAttempts {
return nil, fmt.Errorf("连续查询失败%d次最后错误: %w", maxFailAttempts, err)
}
// 指数回退
backoff = mymin(backoff*2, maxInterval)
sleepTime := backoff + time.Duration(rand.Int63n(int64(backoff/2))) // 随机抖动避免同步
fmt.Printf("任务查询失败(%s), 连续失败%d次将在 %s 后重试\n", err, failAttemptCount, sleepTime)
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-time.After(sleepTime):
continue
}
}
// 查询成功,重置失败计数和回退间隔
failAttemptCount = 0
backoff = baseInterval
pollCount++
// 检查返回的状态数据是否有效
if status == nil || len(status.Data.SubTaskInfos) == 0 {
fmt.Println("任务状态数据为空,继续等待...")
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-time.After(baseInterval):
continue
}
}
// 检查任务状态
taskStatus := status.Data.SubTaskInfos[0].Status
if IsCompleted(taskStatus) {
fmt.Printf("任务完成,状态: %s\n", taskStatus)
return status, nil
}
// 检查任务是否失败
if taskStatus == "Failed" {
return status, fmt.Errorf("任务执行失败,状态: %s", taskStatus)
}
// 任务运行中,继续轮询
fmt.Printf("任务运行中,状态: %s已轮询 %d 次\n", taskStatus, pollCount)
// 等待下一次轮询
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-time.After(baseInterval):
}
}
}
func IsCompleted(status string) bool {
if status == "Completed" || status == "Failed" || status == "Succeed" || status == "Stopped" {
return true
}
return false
}
func mymin(a, b time.Duration) time.Duration {
if a < b {
return a
}
return b
}