Files
gen2d/backend/internal/service/inference.go
T
Gmarker689 6bda5a1f35 fix(inference): 修正文生图 API 调用参数,支持动态分辨率
- 去掉 response_format 参数(API 默认返回 URL,自动下载)
- Resolution 参数覆盖 API 调用尺寸(不再仅 mock 路径生效)
- parseImageResponse 改为动态 width/height 入参
- 新增 tools/gentest.go 生成测试脚本
- .gitignore 忽略 test_output/
2026-05-25 00:21:16 +08:00

306 lines
9.1 KiB
Go
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 生成图片,未配置 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
}
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)
}
size := params.Resolution
if size <= 0 {
size = 64
}
log.Println("[inference] ImageGen 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),
}
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
}