476 lines
16 KiB
Go
476 lines
16 KiB
Go
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)
|
||
}
|