feat(generate): 图片持久化到本地 generation/ 目录,前端对接真实 API

后端:
- POST /api/v1/generate 接收 projectId,生成 taskId,将图片保存至 generation/{projectId}/{taskId}/
- 新增 GET /api/v1/tasks/:taskId 和 GET /api/v1/tasks/:taskId/assets 查询端点
- 添加 /generation 静态文件服务,前端可通过 URL 直接访问生成的图片
- PipelineInput 增加 ProjectID/TaskID 字段

前端:
- generate.ts 对接真实后端 API,移除 mock 模式
- 401 时不再强制登出跳转,改为抛错由调用方处理
- GeneratePage 加载工程风格,传递 projectId 和完整参数
- generation store 简化为同步 API 模式,移除 WebSocket mock
- ResultPage 使用 getTask/getAssets 按 taskId 查询结果
This commit is contained in:
2026-05-25 12:27:07 +08:00
parent 6bda5a1f35
commit 74c1725473
13 changed files with 297 additions and 114 deletions
Regular → Executable
+3
View File
@@ -23,3 +23,6 @@ backend/bin/
backend/data/ backend/data/
backend/.env backend/.env
backend/main backend/main
# Generated output
generation/
Regular → Executable
+1
View File
@@ -1 +1,2 @@
test_output/ test_output/
generation/
Regular → Executable
+9
View File
@@ -4,6 +4,7 @@ package main
import ( import (
"fmt" "fmt"
"log" "log"
"os"
"gen2d/internal/config" "gen2d/internal/config"
"gen2d/internal/db" "gen2d/internal/db"
@@ -20,6 +21,9 @@ func main() {
gin.SetMode(cfg.Server.Mode) gin.SetMode(cfg.Server.Mode)
// 确保 generation 输出目录存在(项目根级别)
_ = os.MkdirAll("../generation", 0755)
// 初始化 SQLite 数据库 // 初始化 SQLite 数据库
if err := db.Init(cfg.Database.DSN, &model.User{}); err != nil { if err := db.Init(cfg.Database.DSN, &model.User{}); err != nil {
log.Fatalf("db init failed: %v", err) log.Fatalf("db init failed: %v", err)
@@ -35,6 +39,9 @@ func main() {
r := gin.New() r := gin.New()
r.Use(gin.Recovery()) // panic 恢复中间件,防止服务因未捕获异常宕机 r.Use(gin.Recovery()) // panic 恢复中间件,防止服务因未捕获异常宕机
// 静态文件服务 — 生成的图片
r.Static("/generation", "../generation")
// API v1 路由组 — 公开接口 // API v1 路由组 — 公开接口
v1 := r.Group("/api/v1") v1 := r.Group("/api/v1")
{ {
@@ -47,6 +54,8 @@ func main() {
v1Auth.Use(mildware.AuthMiddleware(cfg.JWT.Secret)) v1Auth.Use(mildware.AuthMiddleware(cfg.JWT.Secret))
{ {
v1Auth.POST("/generate", handler.Generate) // 素材生成管线 v1Auth.POST("/generate", handler.Generate) // 素材生成管线
v1Auth.GET("/tasks/:taskId", handler.GetTask) // 查询任务
v1Auth.GET("/tasks/:taskId/assets", handler.GetAssets) // 查询任务素材
v1Auth.POST("/images/edit", handler.EditImage) // 图片编辑 v1Auth.POST("/images/edit", handler.EditImage) // 图片编辑
} }
Regular → Executable
+9 -3
View File
@@ -17,9 +17,15 @@ type EditImageRequest struct {
Count int `json:"count"` // 生成数量,默认 1 Count int `json:"count"` // 生成数量,默认 1
} }
// editAssetResponse 编辑结果素材(返回 base64)。
type editAssetResponse struct {
Data string `json:"data"`
Format string `json:"format"`
}
// EditImageResponse 图片编辑响应体。 // EditImageResponse 图片编辑响应体。
type EditImageResponse struct { type EditImageResponse struct {
Assets []AssetResponse `json:"assets"` Assets []editAssetResponse `json:"assets"`
} }
// EditImage 图片编辑接口,基于已有图片和文本指令生成修改后的图片。 // EditImage 图片编辑接口,基于已有图片和文本指令生成修改后的图片。
@@ -47,9 +53,9 @@ func EditImage(c *gin.Context) {
return return
} }
assets := make([]AssetResponse, len(images)) assets := make([]editAssetResponse, len(images))
for i, img := range images { for i, img := range images {
assets[i] = AssetResponse{ assets[i] = editAssetResponse{
Data: base64.StdEncoding.EncodeToString(img.Data), Data: base64.StdEncoding.EncodeToString(img.Data),
Format: img.Format, Format: img.Format,
} }
+105 -7
View File
@@ -1,8 +1,12 @@
package handler package handler
import ( import (
"encoding/base64" "fmt"
"net/http" "net/http"
"os"
"path/filepath"
"sync"
"time"
"gen2d/internal/model" "gen2d/internal/model"
"gen2d/internal/service" "gen2d/internal/service"
@@ -12,6 +16,7 @@ import (
// GenerateRequest 素材生成请求。 // GenerateRequest 素材生成请求。
type GenerateRequest struct { type GenerateRequest struct {
ProjectID string `json:"projectId"`
Prompt string `json:"prompt"` Prompt string `json:"prompt"`
AssetType string `json:"assetType" binding:"required"` AssetType string `json:"assetType" binding:"required"`
Tags []string `json:"tags"` Tags []string `json:"tags"`
@@ -26,18 +31,48 @@ type GenerateRequest struct {
// GenerateResponse 素材生成响应体。 // GenerateResponse 素材生成响应体。
type GenerateResponse struct { type GenerateResponse struct {
TaskID string `json:"taskId"`
Assets []AssetResponse `json:"assets"` Assets []AssetResponse `json:"assets"`
Metadata service.AssetMetadata `json:"metadata"` Metadata service.AssetMetadata `json:"metadata"`
} }
// AssetResponse 单个素材响应(二进制 Data 转 base64)。 // AssetResponse 单个素材响应。
type AssetResponse struct { type AssetResponse struct {
Data string `json:"data"`
Format string `json:"format"`
URL string `json:"url"` URL string `json:"url"`
Format string `json:"format"`
} }
// Generate 素材生成接口,调用完整生成管线(PromptOptimizer → AssetGenerator → QualitySupervisor → FormatAdapter)。 // 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 素材生成接口。
func Generate(c *gin.Context) { func Generate(c *gin.Context) {
var req GenerateRequest var req GenerateRequest
if err := c.ShouldBindJSON(&req); err != nil { if err := c.ShouldBindJSON(&req); err != nil {
@@ -45,7 +80,15 @@ func Generate(c *gin.Context) {
return return
} }
projectID := req.ProjectID
if projectID == "" {
projectID = "default"
}
taskID := fmt.Sprintf("task-%d", time.Now().UnixMilli())
in := service.PipelineInput{ in := service.PipelineInput{
ProjectID: projectID,
TaskID: taskID,
Prompt: req.Prompt, Prompt: req.Prompt,
AssetType: req.AssetType, AssetType: req.AssetType,
Tags: req.Tags, Tags: req.Tags,
@@ -68,17 +111,72 @@ func Generate(c *gin.Context) {
return return
} }
// 保存图片到 ../generation/{projectId}/{taskId}/
genDir := filepath.Join("..", "generation", projectID, taskID)
if err := os.MkdirAll(genDir, 0755); err != nil {
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "创建输出目录失败: "+err.Error()))
return
}
assets := make([]AssetResponse, len(output.Assets)) assets := make([]AssetResponse, len(output.Assets))
for i, a := range 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 {
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "保存图片失败: "+err.Error()))
return
}
assets[i] = AssetResponse{ assets[i] = AssetResponse{
Data: base64.StdEncoding.EncodeToString(a.Data), URL: fmt.Sprintf("/generation/%s/%s/%s", projectID, taskID, filename),
Format: a.Format, Format: a.Format,
URL: a.URL,
} }
} }
// 存储任务记录到内存
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,
})
c.JSON(http.StatusOK, model.OK(GenerateResponse{ c.JSON(http.StatusOK, model.OK(GenerateResponse{
TaskID: taskID,
Assets: assets, Assets: assets,
Metadata: output.Metadata, Metadata: output.Metadata,
})) }))
} }
// 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)
c.JSON(http.StatusOK, model.OK(AssetsResponse{
Assets: r.assets,
Metadata: r.metadata,
}))
}
+2
View File
@@ -2,6 +2,8 @@ package service
// PipelineInput 管线入口输入 // PipelineInput 管线入口输入
type PipelineInput struct { type PipelineInput struct {
ProjectID string // 工程 ID,用于输出目录
TaskID string // 任务 ID,用于输出目录
Prompt string // 用户原始文本 Prompt string // 用户原始文本
AssetType string // 素材类型:sprite / background / ui / animation AssetType string // 素材类型:sprite / background / ui / animation
ProjectStyle map[string]string // 工程风格键值对 ProjectStyle map[string]string // 工程风格键值对
Regular → Executable
-4
View File
@@ -41,10 +41,6 @@ async function request<T>(url: string, options: RequestInit = {}): Promise<T> {
const json: ApiResponse<T> = await res.json() const json: ApiResponse<T> = await res.json()
if (json.code !== 0) { if (json.code !== 0) {
if (json.code === 401) {
clearToken()
window.location.href = '/login'
}
throw new ApiError(json.code, json.message) throw new ApiError(json.code, json.message)
} }
Regular → Executable
+33 -14
View File
@@ -1,23 +1,42 @@
import type { Asset, Task } from './types' import { post, get } from './client'
import { mockGetAssets, mockGetTask, mockSubmitGenerate } from './mock' import type {
Asset,
const USE_MOCK = true AssetsResponse,
GenerateRequest,
GenerateResponse,
Task,
} from './types'
export async function submitGenerate( export async function submitGenerate(
projectId: string, req: GenerateRequest,
prompt: string, ): Promise<GenerateResponse> {
assetType: string return post<GenerateResponse>('/api/v1/generate', req)
): Promise<string> {
if (USE_MOCK) return mockSubmitGenerate(projectId, prompt, assetType)
throw new Error('Not implemented')
} }
export async function getTask(taskId: string): Promise<Task> { export async function getTask(taskId: string): Promise<Task> {
if (USE_MOCK) return mockGetTask(taskId) return get<Task>(`/api/v1/tasks/${taskId}`)
throw new Error('Not implemented')
} }
export async function getAssets(taskId: string): Promise<Asset[]> { export async function getAssets(taskId: string): Promise<Asset[]> {
if (USE_MOCK) return mockGetAssets(taskId) const resp = await get<AssetsResponse>(`/api/v1/tasks/${taskId}/assets`)
throw new Error('Not implemented') return toAssetList(resp)
}
/** 将响应转为 Asset[] 供前端组件使用 */
function toAssetList(
resp: AssetsResponse | GenerateResponse,
): Asset[] {
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,
},
}))
} }
Regular → Executable
+35 -4
View File
@@ -91,16 +91,47 @@ export interface Asset {
} }
} }
// 生成请求 // 生成请求 — 对应 POST /api/v1/generate
export interface GenerateRequest { export interface GenerateRequest {
projectId: string projectId: string
prompt: string prompt?: string
assetType: AssetType assetType: AssetType
tags?: string[]
userNote?: string
projectStyle?: Record<string, string>
taskStyle?: Record<string, string> taskStyle?: Record<string, string>
params?: {
resolution?: number resolution?: number
frames?: { directions?: number; framesPerDirection?: number } directions?: number
framesPerDir?: number
format?: 'spritesheet' | 'individual' format?: 'spritesheet' | 'individual'
}
// 生成响应 — 对应 POST /api/v1/generate 返回
export interface GenerateResponse {
taskId: string
assets: {
url: string
format: string
}[]
metadata: {
frameWidth: number
frameHeight: number
frameCount: number
directions: number
}
}
// 素材列表响应 — 对应 GET /api/v1/tasks/:taskId/assets
export interface AssetsResponse {
assets: {
url: string
format: string
}[]
metadata: {
frameWidth: number
frameHeight: number
frameCount: number
directions: number
} }
} }
-3
View File
@@ -5,11 +5,8 @@ export function useGenerate() {
return { return {
submit: store.submit, submit: store.submit,
taskId: store.taskId, taskId: store.taskId,
stage: store.stage,
progress: store.progress, progress: store.progress,
status: store.status, status: store.status,
retryCount: store.retryCount,
rejectReason: store.rejectReason,
assets: store.assets, assets: store.assets,
error: store.error, error: store.error,
reset: store.reset, reset: store.reset,
+55 -26
View File
@@ -1,27 +1,35 @@
import { useEffect } from 'react' import { useEffect } from 'react'
import { useNavigate, useParams } from 'react-router-dom' import { useNavigate, useParams } from 'react-router-dom'
import { useTaskStore } from '../stores/task' import { useTaskStore } from '../stores/task'
import { useProjectStore } from '../stores/project'
import { useGenerationStore } from '../stores/generation' import { useGenerationStore } from '../stores/generation'
import { useToastStore } from '../stores/toast' 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 GenerateForm from '../components/GenerateForm'
import ProgressBar from '../components/ProgressBar' import ProgressBar from '../components/ProgressBar'
export default function GeneratePage() { export default function GeneratePage() {
const { projectId = 'proj-default' } = useParams() const { projectId = 'proj-default' } = useParams()
const navigate = useNavigate() const navigate = useNavigate()
const { assetType, reset: resetTask } = useTaskStore()
const addToast = useToastStore(s => s.addToast) const addToast = useToastStore(s => s.addToast)
const taskStore = useTaskStore()
const { style: projectStyle, loadProject } = useProjectStore()
const { const {
status, status,
stage,
progress, progress,
retryCount,
rejectReason,
taskId, taskId,
submit, submit,
reset: resetGeneration, reset: resetGeneration,
} = useGenerationStore() } = useGenerationStore()
// 加载工程风格
useEffect(() => {
loadProject(projectId)
}, [projectId, loadProject])
// 组件卸载时重置生成状态 // 组件卸载时重置生成状态
useEffect(() => { useEffect(() => {
return () => resetGeneration() return () => resetGeneration()
@@ -30,48 +38,57 @@ export default function GeneratePage() {
// 失败时显示 toast // 失败时显示 toast
useEffect(() => { useEffect(() => {
if (status === 'failed') { if (status === 'failed') {
addToast({ type: 'error', message: '素材生成失败,请重试' }) const errText = useGenerationStore.getState().error || '未知错误'
addToast({ type: 'error', message: `素材生成失败:${errText}` })
} }
}, [status, addToast]) }, [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) => { 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 = () => { const handleReset = () => {
resetGeneration() resetGeneration()
resetTask() taskStore.reset()
} }
return ( return (
<div className="container page-enter" style={{ paddingTop: 24, paddingBottom: 40 }}> <div className="container page-enter" style={{ paddingTop: 24, paddingBottom: 40 }}>
<h1 style={{ fontSize: 24, marginBottom: 32 }}>新建生成</h1> <h1 style={{ fontSize: 24, marginBottom: 32 }}>新建生成</h1>
{/* 生成表单 */}
{status === 'idle' || status === 'submitting' ? ( {status === 'idle' || status === 'submitting' ? (
<GenerateForm onSubmit={handleSubmit} submitting={status === 'submitting'} /> <GenerateForm onSubmit={handleSubmit} submitting={status === 'submitting'} />
) : ( ) : (
<div style={{ display: 'flex', flexDirection: 'column', gap: 24 }}> <div style={{ display: 'flex', flexDirection: 'column', gap: 24 }}>
{/* 进度条 */}
<ProgressBar <ProgressBar
stage={stage} stage={null}
progress={progress} progress={progress}
status={status} status={status}
retryCount={retryCount} retryCount={0}
rejectReason={rejectReason} rejectReason={null}
/> />
{/* 状态提示 */}
{status === 'running' && ( {status === 'running' && (
<p style={{ textAlign: 'center', color: 'var(--text-secondary)' }}> <p style={{ textAlign: 'center', color: 'var(--text-secondary)' }}>
管线执行中,请稍候... 管线执行中,请稍候...
@@ -79,14 +96,26 @@ export default function GeneratePage() {
)} )}
{status === 'completed' && ( {status === 'completed' && (
<p style={{ textAlign: 'center', color: 'var(--success)' }}> <div style={{ textAlign: 'center' }}>
✓ 生成完成,正在跳转到结果页... <p style={{ color: 'var(--success)', marginBottom: 16 }}>
生成完成
</p> </p>
<div style={{ display: 'flex', gap: 12, justifyContent: 'center' }}>
<button className="btn-primary" onClick={handleViewResult}>
查看结果
</button>
<button className="btn-secondary" onClick={handleReset}>
继续生成
</button>
</div>
</div>
)} )}
{status === 'failed' && ( {status === 'failed' && (
<div style={{ textAlign: 'center' }}> <div style={{ textAlign: 'center' }}>
<p style={{ color: 'var(--error)', marginBottom: 16 }}>生成失败</p> <p style={{ color: 'var(--error)', marginBottom: 16 }}>
{useGenerationStore.getState().error || '生成失败'}
</p>
<button className="btn-primary" onClick={handleReset}> <button className="btn-primary" onClick={handleReset}>
重新开始 重新开始
</button> </button>
View File
+35 -43
View File
@@ -1,82 +1,74 @@
import { create } from 'zustand' import { create } from 'zustand'
import type { Asset, PipelineProgress, PipelineStage } from '../api/types' import type { Asset, GenerateRequest, GenerateResponse } from '../api/types'
import { submitGenerate } from '../api/generate' import { submitGenerate } from '../api/generate'
import { createMockWebSocket } from '../api/mock'
type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed' type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed'
interface GenerationState { interface GenerationState {
taskId: string | null taskId: string | null
stage: PipelineStage | null projectId: string | null
progress: number progress: number
status: Status status: Status
retryCount: number
rejectReason: string | null
assets: Asset[] assets: Asset[]
error: string | null error: string | null
submit: (projectId: string, prompt: string, assetType: string) => Promise<void> submit: (req: GenerateRequest) => Promise<void>
handleProgress: (msg: PipelineProgress) => void
reset: () => void reset: () => void
} }
let cleanupWs: (() => void) | null = null export const useGenerationStore = create<GenerationState>((set) => ({
export const useGenerationStore = create<GenerationState>((set, get) => ({
taskId: null, taskId: null,
stage: null, projectId: null,
progress: 0, progress: 0,
status: 'idle', status: 'idle',
retryCount: 0,
rejectReason: null,
assets: [], assets: [],
error: null, error: null,
submit: async (projectId, prompt, assetType) => { submit: async (req) => {
set({ status: 'submitting', error: null }) set({ status: 'submitting', error: null })
try { try {
const taskId = await submitGenerate(projectId, prompt, assetType) set({ status: 'running', progress: 30 })
set({ taskId, status: 'running', progress: 0 }) const result = await submitGenerate(req)
// 启动 mock WebSocket set({ progress: 80 })
cleanupWs = createMockWebSocket( const assets = mapAssets(result)
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 })
}
},
handleProgress: (msg) => {
set({ set({
stage: msg.stage, taskId: result.taskId,
progress: msg.progress, projectId: req.projectId,
retryCount: msg.retryCount ?? get().retryCount, status: 'completed',
rejectReason: msg.rejectReason ?? null, progress: 100,
assets,
}) })
if (msg.result?.assets) { } catch (err) {
set({ assets: msg.result.assets }) const errMsg = (err as Error).message
set({ status: 'failed', error: errMsg })
} }
}, },
reset: () => { reset: () => {
cleanupWs?.()
cleanupWs = null
set({ set({
taskId: null, taskId: null,
stage: null, projectId: null,
progress: 0, progress: 0,
status: 'idle', status: 'idle',
retryCount: 0,
rejectReason: null,
assets: [], assets: [],
error: null, error: null,
}) })
}, },
})) }))
function mapAssets(resp: GenerateResponse): Asset[] {
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,
},
}))
}