diff --git a/.gitignore b/.gitignore old mode 100644 new mode 100755 index 1fa64c0..bc4fa7d --- a/.gitignore +++ b/.gitignore @@ -23,3 +23,6 @@ backend/bin/ backend/data/ backend/.env backend/main + +# Generated output +generation/ diff --git a/backend/.env.example b/backend/.env.example old mode 100644 new mode 100755 index 3bdd87e..cb4279c --- a/backend/.env.example +++ b/backend/.env.example @@ -19,12 +19,13 @@ GEN2D_LLM_MODEL=gpt-4o GEN2D_LLM_TEMPERATURE=0.7 GEN2D_LLM_MAX_TOKENS=2048 -# 文生图模型 -GEN2D_IMAGE_BASE_URL=https://api.stability.ai/v1 -GEN2D_IMAGE_API_KEY=sk-your-api-key -GEN2D_IMAGE_MODEL=stable-diffusion-xl +# 文生图模型(OpenAI 兼容 Images API) +GEN2D_IMAGE_BASE_URL=https://api.suchuang.vip/v1 +GEN2D_IMAGE_API_KEY= +GEN2D_IMAGE_MODEL=gpt-image-2-token GEN2D_IMAGE_WIDTH=1024 GEN2D_IMAGE_HEIGHT=1024 +GEN2D_IMAGE_QUALITY=low GEN2D_IMAGE_NUM_IMAGES=1 GEN2D_IMAGE_STEPS=30 GEN2D_IMAGE_CFG_SCALE=7.0 diff --git a/backend/.gitignore b/backend/.gitignore new file mode 100755 index 0000000..71eeeb9 --- /dev/null +++ b/backend/.gitignore @@ -0,0 +1,2 @@ +test_output/ +generation/ diff --git a/backend/cmd/main.go b/backend/cmd/main.go old mode 100644 new mode 100755 index 07b34bf..8815037 --- a/backend/cmd/main.go +++ b/backend/cmd/main.go @@ -4,10 +4,12 @@ package main import ( "fmt" "log" + "os" "gen2d/internal/config" "gen2d/internal/db" "gen2d/internal/handler" + "gen2d/internal/mildware" "gen2d/internal/model" "gen2d/internal/service" @@ -19,6 +21,9 @@ func main() { gin.SetMode(cfg.Server.Mode) + // 确保 generation 输出目录存在(项目根级别) + _ = os.MkdirAll("../generation", 0755) + // 初始化 SQLite 数据库 if err := db.Init(cfg.Database.DSN, &model.User{}); err != nil { log.Fatalf("db init failed: %v", err) @@ -38,7 +43,10 @@ func main() { r := gin.New() r.Use(gin.Recovery()) // panic 恢复中间件,防止服务因未捕获异常宕机 - // API v1 路由组 + // 静态文件服务 — 生成的图片 + r.Static("/generation", "../generation") + + // API v1 路由组 — 公开接口 v1 := r.Group("/api/v1") { v1.GET("/health", handler.Health) // 健康检查 @@ -46,6 +54,16 @@ func main() { v1.GET("/assets/download", handler.DownloadAsset) // 素材下载 } + // API v1 路由组 — 需认证 + v1Auth := r.Group("/api/v1") + v1Auth.Use(mildware.AuthMiddleware(cfg.JWT.Secret)) + { + v1Auth.POST("/generate", handler.Generate) // 素材生成管线 + v1Auth.GET("/tasks/:taskId", handler.GetTask) // 查询任务 + v1Auth.GET("/tasks/:taskId/assets", handler.GetAssets) // 查询任务素材 + v1Auth.POST("/images/edit", handler.EditImage) // 图片编辑 + } + // Auth 路由组 auth := r.Group("/auth") { diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go old mode 100644 new mode 100755 index 38d90b2..0810193 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -44,16 +44,14 @@ type LLMConfig struct { MaxTokens int `mapstructure:"max_tokens"` } -// ImageGenConfig 文生图模型配置。 +// ImageGenConfig OpenAI 兼容文生图模型配置。 type ImageGenConfig struct { - BaseURL string `mapstructure:"base_url"` - APIKey string `mapstructure:"api_key"` - Model string `mapstructure:"model"` - Width int `mapstructure:"width"` - Height int `mapstructure:"height"` - NumImages int `mapstructure:"num_images"` - Steps int `mapstructure:"steps"` - CFGScale float64 `mapstructure:"cfg_scale"` + BaseURL string `mapstructure:"base_url"` + APIKey string `mapstructure:"api_key"` + Model string `mapstructure:"model"` + Width int `mapstructure:"width"` + Height int `mapstructure:"height"` + Quality string `mapstructure:"quality"` } // QiniuConfig 七牛云对象存储配置。 @@ -113,11 +111,12 @@ func setDefaults(v *viper.Viper) { v.SetDefault("llm.temperature", 0.7) v.SetDefault("llm.max_tokens", 2048) - v.SetDefault("image_gen.base_url", "https://api.stability.ai/v1") + v.SetDefault("image_gen.base_url", "https://api.suchuang.vip/v1") v.SetDefault("image_gen.api_key", "") - v.SetDefault("image_gen.model", "stable-diffusion-xl") + v.SetDefault("image_gen.model", "gpt-image-2-token") v.SetDefault("image_gen.width", 1024) v.SetDefault("image_gen.height", 1024) + v.SetDefault("image_gen.quality", "low") v.SetDefault("image_gen.num_images", 1) v.SetDefault("image_gen.steps", 30) v.SetDefault("image_gen.cfg_scale", 7.0) @@ -148,6 +147,7 @@ func bindEnvVars(v *viper.Viper) { v.BindEnv("image_gen.model", "GEN2D_IMAGE_MODEL") v.BindEnv("image_gen.width", "GEN2D_IMAGE_WIDTH") v.BindEnv("image_gen.height", "GEN2D_IMAGE_HEIGHT") + v.BindEnv("image_gen.quality", "GEN2D_IMAGE_QUALITY") v.BindEnv("image_gen.num_images", "GEN2D_IMAGE_NUM_IMAGES") v.BindEnv("image_gen.steps", "GEN2D_IMAGE_STEPS") v.BindEnv("image_gen.cfg_scale", "GEN2D_IMAGE_CFG_SCALE") diff --git a/backend/internal/config/config.yml b/backend/internal/config/config.yml old mode 100644 new mode 100755 index 84475e7..aa29e74 --- a/backend/internal/config/config.yml +++ b/backend/internal/config/config.yml @@ -21,11 +21,9 @@ llm: max_tokens: 2048 image_gen: - base_url: "https://api.stability.ai/v1" + base_url: "https://api.suchuang.vip/v1" api_key: "" - model: "stable-diffusion-xl" + model: "gpt-image-2-token" width: 1024 height: 1024 - num_images: 1 - steps: 30 - cfg_scale: 7.0 + quality: "low" diff --git a/backend/internal/handler/edit.go b/backend/internal/handler/edit.go new file mode 100755 index 0000000..f77b05a --- /dev/null +++ b/backend/internal/handler/edit.go @@ -0,0 +1,65 @@ +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 +} + +// editAssetResponse 编辑结果素材(返回 base64)。 +type editAssetResponse struct { + Data string `json:"data"` + Format string `json:"format"` +} + +// EditImageResponse 图片编辑响应体。 +type EditImageResponse struct { + Assets []editAssetResponse `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([]editAssetResponse, len(images)) + for i, img := range images { + assets[i] = editAssetResponse{ + 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 100755 index 0000000..4f54ddc --- /dev/null +++ b/backend/internal/handler/generate.go @@ -0,0 +1,234 @@ +package handler + +import ( + "context" + "fmt" + "log" + "net/http" + "os" + "path/filepath" + "sync" + "time" + + "gen2d/internal/model" + "gen2d/internal/service" + + "github.com/gin-gonic/gin" +) + +// GenerateRequest 素材生成请求。 +type GenerateRequest struct { + ProjectID string `json:"projectId"` + 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 { + TaskID string `json:"taskId"` +} + +// AssetResponse 单个素材响应。 +type AssetResponse struct { + URL string `json:"url"` + Format string `json:"format"` +} + +// TaskResponse 任务查询响应。 +type TaskResponse struct { + ID string `json:"id"` + ProjectID string `json:"projectId"` + Prompt string `json:"prompt"` + AssetType string `json:"assetType"` + Status string `json:"status"` + Progress int `json:"progress"` + RetryCount int `json:"retryCount"` + Error string `json:"error,omitempty"` + CreatedAt string `json:"createdAt"` +} + +// AssetsResponse 素材列表响应。 +type AssetsResponse struct { + Assets []AssetResponse `json:"assets"` + Metadata service.AssetMetadata `json:"metadata"` +} + +// taskRecord 内存中的任务记录。 +type taskRecord struct { + task TaskResponse + assets []AssetResponse + metadata service.AssetMetadata +} + +var ( + taskStore = sync.Map{} // taskID → *taskRecord +) + +// Generate 素材生成接口(异步)。 +// 立即返回 taskId,后台执行管线,前端通过 GET /tasks/:taskId 轮询进度。 +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 + } + + projectID := req.ProjectID + if projectID == "" { + projectID = "default" + } + taskID := fmt.Sprintf("task-%d", time.Now().UnixMilli()) + createdAt := time.Now().Format(time.RFC3339) + + // 存入 pending 状态 + taskStore.Store(taskID, &taskRecord{ + task: TaskResponse{ + ID: taskID, + ProjectID: projectID, + Prompt: req.Prompt, + AssetType: req.AssetType, + Status: "pending", + Progress: 0, + CreatedAt: createdAt, + }, + }) + + // 返回 taskId + c.JSON(http.StatusOK, model.OK(GenerateResponse{TaskID: taskID})) + + // 后台执行管线 + go runPipelineBg(projectID, taskID, req) +} + +// runPipelineBg 后台执行生成管线,更新任务状态。 +func runPipelineBg(projectID, taskID string, req GenerateRequest) { + updateStatus(taskID, "running", 10) + + in := service.PipelineInput{ + ProjectID: projectID, + TaskID: taskID, + 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, + }, + } + + ctx := context.Background() + output, err := service.RunPipeline(ctx, in) + if err != nil { + log.Printf("[generate] task %s failed: %v", taskID, err) + updateFailed(taskID, err.Error()) + return + } + + updateStatus(taskID, "saving", 80) + + // 保存图片到 ../generation/{projectId}/{taskId}/ + genDir := filepath.Join("..", "generation", projectID, taskID) + if err := os.MkdirAll(genDir, 0755); err != nil { + updateFailed(taskID, "创建输出目录失败: "+err.Error()) + return + } + + assets := make([]AssetResponse, len(output.Assets)) + for i, a := range output.Assets { + filename := fmt.Sprintf("%d.%s", i, a.Format) + filePath := filepath.Join(genDir, filename) + if err := os.WriteFile(filePath, a.Data, 0644); err != nil { + updateFailed(taskID, "保存图片失败: "+err.Error()) + return + } + assets[i] = AssetResponse{ + URL: fmt.Sprintf("/generation/%s/%s/%s", projectID, taskID, filename), + Format: a.Format, + } + } + + // 更新为完成状态 + taskStore.Store(taskID, &taskRecord{ + task: TaskResponse{ + ID: taskID, + ProjectID: projectID, + Prompt: req.Prompt, + AssetType: req.AssetType, + Status: "completed", + Progress: 100, + CreatedAt: time.Now().Format(time.RFC3339), + }, + assets: assets, + metadata: output.Metadata, + }) + + log.Printf("[generate] task %s completed, %d assets", taskID, len(assets)) +} + +func updateStatus(taskID, status string, progress int) { + rec, ok := taskStore.Load(taskID) + if !ok { + return + } + r := rec.(*taskRecord) + r.task.Status = status + r.task.Progress = progress + taskStore.Store(taskID, r) +} + +func updateFailed(taskID, errMsg string) { + rec, ok := taskStore.Load(taskID) + if !ok { + return + } + r := rec.(*taskRecord) + r.task.Status = "failed" + r.task.Error = errMsg + taskStore.Store(taskID, r) +} + +// GetTask 查询任务信息。 +func GetTask(c *gin.Context) { + taskID := c.Param("taskId") + rec, ok := taskStore.Load(taskID) + if !ok { + c.JSON(http.StatusNotFound, model.Fail(http.StatusNotFound, "任务不存在")) + return + } + r := rec.(*taskRecord) + c.JSON(http.StatusOK, model.OK(r.task)) +} + +// GetAssets 查询任务素材列表。 +func GetAssets(c *gin.Context) { + taskID := c.Param("taskId") + rec, ok := taskStore.Load(taskID) + if !ok { + c.JSON(http.StatusNotFound, model.Fail(http.StatusNotFound, "任务不存在")) + return + } + r := rec.(*taskRecord) + if r.task.Status != "completed" { + c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "任务尚未完成,当前状态: "+r.task.Status)) + return + } + c.JSON(http.StatusOK, model.OK(AssetsResponse{ + Assets: r.assets, + Metadata: r.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 old mode 100644 new mode 100755 index 30b6ee0..f334a34 --- a/backend/internal/service/inference.go +++ b/backend/internal/service/inference.go @@ -4,10 +4,17 @@ import ( "bytes" "context" "crypto/rand" + "encoding/base64" + "encoding/json" "fmt" "image" "image/color" "image/png" + "io" + "log" + "mime/multipart" + "net/http" + "strings" "gen2d/internal/config" ) @@ -20,50 +27,204 @@ 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 - } +// ======================== OpenAI 兼容 Images API 类型 ======================== +// imageGenRequest OpenAI 兼容文生图请求体。 +type imageGenRequest struct { + Model string `json:"model"` + Prompt string `json:"prompt"` + N int `json:"n,omitempty"` + Size string `json:"size,omitempty"` + Quality string `json:"quality,omitempty"` +} + +// imageGenResponse OpenAI 兼容文生图响应体。 +type imageGenResponse struct { + Data []struct { + URL string `json:"url"` + B64JSON string `json:"b64_json"` + } `json:"data"` +} + +// ======================== 文生图 ======================== + +// GenerateImages 调用 OpenAI 兼容 Images 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 != "" { + width, height := imgCfg.Width, imgCfg.Height + if params.Resolution > 0 { + width = params.Resolution + height = params.Resolution } - images[i] = GeneratedImage{ + log.Println("[inference] calling image gen API") + return callImageAPI(ctx, prompt, count, width, height) + } + + size := params.Resolution + if size <= 0 { + size = 64 + } + log.Println("[inference] image API key not configured, using mock") + return generateMockImages(size, count) +} + +// callImageAPI 调用 OpenAI 兼容 Images API。 +func callImageAPI(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), + Quality: imgCfg.Quality, + } + + body, err := json.Marshal(reqBody) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + + url := strings.TrimRight(imgCfg.BaseURL, "/") + "/images/generations" + 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 api error %d: %s", resp.StatusCode, string(b)) + } + + return parseImageResponse(ctx, resp.Body, width, height) +} + +// parseImageResponse 解析 OpenAI 兼容图片响应(b64_json 或 url)。 +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: size, - Height: size, + Width: width, + Height: height, Format: "png", - } + }) } return images, nil } -// QualityChecker 质检函数,可替换用于测试。 -// 签名:(ctx, images, style) → (pass, reason, error) +// 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) +} + +// ======================== 图片编辑 ======================== + +// EditImages 图片编辑接口,调用 OpenAI 兼容 Images Edits API(multipart/form-data)。 +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") + } + + 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", fmt.Sprintf("%d", 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.TrimRight(imgCfg.BaseURL, "/") + "/images/edits" + 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) +} + +// ======================== 质检 ======================== + 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) { @@ -71,32 +232,40 @@ 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 } } -// generateMockImage 生成一张带随机色块的 PNG 占位图 +// ======================== Mock 回退 ======================== + +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 +} + 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) - 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 @@ -104,7 +273,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) diff --git a/backend/internal/service/types.go b/backend/internal/service/types.go old mode 100644 new mode 100755 index ed684b0..9db9f27 --- a/backend/internal/service/types.go +++ b/backend/internal/service/types.go @@ -2,6 +2,8 @@ package service // PipelineInput 管线入口输入 type PipelineInput struct { + ProjectID string // 工程 ID,用于输出目录 + TaskID string // 任务 ID,用于输出目录 Prompt string // 用户原始文本 AssetType string // 素材类型:sprite / background / ui / animation ProjectStyle map[string]string // 工程风格键值对 diff --git a/backend/tools/gentest.go b/backend/tools/gentest.go new file mode 100644 index 0000000..d6af462 --- /dev/null +++ b/backend/tools/gentest.go @@ -0,0 +1,49 @@ +//go:build ignore + +package main + +import ( + "context" + "fmt" + "os" + + "gen2d/internal/config" + "gen2d/internal/service" +) + +func main() { + cfg := config.Load() + + // 清空 API key 强制走 mock + cfg.ImageGen.APIKey = "" + service.InitImageGenConfig(cfg.ImageGen) + service.InitLLMConfig(cfg.LLM) + + in := service.PipelineInput{ + AssetType: "sprite", + Prompt: "a cute cat warrior with golden armor", + Tags: []string{"pixel art", "fantasy", "16-bit"}, + Params: service.AssetParams{ + Resolution: 816, + Frames: service.FrameParams{Directions: 4, FramesPerDirection: 1}, + Format: "individual", + }, + } + + output, err := service.RunPipeline(context.Background(), in) + if err != nil { + fmt.Fprintf(os.Stderr, "error: %v\n", err) + os.Exit(1) + } + + os.MkdirAll("test_output", 0755) + for i, a := range output.Assets { + fname := fmt.Sprintf("test_output/gen_%d.%s", i, a.Format) + if err := os.WriteFile(fname, a.Data, 0644); err != nil { + fmt.Fprintf(os.Stderr, "write %s: %v\n", fname, err) + os.Exit(1) + } + fmt.Printf("saved %s (%d bytes)\n", fname, len(a.Data)) + } + fmt.Printf("metadata: %+v\n", output.Metadata) +} diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts old mode 100644 new mode 100755 index 9ff77c8..0ec0a26 --- a/frontend/src/api/client.ts +++ b/frontend/src/api/client.ts @@ -41,10 +41,6 @@ async function request(url: string, options: RequestInit = {}): Promise { const json: ApiResponse = await res.json() if (json.code !== 0) { - if (json.code === 401) { - clearToken() - window.location.href = '/login' - } throw new ApiError(json.code, json.message) } diff --git a/frontend/src/api/generate.ts b/frontend/src/api/generate.ts old mode 100644 new mode 100755 index 40670a8..a03b2c6 --- a/frontend/src/api/generate.ts +++ b/frontend/src/api/generate.ts @@ -1,23 +1,35 @@ -import type { Asset, Task } from './types' -import { mockGetAssets, mockGetTask, mockSubmitGenerate } from './mock' - -const USE_MOCK = true +import { post, get } from './client' +import type { + Asset, + AssetsResponse, + GenerateRequest, + GenerateResponse, + Task, +} from './types' export async function submitGenerate( - projectId: string, - prompt: string, - assetType: string -): Promise { - if (USE_MOCK) return mockSubmitGenerate(projectId, prompt, assetType) - throw new Error('Not implemented') + req: GenerateRequest, +): Promise { + return post('/api/v1/generate', req) } export async function getTask(taskId: string): Promise { - if (USE_MOCK) return mockGetTask(taskId) - throw new Error('Not implemented') + return get(`/api/v1/tasks/${taskId}`) } export async function getAssets(taskId: string): Promise { - if (USE_MOCK) return mockGetAssets(taskId) - throw new Error('Not implemented') + const resp = await get(`/api/v1/tasks/${taskId}/assets`) + return resp.assets.map((a, i) => ({ + id: `asset-${i}`, + url: a.url, + format: a.format, + width: resp.metadata.frameWidth, + height: resp.metadata.frameHeight, + metadata: { + frameWidth: resp.metadata.frameWidth, + frameHeight: resp.metadata.frameHeight, + frameCount: resp.metadata.frameCount, + directions: resp.metadata.directions, + }, + })) } diff --git a/frontend/src/api/types.ts b/frontend/src/api/types.ts old mode 100644 new mode 100755 index b9eaa96..3fbdefd --- a/frontend/src/api/types.ts +++ b/frontend/src/api/types.ts @@ -67,7 +67,7 @@ export interface Task { projectId: string prompt: string assetType: string - status: 'pending' | 'running' | 'completed' | 'failed' + status: 'pending' | 'submitted' | 'running' | 'completed' | 'failed' stage?: PipelineStage progress?: number retryCount?: number @@ -91,16 +91,37 @@ export interface Asset { } } -// 生成请求 +// 生成请求 — 对应 POST /api/v1/generate export interface GenerateRequest { projectId: string - prompt: string + prompt?: string assetType: AssetType + tags?: string[] + userNote?: string + projectStyle?: Record taskStyle?: Record - params?: { - resolution?: number - frames?: { directions?: number; framesPerDirection?: number } - format?: 'spritesheet' | 'individual' + resolution?: number + directions?: number + framesPerDir?: number + format?: 'spritesheet' | 'individual' +} + +// 生成响应 — 对应 POST /api/v1/generate 返回(异步,仅含 taskId) +export interface GenerateResponse { + taskId: string +} + +// 素材列表响应 — 对应 GET /api/v1/tasks/:taskId/assets +export interface AssetsResponse { + assets: { + url: string + format: string + }[] + metadata: { + frameWidth: number + frameHeight: number + frameCount: number + directions: number } } diff --git a/frontend/src/hooks/useGenerate.ts b/frontend/src/hooks/useGenerate.ts old mode 100644 new mode 100755 index 49220c7..6bf510e --- a/frontend/src/hooks/useGenerate.ts +++ b/frontend/src/hooks/useGenerate.ts @@ -5,11 +5,8 @@ export function useGenerate() { return { submit: store.submit, taskId: store.taskId, - stage: store.stage, progress: store.progress, status: store.status, - retryCount: store.retryCount, - rejectReason: store.rejectReason, assets: store.assets, error: store.error, reset: store.reset, diff --git a/frontend/src/pages/GeneratePage.tsx b/frontend/src/pages/GeneratePage.tsx old mode 100644 new mode 100755 index bb17ec7..1912994 --- a/frontend/src/pages/GeneratePage.tsx +++ b/frontend/src/pages/GeneratePage.tsx @@ -1,27 +1,36 @@ import { useEffect } from 'react' import { useNavigate, useParams } from 'react-router-dom' import { useTaskStore } from '../stores/task' +import { useProjectStore } from '../stores/project' import { useGenerationStore } from '../stores/generation' import { useToastStore } from '../stores/toast' +import { extractTags } from '../api/prompt' +import { mergeStyles } from '../utils/style' +import type { GenerateRequest } from '../api/types' import GenerateForm from '../components/GenerateForm' import ProgressBar from '../components/ProgressBar' export default function GeneratePage() { const { projectId = 'proj-default' } = useParams() const navigate = useNavigate() - const { assetType, reset: resetTask } = useTaskStore() const addToast = useToastStore(s => s.addToast) + + const taskStore = useTaskStore() + const { style: projectStyle, loadProject } = useProjectStore() const { status, - stage, progress, - retryCount, - rejectReason, taskId, + statusText, submit, reset: resetGeneration, } = useGenerationStore() + // 加载工程风格 + useEffect(() => { + loadProject(projectId) + }, [projectId, loadProject]) + // 组件卸载时重置生成状态 useEffect(() => { return () => resetGeneration() @@ -30,63 +39,84 @@ export default function GeneratePage() { // 失败时显示 toast useEffect(() => { if (status === 'failed') { - addToast({ type: 'error', message: '素材生成失败,请重试' }) + const errText = useGenerationStore.getState().error || '未知错误' + addToast({ type: 'error', message: `素材生成失败:${errText}` }) } }, [status, addToast]) - // 完成后自动跳转 - useEffect(() => { - if (status === 'completed' && taskId) { - const timer = setTimeout(() => { - navigate(`/projects/${projectId}/tasks/${taskId}`) - }, 1500) - return () => clearTimeout(timer) - } - }, [status, taskId, projectId, navigate]) - const handleSubmit = async (finalPrompt: string) => { - await submit(projectId, finalPrompt, assetType) + const { taskStyle, params, enableAI, optimizedPrompt } = taskStore + const mergedStyle = mergeStyles(projectStyle, taskStyle) + const tags = extractTags(mergedStyle) + + const req: GenerateRequest = { + projectId, + prompt: enableAI && optimizedPrompt ? optimizedPrompt : finalPrompt, + assetType: taskStore.assetType, + tags, + projectStyle, + taskStyle, + resolution: params.resolution, + directions: params.frames?.directions, + framesPerDir: params.frames?.framesPerDirection, + format: params.format, + } + + await submit(req) + } + + const handleViewResult = () => { + if (taskId) navigate(`/projects/${projectId}/tasks/${taskId}`) } const handleReset = () => { resetGeneration() - resetTask() + taskStore.reset() } return (

