refactor(inference): 统合 image_gen 为 GPT Image 2 异步 API,删除冗余配置
- 删除 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_* 冗余
This commit is contained in:
+9
-18
@@ -19,21 +19,12 @@ GEN2D_LLM_MODEL=gpt-4o
|
|||||||
GEN2D_LLM_TEMPERATURE=0.7
|
GEN2D_LLM_TEMPERATURE=0.7
|
||||||
GEN2D_LLM_MAX_TOKENS=2048
|
GEN2D_LLM_MAX_TOKENS=2048
|
||||||
|
|
||||||
# 文生图模型(OpenAI 兼容同步 API)
|
# 文生图模型(GPT Image 2 异步 API,如 yuntts 等兼容服务)
|
||||||
GEN2D_IMAGE_BASE_URL=https://api.stability.ai/v1
|
GEN2D_IMAGE_BASE_URL=https://www.yuntts.com/api/v1
|
||||||
GEN2D_IMAGE_API_KEY=sk-your-api-key
|
GEN2D_IMAGE_API_KEY=
|
||||||
GEN2D_IMAGE_MODEL=stable-diffusion-xl
|
GEN2D_IMAGE_MODEL=gpt-image-2
|
||||||
GEN2D_IMAGE_WIDTH=1024
|
GEN2D_IMAGE_QUALITY=low
|
||||||
GEN2D_IMAGE_HEIGHT=1024
|
GEN2D_IMAGE_ASPECT_RATIO=1:1
|
||||||
GEN2D_IMAGE_NUM_IMAGES=1
|
GEN2D_IMAGE_X_CHANNEL=default
|
||||||
GEN2D_IMAGE_STEPS=30
|
GEN2D_IMAGE_POLL_MAX_WAIT=120
|
||||||
GEN2D_IMAGE_CFG_SCALE=7.0
|
GEN2D_IMAGE_POLL_INTERVAL=3
|
||||||
|
|
||||||
# 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
|
|
||||||
|
|||||||
@@ -35,7 +35,6 @@ func main() {
|
|||||||
// 注入 LLM 和文生图配置到 service 层
|
// 注入 LLM 和文生图配置到 service 层
|
||||||
service.InitLLMConfig(cfg.LLM)
|
service.InitLLMConfig(cfg.LLM)
|
||||||
service.InitImageGenConfig(cfg.ImageGen)
|
service.InitImageGenConfig(cfg.ImageGen)
|
||||||
service.InitGptImage2Config(cfg.GptImage2)
|
|
||||||
|
|
||||||
r := gin.New()
|
r := gin.New()
|
||||||
r.Use(gin.Recovery()) // panic 恢复中间件,防止服务因未捕获异常宕机
|
r.Use(gin.Recovery()) // panic 恢复中间件,防止服务因未捕获异常宕机
|
||||||
|
|||||||
@@ -14,7 +14,6 @@ type Config struct {
|
|||||||
JWT JWTConfig `mapstructure:"jwt"`
|
JWT JWTConfig `mapstructure:"jwt"`
|
||||||
LLM LLMConfig `mapstructure:"llm"`
|
LLM LLMConfig `mapstructure:"llm"`
|
||||||
ImageGen ImageGenConfig `mapstructure:"image_gen"`
|
ImageGen ImageGenConfig `mapstructure:"image_gen"`
|
||||||
GptImage2 GptImage2Config `mapstructure:"gpt_image2"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// ServerConfig HTTP 服务配置。
|
// ServerConfig HTTP 服务配置。
|
||||||
@@ -44,26 +43,16 @@ type LLMConfig struct {
|
|||||||
MaxTokens int `mapstructure:"max_tokens"`
|
MaxTokens int `mapstructure:"max_tokens"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// ImageGenConfig 文生图模型配置(OpenAI 兼容同步 API)。
|
// ImageGenConfig GPT Image 2 异步生图模型配置。
|
||||||
type ImageGenConfig struct {
|
type ImageGenConfig struct {
|
||||||
BaseURL string `mapstructure:"base_url"`
|
BaseURL string `mapstructure:"base_url"`
|
||||||
APIKey string `mapstructure:"api_key"`
|
APIKey string `mapstructure:"api_key"`
|
||||||
Model string `mapstructure:"model"`
|
Model string `mapstructure:"model"`
|
||||||
Width int `mapstructure:"width"`
|
Quality string `mapstructure:"quality"`
|
||||||
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"`
|
AspectRatio string `mapstructure:"aspect_ratio"`
|
||||||
XChannel string `mapstructure:"x_channel"`
|
XChannel string `mapstructure:"x_channel"`
|
||||||
PollMaxWait int `mapstructure:"poll_max_wait"` // 轮询最大等待秒数, 默认 120
|
PollMaxWait int `mapstructure:"poll_max_wait"`
|
||||||
PollInterval int `mapstructure:"poll_interval"` // 轮询间隔秒数, 默认 3
|
PollInterval int `mapstructure:"poll_interval"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Load 从 YAML 配置文件和环境变量加载配置。
|
// Load 从 YAML 配置文件和环境变量加载配置。
|
||||||
@@ -114,21 +103,14 @@ func setDefaults(v *viper.Viper) {
|
|||||||
v.SetDefault("llm.temperature", 0.7)
|
v.SetDefault("llm.temperature", 0.7)
|
||||||
v.SetDefault("llm.max_tokens", 2048)
|
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.api_key", "")
|
||||||
v.SetDefault("image_gen.model", "stable-diffusion-xl")
|
v.SetDefault("image_gen.model", "gpt-image-2")
|
||||||
v.SetDefault("image_gen.width", 1024)
|
v.SetDefault("image_gen.quality", "low")
|
||||||
v.SetDefault("image_gen.height", 1024)
|
v.SetDefault("image_gen.aspect_ratio", "1:1")
|
||||||
v.SetDefault("image_gen.num_images", 1)
|
v.SetDefault("image_gen.x_channel", "default")
|
||||||
v.SetDefault("image_gen.steps", 30)
|
v.SetDefault("image_gen.poll_max_wait", 120)
|
||||||
v.SetDefault("image_gen.cfg_scale", 7.0)
|
v.SetDefault("image_gen.poll_interval", 3)
|
||||||
|
|
||||||
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) {
|
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.base_url", "GEN2D_IMAGE_BASE_URL")
|
||||||
v.BindEnv("image_gen.api_key", "GEN2D_IMAGE_API_KEY")
|
v.BindEnv("image_gen.api_key", "GEN2D_IMAGE_API_KEY")
|
||||||
v.BindEnv("image_gen.model", "GEN2D_IMAGE_MODEL")
|
v.BindEnv("image_gen.model", "GEN2D_IMAGE_MODEL")
|
||||||
v.BindEnv("image_gen.width", "GEN2D_IMAGE_WIDTH")
|
v.BindEnv("image_gen.quality", "GEN2D_IMAGE_QUALITY")
|
||||||
v.BindEnv("image_gen.height", "GEN2D_IMAGE_HEIGHT")
|
v.BindEnv("image_gen.aspect_ratio", "GEN2D_IMAGE_ASPECT_RATIO")
|
||||||
v.BindEnv("image_gen.num_images", "GEN2D_IMAGE_NUM_IMAGES")
|
v.BindEnv("image_gen.x_channel", "GEN2D_IMAGE_X_CHANNEL")
|
||||||
v.BindEnv("image_gen.steps", "GEN2D_IMAGE_STEPS")
|
v.BindEnv("image_gen.poll_max_wait", "GEN2D_IMAGE_POLL_MAX_WAIT")
|
||||||
v.BindEnv("image_gen.cfg_scale", "GEN2D_IMAGE_CFG_SCALE")
|
v.BindEnv("image_gen.poll_interval", "GEN2D_IMAGE_POLL_INTERVAL")
|
||||||
|
|
||||||
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")
|
|
||||||
}
|
}
|
||||||
|
|||||||
Regular → Executable
+7
-7
@@ -21,11 +21,11 @@ llm:
|
|||||||
max_tokens: 2048
|
max_tokens: 2048
|
||||||
|
|
||||||
image_gen:
|
image_gen:
|
||||||
base_url: "https://api.stability.ai/v1"
|
base_url: "https://www.yuntts.com/api/v1"
|
||||||
api_key: ""
|
api_key: ""
|
||||||
model: "stable-diffusion-xl"
|
model: "gpt-image-2"
|
||||||
width: 1024
|
quality: "low"
|
||||||
height: 1024
|
aspect_ratio: "1:1"
|
||||||
num_images: 1
|
x_channel: "default"
|
||||||
steps: 30
|
poll_max_wait: 120
|
||||||
cfg_scale: 7.0
|
poll_interval: 3
|
||||||
|
|||||||
@@ -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)
|
|
||||||
}
|
|
||||||
@@ -4,7 +4,6 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"encoding/base64"
|
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"image"
|
"image"
|
||||||
@@ -12,10 +11,9 @@ import (
|
|||||||
"image/png"
|
"image/png"
|
||||||
"io"
|
"io"
|
||||||
"log"
|
"log"
|
||||||
"mime/multipart"
|
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
|
||||||
"strings"
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
"gen2d/internal/config"
|
"gen2d/internal/config"
|
||||||
)
|
)
|
||||||
@@ -28,84 +26,230 @@ func InitImageGenConfig(cfg config.ImageGenConfig) {
|
|||||||
imgCfg = cfg
|
imgCfg = cfg
|
||||||
}
|
}
|
||||||
|
|
||||||
// ======================== 文生图 API 调用层 ========================
|
// ======================== GPT Image 2 API 类型 ========================
|
||||||
|
|
||||||
// imageGenRequest OpenAI 兼容的文生图请求体。
|
// genSubmitReq 提交生图任务请求体。
|
||||||
type imageGenRequest struct {
|
type genSubmitReq struct {
|
||||||
Model string `json:"model"`
|
|
||||||
Prompt string `json:"prompt"`
|
Prompt string `json:"prompt"`
|
||||||
N int `json:"n,omitempty"`
|
AspectRatio string `json:"aspect_ratio,omitempty"`
|
||||||
Size string `json:"size,omitempty"`
|
ReferenceImages []string `json:"reference_images,omitempty"`
|
||||||
Steps int `json:"steps,omitempty"`
|
XChannel string `json:"x_channel,omitempty"`
|
||||||
CFGScale float64 `json:"cfg_scale,omitempty"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// imageGenResponse OpenAI 兼容的文生图响应体。
|
// genSubmitResp 提交生图任务响应体。
|
||||||
type imageGenResponse struct {
|
type genSubmitResp struct {
|
||||||
Data []struct {
|
Code int `json:"code"`
|
||||||
URL string `json:"url"`
|
Message string `json:"message"`
|
||||||
B64JSON string `json:"b64_json"`
|
Data struct {
|
||||||
|
TaskID string `json:"task_id"`
|
||||||
|
Status string `json:"status"`
|
||||||
} `json:"data"`
|
} `json:"data"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// GenerateImages 调用 AI 推理 API 生成图片。
|
// genStatusReq 查询任务状态请求体。
|
||||||
// 优先级:GPT Image 2 > OpenAI 兼容 ImageGen > Mock 回退。
|
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) {
|
func GenerateImages(ctx context.Context, prompt string, params AssetParams) ([]GeneratedImage, error) {
|
||||||
count := 1
|
count := 1
|
||||||
if params.Frames.Directions > 0 && params.Frames.FramesPerDirection > 0 {
|
if params.Frames.Directions > 0 && params.Frames.FramesPerDirection > 0 {
|
||||||
count = params.Frames.Directions * params.Frames.FramesPerDirection
|
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 != "" {
|
if imgCfg.APIKey != "" {
|
||||||
width, height := imgCfg.Width, imgCfg.Height
|
log.Println("[inference] using image gen async API")
|
||||||
if params.Resolution > 0 {
|
return generateAsync(ctx, prompt, count, nil)
|
||||||
width = params.Resolution
|
|
||||||
height = params.Resolution
|
|
||||||
}
|
|
||||||
return callImageGenAPI(ctx, prompt, count, width, height)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// 最终降级:mock 占位图
|
|
||||||
size := params.Resolution
|
size := params.Resolution
|
||||||
if size <= 0 {
|
if size <= 0 {
|
||||||
size = 64
|
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)
|
return generateMockImages(size, count)
|
||||||
}
|
}
|
||||||
|
|
||||||
// callImageGenAPI 调用 OpenAI 兼容的 Images API,返回生成的图片。
|
// ======================== 图片编辑 ========================
|
||||||
func callImageGenAPI(ctx context.Context, prompt string, count, width, height int) ([]GeneratedImage, error) {
|
|
||||||
reqBody := imageGenRequest{
|
// EditImages 图片编辑接口,以参考图模式提交 GPT Image 2 编辑任务。
|
||||||
Model: imgCfg.Model,
|
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,
|
Prompt: prompt,
|
||||||
N: count,
|
AspectRatio: imgCfg.AspectRatio,
|
||||||
Size: fmt.Sprintf("%dx%d", width, height),
|
ReferenceImages: refImages,
|
||||||
}
|
XChannel: imgCfg.XChannel,
|
||||||
if imgCfg.Steps > 0 {
|
|
||||||
reqBody.Steps = imgCfg.Steps
|
|
||||||
}
|
|
||||||
if imgCfg.CFGScale > 0 {
|
|
||||||
reqBody.CFGScale = imgCfg.CFGScale
|
|
||||||
}
|
}
|
||||||
|
|
||||||
body, err := json.Marshal(reqBody)
|
body, err := json.Marshal(reqBody)
|
||||||
if err != nil {
|
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))
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("create request: %w", err)
|
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)
|
resp, err := http.DefaultClient.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("send request: %w", err)
|
return nil, fmt.Errorf("send: %w", err)
|
||||||
}
|
}
|
||||||
defer resp.Body.Close()
|
defer resp.Body.Close()
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
var sr genStatusResp
|
||||||
b, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
|
if err := json.NewDecoder(resp.Body).Decode(&sr); err != nil {
|
||||||
return nil, fmt.Errorf("image gen api error %d: %s", resp.StatusCode, string(b))
|
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)
|
||||||
}
|
}
|
||||||
|
|
||||||
// ======================== 图片编辑 API ========================
|
return &sr, nil
|
||||||
|
|
||||||
// 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)。
|
// downloadResult 下载生成的图片。
|
||||||
func callImageEditAPI(ctx context.Context, imageData []byte, prompt string, count int) ([]GeneratedImage, error) {
|
func downloadResult(ctx context.Context, imageURL string) ([]byte, error) {
|
||||||
var buf bytes.Buffer
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, imageURL, nil)
|
||||||
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)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("create download request: %w", err)
|
return nil, fmt.Errorf("create download request: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
resp, err := http.DefaultClient.Do(req)
|
resp, err := http.DefaultClient.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("download: %w", err)
|
return nil, fmt.Errorf("download: %w", err)
|
||||||
}
|
}
|
||||||
defer resp.Body.Close()
|
defer resp.Body.Close()
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
if resp.StatusCode != http.StatusOK {
|
||||||
return nil, fmt.Errorf("download status %d", resp.StatusCode)
|
return nil, fmt.Errorf("download status %d", resp.StatusCode)
|
||||||
}
|
}
|
||||||
|
|
||||||
return io.ReadAll(resp.Body)
|
return io.ReadAll(resp.Body)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -245,12 +305,10 @@ func CheckQuality(ctx context.Context, images []GeneratedImage, style map[string
|
|||||||
return QualityChecker(ctx, images, style)
|
return QualityChecker(ctx, images, style)
|
||||||
}
|
}
|
||||||
|
|
||||||
// defaultCheckQuality 默认 mock 质检,始终返回 pass。
|
|
||||||
func defaultCheckQuality(ctx context.Context, images []GeneratedImage, style map[string]string) (bool, string, error) {
|
func defaultCheckQuality(ctx context.Context, images []GeneratedImage, style map[string]string) (bool, string, error) {
|
||||||
return true, "", nil
|
return true, "", nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewCountedQualityChecker 创建一个在第 passOnRetry 次调用时返回 pass 的质检函数。
|
|
||||||
func NewCountedQualityChecker(passOnRetry int) func(context.Context, []GeneratedImage, map[string]string) (bool, string, error) {
|
func NewCountedQualityChecker(passOnRetry int) func(context.Context, []GeneratedImage, map[string]string) (bool, string, error) {
|
||||||
var callCount int
|
var callCount int
|
||||||
return func(_ context.Context, _ []GeneratedImage, _ map[string]string) (bool, string, error) {
|
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 {
|
if callCount >= passOnRetry {
|
||||||
return true, "", nil
|
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) {
|
func AlwaysFailQualityChecker() func(context.Context, []GeneratedImage, map[string]string) (bool, string, error) {
|
||||||
return 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 回退 ========================
|
// ======================== Mock 回退 ========================
|
||||||
|
|
||||||
// generateMockImages 批量生成 mock PNG 占位图。
|
|
||||||
func generateMockImages(size, count int) ([]GeneratedImage, error) {
|
func generateMockImages(size, count int) ([]GeneratedImage, error) {
|
||||||
images := make([]GeneratedImage, count)
|
images := make([]GeneratedImage, count)
|
||||||
for i := 0; i < count; i++ {
|
for i := 0; i < count; i++ {
|
||||||
@@ -279,30 +335,21 @@ func generateMockImages(size, count int) ([]GeneratedImage, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("generate mock image %d: %w", i, err)
|
return nil, fmt.Errorf("generate mock image %d: %w", i, err)
|
||||||
}
|
}
|
||||||
images[i] = GeneratedImage{
|
images[i] = GeneratedImage{Data: data, Width: size, Height: size, Format: "png"}
|
||||||
Data: data,
|
|
||||||
Width: size,
|
|
||||||
Height: size,
|
|
||||||
Format: "png",
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return images, nil
|
return images, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// generateMockImage 生成一张带随机色块的 PNG 占位图。
|
|
||||||
func generateMockImage(size int, seed int) ([]byte, error) {
|
func generateMockImage(size int, seed int) ([]byte, error) {
|
||||||
img := image.NewRGBA(image.Rect(0, 0, size, size))
|
img := image.NewRGBA(image.Rect(0, 0, size, size))
|
||||||
|
|
||||||
r := uint8((seed*47 + 13) % 256)
|
r := uint8((seed*47 + 13) % 256)
|
||||||
g := uint8((seed*83 + 37) % 256)
|
g := uint8((seed*83 + 37) % 256)
|
||||||
b := uint8((seed*61 + 71) % 256)
|
b := uint8((seed*61 + 71) % 256)
|
||||||
|
|
||||||
for y := 0; y < size; y++ {
|
for y := 0; y < size; y++ {
|
||||||
for x := 0; x < size; x++ {
|
for x := 0; x < size; x++ {
|
||||||
img.Set(x, y, color.RGBA{R: r, G: g, B: b, A: 255})
|
img.Set(x, y, color.RGBA{R: r, G: g, B: b, A: 255})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var buf bytes.Buffer
|
var buf bytes.Buffer
|
||||||
if err := png.Encode(&buf, img); err != nil {
|
if err := png.Encode(&buf, img); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -310,7 +357,6 @@ func generateMockImage(size int, seed int) ([]byte, error) {
|
|||||||
return buf.Bytes(), nil
|
return buf.Bytes(), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// generateRandomBytes 用于生成随机数据(备用)。
|
|
||||||
func generateRandomBytes(n int) ([]byte, error) {
|
func generateRandomBytes(n int) ([]byte, error) {
|
||||||
b := make([]byte, n)
|
b := make([]byte, n)
|
||||||
_, err := rand.Read(b)
|
_, err := rand.Read(b)
|
||||||
|
|||||||
Reference in New Issue
Block a user