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:
2026-05-25 12:44:39 +08:00
parent d9d3b51262
commit c4dc7394b1
6 changed files with 255 additions and 512 deletions
+213 -167
View File
@@ -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)