新建生成

- {/* 生成表单 */} {status === 'idle' || status === 'submitting' ? ( ) : (
- {/* 进度条 */} - {/* 状态提示 */} {status === 'running' && (

- 管线执行中,请稍候... + {statusText || '管线执行中,请稍候...'}

)} {status === 'completed' && ( -

- ✓ 生成完成,正在跳转到结果页... -

+
+

+ 生成完成 +

+
+ + +
+
)} {status === 'failed' && (
-

生成失败

+

+ {useGenerationStore.getState().error || '生成失败'} +

diff --git a/frontend/src/pages/ResultPage.tsx b/frontend/src/pages/ResultPage.tsx old mode 100644 new mode 100755 index 9031397..85da8b3 --- a/frontend/src/pages/ResultPage.tsx +++ b/frontend/src/pages/ResultPage.tsx @@ -1,4 +1,4 @@ -import { useEffect, useState } from 'react' +import { useEffect, useState, useRef } from 'react' import { Link, useParams } from 'react-router-dom' import { getTask, getAssets } from '../api/generate' import type { Asset, Task } from '../api/types' @@ -10,19 +10,47 @@ export default function ResultPage() { const [task, setTask] = useState(null) const [assets, setAssets] = useState([]) const [loading, setLoading] = useState(true) + const [polling, setPolling] = useState(false) + const pollRef = useRef>() useEffect(() => { if (!taskId) return setLoading(true) - Promise.all([getTask(taskId), getAssets(taskId)]) - .then(([t, a]) => { + + const fetchTask = async () => { + try { + const t = await getTask(taskId) setTask(t) - setAssets(a) - }) - .finally(() => setLoading(false)) + + if (t.status === 'completed') { + if (pollRef.current) clearInterval(pollRef.current) + setPolling(false) + const a = await getAssets(taskId) + setAssets(a) + setLoading(false) + } else if (t.status === 'failed') { + if (pollRef.current) clearInterval(pollRef.current) + setPolling(false) + setLoading(false) + } else if (!pollRef.current) { + // 开始轮询 + setPolling(true) + pollRef.current = setInterval(fetchTask, 2000) + } + } catch { + // 出错也停止加载态 + setLoading(false) + } + } + + fetchTask() + + return () => { + if (pollRef.current) clearInterval(pollRef.current) + } }, [taskId]) - if (loading) { + if (loading || polling) { return (
@@ -30,6 +58,13 @@ export default function ResultPage() {
+ {task && ( +

+ {task.status === 'pending' && '任务排队中...'} + {task.status === 'running' && `生成中... ${task.progress ?? 0}%`} + {task.status === 'submitted' && '已提交,等待处理...'} +

+ )}
@@ -53,27 +88,35 @@ export default function ResultPage() {

生成结果

- {/* 任务信息 */}
-

任务信息

-
+
+

任务信息

+
+ {assets.length > 0 && ( + + )} + + 继续生成 + +
+
+
提示词 {task.prompt} 素材类型 {task.assetType} 状态 - + {task.status === 'completed' ? '已完成' : '失败'} 创建时间 @@ -93,29 +136,29 @@ export default function ResultPage() {
- {/* 素材预览 */}
-
-

素材预览

- {assets.length > 0 && ( - - )} -
+

素材预览

) } + +async function downloadAssets(assets: Asset[]) { + for (const a of assets) { + try { + const res = await fetch(a.url) + const blob = await res.blob() + const blobUrl = URL.createObjectURL(blob) + const link = document.createElement('a') + link.href = blobUrl + link.download = `${a.id}.${a.format}` + document.body.appendChild(link) + link.click() + document.body.removeChild(link) + URL.revokeObjectURL(blobUrl) + } catch { + window.open(a.url, '_blank') + } + } +} diff --git a/frontend/src/stores/generation.ts b/frontend/src/stores/generation.ts old mode 100644 new mode 100755 index b904755..d7346d1 --- a/frontend/src/stores/generation.ts +++ b/frontend/src/stores/generation.ts @@ -1,80 +1,105 @@ import { create } from 'zustand' -import type { Asset, PipelineProgress, PipelineStage } from '../api/types' -import { submitGenerate } from '../api/generate' -import { createMockWebSocket } from '../api/mock' +import type { Asset, GenerateRequest } from '../api/types' +import { submitGenerate, getTask, getAssets } from '../api/generate' type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed' interface GenerationState { taskId: string | null - stage: PipelineStage | null + projectId: string | null progress: number status: Status - retryCount: number - rejectReason: string | null + statusText: string assets: Asset[] error: string | null - submit: (projectId: string, prompt: string, assetType: string) => Promise - handleProgress: (msg: PipelineProgress) => void + submit: (req: GenerateRequest) => Promise reset: () => void } -let cleanupWs: (() => void) | null = null +let pollTimer: ReturnType | null = null + +function stopPolling() { + if (pollTimer) { + clearInterval(pollTimer) + pollTimer = null + } +} export const useGenerationStore = create((set, get) => ({ taskId: null, - stage: null, + projectId: null, progress: 0, status: 'idle', - retryCount: 0, - rejectReason: null, + statusText: '', assets: [], error: null, - submit: async (projectId, prompt, assetType) => { - set({ status: 'submitting', error: null }) + submit: async (req) => { + stopPolling() + set({ status: 'submitting', error: null, statusText: '提交中...' }) try { - const taskId = await submitGenerate(projectId, prompt, assetType) - set({ taskId, status: 'running', progress: 0 }) + const { taskId } = await submitGenerate(req) - // 启动 mock WebSocket - cleanupWs = createMockWebSocket( + set({ taskId, - (msg) => get().handleProgress(msg), - (assets) => { - set({ status: 'completed', assets, progress: 100 }) - }, - (error) => { - set({ status: 'failed', error }) - } - ) - } catch (err) { - set({ status: 'failed', error: (err as Error).message }) - } - }, + projectId: req.projectId, + status: 'running', + progress: 10, + statusText: '任务已提交,等待生成...', + }) - handleProgress: (msg) => { - set({ - stage: msg.stage, - progress: msg.progress, - retryCount: msg.retryCount ?? get().retryCount, - rejectReason: msg.rejectReason ?? null, - }) - if (msg.result?.assets) { - set({ assets: msg.result.assets }) + // 开始轮询进度 + pollTimer = setInterval(async () => { + try { + const task = await getTask(taskId) + + set({ + progress: task.progress ?? get().progress, + statusText: + task.status === 'running' + ? '生成中...' + : task.status === 'pending' + ? '排队中...' + : task.status, + }) + + if (task.status === 'completed') { + stopPolling() + set({ progress: 90, statusText: '获取结果...' }) + + const assets = await getAssets(taskId) + set({ + status: 'completed', + progress: 100, + statusText: '生成完成', + assets, + }) + } else if (task.status === 'failed') { + stopPolling() + set({ + status: 'failed', + error: task.error || '生成失败', + statusText: '生成失败', + }) + } + } catch { + // 网络错误不中断轮询 + } + }, 2000) + } catch (err) { + stopPolling() + set({ status: 'failed', error: (err as Error).message, statusText: '提交失败' }) } }, reset: () => { - cleanupWs?.() - cleanupWs = null + stopPolling() set({ taskId: null, - stage: null, + projectId: null, progress: 0, status: 'idle', - retryCount: 0, - rejectReason: null, + statusText: '', assets: [], error: null, })