Skip to main content

demo_shader_editor/
lib.rs

1#![doc = include_str!("../README.md")]
2#![doc = r#"<link rel="stylesheet" href="../gallery/pkg/demo.css"><script type="module" src="../gallery/pkg/demo-loader.js"></script>"#]
3
4mod colors;
5mod compiler;
6mod default_shader;
7mod shader_graph;
8
9use compiler::ShaderCompiler;
10use demo_common::Demo;
11use iced::{
12    Element, Event, Length, Point, Subscription, Task, Theme, Vector, event, keyboard,
13    widget::{column, container, opaque, stack, text},
14    window,
15};
16use iced_nodegraph::{
17    Ids, PinDirection, PinInfo, PinRef, PinSide, PinStatus, PinStyle, default_pin_style,
18    edge as ng_edge, node as ng_node, node_pin,
19};
20use iced_palette::{
21    Command, command, command_palette, focus_input, get_filtered_command_index, get_filtered_count,
22    navigate_down, navigate_up,
23};
24use shader_graph::ShaderGraph;
25use shader_graph::nodes::ShaderNodeType;
26use std::collections::HashSet;
27
28pub fn main() -> iced::Result {
29    let window_settings = iced::window::Settings::default();
30
31    iced::application(Application::boot, Application::update, Application::view)
32        .subscription(Application::subscription)
33        .title("Visual Shader Editor - iced_nodegraph")
34        .theme(Application::theme)
35        .window(window_settings)
36        .run()
37}
38
39/// The id vocabulary of this demo: indexed nodes and pins, unidentified edges,
40/// and a `TypeId` pin payload carrying the socket type of each pin.
41#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
42struct TypedIds;
43
44impl Ids for TypedIds {
45    type NodeId = usize;
46    type PinId = usize;
47    type EdgeId = ();
48    type AnchorId = usize;
49    type Payload = std::any::TypeId;
50}
51
52#[derive(Debug, Clone)]
53enum Message {
54    EdgeConnected {
55        from: PinRef<TypedIds>,
56        to: PinRef<TypedIds>,
57    },
58    EdgeDisconnected {
59        from: PinRef<TypedIds>,
60        to: PinRef<TypedIds>,
61    },
62    SelectionChanged(Vec<usize>),
63    NodesMoved {
64        delta: Vector,
65        indices: Vec<usize>,
66    },
67    // Command palette messages
68    ToggleCommandPalette,
69    CommandPaletteInput(String),
70    CommandPaletteNavigateUp,
71    CommandPaletteNavigateDown,
72    CommandPaletteNavigate(usize),
73    CommandPaletteSelect(usize),
74    CommandPaletteConfirm,
75    CommandPaletteCancel,
76    // Node spawning
77    SpawnNode(ShaderNodeType),
78    // Theme
79    ChangeTheme(Theme),
80    // Camera/viewport tracking
81    CameraChanged {
82        position: Point,
83        zoom: f32,
84    },
85    WindowResized(iced::Size),
86}
87
88struct Application {
89    shader_graph: ShaderGraph,
90    compiled_shader: Option<String>,
91    compilation_error: Option<String>,
92    visual_edges: Vec<(PinRef<TypedIds>, PinRef<TypedIds>)>,
93    current_theme: Theme,
94    graph_selection: HashSet<usize>,
95    // Command palette state
96    command_palette_open: bool,
97    command_input: String,
98    palette_selected_index: usize,
99    // Camera/viewport tracking for spawn-at-center
100    viewport_size: iced::Size,
101    camera_position: Point,
102    camera_zoom: f32,
103}
104
105impl demo_common::Demo for Application {
106    type Message = Message;
107
108    fn boot() -> (Self, iced::Task<Message>) {
109        let shader_graph = default_shader::create_default_graph();
110
111        // Convert shader graph connections to visual edges
112        // NodeGraph widget uses flat pin indices: [input0, input1, ..., output0, output1, ...]
113        // ShaderGraph uses separate indices: from_socket = output index, to_socket = input index
114        let visual_edges: Vec<(PinRef<TypedIds>, PinRef<TypedIds>)> = shader_graph
115            .connections
116            .iter()
117            .filter_map(|conn| {
118                // Get the nodes to find their input/output counts
119                let from_node = shader_graph.nodes.iter().find(|n| n.id == conn.from_node)?;
120                // Validate target node exists
121                shader_graph.nodes.iter().find(|n| n.id == conn.to_node)?;
122
123                // Find node indices (position in nodes vec)
124                let from_node_idx = shader_graph
125                    .nodes
126                    .iter()
127                    .position(|n| n.id == conn.from_node)?;
128                let to_node_idx = shader_graph
129                    .nodes
130                    .iter()
131                    .position(|n| n.id == conn.to_node)?;
132
133                // from_socket is an output index -> visual pin = num_inputs + output_index
134                let from_visual_pin = from_node.inputs.len() + conn.from_socket;
135
136                // to_socket is an input index -> visual pin = input_index (inputs come first)
137                let to_visual_pin = conn.to_socket;
138
139                Some((
140                    PinRef::new(from_node_idx, from_visual_pin),
141                    PinRef::new(to_node_idx, to_visual_pin),
142                ))
143            })
144            .collect();
145
146        let mut app = Self {
147            shader_graph,
148            compiled_shader: None,
149            compilation_error: None,
150            visual_edges,
151            current_theme: Theme::CatppuccinMocha,
152            graph_selection: HashSet::new(),
153            command_palette_open: false,
154            command_input: String::new(),
155            palette_selected_index: 0,
156            viewport_size: iced::Size::new(800.0, 600.0),
157            camera_position: Point::ORIGIN,
158            camera_zoom: 1.0,
159        };
160
161        app.recompile();
162
163        (app, iced::Task::none())
164    }
165
166    fn update(&mut self, message: Message) -> iced::Task<Message> {
167        match message {
168            Message::EdgeConnected { from, to } => {
169                // Store visual edge as-is
170                self.visual_edges.push((from, to));
171
172                // Convert visual pin indices to shader socket indices
173                // First, gather the info we need
174                let connection_info = {
175                    let from_node_data = self.shader_graph.nodes.get(from.node_id);
176                    let to_node_data = self.shader_graph.nodes.get(to.node_id);
177
178                    if let (Some(from_node_data), Some(to_node_data)) =
179                        (from_node_data, to_node_data)
180                    {
181                        // from.pin_id is visual index, output starts after inputs
182                        let from_socket = from.pin_id.saturating_sub(from_node_data.inputs.len());
183                        // to.pin_id is visual index, inputs come first so it's direct
184                        let to_socket = to.pin_id;
185
186                        if from_socket < from_node_data.outputs.len()
187                            && to_socket < to_node_data.inputs.len()
188                        {
189                            Some((from_node_data.id, from_socket, to_node_data.id, to_socket))
190                        } else {
191                            None
192                        }
193                    } else {
194                        None
195                    }
196                };
197
198                // Now apply the connection
199                if let Some((from_id, from_socket, to_id, to_socket)) = connection_info {
200                    self.shader_graph.add_connection(shader_graph::Connection {
201                        from_node: from_id,
202                        from_socket,
203                        to_node: to_id,
204                        to_socket,
205                    });
206                }
207                self.recompile();
208            }
209            Message::EdgeDisconnected { from, to } => {
210                self.visual_edges.retain(|(f, t)| !(f == &from && t == &to));
211
212                // Convert visual pin indices to shader socket indices and
213                // resolve graph indices to node ids, mirroring EdgeConnected.
214                let from_node = self.shader_graph.nodes.get(from.node_id);
215                let to_node = self.shader_graph.nodes.get(to.node_id);
216                if let (Some(from_node), Some(to_node)) = (from_node, to_node) {
217                    let from_id = from_node.id;
218                    let to_id = to_node.id;
219                    let from_socket = from.pin_id.saturating_sub(from_node.inputs.len());
220                    let to_socket = to.pin_id;
221
222                    self.shader_graph.connections.retain(|c| {
223                        !(c.from_node == from_id
224                            && c.from_socket == from_socket
225                            && c.to_node == to_id
226                            && c.to_socket == to_socket)
227                    });
228                }
229                self.recompile();
230                return Task::none();
231            }
232            Message::SelectionChanged(indices) => {
233                self.graph_selection = indices.into_iter().collect();
234            }
235            Message::NodesMoved { delta, indices } => {
236                for idx in indices {
237                    if let Some(node) = self.shader_graph.get_node_by_index_mut(idx) {
238                        node.position.x += delta.x;
239                        node.position.y += delta.y;
240                    }
241                }
242            }
243            // Command palette
244            Message::ToggleCommandPalette => {
245                self.command_palette_open = !self.command_palette_open;
246                if self.command_palette_open {
247                    self.command_input.clear();
248                    self.palette_selected_index = 0;
249                    return focus_input();
250                }
251            }
252            Message::CommandPaletteInput(input) => {
253                self.command_input = input;
254                self.palette_selected_index = 0;
255            }
256            Message::CommandPaletteNavigateUp => {
257                if !self.command_palette_open {
258                    return Task::none();
259                }
260                let commands = self.build_palette_commands();
261                let count = get_filtered_count(&self.command_input, &commands);
262                self.palette_selected_index = navigate_up(self.palette_selected_index, count);
263            }
264            Message::CommandPaletteNavigateDown => {
265                if !self.command_palette_open {
266                    return Task::none();
267                }
268                let commands = self.build_palette_commands();
269                let count = get_filtered_count(&self.command_input, &commands);
270                self.palette_selected_index = navigate_down(self.palette_selected_index, count);
271            }
272            Message::CommandPaletteNavigate(index) => {
273                if !self.command_palette_open {
274                    return Task::none();
275                }
276                self.palette_selected_index = index;
277            }
278            Message::CommandPaletteSelect(index) => {
279                self.palette_selected_index = index;
280                return self.update(Message::CommandPaletteConfirm);
281            }
282            Message::CommandPaletteConfirm => {
283                if !self.command_palette_open {
284                    return Task::none();
285                }
286                let commands = self.build_palette_commands();
287                if let Some(original_idx) = get_filtered_command_index(
288                    &self.command_input,
289                    &commands,
290                    self.palette_selected_index,
291                ) {
292                    use iced_palette::CommandAction;
293                    if let CommandAction::Message(msg) = &commands[original_idx].action {
294                        let msg = msg.clone();
295                        self.command_palette_open = false;
296                        self.command_input.clear();
297                        self.palette_selected_index = 0;
298                        return self.update(msg);
299                    }
300                }
301            }
302            Message::CommandPaletteCancel => {
303                self.command_palette_open = false;
304                self.command_input.clear();
305                self.palette_selected_index = 0;
306            }
307            Message::SpawnNode(node_type) => {
308                // Spawn node at screen center (converted to world coordinates)
309                let position = self.spawn_position();
310                self.shader_graph.add_node(node_type, position);
311            }
312            Message::ChangeTheme(theme) => {
313                self.current_theme = theme;
314            }
315            Message::CameraChanged { position, zoom } => {
316                self.camera_position = position;
317                self.camera_zoom = zoom;
318            }
319            Message::WindowResized(size) => {
320                self.viewport_size = size;
321            }
322        }
323
324        Task::none()
325    }
326
327    fn view(&self) -> Element<'_, Message> {
328        let graph = ::iced_nodegraph::NodeGraph::<TypedIds, _, _, _>::new()
329            .on_connect(|from, to| Message::EdgeConnected { from, to })
330            .on_move(|delta, indices| Message::NodesMoved { delta, indices })
331            .on_disconnect(|from, to| Message::EdgeDisconnected { from, to })
332            .on_select(Message::SelectionChanged)
333            .on_camera(|position, zoom| Message::CameraChanged { position, zoom })
334            .nodes(
335                self.shader_graph
336                    .nodes
337                    .iter()
338                    .enumerate()
339                    .map(|(node_idx, node)| {
340                        let node_content = create_node_widget(&node.node_type, &self.current_theme);
341                        ng_node(node_idx, node.position, node_content)
342                            .selected(self.graph_selection.contains(&node_idx))
343                            .pin_style(pin_style)
344                    }),
345            )
346            .edges(
347                self.visual_edges
348                    .iter()
349                    .map(|(from, to)| ng_edge((), *from, *to)),
350            );
351
352        let graph_element: Element<Message> = graph.into();
353
354        // Show command palette overlay if open. `opaque` blocks wheel events
355        // from reaching the NodeGraph behind it; the palette's internal
356        // `mouse_area` only captures `on_press`, not scroll.
357        if self.command_palette_open {
358            let commands = self.build_palette_commands();
359            stack![
360                graph_element,
361                opaque(command_palette(
362                    &self.command_input,
363                    &commands,
364                    self.palette_selected_index,
365                    Message::CommandPaletteInput,
366                    Message::CommandPaletteSelect,
367                    Message::CommandPaletteNavigate,
368                    || Message::CommandPaletteCancel,
369                ))
370            ]
371            .into()
372        } else {
373            graph_element
374        }
375    }
376
377    fn theme(&self) -> Theme {
378        self.current_theme.clone()
379    }
380
381    fn set_theme(&mut self, theme: Theme) {
382        self.current_theme = theme;
383    }
384
385    fn subscription(&self) -> Subscription<Message> {
386        use iced::keyboard::key::Named;
387
388        Subscription::batch(vec![
389            // Keyboard events for command palette
390            event::listen_with(|event, _status, _id| {
391                if let Event::Keyboard(keyboard::Event::KeyPressed { key, modifiers, .. }) = event {
392                    // Ctrl+Space or Ctrl+Space to toggle palette
393                    if modifiers.command() && key == keyboard::Key::Named(Named::Space) {
394                        return Some(Message::ToggleCommandPalette);
395                    }
396
397                    // When palette is open, handle navigation
398                    match key {
399                        keyboard::Key::Named(Named::ArrowUp) => {
400                            return Some(Message::CommandPaletteNavigateUp);
401                        }
402                        keyboard::Key::Named(Named::ArrowDown) => {
403                            return Some(Message::CommandPaletteNavigateDown);
404                        }
405                        keyboard::Key::Named(Named::Enter) => {
406                            return Some(Message::CommandPaletteConfirm);
407                        }
408                        keyboard::Key::Named(Named::Escape) => {
409                            return Some(Message::CommandPaletteCancel);
410                        }
411                        _ => {}
412                    }
413                }
414                None
415            }),
416            // Window resize events
417            event::listen_with(|event, _, _| match event {
418                Event::Window(window::Event::Resized(size)) => Some(Message::WindowResized(size)),
419                _ => None,
420            }),
421        ])
422    }
423}
424
425/// Boots this demo for the gallery.
426pub fn scene() -> (
427    Box<dyn demo_common::Scene>,
428    iced::Task<demo_common::SceneMessage>,
429) {
430    demo_common::erase::<Application>()
431}
432
433impl Application {
434    /// Calculate spawn position at screen center, converted to world coordinates.
435    fn spawn_position(&self) -> Point {
436        // Screen center
437        let screen_center_x = self.viewport_size.width / 2.0;
438        let screen_center_y = self.viewport_size.height / 2.0;
439
440        // Convert to world coordinates: world = screen / zoom - camera_position
441        let world_x = screen_center_x / self.camera_zoom - self.camera_position.x;
442        let world_y = screen_center_y / self.camera_zoom - self.camera_position.y;
443
444        // Offset for node size (approximate center)
445        Point::new(world_x - 60.0, world_y - 40.0)
446    }
447
448    fn build_palette_commands(&self) -> Vec<Command<Message>> {
449        let mut commands = Vec::new();
450
451        // Add node spawning commands for all shader node types
452        for node_type in ShaderNodeType::all() {
453            let category = node_type.category();
454            commands.push(
455                command(node_type.name(), node_type.name())
456                    .description(format!("{} node", category))
457                    .action(Message::SpawnNode(*node_type)),
458            );
459        }
460
461        // Add theme switching commands
462        commands.push(
463            command("theme-dark", "Dark Theme")
464                .description("Switch to dark theme")
465                .action(Message::ChangeTheme(Theme::Dark)),
466        );
467        commands.push(
468            command("theme-light", "Light Theme")
469                .description("Switch to light theme")
470                .action(Message::ChangeTheme(Theme::Light)),
471        );
472        commands.push(
473            command("theme-catppuccin", "Catppuccin Mocha")
474                .description("Switch to Catppuccin Mocha theme")
475                .action(Message::ChangeTheme(Theme::CatppuccinMocha)),
476        );
477        commands.push(
478            command("theme-dracula", "Dracula")
479                .description("Switch to Dracula theme")
480                .action(Message::ChangeTheme(Theme::Dracula)),
481        );
482        commands.push(
483            command("theme-nord", "Nord")
484                .description("Switch to Nord theme")
485                .action(Message::ChangeTheme(Theme::Nord)),
486        );
487
488        commands
489    }
490
491    fn recompile(&mut self) {
492        match ShaderCompiler::compile(&self.shader_graph) {
493            Ok(shader) => {
494                self.compiled_shader = Some(shader);
495                self.compilation_error = None;
496            }
497            Err(err) => {
498                self.compiled_shader = None;
499                self.compilation_error = Some(err.to_string());
500            }
501        }
502    }
503}
504
505fn create_node_widget<'a>(
506    node_type: &shader_graph::nodes::ShaderNodeType,
507    theme: &'a Theme,
508) -> iced::Element<'a, Message> {
509    use iced_nodegraph::node_pin;
510
511    let name = node_type.name();
512    let inputs = node_type.inputs();
513    let outputs = node_type.outputs();
514
515    let palette = theme.extended_palette();
516
517    // Title bar - matching hello_world pattern exactly
518    let title_bar = container(text(name).size(13).width(Length::Fill))
519        .width(Length::Fill)
520        .padding([2, 8])
521        .style(move |_theme: &iced::Theme| container::Style {
522            background: None,
523            text_color: Some(palette.background.base.text),
524            ..container::Style::default()
525        });
526
527    // Build pin list - must match hello_world's column![] macro structure
528    // Pin IDs use sequential indices: inputs first (0..n), then outputs (n..n+m)
529    let pin_section = if inputs.is_empty() && outputs.is_empty() {
530        // No pins - minimal output indicator
531        container(
532            column![
533                node_pin(
534                    PinSide::Right,
535                    0usize,
536                    container(text("out").size(11)).padding([0, 8])
537                )
538                .direction(PinDirection::Output)
539            ]
540            .spacing(2),
541        )
542        .padding([6, 0])
543    } else {
544        // Build pins dynamically but wrap in container same way
545        let mut pin_elements: Vec<iced::Element<'a, Message>> = Vec::new();
546        let num_inputs = inputs.len();
547
548        for (i, input) in inputs.into_iter().enumerate() {
549            let label = input.name.clone();
550            pin_elements.push(create_typed_pin(
551                PinSide::Left,
552                i, // Pin ID = input index
553                label,
554                PinDirection::Input,
555                &input.socket_type,
556            ));
557        }
558
559        for (i, output) in outputs.into_iter().enumerate() {
560            let label = output.name.clone();
561            pin_elements.push(create_typed_pin(
562                PinSide::Right,
563                num_inputs + i, // Pin ID = num_inputs + output index
564                label,
565                PinDirection::Output,
566                &output.socket_type,
567            ));
568        }
569
570        container(iced::widget::Column::with_children(pin_elements).spacing(2)).padding([6, 0])
571    };
572
573    column![title_bar, pin_section].width(160.0).into()
574}
575
576/// Colors a node's pins by their socket data-type marker.
577fn pin_style(
578    theme: &Theme,
579    pin: &PinInfo<'_, TypedIds>,
580    _other: Option<&PinInfo<'_, TypedIds>>,
581    status: PinStatus,
582) -> PinStyle {
583    use std::any::TypeId;
584    let ty = *pin.info();
585    let color = if ty == TypeId::of::<colors::Float>() {
586        colors::SOCKET_FLOAT
587    } else if ty == TypeId::of::<colors::Vec2>() {
588        colors::SOCKET_VEC2
589    } else if ty == TypeId::of::<colors::Vec3>() {
590        colors::SOCKET_VEC3
591    } else {
592        colors::SOCKET_VEC4
593    };
594    PinStyle {
595        color: color.into(),
596        ..default_pin_style(theme, status)
597    }
598}
599
600/// Creates a typed pin element based on the socket type.
601/// Uses marker types for TypeId-based connection matching.
602fn create_typed_pin<'a, Message: Clone + 'a>(
603    side: PinSide,
604    pin_id: usize,
605    label: String,
606    direction: PinDirection,
607    socket_type: &shader_graph::sockets::SocketType,
608) -> Element<'a, Message> {
609    use shader_graph::sockets::SocketType;
610
611    let content = container(text(label).size(11)).padding([0, 8]);
612
613    match socket_type {
614        SocketType::Float => node_pin(side, pin_id, content)
615            .direction(direction)
616            .info(::std::any::TypeId::of::<colors::Float>())
617            .into(),
618        SocketType::Vec2 => node_pin(side, pin_id, content)
619            .direction(direction)
620            .info(::std::any::TypeId::of::<colors::Vec2>())
621            .into(),
622        SocketType::Vec3 => node_pin(side, pin_id, content)
623            .direction(direction)
624            .info(::std::any::TypeId::of::<colors::Vec3>())
625            .into(),
626        SocketType::Vec4 => node_pin(side, pin_id, content)
627            .direction(direction)
628            .info(::std::any::TypeId::of::<colors::Vec4>())
629            .into(),
630    }
631}