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

208 lines
6.0 KiB
Go
Raw Normal View History

2025-09-30 08:40:52 +08:00
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))
}
2026-01-26 14:32:32 +08:00
// 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)
}
2025-09-30 08:40:52 +08:00
// HealthCheck 处理健康检查请求
func (h *Handler) HealthCheck(c *gin.Context) {
c.JSON(http.StatusOK, SuccessResponse(gin.H{
"status": "UP",
}))
}