286 lines
7.3 KiB
Go
286 lines
7.3 KiB
Go
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
|
||
}
|