185 lines
5.6 KiB
Go
185 lines
5.6 KiB
Go
package task
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"fmt"
|
||
"os"
|
||
"remote-task-excutor-cli/pkg/client"
|
||
"remote-task-excutor-cli/pkg/models"
|
||
"remote-task-excutor-cli/pkg/service"
|
||
"remote-task-excutor-cli/pkg/service/common"
|
||
"remote-task-excutor-cli/pkg/service/updown"
|
||
"time"
|
||
|
||
"github.com/mholt/archiver/v3"
|
||
)
|
||
|
||
type taskService struct {
|
||
httpClient *client.HTTPClient
|
||
taskSubmitter *common.TaskSubmitter
|
||
updownService service.UpDownService
|
||
}
|
||
|
||
func NewTaskService(baseURL string, timeout time.Duration) service.TaskService {
|
||
httpClient := client.NewHTTPClient(baseURL, timeout)
|
||
return &taskService{
|
||
httpClient: httpClient,
|
||
taskSubmitter: common.NewTaskSubmitter(httpClient, common.TrainingTaskType),
|
||
updownService: updown.NewUpDownService(baseURL, timeout),
|
||
}
|
||
}
|
||
|
||
func (s *taskService) SubmitTask(ctx context.Context, authData *models.AuthData, config *models.RunConfig,
|
||
clusterID string, bindResultSet *models.BindResultSet) (string, error) {
|
||
jobSetID, _, err := s.taskSubmitter.SubmitTask(ctx, authData, config, clusterID, bindResultSet)
|
||
return jobSetID, err
|
||
}
|
||
|
||
func (s *taskService) GetTaskStatus(ctx context.Context, authData *models.AuthData, jobSetID string) (*models.TaskDetailResponse, error) {
|
||
s.httpClient.SetHeader("Authorization", "Bearer "+authData.Token)
|
||
params := map[string]string{
|
||
models.JobSetIDKey: jobSetID,
|
||
models.LocalJobIDKey: models.MainTaskID,
|
||
}
|
||
|
||
for {
|
||
resp, err := s.httpClient.Get("/jsm/v2/jobs/details", params)
|
||
if err != nil {
|
||
fmt.Println("Get task status failed: ", err)
|
||
return nil, err
|
||
}
|
||
|
||
var taskStatus models.TaskDetailResponse
|
||
if err := json.Unmarshal(resp, &taskStatus); err != nil {
|
||
fmt.Println("Get task status response unmarshal failed: ", err)
|
||
return nil, err
|
||
}
|
||
|
||
if taskStatus.Code != models.ResponseOK {
|
||
fmt.Println("Get task status failed: ", taskStatus.Code)
|
||
return nil, fmt.Errorf("get task status failed: code: %s, message: %s", taskStatus.Code, taskStatus.Message)
|
||
}
|
||
fmt.Println("Get task status result: ", string(resp))
|
||
if len(taskStatus.Data.SubTaskInfos) == 0 {
|
||
time.Sleep(5 * time.Second)
|
||
continue
|
||
}
|
||
return &taskStatus, nil
|
||
}
|
||
}
|
||
|
||
func (s *taskService) GetTaskResult(ctx context.Context, authData *models.AuthData, path string, jobSetID string) error {
|
||
result, err := s.getTaskResultDetail(ctx, authData, jobSetID)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
|
||
filename, err := s.updownService.DownLoad(ctx, authData, path, result.PcmJobData.ResultFiles[0].Objects[0].PackageID)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
|
||
if err := archiver.Unarchive(filename, path); err != nil {
|
||
fmt.Println("Unarchive task result failed: ", err)
|
||
return err
|
||
}
|
||
|
||
// 删除filename文件
|
||
if err := os.Remove(filename); err != nil {
|
||
fmt.Println("Remove task result failed: ", err)
|
||
return err
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
func (s *taskService) getTaskResultDetail(ctx context.Context, authData *models.AuthData,
|
||
jobSetID string) (*models.TaskResultDetailData, error) {
|
||
s.httpClient.SetHeader("Authorization", "Bearer "+authData.Token)
|
||
queryTaskResultReq := models.TaskQueryRequest{
|
||
JobSetID: jobSetID,
|
||
LocalJobID: models.MainTaskID,
|
||
}
|
||
|
||
// 设置10分钟超时
|
||
const maxQueryDuration = 10 * time.Minute
|
||
startTime := time.Now()
|
||
queryInterval := 10 * time.Second
|
||
|
||
for {
|
||
// 检查是否超过10分钟
|
||
elapsed := time.Since(startTime)
|
||
if elapsed >= maxQueryDuration {
|
||
return nil, fmt.Errorf("查询任务结果超时:已查询 %v,超过最大查询时间 %v", elapsed, maxQueryDuration)
|
||
}
|
||
|
||
// 检查上下文是否被取消
|
||
if ctx.Err() != nil {
|
||
return nil, fmt.Errorf("查询任务结果被取消: %w", ctx.Err())
|
||
}
|
||
|
||
resp, err := s.httpClient.PostJSON("/jsm/v2/jobs/results", queryTaskResultReq)
|
||
fmt.Println("get result details resp:", string(resp))
|
||
if err != nil {
|
||
fmt.Printf("Get task result details failed: %v (已查询 %v)\n", err, elapsed)
|
||
// 网络错误时等待后重试
|
||
select {
|
||
case <-ctx.Done():
|
||
return nil, fmt.Errorf("查询任务结果被取消: %w", ctx.Err())
|
||
case <-time.After(queryInterval):
|
||
continue
|
||
}
|
||
}
|
||
|
||
var taskResultDetails models.TaskResultDetailResponse
|
||
if err := json.Unmarshal(resp, &taskResultDetails); err != nil {
|
||
fmt.Printf("Get task result details response unmarshal failed: %v (已查询 %v)\n", err, elapsed)
|
||
// 解析错误时等待后重试
|
||
select {
|
||
case <-ctx.Done():
|
||
return nil, fmt.Errorf("查询任务结果被取消: %w", ctx.Err())
|
||
case <-time.After(queryInterval):
|
||
continue
|
||
}
|
||
}
|
||
|
||
if taskResultDetails.Code != models.ResponseOK {
|
||
if taskResultDetails.Message == "query pcm job: record not found" {
|
||
fmt.Printf("任务结果未找到,继续查询... (已查询 %v)\n", elapsed)
|
||
select {
|
||
case <-ctx.Done():
|
||
return nil, fmt.Errorf("查询任务结果被取消: %w", ctx.Err())
|
||
case <-time.After(queryInterval):
|
||
continue
|
||
}
|
||
} else {
|
||
fmt.Printf("Get task result details failed: code=%s, message=%s (已查询 %v)\n",
|
||
taskResultDetails.Code, taskResultDetails.Message, elapsed)
|
||
// 其他错误也等待后重试,直到超时
|
||
select {
|
||
case <-ctx.Done():
|
||
return nil, fmt.Errorf("查询任务结果被取消: %w", ctx.Err())
|
||
case <-time.After(queryInterval):
|
||
continue
|
||
}
|
||
}
|
||
}
|
||
|
||
// 检查任务状态
|
||
if taskResultDetails.Data.PcmJobData.Status == "failed" {
|
||
fmt.Printf("任务状态为 failed,继续查询... (已查询 %v)\n", elapsed)
|
||
select {
|
||
case <-ctx.Done():
|
||
return nil, fmt.Errorf("查询任务结果被取消: %w", ctx.Err())
|
||
case <-time.After(queryInterval):
|
||
continue
|
||
}
|
||
}
|
||
|
||
// 成功获取结果
|
||
fmt.Printf("成功获取任务结果 (耗时 %v)\n", elapsed)
|
||
return &taskResultDetails.Data, nil
|
||
}
|
||
}
|