!34 refactor(config): 精简 image_gen 配置为 OpenAI 兼容同步 API

Merge pull request !34 from 郭永昊/feat/AssertGeneration
This commit is contained in:
2026-05-25 05:30:42 +00:00
committed by Gitee
19 changed files with 905 additions and 193 deletions
Regular → Executable
+3
View File
@@ -23,3 +23,6 @@ backend/bin/
backend/data/
backend/.env
backend/main
# Generated output
generation/
Regular → Executable
+5 -4
View File
@@ -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
+2
View File
@@ -0,0 +1,2 @@
test_output/
generation/
Regular → Executable
+19 -1
View File
@@ -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")
{
+11 -11
View File
@@ -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")
+3 -5
View File
@@ -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"
+65
View File
@@ -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}))
}
+234
View File
@@ -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,
}))
}
+48
View File
@@ -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()
}
}
+198 -30
View File
@@ -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)
+2
View File
@@ -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 // 工程风格键值对
+49
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
}
}
-3
View File
@@ -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,
+58 -28
View File
@@ -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>
+85 -42
View File
@@ -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')
}
}
}
+69 -44
View File
@@ -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,
})