diff --git a/pkg/models/task.go b/pkg/models/task.go index b524772..d2f9372 100644 --- a/pkg/models/task.go +++ b/pkg/models/task.go @@ -285,7 +285,7 @@ type InferenceSubmitTaskResponse struct { // InferenceSubmitTaskResponseData 推理任务提交响应数据 type InferenceSubmitTaskResponseData struct { - TaskInfo InferenceTaskInfo `json:"taskInfo"` // 任务信息 + TaskInfo SubtaskRequest `json:"taskInfo"` // 任务信息 ResultInfo InferenceResultInfo `json:"resultInfo"` // 结果信息 } diff --git a/pkg/service/common/task_submitter.go b/pkg/service/common/task_submitter.go index f678342..088cd72 100644 --- a/pkg/service/common/task_submitter.go +++ b/pkg/service/common/task_submitter.go @@ -22,9 +22,9 @@ func NewTaskSubmitter(httpClient *client.HTTPClient, taskType string) *TaskSubmi } } -// SubmitTask 提交任务的通用方法 +// SubmitTask 提交任务的通用方法,返回 jobSetID、提交请求和错误 func (ts *TaskSubmitter) SubmitTask(ctx context.Context, authData *models.AuthData, config *models.RunConfig, - clusterID string, bindResultSet *models.BindResultSet) (string, error) { + clusterID string, bindResultSet *models.BindResultSet) (string, *models.SubtaskRequest, error) { // 构建任务请求 submitTaskReq := BuildSubmitTaskRequest(ctx, authData, config, clusterID, bindResultSet, ts.taskType) @@ -39,7 +39,7 @@ func (ts *TaskSubmitter) SubmitTask(ctx context.Context, authData *models.AuthDa resp, err := ts.httpClient.PostJSON("/jsm/v2/jobs/submit", submitTaskReq) if err != nil { fmt.Printf("Submit %s task failed: %v\n", getTaskTypeName(ts.taskType), err) - return "", err + return "", nil, err } fmt.Printf("提交%s任务结果:%s\n", getTaskTypeName(ts.taskType), string(resp)) @@ -48,16 +48,16 @@ func (ts *TaskSubmitter) SubmitTask(ctx context.Context, authData *models.AuthDa var submitTaskResp models.SubmitTaskResponse if err := json.Unmarshal(resp, &submitTaskResp); err != nil { fmt.Printf("Submit %s task response unmarshal failed: %v\n", getTaskTypeName(ts.taskType), err) - return "", err + return "", nil, err } if submitTaskResp.Code != models.ResponseOK { fmt.Printf("Submit %s task failed: %s\n", getTaskTypeName(ts.taskType), submitTaskResp.Code) - return "", fmt.Errorf("submit %s task failed: %s", getTaskTypeName(ts.taskType), submitTaskResp.Code) + return "", nil, fmt.Errorf("submit %s task failed: %s", getTaskTypeName(ts.taskType), submitTaskResp.Code) } fmt.Printf("Submit %s task result: %s\n", getTaskTypeName(ts.taskType), string(resp)) - return submitTaskResp.Data.JobSetID, nil + return submitTaskResp.Data.JobSetID, &submitTaskReq, nil } // getTaskTypeName 获取任务类型的中文名称 diff --git a/pkg/service/inference/inference_service.go b/pkg/service/inference/inference_service.go index fea259d..a09c8ca 100644 --- a/pkg/service/inference/inference_service.go +++ b/pkg/service/inference/inference_service.go @@ -37,7 +37,8 @@ func NewInferenceService( // SubmitTask 提交推理任务 func (s *inferenceService) SubmitTask(ctx context.Context, authData *models.AuthData, config *models.RunConfig, clusterID string, bindResultSet *models.BindResultSet) (string, error) { - return s.taskSubmitter.SubmitTask(ctx, authData, config, clusterID, bindResultSet) + jobSetID, _, err := s.taskSubmitter.SubmitTask(ctx, authData, config, clusterID, bindResultSet) + return jobSetID, err } // GetTaskStatus 查询任务状态 @@ -105,22 +106,10 @@ func (s *inferenceService) SubmitInferenceTask(ctx context.Context, authData *mo fmt.Println("增量模型准备成功,ID:", subModelID) } - // 4. 构建推理任务响应数据(用于返回) - taskData := common.BuildInferenceSubmitTaskRequest( - ctx, authData, - codeID, modelID, subModelID, - request.Image.ImageID, - clusterID, - request.ResourceType, - request.Resource.Resources, - request.Command, - request.Description, - ) - - // 5. 使用 taskSubmitter 提交任务 + // 4. 使用 taskSubmitter 提交任务 // 注意:taskSubmitter 会使用 BuildSubmitTaskRequest,它构建的是训练任务格式 // 对于推理任务,我们需要传入 BindDatasetID=0,这样 Files 中的 Dataset 会是空的 - jobSetID, err := s.taskSubmitter.SubmitTask(ctx, authData, runConfig, clusterID, &models.BindResultSet{ + jobSetID, submitTaskReq, err := s.taskSubmitter.SubmitTask(ctx, authData, runConfig, clusterID, &models.BindResultSet{ BindCodeID: codeID, BindModelID: modelID, BindDatasetID: 0, // 推理任务不需要数据集 @@ -129,8 +118,14 @@ func (s *inferenceService) SubmitInferenceTask(ctx context.Context, authData *mo return nil, fmt.Errorf("提交任务失败: %w", err) } - // 更新结果信息 - taskData.ResultInfo.JobSetID = jobSetID + // 5. 构建返回响应数据,使用提交请求中的信息 + taskData := models.InferenceSubmitTaskResponseData{ + TaskInfo: *submitTaskReq, + ResultInfo: models.InferenceResultInfo{ + LocalJobID: models.MainTaskID, + JobSetID: jobSetID, + }, + } return &models.InferenceSubmitTaskResponse{ Code: 200, diff --git a/pkg/service/task/task_service.go b/pkg/service/task/task_service.go index 4984330..7a6755d 100644 --- a/pkg/service/task/task_service.go +++ b/pkg/service/task/task_service.go @@ -32,7 +32,8 @@ func NewTaskService(baseURL string, timeout time.Duration) service.TaskService { func (s *taskService) SubmitTask(ctx context.Context, authData *models.AuthData, config *models.RunConfig, clusterID string, bindResultSet *models.BindResultSet) (string, error) { - return s.taskSubmitter.SubmitTask(ctx, authData, config, clusterID, bindResultSet) + 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) {