package service import ( "context" "fmt" "strings" "github.com/cloudwego/eino/compose" ) // PromptBuilder 节点:接收输入,输出三段式提示词 var promptBuilderNode = compose.InvokableLambda(func(ctx context.Context, in PipelineInput) (string, error) { return buildPrompt(in), nil }) // promptBuilderPreHandler 首次运行时保存输入到 state;重试时注入 RejectReason func promptBuilderPreHandler(ctx context.Context, in PipelineInput, state *PipelineState) (PipelineInput, error) { if state.RetryCount == 0 { state.Input = in } else if state.RejectReason != "" { in.RejectReason = state.RejectReason } return in, nil } // promptBuilderPostHandler 将提示词写入全局状态 func promptBuilderPostHandler(ctx context.Context, out string, state *PipelineState) (string, error) { state.FinalPrompt = out return out, nil } // AssetGenerator 节点:调用 AI 推理 API 出图 var assetGeneratorNode = compose.InvokableLambda(func(ctx context.Context, prompt string) ([]GeneratedImage, error) { var params AssetParams _ = compose.ProcessState[*PipelineState](ctx, func(_ context.Context, state *PipelineState) error { params = state.Input.Params return nil }) return GenerateImages(ctx, prompt, params) }) // assetGeneratorPostHandler 将原始图片写入全局状态 func assetGeneratorPostHandler(ctx context.Context, out []GeneratedImage, state *PipelineState) ([]GeneratedImage, error) { state.RawImages = out return out, nil } // QualitySupervisor 节点:质检,输出 PipelineInput 供下游节点消费。 // 将图片存入 state,设置路由目标 NextNode。 var qualitySupervisorNode = compose.InvokableLambda(func(ctx context.Context, images []GeneratedImage) (PipelineInput, error) { var input PipelineInput err := compose.ProcessState[*PipelineState](ctx, func(_ context.Context, state *PipelineState) error { // 保存图片到状态 state.RawImages = images // 质检 style := mergeStyle(state.Input.ProjectStyle, state.Input.TaskStyle) pass, reason, checkErr := CheckQuality(ctx, images, style) if checkErr != nil { return fmt.Errorf("quality check: %w", checkErr) } state.PassQuality = pass if !pass { state.RejectReason = reason } // 决定路由 if pass { state.NextNode = nodeFormatAdapter } else if state.RetryCount >= 3 { state.NextNode = nodeFormatAdapter // 超过重试次数,降级输出 } else { state.RetryCount++ state.NextNode = nodePromptBuilder // 重生成 } input = state.Input return nil }) if err != nil { return PipelineInput{}, err } return input, nil }) // formatAdapterNode 节点:从 state 读取图片,格式转换,组装输出 var formatAdapterNode = compose.InvokableLambda(func(ctx context.Context, input PipelineInput) (PipelineOutput, error) { var images []GeneratedImage _ = compose.ProcessState[*PipelineState](ctx, func(_ context.Context, state *PipelineState) error { images = state.RawImages return nil }) assets := make([]Asset, len(images)) for i, img := range images { assets[i] = Asset{ Data: img.Data, Format: img.Format, URL: fmt.Sprintf("output/%d.%s", i, img.Format), } } resolution := input.Params.Resolution if resolution <= 0 { resolution = 64 } metadata := AssetMetadata{ FrameWidth: resolution, FrameHeight: resolution, FrameCount: len(images), Directions: input.Params.Frames.Directions, } return PipelineOutput{ Assets: assets, Metadata: metadata, }, nil }) // buildPrompt 构建三段式提示词 func buildPrompt(in PipelineInput) string { var parts []string // 【主题】 parts = append(parts, fmt.Sprintf("【主题】%s", in.Prompt)) // 【约束】 constraints := buildConstraints(in) parts = append(parts, fmt.Sprintf("【约束】%s", constraints)) // 【内容】 content := buildContent(in) parts = append(parts, fmt.Sprintf("【内容】%s", content)) return strings.Join(parts, "\n") } // buildConstraints 合并风格 + 负面提示词 + 重试原因 func buildConstraints(in PipelineInput) string { style := mergeStyle(in.ProjectStyle, in.TaskStyle) var parts []string for k, v := range style { parts = append(parts, fmt.Sprintf("%s: %s", k, v)) } if in.RejectReason != "" { parts = append(parts, fmt.Sprintf("上次质检问题:%s", in.RejectReason)) } if len(parts) == 0 { return "无特殊约束" } return strings.Join(parts, "; ") } // buildContent 构建技术参数段 func buildContent(in PipelineInput) string { var parts []string parts = append(parts, fmt.Sprintf("素材类型: %s", in.AssetType)) if in.Params.Resolution > 0 { parts = append(parts, fmt.Sprintf("分辨率: %d", in.Params.Resolution)) } if in.Params.Frames.Directions > 0 { parts = append(parts, fmt.Sprintf("方向数: %d", in.Params.Frames.Directions)) } if in.Params.Frames.FramesPerDirection > 0 { parts = append(parts, fmt.Sprintf("每方向帧数: %d", in.Params.Frames.FramesPerDirection)) } if in.Params.Format != "" { parts = append(parts, fmt.Sprintf("输出格式: %s", in.Params.Format)) } return strings.Join(parts, "; ") } // mergeStyle 合并工程风格与任务风格覆盖,任务同名键覆盖工程 func mergeStyle(projectStyle, taskStyle map[string]string) map[string]string { result := make(map[string]string) for k, v := range projectStyle { result[k] = v } for k, v := range taskStyle { result[k] = v } return result }