From 322e6aecbbb59f61157843d6c61b9646e738d9ec Mon Sep 17 00:00:00 2001 From: Gmaker689 <1711322114@qq.com> Date: Mon, 25 May 2026 12:51:49 +0800 Subject: [PATCH] =?UTF-8?q?feat(generate):=20=E7=94=9F=E6=88=90=E7=AE=A1?= =?UTF-8?q?=E7=BA=BF=E6=94=B9=E4=B8=BA=E5=BC=82=E6=AD=A5=EF=BC=8C=E5=89=8D?= =?UTF-8?q?=E7=AB=AF=E6=8E=A5=E5=85=A5=E8=BD=AE=E8=AF=A2=E8=BF=9B=E5=BA=A6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 后端: - POST /api/v1/generate 改为异步模式,立即返回 taskId - 管线在后台 goroutine 执行,状态通过 GET /tasks/:taskId 轮询 - 任务状态流转: pending → running → saving → completed/failed 前端: - generation store: submit 后每 2s 轮询 getTask,完成时自动 getAssets - GeneratePage: 实时显示轮询状态文本 + 进度条 - ResultPage: 挂载时轮询任务,未完成显示骨架屏+进度,完成自动加载素材 - Types: Task.status 增加 submitted / saving 状态 --- backend/internal/handler/generate.go | 102 +++++++++++++++------ frontend/src/api/generate.ts | 7 -- frontend/src/api/types.ts | 14 +-- frontend/src/pages/GeneratePage.tsx | 3 +- frontend/src/pages/ResultPage.tsx | 127 ++++++++++++++++++--------- frontend/src/stores/generation.ts | 95 +++++++++++++------- 6 files changed, 230 insertions(+), 118 deletions(-) diff --git a/backend/internal/handler/generate.go b/backend/internal/handler/generate.go index ebf5ca7..4f54ddc 100755 --- a/backend/internal/handler/generate.go +++ b/backend/internal/handler/generate.go @@ -1,7 +1,9 @@ package handler import ( + "context" "fmt" + "log" "net/http" "os" "path/filepath" @@ -31,9 +33,7 @@ type GenerateRequest struct { // GenerateResponse 素材生成响应体。 type GenerateResponse struct { - TaskID string `json:"taskId"` - Assets []AssetResponse `json:"assets"` - Metadata service.AssetMetadata `json:"metadata"` + TaskID string `json:"taskId"` } // AssetResponse 单个素材响应。 @@ -44,21 +44,21 @@ type AssetResponse struct { // 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"` + 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"` + Assets []AssetResponse `json:"assets"` + Metadata service.AssetMetadata `json:"metadata"` } // taskRecord 内存中的任务记录。 @@ -72,7 +72,8 @@ var ( taskStore = sync.Map{} // taskID → *taskRecord ) -// Generate 素材生成接口。 +// Generate 素材生成接口(异步)。 +// 立即返回 taskId,后台执行管线,前端通过 GET /tasks/:taskId 轮询进度。 func Generate(c *gin.Context) { var req GenerateRequest if err := c.ShouldBindJSON(&req); err != nil { @@ -85,6 +86,31 @@ func Generate(c *gin.Context) { 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, @@ -105,16 +131,20 @@ func Generate(c *gin.Context) { }, } - output, err := service.RunPipeline(c.Request.Context(), in) + ctx := context.Background() + output, err := service.RunPipeline(ctx, in) if err != nil { - c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "素材生成失败: "+err.Error())) + 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 { - c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "创建输出目录失败: "+err.Error())) + updateFailed(taskID, "创建输出目录失败: "+err.Error()) return } @@ -123,7 +153,7 @@ func Generate(c *gin.Context) { 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())) + updateFailed(taskID, "保存图片失败: "+err.Error()) return } assets[i] = AssetResponse{ @@ -132,7 +162,7 @@ func Generate(c *gin.Context) { } } - // 存储任务记录到内存 + // 更新为完成状态 taskStore.Store(taskID, &taskRecord{ task: TaskResponse{ ID: taskID, @@ -147,11 +177,29 @@ func Generate(c *gin.Context) { metadata: output.Metadata, }) - c.JSON(http.StatusOK, model.OK(GenerateResponse{ - TaskID: taskID, - 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 查询任务信息。 @@ -175,6 +223,10 @@ func GetAssets(c *gin.Context) { 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/frontend/src/api/generate.ts b/frontend/src/api/generate.ts index 971d9b2..a03b2c6 100755 --- a/frontend/src/api/generate.ts +++ b/frontend/src/api/generate.ts @@ -19,13 +19,6 @@ export async function getTask(taskId: string): Promise { export async function getAssets(taskId: string): Promise { const resp = await get(`/api/v1/tasks/${taskId}/assets`) - return toAssetList(resp) -} - -/** 将响应转为 Asset[] 供前端组件使用 */ -function toAssetList( - resp: AssetsResponse | GenerateResponse, -): Asset[] { return resp.assets.map((a, i) => ({ id: `asset-${i}`, url: a.url, diff --git a/frontend/src/api/types.ts b/frontend/src/api/types.ts index 5c5f8a9..3fbdefd 100755 --- 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 @@ -106,19 +106,9 @@ export interface GenerateRequest { format?: 'spritesheet' | 'individual' } -// 生成响应 — 对应 POST /api/v1/generate 返回 +// 生成响应 — 对应 POST /api/v1/generate 返回(异步,仅含 taskId) 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 diff --git a/frontend/src/pages/GeneratePage.tsx b/frontend/src/pages/GeneratePage.tsx index d21356c..1912994 100755 --- a/frontend/src/pages/GeneratePage.tsx +++ b/frontend/src/pages/GeneratePage.tsx @@ -21,6 +21,7 @@ export default function GeneratePage() { status, progress, taskId, + statusText, submit, reset: resetGeneration, } = useGenerationStore() @@ -91,7 +92,7 @@ export default function GeneratePage() { {status === 'running' && (

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

)} diff --git a/frontend/src/pages/ResultPage.tsx b/frontend/src/pages/ResultPage.tsx index 9031397..85da8b3 100755 --- 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 index bb3ec9c..d7346d1 100755 --- a/frontend/src/stores/generation.ts +++ b/frontend/src/stores/generation.ts @@ -1,6 +1,6 @@ import { create } from 'zustand' -import type { Asset, GenerateRequest, GenerateResponse } from '../api/types' -import { submitGenerate } from '../api/generate' +import type { Asset, GenerateRequest } from '../api/types' +import { submitGenerate, getTask, getAssets } from '../api/generate' type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed' @@ -9,66 +9,99 @@ interface GenerationState { projectId: string | null progress: number status: Status + statusText: string assets: Asset[] error: string | null submit: (req: GenerateRequest) => Promise reset: () => void } -export const useGenerationStore = create((set) => ({ +let pollTimer: ReturnType | null = null + +function stopPolling() { + if (pollTimer) { + clearInterval(pollTimer) + pollTimer = null + } +} + +export const useGenerationStore = create((set, get) => ({ taskId: null, projectId: null, progress: 0, status: 'idle', + statusText: '', assets: [], error: null, submit: async (req) => { - set({ status: 'submitting', error: null }) + stopPolling() + set({ status: 'submitting', error: null, statusText: '提交中...' }) try { - set({ status: 'running', progress: 30 }) - const result = await submitGenerate(req) - - set({ progress: 80 }) - const assets = mapAssets(result) + const { taskId } = await submitGenerate(req) set({ - taskId: result.taskId, + taskId, projectId: req.projectId, - status: 'completed', - progress: 100, - assets, + status: 'running', + progress: 10, + statusText: '任务已提交,等待生成...', }) + + // 开始轮询进度 + 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) { - const errMsg = (err as Error).message - set({ status: 'failed', error: errMsg }) + stopPolling() + set({ status: 'failed', error: (err as Error).message, statusText: '提交失败' }) } }, reset: () => { + stopPolling() set({ taskId: null, projectId: null, progress: 0, status: 'idle', + statusText: '', assets: [], 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, - }, - })) -}