Merge branch 'develop' of gitee.com:hezhaohui123/gen2d into feat/async-generate-pipeline
Signed-off-by: 何朝晖 <hezhaohui0807@163.com>
This commit is contained in:
@@ -4,6 +4,7 @@ import (
|
||||
"encoding/base64"
|
||||
"net/http"
|
||||
|
||||
"gen2d/internal/logger"
|
||||
"gen2d/internal/model"
|
||||
"gen2d/internal/service"
|
||||
|
||||
@@ -12,9 +13,9 @@ import (
|
||||
|
||||
// EditImageRequest 图片编辑请求。
|
||||
type EditImageRequest struct {
|
||||
Image string `json:"image" binding:"required"` // 底图 base64 编码
|
||||
Prompt string `json:"prompt" binding:"required"` // 编辑指令
|
||||
Count int `json:"count"` // 生成数量,默认 1
|
||||
Image string `json:"image" binding:"required"` // 底图 base64 编码
|
||||
Prompt string `json:"prompt" binding:"required"` // 编辑指令
|
||||
Count int `json:"count"` // 生成数量,默认 1
|
||||
}
|
||||
|
||||
// editAssetResponse 编辑结果素材(返回 base64)。
|
||||
@@ -38,7 +39,7 @@ func EditImage(c *gin.Context) {
|
||||
|
||||
imageData, err := base64.StdEncoding.DecodeString(req.Image)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "图片 base64 解码失败: "+err.Error()))
|
||||
c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "图片 base64 解码失败"))
|
||||
return
|
||||
}
|
||||
|
||||
@@ -49,7 +50,8 @@ func EditImage(c *gin.Context) {
|
||||
|
||||
images, err := service.EditImages(c.Request.Context(), imageData, req.Prompt, count)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "图片编辑失败: "+err.Error()))
|
||||
logger.FromCtx(c.Request.Context()).Error("图片编辑失败", "error", err)
|
||||
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "图片编辑失败"))
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -3,13 +3,11 @@ package handler
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gen2d/internal/logger"
|
||||
"gen2d/internal/model"
|
||||
"gen2d/internal/service"
|
||||
|
||||
@@ -58,8 +56,8 @@ type TaskResponse struct {
|
||||
|
||||
// AssetsResponse 素材列表响应。
|
||||
type AssetsResponse struct {
|
||||
Assets []AssetResponse `json:"assets"`
|
||||
Metadata service.AssetMetadata `json:"metadata"`
|
||||
Assets []AssetResponse `json:"assets"`
|
||||
Metadata service.AssetMetadata `json:"metadata"`
|
||||
}
|
||||
|
||||
// taskRecord 内存中的任务记录。
|
||||
@@ -106,7 +104,7 @@ func Generate(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, model.OK(GenerateResponse{TaskID: taskID}))
|
||||
|
||||
// 后台执行管线
|
||||
go runPipelineBg(projectID, taskID, req)
|
||||
go runPipelineBg(c.Request.Context(), projectID, taskID, req)
|
||||
}
|
||||
|
||||
// runPipelineBg 后台执行生成管线,更新任务状态。
|
||||
@@ -116,6 +114,8 @@ func runPipelineBg(projectID, taskID string, req GenerateRequest) {
|
||||
updateTaskProgress(taskID, "running", stage, progress)
|
||||
})
|
||||
|
||||
l := logger.With("task_id", taskID, "project_id", projectID)
|
||||
|
||||
updateTaskProgress(taskID, "running", "prompt_builder", 5)
|
||||
|
||||
in := service.PipelineInput{
|
||||
@@ -139,30 +139,24 @@ func runPipelineBg(projectID, taskID string, req GenerateRequest) {
|
||||
|
||||
output, err := service.RunPipeline(ctx, in)
|
||||
if err != nil {
|
||||
log.Printf("[generate] task %s failed: %v", taskID, err)
|
||||
updateFailed(taskID, err.Error())
|
||||
l.Error("task pipeline failed", "error", err)
|
||||
updateFailed(taskID, "生成管线执行失败")
|
||||
return
|
||||
}
|
||||
|
||||
updateTaskProgress(taskID, "saving", "format_adapter", 90)
|
||||
|
||||
// 保存图片到 ../generation/{projectId}/{taskId}/
|
||||
genDir := filepath.Join("..", "generation", projectID, taskID)
|
||||
if err := os.MkdirAll(genDir, 0755); err != nil {
|
||||
updateFailed(taskID, "创建输出目录失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
assets := make([]AssetResponse, len(output.Assets))
|
||||
for i, a := range output.Assets {
|
||||
filename := fmt.Sprintf("%d.%s", i, a.Format)
|
||||
filePath := filepath.Join(genDir, filename)
|
||||
if err := os.WriteFile(filePath, a.Data, 0644); err != nil {
|
||||
updateFailed(taskID, "保存图片失败: "+err.Error())
|
||||
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{
|
||||
URL: fmt.Sprintf("/generation/%s/%s/%s", projectID, taskID, filename),
|
||||
URL: cdnURL,
|
||||
Format: a.Format,
|
||||
}
|
||||
}
|
||||
@@ -182,7 +176,7 @@ func runPipelineBg(projectID, taskID string, req GenerateRequest) {
|
||||
metadata: output.Metadata,
|
||||
})
|
||||
|
||||
log.Printf("[generate] task %s completed, %d assets", taskID, len(assets))
|
||||
l.Info("task completed", "asset_count", len(assets))
|
||||
}
|
||||
|
||||
func updateTaskProgress(taskID, status, stage string, progress int) {
|
||||
|
||||
@@ -3,6 +3,7 @@ package handler
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"gen2d/internal/logger"
|
||||
"gen2d/internal/model"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -31,6 +32,7 @@ func Login(c *gin.Context) {
|
||||
|
||||
token, expiresIn, user, err := authSvc.Login(c.Request.Context(), req.Username, req.Password)
|
||||
if err != nil {
|
||||
logger.FromCtx(c.Request.Context()).Warn("登录失败", "username", req.Username, "error", err)
|
||||
c.JSON(http.StatusUnauthorized, model.Fail(http.StatusUnauthorized, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package handler
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"gen2d/internal/logger"
|
||||
"gen2d/internal/model"
|
||||
"gen2d/internal/service"
|
||||
|
||||
@@ -34,7 +35,8 @@ func PromptOptimize(c *gin.Context) {
|
||||
|
||||
output, err := service.RunPromptAgent(c.Request.Context(), in)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "提示词优化失败: "+err.Error()))
|
||||
logger.FromCtx(c.Request.Context()).Error("提示词优化失败", "error", err)
|
||||
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "提示词优化失败"))
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -3,6 +3,7 @@ package handler
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"gen2d/internal/logger"
|
||||
"gen2d/internal/model"
|
||||
"gen2d/internal/service"
|
||||
|
||||
@@ -39,6 +40,7 @@ func Register(c *gin.Context) {
|
||||
|
||||
user, err := authSvc.Register(c.Request.Context(), req.Username, req.Password, req.Email)
|
||||
if err != nil {
|
||||
logger.FromCtx(c.Request.Context()).Warn("注册失败", "username", req.Username, "error", err)
|
||||
c.JSON(http.StatusConflict, model.Fail(http.StatusConflict, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"gen2d/internal/logger"
|
||||
"gen2d/internal/model"
|
||||
"gen2d/internal/service"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
var storageSvc *service.StorageService
|
||||
|
||||
// InitStorageService 由 main 在启动时调用,注入七牛云存储服务。
|
||||
func InitStorageService(svc *service.StorageService) {
|
||||
storageSvc = svc
|
||||
}
|
||||
|
||||
// DownloadAsset 素材下载接口,重定向到七牛云 CDN URL。
|
||||
// GET /api/v1/assets/download?key=...
|
||||
func DownloadAsset(c *gin.Context) {
|
||||
key := c.Query("key")
|
||||
if key == "" {
|
||||
c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "缺少 key 参数"))
|
||||
return
|
||||
}
|
||||
|
||||
downloadURL, err := storageSvc.GetDownloadURL(c.Request.Context(), key)
|
||||
if err != nil {
|
||||
logger.FromCtx(c.Request.Context()).Error("生成下载链接失败", "key", key, "error", err)
|
||||
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "生成下载链接失败"))
|
||||
return
|
||||
}
|
||||
|
||||
c.Redirect(http.StatusFound, downloadURL)
|
||||
}
|
||||
Reference in New Issue
Block a user