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

208 lines
6.0 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 (
"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 处理模型上传请求
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
}
// 获取认证信息
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 {
c.JSON(http.StatusOK, InternalServerErrorResponse("准备模型失败: "+err.Error()))
return
}
c.JSON(http.StatusOK, SuccessResponse(gin.H{
"modelID": modelID,
}))
}
// UploadCode 处理代码上传请求
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
}
// 获取认证信息
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 {
c.JSON(http.StatusOK, InternalServerErrorResponse("准备代码失败: "+err.Error()))
return
}
c.JSON(http.StatusOK, SuccessResponse(gin.H{
"codeID": codeID,
}))
}
// SubmitTask 处理任务提交请求
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
}
// 获取认证信息
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 {
c.JSON(http.StatusOK, InternalServerErrorResponse("提交任务失败: "+err.Error()))
return
}
c.JSON(http.StatusOK, SuccessResponse(gin.H{
"jobSetID": jobSetID,
}))
}
// QueryStatus 处理状态查询请求
func (h *Handler) QueryStatus(c *gin.Context) {
var request struct {
JobSetID string `json:"jobSetID" binding:"required"`
}
if err := c.ShouldBindJSON(&request); err != nil {
c.JSON(http.StatusOK, BadRequestResponse("请求参数错误: "+err.Error()))
return
}
// 获取认证信息
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.JobSetID)
if err != nil {
c.JSON(http.StatusOK, InternalServerErrorResponse("查询任务状态失败: "+err.Error()))
return
}
c.JSON(http.StatusOK, SuccessResponse(status))
}
// SubmitInferenceTask 处理推理任务提交请求
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
}
// 获取认证信息
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 {
c.JSON(http.StatusOK, InternalServerErrorResponse("提交推理任务失败: "+err.Error()))
return
}
c.JSON(http.StatusOK, response)
}
// HealthCheck 处理健康检查请求
func (h *Handler) HealthCheck(c *gin.Context) {
c.JSON(http.StatusOK, SuccessResponse(gin.H{
"status": "UP",
}))
}