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

265 lines
6.4 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"
"github.com/otiai10/copy"
"math/rand"
"os"
"os/exec"
"path"
"path/filepath"
"remote-task-excutor-cli/pkg/models"
"remote-task-excutor-cli/pkg/service"
"strings"
"time"
)
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 = 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 = mymin(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 mymin(a, b time.Duration) time.Duration {
if a < b {
return a
}
return b
}