!60 !58 refactor: 移除内存存储,全部改为数据库持久化

Merge pull request !60 from 何朝晖/develop
This commit is contained in:
2026-05-25 10:04:02 +00:00
committed by Gitee
4 changed files with 198 additions and 120 deletions
+1 -1
View File
@@ -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
}
+164 -110
View File
@@ -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"),
}
}
+22
View File
@@ -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"`
}
+11 -9
View File
@@ -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"`
}