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() { ) : (