package handler import ( "context" "encoding/json" "fmt" "net/http" "strconv" "time" "gen2d/internal/db" "gen2d/internal/logger" "gen2d/internal/model" "gen2d/internal/service" "github.com/gin-gonic/gin" ) // GenerateRequest 素材生成请求。 type GenerateRequest struct { ProjectID string `json:"projectId"` Prompt string `json:"prompt"` AssetType string `json:"assetType" binding:"required"` Tags []string `json:"tags"` UserNote string `json:"userNote"` ProjectStyle map[string]string `json:"projectStyle"` TaskStyle map[string]string `json:"taskStyle"` Resolution int `json:"resolution"` Directions int `json:"directions"` FramesPerDir int `json:"framesPerDir"` Format string `json:"format"` } // GenerateResponse 素材生成响应体。 type GenerateResponse struct { TaskID string `json:"taskId"` } // AssetsResponse 素材列表响应。 type AssetsResponse struct { Assets []model.AssetResponse `json:"assets"` Metadata service.AssetMetadata `json:"metadata"` } // Generate 素材生成接口(异步)。 // 立即返回 taskId,后台执行管线,前端通过 GET /tasks/:taskId 轮询进度。 func Generate(c *gin.Context) { var req GenerateRequest if err := c.ShouldBindJSON(&req); err != nil { c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "参数错误: "+err.Error())) return } projectID := req.ProjectID if projectID == "" { projectID = "default" } taskID := fmt.Sprintf("task-%d", time.Now().UnixMilli()) // 保存任务到数据库 if err := saveTaskToDB(c.Request.Context(), projectID, taskID, req); err != nil { logger.FromCtx(c.Request.Context()).Error("failed to save task", "error", err) c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "创建任务失败")) return } // 返回 taskId c.JSON(http.StatusOK, model.OK(GenerateResponse{TaskID: taskID})) // 后台执行管线 go runPipelineBg(c.Request.Context(), projectID, taskID, req) } // runPipelineBg 后台执行生成管线,更新任务状态。 func runPipelineBg(ctx context.Context, projectID, taskID string, req GenerateRequest) { l := logger.With("task_id", taskID, "project_id", projectID) l.Info("task started", "asset_type", req.AssetType, "prompt", req.Prompt) // 注入进度上报回调 ctx = service.WithProgressReporter(ctx, func(stage string, progress int) { l.Info("progress update", "stage", stage, "progress", progress) updateTaskInDB(ctx, taskID, "running", stage, "", progress) }) updateTaskInDB(ctx, taskID, "running", "prompt_builder", "", 5) in := service.PipelineInput{ ProjectID: projectID, TaskID: taskID, Prompt: req.Prompt, AssetType: req.AssetType, Tags: req.Tags, UserNote: req.UserNote, ProjectStyle: req.ProjectStyle, TaskStyle: req.TaskStyle, Params: service.AssetParams{ Resolution: req.Resolution, Frames: service.FrameParams{ Directions: req.Directions, FramesPerDirection: req.FramesPerDir, }, Format: req.Format, }, } l.Info("calling pipeline", "input_tags", in.Tags, "user_note", in.UserNote) output, err := service.RunPipeline(ctx, in) if err != nil { l.Error("task pipeline failed", "error", err) updateTaskInDB(ctx, taskID, "failed", "", err.Error(), 0) return } l.Info("pipeline completed", "asset_count", len(output.Assets)) updateTaskInDB(ctx, taskID, "saving", "format_adapter", "", 90) // 上传素材并保存到数据库 for i, a := range output.Assets { key := fmt.Sprintf("generation/%s/%s/%d.%s", projectID, taskID, i, a.Format) cdnURL, err := storageSvc.Upload(ctx, key, a.Data) if err != nil { l.Error("upload asset failed", "index", i, "error", err) updateTaskInDB(ctx, taskID, "failed", "", "上传素材失败: "+err.Error(), 0) return } // 序列化单个素材的元数据 var metadata map[string]interface{} if i < len(output.Assets) { metadata = map[string]interface{}{ "index": i, } } metadataJSON, _ := json.Marshal(metadata) asset := &model.Asset{ TaskID: getTaskDBID(ctx, taskID), Key: key, URL: cdnURL, Format: a.Format, Metadata: string(metadataJSON), } if err := db.GetDB().WithContext(ctx).Create(asset).Error; err != nil { l.Error("save asset to db failed", "index", i, "error", err) } } // 更新为完成状态 var fullMetadata string if metadataJSON, err := json.Marshal(output.Metadata); err == nil { fullMetadata = string(metadataJSON) } db.GetDB().WithContext(ctx).Model(&model.Task{}). Where("external_id = ?", taskID). Updates(map[string]interface{}{ "status": "completed", "progress": 100, "metadata": fullMetadata, }) l.Info("task completed") } // GetTask 查询任务信息。 func GetTask(c *gin.Context) { taskID := c.Param("taskId") var task model.Task if err := db.GetDB().WithContext(c.Request.Context()). Where("external_id = ?", taskID). First(&task).Error; err != nil { c.JSON(http.StatusNotFound, model.Fail(http.StatusNotFound, "任务不存在")) return } c.JSON(http.StatusOK, model.OK(toTaskResponse(&task))) } // GetAssets 查询任务素材列表。 func GetAssets(c *gin.Context) { taskID := c.Param("taskId") var task model.Task if err := db.GetDB().WithContext(c.Request.Context()). Where("external_id = ?", taskID). First(&task).Error; err != nil { c.JSON(http.StatusNotFound, model.Fail(http.StatusNotFound, "任务不存在")) return } if task.Status != "completed" { c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "任务尚未完成,当前状态: "+task.Status)) return } var assets []model.Asset if err := db.GetDB().WithContext(c.Request.Context()). Where("task_id = ?", task.ID). Find(&assets).Error; err != nil { c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "查询素材失败")) return } response := []model.AssetResponse{} for _, a := range assets { response = append(response, model.AssetResponse{ Key: a.Key, URL: a.URL, Format: a.Format, }) } var metadata service.AssetMetadata if task.Metadata != "" { json.Unmarshal([]byte(task.Metadata), &metadata) } c.JSON(http.StatusOK, model.OK(AssetsResponse{ Assets: response, Metadata: metadata, })) } // saveTaskToDB 将任务记录保存到数据库。 func saveTaskToDB(ctx context.Context, projectID, taskID string, req GenerateRequest) error { // "default" 工程没有数据库记录,跳过 if projectID == "default" { return nil } projectIDUint, err := strconv.ParseUint(projectID, 10, 32) if err != nil { return fmt.Errorf("failed to parse projectID: %w", err) } task := &model.Task{ ExternalID: taskID, ProjectID: uint(projectIDUint), Prompt: req.Prompt, AssetType: req.AssetType, Status: "running", Progress: 5, Stage: "prompt_builder", RetryCount: 0, CreatedAt: time.Now(), UpdatedAt: time.Now(), } return db.GetDB().WithContext(ctx).Create(task).Error } // updateTaskInDB 更新数据库中的任务状态。 func updateTaskInDB(ctx context.Context, taskID, status, stage, error string, progress int) { updates := map[string]interface{}{ "status": status, "progress": progress, "updated_at": time.Now(), } if stage != "" { updates["stage"] = stage } if error != "" { updates["error"] = error } db.GetDB().WithContext(ctx).Model(&model.Task{}). Where("external_id = ?", taskID). Updates(updates) } // getTaskDBID 根据外部taskID获取数据库中的任务ID。 func getTaskDBID(ctx context.Context, taskID string) uint { var task model.Task db.GetDB().WithContext(ctx).Select("id").Where("external_id = ?", taskID).First(&task) return task.ID } // toTaskResponse 转换任务响应格式。 func toTaskResponse(task *model.Task) model.TaskResponse { return model.TaskResponse{ ID: task.ExternalID, ProjectID: fmt.Sprintf("%d", task.ProjectID), Prompt: task.Prompt, AssetType: task.AssetType, Status: task.Status, Stage: task.Stage, Progress: task.Progress, RetryCount: task.RetryCount, Error: task.Error, CreatedAt: task.CreatedAt.Format("2006-01-02T15:04:05Z07:00"), UpdatedAt: task.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"), } }