Skip to main content

proof_engine/editor/
node_editor.rs

1
2//! Generic node editor foundation — ports, connections, layout, zoom/pan, minimap.
3
4use glam::{Vec2, Vec4};
5use std::collections::HashMap;
6
7// ---------------------------------------------------------------------------
8// Port / pin definitions
9// ---------------------------------------------------------------------------
10
11#[derive(Debug, Clone, Copy, PartialEq)]
12pub enum PortDirection { Input, Output }
13
14#[derive(Debug, Clone, Copy, PartialEq)]
15pub enum PortKind {
16    Float, Float2, Float3, Float4,
17    Int, Bool, Color,
18    Texture2D, TextureCube, Sampler,
19    Matrix2, Matrix3, Matrix4,
20    String, Any, Flow, Object,
21}
22
23impl PortKind {
24    pub fn color(self) -> Vec4 {
25        match self {
26            PortKind::Float => Vec4::new(0.6, 0.6, 0.6, 1.0),
27            PortKind::Float2 => Vec4::new(0.4, 0.8, 0.4, 1.0),
28            PortKind::Float3 => Vec4::new(0.4, 0.4, 0.9, 1.0),
29            PortKind::Float4 => Vec4::new(0.8, 0.4, 0.8, 1.0),
30            PortKind::Int => Vec4::new(0.4, 0.7, 0.9, 1.0),
31            PortKind::Bool => Vec4::new(0.9, 0.7, 0.3, 1.0),
32            PortKind::Color => Vec4::new(1.0, 0.8, 0.2, 1.0),
33            PortKind::Texture2D | PortKind::TextureCube => Vec4::new(0.8, 0.5, 0.2, 1.0),
34            PortKind::Sampler => Vec4::new(0.7, 0.4, 0.1, 1.0),
35            PortKind::Matrix2 | PortKind::Matrix3 | PortKind::Matrix4 => Vec4::new(0.9, 0.2, 0.4, 1.0),
36            PortKind::Flow => Vec4::new(1.0, 1.0, 1.0, 1.0),
37            PortKind::Object => Vec4::new(0.5, 0.9, 0.9, 1.0),
38            _ => Vec4::new(0.5, 0.5, 0.5, 1.0),
39        }
40    }
41
42    pub fn label(self) -> &'static str {
43        match self {
44            PortKind::Float => "float",
45            PortKind::Float2 => "vec2",
46            PortKind::Float3 => "vec3",
47            PortKind::Float4 => "vec4",
48            PortKind::Int => "int",
49            PortKind::Bool => "bool",
50            PortKind::Color => "color",
51            PortKind::Texture2D => "Texture2D",
52            PortKind::TextureCube => "TextureCube",
53            PortKind::Sampler => "Sampler",
54            PortKind::Matrix2 => "mat2",
55            PortKind::Matrix3 => "mat3",
56            PortKind::Matrix4 => "mat4",
57            PortKind::String => "string",
58            PortKind::Any => "any",
59            PortKind::Flow => "flow",
60            PortKind::Object => "object",
61        }
62    }
63
64    pub fn can_connect_to(self, other: PortKind) -> bool {
65        if self == other { return true; }
66        if self == PortKind::Any || other == PortKind::Any { return true; }
67        // Numeric promotions
68        matches!((self, other),
69            (PortKind::Float, PortKind::Float2) |
70            (PortKind::Float, PortKind::Float3) |
71            (PortKind::Float, PortKind::Float4) |
72            (PortKind::Float, PortKind::Color) |
73            (PortKind::Float3, PortKind::Color) |
74            (PortKind::Float4, PortKind::Color) |
75            (PortKind::Color, PortKind::Float3) |
76            (PortKind::Color, PortKind::Float4) |
77            (PortKind::Int, PortKind::Float)
78        )
79    }
80}
81
82#[derive(Debug, Clone)]
83pub struct Port {
84    pub id: u32,
85    pub name: String,
86    pub kind: PortKind,
87    pub direction: PortDirection,
88    pub optional: bool,
89    pub default_value: Option<Vec4>,
90    pub tooltip: String,
91}
92
93impl Port {
94    pub fn input(id: u32, name: impl Into<String>, kind: PortKind) -> Self {
95        Self {
96            id, name: name.into(), kind, direction: PortDirection::Input,
97            optional: false, default_value: None, tooltip: String::new(),
98        }
99    }
100    pub fn output(id: u32, name: impl Into<String>, kind: PortKind) -> Self {
101        Self {
102            id, name: name.into(), kind, direction: PortDirection::Output,
103            optional: false, default_value: None, tooltip: String::new(),
104        }
105    }
106    pub fn with_default(mut self, v: Vec4) -> Self { self.default_value = Some(v); self }
107    pub fn optional(mut self) -> Self { self.optional = true; self }
108}
109
110// ---------------------------------------------------------------------------
111// Node
112// ---------------------------------------------------------------------------
113
114#[derive(Debug, Clone, Copy, PartialEq)]
115pub enum NodeState {
116    Normal,
117    Selected,
118    Hovered,
119    Error,
120    Warning,
121    Disabled,
122    Processing,
123    Dirty,
124}
125
126#[derive(Debug, Clone)]
127pub struct Node {
128    pub id: u32,
129    pub title: String,
130    pub category: String,
131    pub position: Vec2,
132    pub size: Vec2,
133    pub ports: Vec<Port>,
134    pub state: NodeState,
135    pub collapsed: bool,
136    pub comment: Option<String>,
137    pub color: Vec4,
138    pub pinned: bool,
139    pub metadata: HashMap<String, String>,
140    pub execution_order: u32,
141    pub last_processed_frame: u64,
142    pub processing_time_us: u64,
143}
144
145impl Node {
146    pub fn new(id: u32, title: impl Into<String>, category: impl Into<String>) -> Self {
147        Self {
148            id,
149            title: title.into(),
150            category: category.into(),
151            position: Vec2::ZERO,
152            size: Vec2::new(150.0, 60.0),
153            ports: Vec::new(),
154            state: NodeState::Normal,
155            collapsed: false,
156            comment: None,
157            color: Vec4::new(0.2, 0.2, 0.2, 1.0),
158            pinned: false,
159            metadata: HashMap::new(),
160            execution_order: 0,
161            last_processed_frame: 0,
162            processing_time_us: 0,
163        }
164    }
165
166    pub fn with_position(mut self, pos: Vec2) -> Self { self.position = pos; self }
167    pub fn with_color(mut self, color: Vec4) -> Self { self.color = color; self }
168
169    pub fn add_port(&mut self, port: Port) {
170        self.ports.push(port);
171        self.recalculate_size();
172    }
173
174    pub fn recalculate_size(&mut self) {
175        let input_count = self.ports.iter().filter(|p| p.direction == PortDirection::Input).count();
176        let output_count = self.ports.iter().filter(|p| p.direction == PortDirection::Output).count();
177        let rows = input_count.max(output_count);
178        self.size = Vec2::new(160.0, 40.0 + rows as f32 * 22.0);
179    }
180
181    pub fn port_position(&self, port_id: u32) -> Option<Vec2> {
182        let inputs: Vec<u32> = self.ports.iter().filter(|p| p.direction == PortDirection::Input).map(|p| p.id).collect();
183        let outputs: Vec<u32> = self.ports.iter().filter(|p| p.direction == PortDirection::Output).map(|p| p.id).collect();
184        if let Some(i) = inputs.iter().position(|&id| id == port_id) {
185            return Some(self.position + Vec2::new(0.0, 40.0 + i as f32 * 22.0));
186        }
187        if let Some(i) = outputs.iter().position(|&id| id == port_id) {
188            return Some(self.position + Vec2::new(self.size.x, 40.0 + i as f32 * 22.0));
189        }
190        None
191    }
192
193    pub fn input_ports(&self) -> impl Iterator<Item = &Port> {
194        self.ports.iter().filter(|p| p.direction == PortDirection::Input)
195    }
196    pub fn output_ports(&self) -> impl Iterator<Item = &Port> {
197        self.ports.iter().filter(|p| p.direction == PortDirection::Output)
198    }
199
200    pub fn bounds(&self) -> [Vec2; 2] {
201        [self.position, self.position + self.size]
202    }
203
204    pub fn contains_point(&self, p: Vec2) -> bool {
205        p.x >= self.position.x && p.x <= self.position.x + self.size.x &&
206        p.y >= self.position.y && p.y <= self.position.y + self.size.y
207    }
208}
209
210// ---------------------------------------------------------------------------
211// Connection / edge
212// ---------------------------------------------------------------------------
213
214#[derive(Debug, Clone, Copy)]
215pub struct Connection {
216    pub id: u32,
217    pub from_node: u32,
218    pub from_port: u32,
219    pub to_node: u32,
220    pub to_port: u32,
221}
222
223impl Connection {
224    pub fn new(id: u32, from_node: u32, from_port: u32, to_node: u32, to_port: u32) -> Self {
225        Self { id, from_node, from_port, to_node, to_port }
226    }
227}
228
229// ---------------------------------------------------------------------------
230// Node graph
231// ---------------------------------------------------------------------------
232
233#[derive(Debug, Clone)]
234pub struct NodeGraph {
235    pub name: String,
236    pub nodes: Vec<Node>,
237    pub connections: Vec<Connection>,
238    pub next_node_id: u32,
239    pub next_conn_id: u32,
240    pub next_port_id: u32,
241    pub dirty: bool,
242}
243
244impl NodeGraph {
245    pub fn new(name: impl Into<String>) -> Self {
246        Self {
247            name: name.into(),
248            nodes: Vec::new(),
249            connections: Vec::new(),
250            next_node_id: 1,
251            next_conn_id: 1,
252            next_port_id: 1,
253            dirty: false,
254        }
255    }
256
257    pub fn add_node(&mut self, mut node: Node) -> u32 {
258        let id = self.next_node_id;
259        node.id = id;
260        self.next_node_id += 1;
261        self.nodes.push(node);
262        self.dirty = true;
263        id
264    }
265
266    pub fn remove_node(&mut self, id: u32) {
267        self.nodes.retain(|n| n.id != id);
268        self.connections.retain(|c| c.from_node != id && c.to_node != id);
269        self.dirty = true;
270    }
271
272    pub fn connect(&mut self, from_node: u32, from_port: u32, to_node: u32, to_port: u32) -> Result<u32, &'static str> {
273        // Validate nodes exist
274        if self.node(from_node).is_none() { return Err("source node not found"); }
275        if self.node(to_node).is_none() { return Err("target node not found"); }
276        // Check for existing input connection (inputs can only have one source)
277        self.connections.retain(|c| !(c.to_node == to_node && c.to_port == to_port));
278        // Cycle check
279        if self.would_create_cycle(from_node, to_node) {
280            return Err("connection would create cycle");
281        }
282        let id = self.next_conn_id;
283        self.next_conn_id += 1;
284        self.connections.push(Connection::new(id, from_node, from_port, to_node, to_port));
285        self.dirty = true;
286        Ok(id)
287    }
288
289    pub fn disconnect(&mut self, conn_id: u32) {
290        self.connections.retain(|c| c.id != conn_id);
291        self.dirty = true;
292    }
293
294    pub fn would_create_cycle(&self, from: u32, to: u32) -> bool {
295        // BFS from `to` — does it reach `from`?
296        let mut visited = std::collections::HashSet::new();
297        let mut queue = vec![to];
298        while let Some(n) = queue.pop() {
299            if n == from { return true; }
300            if visited.contains(&n) { continue; }
301            visited.insert(n);
302            for c in &self.connections {
303                if c.from_node == n { queue.push(c.to_node); }
304            }
305        }
306        false
307    }
308
309    pub fn topological_order(&self) -> Vec<u32> {
310        let mut in_degree: HashMap<u32, usize> = self.nodes.iter().map(|n| (n.id, 0)).collect();
311        for conn in &self.connections {
312            *in_degree.entry(conn.to_node).or_insert(0) += 1;
313        }
314        let mut queue: Vec<u32> = in_degree.iter().filter(|(_, &d)| d == 0).map(|(&id, _)| id).collect();
315        let mut order = Vec::new();
316        while let Some(n) = queue.pop() {
317            order.push(n);
318            for conn in &self.connections {
319                if conn.from_node == n {
320                    let d = in_degree.entry(conn.to_node).or_insert(1);
321                    *d -= 1;
322                    if *d == 0 { queue.push(conn.to_node); }
323                }
324            }
325        }
326        order
327    }
328
329    pub fn node(&self, id: u32) -> Option<&Node> {
330        self.nodes.iter().find(|n| n.id == id)
331    }
332    pub fn node_mut(&mut self, id: u32) -> Option<&mut Node> {
333        self.nodes.iter_mut().find(|n| n.id == id)
334    }
335
336    pub fn connections_from(&self, node_id: u32) -> impl Iterator<Item = &Connection> {
337        self.connections.iter().filter(move |c| c.from_node == node_id)
338    }
339
340    pub fn connections_to(&self, node_id: u32) -> impl Iterator<Item = &Connection> {
341        self.connections.iter().filter(move |c| c.to_node == node_id)
342    }
343
344    pub fn upstream_nodes(&self, node_id: u32) -> Vec<u32> {
345        self.connections.iter()
346            .filter(|c| c.to_node == node_id)
347            .map(|c| c.from_node)
348            .collect()
349    }
350
351    pub fn downstream_nodes(&self, node_id: u32) -> Vec<u32> {
352        self.connections.iter()
353            .filter(|c| c.from_node == node_id)
354            .map(|c| c.to_node)
355            .collect()
356    }
357
358    pub fn selected_nodes(&self) -> impl Iterator<Item = &Node> {
359        self.nodes.iter().filter(|n| n.state == NodeState::Selected)
360    }
361
362    pub fn auto_layout(&mut self) {
363        // Simple layered layout (Sugiyama-style approximation)
364        let order = self.topological_order();
365        let mut layers: HashMap<u32, u32> = HashMap::new();
366        for &n in &order {
367            let layer = self.connections_to(n)
368                .filter_map(|c| layers.get(&c.from_node))
369                .max()
370                .map(|&l| l + 1)
371                .unwrap_or(0);
372            layers.insert(n, layer);
373        }
374        let mut layer_counts: HashMap<u32, u32> = HashMap::new();
375        for n in self.nodes.iter_mut() {
376            let layer = layers.get(&n.id).copied().unwrap_or(0);
377            let row = *layer_counts.entry(layer).or_insert(0);
378            *layer_counts.get_mut(&layer).unwrap() += 1;
379            n.position = Vec2::new(layer as f32 * 220.0 + 20.0, row as f32 * 140.0 + 20.0);
380        }
381    }
382}
383
384// ---------------------------------------------------------------------------
385// Graph view / viewport
386// ---------------------------------------------------------------------------
387
388#[derive(Debug, Clone, Copy, PartialEq)]
389pub enum NodeEditorAction {
390    None,
391    DragNode(u32),
392    DragConnection(u32, u32),   // from_node, from_port
393    DragSelect,
394    Pan,
395    Zoom,
396}
397
398#[derive(Debug, Clone)]
399pub struct NodeEditorView {
400    pub offset: Vec2,
401    pub zoom: f32,
402    pub canvas_size: Vec2,
403    pub selection_rect: Option<[Vec2; 2]>,
404    pub active_action: NodeEditorAction,
405    pub hovered_node: Option<u32>,
406    pub hovered_port: Option<(u32, u32)>,
407    pub pending_connection: Option<(u32, u32)>,
408    pub show_minimap: bool,
409    pub show_grid: bool,
410    pub snap_to_grid: bool,
411    pub grid_size: f32,
412    pub minimap_rect: [Vec2; 2],
413}
414
415impl NodeEditorView {
416    pub fn new(canvas_size: Vec2) -> Self {
417        Self {
418            offset: Vec2::ZERO,
419            zoom: 1.0,
420            canvas_size,
421            selection_rect: None,
422            active_action: NodeEditorAction::None,
423            hovered_node: None,
424            hovered_port: None,
425            pending_connection: None,
426            show_minimap: true,
427            show_grid: true,
428            snap_to_grid: false,
429            grid_size: 20.0,
430            minimap_rect: [
431                Vec2::new(canvas_size.x - 200.0, canvas_size.y - 150.0),
432                Vec2::new(canvas_size.x - 10.0, canvas_size.y - 10.0),
433            ],
434        }
435    }
436
437    pub fn screen_to_canvas(&self, screen: Vec2) -> Vec2 {
438        (screen - self.canvas_size * 0.5 - self.offset) / self.zoom
439    }
440
441    pub fn canvas_to_screen(&self, canvas: Vec2) -> Vec2 {
442        canvas * self.zoom + self.canvas_size * 0.5 + self.offset
443    }
444
445    pub fn zoom_at(&mut self, screen_pivot: Vec2, factor: f32) {
446        let canvas_pivot = self.screen_to_canvas(screen_pivot);
447        self.zoom = (self.zoom * factor).clamp(0.1, 8.0);
448        let new_screen = self.canvas_to_screen(canvas_pivot);
449        self.offset += screen_pivot - new_screen;
450    }
451
452    pub fn pan(&mut self, delta: Vec2) {
453        self.offset += delta;
454    }
455
456    pub fn fit_to_graph(&mut self, graph: &NodeGraph) {
457        if graph.nodes.is_empty() { return; }
458        let mut min = Vec2::splat(f32::INFINITY);
459        let mut max = Vec2::splat(f32::NEG_INFINITY);
460        for node in &graph.nodes {
461            min = min.min(node.position);
462            max = max.max(node.position + node.size);
463        }
464        let content_size = max - min;
465        let scale_x = (self.canvas_size.x - 80.0) / content_size.x.max(1.0);
466        let scale_y = (self.canvas_size.y - 80.0) / content_size.y.max(1.0);
467        self.zoom = scale_x.min(scale_y).min(1.5);
468        let center = (min + max) * 0.5;
469        self.offset = -center * self.zoom;
470    }
471
472    pub fn snap_position(&self, pos: Vec2) -> Vec2 {
473        if !self.snap_to_grid { return pos; }
474        let g = self.grid_size;
475        Vec2::new((pos.x / g).round() * g, (pos.y / g).round() * g)
476    }
477
478    pub fn nodes_in_selection(&self, graph: &NodeGraph) -> Vec<u32> {
479        if let Some([a, b]) = self.selection_rect {
480            let min = a.min(b);
481            let max = a.max(b);
482            let canvas_min = self.screen_to_canvas(min);
483            let canvas_max = self.screen_to_canvas(max);
484            graph.nodes.iter()
485                .filter(|n| n.position.x < canvas_max.x && n.position.x + n.size.x > canvas_min.x &&
486                            n.position.y < canvas_max.y && n.position.y + n.size.y > canvas_min.y)
487                .map(|n| n.id)
488                .collect()
489        } else {
490            Vec::new()
491        }
492    }
493
494    pub fn cubic_bezier_connection(from: Vec2, to: Vec2) -> [Vec2; 4] {
495        let dx = (to.x - from.x).abs().max(50.0) * 0.5;
496        let cp1 = Vec2::new(from.x + dx, from.y);
497        let cp2 = Vec2::new(to.x - dx, to.y);
498        [from, cp1, cp2, to]
499    }
500}
501
502// ---------------------------------------------------------------------------
503// Context menu
504// ---------------------------------------------------------------------------
505
506#[derive(Debug, Clone)]
507pub struct ContextMenuItem {
508    pub label: String,
509    pub shortcut: Option<String>,
510    pub enabled: bool,
511    pub separator_after: bool,
512    pub action: String,
513    pub children: Vec<ContextMenuItem>,
514}
515
516impl ContextMenuItem {
517    pub fn action(label: impl Into<String>, action: impl Into<String>) -> Self {
518        Self {
519            label: label.into(),
520            shortcut: None,
521            enabled: true,
522            separator_after: false,
523            action: action.into(),
524            children: Vec::new(),
525        }
526    }
527    pub fn separator() -> Self {
528        Self { label: "---".into(), shortcut: None, enabled: false, separator_after: false, action: String::new(), children: Vec::new() }
529    }
530    pub fn submenu(label: impl Into<String>, children: Vec<ContextMenuItem>) -> Self {
531        Self { label: label.into(), shortcut: None, enabled: true, separator_after: false, action: String::new(), children }
532    }
533}
534
535// ---------------------------------------------------------------------------
536// Node editor
537// ---------------------------------------------------------------------------
538
539#[derive(Debug, Clone)]
540pub struct NodeEditor {
541    pub graphs: Vec<NodeGraph>,
542    pub active_graph: usize,
543    pub view: NodeEditorView,
544    pub clipboard: Vec<Node>,
545    pub search_query: String,
546    pub show_context_menu: bool,
547    pub context_menu_pos: Vec2,
548    pub context_menu_items: Vec<ContextMenuItem>,
549    pub undo_stack: Vec<NodeGraph>,
550    pub undo_pos: usize,
551    pub status_message: String,
552    pub node_categories: Vec<String>,
553}
554
555impl NodeEditor {
556    pub fn new() -> Self {
557        let mut ed = Self {
558            graphs: vec![NodeGraph::new("Main")],
559            active_graph: 0,
560            view: NodeEditorView::new(Vec2::new(1200.0, 800.0)),
561            clipboard: Vec::new(),
562            search_query: String::new(),
563            show_context_menu: false,
564            context_menu_pos: Vec2::ZERO,
565            context_menu_items: Vec::new(),
566            undo_stack: Vec::new(),
567            undo_pos: 0,
568            status_message: String::new(),
569            node_categories: vec!["Math".into(), "Logic".into(), "Color".into(), "Texture".into(), "Utility".into()],
570        };
571        ed.populate_demo_graph();
572        ed
573    }
574
575    fn populate_demo_graph(&mut self) {
576        let g = &mut self.graphs[0];
577        let mut n1 = Node::new(0, "Add", "Math");
578        n1.add_port(Port::input(1, "A", PortKind::Float).with_default(Vec4::ZERO));
579        n1.add_port(Port::input(2, "B", PortKind::Float).with_default(Vec4::ZERO));
580        n1.add_port(Port::output(3, "Result", PortKind::Float));
581        n1.position = Vec2::new(200.0, 100.0);
582        let id1 = g.add_node(n1);
583
584        let mut n2 = Node::new(0, "Multiply", "Math");
585        n2.add_port(Port::input(1, "A", PortKind::Float));
586        n2.add_port(Port::input(2, "B", PortKind::Float));
587        n2.add_port(Port::output(3, "Result", PortKind::Float));
588        n2.position = Vec2::new(400.0, 100.0);
589        let id2 = g.add_node(n2);
590
591        let mut n3 = Node::new(0, "Output", "Output");
592        n3.add_port(Port::input(1, "Value", PortKind::Float));
593        n3.position = Vec2::new(600.0, 100.0);
594        n3.color = Vec4::new(0.1, 0.35, 0.1, 1.0);
595        let id3 = g.add_node(n3);
596
597        let _ = g.connect(id1, 3, id2, 1);
598        let _ = g.connect(id2, 3, id3, 1);
599    }
600
601    pub fn active_graph(&self) -> &NodeGraph {
602        &self.graphs[self.active_graph]
603    }
604
605    pub fn active_graph_mut(&mut self) -> &mut NodeGraph {
606        &mut self.graphs[self.active_graph]
607    }
608
609    pub fn snapshot(&mut self) {
610        let graph = self.active_graph().clone();
611        self.undo_stack.truncate(self.undo_pos);
612        self.undo_stack.push(graph);
613        self.undo_pos = self.undo_stack.len();
614    }
615
616    pub fn undo(&mut self) {
617        if self.undo_pos > 1 {
618            self.undo_pos -= 1;
619            self.graphs[self.active_graph] = self.undo_stack[self.undo_pos - 1].clone();
620        }
621    }
622
623    pub fn redo(&mut self) {
624        if self.undo_pos < self.undo_stack.len() {
625            self.graphs[self.active_graph] = self.undo_stack[self.undo_pos].clone();
626            self.undo_pos += 1;
627        }
628    }
629
630    pub fn copy_selected(&mut self) {
631        self.clipboard = self.active_graph().selected_nodes().cloned().collect();
632    }
633
634    pub fn paste(&mut self) {
635        if self.clipboard.is_empty() { return; }
636        self.snapshot();
637        let nodes: Vec<Node> = self.clipboard.clone();
638        let g = self.active_graph_mut();
639        for node in nodes {
640            let mut n = node;
641            n.position += Vec2::new(20.0, 20.0);
642            n.state = NodeState::Selected;
643            g.add_node(n);
644        }
645    }
646
647    pub fn delete_selected(&mut self) {
648        self.snapshot();
649        let to_remove: Vec<u32> = self.active_graph().selected_nodes().map(|n| n.id).collect();
650        let g = self.active_graph_mut();
651        for id in to_remove {
652            g.remove_node(id);
653        }
654    }
655
656    pub fn select_all(&mut self) {
657        for n in self.active_graph_mut().nodes.iter_mut() {
658            n.state = NodeState::Selected;
659        }
660    }
661
662    pub fn deselect_all(&mut self) {
663        for n in self.active_graph_mut().nodes.iter_mut() {
664            if n.state == NodeState::Selected {
665                n.state = NodeState::Normal;
666            }
667        }
668    }
669
670    pub fn auto_layout(&mut self) {
671        self.snapshot();
672        self.active_graph_mut().auto_layout();
673        let ai = self.active_graph;
674        self.view.fit_to_graph(&self.graphs[ai]);
675    }
676
677    pub fn build_context_menu(&mut self, pos: Vec2) {
678        self.context_menu_pos = pos;
679        self.show_context_menu = true;
680        let canvas_pos = self.view.screen_to_canvas(pos);
681        let on_node = self.active_graph().nodes.iter().any(|n| n.contains_point(canvas_pos));
682        if on_node {
683            self.context_menu_items = vec![
684                ContextMenuItem::action("Delete", "delete_selected"),
685                ContextMenuItem::action("Duplicate", "duplicate"),
686                ContextMenuItem::separator(),
687                ContextMenuItem::action("Collapse", "collapse"),
688                ContextMenuItem::action("Add Comment", "add_comment"),
689                ContextMenuItem::separator(),
690                ContextMenuItem::action("Properties", "properties"),
691            ];
692        } else {
693            self.context_menu_items = vec![
694                ContextMenuItem::submenu("Add Node", vec![
695                    ContextMenuItem::action("Add", "add_add_node"),
696                    ContextMenuItem::action("Multiply", "add_mul_node"),
697                    ContextMenuItem::action("Constant", "add_const_node"),
698                ]),
699                ContextMenuItem::separator(),
700                ContextMenuItem::action("Select All", "select_all"),
701                ContextMenuItem::action("Auto Layout", "auto_layout"),
702                ContextMenuItem::separator(),
703                ContextMenuItem::action("Paste", "paste"),
704            ];
705        }
706    }
707}
708
709// ---------------------------------------------------------------------------
710// Tests
711// ---------------------------------------------------------------------------
712#[cfg(test)]
713mod tests {
714    use super::*;
715
716    #[test]
717    fn test_port_compatibility() {
718        assert!(PortKind::Float.can_connect_to(PortKind::Float3));
719        assert!(!PortKind::Texture2D.can_connect_to(PortKind::Float));
720        assert!(PortKind::Any.can_connect_to(PortKind::Matrix4));
721    }
722
723    #[test]
724    fn test_graph_cycle_detection() {
725        let mut g = NodeGraph::new("test");
726        let n1 = g.add_node(Node::new(0, "A", "test"));
727        let n2 = g.add_node(Node::new(0, "B", "test"));
728        let _ = g.connect(n1, 0, n2, 0);
729        assert!(g.would_create_cycle(n2, n1));
730        assert!(!g.would_create_cycle(n1, n2));
731    }
732
733    #[test]
734    fn test_topo_sort() {
735        let mut g = NodeGraph::new("test");
736        let n1 = g.add_node(Node::new(0, "A", "test"));
737        let n2 = g.add_node(Node::new(0, "B", "test"));
738        let n3 = g.add_node(Node::new(0, "C", "test"));
739        let _ = g.connect(n1, 0, n2, 0);
740        let _ = g.connect(n2, 0, n3, 0);
741        let order = g.topological_order();
742        let pos1 = order.iter().position(|&x| x == n1).unwrap();
743        let pos2 = order.iter().position(|&x| x == n2).unwrap();
744        assert!(pos1 < pos2);
745    }
746
747    #[test]
748    fn test_view_transform() {
749        let mut view = NodeEditorView::new(Vec2::new(800.0, 600.0));
750        view.zoom = 2.0;
751        let canvas = Vec2::new(10.0, 20.0);
752        let screen = view.canvas_to_screen(canvas);
753        let back = view.screen_to_canvas(screen);
754        assert!((back.x - canvas.x).abs() < 0.001);
755    }
756
757    #[test]
758    fn test_node_editor() {
759        let mut ed = NodeEditor::new();
760        assert!(!ed.active_graph().nodes.is_empty());
761        ed.select_all();
762        ed.copy_selected();
763        ed.paste();
764        let count = ed.active_graph().nodes.len();
765        assert!(count > 3);
766    }
767}