!34 refactor(config): 精简 image_gen 配置为 OpenAI 兼容同步 API
Merge pull request !34 from 郭永昊/feat/AssertGeneration
This commit is contained in:
Regular → Executable
+3
@@ -23,3 +23,6 @@ backend/bin/
|
||||
backend/data/
|
||||
backend/.env
|
||||
backend/main
|
||||
|
||||
# Generated output
|
||||
generation/
|
||||
|
||||
Regular → Executable
+5
-4
@@ -19,12 +19,13 @@ GEN2D_LLM_MODEL=gpt-4o
|
||||
GEN2D_LLM_TEMPERATURE=0.7
|
||||
GEN2D_LLM_MAX_TOKENS=2048
|
||||
|
||||
# 文生图模型
|
||||
GEN2D_IMAGE_BASE_URL=https://api.stability.ai/v1
|
||||
GEN2D_IMAGE_API_KEY=sk-your-api-key
|
||||
GEN2D_IMAGE_MODEL=stable-diffusion-xl
|
||||
# 文生图模型(OpenAI 兼容 Images API)
|
||||
GEN2D_IMAGE_BASE_URL=https://api.suchuang.vip/v1
|
||||
GEN2D_IMAGE_API_KEY=
|
||||
GEN2D_IMAGE_MODEL=gpt-image-2-token
|
||||
GEN2D_IMAGE_WIDTH=1024
|
||||
GEN2D_IMAGE_HEIGHT=1024
|
||||
GEN2D_IMAGE_QUALITY=low
|
||||
GEN2D_IMAGE_NUM_IMAGES=1
|
||||
GEN2D_IMAGE_STEPS=30
|
||||
GEN2D_IMAGE_CFG_SCALE=7.0
|
||||
|
||||
Executable
+2
@@ -0,0 +1,2 @@
|
||||
test_output/
|
||||
generation/
|
||||
Regular → Executable
+19
-1
@@ -4,10 +4,12 @@ package main
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
|
||||
"gen2d/internal/config"
|
||||
"gen2d/internal/db"
|
||||
"gen2d/internal/handler"
|
||||
"gen2d/internal/mildware"
|
||||
"gen2d/internal/model"
|
||||
"gen2d/internal/service"
|
||||
|
||||
@@ -19,6 +21,9 @@ func main() {
|
||||
|
||||
gin.SetMode(cfg.Server.Mode)
|
||||
|
||||
// 确保 generation 输出目录存在(项目根级别)
|
||||
_ = os.MkdirAll("../generation", 0755)
|
||||
|
||||
// 初始化 SQLite 数据库
|
||||
if err := db.Init(cfg.Database.DSN, &model.User{}); err != nil {
|
||||
log.Fatalf("db init failed: %v", err)
|
||||
@@ -38,7 +43,10 @@ func main() {
|
||||
r := gin.New()
|
||||
r.Use(gin.Recovery()) // panic 恢复中间件,防止服务因未捕获异常宕机
|
||||
|
||||
// API v1 路由组
|
||||
// 静态文件服务 — 生成的图片
|
||||
r.Static("/generation", "../generation")
|
||||
|
||||
// API v1 路由组 — 公开接口
|
||||
v1 := r.Group("/api/v1")
|
||||
{
|
||||
v1.GET("/health", handler.Health) // 健康检查
|
||||
@@ -46,6 +54,16 @@ func main() {
|
||||
v1.GET("/assets/download", handler.DownloadAsset) // 素材下载
|
||||
}
|
||||
|
||||
// API v1 路由组 — 需认证
|
||||
v1Auth := r.Group("/api/v1")
|
||||
v1Auth.Use(mildware.AuthMiddleware(cfg.JWT.Secret))
|
||||
{
|
||||
v1Auth.POST("/generate", handler.Generate) // 素材生成管线
|
||||
v1Auth.GET("/tasks/:taskId", handler.GetTask) // 查询任务
|
||||
v1Auth.GET("/tasks/:taskId/assets", handler.GetAssets) // 查询任务素材
|
||||
v1Auth.POST("/images/edit", handler.EditImage) // 图片编辑
|
||||
}
|
||||
|
||||
// Auth 路由组
|
||||
auth := r.Group("/auth")
|
||||
{
|
||||
|
||||
Regular → Executable
+11
-11
@@ -44,16 +44,14 @@ type LLMConfig struct {
|
||||
MaxTokens int `mapstructure:"max_tokens"`
|
||||
}
|
||||
|
||||
// ImageGenConfig 文生图模型配置。
|
||||
// ImageGenConfig OpenAI 兼容文生图模型配置。
|
||||
type ImageGenConfig struct {
|
||||
BaseURL string `mapstructure:"base_url"`
|
||||
APIKey string `mapstructure:"api_key"`
|
||||
Model string `mapstructure:"model"`
|
||||
Width int `mapstructure:"width"`
|
||||
Height int `mapstructure:"height"`
|
||||
NumImages int `mapstructure:"num_images"`
|
||||
Steps int `mapstructure:"steps"`
|
||||
CFGScale float64 `mapstructure:"cfg_scale"`
|
||||
BaseURL string `mapstructure:"base_url"`
|
||||
APIKey string `mapstructure:"api_key"`
|
||||
Model string `mapstructure:"model"`
|
||||
Width int `mapstructure:"width"`
|
||||
Height int `mapstructure:"height"`
|
||||
Quality string `mapstructure:"quality"`
|
||||
}
|
||||
|
||||
// QiniuConfig 七牛云对象存储配置。
|
||||
@@ -113,11 +111,12 @@ func setDefaults(v *viper.Viper) {
|
||||
v.SetDefault("llm.temperature", 0.7)
|
||||
v.SetDefault("llm.max_tokens", 2048)
|
||||
|
||||
v.SetDefault("image_gen.base_url", "https://api.stability.ai/v1")
|
||||
v.SetDefault("image_gen.base_url", "https://api.suchuang.vip/v1")
|
||||
v.SetDefault("image_gen.api_key", "")
|
||||
v.SetDefault("image_gen.model", "stable-diffusion-xl")
|
||||
v.SetDefault("image_gen.model", "gpt-image-2-token")
|
||||
v.SetDefault("image_gen.width", 1024)
|
||||
v.SetDefault("image_gen.height", 1024)
|
||||
v.SetDefault("image_gen.quality", "low")
|
||||
v.SetDefault("image_gen.num_images", 1)
|
||||
v.SetDefault("image_gen.steps", 30)
|
||||
v.SetDefault("image_gen.cfg_scale", 7.0)
|
||||
@@ -148,6 +147,7 @@ func bindEnvVars(v *viper.Viper) {
|
||||
v.BindEnv("image_gen.model", "GEN2D_IMAGE_MODEL")
|
||||
v.BindEnv("image_gen.width", "GEN2D_IMAGE_WIDTH")
|
||||
v.BindEnv("image_gen.height", "GEN2D_IMAGE_HEIGHT")
|
||||
v.BindEnv("image_gen.quality", "GEN2D_IMAGE_QUALITY")
|
||||
v.BindEnv("image_gen.num_images", "GEN2D_IMAGE_NUM_IMAGES")
|
||||
v.BindEnv("image_gen.steps", "GEN2D_IMAGE_STEPS")
|
||||
v.BindEnv("image_gen.cfg_scale", "GEN2D_IMAGE_CFG_SCALE")
|
||||
|
||||
Regular → Executable
+3
-5
@@ -21,11 +21,9 @@ llm:
|
||||
max_tokens: 2048
|
||||
|
||||
image_gen:
|
||||
base_url: "https://api.stability.ai/v1"
|
||||
base_url: "https://api.suchuang.vip/v1"
|
||||
api_key: ""
|
||||
model: "stable-diffusion-xl"
|
||||
model: "gpt-image-2-token"
|
||||
width: 1024
|
||||
height: 1024
|
||||
num_images: 1
|
||||
steps: 30
|
||||
cfg_scale: 7.0
|
||||
quality: "low"
|
||||
|
||||
Executable
+65
@@ -0,0 +1,65 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"net/http"
|
||||
|
||||
"gen2d/internal/model"
|
||||
"gen2d/internal/service"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// EditImageRequest 图片编辑请求。
|
||||
type EditImageRequest struct {
|
||||
Image string `json:"image" binding:"required"` // 底图 base64 编码
|
||||
Prompt string `json:"prompt" binding:"required"` // 编辑指令
|
||||
Count int `json:"count"` // 生成数量,默认 1
|
||||
}
|
||||
|
||||
// editAssetResponse 编辑结果素材(返回 base64)。
|
||||
type editAssetResponse struct {
|
||||
Data string `json:"data"`
|
||||
Format string `json:"format"`
|
||||
}
|
||||
|
||||
// EditImageResponse 图片编辑响应体。
|
||||
type EditImageResponse struct {
|
||||
Assets []editAssetResponse `json:"assets"`
|
||||
}
|
||||
|
||||
// EditImage 图片编辑接口,基于已有图片和文本指令生成修改后的图片。
|
||||
func EditImage(c *gin.Context) {
|
||||
var req EditImageRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "参数错误: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
imageData, err := base64.StdEncoding.DecodeString(req.Image)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "图片 base64 解码失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
count := req.Count
|
||||
if count <= 0 {
|
||||
count = 1
|
||||
}
|
||||
|
||||
images, err := service.EditImages(c.Request.Context(), imageData, req.Prompt, count)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "图片编辑失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
assets := make([]editAssetResponse, len(images))
|
||||
for i, img := range images {
|
||||
assets[i] = editAssetResponse{
|
||||
Data: base64.StdEncoding.EncodeToString(img.Data),
|
||||
Format: img.Format,
|
||||
}
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, model.OK(EditImageResponse{Assets: assets}))
|
||||
}
|
||||
Executable
+234
@@ -0,0 +1,234 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"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 {
|
||||
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"`
|
||||
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) {
|
||||
updateStatus(taskID, "running", 10)
|
||||
|
||||
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,
|
||||
},
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
output, err := service.RunPipeline(ctx, in)
|
||||
if err != nil {
|
||||
log.Printf("[generate] task %s failed: %v", taskID, err)
|
||||
updateFailed(taskID, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
updateStatus(taskID, "saving", 80)
|
||||
|
||||
// 保存图片到 ../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())
|
||||
return
|
||||
}
|
||||
assets[i] = AssetResponse{
|
||||
URL: fmt.Sprintf("/generation/%s/%s/%s", projectID, taskID, filename),
|
||||
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,
|
||||
})
|
||||
|
||||
log.Printf("[generate] task %s completed, %d assets", taskID, len(assets))
|
||||
}
|
||||
|
||||
func updateStatus(taskID, status string, progress int) {
|
||||
rec, ok := taskStore.Load(taskID)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
r := rec.(*taskRecord)
|
||||
r.task.Status = status
|
||||
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,
|
||||
}))
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package mildware
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"gen2d/internal/model"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
// AuthMiddleware 返回 JWT 认证中间件,校验 Bearer token 并注入 userID 到上下文。
|
||||
func AuthMiddleware(jwtSecret string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if authHeader == "" {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, model.Fail(http.StatusUnauthorized, "未提供认证令牌"))
|
||||
return
|
||||
}
|
||||
|
||||
tokenString := strings.TrimPrefix(authHeader, "Bearer ")
|
||||
if tokenString == authHeader {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, model.Fail(http.StatusUnauthorized, "认证格式错误,需为 Bearer <token>"))
|
||||
return
|
||||
}
|
||||
|
||||
token, err := jwt.Parse(tokenString, func(token *jwt.Token) (interface{}, error) {
|
||||
return []byte(jwtSecret), nil
|
||||
})
|
||||
if err != nil || !token.Valid {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, model.Fail(http.StatusUnauthorized, "令牌无效或已过期"))
|
||||
return
|
||||
}
|
||||
|
||||
claims, ok := token.Claims.(jwt.MapClaims)
|
||||
if !ok {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, model.Fail(http.StatusUnauthorized, "令牌解析失败"))
|
||||
return
|
||||
}
|
||||
|
||||
if sub, ok := claims["sub"]; ok {
|
||||
c.Set("userID", sub)
|
||||
}
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
Regular → Executable
+198
-30
@@ -4,10 +4,17 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"image"
|
||||
"image/color"
|
||||
"image/png"
|
||||
"io"
|
||||
"log"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"gen2d/internal/config"
|
||||
)
|
||||
@@ -20,50 +27,204 @@ func InitImageGenConfig(cfg config.ImageGenConfig) {
|
||||
imgCfg = cfg
|
||||
}
|
||||
|
||||
// GenerateImages 调用 AI 推理 API 生成图片。
|
||||
// MVP 阶段返回 mock 占位图。
|
||||
func GenerateImages(ctx context.Context, prompt string, params AssetParams) ([]GeneratedImage, error) {
|
||||
size := params.Resolution
|
||||
if size <= 0 {
|
||||
size = 64
|
||||
}
|
||||
// ======================== OpenAI 兼容 Images API 类型 ========================
|
||||
|
||||
// imageGenRequest OpenAI 兼容文生图请求体。
|
||||
type imageGenRequest struct {
|
||||
Model string `json:"model"`
|
||||
Prompt string `json:"prompt"`
|
||||
N int `json:"n,omitempty"`
|
||||
Size string `json:"size,omitempty"`
|
||||
Quality string `json:"quality,omitempty"`
|
||||
}
|
||||
|
||||
// imageGenResponse OpenAI 兼容文生图响应体。
|
||||
type imageGenResponse struct {
|
||||
Data []struct {
|
||||
URL string `json:"url"`
|
||||
B64JSON string `json:"b64_json"`
|
||||
} `json:"data"`
|
||||
}
|
||||
|
||||
// ======================== 文生图 ========================
|
||||
|
||||
// GenerateImages 调用 OpenAI 兼容 Images API 生成图片,未配置 key 时回退 mock。
|
||||
func GenerateImages(ctx context.Context, prompt string, params AssetParams) ([]GeneratedImage, error) {
|
||||
count := 1
|
||||
if params.Frames.Directions > 0 && params.Frames.FramesPerDirection > 0 {
|
||||
count = params.Frames.Directions * params.Frames.FramesPerDirection
|
||||
}
|
||||
|
||||
images := make([]GeneratedImage, count)
|
||||
for i := 0; i < count; i++ {
|
||||
data, err := generateMockImage(size, i)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("generate mock image %d: %w", i, err)
|
||||
if imgCfg.APIKey != "" {
|
||||
width, height := imgCfg.Width, imgCfg.Height
|
||||
if params.Resolution > 0 {
|
||||
width = params.Resolution
|
||||
height = params.Resolution
|
||||
}
|
||||
images[i] = GeneratedImage{
|
||||
log.Println("[inference] calling image gen API")
|
||||
return callImageAPI(ctx, prompt, count, width, height)
|
||||
}
|
||||
|
||||
size := params.Resolution
|
||||
if size <= 0 {
|
||||
size = 64
|
||||
}
|
||||
log.Println("[inference] image API key not configured, using mock")
|
||||
return generateMockImages(size, count)
|
||||
}
|
||||
|
||||
// callImageAPI 调用 OpenAI 兼容 Images API。
|
||||
func callImageAPI(ctx context.Context, prompt string, count, width, height int) ([]GeneratedImage, error) {
|
||||
reqBody := imageGenRequest{
|
||||
Model: imgCfg.Model,
|
||||
Prompt: prompt,
|
||||
N: count,
|
||||
Size: fmt.Sprintf("%dx%d", width, height),
|
||||
Quality: imgCfg.Quality,
|
||||
}
|
||||
|
||||
body, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal request: %w", err)
|
||||
}
|
||||
|
||||
url := strings.TrimRight(imgCfg.BaseURL, "/") + "/images/generations"
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+imgCfg.APIKey)
|
||||
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("send request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
b, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
|
||||
return nil, fmt.Errorf("image api error %d: %s", resp.StatusCode, string(b))
|
||||
}
|
||||
|
||||
return parseImageResponse(ctx, resp.Body, width, height)
|
||||
}
|
||||
|
||||
// parseImageResponse 解析 OpenAI 兼容图片响应(b64_json 或 url)。
|
||||
func parseImageResponse(ctx context.Context, r io.Reader, width, height int) ([]GeneratedImage, error) {
|
||||
var genResp imageGenResponse
|
||||
if err := json.NewDecoder(r).Decode(&genResp); err != nil {
|
||||
return nil, fmt.Errorf("decode response: %w", err)
|
||||
}
|
||||
|
||||
images := make([]GeneratedImage, 0, len(genResp.Data))
|
||||
for i, d := range genResp.Data {
|
||||
var data []byte
|
||||
switch {
|
||||
case d.B64JSON != "":
|
||||
var err error
|
||||
data, err = base64.StdEncoding.DecodeString(d.B64JSON)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode base64 image %d: %w", i, err)
|
||||
}
|
||||
case d.URL != "":
|
||||
var err error
|
||||
data, err = downloadImage(ctx, d.URL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("download image %d: %w", i, err)
|
||||
}
|
||||
default:
|
||||
return nil, fmt.Errorf("image %d: no data or url in response", i)
|
||||
}
|
||||
images = append(images, GeneratedImage{
|
||||
Data: data,
|
||||
Width: size,
|
||||
Height: size,
|
||||
Width: width,
|
||||
Height: height,
|
||||
Format: "png",
|
||||
}
|
||||
})
|
||||
}
|
||||
return images, nil
|
||||
}
|
||||
|
||||
// QualityChecker 质检函数,可替换用于测试。
|
||||
// 签名:(ctx, images, style) → (pass, reason, error)
|
||||
// downloadImage 从 URL 下载图片数据。
|
||||
func downloadImage(ctx context.Context, url string) ([]byte, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create download request: %w", err)
|
||||
}
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("download: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("download status %d", resp.StatusCode)
|
||||
}
|
||||
return io.ReadAll(resp.Body)
|
||||
}
|
||||
|
||||
// ======================== 图片编辑 ========================
|
||||
|
||||
// EditImages 图片编辑接口,调用 OpenAI 兼容 Images Edits API(multipart/form-data)。
|
||||
func EditImages(ctx context.Context, imageData []byte, prompt string, count int) ([]GeneratedImage, error) {
|
||||
if imgCfg.APIKey == "" {
|
||||
return nil, fmt.Errorf("image API key not configured")
|
||||
}
|
||||
|
||||
var buf bytes.Buffer
|
||||
writer := multipart.NewWriter(&buf)
|
||||
|
||||
part, err := writer.CreateFormFile("image", "image.png")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create form file: %w", err)
|
||||
}
|
||||
if _, err := part.Write(imageData); err != nil {
|
||||
return nil, fmt.Errorf("write image data: %w", err)
|
||||
}
|
||||
|
||||
writer.WriteField("prompt", prompt)
|
||||
writer.WriteField("model", imgCfg.Model)
|
||||
writer.WriteField("n", fmt.Sprintf("%d", count))
|
||||
writer.WriteField("size", fmt.Sprintf("%dx%d", imgCfg.Width, imgCfg.Height))
|
||||
|
||||
if err := writer.Close(); err != nil {
|
||||
return nil, fmt.Errorf("close multipart writer: %w", err)
|
||||
}
|
||||
|
||||
url := strings.TrimRight(imgCfg.BaseURL, "/") + "/images/edits"
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, &buf)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer "+imgCfg.APIKey)
|
||||
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("send request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
b, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
|
||||
return nil, fmt.Errorf("image edit api error %d: %s", resp.StatusCode, string(b))
|
||||
}
|
||||
|
||||
return parseImageResponse(ctx, resp.Body, imgCfg.Width, imgCfg.Height)
|
||||
}
|
||||
|
||||
// ======================== 质检 ========================
|
||||
|
||||
var QualityChecker = defaultCheckQuality
|
||||
|
||||
// CheckQuality 调用当前 QualityChecker。
|
||||
func CheckQuality(ctx context.Context, images []GeneratedImage, style map[string]string) (bool, string, error) {
|
||||
return QualityChecker(ctx, images, style)
|
||||
}
|
||||
|
||||
// defaultCheckQuality 默认 mock 质检,始终返回 pass。
|
||||
func defaultCheckQuality(ctx context.Context, images []GeneratedImage, style map[string]string) (bool, string, error) {
|
||||
return true, "", nil
|
||||
}
|
||||
|
||||
// NewCountedQualityChecker 创建一个在第 passOnRetry 次调用时返回 pass 的质检函数。
|
||||
func NewCountedQualityChecker(passOnRetry int) func(context.Context, []GeneratedImage, map[string]string) (bool, string, error) {
|
||||
var callCount int
|
||||
return func(_ context.Context, _ []GeneratedImage, _ map[string]string) (bool, string, error) {
|
||||
@@ -71,32 +232,40 @@ func NewCountedQualityChecker(passOnRetry int) func(context.Context, []Generated
|
||||
if callCount >= passOnRetry {
|
||||
return true, "", nil
|
||||
}
|
||||
return false, fmt.Sprintf("风格不一致(第 %d 次质检)", callCount), nil
|
||||
return false, fmt.Sprintf("style inconsistent (attempt %d)", callCount), nil
|
||||
}
|
||||
}
|
||||
|
||||
// AlwaysFailQualityChecker 始终返回 fail 的质检函数。
|
||||
func AlwaysFailQualityChecker() func(context.Context, []GeneratedImage, map[string]string) (bool, string, error) {
|
||||
return func(_ context.Context, _ []GeneratedImage, _ map[string]string) (bool, string, error) {
|
||||
return false, "风格不一致", nil
|
||||
return false, "style inconsistent", nil
|
||||
}
|
||||
}
|
||||
|
||||
// generateMockImage 生成一张带随机色块的 PNG 占位图
|
||||
// ======================== Mock 回退 ========================
|
||||
|
||||
func generateMockImages(size, count int) ([]GeneratedImage, error) {
|
||||
images := make([]GeneratedImage, count)
|
||||
for i := 0; i < count; i++ {
|
||||
data, err := generateMockImage(size, i)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("generate mock image %d: %w", i, err)
|
||||
}
|
||||
images[i] = GeneratedImage{Data: data, Width: size, Height: size, Format: "png"}
|
||||
}
|
||||
return images, nil
|
||||
}
|
||||
|
||||
func generateMockImage(size int, seed int) ([]byte, error) {
|
||||
img := image.NewRGBA(image.Rect(0, 0, size, size))
|
||||
|
||||
// 用 seed 生成不同颜色
|
||||
r := uint8((seed*47 + 13) % 256)
|
||||
g := uint8((seed*83 + 37) % 256)
|
||||
b := uint8((seed*61 + 71) % 256)
|
||||
|
||||
for y := 0; y < size; y++ {
|
||||
for x := 0; x < size; x++ {
|
||||
img.Set(x, y, color.RGBA{R: r, G: g, B: b, A: 255})
|
||||
}
|
||||
}
|
||||
|
||||
var buf bytes.Buffer
|
||||
if err := png.Encode(&buf, img); err != nil {
|
||||
return nil, err
|
||||
@@ -104,7 +273,6 @@ func generateMockImage(size int, seed int) ([]byte, error) {
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
// generateRandomBytes 用于生成随机数据(备用)
|
||||
func generateRandomBytes(n int) ([]byte, error) {
|
||||
b := make([]byte, n)
|
||||
_, err := rand.Read(b)
|
||||
|
||||
Regular → Executable
+2
@@ -2,6 +2,8 @@ package service
|
||||
|
||||
// PipelineInput 管线入口输入
|
||||
type PipelineInput struct {
|
||||
ProjectID string // 工程 ID,用于输出目录
|
||||
TaskID string // 任务 ID,用于输出目录
|
||||
Prompt string // 用户原始文本
|
||||
AssetType string // 素材类型:sprite / background / ui / animation
|
||||
ProjectStyle map[string]string // 工程风格键值对
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
//go:build ignore
|
||||
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"gen2d/internal/config"
|
||||
"gen2d/internal/service"
|
||||
)
|
||||
|
||||
func main() {
|
||||
cfg := config.Load()
|
||||
|
||||
// 清空 API key 强制走 mock
|
||||
cfg.ImageGen.APIKey = ""
|
||||
service.InitImageGenConfig(cfg.ImageGen)
|
||||
service.InitLLMConfig(cfg.LLM)
|
||||
|
||||
in := service.PipelineInput{
|
||||
AssetType: "sprite",
|
||||
Prompt: "a cute cat warrior with golden armor",
|
||||
Tags: []string{"pixel art", "fantasy", "16-bit"},
|
||||
Params: service.AssetParams{
|
||||
Resolution: 816,
|
||||
Frames: service.FrameParams{Directions: 4, FramesPerDirection: 1},
|
||||
Format: "individual",
|
||||
},
|
||||
}
|
||||
|
||||
output, err := service.RunPipeline(context.Background(), in)
|
||||
if err != nil {
|
||||
fmt.Fprintf(os.Stderr, "error: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
os.MkdirAll("test_output", 0755)
|
||||
for i, a := range output.Assets {
|
||||
fname := fmt.Sprintf("test_output/gen_%d.%s", i, a.Format)
|
||||
if err := os.WriteFile(fname, a.Data, 0644); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "write %s: %v\n", fname, err)
|
||||
os.Exit(1)
|
||||
}
|
||||
fmt.Printf("saved %s (%d bytes)\n", fname, len(a.Data))
|
||||
}
|
||||
fmt.Printf("metadata: %+v\n", output.Metadata)
|
||||
}
|
||||
Regular → Executable
-4
@@ -41,10 +41,6 @@ async function request<T>(url: string, options: RequestInit = {}): Promise<T> {
|
||||
const json: ApiResponse<T> = await res.json()
|
||||
|
||||
if (json.code !== 0) {
|
||||
if (json.code === 401) {
|
||||
clearToken()
|
||||
window.location.href = '/login'
|
||||
}
|
||||
throw new ApiError(json.code, json.message)
|
||||
}
|
||||
|
||||
|
||||
Regular → Executable
+26
-14
@@ -1,23 +1,35 @@
|
||||
import type { Asset, Task } from './types'
|
||||
import { mockGetAssets, mockGetTask, mockSubmitGenerate } from './mock'
|
||||
|
||||
const USE_MOCK = true
|
||||
import { post, get } from './client'
|
||||
import type {
|
||||
Asset,
|
||||
AssetsResponse,
|
||||
GenerateRequest,
|
||||
GenerateResponse,
|
||||
Task,
|
||||
} from './types'
|
||||
|
||||
export async function submitGenerate(
|
||||
projectId: string,
|
||||
prompt: string,
|
||||
assetType: string
|
||||
): Promise<string> {
|
||||
if (USE_MOCK) return mockSubmitGenerate(projectId, prompt, assetType)
|
||||
throw new Error('Not implemented')
|
||||
req: GenerateRequest,
|
||||
): Promise<GenerateResponse> {
|
||||
return post<GenerateResponse>('/api/v1/generate', req)
|
||||
}
|
||||
|
||||
export async function getTask(taskId: string): Promise<Task> {
|
||||
if (USE_MOCK) return mockGetTask(taskId)
|
||||
throw new Error('Not implemented')
|
||||
return get<Task>(`/api/v1/tasks/${taskId}`)
|
||||
}
|
||||
|
||||
export async function getAssets(taskId: string): Promise<Asset[]> {
|
||||
if (USE_MOCK) return mockGetAssets(taskId)
|
||||
throw new Error('Not implemented')
|
||||
const resp = await get<AssetsResponse>(`/api/v1/tasks/${taskId}/assets`)
|
||||
return resp.assets.map((a, i) => ({
|
||||
id: `asset-${i}`,
|
||||
url: a.url,
|
||||
format: a.format,
|
||||
width: resp.metadata.frameWidth,
|
||||
height: resp.metadata.frameHeight,
|
||||
metadata: {
|
||||
frameWidth: resp.metadata.frameWidth,
|
||||
frameHeight: resp.metadata.frameHeight,
|
||||
frameCount: resp.metadata.frameCount,
|
||||
directions: resp.metadata.directions,
|
||||
},
|
||||
}))
|
||||
}
|
||||
|
||||
Regular → Executable
+28
-7
@@ -67,7 +67,7 @@ export interface Task {
|
||||
projectId: string
|
||||
prompt: string
|
||||
assetType: string
|
||||
status: 'pending' | 'running' | 'completed' | 'failed'
|
||||
status: 'pending' | 'submitted' | 'running' | 'completed' | 'failed'
|
||||
stage?: PipelineStage
|
||||
progress?: number
|
||||
retryCount?: number
|
||||
@@ -91,16 +91,37 @@ export interface Asset {
|
||||
}
|
||||
}
|
||||
|
||||
// 生成请求
|
||||
// 生成请求 — 对应 POST /api/v1/generate
|
||||
export interface GenerateRequest {
|
||||
projectId: string
|
||||
prompt: string
|
||||
prompt?: string
|
||||
assetType: AssetType
|
||||
tags?: string[]
|
||||
userNote?: string
|
||||
projectStyle?: Record<string, string>
|
||||
taskStyle?: Record<string, string>
|
||||
params?: {
|
||||
resolution?: number
|
||||
frames?: { directions?: number; framesPerDirection?: number }
|
||||
format?: 'spritesheet' | 'individual'
|
||||
resolution?: number
|
||||
directions?: number
|
||||
framesPerDir?: number
|
||||
format?: 'spritesheet' | 'individual'
|
||||
}
|
||||
|
||||
// 生成响应 — 对应 POST /api/v1/generate 返回(异步,仅含 taskId)
|
||||
export interface GenerateResponse {
|
||||
taskId: string
|
||||
}
|
||||
|
||||
// 素材列表响应 — 对应 GET /api/v1/tasks/:taskId/assets
|
||||
export interface AssetsResponse {
|
||||
assets: {
|
||||
url: string
|
||||
format: string
|
||||
}[]
|
||||
metadata: {
|
||||
frameWidth: number
|
||||
frameHeight: number
|
||||
frameCount: number
|
||||
directions: number
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Regular → Executable
-3
@@ -5,11 +5,8 @@ export function useGenerate() {
|
||||
return {
|
||||
submit: store.submit,
|
||||
taskId: store.taskId,
|
||||
stage: store.stage,
|
||||
progress: store.progress,
|
||||
status: store.status,
|
||||
retryCount: store.retryCount,
|
||||
rejectReason: store.rejectReason,
|
||||
assets: store.assets,
|
||||
error: store.error,
|
||||
reset: store.reset,
|
||||
|
||||
Regular → Executable
+58
-28
@@ -1,27 +1,36 @@
|
||||
import { useEffect } from 'react'
|
||||
import { useNavigate, useParams } from 'react-router-dom'
|
||||
import { useTaskStore } from '../stores/task'
|
||||
import { useProjectStore } from '../stores/project'
|
||||
import { useGenerationStore } from '../stores/generation'
|
||||
import { useToastStore } from '../stores/toast'
|
||||
import { extractTags } from '../api/prompt'
|
||||
import { mergeStyles } from '../utils/style'
|
||||
import type { GenerateRequest } from '../api/types'
|
||||
import GenerateForm from '../components/GenerateForm'
|
||||
import ProgressBar from '../components/ProgressBar'
|
||||
|
||||
export default function GeneratePage() {
|
||||
const { projectId = 'proj-default' } = useParams()
|
||||
const navigate = useNavigate()
|
||||
const { assetType, reset: resetTask } = useTaskStore()
|
||||
const addToast = useToastStore(s => s.addToast)
|
||||
|
||||
const taskStore = useTaskStore()
|
||||
const { style: projectStyle, loadProject } = useProjectStore()
|
||||
const {
|
||||
status,
|
||||
stage,
|
||||
progress,
|
||||
retryCount,
|
||||
rejectReason,
|
||||
taskId,
|
||||
statusText,
|
||||
submit,
|
||||
reset: resetGeneration,
|
||||
} = useGenerationStore()
|
||||
|
||||
// 加载工程风格
|
||||
useEffect(() => {
|
||||
loadProject(projectId)
|
||||
}, [projectId, loadProject])
|
||||
|
||||
// 组件卸载时重置生成状态
|
||||
useEffect(() => {
|
||||
return () => resetGeneration()
|
||||
@@ -30,63 +39,84 @@ export default function GeneratePage() {
|
||||
// 失败时显示 toast
|
||||
useEffect(() => {
|
||||
if (status === 'failed') {
|
||||
addToast({ type: 'error', message: '素材生成失败,请重试' })
|
||||
const errText = useGenerationStore.getState().error || '未知错误'
|
||||
addToast({ type: 'error', message: `素材生成失败:${errText}` })
|
||||
}
|
||||
}, [status, addToast])
|
||||
|
||||
// 完成后自动跳转
|
||||
useEffect(() => {
|
||||
if (status === 'completed' && taskId) {
|
||||
const timer = setTimeout(() => {
|
||||
navigate(`/projects/${projectId}/tasks/${taskId}`)
|
||||
}, 1500)
|
||||
return () => clearTimeout(timer)
|
||||
}
|
||||
}, [status, taskId, projectId, navigate])
|
||||
|
||||
const handleSubmit = async (finalPrompt: string) => {
|
||||
await submit(projectId, finalPrompt, assetType)
|
||||
const { taskStyle, params, enableAI, optimizedPrompt } = taskStore
|
||||
const mergedStyle = mergeStyles(projectStyle, taskStyle)
|
||||
const tags = extractTags(mergedStyle)
|
||||
|
||||
const req: GenerateRequest = {
|
||||
projectId,
|
||||
prompt: enableAI && optimizedPrompt ? optimizedPrompt : finalPrompt,
|
||||
assetType: taskStore.assetType,
|
||||
tags,
|
||||
projectStyle,
|
||||
taskStyle,
|
||||
resolution: params.resolution,
|
||||
directions: params.frames?.directions,
|
||||
framesPerDir: params.frames?.framesPerDirection,
|
||||
format: params.format,
|
||||
}
|
||||
|
||||
await submit(req)
|
||||
}
|
||||
|
||||
const handleViewResult = () => {
|
||||
if (taskId) navigate(`/projects/${projectId}/tasks/${taskId}`)
|
||||
}
|
||||
|
||||
const handleReset = () => {
|
||||
resetGeneration()
|
||||
resetTask()
|
||||
taskStore.reset()
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="container page-enter" style={{ paddingTop: 24, paddingBottom: 40 }}>
|
||||
<h1 style={{ fontSize: 24, marginBottom: 32 }}>新建生成</h1>
|
||||
|
||||
{/* 生成表单 */}
|
||||
{status === 'idle' || status === 'submitting' ? (
|
||||
<GenerateForm onSubmit={handleSubmit} submitting={status === 'submitting'} />
|
||||
) : (
|
||||
<div style={{ display: 'flex', flexDirection: 'column', gap: 24 }}>
|
||||
{/* 进度条 */}
|
||||
<ProgressBar
|
||||
stage={stage}
|
||||
stage={null}
|
||||
progress={progress}
|
||||
status={status}
|
||||
retryCount={retryCount}
|
||||
rejectReason={rejectReason}
|
||||
retryCount={0}
|
||||
rejectReason={null}
|
||||
/>
|
||||
|
||||
{/* 状态提示 */}
|
||||
{status === 'running' && (
|
||||
<p style={{ textAlign: 'center', color: 'var(--text-secondary)' }}>
|
||||
管线执行中,请稍候...
|
||||
{statusText || '管线执行中,请稍候...'}
|
||||
</p>
|
||||
)}
|
||||
|
||||
{status === 'completed' && (
|
||||
<p style={{ textAlign: 'center', color: 'var(--success)' }}>
|
||||
✓ 生成完成,正在跳转到结果页...
|
||||
</p>
|
||||
<div style={{ textAlign: 'center' }}>
|
||||
<p style={{ color: 'var(--success)', marginBottom: 16 }}>
|
||||
生成完成
|
||||
</p>
|
||||
<div style={{ display: 'flex', gap: 12, justifyContent: 'center' }}>
|
||||
<button className="btn-primary" onClick={handleViewResult}>
|
||||
查看结果
|
||||
</button>
|
||||
<button className="btn-secondary" onClick={handleReset}>
|
||||
继续生成
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{status === 'failed' && (
|
||||
<div style={{ textAlign: 'center' }}>
|
||||
<p style={{ color: 'var(--error)', marginBottom: 16 }}>生成失败</p>
|
||||
<p style={{ color: 'var(--error)', marginBottom: 16 }}>
|
||||
{useGenerationStore.getState().error || '生成失败'}
|
||||
</p>
|
||||
<button className="btn-primary" onClick={handleReset}>
|
||||
重新开始
|
||||
</button>
|
||||
|
||||
Regular → Executable
+85
-42
@@ -1,4 +1,4 @@
|
||||
import { useEffect, useState } from 'react'
|
||||
import { useEffect, useState, useRef } from 'react'
|
||||
import { Link, useParams } from 'react-router-dom'
|
||||
import { getTask, getAssets } from '../api/generate'
|
||||
import type { Asset, Task } from '../api/types'
|
||||
@@ -10,19 +10,47 @@ export default function ResultPage() {
|
||||
const [task, setTask] = useState<Task | null>(null)
|
||||
const [assets, setAssets] = useState<Asset[]>([])
|
||||
const [loading, setLoading] = useState(true)
|
||||
const [polling, setPolling] = useState(false)
|
||||
const pollRef = useRef<ReturnType<typeof setInterval>>()
|
||||
|
||||
useEffect(() => {
|
||||
if (!taskId) return
|
||||
setLoading(true)
|
||||
Promise.all([getTask(taskId), getAssets(taskId)])
|
||||
.then(([t, a]) => {
|
||||
|
||||
const fetchTask = async () => {
|
||||
try {
|
||||
const t = await getTask(taskId)
|
||||
setTask(t)
|
||||
setAssets(a)
|
||||
})
|
||||
.finally(() => setLoading(false))
|
||||
|
||||
if (t.status === 'completed') {
|
||||
if (pollRef.current) clearInterval(pollRef.current)
|
||||
setPolling(false)
|
||||
const a = await getAssets(taskId)
|
||||
setAssets(a)
|
||||
setLoading(false)
|
||||
} else if (t.status === 'failed') {
|
||||
if (pollRef.current) clearInterval(pollRef.current)
|
||||
setPolling(false)
|
||||
setLoading(false)
|
||||
} else if (!pollRef.current) {
|
||||
// 开始轮询
|
||||
setPolling(true)
|
||||
pollRef.current = setInterval(fetchTask, 2000)
|
||||
}
|
||||
} catch {
|
||||
// 出错也停止加载态
|
||||
setLoading(false)
|
||||
}
|
||||
}
|
||||
|
||||
fetchTask()
|
||||
|
||||
return () => {
|
||||
if (pollRef.current) clearInterval(pollRef.current)
|
||||
}
|
||||
}, [taskId])
|
||||
|
||||
if (loading) {
|
||||
if (loading || polling) {
|
||||
return (
|
||||
<div className="container page-enter" style={{ paddingTop: 40 }}>
|
||||
<div className="card" style={{ marginBottom: 24 }}>
|
||||
@@ -30,6 +58,13 @@ export default function ResultPage() {
|
||||
<div style={{ marginTop: 16 }}>
|
||||
<Skeleton variant="text" lines={4} />
|
||||
</div>
|
||||
{task && (
|
||||
<p style={{ textAlign: 'center', color: 'var(--text-secondary)', marginTop: 16 }}>
|
||||
{task.status === 'pending' && '任务排队中...'}
|
||||
{task.status === 'running' && `生成中... ${task.progress ?? 0}%`}
|
||||
{task.status === 'submitted' && '已提交,等待处理...'}
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
<div className="card">
|
||||
<Skeleton variant="card" />
|
||||
@@ -53,27 +88,35 @@ export default function ResultPage() {
|
||||
<div className="container page-enter" style={{ paddingTop: 24, paddingBottom: 40 }}>
|
||||
<h1 style={{ fontSize: 24, marginBottom: 32 }}>生成结果</h1>
|
||||
|
||||
{/* 任务信息 */}
|
||||
<section className="card" style={{ marginBottom: 24 }}>
|
||||
<h2 style={{ fontSize: 16, marginBottom: 16 }}>任务信息</h2>
|
||||
<div
|
||||
style={{
|
||||
display: 'grid',
|
||||
gridTemplateColumns: '120px 1fr',
|
||||
gap: '8px 16px',
|
||||
fontSize: 13,
|
||||
}}
|
||||
>
|
||||
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'center', marginBottom: 16 }}>
|
||||
<h2 style={{ fontSize: 16, marginBottom: 0 }}>任务信息</h2>
|
||||
<div style={{ display: 'flex', gap: 12 }}>
|
||||
{assets.length > 0 && (
|
||||
<button
|
||||
className="btn-primary"
|
||||
onClick={() => downloadAssets(assets)}
|
||||
style={{ padding: '8px 20px', fontSize: 13 }}
|
||||
>
|
||||
下载全部素材
|
||||
</button>
|
||||
)}
|
||||
<Link
|
||||
to={`/projects/${projectId}/generate`}
|
||||
className="btn-secondary"
|
||||
style={{ padding: '8px 20px', fontSize: 13, borderRadius: 'var(--radius)', display: 'inline-block' }}
|
||||
>
|
||||
继续生成
|
||||
</Link>
|
||||
</div>
|
||||
</div>
|
||||
<div style={{ display: 'grid', gridTemplateColumns: '120px 1fr', gap: '8px 16px', fontSize: 13 }}>
|
||||
<span style={{ color: 'var(--text-secondary)' }}>提示词</span>
|
||||
<span>{task.prompt}</span>
|
||||
<span style={{ color: 'var(--text-secondary)' }}>素材类型</span>
|
||||
<span>{task.assetType}</span>
|
||||
<span style={{ color: 'var(--text-secondary)' }}>状态</span>
|
||||
<span
|
||||
style={{
|
||||
color: task.status === 'completed' ? 'var(--success)' : 'var(--error)',
|
||||
}}
|
||||
>
|
||||
<span style={{ color: task.status === 'completed' ? 'var(--success)' : 'var(--error)' }}>
|
||||
{task.status === 'completed' ? '已完成' : '失败'}
|
||||
</span>
|
||||
<span style={{ color: 'var(--text-secondary)' }}>创建时间</span>
|
||||
@@ -93,29 +136,29 @@ export default function ResultPage() {
|
||||
</div>
|
||||
</section>
|
||||
|
||||
{/* 素材预览 */}
|
||||
<section className="card" style={{ marginBottom: 24 }}>
|
||||
<div
|
||||
style={{
|
||||
display: 'flex',
|
||||
justifyContent: 'space-between',
|
||||
alignItems: 'center',
|
||||
marginBottom: 16,
|
||||
}}
|
||||
>
|
||||
<h2 style={{ fontSize: 16 }}>素材预览</h2>
|
||||
{assets.length > 0 && (
|
||||
<button
|
||||
className="btn-primary"
|
||||
onClick={() => alert('下载功能即将上线')}
|
||||
style={{ padding: '8px 20px', fontSize: 13 }}
|
||||
>
|
||||
下载素材
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
<h2 style={{ fontSize: 16, marginBottom: 16 }}>素材预览</h2>
|
||||
<AssetPreview assets={assets} />
|
||||
</section>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
async function downloadAssets(assets: Asset[]) {
|
||||
for (const a of assets) {
|
||||
try {
|
||||
const res = await fetch(a.url)
|
||||
const blob = await res.blob()
|
||||
const blobUrl = URL.createObjectURL(blob)
|
||||
const link = document.createElement('a')
|
||||
link.href = blobUrl
|
||||
link.download = `${a.id}.${a.format}`
|
||||
document.body.appendChild(link)
|
||||
link.click()
|
||||
document.body.removeChild(link)
|
||||
URL.revokeObjectURL(blobUrl)
|
||||
} catch {
|
||||
window.open(a.url, '_blank')
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Regular → Executable
+69
-44
@@ -1,80 +1,105 @@
|
||||
import { create } from 'zustand'
|
||||
import type { Asset, PipelineProgress, PipelineStage } from '../api/types'
|
||||
import { submitGenerate } from '../api/generate'
|
||||
import { createMockWebSocket } from '../api/mock'
|
||||
import type { Asset, GenerateRequest } from '../api/types'
|
||||
import { submitGenerate, getTask, getAssets } from '../api/generate'
|
||||
|
||||
type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed'
|
||||
|
||||
interface GenerationState {
|
||||
taskId: string | null
|
||||
stage: PipelineStage | null
|
||||
projectId: string | null
|
||||
progress: number
|
||||
status: Status
|
||||
retryCount: number
|
||||
rejectReason: string | null
|
||||
statusText: string
|
||||
assets: Asset[]
|
||||
error: string | null
|
||||
submit: (projectId: string, prompt: string, assetType: string) => Promise<void>
|
||||
handleProgress: (msg: PipelineProgress) => void
|
||||
submit: (req: GenerateRequest) => Promise<void>
|
||||
reset: () => void
|
||||
}
|
||||
|
||||
let cleanupWs: (() => void) | null = null
|
||||
let pollTimer: ReturnType<typeof setInterval> | null = null
|
||||
|
||||
function stopPolling() {
|
||||
if (pollTimer) {
|
||||
clearInterval(pollTimer)
|
||||
pollTimer = null
|
||||
}
|
||||
}
|
||||
|
||||
export const useGenerationStore = create<GenerationState>((set, get) => ({
|
||||
taskId: null,
|
||||
stage: null,
|
||||
projectId: null,
|
||||
progress: 0,
|
||||
status: 'idle',
|
||||
retryCount: 0,
|
||||
rejectReason: null,
|
||||
statusText: '',
|
||||
assets: [],
|
||||
error: null,
|
||||
|
||||
submit: async (projectId, prompt, assetType) => {
|
||||
set({ status: 'submitting', error: null })
|
||||
submit: async (req) => {
|
||||
stopPolling()
|
||||
set({ status: 'submitting', error: null, statusText: '提交中...' })
|
||||
try {
|
||||
const taskId = await submitGenerate(projectId, prompt, assetType)
|
||||
set({ taskId, status: 'running', progress: 0 })
|
||||
const { taskId } = await submitGenerate(req)
|
||||
|
||||
// 启动 mock WebSocket
|
||||
cleanupWs = createMockWebSocket(
|
||||
set({
|
||||
taskId,
|
||||
(msg) => get().handleProgress(msg),
|
||||
(assets) => {
|
||||
set({ status: 'completed', assets, progress: 100 })
|
||||
},
|
||||
(error) => {
|
||||
set({ status: 'failed', error })
|
||||
}
|
||||
)
|
||||
} catch (err) {
|
||||
set({ status: 'failed', error: (err as Error).message })
|
||||
}
|
||||
},
|
||||
projectId: req.projectId,
|
||||
status: 'running',
|
||||
progress: 10,
|
||||
statusText: '任务已提交,等待生成...',
|
||||
})
|
||||
|
||||
handleProgress: (msg) => {
|
||||
set({
|
||||
stage: msg.stage,
|
||||
progress: msg.progress,
|
||||
retryCount: msg.retryCount ?? get().retryCount,
|
||||
rejectReason: msg.rejectReason ?? null,
|
||||
})
|
||||
if (msg.result?.assets) {
|
||||
set({ assets: msg.result.assets })
|
||||
// 开始轮询进度
|
||||
pollTimer = setInterval(async () => {
|
||||
try {
|
||||
const task = await getTask(taskId)
|
||||
|
||||
set({
|
||||
progress: task.progress ?? get().progress,
|
||||
statusText:
|
||||
task.status === 'running'
|
||||
? '生成中...'
|
||||
: task.status === 'pending'
|
||||
? '排队中...'
|
||||
: task.status,
|
||||
})
|
||||
|
||||
if (task.status === 'completed') {
|
||||
stopPolling()
|
||||
set({ progress: 90, statusText: '获取结果...' })
|
||||
|
||||
const assets = await getAssets(taskId)
|
||||
set({
|
||||
status: 'completed',
|
||||
progress: 100,
|
||||
statusText: '生成完成',
|
||||
assets,
|
||||
})
|
||||
} else if (task.status === 'failed') {
|
||||
stopPolling()
|
||||
set({
|
||||
status: 'failed',
|
||||
error: task.error || '生成失败',
|
||||
statusText: '生成失败',
|
||||
})
|
||||
}
|
||||
} catch {
|
||||
// 网络错误不中断轮询
|
||||
}
|
||||
}, 2000)
|
||||
} catch (err) {
|
||||
stopPolling()
|
||||
set({ status: 'failed', error: (err as Error).message, statusText: '提交失败' })
|
||||
}
|
||||
},
|
||||
|
||||
reset: () => {
|
||||
cleanupWs?.()
|
||||
cleanupWs = null
|
||||
stopPolling()
|
||||
set({
|
||||
taskId: null,
|
||||
stage: null,
|
||||
projectId: null,
|
||||
progress: 0,
|
||||
status: 'idle',
|
||||
retryCount: 0,
|
||||
rejectReason: null,
|
||||
statusText: '',
|
||||
assets: [],
|
||||
error: null,
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user