From 0a9a923d55cfc66394ecd475cd520f3a768b95dd Mon Sep 17 00:00:00 2001 From: Gmaker689 <1711322114@qq.com> Date: Sun, 24 May 2026 22:47:43 +0800 Subject: [PATCH 1/8] =?UTF-8?q?feat(inference):=20=E5=AE=9E=E7=8E=B0?= =?UTF-8?q?=E6=96=87=E7=94=9F=E5=9B=BE=E4=B8=8E=E5=9B=BE=E7=89=87=E7=BC=96?= =?UTF-8?q?=E8=BE=91=20API=20=E8=B0=83=E7=94=A8=EF=BC=8C=E6=96=B0=E5=A2=9E?= =?UTF-8?q?=E7=AE=A1=E7=BA=BF=E4=B8=8E=E7=BC=96=E8=BE=91=20HTTP=20?= =?UTF-8?q?=E7=AB=AF=E7=82=B9=EF=BC=8CJWT=20=E8=AE=A4=E8=AF=81=E4=B8=AD?= =?UTF-8?q?=E9=97=B4=E4=BB=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - inference.go: GenerateImages 改为 API 优先(callImageGenAPI),无 key 回退 mock;新增 EditImages/callImageEditAPI(multipart 上传编辑);提取 parseImageResponse 共享响应解析 - handler/generate.go: POST /api/v1/generate 触发生成管线,返回 base64 图片 - handler/edit.go: POST /api/v1/images/edit 图片编辑端点 - mildware/auth.go: JWT Bearer token 认证中间件 - main.go: 路由拆分公开/认证组,generate 与 images/edit 需鉴权 --- backend/cmd/main.go | 11 +- backend/internal/handler/edit.go | 59 +++++++ backend/internal/handler/generate.go | 84 ++++++++++ backend/internal/mildware/auth.go | 48 ++++++ backend/internal/service/inference.go | 231 +++++++++++++++++++++++--- 5 files changed, 412 insertions(+), 21 deletions(-) create mode 100644 backend/internal/handler/edit.go create mode 100644 backend/internal/handler/generate.go create mode 100644 backend/internal/mildware/auth.go diff --git a/backend/cmd/main.go b/backend/cmd/main.go index 0d4c328..2f2e859 100644 --- a/backend/cmd/main.go +++ b/backend/cmd/main.go @@ -8,6 +8,7 @@ import ( "gen2d/internal/config" "gen2d/internal/db" "gen2d/internal/handler" + "gen2d/internal/mildware" "gen2d/internal/model" "gen2d/internal/service" @@ -34,13 +35,21 @@ func main() { r := gin.New() r.Use(gin.Recovery()) // panic 恢复中间件,防止服务因未捕获异常宕机 - // API v1 路由组 + // API v1 路由组 — 公开接口 v1 := r.Group("/api/v1") { v1.GET("/health", handler.Health) // 健康检查 v1.POST("/prompt/optimize", handler.PromptOptimize) // 提示词优化 } + // API v1 路由组 — 需认证 + v1Auth := r.Group("/api/v1") + v1Auth.Use(mildware.AuthMiddleware(cfg.JWT.Secret)) + { + v1Auth.POST("/generate", handler.Generate) // 素材生成管线 + v1Auth.POST("/images/edit", handler.EditImage) // 图片编辑 + } + // Auth 路由组 auth := r.Group("/auth") { diff --git a/backend/internal/handler/edit.go b/backend/internal/handler/edit.go new file mode 100644 index 0000000..b17d052 --- /dev/null +++ b/backend/internal/handler/edit.go @@ -0,0 +1,59 @@ +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 +} + +// EditImageResponse 图片编辑响应体。 +type EditImageResponse struct { + Assets []AssetResponse `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([]AssetResponse, len(images)) + for i, img := range images { + assets[i] = AssetResponse{ + Data: base64.StdEncoding.EncodeToString(img.Data), + Format: img.Format, + } + } + + c.JSON(http.StatusOK, model.OK(EditImageResponse{Assets: assets})) +} diff --git a/backend/internal/handler/generate.go b/backend/internal/handler/generate.go new file mode 100644 index 0000000..52b38af --- /dev/null +++ b/backend/internal/handler/generate.go @@ -0,0 +1,84 @@ +package handler + +import ( + "encoding/base64" + "net/http" + + "gen2d/internal/model" + "gen2d/internal/service" + + "github.com/gin-gonic/gin" +) + +// GenerateRequest 素材生成请求。 +type GenerateRequest struct { + 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 { + Assets []AssetResponse `json:"assets"` + Metadata service.AssetMetadata `json:"metadata"` +} + +// AssetResponse 单个素材响应(二进制 Data 转 base64)。 +type AssetResponse struct { + Data string `json:"data"` + Format string `json:"format"` + URL string `json:"url"` +} + +// Generate 素材生成接口,调用完整生成管线(PromptOptimizer → AssetGenerator → QualitySupervisor → FormatAdapter)。 +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 + } + + in := service.PipelineInput{ + Prompt: req.Prompt, + AssetType: req.AssetType, + Tags: req.Tags, + UserNote: req.UserNote, + ProjectStyle: req.ProjectStyle, + TaskStyle: req.TaskStyle, + Params: service.AssetParams{ + Resolution: req.Resolution, + Frames: service.FrameParams{ + Directions: req.Directions, + FramesPerDirection: req.FramesPerDir, + }, + Format: req.Format, + }, + } + + output, err := service.RunPipeline(c.Request.Context(), in) + if err != nil { + c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "素材生成失败: "+err.Error())) + return + } + + assets := make([]AssetResponse, len(output.Assets)) + for i, a := range output.Assets { + assets[i] = AssetResponse{ + Data: base64.StdEncoding.EncodeToString(a.Data), + Format: a.Format, + URL: a.URL, + } + } + + c.JSON(http.StatusOK, model.OK(GenerateResponse{ + Assets: assets, + Metadata: output.Metadata, + })) +} diff --git a/backend/internal/mildware/auth.go b/backend/internal/mildware/auth.go new file mode 100644 index 0000000..9a0226a --- /dev/null +++ b/backend/internal/mildware/auth.go @@ -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 ")) + 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() + } +} diff --git a/backend/internal/service/inference.go b/backend/internal/service/inference.go index 30b6ee0..ac9c896 100644 --- a/backend/internal/service/inference.go +++ b/backend/internal/service/inference.go @@ -4,10 +4,18 @@ import ( "bytes" "context" "crypto/rand" + "encoding/base64" + "encoding/json" "fmt" "image" "image/color" "image/png" + "io" + "log" + "mime/multipart" + "net/http" + "strconv" + "strings" "gen2d/internal/config" ) @@ -20,37 +28,201 @@ 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 - } +// ======================== 文生图 API 调用层 ======================== +// imageGenRequest OpenAI 兼容的文生图请求体。 +type imageGenRequest struct { + Model string `json:"model"` + Prompt string `json:"prompt"` + N int `json:"n,omitempty"` + Size string `json:"size,omitempty"` + ResponseFormat string `json:"response_format,omitempty"` + Steps int `json:"steps,omitempty"` + CFGScale float64 `json:"cfg_scale,omitempty"` +} + +// imageGenResponse OpenAI 兼容的文生图响应体。 +type imageGenResponse struct { + Data []struct { + URL string `json:"url"` + B64JSON string `json:"b64_json"` + } `json:"data"` +} + +// GenerateImages 调用 AI 推理 API 生成图片,未配置 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 != "" { + return callImageGenAPI(ctx, prompt, count) + } + + size := params.Resolution + if size <= 0 { + size = 64 + } + log.Println("[inference] ImageGen API key not configured, using mock") + return generateMockImages(size, count) +} + +// callImageGenAPI 调用 OpenAI 兼容的 Images API,返回生成的图片。 +func callImageGenAPI(ctx context.Context, prompt string, count int) ([]GeneratedImage, error) { + reqBody := imageGenRequest{ + Model: imgCfg.Model, + Prompt: prompt, + N: count, + Size: fmt.Sprintf("%dx%d", imgCfg.Width, imgCfg.Height), + ResponseFormat: "b64_json", + } + if imgCfg.Steps > 0 { + reqBody.Steps = imgCfg.Steps + } + if imgCfg.CFGScale > 0 { + reqBody.CFGScale = imgCfg.CFGScale + } + + body, err := json.Marshal(reqBody) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + + url := strings.TrimRight(imgCfg.BaseURL, "/") + 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 gen api error %d: %s", resp.StatusCode, string(b)) + } + + return parseImageResponse(ctx, resp.Body) +} + +// ======================== 图片编辑 API ======================== + +// EditImages 调用图片编辑 API,基于已有图片和文本提示词生成修改后的图片。 +func EditImages(ctx context.Context, imageData []byte, prompt string, count int) ([]GeneratedImage, error) { + if imgCfg.APIKey == "" { + return nil, fmt.Errorf("image edit API key not configured") + } + return callImageEditAPI(ctx, imageData, prompt, count) +} + +// callImageEditAPI 调用 OpenAI 兼容的 Images Edits API(multipart/form-data)。 +func callImageEditAPI(ctx context.Context, imageData []byte, prompt string, count int) ([]GeneratedImage, error) { + 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", strconv.Itoa(count)) + writer.WriteField("size", fmt.Sprintf("%dx%d", imgCfg.Width, imgCfg.Height)) + writer.WriteField("response_format", "b64_json") + + if err := writer.Close(); err != nil { + return nil, fmt.Errorf("close multipart writer: %w", err) + } + + url := strings.Replace(strings.TrimRight(imgCfg.BaseURL, "/"), "generations", "edits", 1) + + 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) +} + +// parseImageResponse 解析 OpenAI 兼容的图片生成/编辑响应体。 +func parseImageResponse(ctx context.Context, r io.Reader) ([]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[i] = GeneratedImage{ + images = append(images, GeneratedImage{ Data: data, - Width: size, - Height: size, + Width: imgCfg.Width, + Height: imgCfg.Height, Format: "png", - } + }) } return images, nil } +// 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) +} + +// ======================== 质检 ======================== + // QualityChecker 质检函数,可替换用于测试。 -// 签名:(ctx, images, style) → (pass, reason, error) var QualityChecker = defaultCheckQuality // CheckQuality 调用当前 QualityChecker。 @@ -82,11 +254,30 @@ func AlwaysFailQualityChecker() func(context.Context, []GeneratedImage, map[stri } } -// generateMockImage 生成一张带随机色块的 PNG 占位图 +// ======================== Mock 回退 ======================== + +// generateMockImages 批量生成 mock PNG 占位图。 +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 +} + +// generateMockImage 生成一张带随机色块的 PNG 占位图。 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) @@ -104,7 +295,7 @@ func generateMockImage(size int, seed int) ([]byte, error) { return buf.Bytes(), nil } -// generateRandomBytes 用于生成随机数据(备用) +// generateRandomBytes 用于生成随机数据(备用)。 func generateRandomBytes(n int) ([]byte, error) { b := make([]byte, n) _, err := rand.Read(b) From 6bda5a1f351ed459eb0e068b421b66e83c8e6bf5 Mon Sep 17 00:00:00 2001 From: Gmaker689 <1711322114@qq.com> Date: Mon, 25 May 2026 00:21:16 +0800 Subject: [PATCH 2/8] =?UTF-8?q?fix(inference):=20=E4=BF=AE=E6=AD=A3?= =?UTF-8?q?=E6=96=87=E7=94=9F=E5=9B=BE=20API=20=E8=B0=83=E7=94=A8=E5=8F=82?= =?UTF-8?q?=E6=95=B0=EF=BC=8C=E6=94=AF=E6=8C=81=E5=8A=A8=E6=80=81=E5=88=86?= =?UTF-8?q?=E8=BE=A8=E7=8E=87?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 去掉 response_format 参数(API 默认返回 URL,自动下载) - Resolution 参数覆盖 API 调用尺寸(不再仅 mock 路径生效) - parseImageResponse 改为动态 width/height 入参 - 新增 tools/gentest.go 生成测试脚本 - .gitignore 忽略 test_output/ --- backend/.gitignore | 1 + backend/internal/service/inference.go | 24 +++++++------ backend/tools/gentest.go | 49 +++++++++++++++++++++++++++ 3 files changed, 63 insertions(+), 11 deletions(-) create mode 100644 backend/.gitignore create mode 100644 backend/tools/gentest.go diff --git a/backend/.gitignore b/backend/.gitignore new file mode 100644 index 0000000..46c8ed0 --- /dev/null +++ b/backend/.gitignore @@ -0,0 +1 @@ +test_output/ diff --git a/backend/internal/service/inference.go b/backend/internal/service/inference.go index ac9c896..bd55b0c 100644 --- a/backend/internal/service/inference.go +++ b/backend/internal/service/inference.go @@ -36,7 +36,6 @@ type imageGenRequest struct { Prompt string `json:"prompt"` N int `json:"n,omitempty"` Size string `json:"size,omitempty"` - ResponseFormat string `json:"response_format,omitempty"` Steps int `json:"steps,omitempty"` CFGScale float64 `json:"cfg_scale,omitempty"` } @@ -57,7 +56,12 @@ func GenerateImages(ctx context.Context, prompt string, params AssetParams) ([]G } if imgCfg.APIKey != "" { - return callImageGenAPI(ctx, prompt, count) + width, height := imgCfg.Width, imgCfg.Height + if params.Resolution > 0 { + width = params.Resolution + height = params.Resolution + } + return callImageGenAPI(ctx, prompt, count, width, height) } size := params.Resolution @@ -69,13 +73,12 @@ func GenerateImages(ctx context.Context, prompt string, params AssetParams) ([]G } // callImageGenAPI 调用 OpenAI 兼容的 Images API,返回生成的图片。 -func callImageGenAPI(ctx context.Context, prompt string, count int) ([]GeneratedImage, error) { +func callImageGenAPI(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", imgCfg.Width, imgCfg.Height), - ResponseFormat: "b64_json", + Size: fmt.Sprintf("%dx%d", width, height), } if imgCfg.Steps > 0 { reqBody.Steps = imgCfg.Steps @@ -108,7 +111,7 @@ func callImageGenAPI(ctx context.Context, prompt string, count int) ([]Generated return nil, fmt.Errorf("image gen api error %d: %s", resp.StatusCode, string(b)) } - return parseImageResponse(ctx, resp.Body) + return parseImageResponse(ctx, resp.Body, width, height) } // ======================== 图片编辑 API ======================== @@ -138,7 +141,6 @@ func callImageEditAPI(ctx context.Context, imageData []byte, prompt string, coun writer.WriteField("model", imgCfg.Model) writer.WriteField("n", strconv.Itoa(count)) writer.WriteField("size", fmt.Sprintf("%dx%d", imgCfg.Width, imgCfg.Height)) - writer.WriteField("response_format", "b64_json") if err := writer.Close(); err != nil { return nil, fmt.Errorf("close multipart writer: %w", err) @@ -164,11 +166,11 @@ func callImageEditAPI(ctx context.Context, imageData []byte, prompt string, coun return nil, fmt.Errorf("image edit api error %d: %s", resp.StatusCode, string(b)) } - return parseImageResponse(ctx, resp.Body) + return parseImageResponse(ctx, resp.Body, imgCfg.Width, imgCfg.Height) } // parseImageResponse 解析 OpenAI 兼容的图片生成/编辑响应体。 -func parseImageResponse(ctx context.Context, r io.Reader) ([]GeneratedImage, error) { +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) @@ -195,8 +197,8 @@ func parseImageResponse(ctx context.Context, r io.Reader) ([]GeneratedImage, err } images = append(images, GeneratedImage{ Data: data, - Width: imgCfg.Width, - Height: imgCfg.Height, + Width: width, + Height: height, Format: "png", }) } diff --git a/backend/tools/gentest.go b/backend/tools/gentest.go new file mode 100644 index 0000000..d6af462 --- /dev/null +++ b/backend/tools/gentest.go @@ -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) +} From 74c17254739c504aad49e4f0d1a995efd0426bc3 Mon Sep 17 00:00:00 2001 From: Gmaker689 <1711322114@qq.com> Date: Mon, 25 May 2026 12:27:07 +0800 Subject: [PATCH 3/8] =?UTF-8?q?feat(generate):=20=E5=9B=BE=E7=89=87?= =?UTF-8?q?=E6=8C=81=E4=B9=85=E5=8C=96=E5=88=B0=E6=9C=AC=E5=9C=B0=20genera?= =?UTF-8?q?tion/=20=E7=9B=AE=E5=BD=95=EF=BC=8C=E5=89=8D=E7=AB=AF=E5=AF=B9?= =?UTF-8?q?=E6=8E=A5=E7=9C=9F=E5=AE=9E=20API?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 后端: - POST /api/v1/generate 接收 projectId,生成 taskId,将图片保存至 generation/{projectId}/{taskId}/ - 新增 GET /api/v1/tasks/:taskId 和 GET /api/v1/tasks/:taskId/assets 查询端点 - 添加 /generation 静态文件服务,前端可通过 URL 直接访问生成的图片 - PipelineInput 增加 ProjectID/TaskID 字段 前端: - generate.ts 对接真实后端 API,移除 mock 模式 - 401 时不再强制登出跳转,改为抛错由调用方处理 - GeneratePage 加载工程风格,传递 projectId 和完整参数 - generation store 简化为同步 API 模式,移除 WebSocket mock - ResultPage 使用 getTask/getAssets 按 taskId 查询结果 --- .gitignore | 3 + backend/.gitignore | 1 + backend/cmd/main.go | 13 ++- backend/internal/handler/edit.go | 12 ++- backend/internal/handler/generate.go | 118 ++++++++++++++++++++++++--- backend/internal/service/types.go | 2 + frontend/src/api/client.ts | 4 - frontend/src/api/generate.ts | 47 +++++++---- frontend/src/api/types.ts | 43 ++++++++-- frontend/src/hooks/useGenerate.ts | 3 - frontend/src/pages/GeneratePage.tsx | 83 +++++++++++++------ frontend/src/pages/ResultPage.tsx | 0 frontend/src/stores/generation.ts | 82 +++++++++---------- 13 files changed, 297 insertions(+), 114 deletions(-) mode change 100644 => 100755 .gitignore mode change 100644 => 100755 backend/.gitignore mode change 100644 => 100755 backend/cmd/main.go mode change 100644 => 100755 backend/internal/handler/edit.go mode change 100644 => 100755 backend/internal/handler/generate.go mode change 100644 => 100755 backend/internal/service/types.go mode change 100644 => 100755 frontend/src/api/client.ts mode change 100644 => 100755 frontend/src/api/generate.ts mode change 100644 => 100755 frontend/src/api/types.ts mode change 100644 => 100755 frontend/src/hooks/useGenerate.ts mode change 100644 => 100755 frontend/src/pages/GeneratePage.tsx mode change 100644 => 100755 frontend/src/pages/ResultPage.tsx mode change 100644 => 100755 frontend/src/stores/generation.ts diff --git a/.gitignore b/.gitignore old mode 100644 new mode 100755 index 1fa64c0..bc4fa7d --- a/.gitignore +++ b/.gitignore @@ -23,3 +23,6 @@ backend/bin/ backend/data/ backend/.env backend/main + +# Generated output +generation/ diff --git a/backend/.gitignore b/backend/.gitignore old mode 100644 new mode 100755 index 46c8ed0..71eeeb9 --- a/backend/.gitignore +++ b/backend/.gitignore @@ -1 +1,2 @@ test_output/ +generation/ diff --git a/backend/cmd/main.go b/backend/cmd/main.go old mode 100644 new mode 100755 index 2f2e859..2a447dd --- a/backend/cmd/main.go +++ b/backend/cmd/main.go @@ -4,6 +4,7 @@ package main import ( "fmt" "log" + "os" "gen2d/internal/config" "gen2d/internal/db" @@ -20,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) @@ -35,6 +39,9 @@ func main() { r := gin.New() r.Use(gin.Recovery()) // panic 恢复中间件,防止服务因未捕获异常宕机 + // 静态文件服务 — 生成的图片 + r.Static("/generation", "../generation") + // API v1 路由组 — 公开接口 v1 := r.Group("/api/v1") { @@ -46,8 +53,10 @@ func main() { v1Auth := r.Group("/api/v1") v1Auth.Use(mildware.AuthMiddleware(cfg.JWT.Secret)) { - v1Auth.POST("/generate", handler.Generate) // 素材生成管线 - v1Auth.POST("/images/edit", handler.EditImage) // 图片编辑 + 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 路由组 diff --git a/backend/internal/handler/edit.go b/backend/internal/handler/edit.go old mode 100644 new mode 100755 index b17d052..f77b05a --- a/backend/internal/handler/edit.go +++ b/backend/internal/handler/edit.go @@ -17,9 +17,15 @@ type EditImageRequest struct { Count int `json:"count"` // 生成数量,默认 1 } +// editAssetResponse 编辑结果素材(返回 base64)。 +type editAssetResponse struct { + Data string `json:"data"` + Format string `json:"format"` +} + // EditImageResponse 图片编辑响应体。 type EditImageResponse struct { - Assets []AssetResponse `json:"assets"` + Assets []editAssetResponse `json:"assets"` } // EditImage 图片编辑接口,基于已有图片和文本指令生成修改后的图片。 @@ -47,9 +53,9 @@ func EditImage(c *gin.Context) { return } - assets := make([]AssetResponse, len(images)) + assets := make([]editAssetResponse, len(images)) for i, img := range images { - assets[i] = AssetResponse{ + assets[i] = editAssetResponse{ Data: base64.StdEncoding.EncodeToString(img.Data), Format: img.Format, } diff --git a/backend/internal/handler/generate.go b/backend/internal/handler/generate.go old mode 100644 new mode 100755 index 52b38af..ebf5ca7 --- a/backend/internal/handler/generate.go +++ b/backend/internal/handler/generate.go @@ -1,8 +1,12 @@ package handler import ( - "encoding/base64" + "fmt" "net/http" + "os" + "path/filepath" + "sync" + "time" "gen2d/internal/model" "gen2d/internal/service" @@ -12,6 +16,7 @@ import ( // GenerateRequest 素材生成请求。 type GenerateRequest struct { + ProjectID string `json:"projectId"` Prompt string `json:"prompt"` AssetType string `json:"assetType" binding:"required"` Tags []string `json:"tags"` @@ -26,18 +31,48 @@ type GenerateRequest struct { // GenerateResponse 素材生成响应体。 type GenerateResponse struct { - Assets []AssetResponse `json:"assets"` - Metadata service.AssetMetadata `json:"metadata"` + TaskID string `json:"taskId"` + Assets []AssetResponse `json:"assets"` + Metadata service.AssetMetadata `json:"metadata"` } -// AssetResponse 单个素材响应(二进制 Data 转 base64)。 +// AssetResponse 单个素材响应。 type AssetResponse struct { - Data string `json:"data"` - Format string `json:"format"` URL string `json:"url"` + Format string `json:"format"` } -// Generate 素材生成接口,调用完整生成管线(PromptOptimizer → AssetGenerator → QualitySupervisor → FormatAdapter)。 +// 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 素材生成接口。 func Generate(c *gin.Context) { var req GenerateRequest if err := c.ShouldBindJSON(&req); err != nil { @@ -45,7 +80,15 @@ func Generate(c *gin.Context) { return } + projectID := req.ProjectID + if projectID == "" { + projectID = "default" + } + taskID := fmt.Sprintf("task-%d", time.Now().UnixMilli()) + in := service.PipelineInput{ + ProjectID: projectID, + TaskID: taskID, Prompt: req.Prompt, AssetType: req.AssetType, Tags: req.Tags, @@ -55,7 +98,7 @@ func Generate(c *gin.Context) { Params: service.AssetParams{ Resolution: req.Resolution, Frames: service.FrameParams{ - Directions: req.Directions, + Directions: req.Directions, FramesPerDirection: req.FramesPerDir, }, Format: req.Format, @@ -68,17 +111,72 @@ func Generate(c *gin.Context) { return } + // 保存图片到 ../generation/{projectId}/{taskId}/ + genDir := filepath.Join("..", "generation", projectID, taskID) + if err := os.MkdirAll(genDir, 0755); err != nil { + c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "创建输出目录失败: "+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 { + c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "保存图片失败: "+err.Error())) + return + } assets[i] = AssetResponse{ - Data: base64.StdEncoding.EncodeToString(a.Data), + URL: fmt.Sprintf("/generation/%s/%s/%s", projectID, taskID, filename), Format: a.Format, - URL: a.URL, } } + // 存储任务记录到内存 + 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, + }) + c.JSON(http.StatusOK, model.OK(GenerateResponse{ + TaskID: taskID, Assets: assets, Metadata: output.Metadata, })) } + +// 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) + c.JSON(http.StatusOK, model.OK(AssetsResponse{ + Assets: r.assets, + Metadata: r.metadata, + })) +} diff --git a/backend/internal/service/types.go b/backend/internal/service/types.go old mode 100644 new mode 100755 index ed684b0..9db9f27 --- a/backend/internal/service/types.go +++ b/backend/internal/service/types.go @@ -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 // 工程风格键值对 diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts old mode 100644 new mode 100755 index 9ff77c8..0ec0a26 --- a/frontend/src/api/client.ts +++ b/frontend/src/api/client.ts @@ -41,10 +41,6 @@ async function request(url: string, options: RequestInit = {}): Promise { const json: ApiResponse = await res.json() if (json.code !== 0) { - if (json.code === 401) { - clearToken() - window.location.href = '/login' - } throw new ApiError(json.code, json.message) } diff --git a/frontend/src/api/generate.ts b/frontend/src/api/generate.ts old mode 100644 new mode 100755 index 40670a8..971d9b2 --- a/frontend/src/api/generate.ts +++ b/frontend/src/api/generate.ts @@ -1,23 +1,42 @@ -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 { - if (USE_MOCK) return mockSubmitGenerate(projectId, prompt, assetType) - throw new Error('Not implemented') + req: GenerateRequest, +): Promise { + return post('/api/v1/generate', req) } export async function getTask(taskId: string): Promise { - if (USE_MOCK) return mockGetTask(taskId) - throw new Error('Not implemented') + return get(`/api/v1/tasks/${taskId}`) } export async function getAssets(taskId: string): Promise { - if (USE_MOCK) return mockGetAssets(taskId) - throw new Error('Not implemented') + const resp = await get(`/api/v1/tasks/${taskId}/assets`) + return toAssetList(resp) +} + +/** 将响应转为 Asset[] 供前端组件使用 */ +function toAssetList( + resp: AssetsResponse | GenerateResponse, +): Asset[] { + 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, + }, + })) } diff --git a/frontend/src/api/types.ts b/frontend/src/api/types.ts old mode 100644 new mode 100755 index b9eaa96..5c5f8a9 --- a/frontend/src/api/types.ts +++ b/frontend/src/api/types.ts @@ -91,16 +91,47 @@ export interface Asset { } } -// 生成请求 +// 生成请求 — 对应 POST /api/v1/generate export interface GenerateRequest { projectId: string - prompt: string + prompt?: string assetType: AssetType + tags?: string[] + userNote?: string + projectStyle?: Record taskStyle?: Record - params?: { - resolution?: number - frames?: { directions?: number; framesPerDirection?: number } - format?: 'spritesheet' | 'individual' + resolution?: number + directions?: number + framesPerDir?: number + format?: 'spritesheet' | 'individual' +} + +// 生成响应 — 对应 POST /api/v1/generate 返回 +export interface GenerateResponse { + taskId: string + assets: { + url: string + format: string + }[] + metadata: { + frameWidth: number + frameHeight: number + frameCount: number + directions: number + } +} + +// 素材列表响应 — 对应 GET /api/v1/tasks/:taskId/assets +export interface AssetsResponse { + assets: { + url: string + format: string + }[] + metadata: { + frameWidth: number + frameHeight: number + frameCount: number + directions: number } } diff --git a/frontend/src/hooks/useGenerate.ts b/frontend/src/hooks/useGenerate.ts old mode 100644 new mode 100755 index 49220c7..6bf510e --- a/frontend/src/hooks/useGenerate.ts +++ b/frontend/src/hooks/useGenerate.ts @@ -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, diff --git a/frontend/src/pages/GeneratePage.tsx b/frontend/src/pages/GeneratePage.tsx old mode 100644 new mode 100755 index bb17ec7..d21356c --- a/frontend/src/pages/GeneratePage.tsx +++ b/frontend/src/pages/GeneratePage.tsx @@ -1,27 +1,35 @@ 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, submit, reset: resetGeneration, } = useGenerationStore() + // 加载工程风格 + useEffect(() => { + loadProject(projectId) + }, [projectId, loadProject]) + // 组件卸载时重置生成状态 useEffect(() => { return () => resetGeneration() @@ -30,48 +38,57 @@ 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 (

新建生成

- {/* 生成表单 */} {status === 'idle' || status === 'submitting' ? ( ) : (
- {/* 进度条 */} - {/* 状态提示 */} {status === 'running' && (

管线执行中,请稍候... @@ -79,14 +96,26 @@ export default function GeneratePage() { )} {status === 'completed' && ( -

- ✓ 生成完成,正在跳转到结果页... -

+
+

+ 生成完成 +

+
+ + +
+
)} {status === 'failed' && (
-

生成失败

+

+ {useGenerationStore.getState().error || '生成失败'} +

diff --git a/frontend/src/pages/ResultPage.tsx b/frontend/src/pages/ResultPage.tsx old mode 100644 new mode 100755 diff --git a/frontend/src/stores/generation.ts b/frontend/src/stores/generation.ts old mode 100644 new mode 100755 index b904755..bb3ec9c --- a/frontend/src/stores/generation.ts +++ b/frontend/src/stores/generation.ts @@ -1,82 +1,74 @@ import { create } from 'zustand' -import type { Asset, PipelineProgress, PipelineStage } from '../api/types' +import type { Asset, GenerateRequest, GenerateResponse } from '../api/types' import { submitGenerate } from '../api/generate' -import { createMockWebSocket } from '../api/mock' 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 assets: Asset[] error: string | null - submit: (projectId: string, prompt: string, assetType: string) => Promise - handleProgress: (msg: PipelineProgress) => void + submit: (req: GenerateRequest) => Promise reset: () => void } -let cleanupWs: (() => void) | null = null - -export const useGenerationStore = create((set, get) => ({ +export const useGenerationStore = create((set) => ({ taskId: null, - stage: null, + projectId: null, progress: 0, status: 'idle', - retryCount: 0, - rejectReason: null, assets: [], error: null, - submit: async (projectId, prompt, assetType) => { + submit: async (req) => { set({ status: 'submitting', error: null }) try { - const taskId = await submitGenerate(projectId, prompt, assetType) - set({ taskId, status: 'running', progress: 0 }) + set({ status: 'running', progress: 30 }) + const result = await submitGenerate(req) - // 启动 mock WebSocket - cleanupWs = createMockWebSocket( - taskId, - (msg) => get().handleProgress(msg), - (assets) => { - set({ status: 'completed', assets, progress: 100 }) - }, - (error) => { - set({ status: 'failed', error }) - } - ) + set({ progress: 80 }) + const assets = mapAssets(result) + + set({ + taskId: result.taskId, + projectId: req.projectId, + status: 'completed', + progress: 100, + assets, + }) } catch (err) { - set({ status: 'failed', error: (err as Error).message }) - } - }, - - 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 }) + const errMsg = (err as Error).message + set({ status: 'failed', error: errMsg }) } }, reset: () => { - cleanupWs?.() - cleanupWs = null set({ taskId: null, - stage: null, + projectId: null, progress: 0, status: 'idle', - retryCount: 0, - rejectReason: null, assets: [], error: null, }) }, })) + +function mapAssets(resp: GenerateResponse): Asset[] { + 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, + }, + })) +} From d9d3b51262ed16ec1f40c7be1a64bc3e9e732e1e Mon Sep 17 00:00:00 2001 From: Gmaker689 <1711322114@qq.com> Date: Mon, 25 May 2026 12:36:37 +0800 Subject: [PATCH 4/8] =?UTF-8?q?feat(inference):=20=E9=9B=86=E6=88=90=20GPT?= =?UTF-8?q?=20Image=202=20=E5=BC=82=E6=AD=A5=E6=96=87=E7=94=9F=E5=9B=BE=20?= =?UTF-8?q?API?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 GptImage2Config 配置(yuntts 等 GPT Image 2 兼容服务) - 新增 gptimage.go 异步客户端:提交任务 → 轮询状态 → 下载图片 - GenerateImages 优先级调整为:GPT Image 2 > OpenAI 兼容 ImageGen > Mock - 支持环境变量 GEN2D_GPT_IMAGE2_* 系列配置 --- backend/.env.example | 11 +- backend/cmd/main.go | 1 + backend/internal/config/config.go | 37 +++- backend/internal/service/gptimage.go | 268 ++++++++++++++++++++++++++ backend/internal/service/inference.go | 17 +- 5 files changed, 325 insertions(+), 9 deletions(-) mode change 100644 => 100755 backend/.env.example mode change 100644 => 100755 backend/internal/config/config.go create mode 100755 backend/internal/service/gptimage.go mode change 100644 => 100755 backend/internal/service/inference.go diff --git a/backend/.env.example b/backend/.env.example old mode 100644 new mode 100755 index 97bf0b8..2501aa8 --- a/backend/.env.example +++ b/backend/.env.example @@ -19,7 +19,7 @@ GEN2D_LLM_MODEL=gpt-4o GEN2D_LLM_TEMPERATURE=0.7 GEN2D_LLM_MAX_TOKENS=2048 -# 文生图模型 +# 文生图模型(OpenAI 兼容同步 API) GEN2D_IMAGE_BASE_URL=https://api.stability.ai/v1 GEN2D_IMAGE_API_KEY=sk-your-api-key GEN2D_IMAGE_MODEL=stable-diffusion-xl @@ -28,3 +28,12 @@ GEN2D_IMAGE_HEIGHT=1024 GEN2D_IMAGE_NUM_IMAGES=1 GEN2D_IMAGE_STEPS=30 GEN2D_IMAGE_CFG_SCALE=7.0 + +# GPT Image 2 文生图(异步 API,如 yuntts 等兼容服务) +# 优先级高于 IMAGE API,留空则使用上方的 IMAGE API +GEN2D_GPT_IMAGE2_BASE_URL=https://www.yuntts.com/api/v1 +GEN2D_GPT_IMAGE2_API_KEY= +GEN2D_GPT_IMAGE2_ASPECT_RATIO=1:1 +GEN2D_GPT_IMAGE2_X_CHANNEL=default +GEN2D_GPT_IMAGE2_POLL_MAX_WAIT=120 +GEN2D_GPT_IMAGE2_POLL_INTERVAL=3 diff --git a/backend/cmd/main.go b/backend/cmd/main.go index 2a447dd..af7578a 100755 --- a/backend/cmd/main.go +++ b/backend/cmd/main.go @@ -35,6 +35,7 @@ func main() { // 注入 LLM 和文生图配置到 service 层 service.InitLLMConfig(cfg.LLM) service.InitImageGenConfig(cfg.ImageGen) + service.InitGptImage2Config(cfg.GptImage2) r := gin.New() r.Use(gin.Recovery()) // panic 恢复中间件,防止服务因未捕获异常宕机 diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go old mode 100644 new mode 100755 index 62359f3..12dea7e --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -9,11 +9,12 @@ import ( // Config 应用全局配置。 type Config struct { - Server ServerConfig `mapstructure:"server"` - Database DatabaseConfig `mapstructure:"database"` - JWT JWTConfig `mapstructure:"jwt"` - LLM LLMConfig `mapstructure:"llm"` - ImageGen ImageGenConfig `mapstructure:"image_gen"` + Server ServerConfig `mapstructure:"server"` + Database DatabaseConfig `mapstructure:"database"` + JWT JWTConfig `mapstructure:"jwt"` + LLM LLMConfig `mapstructure:"llm"` + ImageGen ImageGenConfig `mapstructure:"image_gen"` + GptImage2 GptImage2Config `mapstructure:"gpt_image2"` } // ServerConfig HTTP 服务配置。 @@ -43,7 +44,7 @@ type LLMConfig struct { MaxTokens int `mapstructure:"max_tokens"` } -// ImageGenConfig 文生图模型配置。 +// ImageGenConfig 文生图模型配置(OpenAI 兼容同步 API)。 type ImageGenConfig struct { BaseURL string `mapstructure:"base_url"` APIKey string `mapstructure:"api_key"` @@ -55,6 +56,16 @@ type ImageGenConfig struct { CFGScale float64 `mapstructure:"cfg_scale"` } +// GptImage2Config GPT Image 2 异步生图模型配置(yuntts 等兼容服务)。 +type GptImage2Config struct { + BaseURL string `mapstructure:"base_url"` + APIKey string `mapstructure:"api_key"` + AspectRatio string `mapstructure:"aspect_ratio"` + XChannel string `mapstructure:"x_channel"` + PollMaxWait int `mapstructure:"poll_max_wait"` // 轮询最大等待秒数, 默认 120 + PollInterval int `mapstructure:"poll_interval"` // 轮询间隔秒数, 默认 3 +} + // Load 从 YAML 配置文件和环境变量加载配置。 // 优先级:环境变量 > YAML 文件 > 默认值。 func Load() *Config { @@ -111,6 +122,13 @@ func setDefaults(v *viper.Viper) { v.SetDefault("image_gen.num_images", 1) v.SetDefault("image_gen.steps", 30) v.SetDefault("image_gen.cfg_scale", 7.0) + + v.SetDefault("gpt_image2.base_url", "https://www.yuntts.com/api/v1") + v.SetDefault("gpt_image2.api_key", "") + v.SetDefault("gpt_image2.aspect_ratio", "1:1") + v.SetDefault("gpt_image2.x_channel", "default") + v.SetDefault("gpt_image2.poll_max_wait", 120) + v.SetDefault("gpt_image2.poll_interval", 3) } func bindEnvVars(v *viper.Viper) { @@ -135,4 +153,11 @@ func bindEnvVars(v *viper.Viper) { 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") + + v.BindEnv("gpt_image2.base_url", "GEN2D_GPT_IMAGE2_BASE_URL") + v.BindEnv("gpt_image2.api_key", "GEN2D_GPT_IMAGE2_API_KEY") + v.BindEnv("gpt_image2.aspect_ratio", "GEN2D_GPT_IMAGE2_ASPECT_RATIO") + v.BindEnv("gpt_image2.x_channel", "GEN2D_GPT_IMAGE2_X_CHANNEL") + v.BindEnv("gpt_image2.poll_max_wait", "GEN2D_GPT_IMAGE2_POLL_MAX_WAIT") + v.BindEnv("gpt_image2.poll_interval", "GEN2D_GPT_IMAGE2_POLL_INTERVAL") } diff --git a/backend/internal/service/gptimage.go b/backend/internal/service/gptimage.go new file mode 100755 index 0000000..a684b90 --- /dev/null +++ b/backend/internal/service/gptimage.go @@ -0,0 +1,268 @@ +package service + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "log" + "net/http" + "strings" + "time" + + "gen2d/internal/config" +) + +// gptImage2Cfg 保存 GPT Image 2 配置。 +var gptImage2Cfg config.GptImage2Config + +// InitGptImage2Config 注入 GPT Image 2 配置。 +func InitGptImage2Config(cfg config.GptImage2Config) { + gptImage2Cfg = cfg +} + +// ======================== GPT Image 2 API 类型 ======================== + +// gpt2SubmitRequest 提交生图任务请求体。 +type gpt2SubmitRequest struct { + Prompt string `json:"prompt"` + AspectRatio string `json:"aspect_ratio,omitempty"` + ReferenceImages []string `json:"reference_images,omitempty"` + XChannel string `json:"x_channel,omitempty"` +} + +// gpt2SubmitResponse 提交生图任务响应体。 +type gpt2SubmitResponse struct { + Code int `json:"code"` + Message string `json:"message"` + Data struct { + TaskID string `json:"task_id"` + Status string `json:"status"` + } `json:"data"` +} + +// gpt2StatusRequest 查询任务状态请求体。 +type gpt2StatusRequest struct { + TaskID string `json:"task_id"` +} + +// gpt2StatusResponse 查询任务状态响应体。 +type gpt2StatusResponse struct { + Code int `json:"code"` + Message string `json:"message"` + Data struct { + TaskID string `json:"task_id"` + Status string `json:"status"` + Progress int `json:"progress"` + ResultImageURL string `json:"result_image_url"` + ErrorMessage string `json:"error_message"` + } `json:"data"` +} + +// ======================== 公开接口 ======================== + +// GenerateImagesGpt2 通过 GPT Image 2 异步 API 生成图片。 +// count 张图通过并行提交+轮询实现。 +func GenerateImagesGpt2(ctx context.Context, prompt string, count int) ([]GeneratedImage, error) { + if gptImage2Cfg.APIKey == "" { + return nil, fmt.Errorf("gpt_image2 api_key not configured") + } + + // 并行提交任务 + type submitResult struct { + index int + taskID string + err error + } + results := make(chan submitResult, count) + for i := 0; i < count; i++ { + go func(idx int) { + taskID, err := submitGpt2Task(ctx, prompt) + results <- submitResult{index: idx, taskID: taskID, err: err} + }(i) + } + + // 收集 taskID + taskIDs := make([]string, count) + for i := 0; i < count; i++ { + r := <-results + if r.err != nil { + return nil, fmt.Errorf("submit task %d: %w", r.index, r.err) + } + taskIDs[r.index] = r.taskID + } + + // 并行轮询+下载 + type imageResult struct { + index int + image GeneratedImage + err error + } + imgResults := make(chan imageResult, count) + for i, tid := range taskIDs { + go func(idx int, taskID string) { + img, err := pollAndDownloadGpt2(ctx, taskID, prompt) + imgResults <- imageResult{index: idx, image: img, err: err} + }(i, tid) + } + + images := make([]GeneratedImage, count) + for i := 0; i < count; i++ { + r := <-imgResults + if r.err != nil { + return nil, fmt.Errorf("task %d: %w", r.index, r.err) + } + images[r.index] = r.image + } + + return images, nil +} + +// ======================== 内部实现 ======================== + +// submitGpt2Task 提交生图任务,返回 taskID。 +func submitGpt2Task(ctx context.Context, prompt string) (string, error) { + reqBody := gpt2SubmitRequest{ + Prompt: prompt, + AspectRatio: gptImage2Cfg.AspectRatio, + XChannel: gptImage2Cfg.XChannel, + } + + body, err := json.Marshal(reqBody) + if err != nil { + return "", fmt.Errorf("marshal: %w", err) + } + + url := strings.TrimRight(gptImage2Cfg.BaseURL, "/") + "/gpt-image2/generate" + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) + if err != nil { + return "", fmt.Errorf("create request: %w", err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+gptImage2Cfg.APIKey) + + resp, err := http.DefaultClient.Do(req) + if err != nil { + return "", fmt.Errorf("send: %w", err) + } + defer resp.Body.Close() + + var submitResp gpt2SubmitResponse + if err := json.NewDecoder(resp.Body).Decode(&submitResp); err != nil { + return "", fmt.Errorf("decode: %w", err) + } + + if submitResp.Code != 200 { + return "", fmt.Errorf("submit failed: %s", submitResp.Message) + } + + log.Printf("[gpt_image2] task submitted: %s", submitResp.Data.TaskID) + return submitResp.Data.TaskID, nil +} + +// pollAndDownloadGpt2 轮询任务状态直到完成,然后下载图片。 +func pollAndDownloadGpt2(ctx context.Context, taskID, prompt string) (GeneratedImage, error) { + pollInterval := gptImage2Cfg.PollInterval + if pollInterval <= 0 { + pollInterval = 3 + } + maxWait := gptImage2Cfg.PollMaxWait + if maxWait <= 0 { + maxWait = 120 + } + + deadline := time.Now().Add(time.Duration(maxWait) * time.Second) + ticker := time.NewTicker(time.Duration(pollInterval) * time.Second) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return GeneratedImage{}, ctx.Err() + case <-ticker.C: + status, err := queryGpt2Status(ctx, taskID) + if err != nil { + return GeneratedImage{}, fmt.Errorf("query status: %w", err) + } + + switch status.Data.Status { + case "completed": + log.Printf("[gpt_image2] task %s completed, downloading from %s", taskID, status.Data.ResultImageURL) + data, err := downloadGpt2Image(ctx, status.Data.ResultImageURL) + if err != nil { + return GeneratedImage{}, fmt.Errorf("download: %w", err) + } + return GeneratedImage{Data: data, Format: "png"}, nil + + case "failed": + errMsg := status.Data.ErrorMessage + if errMsg == "" { + errMsg = "unknown error" + } + return GeneratedImage{}, fmt.Errorf("generation failed: %s", errMsg) + + default: + log.Printf("[gpt_image2] task %s status=%s progress=%d", taskID, status.Data.Status, status.Data.Progress) + } + + if time.Now().After(deadline) { + return GeneratedImage{}, fmt.Errorf("poll timeout after %ds", maxWait) + } + } + } +} + +// queryGpt2Status 查询任务状态。 +func queryGpt2Status(ctx context.Context, taskID string) (*gpt2StatusResponse, error) { + reqBody := gpt2StatusRequest{TaskID: taskID} + body, err := json.Marshal(reqBody) + if err != nil { + return nil, fmt.Errorf("marshal: %w", err) + } + + url := strings.TrimRight(gptImage2Cfg.BaseURL, "/") + "/gpt-image2/status" + 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 "+gptImage2Cfg.APIKey) + + resp, err := http.DefaultClient.Do(req) + if err != nil { + return nil, fmt.Errorf("send: %w", err) + } + defer resp.Body.Close() + + var statusResp gpt2StatusResponse + if err := json.NewDecoder(resp.Body).Decode(&statusResp); err != nil { + return nil, fmt.Errorf("decode: %w", err) + } + + if statusResp.Code != 200 { + return nil, fmt.Errorf("status query failed: %s", statusResp.Message) + } + + return &statusResp, nil +} + +// downloadGpt2Image 下载生成的图片。 +func downloadGpt2Image(ctx context.Context, imageURL string) ([]byte, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, imageURL, 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) +} diff --git a/backend/internal/service/inference.go b/backend/internal/service/inference.go old mode 100644 new mode 100755 index bd55b0c..603b8d3 --- a/backend/internal/service/inference.go +++ b/backend/internal/service/inference.go @@ -48,13 +48,25 @@ type imageGenResponse struct { } `json:"data"` } -// GenerateImages 调用 AI 推理 API 生成图片,未配置 API key 时回退到 mock。 +// GenerateImages 调用 AI 推理 API 生成图片。 +// 优先级:GPT Image 2 > OpenAI 兼容 ImageGen > 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 } + // 优先使用 GPT Image 2 异步 API + if gptImage2Cfg.APIKey != "" { + log.Println("[inference] using GPT Image 2 async API") + images, err := GenerateImagesGpt2(ctx, prompt, count) + if err != nil { + return nil, fmt.Errorf("gpt_image2: %w", err) + } + return images, nil + } + + // 降级:OpenAI 兼容 Images API if imgCfg.APIKey != "" { width, height := imgCfg.Width, imgCfg.Height if params.Resolution > 0 { @@ -64,11 +76,12 @@ func GenerateImages(ctx context.Context, prompt string, params AssetParams) ([]G return callImageGenAPI(ctx, prompt, count, width, height) } + // 最终降级:mock 占位图 size := params.Resolution if size <= 0 { size = 64 } - log.Println("[inference] ImageGen API key not configured, using mock") + log.Println("[inference] no image API key configured, using mock") return generateMockImages(size, count) } From c4dc7394b1e7e36de60b81c942454a3e2c1f8a99 Mon Sep 17 00:00:00 2001 From: Gmaker689 <1711322114@qq.com> Date: Mon, 25 May 2026 12:44:39 +0800 Subject: [PATCH 5/8] =?UTF-8?q?refactor(inference):=20=E7=BB=9F=E5=90=88?= =?UTF-8?q?=20image=5Fgen=20=E4=B8=BA=20GPT=20Image=202=20=E5=BC=82?= =?UTF-8?q?=E6=AD=A5=20API=EF=BC=8C=E5=88=A0=E9=99=A4=E5=86=97=E4=BD=99?= =?UTF-8?q?=E9=85=8D=E7=BD=AE?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 删除 gptimage.go,异步逻辑直接内聚到 inference.go - ImageGenConfig 改为 GPT Image 2 参数:aspect_ratio / x_channel / poll_max_wait / poll_interval - GenerateImages 流程:submit → poll → download,支持并行多图(spritesheet) - 移除旧的 OpenAI 兼容同步 API 代码(callImageGenAPI / parseImageResponse 等) - 统一 env:GEN2D_IMAGE_* 系列,移除 GEN2D_GPT_IMAGE2_* 冗余 --- backend/.env.example | 27 +- backend/cmd/main.go | 1 - backend/internal/config/config.go | 77 ++---- backend/internal/config/config.yml | 14 +- backend/internal/service/gptimage.go | 268 ------------------ backend/internal/service/inference.go | 380 +++++++++++++++----------- 6 files changed, 255 insertions(+), 512 deletions(-) mode change 100644 => 100755 backend/internal/config/config.yml delete mode 100755 backend/internal/service/gptimage.go diff --git a/backend/.env.example b/backend/.env.example index 2501aa8..6789aac 100755 --- a/backend/.env.example +++ b/backend/.env.example @@ -19,21 +19,12 @@ GEN2D_LLM_MODEL=gpt-4o GEN2D_LLM_TEMPERATURE=0.7 GEN2D_LLM_MAX_TOKENS=2048 -# 文生图模型(OpenAI 兼容同步 API) -GEN2D_IMAGE_BASE_URL=https://api.stability.ai/v1 -GEN2D_IMAGE_API_KEY=sk-your-api-key -GEN2D_IMAGE_MODEL=stable-diffusion-xl -GEN2D_IMAGE_WIDTH=1024 -GEN2D_IMAGE_HEIGHT=1024 -GEN2D_IMAGE_NUM_IMAGES=1 -GEN2D_IMAGE_STEPS=30 -GEN2D_IMAGE_CFG_SCALE=7.0 - -# GPT Image 2 文生图(异步 API,如 yuntts 等兼容服务) -# 优先级高于 IMAGE API,留空则使用上方的 IMAGE API -GEN2D_GPT_IMAGE2_BASE_URL=https://www.yuntts.com/api/v1 -GEN2D_GPT_IMAGE2_API_KEY= -GEN2D_GPT_IMAGE2_ASPECT_RATIO=1:1 -GEN2D_GPT_IMAGE2_X_CHANNEL=default -GEN2D_GPT_IMAGE2_POLL_MAX_WAIT=120 -GEN2D_GPT_IMAGE2_POLL_INTERVAL=3 +# 文生图模型(GPT Image 2 异步 API,如 yuntts 等兼容服务) +GEN2D_IMAGE_BASE_URL=https://www.yuntts.com/api/v1 +GEN2D_IMAGE_API_KEY= +GEN2D_IMAGE_MODEL=gpt-image-2 +GEN2D_IMAGE_QUALITY=low +GEN2D_IMAGE_ASPECT_RATIO=1:1 +GEN2D_IMAGE_X_CHANNEL=default +GEN2D_IMAGE_POLL_MAX_WAIT=120 +GEN2D_IMAGE_POLL_INTERVAL=3 diff --git a/backend/cmd/main.go b/backend/cmd/main.go index af7578a..2a447dd 100755 --- a/backend/cmd/main.go +++ b/backend/cmd/main.go @@ -35,7 +35,6 @@ func main() { // 注入 LLM 和文生图配置到 service 层 service.InitLLMConfig(cfg.LLM) service.InitImageGenConfig(cfg.ImageGen) - service.InitGptImage2Config(cfg.GptImage2) r := gin.New() r.Use(gin.Recovery()) // panic 恢复中间件,防止服务因未捕获异常宕机 diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index 12dea7e..be4deb3 100755 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -9,12 +9,11 @@ import ( // Config 应用全局配置。 type Config struct { - Server ServerConfig `mapstructure:"server"` - Database DatabaseConfig `mapstructure:"database"` - JWT JWTConfig `mapstructure:"jwt"` - LLM LLMConfig `mapstructure:"llm"` - ImageGen ImageGenConfig `mapstructure:"image_gen"` - GptImage2 GptImage2Config `mapstructure:"gpt_image2"` + Server ServerConfig `mapstructure:"server"` + Database DatabaseConfig `mapstructure:"database"` + JWT JWTConfig `mapstructure:"jwt"` + LLM LLMConfig `mapstructure:"llm"` + ImageGen ImageGenConfig `mapstructure:"image_gen"` } // ServerConfig HTTP 服务配置。 @@ -44,26 +43,16 @@ type LLMConfig struct { MaxTokens int `mapstructure:"max_tokens"` } -// ImageGenConfig 文生图模型配置(OpenAI 兼容同步 API)。 +// ImageGenConfig GPT Image 2 异步生图模型配置。 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"` -} - -// GptImage2Config GPT Image 2 异步生图模型配置(yuntts 等兼容服务)。 -type GptImage2Config struct { - BaseURL string `mapstructure:"base_url"` - APIKey string `mapstructure:"api_key"` - AspectRatio string `mapstructure:"aspect_ratio"` - XChannel string `mapstructure:"x_channel"` - PollMaxWait int `mapstructure:"poll_max_wait"` // 轮询最大等待秒数, 默认 120 - PollInterval int `mapstructure:"poll_interval"` // 轮询间隔秒数, 默认 3 + BaseURL string `mapstructure:"base_url"` + APIKey string `mapstructure:"api_key"` + Model string `mapstructure:"model"` + Quality string `mapstructure:"quality"` + AspectRatio string `mapstructure:"aspect_ratio"` + XChannel string `mapstructure:"x_channel"` + PollMaxWait int `mapstructure:"poll_max_wait"` + PollInterval int `mapstructure:"poll_interval"` } // Load 从 YAML 配置文件和环境变量加载配置。 @@ -114,21 +103,14 @@ 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://www.yuntts.com/api/v1") v.SetDefault("image_gen.api_key", "") - v.SetDefault("image_gen.model", "stable-diffusion-xl") - v.SetDefault("image_gen.width", 1024) - v.SetDefault("image_gen.height", 1024) - v.SetDefault("image_gen.num_images", 1) - v.SetDefault("image_gen.steps", 30) - v.SetDefault("image_gen.cfg_scale", 7.0) - - v.SetDefault("gpt_image2.base_url", "https://www.yuntts.com/api/v1") - v.SetDefault("gpt_image2.api_key", "") - v.SetDefault("gpt_image2.aspect_ratio", "1:1") - v.SetDefault("gpt_image2.x_channel", "default") - v.SetDefault("gpt_image2.poll_max_wait", 120) - v.SetDefault("gpt_image2.poll_interval", 3) + v.SetDefault("image_gen.model", "gpt-image-2") + v.SetDefault("image_gen.quality", "low") + v.SetDefault("image_gen.aspect_ratio", "1:1") + v.SetDefault("image_gen.x_channel", "default") + v.SetDefault("image_gen.poll_max_wait", 120) + v.SetDefault("image_gen.poll_interval", 3) } func bindEnvVars(v *viper.Viper) { @@ -148,16 +130,9 @@ func bindEnvVars(v *viper.Viper) { v.BindEnv("image_gen.base_url", "GEN2D_IMAGE_BASE_URL") v.BindEnv("image_gen.api_key", "GEN2D_IMAGE_API_KEY") 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.num_images", "GEN2D_IMAGE_NUM_IMAGES") - v.BindEnv("image_gen.steps", "GEN2D_IMAGE_STEPS") - v.BindEnv("image_gen.cfg_scale", "GEN2D_IMAGE_CFG_SCALE") - - v.BindEnv("gpt_image2.base_url", "GEN2D_GPT_IMAGE2_BASE_URL") - v.BindEnv("gpt_image2.api_key", "GEN2D_GPT_IMAGE2_API_KEY") - v.BindEnv("gpt_image2.aspect_ratio", "GEN2D_GPT_IMAGE2_ASPECT_RATIO") - v.BindEnv("gpt_image2.x_channel", "GEN2D_GPT_IMAGE2_X_CHANNEL") - v.BindEnv("gpt_image2.poll_max_wait", "GEN2D_GPT_IMAGE2_POLL_MAX_WAIT") - v.BindEnv("gpt_image2.poll_interval", "GEN2D_GPT_IMAGE2_POLL_INTERVAL") + v.BindEnv("image_gen.quality", "GEN2D_IMAGE_QUALITY") + v.BindEnv("image_gen.aspect_ratio", "GEN2D_IMAGE_ASPECT_RATIO") + v.BindEnv("image_gen.x_channel", "GEN2D_IMAGE_X_CHANNEL") + v.BindEnv("image_gen.poll_max_wait", "GEN2D_IMAGE_POLL_MAX_WAIT") + v.BindEnv("image_gen.poll_interval", "GEN2D_IMAGE_POLL_INTERVAL") } diff --git a/backend/internal/config/config.yml b/backend/internal/config/config.yml old mode 100644 new mode 100755 index 84475e7..82d076b --- a/backend/internal/config/config.yml +++ b/backend/internal/config/config.yml @@ -21,11 +21,11 @@ llm: max_tokens: 2048 image_gen: - base_url: "https://api.stability.ai/v1" + base_url: "https://www.yuntts.com/api/v1" api_key: "" - model: "stable-diffusion-xl" - width: 1024 - height: 1024 - num_images: 1 - steps: 30 - cfg_scale: 7.0 + model: "gpt-image-2" + quality: "low" + aspect_ratio: "1:1" + x_channel: "default" + poll_max_wait: 120 + poll_interval: 3 diff --git a/backend/internal/service/gptimage.go b/backend/internal/service/gptimage.go deleted file mode 100755 index a684b90..0000000 --- a/backend/internal/service/gptimage.go +++ /dev/null @@ -1,268 +0,0 @@ -package service - -import ( - "bytes" - "context" - "encoding/json" - "fmt" - "io" - "log" - "net/http" - "strings" - "time" - - "gen2d/internal/config" -) - -// gptImage2Cfg 保存 GPT Image 2 配置。 -var gptImage2Cfg config.GptImage2Config - -// InitGptImage2Config 注入 GPT Image 2 配置。 -func InitGptImage2Config(cfg config.GptImage2Config) { - gptImage2Cfg = cfg -} - -// ======================== GPT Image 2 API 类型 ======================== - -// gpt2SubmitRequest 提交生图任务请求体。 -type gpt2SubmitRequest struct { - Prompt string `json:"prompt"` - AspectRatio string `json:"aspect_ratio,omitempty"` - ReferenceImages []string `json:"reference_images,omitempty"` - XChannel string `json:"x_channel,omitempty"` -} - -// gpt2SubmitResponse 提交生图任务响应体。 -type gpt2SubmitResponse struct { - Code int `json:"code"` - Message string `json:"message"` - Data struct { - TaskID string `json:"task_id"` - Status string `json:"status"` - } `json:"data"` -} - -// gpt2StatusRequest 查询任务状态请求体。 -type gpt2StatusRequest struct { - TaskID string `json:"task_id"` -} - -// gpt2StatusResponse 查询任务状态响应体。 -type gpt2StatusResponse struct { - Code int `json:"code"` - Message string `json:"message"` - Data struct { - TaskID string `json:"task_id"` - Status string `json:"status"` - Progress int `json:"progress"` - ResultImageURL string `json:"result_image_url"` - ErrorMessage string `json:"error_message"` - } `json:"data"` -} - -// ======================== 公开接口 ======================== - -// GenerateImagesGpt2 通过 GPT Image 2 异步 API 生成图片。 -// count 张图通过并行提交+轮询实现。 -func GenerateImagesGpt2(ctx context.Context, prompt string, count int) ([]GeneratedImage, error) { - if gptImage2Cfg.APIKey == "" { - return nil, fmt.Errorf("gpt_image2 api_key not configured") - } - - // 并行提交任务 - type submitResult struct { - index int - taskID string - err error - } - results := make(chan submitResult, count) - for i := 0; i < count; i++ { - go func(idx int) { - taskID, err := submitGpt2Task(ctx, prompt) - results <- submitResult{index: idx, taskID: taskID, err: err} - }(i) - } - - // 收集 taskID - taskIDs := make([]string, count) - for i := 0; i < count; i++ { - r := <-results - if r.err != nil { - return nil, fmt.Errorf("submit task %d: %w", r.index, r.err) - } - taskIDs[r.index] = r.taskID - } - - // 并行轮询+下载 - type imageResult struct { - index int - image GeneratedImage - err error - } - imgResults := make(chan imageResult, count) - for i, tid := range taskIDs { - go func(idx int, taskID string) { - img, err := pollAndDownloadGpt2(ctx, taskID, prompt) - imgResults <- imageResult{index: idx, image: img, err: err} - }(i, tid) - } - - images := make([]GeneratedImage, count) - for i := 0; i < count; i++ { - r := <-imgResults - if r.err != nil { - return nil, fmt.Errorf("task %d: %w", r.index, r.err) - } - images[r.index] = r.image - } - - return images, nil -} - -// ======================== 内部实现 ======================== - -// submitGpt2Task 提交生图任务,返回 taskID。 -func submitGpt2Task(ctx context.Context, prompt string) (string, error) { - reqBody := gpt2SubmitRequest{ - Prompt: prompt, - AspectRatio: gptImage2Cfg.AspectRatio, - XChannel: gptImage2Cfg.XChannel, - } - - body, err := json.Marshal(reqBody) - if err != nil { - return "", fmt.Errorf("marshal: %w", err) - } - - url := strings.TrimRight(gptImage2Cfg.BaseURL, "/") + "/gpt-image2/generate" - req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) - if err != nil { - return "", fmt.Errorf("create request: %w", err) - } - req.Header.Set("Content-Type", "application/json") - req.Header.Set("Authorization", "Bearer "+gptImage2Cfg.APIKey) - - resp, err := http.DefaultClient.Do(req) - if err != nil { - return "", fmt.Errorf("send: %w", err) - } - defer resp.Body.Close() - - var submitResp gpt2SubmitResponse - if err := json.NewDecoder(resp.Body).Decode(&submitResp); err != nil { - return "", fmt.Errorf("decode: %w", err) - } - - if submitResp.Code != 200 { - return "", fmt.Errorf("submit failed: %s", submitResp.Message) - } - - log.Printf("[gpt_image2] task submitted: %s", submitResp.Data.TaskID) - return submitResp.Data.TaskID, nil -} - -// pollAndDownloadGpt2 轮询任务状态直到完成,然后下载图片。 -func pollAndDownloadGpt2(ctx context.Context, taskID, prompt string) (GeneratedImage, error) { - pollInterval := gptImage2Cfg.PollInterval - if pollInterval <= 0 { - pollInterval = 3 - } - maxWait := gptImage2Cfg.PollMaxWait - if maxWait <= 0 { - maxWait = 120 - } - - deadline := time.Now().Add(time.Duration(maxWait) * time.Second) - ticker := time.NewTicker(time.Duration(pollInterval) * time.Second) - defer ticker.Stop() - - for { - select { - case <-ctx.Done(): - return GeneratedImage{}, ctx.Err() - case <-ticker.C: - status, err := queryGpt2Status(ctx, taskID) - if err != nil { - return GeneratedImage{}, fmt.Errorf("query status: %w", err) - } - - switch status.Data.Status { - case "completed": - log.Printf("[gpt_image2] task %s completed, downloading from %s", taskID, status.Data.ResultImageURL) - data, err := downloadGpt2Image(ctx, status.Data.ResultImageURL) - if err != nil { - return GeneratedImage{}, fmt.Errorf("download: %w", err) - } - return GeneratedImage{Data: data, Format: "png"}, nil - - case "failed": - errMsg := status.Data.ErrorMessage - if errMsg == "" { - errMsg = "unknown error" - } - return GeneratedImage{}, fmt.Errorf("generation failed: %s", errMsg) - - default: - log.Printf("[gpt_image2] task %s status=%s progress=%d", taskID, status.Data.Status, status.Data.Progress) - } - - if time.Now().After(deadline) { - return GeneratedImage{}, fmt.Errorf("poll timeout after %ds", maxWait) - } - } - } -} - -// queryGpt2Status 查询任务状态。 -func queryGpt2Status(ctx context.Context, taskID string) (*gpt2StatusResponse, error) { - reqBody := gpt2StatusRequest{TaskID: taskID} - body, err := json.Marshal(reqBody) - if err != nil { - return nil, fmt.Errorf("marshal: %w", err) - } - - url := strings.TrimRight(gptImage2Cfg.BaseURL, "/") + "/gpt-image2/status" - 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 "+gptImage2Cfg.APIKey) - - resp, err := http.DefaultClient.Do(req) - if err != nil { - return nil, fmt.Errorf("send: %w", err) - } - defer resp.Body.Close() - - var statusResp gpt2StatusResponse - if err := json.NewDecoder(resp.Body).Decode(&statusResp); err != nil { - return nil, fmt.Errorf("decode: %w", err) - } - - if statusResp.Code != 200 { - return nil, fmt.Errorf("status query failed: %s", statusResp.Message) - } - - return &statusResp, nil -} - -// downloadGpt2Image 下载生成的图片。 -func downloadGpt2Image(ctx context.Context, imageURL string) ([]byte, error) { - req, err := http.NewRequestWithContext(ctx, http.MethodGet, imageURL, 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) -} diff --git a/backend/internal/service/inference.go b/backend/internal/service/inference.go index 603b8d3..c52a24e 100755 --- a/backend/internal/service/inference.go +++ b/backend/internal/service/inference.go @@ -4,7 +4,6 @@ import ( "bytes" "context" "crypto/rand" - "encoding/base64" "encoding/json" "fmt" "image" @@ -12,10 +11,9 @@ import ( "image/png" "io" "log" - "mime/multipart" "net/http" - "strconv" "strings" + "time" "gen2d/internal/config" ) @@ -28,84 +26,230 @@ func InitImageGenConfig(cfg config.ImageGenConfig) { imgCfg = cfg } -// ======================== 文生图 API 调用层 ======================== +// ======================== GPT Image 2 API 类型 ======================== -// imageGenRequest OpenAI 兼容的文生图请求体。 -type imageGenRequest struct { - Model string `json:"model"` - Prompt string `json:"prompt"` - N int `json:"n,omitempty"` - Size string `json:"size,omitempty"` - Steps int `json:"steps,omitempty"` - CFGScale float64 `json:"cfg_scale,omitempty"` +// genSubmitReq 提交生图任务请求体。 +type genSubmitReq struct { + Prompt string `json:"prompt"` + AspectRatio string `json:"aspect_ratio,omitempty"` + ReferenceImages []string `json:"reference_images,omitempty"` + XChannel string `json:"x_channel,omitempty"` } -// imageGenResponse OpenAI 兼容的文生图响应体。 -type imageGenResponse struct { - Data []struct { - URL string `json:"url"` - B64JSON string `json:"b64_json"` +// genSubmitResp 提交生图任务响应体。 +type genSubmitResp struct { + Code int `json:"code"` + Message string `json:"message"` + Data struct { + TaskID string `json:"task_id"` + Status string `json:"status"` } `json:"data"` } -// GenerateImages 调用 AI 推理 API 生成图片。 -// 优先级:GPT Image 2 > OpenAI 兼容 ImageGen > Mock 回退。 +// genStatusReq 查询任务状态请求体。 +type genStatusReq struct { + TaskID string `json:"task_id"` +} + +// genStatusResp 查询任务状态响应体。 +type genStatusResp struct { + Code int `json:"code"` + Message string `json:"message"` + Data struct { + TaskID string `json:"task_id"` + Status string `json:"status"` + Progress int `json:"progress"` + ResultImageURL string `json:"result_image_url"` + ErrorMessage string `json:"error_message"` + } `json:"data"` +} + +// ======================== 文生图 ======================== + +// GenerateImages 通过 GPT Image 2 异步 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 } - // 优先使用 GPT Image 2 异步 API - if gptImage2Cfg.APIKey != "" { - log.Println("[inference] using GPT Image 2 async API") - images, err := GenerateImagesGpt2(ctx, prompt, count) - if err != nil { - return nil, fmt.Errorf("gpt_image2: %w", err) - } - return images, nil - } - - // 降级:OpenAI 兼容 Images API if imgCfg.APIKey != "" { - width, height := imgCfg.Width, imgCfg.Height - if params.Resolution > 0 { - width = params.Resolution - height = params.Resolution - } - return callImageGenAPI(ctx, prompt, count, width, height) + log.Println("[inference] using image gen async API") + return generateAsync(ctx, prompt, count, nil) } - // 最终降级:mock 占位图 size := params.Resolution if size <= 0 { size = 64 } - log.Println("[inference] no image API key configured, using mock") + log.Println("[inference] image API key not configured, using mock") return generateMockImages(size, count) } -// callImageGenAPI 调用 OpenAI 兼容的 Images API,返回生成的图片。 -func callImageGenAPI(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), +// ======================== 图片编辑 ======================== + +// EditImages 图片编辑接口,以参考图模式提交 GPT Image 2 编辑任务。 +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") } - if imgCfg.Steps > 0 { - reqBody.Steps = imgCfg.Steps + log.Println("[inference] using image gen async API for edit") + // 暂不支持参考图 URL 模式,提示用户 + return nil, fmt.Errorf("image edit with reference image: not yet implemented for async API") +} + +// ======================== 异步核心 ======================== + +// generateAsync 并行提交 count 个任务,轮询完成后下载图片。 +func generateAsync(ctx context.Context, prompt string, count int, refImages []string) ([]GeneratedImage, error) { + type submitResult struct { + index int + taskID string + err error } - if imgCfg.CFGScale > 0 { - reqBody.CFGScale = imgCfg.CFGScale + results := make(chan submitResult, count) + for i := 0; i < count; i++ { + go func(idx int) { + taskID, err := submitTask(ctx, prompt, refImages) + results <- submitResult{index: idx, taskID: taskID, err: err} + }(i) + } + + taskIDs := make([]string, count) + for i := 0; i < count; i++ { + r := <-results + if r.err != nil { + return nil, fmt.Errorf("submit task %d: %w", r.index, r.err) + } + taskIDs[r.index] = r.taskID + } + + type imageResult struct { + index int + image GeneratedImage + err error + } + imgResults := make(chan imageResult, count) + for i, tid := range taskIDs { + go func(idx int, taskID string) { + img, err := pollAndDownload(ctx, taskID) + imgResults <- imageResult{index: idx, image: img, err: err} + }(i, tid) + } + + images := make([]GeneratedImage, count) + for i := 0; i < count; i++ { + r := <-imgResults + if r.err != nil { + return nil, fmt.Errorf("task %d: %w", r.index, r.err) + } + images[r.index] = r.image + } + + return images, nil +} + +// submitTask 提交生图任务,返回 taskID。 +func submitTask(ctx context.Context, prompt string, refImages []string) (string, error) { + reqBody := genSubmitReq{ + Prompt: prompt, + AspectRatio: imgCfg.AspectRatio, + ReferenceImages: refImages, + XChannel: imgCfg.XChannel, } body, err := json.Marshal(reqBody) if err != nil { - return nil, fmt.Errorf("marshal request: %w", err) + return "", fmt.Errorf("marshal: %w", err) } - url := strings.TrimRight(imgCfg.BaseURL, "/") + url := strings.TrimRight(imgCfg.BaseURL, "/") + "/gpt-image2/generate" + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) + if err != nil { + return "", 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 "", fmt.Errorf("send: %w", err) + } + defer resp.Body.Close() + + var sr genSubmitResp + if err := json.NewDecoder(resp.Body).Decode(&sr); err != nil { + return "", fmt.Errorf("decode: %w", err) + } + + if sr.Code != 200 { + return "", fmt.Errorf("submit failed: %s", sr.Message) + } + + log.Printf("[inference] task submitted: %s", sr.Data.TaskID) + return sr.Data.TaskID, nil +} + +// pollAndDownload 轮询任务直到完成,下载图片。 +func pollAndDownload(ctx context.Context, taskID string) (GeneratedImage, error) { + pollInterval := imgCfg.PollInterval + if pollInterval <= 0 { + pollInterval = 3 + } + maxWait := imgCfg.PollMaxWait + if maxWait <= 0 { + maxWait = 120 + } + + deadline := time.Now().Add(time.Duration(maxWait) * time.Second) + ticker := time.NewTicker(time.Duration(pollInterval) * time.Second) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return GeneratedImage{}, ctx.Err() + case <-ticker.C: + st, err := queryStatus(ctx, taskID) + if err != nil { + return GeneratedImage{}, fmt.Errorf("query status: %w", err) + } + + switch st.Data.Status { + case "completed": + log.Printf("[inference] task %s completed, downloading", taskID) + data, err := downloadResult(ctx, st.Data.ResultImageURL) + if err != nil { + return GeneratedImage{}, fmt.Errorf("download: %w", err) + } + return GeneratedImage{Data: data, Format: "png"}, nil + + case "failed": + errMsg := st.Data.ErrorMessage + if errMsg == "" { + errMsg = "unknown error" + } + return GeneratedImage{}, fmt.Errorf("generation failed: %s", errMsg) + + default: + log.Printf("[inference] task %s status=%s progress=%d", taskID, st.Data.Status, st.Data.Progress) + } + + if time.Now().After(deadline) { + return GeneratedImage{}, fmt.Errorf("poll timeout after %ds", maxWait) + } + } + } +} + +// queryStatus 查询任务状态。 +func queryStatus(ctx context.Context, taskID string) (*genStatusResp, error) { + body, err := json.Marshal(genStatusReq{TaskID: taskID}) + if err != nil { + return nil, fmt.Errorf("marshal: %w", err) + } + + url := strings.TrimRight(imgCfg.BaseURL, "/") + "/gpt-image2/status" req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) if err != nil { return nil, fmt.Errorf("create request: %w", err) @@ -115,123 +259,39 @@ func callImageGenAPI(ctx context.Context, prompt string, count, width, height in resp, err := http.DefaultClient.Do(req) if err != nil { - return nil, fmt.Errorf("send request: %w", err) + return nil, fmt.Errorf("send: %w", err) } defer resp.Body.Close() - if resp.StatusCode != http.StatusOK { - b, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) - return nil, fmt.Errorf("image gen api error %d: %s", resp.StatusCode, string(b)) + var sr genStatusResp + if err := json.NewDecoder(resp.Body).Decode(&sr); err != nil { + return nil, fmt.Errorf("decode: %w", err) } - return parseImageResponse(ctx, resp.Body, width, height) + if sr.Code != 200 { + return nil, fmt.Errorf("status query failed: %s", sr.Message) + } + + return &sr, nil } -// ======================== 图片编辑 API ======================== - -// EditImages 调用图片编辑 API,基于已有图片和文本提示词生成修改后的图片。 -func EditImages(ctx context.Context, imageData []byte, prompt string, count int) ([]GeneratedImage, error) { - if imgCfg.APIKey == "" { - return nil, fmt.Errorf("image edit API key not configured") - } - return callImageEditAPI(ctx, imageData, prompt, count) -} - -// callImageEditAPI 调用 OpenAI 兼容的 Images Edits API(multipart/form-data)。 -func callImageEditAPI(ctx context.Context, imageData []byte, prompt string, count int) ([]GeneratedImage, error) { - 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", strconv.Itoa(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.Replace(strings.TrimRight(imgCfg.BaseURL, "/"), "generations", "edits", 1) - - 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) -} - -// parseImageResponse 解析 OpenAI 兼容的图片生成/编辑响应体。 -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: width, - Height: height, - Format: "png", - }) - } - return images, nil -} - -// downloadImage 从 URL 下载图片数据。 -func downloadImage(ctx context.Context, url string) ([]byte, error) { - req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) +// downloadResult 下载生成的图片。 +func downloadResult(ctx context.Context, imageURL string) ([]byte, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, imageURL, 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) } @@ -245,12 +305,10 @@ func CheckQuality(ctx context.Context, images []GeneratedImage, style map[string 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) { @@ -258,20 +316,18 @@ 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 } } // ======================== Mock 回退 ======================== -// generateMockImages 批量生成 mock PNG 占位图。 func generateMockImages(size, count int) ([]GeneratedImage, error) { images := make([]GeneratedImage, count) for i := 0; i < count; i++ { @@ -279,30 +335,21 @@ func generateMockImages(size, count int) ([]GeneratedImage, error) { 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", - } + images[i] = GeneratedImage{Data: data, Width: size, Height: size, Format: "png"} } return images, nil } -// generateMockImage 生成一张带随机色块的 PNG 占位图。 func generateMockImage(size int, seed int) ([]byte, error) { img := image.NewRGBA(image.Rect(0, 0, size, size)) - 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 @@ -310,7 +357,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) From 322e6aecbbb59f61157843d6c61b9646e738d9ec Mon Sep 17 00:00:00 2001 From: Gmaker689 <1711322114@qq.com> Date: Mon, 25 May 2026 12:51:49 +0800 Subject: [PATCH 6/8] =?UTF-8?q?feat(generate):=20=E7=94=9F=E6=88=90?= =?UTF-8?q?=E7=AE=A1=E7=BA=BF=E6=94=B9=E4=B8=BA=E5=BC=82=E6=AD=A5=EF=BC=8C?= =?UTF-8?q?=E5=89=8D=E7=AB=AF=E6=8E=A5=E5=85=A5=E8=BD=AE=E8=AF=A2=E8=BF=9B?= =?UTF-8?q?=E5=BA=A6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 后端: - POST /api/v1/generate 改为异步模式,立即返回 taskId - 管线在后台 goroutine 执行,状态通过 GET /tasks/:taskId 轮询 - 任务状态流转: pending → running → saving → completed/failed 前端: - generation store: submit 后每 2s 轮询 getTask,完成时自动 getAssets - GeneratePage: 实时显示轮询状态文本 + 进度条 - ResultPage: 挂载时轮询任务,未完成显示骨架屏+进度,完成自动加载素材 - Types: Task.status 增加 submitted / saving 状态 --- backend/internal/handler/generate.go | 102 +++++++++++++++------ frontend/src/api/generate.ts | 7 -- frontend/src/api/types.ts | 14 +-- frontend/src/pages/GeneratePage.tsx | 3 +- frontend/src/pages/ResultPage.tsx | 127 ++++++++++++++++++--------- frontend/src/stores/generation.ts | 95 +++++++++++++------- 6 files changed, 230 insertions(+), 118 deletions(-) diff --git a/backend/internal/handler/generate.go b/backend/internal/handler/generate.go index ebf5ca7..4f54ddc 100755 --- a/backend/internal/handler/generate.go +++ b/backend/internal/handler/generate.go @@ -1,7 +1,9 @@ package handler import ( + "context" "fmt" + "log" "net/http" "os" "path/filepath" @@ -31,9 +33,7 @@ type GenerateRequest struct { // GenerateResponse 素材生成响应体。 type GenerateResponse struct { - TaskID string `json:"taskId"` - Assets []AssetResponse `json:"assets"` - Metadata service.AssetMetadata `json:"metadata"` + TaskID string `json:"taskId"` } // AssetResponse 单个素材响应。 @@ -44,21 +44,21 @@ type AssetResponse struct { // 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"` + 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"` + Assets []AssetResponse `json:"assets"` + Metadata service.AssetMetadata `json:"metadata"` } // taskRecord 内存中的任务记录。 @@ -72,7 +72,8 @@ var ( taskStore = sync.Map{} // taskID → *taskRecord ) -// Generate 素材生成接口。 +// Generate 素材生成接口(异步)。 +// 立即返回 taskId,后台执行管线,前端通过 GET /tasks/:taskId 轮询进度。 func Generate(c *gin.Context) { var req GenerateRequest if err := c.ShouldBindJSON(&req); err != nil { @@ -85,6 +86,31 @@ func Generate(c *gin.Context) { 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, @@ -105,16 +131,20 @@ func Generate(c *gin.Context) { }, } - output, err := service.RunPipeline(c.Request.Context(), in) + ctx := context.Background() + output, err := service.RunPipeline(ctx, in) if err != nil { - c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "素材生成失败: "+err.Error())) + 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 { - c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "创建输出目录失败: "+err.Error())) + updateFailed(taskID, "创建输出目录失败: "+err.Error()) return } @@ -123,7 +153,7 @@ func Generate(c *gin.Context) { filename := fmt.Sprintf("%d.%s", i, a.Format) filePath := filepath.Join(genDir, filename) if err := os.WriteFile(filePath, a.Data, 0644); err != nil { - c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "保存图片失败: "+err.Error())) + updateFailed(taskID, "保存图片失败: "+err.Error()) return } assets[i] = AssetResponse{ @@ -132,7 +162,7 @@ func Generate(c *gin.Context) { } } - // 存储任务记录到内存 + // 更新为完成状态 taskStore.Store(taskID, &taskRecord{ task: TaskResponse{ ID: taskID, @@ -147,11 +177,29 @@ func Generate(c *gin.Context) { metadata: output.Metadata, }) - c.JSON(http.StatusOK, model.OK(GenerateResponse{ - TaskID: taskID, - 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 查询任务信息。 @@ -175,6 +223,10 @@ func GetAssets(c *gin.Context) { 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, diff --git a/frontend/src/api/generate.ts b/frontend/src/api/generate.ts index 971d9b2..a03b2c6 100755 --- a/frontend/src/api/generate.ts +++ b/frontend/src/api/generate.ts @@ -19,13 +19,6 @@ export async function getTask(taskId: string): Promise { export async function getAssets(taskId: string): Promise { const resp = await get(`/api/v1/tasks/${taskId}/assets`) - return toAssetList(resp) -} - -/** 将响应转为 Asset[] 供前端组件使用 */ -function toAssetList( - resp: AssetsResponse | GenerateResponse, -): Asset[] { return resp.assets.map((a, i) => ({ id: `asset-${i}`, url: a.url, diff --git a/frontend/src/api/types.ts b/frontend/src/api/types.ts index 5c5f8a9..3fbdefd 100755 --- a/frontend/src/api/types.ts +++ b/frontend/src/api/types.ts @@ -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 @@ -106,19 +106,9 @@ export interface GenerateRequest { format?: 'spritesheet' | 'individual' } -// 生成响应 — 对应 POST /api/v1/generate 返回 +// 生成响应 — 对应 POST /api/v1/generate 返回(异步,仅含 taskId) export interface GenerateResponse { taskId: string - assets: { - url: string - format: string - }[] - metadata: { - frameWidth: number - frameHeight: number - frameCount: number - directions: number - } } // 素材列表响应 — 对应 GET /api/v1/tasks/:taskId/assets diff --git a/frontend/src/pages/GeneratePage.tsx b/frontend/src/pages/GeneratePage.tsx index d21356c..1912994 100755 --- a/frontend/src/pages/GeneratePage.tsx +++ b/frontend/src/pages/GeneratePage.tsx @@ -21,6 +21,7 @@ export default function GeneratePage() { status, progress, taskId, + statusText, submit, reset: resetGeneration, } = useGenerationStore() @@ -91,7 +92,7 @@ export default function GeneratePage() { {status === 'running' && (

- 管线执行中,请稍候... + {statusText || '管线执行中,请稍候...'}

)} diff --git a/frontend/src/pages/ResultPage.tsx b/frontend/src/pages/ResultPage.tsx index 9031397..85da8b3 100755 --- a/frontend/src/pages/ResultPage.tsx +++ b/frontend/src/pages/ResultPage.tsx @@ -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(null) const [assets, setAssets] = useState([]) const [loading, setLoading] = useState(true) + const [polling, setPolling] = useState(false) + const pollRef = useRef>() 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 (
@@ -30,6 +58,13 @@ export default function ResultPage() {
+ {task && ( +

+ {task.status === 'pending' && '任务排队中...'} + {task.status === 'running' && `生成中... ${task.progress ?? 0}%`} + {task.status === 'submitted' && '已提交,等待处理...'} +

+ )}
@@ -53,27 +88,35 @@ export default function ResultPage() {

生成结果

- {/* 任务信息 */}
-

任务信息

-
+
+

任务信息

+
+ {assets.length > 0 && ( + + )} + + 继续生成 + +
+
+
提示词 {task.prompt} 素材类型 {task.assetType} 状态 - + {task.status === 'completed' ? '已完成' : '失败'} 创建时间 @@ -93,29 +136,29 @@ export default function ResultPage() {
- {/* 素材预览 */}
-
-

素材预览

- {assets.length > 0 && ( - - )} -
+

素材预览

) } + +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') + } + } +} diff --git a/frontend/src/stores/generation.ts b/frontend/src/stores/generation.ts index bb3ec9c..d7346d1 100755 --- a/frontend/src/stores/generation.ts +++ b/frontend/src/stores/generation.ts @@ -1,6 +1,6 @@ import { create } from 'zustand' -import type { Asset, GenerateRequest, GenerateResponse } from '../api/types' -import { submitGenerate } from '../api/generate' +import type { Asset, GenerateRequest } from '../api/types' +import { submitGenerate, getTask, getAssets } from '../api/generate' type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed' @@ -9,66 +9,99 @@ interface GenerationState { projectId: string | null progress: number status: Status + statusText: string assets: Asset[] error: string | null submit: (req: GenerateRequest) => Promise reset: () => void } -export const useGenerationStore = create((set) => ({ +let pollTimer: ReturnType | null = null + +function stopPolling() { + if (pollTimer) { + clearInterval(pollTimer) + pollTimer = null + } +} + +export const useGenerationStore = create((set, get) => ({ taskId: null, projectId: null, progress: 0, status: 'idle', + statusText: '', assets: [], error: null, submit: async (req) => { - set({ status: 'submitting', error: null }) + stopPolling() + set({ status: 'submitting', error: null, statusText: '提交中...' }) try { - set({ status: 'running', progress: 30 }) - const result = await submitGenerate(req) - - set({ progress: 80 }) - const assets = mapAssets(result) + const { taskId } = await submitGenerate(req) set({ - taskId: result.taskId, + taskId, projectId: req.projectId, - status: 'completed', - progress: 100, - assets, + status: 'running', + progress: 10, + statusText: '任务已提交,等待生成...', }) + + // 开始轮询进度 + 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) { - const errMsg = (err as Error).message - set({ status: 'failed', error: errMsg }) + stopPolling() + set({ status: 'failed', error: (err as Error).message, statusText: '提交失败' }) } }, reset: () => { + stopPolling() set({ taskId: null, projectId: null, progress: 0, status: 'idle', + statusText: '', assets: [], error: null, }) }, })) - -function mapAssets(resp: GenerateResponse): Asset[] { - 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, - }, - })) -} From 8c027d9b52829487b7876ed66e2331a66663afec Mon Sep 17 00:00:00 2001 From: Gmaker689 <1711322114@qq.com> Date: Mon, 25 May 2026 13:01:13 +0800 Subject: [PATCH 7/8] =?UTF-8?q?fix(inference):=20=E5=9B=9E=E9=80=80?= =?UTF-8?q?=E4=B8=BA=20OpenAI=20=E5=85=BC=E5=AE=B9=E5=90=8C=E6=AD=A5=20Ima?= =?UTF-8?q?ges=20API=EF=BC=8C=E6=81=A2=E5=A4=8D=E5=9B=BE=E7=89=87=E7=BC=96?= =?UTF-8?q?=E8=BE=91=E5=8A=9F=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 使用标准 POST /images/generations 端点(suchuang.vip 兼容) - 恢复 EditImages 的 multipart/form-data 实现 - 移除 yuntts 异步 submit/poll 流程 - 配置精简: base_url / api_key / model / width / height / quality --- backend/internal/service/inference.go | 338 ++++++++++---------------- 1 file changed, 127 insertions(+), 211 deletions(-) diff --git a/backend/internal/service/inference.go b/backend/internal/service/inference.go index c52a24e..f334a34 100755 --- a/backend/internal/service/inference.go +++ b/backend/internal/service/inference.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "crypto/rand" + "encoding/base64" "encoding/json" "fmt" "image" @@ -11,9 +12,9 @@ import ( "image/png" "io" "log" + "mime/multipart" "net/http" "strings" - "time" "gen2d/internal/config" ) @@ -26,47 +27,28 @@ func InitImageGenConfig(cfg config.ImageGenConfig) { imgCfg = cfg } -// ======================== GPT Image 2 API 类型 ======================== +// ======================== OpenAI 兼容 Images API 类型 ======================== -// genSubmitReq 提交生图任务请求体。 -type genSubmitReq struct { - Prompt string `json:"prompt"` - AspectRatio string `json:"aspect_ratio,omitempty"` - ReferenceImages []string `json:"reference_images,omitempty"` - XChannel string `json:"x_channel,omitempty"` +// 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"` } -// genSubmitResp 提交生图任务响应体。 -type genSubmitResp struct { - Code int `json:"code"` - Message string `json:"message"` - Data struct { - TaskID string `json:"task_id"` - Status string `json:"status"` - } `json:"data"` -} - -// genStatusReq 查询任务状态请求体。 -type genStatusReq struct { - TaskID string `json:"task_id"` -} - -// genStatusResp 查询任务状态响应体。 -type genStatusResp struct { - Code int `json:"code"` - Message string `json:"message"` - Data struct { - TaskID string `json:"task_id"` - Status string `json:"status"` - Progress int `json:"progress"` - ResultImageURL string `json:"result_image_url"` - ErrorMessage string `json:"error_message"` +// imageGenResponse OpenAI 兼容文生图响应体。 +type imageGenResponse struct { + Data []struct { + URL string `json:"url"` + B64JSON string `json:"b64_json"` } `json:"data"` } // ======================== 文生图 ======================== -// GenerateImages 通过 GPT Image 2 异步 API 生成图片,未配置 key 时回退到 mock。 +// 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 { @@ -74,8 +56,13 @@ func GenerateImages(ctx context.Context, prompt string, params AssetParams) ([]G } if imgCfg.APIKey != "" { - log.Println("[inference] using image gen async API") - return generateAsync(ctx, prompt, count, nil) + width, height := imgCfg.Width, imgCfg.Height + if params.Resolution > 0 { + width = params.Resolution + height = params.Resolution + } + log.Println("[inference] calling image gen API") + return callImageAPI(ctx, prompt, count, width, height) } size := params.Resolution @@ -86,170 +73,22 @@ func GenerateImages(ctx context.Context, prompt string, params AssetParams) ([]G return generateMockImages(size, count) } -// ======================== 图片编辑 ======================== - -// EditImages 图片编辑接口,以参考图模式提交 GPT Image 2 编辑任务。 -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") - } - log.Println("[inference] using image gen async API for edit") - // 暂不支持参考图 URL 模式,提示用户 - return nil, fmt.Errorf("image edit with reference image: not yet implemented for async API") -} - -// ======================== 异步核心 ======================== - -// generateAsync 并行提交 count 个任务,轮询完成后下载图片。 -func generateAsync(ctx context.Context, prompt string, count int, refImages []string) ([]GeneratedImage, error) { - type submitResult struct { - index int - taskID string - err error - } - results := make(chan submitResult, count) - for i := 0; i < count; i++ { - go func(idx int) { - taskID, err := submitTask(ctx, prompt, refImages) - results <- submitResult{index: idx, taskID: taskID, err: err} - }(i) - } - - taskIDs := make([]string, count) - for i := 0; i < count; i++ { - r := <-results - if r.err != nil { - return nil, fmt.Errorf("submit task %d: %w", r.index, r.err) - } - taskIDs[r.index] = r.taskID - } - - type imageResult struct { - index int - image GeneratedImage - err error - } - imgResults := make(chan imageResult, count) - for i, tid := range taskIDs { - go func(idx int, taskID string) { - img, err := pollAndDownload(ctx, taskID) - imgResults <- imageResult{index: idx, image: img, err: err} - }(i, tid) - } - - images := make([]GeneratedImage, count) - for i := 0; i < count; i++ { - r := <-imgResults - if r.err != nil { - return nil, fmt.Errorf("task %d: %w", r.index, r.err) - } - images[r.index] = r.image - } - - return images, nil -} - -// submitTask 提交生图任务,返回 taskID。 -func submitTask(ctx context.Context, prompt string, refImages []string) (string, error) { - reqBody := genSubmitReq{ - Prompt: prompt, - AspectRatio: imgCfg.AspectRatio, - ReferenceImages: refImages, - XChannel: imgCfg.XChannel, +// 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 "", fmt.Errorf("marshal: %w", err) + return nil, fmt.Errorf("marshal request: %w", err) } - url := strings.TrimRight(imgCfg.BaseURL, "/") + "/gpt-image2/generate" - req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) - if err != nil { - return "", 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 "", fmt.Errorf("send: %w", err) - } - defer resp.Body.Close() - - var sr genSubmitResp - if err := json.NewDecoder(resp.Body).Decode(&sr); err != nil { - return "", fmt.Errorf("decode: %w", err) - } - - if sr.Code != 200 { - return "", fmt.Errorf("submit failed: %s", sr.Message) - } - - log.Printf("[inference] task submitted: %s", sr.Data.TaskID) - return sr.Data.TaskID, nil -} - -// pollAndDownload 轮询任务直到完成,下载图片。 -func pollAndDownload(ctx context.Context, taskID string) (GeneratedImage, error) { - pollInterval := imgCfg.PollInterval - if pollInterval <= 0 { - pollInterval = 3 - } - maxWait := imgCfg.PollMaxWait - if maxWait <= 0 { - maxWait = 120 - } - - deadline := time.Now().Add(time.Duration(maxWait) * time.Second) - ticker := time.NewTicker(time.Duration(pollInterval) * time.Second) - defer ticker.Stop() - - for { - select { - case <-ctx.Done(): - return GeneratedImage{}, ctx.Err() - case <-ticker.C: - st, err := queryStatus(ctx, taskID) - if err != nil { - return GeneratedImage{}, fmt.Errorf("query status: %w", err) - } - - switch st.Data.Status { - case "completed": - log.Printf("[inference] task %s completed, downloading", taskID) - data, err := downloadResult(ctx, st.Data.ResultImageURL) - if err != nil { - return GeneratedImage{}, fmt.Errorf("download: %w", err) - } - return GeneratedImage{Data: data, Format: "png"}, nil - - case "failed": - errMsg := st.Data.ErrorMessage - if errMsg == "" { - errMsg = "unknown error" - } - return GeneratedImage{}, fmt.Errorf("generation failed: %s", errMsg) - - default: - log.Printf("[inference] task %s status=%s progress=%d", taskID, st.Data.Status, st.Data.Progress) - } - - if time.Now().After(deadline) { - return GeneratedImage{}, fmt.Errorf("poll timeout after %ds", maxWait) - } - } - } -} - -// queryStatus 查询任务状态。 -func queryStatus(ctx context.Context, taskID string) (*genStatusResp, error) { - body, err := json.Marshal(genStatusReq{TaskID: taskID}) - if err != nil { - return nil, fmt.Errorf("marshal: %w", err) - } - - url := strings.TrimRight(imgCfg.BaseURL, "/") + "/gpt-image2/status" + 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) @@ -259,48 +98,125 @@ func queryStatus(ctx context.Context, taskID string) (*genStatusResp, error) { resp, err := http.DefaultClient.Do(req) if err != nil { - return nil, fmt.Errorf("send: %w", err) + return nil, fmt.Errorf("send request: %w", err) } defer resp.Body.Close() - var sr genStatusResp - if err := json.NewDecoder(resp.Body).Decode(&sr); err != nil { - return nil, fmt.Errorf("decode: %w", err) + 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)) } - if sr.Code != 200 { - return nil, fmt.Errorf("status query failed: %s", sr.Message) - } - - return &sr, nil + return parseImageResponse(ctx, resp.Body, width, height) } -// downloadResult 下载生成的图片。 -func downloadResult(ctx context.Context, imageURL string) ([]byte, error) { - req, err := http.NewRequestWithContext(ctx, http.MethodGet, imageURL, nil) +// 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: width, + Height: height, + Format: "png", + }) + } + return images, nil +} + +// 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) +} + // ======================== 质检 ======================== -// QualityChecker 质检函数,可替换用于测试。 var QualityChecker = defaultCheckQuality -// CheckQuality 调用当前 QualityChecker。 func CheckQuality(ctx context.Context, images []GeneratedImage, style map[string]string) (bool, string, error) { return QualityChecker(ctx, images, style) } From d8bfb81c4ddddf40b0ccf82c895ddd7712514cf0 Mon Sep 17 00:00:00 2001 From: Gmaker689 <1711322114@qq.com> Date: Mon, 25 May 2026 13:17:57 +0800 Subject: [PATCH 8/8] =?UTF-8?q?refactor(config):=20=E7=B2=BE=E7=AE=80=20im?= =?UTF-8?q?age=5Fgen=20=E9=85=8D=E7=BD=AE=E4=B8=BA=20OpenAI=20=E5=85=BC?= =?UTF-8?q?=E5=AE=B9=E5=90=8C=E6=AD=A5=20API?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 移除 yuntts 异步字段: aspect_ratio / x_channel / poll_max_wait / poll_interval - 改用标准 OpenAI Images API 字段: width / height / quality - 默认 base_url 改为 api.suchuang.vip/v1 - 模型默认值: gpt-image-2-token(此模型在 suchuang 上需换 gpt-image-1) --- backend/.env.example | 12 +++++------ backend/internal/config/config.go | 32 ++++++++++++------------------ backend/internal/config/config.yml | 10 ++++------ 3 files changed, 22 insertions(+), 32 deletions(-) diff --git a/backend/.env.example b/backend/.env.example index 6789aac..0a64dad 100755 --- a/backend/.env.example +++ b/backend/.env.example @@ -19,12 +19,10 @@ GEN2D_LLM_MODEL=gpt-4o GEN2D_LLM_TEMPERATURE=0.7 GEN2D_LLM_MAX_TOKENS=2048 -# 文生图模型(GPT Image 2 异步 API,如 yuntts 等兼容服务) -GEN2D_IMAGE_BASE_URL=https://www.yuntts.com/api/v1 +# 文生图模型(OpenAI 兼容 Images API) +GEN2D_IMAGE_BASE_URL=https://api.suchuang.vip/v1 GEN2D_IMAGE_API_KEY= -GEN2D_IMAGE_MODEL=gpt-image-2 +GEN2D_IMAGE_MODEL=gpt-image-2-token +GEN2D_IMAGE_WIDTH=1024 +GEN2D_IMAGE_HEIGHT=1024 GEN2D_IMAGE_QUALITY=low -GEN2D_IMAGE_ASPECT_RATIO=1:1 -GEN2D_IMAGE_X_CHANNEL=default -GEN2D_IMAGE_POLL_MAX_WAIT=120 -GEN2D_IMAGE_POLL_INTERVAL=3 diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index be4deb3..a15d38a 100755 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -43,16 +43,14 @@ type LLMConfig struct { MaxTokens int `mapstructure:"max_tokens"` } -// ImageGenConfig GPT Image 2 异步生图模型配置。 +// ImageGenConfig OpenAI 兼容文生图模型配置。 type ImageGenConfig struct { - BaseURL string `mapstructure:"base_url"` - APIKey string `mapstructure:"api_key"` - Model string `mapstructure:"model"` - Quality string `mapstructure:"quality"` - AspectRatio string `mapstructure:"aspect_ratio"` - XChannel string `mapstructure:"x_channel"` - PollMaxWait int `mapstructure:"poll_max_wait"` - PollInterval int `mapstructure:"poll_interval"` + 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"` } // Load 从 YAML 配置文件和环境变量加载配置。 @@ -103,14 +101,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://www.yuntts.com/api/v1") + v.SetDefault("image_gen.base_url", "https://api.suchuang.vip/v1") v.SetDefault("image_gen.api_key", "") - v.SetDefault("image_gen.model", "gpt-image-2") + 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.aspect_ratio", "1:1") - v.SetDefault("image_gen.x_channel", "default") - v.SetDefault("image_gen.poll_max_wait", 120) - v.SetDefault("image_gen.poll_interval", 3) } func bindEnvVars(v *viper.Viper) { @@ -130,9 +126,7 @@ func bindEnvVars(v *viper.Viper) { v.BindEnv("image_gen.base_url", "GEN2D_IMAGE_BASE_URL") v.BindEnv("image_gen.api_key", "GEN2D_IMAGE_API_KEY") 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.aspect_ratio", "GEN2D_IMAGE_ASPECT_RATIO") - v.BindEnv("image_gen.x_channel", "GEN2D_IMAGE_X_CHANNEL") - v.BindEnv("image_gen.poll_max_wait", "GEN2D_IMAGE_POLL_MAX_WAIT") - v.BindEnv("image_gen.poll_interval", "GEN2D_IMAGE_POLL_INTERVAL") } diff --git a/backend/internal/config/config.yml b/backend/internal/config/config.yml index 82d076b..aa29e74 100755 --- a/backend/internal/config/config.yml +++ b/backend/internal/config/config.yml @@ -21,11 +21,9 @@ llm: max_tokens: 2048 image_gen: - base_url: "https://www.yuntts.com/api/v1" + base_url: "https://api.suchuang.vip/v1" api_key: "" - model: "gpt-image-2" + model: "gpt-image-2-token" + width: 1024 + height: 1024 quality: "low" - aspect_ratio: "1:1" - x_channel: "default" - poll_max_wait: 120 - poll_interval: 3