remote-task-excutor-cli/pkg/service/task/task_service.go

185 lines
5.6 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 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
}
}