feat(inference): 实现文生图与图片编辑 API 调用,新增管线与编辑 HTTP 端点,JWT 认证中间件

- inference.go: GenerateImages 改为 API 优先(callImageGenAPI),无 key 回退 mock;新增 EditImages/callImageEditAPI(multipart 上传编辑);提取 parseImageResponse 共享响应解析
- handler/generate.go: POST /api/v1/generate 触发生成管线,返回 base64 图片
- handler/edit.go: POST /api/v1/images/edit 图片编辑端点
- mildware/auth.go: JWT Bearer token 认证中间件
- main.go: 路由拆分公开/认证组,generate 与 images/edit 需鉴权
This commit is contained in:
2026-05-24 22:47:43 +08:00
parent 5e54591b1f
commit 0a9a923d55
5 changed files with 412 additions and 21 deletions
+10 -1
View File
@@ -8,6 +8,7 @@ import (
"gen2d/internal/config"
"gen2d/internal/db"
"gen2d/internal/handler"
"gen2d/internal/mildware"
"gen2d/internal/model"
"gen2d/internal/service"
@@ -34,13 +35,21 @@ func main() {
r := gin.New()
r.Use(gin.Recovery()) // panic 恢复中间件,防止服务因未捕获异常宕机
// API v1 路由组
// API v1 路由组 — 公开接口
v1 := r.Group("/api/v1")
{
v1.GET("/health", handler.Health) // 健康检查
v1.POST("/prompt/optimize", handler.PromptOptimize) // 提示词优化
}
// API v1 路由组 — 需认证
v1Auth := r.Group("/api/v1")
v1Auth.Use(mildware.AuthMiddleware(cfg.JWT.Secret))
{
v1Auth.POST("/generate", handler.Generate) // 素材生成管线
v1Auth.POST("/images/edit", handler.EditImage) // 图片编辑
}
// Auth 路由组
auth := r.Group("/auth")
{
+59
View File
@@ -0,0 +1,59 @@
package handler
import (
"encoding/base64"
"net/http"
"gen2d/internal/model"
"gen2d/internal/service"
"github.com/gin-gonic/gin"
)
// EditImageRequest 图片编辑请求。
type EditImageRequest struct {
Image string `json:"image" binding:"required"` // 底图 base64 编码
Prompt string `json:"prompt" binding:"required"` // 编辑指令
Count int `json:"count"` // 生成数量,默认 1
}
// EditImageResponse 图片编辑响应体。
type EditImageResponse struct {
Assets []AssetResponse `json:"assets"`
}
// EditImage 图片编辑接口,基于已有图片和文本指令生成修改后的图片。
func EditImage(c *gin.Context) {
var req EditImageRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "参数错误: "+err.Error()))
return
}
imageData, err := base64.StdEncoding.DecodeString(req.Image)
if err != nil {
c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "图片 base64 解码失败: "+err.Error()))
return
}
count := req.Count
if count <= 0 {
count = 1
}
images, err := service.EditImages(c.Request.Context(), imageData, req.Prompt, count)
if err != nil {
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "图片编辑失败: "+err.Error()))
return
}
assets := make([]AssetResponse, len(images))
for i, img := range images {
assets[i] = AssetResponse{
Data: base64.StdEncoding.EncodeToString(img.Data),
Format: img.Format,
}
}
c.JSON(http.StatusOK, model.OK(EditImageResponse{Assets: assets}))
}
+84
View File
@@ -0,0 +1,84 @@
package handler
import (
"encoding/base64"
"net/http"
"gen2d/internal/model"
"gen2d/internal/service"
"github.com/gin-gonic/gin"
)
// GenerateRequest 素材生成请求。
type GenerateRequest struct {
Prompt string `json:"prompt"`
AssetType string `json:"assetType" binding:"required"`
Tags []string `json:"tags"`
UserNote string `json:"userNote"`
ProjectStyle map[string]string `json:"projectStyle"`
TaskStyle map[string]string `json:"taskStyle"`
Resolution int `json:"resolution"`
Directions int `json:"directions"`
FramesPerDir int `json:"framesPerDir"`
Format string `json:"format"`
}
// GenerateResponse 素材生成响应体。
type GenerateResponse struct {
Assets []AssetResponse `json:"assets"`
Metadata service.AssetMetadata `json:"metadata"`
}
// AssetResponse 单个素材响应(二进制 Data 转 base64)。
type AssetResponse struct {
Data string `json:"data"`
Format string `json:"format"`
URL string `json:"url"`
}
// Generate 素材生成接口,调用完整生成管线(PromptOptimizer → AssetGenerator → QualitySupervisor → FormatAdapter)。
func Generate(c *gin.Context) {
var req GenerateRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "参数错误: "+err.Error()))
return
}
in := service.PipelineInput{
Prompt: req.Prompt,
AssetType: req.AssetType,
Tags: req.Tags,
UserNote: req.UserNote,
ProjectStyle: req.ProjectStyle,
TaskStyle: req.TaskStyle,
Params: service.AssetParams{
Resolution: req.Resolution,
Frames: service.FrameParams{
Directions: req.Directions,
FramesPerDirection: req.FramesPerDir,
},
Format: req.Format,
},
}
output, err := service.RunPipeline(c.Request.Context(), in)
if err != nil {
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "素材生成失败: "+err.Error()))
return
}
assets := make([]AssetResponse, len(output.Assets))
for i, a := range output.Assets {
assets[i] = AssetResponse{
Data: base64.StdEncoding.EncodeToString(a.Data),
Format: a.Format,
URL: a.URL,
}
}
c.JSON(http.StatusOK, model.OK(GenerateResponse{
Assets: assets,
Metadata: output.Metadata,
}))
}
+48
View File
@@ -0,0 +1,48 @@
package mildware
import (
"net/http"
"strings"
"gen2d/internal/model"
"github.com/gin-gonic/gin"
"github.com/golang-jwt/jwt/v5"
)
// AuthMiddleware 返回 JWT 认证中间件,校验 Bearer token 并注入 userID 到上下文。
func AuthMiddleware(jwtSecret string) gin.HandlerFunc {
return func(c *gin.Context) {
authHeader := c.GetHeader("Authorization")
if authHeader == "" {
c.AbortWithStatusJSON(http.StatusUnauthorized, model.Fail(http.StatusUnauthorized, "未提供认证令牌"))
return
}
tokenString := strings.TrimPrefix(authHeader, "Bearer ")
if tokenString == authHeader {
c.AbortWithStatusJSON(http.StatusUnauthorized, model.Fail(http.StatusUnauthorized, "认证格式错误,需为 Bearer <token>"))
return
}
token, err := jwt.Parse(tokenString, func(token *jwt.Token) (interface{}, error) {
return []byte(jwtSecret), nil
})
if err != nil || !token.Valid {
c.AbortWithStatusJSON(http.StatusUnauthorized, model.Fail(http.StatusUnauthorized, "令牌无效或已过期"))
return
}
claims, ok := token.Claims.(jwt.MapClaims)
if !ok {
c.AbortWithStatusJSON(http.StatusUnauthorized, model.Fail(http.StatusUnauthorized, "令牌解析失败"))
return
}
if sub, ok := claims["sub"]; ok {
c.Set("userID", sub)
}
c.Next()
}
}
+211 -20
View File
@@ -4,10 +4,18 @@ 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"
)
@@ -20,37 +28,201 @@ func InitImageGenConfig(cfg config.ImageGenConfig) {
imgCfg = cfg
}
// GenerateImages 调用 AI 推理 API 生成图片。
// MVP 阶段返回 mock 占位图。
func GenerateImages(ctx context.Context, prompt string, params AssetParams) ([]GeneratedImage, error) {
size := params.Resolution
if size <= 0 {
size = 64
}
// ======================== 文生图 API 调用层 ========================
// imageGenRequest OpenAI 兼容的文生图请求体。
type imageGenRequest struct {
Model string `json:"model"`
Prompt string `json:"prompt"`
N int `json:"n,omitempty"`
Size string `json:"size,omitempty"`
ResponseFormat string `json:"response_format,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
}
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)
if imgCfg.APIKey != "" {
return callImageGenAPI(ctx, prompt, count)
}
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 int) ([]GeneratedImage, error) {
reqBody := imageGenRequest{
Model: imgCfg.Model,
Prompt: prompt,
N: count,
Size: fmt.Sprintf("%dx%d", imgCfg.Width, imgCfg.Height),
ResponseFormat: "b64_json",
}
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)
}
// ======================== 图片编辑 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))
writer.WriteField("response_format", "b64_json")
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)
}
// parseImageResponse 解析 OpenAI 兼容的图片生成/编辑响应体。
func parseImageResponse(ctx context.Context, r io.Reader) ([]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[i] = GeneratedImage{
images = append(images, GeneratedImage{
Data: data,
Width: size,
Height: size,
Width: imgCfg.Width,
Height: imgCfg.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 质检函数,可替换用于测试。
// 签名:(ctx, images, style) → (pass, reason, error)
var QualityChecker = defaultCheckQuality
// CheckQuality 调用当前 QualityChecker。
@@ -82,11 +254,30 @@ func AlwaysFailQualityChecker() func(context.Context, []GeneratedImage, map[stri
}
}
// generateMockImage 生成一张带随机色块的 PNG 占位图
// ======================== 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))
// 用 seed 生成不同颜色
r := uint8((seed*47 + 13) % 256)
g := uint8((seed*83 + 37) % 256)
b := uint8((seed*61 + 71) % 256)
@@ -104,7 +295,7 @@ func generateMockImage(size int, seed int) ([]byte, error) {
return buf.Bytes(), nil
}
// generateRandomBytes 用于生成随机数据(备用)
// generateRandomBytes 用于生成随机数据(备用)。
func generateRandomBytes(n int) ([]byte, error) {
b := make([]byte, n)
_, err := rand.Read(b)