推理接口返回请求体修改

This commit is contained in:
somunslotus 2026-01-26 14:44:01 +08:00
parent e98bc40203
commit 3b7a8054c1
4 changed files with 21 additions and 25 deletions

View File

@ -285,7 +285,7 @@ type InferenceSubmitTaskResponse struct {
// InferenceSubmitTaskResponseData 推理任务提交响应数据
type InferenceSubmitTaskResponseData struct {
TaskInfo InferenceTaskInfo `json:"taskInfo"` // 任务信息
TaskInfo SubtaskRequest `json:"taskInfo"` // 任务信息
ResultInfo InferenceResultInfo `json:"resultInfo"` // 结果信息
}

View File

@ -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 获取任务类型的中文名称

View File

@ -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,

View File

@ -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) {