Skip to main content

lumen_engine/node/processing/
curves.rs

1use crate::node::{NodeId, NodeProperty, PortRef};
2
3use crate::gpu::{
4    BoundFrame, CompiledOutput, FrameBindContext, FrameBinding, GpuCompileNode, GpuFrameBindNode,
5    RasterHandle, compiler,
6};
7
8pub(crate) const SHADER: &str = include_str!("curves.wgsl");
9
10/// Applies a 1D RGB curve table to a raster.
11#[derive(Debug, Clone, lumen_macros::Node)]
12#[node(kind = "curves", name = "Curves", category = "processing")]
13pub struct Curves {
14    pub id: NodeId,
15    /// Curve table data source or named curve preset.
16    #[property(
17        kind = "string",
18        name = "Curve",
19        role = "curve_source",
20        multiline,
21        recommended_rows = 4
22    )]
23    pub curve_source: NodeProperty,
24    /// Blend amount for the curve adjustment.
25    #[property(kind = "float", min = 0, max = 1, step = 0.01)]
26    pub strength: NodeProperty,
27    #[input()]
28    pub source: PortRef,
29}
30
31impl Default for Curves {
32    fn default() -> Self {
33        Self {
34            id: NodeId::new(0),
35            curve_source: NodeProperty::String("identity".to_string()),
36            strength: NodeProperty::Float(1.0),
37            source: PortRef::empty(),
38        }
39    }
40}
41
42impl GpuCompileNode for Curves {
43    fn compile_gpu(
44        &self,
45        ctx: &mut crate::gpu::CompileContext<'_>,
46        port: &PortRef,
47    ) -> crate::Result<CompiledOutput> {
48        if port.port != "output" {
49            return Err(ctx.missing_output(self.id, &port.port));
50        }
51
52        let source = ctx
53            .compile_port(&self.source)?
54            .into_raster(self.source.id, &self.source.port)?;
55        let size = source.domain.storage_size;
56        let texture = ctx.builder_mut().texture_for(
57            lumen_gpu::NodeKey(self.id.0),
58            Some(format!("curves:{}:output", self.id.0)),
59            lumen_gpu::TextureDesc::storage(size, lumen_gpu::wgpu::TextureFormat::Rgba8Unorm),
60        );
61        let params = ctx.builder_mut().buffer_for(
62            lumen_gpu::NodeKey(self.id.0),
63            Some(format!("curves:{}:params", self.id.0)),
64            lumen_gpu::BufferDesc::uniform(std::mem::size_of::<compiler::CurvesParams>() as u64),
65        );
66        let curve = ctx.builder_mut().buffer_for(
67            lumen_gpu::NodeKey(self.id.0),
68            Some(format!("curves:{}:table", self.id.0)),
69            lumen_gpu::BufferDesc::storage(std::mem::size_of::<compiler::CurvesTable>() as u64),
70        );
71        let program = ctx.builder_mut().program_for(
72            lumen_gpu::NodeKey(self.id.0),
73            lumen_gpu::ProgramDesc::Compute(lumen_gpu::ComputeProgramDesc {
74                label: Some("curves".to_string()),
75                shader: SHADER.to_string(),
76                entry: "cs_main".to_string(),
77                bind_groups: lumen_gpu::BindGroupLayoutSpec::single(vec![
78                    lumen_gpu::BindingLayoutEntry::texture(
79                        0,
80                        lumen_gpu::wgpu::ShaderStages::COMPUTE,
81                    ),
82                    lumen_gpu::BindingLayoutEntry::uniform(
83                        1,
84                        lumen_gpu::wgpu::ShaderStages::COMPUTE,
85                    ),
86                    lumen_gpu::BindingLayoutEntry::storage(
87                        2,
88                        lumen_gpu::wgpu::ShaderStages::COMPUTE,
89                        true,
90                    ),
91                    lumen_gpu::BindingLayoutEntry::storage_texture(
92                        3,
93                        lumen_gpu::wgpu::ShaderStages::COMPUTE,
94                        lumen_gpu::wgpu::TextureFormat::Rgba8Unorm,
95                        lumen_gpu::wgpu::StorageTextureAccess::WriteOnly,
96                    ),
97                ]),
98            }),
99        );
100        ctx.builder_mut().compute_pass(lumen_gpu::ComputePassDesc {
101            label: Some(format!("curves:{}:apply", self.id.0)),
102            owner: Some(lumen_gpu::NodeKey(self.id.0)),
103            program,
104            bindings: vec![
105                lumen_gpu::Binding::sampled_texture(0, 0, source.texture),
106                lumen_gpu::Binding::uniform(0, 1, params),
107                lumen_gpu::Binding::storage_buffer(0, 2, curve),
108                lumen_gpu::Binding::storage_texture(0, 3, texture),
109            ],
110            dispatch: compiler::dispatch_for(size).into(),
111        });
112        ctx.builder_mut().param(
113            lumen_gpu::ParamKey {
114                owner: lumen_gpu::NodeKey(self.id.0),
115                slot: 0,
116            },
117            lumen_gpu::ParamTarget::Buffer(params),
118        );
119        ctx.builder_mut().param(
120            lumen_gpu::ParamKey {
121                owner: lumen_gpu::NodeKey(self.id.0),
122                slot: 1,
123            },
124            lumen_gpu::ParamTarget::Buffer(curve),
125        );
126        ctx.push_frame_binding(FrameBinding::Curves {
127            node_id: self.id,
128            curve_source: self.curve_source.clone(),
129            strength: self.strength.clone(),
130            params_buffer: params,
131            curve_buffer: curve,
132        });
133
134        Ok(CompiledOutput::Raster(RasterHandle {
135            texture,
136            domain: source.domain,
137            metadata: source.metadata,
138        }))
139    }
140}
141
142impl GpuFrameBindNode for Curves {
143    fn bind_gpu_frame(
144        &self,
145        ctx: &FrameBindContext<'_>,
146        binding: &FrameBinding,
147        bound: &mut BoundFrame,
148    ) -> crate::Result<()> {
149        let FrameBinding::Curves {
150            node_id,
151            curve_source,
152            strength,
153            params_buffer,
154            curve_buffer,
155        } = binding
156        else {
157            return Ok(());
158        };
159        let curve_source = curve_source.resolve_string(
160            *node_id,
161            "curve_source",
162            &ctx.expr_context(*node_id, "curve_source"),
163        )?;
164        let params = compiler::CurvesParams {
165            values: [
166                strength.resolve_float(
167                    *node_id,
168                    "strength",
169                    &ctx.expr_context(*node_id, "strength"),
170                )? as f32,
171                0.0,
172                0.0,
173                0.0,
174            ],
175        };
176        let curve = compiler::CurvesTable::parse(*node_id, ctx.frame(), &curve_source)?;
177        bound.write_buffer(*params_buffer, 0, bytemuck::bytes_of(&params));
178        bound.write_buffer(*curve_buffer, 0, bytemuck::bytes_of(&curve));
179        Ok(())
180    }
181}