From 0a9a923d55cfc66394ecd475cd520f3a768b95dd Mon Sep 17 00:00:00 2001 From: Gmaker689 <1711322114@qq.com> Date: Sun, 24 May 2026 22:47:43 +0800 Subject: [PATCH] =?UTF-8?q?feat(inference):=20=E5=AE=9E=E7=8E=B0=E6=96=87?= =?UTF-8?q?=E7=94=9F=E5=9B=BE=E4=B8=8E=E5=9B=BE=E7=89=87=E7=BC=96=E8=BE=91?= =?UTF-8?q?=20API=20=E8=B0=83=E7=94=A8=EF=BC=8C=E6=96=B0=E5=A2=9E=E7=AE=A1?= =?UTF-8?q?=E7=BA=BF=E4=B8=8E=E7=BC=96=E8=BE=91=20HTTP=20=E7=AB=AF?= =?UTF-8?q?=E7=82=B9=EF=BC=8CJWT=20=E8=AE=A4=E8=AF=81=E4=B8=AD=E9=97=B4?= =?UTF-8?q?=E4=BB=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 需鉴权 --- backend/cmd/main.go | 11 +- backend/internal/handler/edit.go | 59 +++++++ backend/internal/handler/generate.go | 84 ++++++++++ backend/internal/mildware/auth.go | 48 ++++++ backend/internal/service/inference.go | 231 +++++++++++++++++++++++--- 5 files changed, 412 insertions(+), 21 deletions(-) create mode 100644 backend/internal/handler/edit.go create mode 100644 backend/internal/handler/generate.go create mode 100644 backend/internal/mildware/auth.go diff --git a/backend/cmd/main.go b/backend/cmd/main.go index 0d4c328..2f2e859 100644 --- a/backend/cmd/main.go +++ b/backend/cmd/main.go @@ -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") { diff --git a/backend/internal/handler/edit.go b/backend/internal/handler/edit.go new file mode 100644 index 0000000..b17d052 --- /dev/null +++ b/backend/internal/handler/edit.go @@ -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})) +} diff --git a/backend/internal/handler/generate.go b/backend/internal/handler/generate.go new file mode 100644 index 0000000..52b38af --- /dev/null +++ b/backend/internal/handler/generate.go @@ -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, + })) +} diff --git a/backend/internal/mildware/auth.go b/backend/internal/mildware/auth.go new file mode 100644 index 0000000..9a0226a --- /dev/null +++ b/backend/internal/mildware/auth.go @@ -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 ")) + 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() + } +} diff --git a/backend/internal/service/inference.go b/backend/internal/service/inference.go index 30b6ee0..ac9c896 100644 --- a/backend/internal/service/inference.go +++ b/backend/internal/service/inference.go @@ -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)