remote-task-excutor-cli/pkg/handler/handler.go

476 lines
16 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 handler
import (
"fmt"
"net/http"
"remote-task-excutor-cli/pkg/models"
"remote-task-excutor-cli/pkg/service"
"github.com/gin-gonic/gin"
)
// Handler 包含所有HTTP处理函数组合auth、preparation和inference三个service
type Handler struct {
authService service.AuthService
preparationService service.PreparationService
inferenceService service.InferenceService
}
// NewHandler 创建新的Handler实例
func NewHandler(
authService service.AuthService,
preparationService service.PreparationService,
inferenceService service.InferenceService,
) *Handler {
return &Handler{
authService: authService,
preparationService: preparationService,
inferenceService: inferenceService,
}
}
// UploadModel 处理模型上传请求
// @Summary 上传模型资源
// @Description 上传模型资源到远程存储
// @Tags 资源上传
// @Accept json
// @Produce json
// @Param request body models.DataResourceConfig true "模型配置"
// @Success 200 {object} Response{data=object{modelID=int}}
// @Failure 400 {object} Response
// @Failure 500 {object} Response
// @Router /api/v1/uploadModel [post]
func (h *Handler) UploadModel(c *gin.Context) {
var config models.DataResourceConfig
if err := c.ShouldBindJSON(&config); err != nil {
c.JSON(http.StatusOK, BadRequestResponse("请求参数错误: "+err.Error()))
return
}
fmt.Printf("[UploadModel] 请求参数: %+v\n", config)
// 获取认证信息
authData, err := h.authService.GetToken(c.Request.Context())
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取认证信息失败: "+err.Error()))
return
}
// 获取集群ID
clusterID, err := h.authService.GetClusterID(c.Request.Context(), "default")
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取集群ID失败: "+err.Error()))
return
}
// 准备模型
modelID, err := h.preparationService.PrepareModel(c.Request.Context(), authData, config, clusterID)
if err != nil {
response := InternalServerErrorResponse("准备模型失败: " + err.Error())
fmt.Printf("[UploadModel] 返回错误: %+v\n", response)
c.JSON(http.StatusOK, response)
return
}
response := SuccessResponse(gin.H{
"modelID": modelID,
})
fmt.Printf("[UploadModel] 返回结果: %+v\n", response)
c.JSON(http.StatusOK, response)
}
// UploadCode 处理代码上传请求
// @Summary 上传代码资源
// @Description 上传代码资源到远程存储
// @Tags 资源上传
// @Accept json
// @Produce json
// @Param request body models.UploadCodeRequest true "代码配置"
// @Success 200 {object} Response{data=object{codeID=int}}
// @Failure 400 {object} Response
// @Failure 500 {object} Response
// @Router /api/v1/uploadCode [post]
func (h *Handler) UploadCode(c *gin.Context) {
var request models.UploadCodeRequest
if err := c.ShouldBindJSON(&request); err != nil {
c.JSON(http.StatusOK, BadRequestResponse("请求参数错误: "+err.Error()))
return
}
fmt.Printf("[UploadCode] 请求参数: %+v\n", request)
// 获取认证信息
authData, err := h.authService.GetToken(c.Request.Context())
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取认证信息失败: "+err.Error()))
return
}
// 获取集群ID
clusterID, err := h.authService.GetClusterID(c.Request.Context(), "default")
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取集群ID失败: "+err.Error()))
return
}
// 准备代码
codeID, err := h.preparationService.PrepareCode(c.Request.Context(), authData, &request.RunConfig, clusterID)
if err != nil {
response := InternalServerErrorResponse("准备代码失败: " + err.Error())
fmt.Printf("[UploadCode] 返回错误: %+v\n", response)
c.JSON(http.StatusOK, response)
return
}
response := SuccessResponse(gin.H{
"codeID": codeID,
})
fmt.Printf("[UploadCode] 返回结果: %+v\n", response)
c.JSON(http.StatusOK, response)
}
// SubmitTask 处理任务提交请求
// @Summary 提交任务(已废弃,请使用 submitTask
// @Description 提交任务使用已上传的代码和模型ID
// @Tags 任务管理
// @Accept json
// @Produce json
// @Param request body models.SubmitTaskRequest true "任务配置"
// @Success 200 {object} Response{data=object{jobSetID=string}}
// @Failure 400 {object} Response
// @Failure 500 {object} Response
// @Router /api/v1/submitTask [post]
// @Deprecated
func (h *Handler) SubmitTask(c *gin.Context) {
var request models.SubmitTaskRequest
if err := c.ShouldBindJSON(&request); err != nil {
c.JSON(http.StatusOK, BadRequestResponse("请求参数错误: "+err.Error()))
return
}
fmt.Printf("[SubmitTask] 请求参数: %+v\n", request)
// 获取认证信息
authData, err := h.authService.GetToken(c.Request.Context())
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取认证信息失败: "+err.Error()))
return
}
// 获取集群ID
clusterID, err := h.authService.GetClusterID(c.Request.Context(), "default")
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取集群ID失败: "+err.Error()))
return
}
// 创建BindResultSet使用传入的CodeID、ModelID和DatasetID
bindResultSet := &models.BindResultSet{
BindCodeID: request.CodeID,
BindDatasetID: request.DatasetID,
BindModelID: request.ModelID,
}
// 提交任务
jobSetID, err := h.inferenceService.SubmitTask(c.Request.Context(), authData, &request.RunConfig, clusterID, bindResultSet)
if err != nil {
response := InternalServerErrorResponse("提交任务失败: " + err.Error())
fmt.Printf("[SubmitTask] 返回错误: %+v\n", response)
c.JSON(http.StatusOK, response)
return
}
response := SuccessResponse(gin.H{
"jobSetID": jobSetID,
})
fmt.Printf("[SubmitTask] 返回结果: %+v\n", response)
c.JSON(http.StatusOK, response)
}
// QueryStatus 处理状态查询请求
// @Summary 查询任务状态
// @Description 查询推理任务的执行状态和推理URL
// @Tags 任务查询
// @Accept json
// @Produce json
// @Param request body models.InferenceTaskStatusRequest true "查询参数"
// @Success 200 {object} models.InferenceTaskStatusResponse
// @Failure 400 {object} Response
// @Failure 500 {object} Response
// @Router /api/v1/getTaskStatus [post]
func (h *Handler) QueryStatus(c *gin.Context) {
var request models.InferenceTaskStatusRequest
if err := c.ShouldBindJSON(&request); err != nil {
c.JSON(http.StatusOK, BadRequestResponse("请求参数错误: "+err.Error()))
return
}
fmt.Printf("[QueryStatus] 请求参数: %+v\n", request)
// 获取认证信息
authData, err := h.authService.GetToken(c.Request.Context())
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取认证信息失败: "+err.Error()))
return
}
// 查询任务状态
status, err := h.inferenceService.GetTaskStatus(c.Request.Context(), authData, request.LocalJobID, request.JobSetID)
if err != nil {
response := InternalServerErrorResponse("查询任务状态失败: " + err.Error())
fmt.Printf("[QueryStatus] 返回错误: %+v\n", response)
c.JSON(http.StatusOK, response)
return
}
fmt.Printf("[QueryStatus] 返回结果: %+v\n", status)
c.JSON(http.StatusOK, status)
}
// StopInferenceTask 处理停止推理任务请求
// @Summary 停止推理任务
// @Description 停止正在运行的推理任务
// @Tags 任务控制
// @Accept json
// @Produce json
// @Param request body models.StopInferenceTaskRequest true "停止任务参数"
// @Success 200 {object} models.StopInferenceTaskResponse
// @Failure 400 {object} Response
// @Failure 500 {object} Response
// @Router /api/v1/stopInferenceTask [post]
func (h *Handler) StopInferenceTask(c *gin.Context) {
var request models.StopInferenceTaskRequest
if err := c.ShouldBindJSON(&request); err != nil {
c.JSON(http.StatusOK, BadRequestResponse("请求参数错误: "+err.Error()))
return
}
fmt.Printf("[StopInferenceTask] 请求参数: %+v\n", request)
// 获取认证信息
authData, err := h.authService.GetToken(c.Request.Context())
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取认证信息失败: "+err.Error()))
return
}
// 停止任务
response, err := h.inferenceService.StopTask(c.Request.Context(), authData, request.LocalJobID, request.JobSetID)
if err != nil {
errResponse := InternalServerErrorResponse("停止任务失败: " + err.Error())
fmt.Printf("[StopInferenceTask] 返回错误: %+v\n", errResponse)
c.JSON(http.StatusOK, errResponse)
return
}
fmt.Printf("[StopInferenceTask] 返回结果: %+v\n", response)
c.JSON(http.StatusOK, response)
}
// ReSubmitTask 处理重启推理任务请求
// @Summary 重启推理任务
// @Description 将入参原样转发给第三方 /jsm/v2/jobs/submit返回新的 jobSetID 与入参中的 localJobID
// @Tags 任务控制
// @Accept json
// @Produce json
// @Param request body models.SubtaskRequest true "重启任务参数userID、jobSetInfo.jobs"
// @Success 200 {object} models.ReSubmitTaskResponse
// @Failure 400 {object} Response
// @Failure 500 {object} Response
// @Router /api/v1/reSubmitTask [post]
func (h *Handler) ReSubmitTask(c *gin.Context) {
var request models.SubtaskRequest
if err := c.ShouldBindJSON(&request); err != nil {
c.JSON(http.StatusOK, BadRequestResponse("请求参数错误: "+err.Error()))
return
}
fmt.Printf("[ReSubmitTask] 请求参数: %+v\n", request)
if len(request.JobSetInfo.Jobs) == 0 {
c.JSON(http.StatusOK, BadRequestResponse("jobSetInfo.jobs 不能为空"))
return
}
// 获取认证信息
authData, err := h.authService.GetToken(c.Request.Context())
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取认证信息失败: "+err.Error()))
return
}
response, err := h.inferenceService.ReSubmitTask(c.Request.Context(), authData, &request)
if err != nil {
errResponse := InternalServerErrorResponse("重启任务失败: " + err.Error())
fmt.Printf("[ReSubmitTask] 返回错误: %+v\n", errResponse)
c.JSON(http.StatusOK, errResponse)
return
}
fmt.Printf("[ReSubmitTask] 返回结果: %+v\n", response)
c.JSON(http.StatusOK, response)
}
// SubmitInferenceTask 处理推理任务提交请求
// @Summary 提交推理任务(同步)
// @Description 同步提交推理任务,包含代码、模型、增量模型的准备步骤,等待完成后返回结果
// @Tags 任务管理
// @Accept json
// @Produce json
// @Param request body models.InferenceSubmitTaskRequest true "推理任务配置"
// @Success 200 {object} models.InferenceSubmitTaskResponse
// @Failure 400 {object} Response
// @Failure 500 {object} Response
// @Router /api/v1/submitTask [post]
func (h *Handler) SubmitInferenceTask(c *gin.Context) {
var request models.InferenceSubmitTaskRequest
if err := c.ShouldBindJSON(&request); err != nil {
c.JSON(http.StatusOK, BadRequestResponse("请求参数错误: "+err.Error()))
return
}
fmt.Printf("[SubmitInferenceTask] 请求参数: %+v\n", request)
// 参数校验:检查 model.Path 是否为空
if request.Model.Path == "" {
c.JSON(http.StatusOK, BadRequestResponse("model.Path 不能为空"))
return
}
// 参数校验:如果提供了 subModel检查 subModel.Path 是否为空
if request.SubModel != nil && request.SubModel.Path == "" {
c.JSON(http.StatusOK, BadRequestResponse("subModel.Path 不能为空"))
return
}
// 获取认证信息
authData, err := h.authService.GetToken(c.Request.Context())
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取认证信息失败: "+err.Error()))
return
}
// 获取集群ID从 resource 中获取,如果没有则使用默认值)
clusterID := request.Resource.ClusterID
if clusterID == "" {
clusterID, err = h.authService.GetClusterID(c.Request.Context(), "default")
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取集群ID失败: "+err.Error()))
return
}
}
// 提交推理任务(包含准备步骤)
response, err := h.inferenceService.SubmitInferenceTask(c.Request.Context(), authData, &request, clusterID)
if err != nil {
errResponse := InternalServerErrorResponse("提交推理任务失败: " + err.Error())
fmt.Printf("[SubmitInferenceTask] 返回错误: %+v\n", errResponse)
c.JSON(http.StatusOK, errResponse)
return
}
fmt.Printf("[SubmitInferenceTask] 返回结果: %+v\n", response)
c.JSON(http.StatusOK, response)
}
// SubmitInferenceTaskAsync 异步提交推理任务:立刻返回 task_id
// @Summary 异步提交推理任务
// @Description 异步提交推理任务,立即返回 task_id任务在后台处理
// @Tags 任务管理
// @Accept json
// @Produce json
// @Param request body models.InferenceSubmitTaskRequest true "推理任务配置"
// @Success 200 {object} models.InferenceAsyncSubmitResponse
// @Failure 400 {object} Response
// @Failure 500 {object} Response
// @Router /api/v1/submitInferenceTaskAsync [post]
func (h *Handler) SubmitInferenceTaskAsync(c *gin.Context) {
var request models.InferenceSubmitTaskRequest
if err := c.ShouldBindJSON(&request); err != nil {
c.JSON(http.StatusOK, BadRequestResponse("请求参数错误: "+err.Error()))
return
}
fmt.Printf("[SubmitInferenceTaskAsync] 请求参数: %+v\n", request)
// 参数校验:检查 model.Path 是否为空
if request.Model.Path == "" {
c.JSON(http.StatusOK, BadRequestResponse("model.Path 不能为空"))
return
}
// 参数校验:如果提供了 subModel检查 subModel.Path 是否为空
if request.SubModel != nil && request.SubModel.Path == "" {
c.JSON(http.StatusOK, BadRequestResponse("subModel.Path 不能为空"))
return
}
authData, err := h.authService.GetToken(c.Request.Context())
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取认证信息失败: "+err.Error()))
return
}
clusterID := request.Resource.ClusterID
if clusterID == "" {
clusterID, err = h.authService.GetClusterID(c.Request.Context(), "default")
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("获取集群ID失败: "+err.Error()))
return
}
}
taskID, err := h.inferenceService.SubmitInferenceTaskAsync(c.Request.Context(), authData, &request, clusterID)
if err != nil {
errResponse := InternalServerErrorResponse("异步提交推理任务失败: " + err.Error())
fmt.Printf("[SubmitInferenceTaskAsync] 返回错误: %+v\n", errResponse)
c.JSON(http.StatusOK, errResponse)
return
}
response := models.InferenceAsyncSubmitResponse{
Code: http.StatusOK,
Msg: "",
Data: models.InferenceAsyncSubmitRespData{TaskID: taskID},
}
fmt.Printf("[SubmitInferenceTaskAsync] 返回结果: %+v\n", response)
c.JSON(http.StatusOK, response)
}
// GetInferenceTaskAsync 查询异步推理任务状态/结果
// @Summary 查询异步任务状态
// @Description 查询异步推理任务的处理状态、进度和最终结果
// @Tags 任务查询
// @Accept json
// @Produce json
// @Param request body models.InferenceAsyncTaskStatusRequest true "查询参数"
// @Success 200 {object} models.InferenceAsyncTaskStatusResponse
// @Failure 400 {object} Response
// @Failure 500 {object} Response
// @Router /api/v1/getInferenceTaskAsync [post]
func (h *Handler) GetInferenceTaskAsync(c *gin.Context) {
var request models.InferenceAsyncTaskStatusRequest
if err := c.ShouldBindJSON(&request); err != nil {
c.JSON(http.StatusOK, BadRequestResponse("请求参数错误: "+err.Error()))
return
}
fmt.Printf("[GetInferenceTaskAsync] 请求参数: %+v\n", request)
resp, err := h.inferenceService.GetInferenceTaskAsync(c.Request.Context(), request.TaskID)
if err != nil {
errResponse := InternalServerErrorResponse("查询异步任务失败: " + err.Error())
fmt.Printf("[GetInferenceTaskAsync] 返回错误: %+v\n", errResponse)
c.JSON(http.StatusOK, errResponse)
return
}
fmt.Printf("[GetInferenceTaskAsync] 返回结果: %+v\n", resp)
c.JSON(http.StatusOK, resp)
}
// HealthCheck 处理健康检查请求
// @Summary 健康检查
// @Description 检查服务健康状态
// @Tags 系统
// @Produce json
// @Success 200 {object} Response{data=object{status=string}}
// @Router /health [get]
func (h *Handler) HealthCheck(c *gin.Context) {
//fmt.Printf("[HealthCheck] 请求参数: 无\n")
response := SuccessResponse(gin.H{
"status": "UP",
})
//fmt.Printf("[HealthCheck] 返回结果: %+v\n", response)
c.JSON(http.StatusOK, response)
}