fix: 为图片生成 API 添加超时与 5xx 重试机制

- ImageGenConfig 新增 timeout/max_retries/retry_delay 配置项(默认 120s/2次/5s)
- 替换 http.DefaultClient 为带超时的 imageHTTPClient
- callImageAPI 对 5xx 错误自动重试,避免 504 等瞬时故障直接失败
This commit is contained in:
2026-05-25 14:05:14 +08:00
parent 83d3d2663d
commit 50f95adb98
3 changed files with 71 additions and 21 deletions
+3
View File
@@ -26,6 +26,9 @@ GEN2D_IMAGE_MODEL=gpt-image-2-token
GEN2D_IMAGE_WIDTH=1024 GEN2D_IMAGE_WIDTH=1024
GEN2D_IMAGE_HEIGHT=1024 GEN2D_IMAGE_HEIGHT=1024
GEN2D_IMAGE_QUALITY=low GEN2D_IMAGE_QUALITY=low
GEN2D_IMAGE_TIMEOUT=120
GEN2D_IMAGE_MAX_RETRIES=2
GEN2D_IMAGE_RETRY_DELAY=5
GEN2D_IMAGE_NUM_IMAGES=1 GEN2D_IMAGE_NUM_IMAGES=1
GEN2D_IMAGE_STEPS=30 GEN2D_IMAGE_STEPS=30
GEN2D_IMAGE_CFG_SCALE=7.0 GEN2D_IMAGE_CFG_SCALE=7.0
+9
View File
@@ -52,6 +52,9 @@ type ImageGenConfig struct {
Width int `mapstructure:"width"` Width int `mapstructure:"width"`
Height int `mapstructure:"height"` Height int `mapstructure:"height"`
Quality string `mapstructure:"quality"` Quality string `mapstructure:"quality"`
Timeout int `mapstructure:"timeout"` // HTTP 请求超时秒数,默认 120
MaxRetries int `mapstructure:"max_retries"` // 5xx 错误重试次数,默认 2
RetryDelay int `mapstructure:"retry_delay"` // 重试间隔秒数,默认 5
} }
// QiniuConfig 七牛云对象存储配置。 // QiniuConfig 七牛云对象存储配置。
@@ -117,6 +120,9 @@ func setDefaults(v *viper.Viper) {
v.SetDefault("image_gen.width", 1024) v.SetDefault("image_gen.width", 1024)
v.SetDefault("image_gen.height", 1024) v.SetDefault("image_gen.height", 1024)
v.SetDefault("image_gen.quality", "low") v.SetDefault("image_gen.quality", "low")
v.SetDefault("image_gen.timeout", 120)
v.SetDefault("image_gen.max_retries", 2)
v.SetDefault("image_gen.retry_delay", 5)
v.SetDefault("image_gen.num_images", 1) v.SetDefault("image_gen.num_images", 1)
v.SetDefault("image_gen.steps", 30) v.SetDefault("image_gen.steps", 30)
v.SetDefault("image_gen.cfg_scale", 7.0) v.SetDefault("image_gen.cfg_scale", 7.0)
@@ -148,6 +154,9 @@ func bindEnvVars(v *viper.Viper) {
v.BindEnv("image_gen.width", "GEN2D_IMAGE_WIDTH") v.BindEnv("image_gen.width", "GEN2D_IMAGE_WIDTH")
v.BindEnv("image_gen.height", "GEN2D_IMAGE_HEIGHT") v.BindEnv("image_gen.height", "GEN2D_IMAGE_HEIGHT")
v.BindEnv("image_gen.quality", "GEN2D_IMAGE_QUALITY") v.BindEnv("image_gen.quality", "GEN2D_IMAGE_QUALITY")
v.BindEnv("image_gen.timeout", "GEN2D_IMAGE_TIMEOUT")
v.BindEnv("image_gen.max_retries", "GEN2D_IMAGE_MAX_RETRIES")
v.BindEnv("image_gen.retry_delay", "GEN2D_IMAGE_RETRY_DELAY")
v.BindEnv("image_gen.num_images", "GEN2D_IMAGE_NUM_IMAGES") v.BindEnv("image_gen.num_images", "GEN2D_IMAGE_NUM_IMAGES")
v.BindEnv("image_gen.steps", "GEN2D_IMAGE_STEPS") v.BindEnv("image_gen.steps", "GEN2D_IMAGE_STEPS")
v.BindEnv("image_gen.cfg_scale", "GEN2D_IMAGE_CFG_SCALE") v.BindEnv("image_gen.cfg_scale", "GEN2D_IMAGE_CFG_SCALE")
+46 -8
View File
@@ -15,6 +15,7 @@ import (
"mime/multipart" "mime/multipart"
"net/http" "net/http"
"strings" "strings"
"time"
"gen2d/internal/config" "gen2d/internal/config"
) )
@@ -22,9 +23,17 @@ import (
// imgCfg 保存文生图配置,由 main 通过 InitImageGenConfig 注入。 // imgCfg 保存文生图配置,由 main 通过 InitImageGenConfig 注入。
var imgCfg config.ImageGenConfig var imgCfg config.ImageGenConfig
// imageHTTPClient 带超时的 HTTP 客户端,由 InitImageGenConfig 初始化。
var imageHTTPClient *http.Client
// InitImageGenConfig 注入文生图配置。 // InitImageGenConfig 注入文生图配置。
func InitImageGenConfig(cfg config.ImageGenConfig) { func InitImageGenConfig(cfg config.ImageGenConfig) {
imgCfg = cfg imgCfg = cfg
timeout := time.Duration(cfg.Timeout) * time.Second
if timeout <= 0 {
timeout = 120 * time.Second
}
imageHTTPClient = &http.Client{Timeout: timeout}
} }
// ======================== OpenAI 兼容 Images API 类型 ======================== // ======================== OpenAI 兼容 Images API 类型 ========================
@@ -73,7 +82,7 @@ func GenerateImages(ctx context.Context, prompt string, params AssetParams) ([]G
return generateMockImages(size, count) return generateMockImages(size, count)
} }
// callImageAPI 调用 OpenAI 兼容 Images API。 // callImageAPI 调用 OpenAI 兼容 Images API,5xx 错误自动重试。
func callImageAPI(ctx context.Context, prompt string, count, width, height int) ([]GeneratedImage, error) { func callImageAPI(ctx context.Context, prompt string, count, width, height int) ([]GeneratedImage, error) {
reqBody := imageGenRequest{ reqBody := imageGenRequest{
Model: imgCfg.Model, Model: imgCfg.Model,
@@ -89,6 +98,26 @@ func callImageAPI(ctx context.Context, prompt string, count, width, height int)
} }
url := strings.TrimRight(imgCfg.BaseURL, "/") + "/images/generations" url := strings.TrimRight(imgCfg.BaseURL, "/") + "/images/generations"
maxRetries := imgCfg.MaxRetries
if maxRetries <= 0 {
maxRetries = 2
}
retryDelay := time.Duration(imgCfg.RetryDelay) * time.Second
if retryDelay <= 0 {
retryDelay = 5 * time.Second
}
var lastErr error
for attempt := 0; attempt <= maxRetries; attempt++ {
if attempt > 0 {
log.Printf("[inference] retrying image API (attempt %d/%d)", attempt, maxRetries)
select {
case <-ctx.Done():
return nil, fmt.Errorf("context cancelled during retry: %w", ctx.Err())
case <-time.After(retryDelay):
}
}
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)
@@ -96,18 +125,27 @@ func callImageAPI(ctx context.Context, prompt string, count, width, height int)
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+imgCfg.APIKey) req.Header.Set("Authorization", "Bearer "+imgCfg.APIKey)
resp, err := http.DefaultClient.Do(req) resp, err := imageHTTPClient.Do(req)
if err != nil { if err != nil {
return nil, fmt.Errorf("send request: %w", err) lastErr = fmt.Errorf("send request: %w", err)
continue
}
if resp.StatusCode == http.StatusOK {
return parseImageResponse(ctx, resp.Body, width, height)
} }
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
b, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) b, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
resp.Body.Close()
if resp.StatusCode >= 500 {
lastErr = fmt.Errorf("image api error %d: %s", resp.StatusCode, string(b))
continue
}
return nil, fmt.Errorf("image api error %d: %s", resp.StatusCode, string(b)) return nil, fmt.Errorf("image api error %d: %s", resp.StatusCode, string(b))
} }
return parseImageResponse(ctx, resp.Body, width, height) return nil, lastErr
} }
// parseImageResponse 解析 OpenAI 兼容图片响应(b64_json 或 url)。 // parseImageResponse 解析 OpenAI 兼容图片响应(b64_json 或 url)。
@@ -152,7 +190,7 @@ func downloadImage(ctx context.Context, url string) ([]byte, error) {
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 := imageHTTPClient.Do(req)
if err != nil { if err != nil {
return nil, fmt.Errorf("download: %w", err) return nil, fmt.Errorf("download: %w", err)
} }
@@ -199,7 +237,7 @@ func EditImages(ctx context.Context, imageData []byte, prompt string, count int)
req.Header.Set("Content-Type", writer.FormDataContentType()) req.Header.Set("Content-Type", writer.FormDataContentType())
req.Header.Set("Authorization", "Bearer "+imgCfg.APIKey) req.Header.Set("Authorization", "Bearer "+imgCfg.APIKey)
resp, err := http.DefaultClient.Do(req) resp, err := imageHTTPClient.Do(req)
if err != nil { if err != nil {
return nil, fmt.Errorf("send request: %w", err) return nil, fmt.Errorf("send request: %w", err)
} }