Files
gen2d/backend/internal/service/inference.go
T
Gmarker689 d9d3b51262 feat(inference): 集成 GPT Image 2 异步文生图 API
- 新增 GptImage2Config 配置(yuntts 等 GPT Image 2 兼容服务)
- 新增 gptimage.go 异步客户端:提交任务 → 轮询状态 → 下载图片
- GenerateImages 优先级调整为:GPT Image 2 > OpenAI 兼容 ImageGen > Mock
- 支持环境变量 GEN2D_GPT_IMAGE2_* 系列配置
2026-05-25 12:36:37 +08:00

319 lines
9.5 KiB
Go
Executable File
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package service
import (
"bytes"
"context"
"crypto/rand"
"encoding/base64"
"encoding/json"
"fmt"
"image"
"image/color"
"image/png"
"io"
"log"
"mime/multipart"
"net/http"
"strconv"
"strings"
"gen2d/internal/config"
)
// imgCfg 保存文生图配置,由 main 通过 InitImageGenConfig 注入。
var imgCfg config.ImageGenConfig
// InitImageGenConfig 注入文生图配置。
func InitImageGenConfig(cfg config.ImageGenConfig) {
imgCfg = cfg
}
// ======================== 文生图 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"`
}
// imageGenResponse OpenAI 兼容的文生图响应体。
type imageGenResponse struct {
Data []struct {
URL string `json:"url"`
B64JSON string `json:"b64_json"`
} `json:"data"`
}
// GenerateImages 调用 AI 推理 API 生成图片。
// 优先级:GPT Image 2 > OpenAI 兼容 ImageGen > Mock 回退。
func GenerateImages(ctx context.Context, prompt string, params AssetParams) ([]GeneratedImage, error) {
count := 1
if params.Frames.Directions > 0 && params.Frames.FramesPerDirection > 0 {
count = params.Frames.Directions * params.Frames.FramesPerDirection
}
// 优先使用 GPT Image 2 异步 API
if gptImage2Cfg.APIKey != "" {
log.Println("[inference] using GPT Image 2 async API")
images, err := GenerateImagesGpt2(ctx, prompt, count)
if err != nil {
return nil, fmt.Errorf("gpt_image2: %w", err)
}
return images, nil
}
// 降级:OpenAI 兼容 Images API
if imgCfg.APIKey != "" {
width, height := imgCfg.Width, imgCfg.Height
if params.Resolution > 0 {
width = params.Resolution
height = params.Resolution
}
return callImageGenAPI(ctx, prompt, count, width, height)
}
// 最终降级:mock 占位图
size := params.Resolution
if size <= 0 {
size = 64
}
log.Println("[inference] no image API key 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),
}
if imgCfg.Steps > 0 {
reqBody.Steps = imgCfg.Steps
}
if imgCfg.CFGScale > 0 {
reqBody.CFGScale = imgCfg.CFGScale
}
body, err := json.Marshal(reqBody)
if err != nil {
return nil, fmt.Errorf("marshal request: %w", err)
}
url := strings.TrimRight(imgCfg.BaseURL, "/")
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("create request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+imgCfg.APIKey)
resp, err := http.DefaultClient.Do(req)
if err != nil {
return nil, fmt.Errorf("send request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
b, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
return nil, fmt.Errorf("image gen api error %d: %s", resp.StatusCode, string(b))
}
return parseImageResponse(ctx, resp.Body, width, height)
}
// ======================== 图片编辑 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)
if err != nil {
return nil, fmt.Errorf("create download request: %w", err)
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return nil, fmt.Errorf("download: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("download status %d", resp.StatusCode)
}
return io.ReadAll(resp.Body)
}
// ======================== 质检 ========================
// QualityChecker 质检函数,可替换用于测试。
var QualityChecker = defaultCheckQuality
// CheckQuality 调用当前 QualityChecker。
func CheckQuality(ctx context.Context, images []GeneratedImage, style map[string]string) (bool, string, error) {
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) {
callCount++
if callCount >= passOnRetry {
return true, "", nil
}
return false, fmt.Sprintf("风格不一致(第 %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
}
}
// ======================== Mock 回退 ========================
// generateMockImages 批量生成 mock PNG 占位图。
func generateMockImages(size, count int) ([]GeneratedImage, error) {
images := make([]GeneratedImage, count)
for i := 0; i < count; i++ {
data, err := generateMockImage(size, i)
if err != nil {
return nil, fmt.Errorf("generate mock image %d: %w", i, err)
}
images[i] = GeneratedImage{
Data: data,
Width: size,
Height: size,
Format: "png",
}
}
return images, nil
}
// generateMockImage 生成一张带随机色块的 PNG 占位图。
func generateMockImage(size int, seed int) ([]byte, error) {
img := image.NewRGBA(image.Rect(0, 0, size, size))
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
}
return buf.Bytes(), nil
}
// generateRandomBytes 用于生成随机数据(备用)。
func generateRandomBytes(n int) ([]byte, error) {
b := make([]byte, n)
_, err := rand.Read(b)
return b, err
}