From 5d5fe60426171a14256e6c3a63169de20c701b4a Mon Sep 17 00:00:00 2001 From: Gmaker689 <1711322114@qq.com> Date: Mon, 25 May 2026 14:26:26 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E7=AE=A1=E7=BA=BF=E5=9B=9B=E9=98=B6?= =?UTF-8?q?=E6=AE=B5=E8=BF=9B=E5=BA=A6=E4=B8=8A=E6=8A=A5=20+=20=E5=9B=BE?= =?UTF-8?q?=E7=89=87=E9=A2=84=E8=A7=88=20+=20Vite=20=E4=BB=A3=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 后端: - pipeline.go: ProgressReporter 回调类型 + WithProgressReporter 注入 context - nodes.go: 各节点 pre/post handler 调用 reportProgress() 上报阶段进度 - generate.go: TaskResponse 新增 stage 字段,runPipelineBg 注入进度回调 前端: - vite.config.ts: 添加 /generation 代理到后端静态文件服务 - generation.ts: 轮询读取 stage 字段,暴露 stage/retryCount/rejectReason - GeneratePage.tsx: ProgressBar 接收真实管线阶段数据 --- backend/internal/handler/generate.go | 14 +++++++--- backend/internal/service/nodes.go | 7 +++++ backend/internal/service/pipeline.go | 19 +++++++++++++ frontend/src/pages/GeneratePage.tsx | 9 ++++-- frontend/src/stores/generation.ts | 42 +++++++++++++++++++++------- frontend/vite.config.ts | 4 +++ 6 files changed, 78 insertions(+), 17 deletions(-) diff --git a/backend/internal/handler/generate.go b/backend/internal/handler/generate.go index 4f54ddc..c3ee063 100755 --- a/backend/internal/handler/generate.go +++ b/backend/internal/handler/generate.go @@ -49,6 +49,7 @@ type TaskResponse struct { Prompt string `json:"prompt"` AssetType string `json:"assetType"` Status string `json:"status"` + Stage string `json:"stage,omitempty"` Progress int `json:"progress"` RetryCount int `json:"retryCount"` Error string `json:"error,omitempty"` @@ -110,7 +111,12 @@ func Generate(c *gin.Context) { // runPipelineBg 后台执行生成管线,更新任务状态。 func runPipelineBg(projectID, taskID string, req GenerateRequest) { - updateStatus(taskID, "running", 10) + // 注入进度上报回调 + ctx := service.WithProgressReporter(context.Background(), func(stage string, progress int) { + updateTaskProgress(taskID, "running", stage, progress) + }) + + updateTaskProgress(taskID, "running", "prompt_builder", 5) in := service.PipelineInput{ ProjectID: projectID, @@ -131,7 +137,6 @@ func runPipelineBg(projectID, taskID string, req GenerateRequest) { }, } - ctx := context.Background() output, err := service.RunPipeline(ctx, in) if err != nil { log.Printf("[generate] task %s failed: %v", taskID, err) @@ -139,7 +144,7 @@ func runPipelineBg(projectID, taskID string, req GenerateRequest) { return } - updateStatus(taskID, "saving", 80) + updateTaskProgress(taskID, "saving", "format_adapter", 90) // 保存图片到 ../generation/{projectId}/{taskId}/ genDir := filepath.Join("..", "generation", projectID, taskID) @@ -180,13 +185,14 @@ func runPipelineBg(projectID, taskID string, req GenerateRequest) { log.Printf("[generate] task %s completed, %d assets", taskID, len(assets)) } -func updateStatus(taskID, status string, progress int) { +func updateTaskProgress(taskID, status, stage string, progress int) { rec, ok := taskStore.Load(taskID) if !ok { return } r := rec.(*taskRecord) r.task.Status = status + r.task.Stage = stage r.task.Progress = progress taskStore.Store(taskID, r) } diff --git a/backend/internal/service/nodes.go b/backend/internal/service/nodes.go index 1bac3ed..c64156d 100755 --- a/backend/internal/service/nodes.go +++ b/backend/internal/service/nodes.go @@ -57,8 +57,10 @@ var promptOptimizerNode = compose.InvokableLambda(func(ctx context.Context, in P func promptOptimizerPreHandler(ctx context.Context, in PipelineInput, state *PipelineState) (PipelineInput, error) { if state.RetryCount == 0 { state.Input = in + reportProgress(ctx, "prompt_builder", 10) } else if state.RejectReason != "" { in.RejectReason = state.RejectReason + reportProgress(ctx, "prompt_builder", 30+state.RetryCount*10) } return in, nil } @@ -66,6 +68,7 @@ func promptOptimizerPreHandler(ctx context.Context, in PipelineInput, state *Pip // promptOptimizerPostHandler 将最终提示词写入全局状态。 func promptOptimizerPostHandler(ctx context.Context, out string, state *PipelineState) (string, error) { state.FinalPrompt = out + reportProgress(ctx, "asset_generator", 35) return out, nil } @@ -82,6 +85,7 @@ var assetGeneratorNode = compose.InvokableLambda(func(ctx context.Context, promp // assetGeneratorPostHandler 将原始图片写入全局状态。 func assetGeneratorPostHandler(ctx context.Context, out []GeneratedImage, state *PipelineState) ([]GeneratedImage, error) { state.RawImages = out + reportProgress(ctx, "quality_supervisor", 60) return out, nil } @@ -104,11 +108,14 @@ var qualitySupervisorNode = compose.InvokableLambda(func(ctx context.Context, im if pass { state.NextNode = nodeFormatAdapter + reportProgress(ctx, "format_adapter", 85) } else if state.RetryCount >= 3 { state.NextNode = nodeFormatAdapter + reportProgress(ctx, "format_adapter", 85) } else { state.RetryCount++ state.NextNode = nodePromptOptimizer + reportProgress(ctx, "quality_supervisor", 50+state.RetryCount*10) } input = state.Input diff --git a/backend/internal/service/pipeline.go b/backend/internal/service/pipeline.go index 6ce8a23..f37823e 100755 --- a/backend/internal/service/pipeline.go +++ b/backend/internal/service/pipeline.go @@ -14,6 +14,25 @@ const ( nodeFormatAdapter = "format_adapter" ) +// ProgressReporter 管线进度回调:stage 为当前节点名,progress 为 0-100。 +type ProgressReporter func(stage string, progress int) + +type progressKeyType struct{} + +var progressCtxKey progressKeyType + +// WithProgressReporter 将进度回调注入 context。 +func WithProgressReporter(ctx context.Context, r ProgressReporter) context.Context { + return context.WithValue(ctx, progressCtxKey, r) +} + +// reportProgress 从 context 取出回调上报进度。 +func reportProgress(ctx context.Context, stage string, progress int) { + if r, ok := ctx.Value(progressCtxKey).(ProgressReporter); ok { + r(stage, progress) + } +} + // NewGenerateGraph 创建生成管线 Graph(PromptOptimizer → AssetGenerator → QualitySupervisor → FormatAdapter)。 // // START → PromptOptimizer → AssetGenerator → QualitySupervisor diff --git a/frontend/src/pages/GeneratePage.tsx b/frontend/src/pages/GeneratePage.tsx index 1912994..a5581e6 100755 --- a/frontend/src/pages/GeneratePage.tsx +++ b/frontend/src/pages/GeneratePage.tsx @@ -19,9 +19,12 @@ export default function GeneratePage() { const { style: projectStyle, loadProject } = useProjectStore() const { status, + stage, progress, taskId, statusText, + retryCount, + rejectReason, submit, reset: resetGeneration, } = useGenerationStore() @@ -83,11 +86,11 @@ export default function GeneratePage() { ) : (
{status === 'running' && ( diff --git a/frontend/src/stores/generation.ts b/frontend/src/stores/generation.ts index d7346d1..6686a2f 100755 --- a/frontend/src/stores/generation.ts +++ b/frontend/src/stores/generation.ts @@ -1,5 +1,5 @@ import { create } from 'zustand' -import type { Asset, GenerateRequest } from '../api/types' +import type { Asset, GenerateRequest, PipelineStage } from '../api/types' import { submitGenerate, getTask, getAssets } from '../api/generate' type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed' @@ -7,9 +7,12 @@ type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed' interface GenerationState { taskId: string | null projectId: string | null + stage: PipelineStage | null progress: number status: Status statusText: string + retryCount: number + rejectReason: string | null assets: Asset[] error: string | null submit: (req: GenerateRequest) => Promise @@ -25,18 +28,21 @@ function stopPolling() { } } -export const useGenerationStore = create((set, get) => ({ +export const useGenerationStore = create((set) => ({ taskId: null, projectId: null, + stage: null, progress: 0, status: 'idle', statusText: '', + retryCount: 0, + rejectReason: null, assets: [], error: null, submit: async (req) => { stopPolling() - set({ status: 'submitting', error: null, statusText: '提交中...' }) + set({ status: 'submitting', error: null, statusText: '提交中...', stage: null }) try { const { taskId } = await submitGenerate(req) @@ -44,24 +50,26 @@ export const useGenerationStore = create((set, get) => ({ taskId, projectId: req.projectId, status: 'running', - progress: 10, + progress: 5, + stage: 'prompt_builder', statusText: '任务已提交,等待生成...', }) - // 开始轮询进度 pollTimer = setInterval(async () => { try { const task = await getTask(taskId) - set({ - progress: task.progress ?? get().progress, + set((s) => ({ + stage: task.stage ?? s.stage, + progress: task.progress ?? s.progress, + retryCount: task.retryCount ?? s.retryCount, statusText: task.status === 'running' - ? '生成中...' + ? stageLabel(task.stage) : task.status === 'pending' ? '排队中...' : task.status, - }) + })) if (task.status === 'completed') { stopPolling() @@ -71,6 +79,7 @@ export const useGenerationStore = create((set, get) => ({ set({ status: 'completed', progress: 100, + stage: 'format_adapter', statusText: '生成完成', assets, }) @@ -85,7 +94,7 @@ export const useGenerationStore = create((set, get) => ({ } catch { // 网络错误不中断轮询 } - }, 2000) + }, 1500) } catch (err) { stopPolling() set({ status: 'failed', error: (err as Error).message, statusText: '提交失败' }) @@ -97,11 +106,24 @@ export const useGenerationStore = create((set, get) => ({ set({ taskId: null, projectId: null, + stage: null, progress: 0, status: 'idle', statusText: '', + retryCount: 0, + rejectReason: null, assets: [], error: null, }) }, })) + +function stageLabel(stage?: string): string { + switch (stage) { + case 'prompt_builder': return '优化提示词...' + case 'asset_generator': return '生成素材中...' + case 'quality_supervisor': return '质检中...' + case 'format_adapter': return '格式转换中...' + default: return '生成中...' + } +} diff --git a/frontend/vite.config.ts b/frontend/vite.config.ts index 7f74e10..2f4f4e8 100755 --- a/frontend/vite.config.ts +++ b/frontend/vite.config.ts @@ -14,6 +14,10 @@ export default defineConfig({ target: 'http://localhost:8080', changeOrigin: true, }, + '/generation': { + target: 'http://localhost:8080', + changeOrigin: true, + }, }, }, })