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] =?UTF-8?q?refactor(inference):=20=E7=BB=9F=E5=90=88=20ima?= =?UTF-8?q?ge=5Fgen=20=E4=B8=BA=20GPT=20Image=202=20=E5=BC=82=E6=AD=A5=20A?= =?UTF-8?q?PI=EF=BC=8C=E5=88=A0=E9=99=A4=E5=86=97=E4=BD=99=E9=85=8D?= =?UTF-8?q?=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)