diff --git a/backend/cmd/main.go b/backend/cmd/main.go index 820646d..967a175 100755 --- a/backend/cmd/main.go +++ b/backend/cmd/main.go @@ -25,7 +25,7 @@ func main() { gin.SetMode(cfg.Server.Mode) // 初始化 SQLite 数据库 - if err := db.Init(cfg.Database.DSN, &model.User{}, &model.Project{}, &model.ProjectStyleRecord{}, &model.Task{}); err != nil { + if err := db.Init(cfg.Database.DSN, &model.User{}, &model.Project{}, &model.ProjectStyleRecord{}, &model.Task{}, &model.Asset{}); err != nil { slog.Error("db init failed", "error", err) return } diff --git a/backend/internal/handler/generate.go b/backend/internal/handler/generate.go index 7ec2df0..0c5864b 100755 --- a/backend/internal/handler/generate.go +++ b/backend/internal/handler/generate.go @@ -2,11 +2,13 @@ package handler import ( "context" + "encoding/json" "fmt" "net/http" - "sync" + "strconv" "time" + "gen2d/internal/db" "gen2d/internal/logger" "gen2d/internal/model" "gen2d/internal/service" @@ -34,44 +36,12 @@ type GenerateResponse struct { TaskID string `json:"taskId"` } -// AssetResponse 单个素材响应。 -type AssetResponse struct { - Key string `json:"key"` - URL string `json:"url"` - Format string `json:"format"` -} - -// TaskResponse 任务查询响应。 -type TaskResponse struct { - ID string `json:"id"` - ProjectID string `json:"projectId"` - Prompt string `json:"prompt"` - AssetType string `json:"assetType"` - Status string `json:"status"` - Stage string `json:"stage,omitempty"` - Progress int `json:"progress"` - RetryCount int `json:"retryCount"` - Error string `json:"error,omitempty"` - CreatedAt string `json:"createdAt"` -} - // AssetsResponse 素材列表响应。 type AssetsResponse struct { - Assets []AssetResponse `json:"assets"` + Assets []model.AssetResponse `json:"assets"` Metadata service.AssetMetadata `json:"metadata"` } -// taskRecord 内存中的任务记录。 -type taskRecord struct { - task TaskResponse - assets []AssetResponse - metadata service.AssetMetadata -} - -var ( - taskStore = sync.Map{} // taskID → *taskRecord -) - // Generate 素材生成接口(异步)。 // 立即返回 taskId,后台执行管线,前端通过 GET /tasks/:taskId 轮询进度。 func Generate(c *gin.Context) { @@ -86,38 +56,31 @@ func Generate(c *gin.Context) { projectID = "default" } taskID := fmt.Sprintf("task-%d", time.Now().UnixMilli()) - createdAt := time.Now().Format(time.RFC3339) - // 存入 pending 状态 - taskStore.Store(taskID, &taskRecord{ - task: TaskResponse{ - ID: taskID, - ProjectID: projectID, - Prompt: req.Prompt, - AssetType: req.AssetType, - Status: "pending", - Progress: 0, - CreatedAt: createdAt, - }, - }) + // 保存任务到数据库 + 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(projectID, taskID, req) + go runPipelineBg(c.Request.Context(), projectID, taskID, req) } // runPipelineBg 后台执行生成管线,更新任务状态。 -func runPipelineBg(projectID, taskID string, req GenerateRequest) { - // 注入进度上报回调 - ctx := service.WithProgressReporter(context.Background(), func(stage string, progress int) { - updateTaskProgress(taskID, "running", stage, progress) - }) - +func runPipelineBg(ctx context.Context, projectID, taskID string, req GenerateRequest) { l := logger.With("task_id", taskID, "project_id", projectID) - updateTaskProgress(taskID, "running", "prompt_builder", 5) + // 注入进度上报回调 + ctx = service.WithProgressReporter(ctx, func(stage string, progress int) { + updateTaskInDB(ctx, taskID, "running", stage, "", progress) + }) + + updateTaskInDB(ctx, taskID, "running", "prompt_builder", "", 5) in := service.PipelineInput{ ProjectID: projectID, @@ -141,96 +104,187 @@ func runPipelineBg(projectID, taskID string, req GenerateRequest) { output, err := service.RunPipeline(ctx, in) if err != nil { l.Error("task pipeline failed", "error", err) - updateFailed(taskID, "生成管线执行失败") + updateTaskInDB(ctx, taskID, "failed", "", err.Error(), 0) return } - updateTaskProgress(taskID, "saving", "format_adapter", 90) + updateTaskInDB(ctx, taskID, "saving", "format_adapter", "", 90) - assets := make([]AssetResponse, len(output.Assets)) + // 上传素材并保存到数据库 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) - updateFailed(taskID, "上传素材失败") + updateTaskInDB(ctx, taskID, "failed", "", "上传素材失败: "+err.Error(), 0) return } - assets[i] = AssetResponse{ - Key: key, - URL: cdnURL, - Format: a.Format, + + // 序列化单个素材的元数据 + 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) } } // 更新为完成状态 - taskStore.Store(taskID, &taskRecord{ - task: TaskResponse{ - ID: taskID, - ProjectID: projectID, - Prompt: req.Prompt, - AssetType: req.AssetType, - Status: "completed", - Progress: 100, - CreatedAt: time.Now().Format(time.RFC3339), - }, - assets: assets, - metadata: output.Metadata, - }) - - l.Info("task completed", "asset_count", len(assets)) -} - -func updateTaskProgress(taskID, status, stage string, progress int) { - rec, ok := taskStore.Load(taskID) - if !ok { - return + var fullMetadata string + if metadataJSON, err := json.Marshal(output.Metadata); err == nil { + fullMetadata = string(metadataJSON) } - r := rec.(*taskRecord) - r.task.Status = status - r.task.Stage = stage - r.task.Progress = progress - taskStore.Store(taskID, r) -} -func updateFailed(taskID, errMsg string) { - rec, ok := taskStore.Load(taskID) - if !ok { - return - } - r := rec.(*taskRecord) - r.task.Status = "failed" - r.task.Error = errMsg - taskStore.Store(taskID, r) + 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") - rec, ok := taskStore.Load(taskID) - if !ok { + + 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 } - r := rec.(*taskRecord) - c.JSON(http.StatusOK, model.OK(r.task)) + + c.JSON(http.StatusOK, model.OK(toTaskResponse(&task))) } // GetAssets 查询任务素材列表。 func GetAssets(c *gin.Context) { taskID := c.Param("taskId") - rec, ok := taskStore.Load(taskID) - if !ok { + + 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 } - r := rec.(*taskRecord) - if r.task.Status != "completed" { - c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "任务尚未完成,当前状态: "+r.task.Status)) + + 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: r.assets, - Metadata: r.metadata, + 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: "pending", + Progress: 0, + 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"), + } +} diff --git a/backend/internal/model/asset.go b/backend/internal/model/asset.go new file mode 100644 index 0000000..c9c12b1 --- /dev/null +++ b/backend/internal/model/asset.go @@ -0,0 +1,22 @@ +package model + +import "time" + +// Asset 素材数据模型,对应数据库表。 +type Asset struct { + ID uint `gorm:"primaryKey"` + TaskID uint `gorm:"index;not null"` + Key string `gorm:"size:255;not null"` // 七牛云对象键 + URL string `gorm:"size:255;not null"` // CDN访问URL + Format string `gorm:"size:20;not null"` // 文件格式 + Metadata string `gorm:"type:text"` // 元数据(JSON格式) + CreatedAt time.Time + UpdatedAt time.Time +} + +// AssetResponse 素材响应。 +type AssetResponse struct { + Key string `json:"key"` + URL string `json:"url"` + Format string `json:"format"` +} diff --git a/backend/internal/model/task.go b/backend/internal/model/task.go index a62ab60..a2160d0 100644 --- a/backend/internal/model/task.go +++ b/backend/internal/model/task.go @@ -4,15 +4,17 @@ import "time" // Task 任务数据模型,对应数据库表。 type Task struct { - ID uint `gorm:"primaryKey" json:"id"` - ProjectID uint `gorm:"index" json:"-"` - Prompt string `gorm:"type:text" json:"prompt"` - AssetType string `gorm:"size:50" json:"assetType"` - Status string `gorm:"size:20;default:'pending'" json:"status"` - Stage string `gorm:"size:50" json:"stage,omitempty"` - Progress int `gorm:"default:0" json:"progress"` - RetryCount int `gorm:"default:0" json:"retryCount"` - Error string `gorm:"type:text" json:"error,omitempty"` + ID uint `gorm:"primaryKey" json:"id"` + ExternalID string `gorm:"size:100;index" json:"externalId,omitempty"` // 外部任务ID,关联taskID + ProjectID uint `gorm:"index" json:"-"` + Prompt string `gorm:"type:text" json:"prompt"` + AssetType string `gorm:"size:50" json:"assetType"` + Status string `gorm:"size:20;default:'pending'" json:"status"` + Stage string `gorm:"size:50" json:"stage,omitempty"` + Progress int `gorm:"default:0" json:"progress"` + RetryCount int `gorm:"default:0" json:"retryCount"` + Error string `gorm:"type:text" json:"error,omitempty"` + Metadata string `gorm:"type:text" json:"-"` // 元数据(JSON格式) CreatedAt time.Time `json:"createdAt"` UpdatedAt time.Time `json:"updatedAt,omitempty"` }