173 lines
4.9 KiB
Go
173 lines
4.9 KiB
Go
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))
|
||
}
|
||
|
||
// HealthCheck 处理健康检查请求
|
||
func (h *Handler) HealthCheck(c *gin.Context) {
|
||
c.JSON(http.StatusOK, SuccessResponse(gin.H{
|
||
"status": "UP",
|
||
}))
|
||
}
|