Merge pull request !60 from 何朝晖/develop
This commit is contained in:
+1
-1
@@ -25,7 +25,7 @@ func main() {
|
|||||||
gin.SetMode(cfg.Server.Mode)
|
gin.SetMode(cfg.Server.Mode)
|
||||||
|
|
||||||
// 初始化 SQLite 数据库
|
// 初始化 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)
|
slog.Error("db init failed", "error", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,11 +2,13 @@ package handler
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sync"
|
"strconv"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"gen2d/internal/db"
|
||||||
"gen2d/internal/logger"
|
"gen2d/internal/logger"
|
||||||
"gen2d/internal/model"
|
"gen2d/internal/model"
|
||||||
"gen2d/internal/service"
|
"gen2d/internal/service"
|
||||||
@@ -34,44 +36,12 @@ type GenerateResponse struct {
|
|||||||
TaskID string `json:"taskId"`
|
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 素材列表响应。
|
// AssetsResponse 素材列表响应。
|
||||||
type AssetsResponse struct {
|
type AssetsResponse struct {
|
||||||
Assets []AssetResponse `json:"assets"`
|
Assets []model.AssetResponse `json:"assets"`
|
||||||
Metadata service.AssetMetadata `json:"metadata"`
|
Metadata service.AssetMetadata `json:"metadata"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// taskRecord 内存中的任务记录。
|
|
||||||
type taskRecord struct {
|
|
||||||
task TaskResponse
|
|
||||||
assets []AssetResponse
|
|
||||||
metadata service.AssetMetadata
|
|
||||||
}
|
|
||||||
|
|
||||||
var (
|
|
||||||
taskStore = sync.Map{} // taskID → *taskRecord
|
|
||||||
)
|
|
||||||
|
|
||||||
// Generate 素材生成接口(异步)。
|
// Generate 素材生成接口(异步)。
|
||||||
// 立即返回 taskId,后台执行管线,前端通过 GET /tasks/:taskId 轮询进度。
|
// 立即返回 taskId,后台执行管线,前端通过 GET /tasks/:taskId 轮询进度。
|
||||||
func Generate(c *gin.Context) {
|
func Generate(c *gin.Context) {
|
||||||
@@ -86,38 +56,31 @@ func Generate(c *gin.Context) {
|
|||||||
projectID = "default"
|
projectID = "default"
|
||||||
}
|
}
|
||||||
taskID := fmt.Sprintf("task-%d", time.Now().UnixMilli())
|
taskID := fmt.Sprintf("task-%d", time.Now().UnixMilli())
|
||||||
createdAt := time.Now().Format(time.RFC3339)
|
|
||||||
|
|
||||||
// 存入 pending 状态
|
// 保存任务到数据库
|
||||||
taskStore.Store(taskID, &taskRecord{
|
if err := saveTaskToDB(c.Request.Context(), projectID, taskID, req); err != nil {
|
||||||
task: TaskResponse{
|
logger.FromCtx(c.Request.Context()).Error("failed to save task", "error", err)
|
||||||
ID: taskID,
|
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "创建任务失败"))
|
||||||
ProjectID: projectID,
|
return
|
||||||
Prompt: req.Prompt,
|
}
|
||||||
AssetType: req.AssetType,
|
|
||||||
Status: "pending",
|
|
||||||
Progress: 0,
|
|
||||||
CreatedAt: createdAt,
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
// 返回 taskId
|
// 返回 taskId
|
||||||
c.JSON(http.StatusOK, model.OK(GenerateResponse{TaskID: taskID}))
|
c.JSON(http.StatusOK, model.OK(GenerateResponse{TaskID: taskID}))
|
||||||
|
|
||||||
// 后台执行管线
|
// 后台执行管线
|
||||||
go runPipelineBg(projectID, taskID, req)
|
go runPipelineBg(c.Request.Context(), projectID, taskID, req)
|
||||||
}
|
}
|
||||||
|
|
||||||
// runPipelineBg 后台执行生成管线,更新任务状态。
|
// runPipelineBg 后台执行生成管线,更新任务状态。
|
||||||
func runPipelineBg(projectID, taskID string, req GenerateRequest) {
|
func runPipelineBg(ctx context.Context, projectID, taskID string, req GenerateRequest) {
|
||||||
// 注入进度上报回调
|
|
||||||
ctx := service.WithProgressReporter(context.Background(), func(stage string, progress int) {
|
|
||||||
updateTaskProgress(taskID, "running", stage, progress)
|
|
||||||
})
|
|
||||||
|
|
||||||
l := logger.With("task_id", taskID, "project_id", projectID)
|
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{
|
in := service.PipelineInput{
|
||||||
ProjectID: projectID,
|
ProjectID: projectID,
|
||||||
@@ -141,96 +104,187 @@ func runPipelineBg(projectID, taskID string, req GenerateRequest) {
|
|||||||
output, err := service.RunPipeline(ctx, in)
|
output, err := service.RunPipeline(ctx, in)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
l.Error("task pipeline failed", "error", err)
|
l.Error("task pipeline failed", "error", err)
|
||||||
updateFailed(taskID, "生成管线执行失败")
|
updateTaskInDB(ctx, taskID, "failed", "", err.Error(), 0)
|
||||||
return
|
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 {
|
for i, a := range output.Assets {
|
||||||
key := fmt.Sprintf("generation/%s/%s/%d.%s", projectID, taskID, i, a.Format)
|
key := fmt.Sprintf("generation/%s/%s/%d.%s", projectID, taskID, i, a.Format)
|
||||||
cdnURL, err := storageSvc.Upload(ctx, key, a.Data)
|
cdnURL, err := storageSvc.Upload(ctx, key, a.Data)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
l.Error("upload asset failed", "index", i, "error", err)
|
l.Error("upload asset failed", "index", i, "error", err)
|
||||||
updateFailed(taskID, "上传素材失败")
|
updateTaskInDB(ctx, taskID, "failed", "", "上传素材失败: "+err.Error(), 0)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
assets[i] = AssetResponse{
|
|
||||||
Key: key,
|
// 序列化单个素材的元数据
|
||||||
URL: cdnURL,
|
var metadata map[string]interface{}
|
||||||
Format: a.Format,
|
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{
|
var fullMetadata string
|
||||||
task: TaskResponse{
|
if metadataJSON, err := json.Marshal(output.Metadata); err == nil {
|
||||||
ID: taskID,
|
fullMetadata = string(metadataJSON)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
r := rec.(*taskRecord)
|
|
||||||
r.task.Status = status
|
|
||||||
r.task.Stage = stage
|
|
||||||
r.task.Progress = progress
|
|
||||||
taskStore.Store(taskID, r)
|
|
||||||
}
|
|
||||||
|
|
||||||
func updateFailed(taskID, errMsg string) {
|
db.GetDB().WithContext(ctx).Model(&model.Task{}).
|
||||||
rec, ok := taskStore.Load(taskID)
|
Where("external_id = ?", taskID).
|
||||||
if !ok {
|
Updates(map[string]interface{}{
|
||||||
return
|
"status": "completed",
|
||||||
}
|
"progress": 100,
|
||||||
r := rec.(*taskRecord)
|
"metadata": fullMetadata,
|
||||||
r.task.Status = "failed"
|
})
|
||||||
r.task.Error = errMsg
|
|
||||||
taskStore.Store(taskID, r)
|
l.Info("task completed")
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetTask 查询任务信息。
|
// GetTask 查询任务信息。
|
||||||
func GetTask(c *gin.Context) {
|
func GetTask(c *gin.Context) {
|
||||||
taskID := c.Param("taskId")
|
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, "任务不存在"))
|
c.JSON(http.StatusNotFound, model.Fail(http.StatusNotFound, "任务不存在"))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
r := rec.(*taskRecord)
|
|
||||||
c.JSON(http.StatusOK, model.OK(r.task))
|
c.JSON(http.StatusOK, model.OK(toTaskResponse(&task)))
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetAssets 查询任务素材列表。
|
// GetAssets 查询任务素材列表。
|
||||||
func GetAssets(c *gin.Context) {
|
func GetAssets(c *gin.Context) {
|
||||||
taskID := c.Param("taskId")
|
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, "任务不存在"))
|
c.JSON(http.StatusNotFound, model.Fail(http.StatusNotFound, "任务不存在"))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
r := rec.(*taskRecord)
|
|
||||||
if r.task.Status != "completed" {
|
if task.Status != "completed" {
|
||||||
c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "任务尚未完成,当前状态: "+r.task.Status))
|
c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "任务尚未完成,当前状态: "+task.Status))
|
||||||
return
|
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{
|
c.JSON(http.StatusOK, model.OK(AssetsResponse{
|
||||||
Assets: r.assets,
|
Assets: response,
|
||||||
Metadata: r.metadata,
|
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"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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"`
|
||||||
|
}
|
||||||
@@ -4,15 +4,17 @@ import "time"
|
|||||||
|
|
||||||
// Task 任务数据模型,对应数据库表。
|
// Task 任务数据模型,对应数据库表。
|
||||||
type Task struct {
|
type Task struct {
|
||||||
ID uint `gorm:"primaryKey" json:"id"`
|
ID uint `gorm:"primaryKey" json:"id"`
|
||||||
ProjectID uint `gorm:"index" json:"-"`
|
ExternalID string `gorm:"size:100;index" json:"externalId,omitempty"` // 外部任务ID,关联taskID
|
||||||
Prompt string `gorm:"type:text" json:"prompt"`
|
ProjectID uint `gorm:"index" json:"-"`
|
||||||
AssetType string `gorm:"size:50" json:"assetType"`
|
Prompt string `gorm:"type:text" json:"prompt"`
|
||||||
Status string `gorm:"size:20;default:'pending'" json:"status"`
|
AssetType string `gorm:"size:50" json:"assetType"`
|
||||||
Stage string `gorm:"size:50" json:"stage,omitempty"`
|
Status string `gorm:"size:20;default:'pending'" json:"status"`
|
||||||
Progress int `gorm:"default:0" json:"progress"`
|
Stage string `gorm:"size:50" json:"stage,omitempty"`
|
||||||
RetryCount int `gorm:"default:0" json:"retryCount"`
|
Progress int `gorm:"default:0" json:"progress"`
|
||||||
Error string `gorm:"type:text" json:"error,omitempty"`
|
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"`
|
CreatedAt time.Time `json:"createdAt"`
|
||||||
UpdatedAt time.Time `json:"updatedAt,omitempty"`
|
UpdatedAt time.Time `json:"updatedAt,omitempty"`
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user