From bb5d3431c19fb06780c221b437ddc8f76707171d Mon Sep 17 00:00:00 2001 From: wonder Date: Mon, 25 May 2026 17:48:31 +0800 Subject: [PATCH 1/2] =?UTF-8?q?fix:=20=E4=BF=AE=E5=A4=8D=E6=96=B0=E5=BB=BA?= =?UTF-8?q?=E4=BB=BB=E5=8A=A1=E5=90=8E=E4=BB=BB=E5=8A=A1=E5=88=97=E8=A1=A8?= =?UTF-8?q?=E8=BF=94=E5=9B=9E=E4=B8=BA=E7=A9=BA=E7=9A=84=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 在 Task 模型中添加 ExternalID 字段用于关联内存中的 taskID - 修改 Generate 函数,将任务记录同时存入内存和数据库 - 添加 saveTaskToDB 和 updateTaskInDB 函数处理任务持久化 - 同步更新任务状态到数据库,确保 GetProjectTasks 能正确查询 - 处理 default 工程的特殊情况(跳过数据库存储) --- backend/internal/handler/generate.go | 54 ++++++++++++++++++++++++++++ backend/internal/model/task.go | 19 +++++----- 2 files changed, 64 insertions(+), 9 deletions(-) diff --git a/backend/internal/handler/generate.go b/backend/internal/handler/generate.go index 7ec2df0..458eb4a 100755 --- a/backend/internal/handler/generate.go +++ b/backend/internal/handler/generate.go @@ -4,9 +4,11 @@ import ( "context" "fmt" "net/http" + "strconv" "sync" "time" + "gen2d/internal/db" "gen2d/internal/logger" "gen2d/internal/model" "gen2d/internal/service" @@ -101,6 +103,9 @@ func Generate(c *gin.Context) { }, }) + // 将任务记录写入数据库(用于任务列表查询) + go saveTaskToDB(c.Request.Context(), projectID, taskID, req) + // 返回 taskId c.JSON(http.StatusOK, model.OK(GenerateResponse{TaskID: taskID})) @@ -191,6 +196,9 @@ func updateTaskProgress(taskID, status, stage string, progress int) { r.task.Stage = stage r.task.Progress = progress taskStore.Store(taskID, r) + + // 同步更新数据库状态 + updateTaskInDB(taskID, status, stage, "", progress) } func updateFailed(taskID, errMsg string) { @@ -202,6 +210,9 @@ func updateFailed(taskID, errMsg string) { r.task.Status = "failed" r.task.Error = errMsg taskStore.Store(taskID, r) + + // 同步更新数据库状态 + updateTaskInDB(taskID, "failed", r.task.Stage, errMsg, r.task.Progress) } // GetTask 查询任务信息。 @@ -234,3 +245,46 @@ func GetAssets(c *gin.Context) { Metadata: r.metadata, })) } + +// saveTaskToDB 将任务记录保存到数据库。 +func saveTaskToDB(ctx context.Context, projectID, taskID string, req GenerateRequest) { + // "default" 工程没有数据库记录,跳过 + if projectID == "default" { + return + } + + projectIDUint, err := strconv.ParseUint(projectID, 10, 32) + if err != nil { + logger.FromCtx(ctx).Error("failed to parse projectID", "projectID", projectID, "error", err) + return + } + + 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(), + } + + if err := db.GetDB().WithContext(ctx).Create(task).Error; err != nil { + logger.FromCtx(ctx).Error("failed to save task to db", "taskID", taskID, "error", err) + } +} + +// updateTaskInDB 更新数据库中的任务状态。 +func updateTaskInDB(taskID string, status, stage, error string, progress int) { + db.GetDB().Model(&model.Task{}). + Where("external_id = ?", taskID). + Updates(map[string]interface{}{ + "status": status, + "stage": stage, + "progress": progress, + "error": error, + "updated_at": time.Now(), + }) +} diff --git a/backend/internal/model/task.go b/backend/internal/model/task.go index a62ab60..a34558b 100644 --- a/backend/internal/model/task.go +++ b/backend/internal/model/task.go @@ -4,15 +4,16 @@ 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"` CreatedAt time.Time `json:"createdAt"` UpdatedAt time.Time `json:"updatedAt,omitempty"` } From 3081066d35a52dc7f3c8a2d3539103a584176053 Mon Sep 17 00:00:00 2001 From: wonder Date: Mon, 25 May 2026 17:54:18 +0800 Subject: [PATCH 2/2] =?UTF-8?q?refactor:=20=E7=A7=BB=E9=99=A4=E5=86=85?= =?UTF-8?q?=E5=AD=98=E5=AD=98=E5=82=A8=EF=BC=8C=E5=85=A8=E9=83=A8=E6=94=B9?= =?UTF-8?q?=E4=B8=BA=E6=95=B0=E6=8D=AE=E5=BA=93=E6=8C=81=E4=B9=85=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新建 Asset 模型并添加自动迁移 - 移除 taskStore 内存中的任务缓存 - 新任务直接写入数据库 Task 表 - 任务状态更新同步写入数据库 - 素材上传结果保存到 Asset 表 - GetTask/GetAssets 从数据库查询 - 添加 Metadata 字段到 Task 模型存储完整元数据 --- backend/cmd/main.go | 2 +- backend/internal/handler/generate.go | 268 +++++++++++++-------------- backend/internal/model/asset.go | 22 +++ backend/internal/model/task.go | 3 +- 4 files changed, 159 insertions(+), 136 deletions(-) create mode 100644 backend/internal/model/asset.go 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 458eb4a..0c5864b 100755 --- a/backend/internal/handler/generate.go +++ b/backend/internal/handler/generate.go @@ -2,10 +2,10 @@ package handler import ( "context" + "encoding/json" "fmt" "net/http" "strconv" - "sync" "time" "gen2d/internal/db" @@ -36,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) { @@ -88,41 +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, - }, - }) - - // 将任务记录写入数据库(用于任务列表查询) - go saveTaskToDB(c.Request.Context(), projectID, taskID, req) + // 保存任务到数据库 + 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, @@ -146,117 +104,131 @@ 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) - // 同步更新数据库状态 - updateTaskInDB(taskID, status, stage, "", progress) -} + db.GetDB().WithContext(ctx).Model(&model.Task{}). + Where("external_id = ?", taskID). + Updates(map[string]interface{}{ + "status": "completed", + "progress": 100, + "metadata": fullMetadata, + }) -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) - - // 同步更新数据库状态 - updateTaskInDB(taskID, "failed", r.task.Stage, errMsg, r.task.Progress) + 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) { +func saveTaskToDB(ctx context.Context, projectID, taskID string, req GenerateRequest) error { // "default" 工程没有数据库记录,跳过 if projectID == "default" { - return + return nil } projectIDUint, err := strconv.ParseUint(projectID, 10, 32) if err != nil { - logger.FromCtx(ctx).Error("failed to parse projectID", "projectID", projectID, "error", err) - return + return fmt.Errorf("failed to parse projectID: %w", err) } task := &model.Task{ @@ -271,20 +243,48 @@ func saveTaskToDB(ctx context.Context, projectID, taskID string, req GenerateReq UpdatedAt: time.Now(), } - if err := db.GetDB().WithContext(ctx).Create(task).Error; err != nil { - logger.FromCtx(ctx).Error("failed to save task to db", "taskID", taskID, "error", err) - } + return db.GetDB().WithContext(ctx).Create(task).Error } // updateTaskInDB 更新数据库中的任务状态。 -func updateTaskInDB(taskID string, status, stage, error string, progress int) { - db.GetDB().Model(&model.Task{}). +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(map[string]interface{}{ - "status": status, - "stage": stage, - "progress": progress, - "error": error, - "updated_at": time.Now(), - }) + 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 a34558b..a2160d0 100644 --- a/backend/internal/model/task.go +++ b/backend/internal/model/task.go @@ -5,7 +5,7 @@ import "time" // Task 任务数据模型,对应数据库表。 type Task struct { ID uint `gorm:"primaryKey" json:"id"` - ExternalID string `gorm:"size:100;index" json:"externalId,omitempty"` // 外部任务ID,关联内存中的taskID + 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"` @@ -14,6 +14,7 @@ type Task struct { 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"` }