2f8e7ead17
1. 修复 buildURL 中 url.PathEscape 对路径分隔符 / 的错误编码(/→%2F), 改为逐段编码后拼接 2. AssetResponse 新增 key 字段,返回七牛云对象存储 Key 3. 前端改用 /api/v1/assets/download 接口加载图片,通过后端 302 重定向 访问 CDN,避免浏览器直接访问 CDN 可能出现的跨域或编码问题 4. 更新 api.md 文档,补充 key 字段说明和 download 接口公开属性
237 lines
6.0 KiB
Go
Executable File
237 lines
6.0 KiB
Go
Executable File
package handler
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"net/http"
|
|
"sync"
|
|
"time"
|
|
|
|
"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"`
|
|
}
|
|
|
|
// 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"`
|
|
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) {
|
|
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())
|
|
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,
|
|
},
|
|
})
|
|
|
|
// 返回 taskId
|
|
c.JSON(http.StatusOK, model.OK(GenerateResponse{TaskID: taskID}))
|
|
|
|
// 后台执行管线
|
|
go runPipelineBg(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)
|
|
})
|
|
|
|
l := logger.With("task_id", taskID, "project_id", projectID)
|
|
|
|
updateTaskProgress(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,
|
|
},
|
|
}
|
|
|
|
output, err := service.RunPipeline(ctx, in)
|
|
if err != nil {
|
|
l.Error("task pipeline failed", "error", err)
|
|
updateFailed(taskID, "生成管线执行失败")
|
|
return
|
|
}
|
|
|
|
updateTaskProgress(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, "上传素材失败")
|
|
return
|
|
}
|
|
assets[i] = AssetResponse{
|
|
Key: key,
|
|
URL: cdnURL,
|
|
Format: a.Format,
|
|
}
|
|
}
|
|
|
|
// 更新为完成状态
|
|
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
|
|
}
|
|
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)
|
|
}
|
|
|
|
// GetTask 查询任务信息。
|
|
func GetTask(c *gin.Context) {
|
|
taskID := c.Param("taskId")
|
|
rec, ok := taskStore.Load(taskID)
|
|
if !ok {
|
|
c.JSON(http.StatusNotFound, model.Fail(http.StatusNotFound, "任务不存在"))
|
|
return
|
|
}
|
|
r := rec.(*taskRecord)
|
|
c.JSON(http.StatusOK, model.OK(r.task))
|
|
}
|
|
|
|
// GetAssets 查询任务素材列表。
|
|
func GetAssets(c *gin.Context) {
|
|
taskID := c.Param("taskId")
|
|
rec, ok := taskStore.Load(taskID)
|
|
if !ok {
|
|
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))
|
|
return
|
|
}
|
|
c.JSON(http.StatusOK, model.OK(AssetsResponse{
|
|
Assets: r.assets,
|
|
Metadata: r.metadata,
|
|
}))
|
|
}
|