265 lines
6.4 KiB
Go
265 lines
6.4 KiB
Go
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
|
||
}
|