Skip to main content

pedalkernel_layout/
graph.rs

1//! Phase 1: Directed graph construction from the `.pedal` netlist.
2//!
3//! Builds a [`LayoutGraph`] where each component is a node and each net
4//! creates edges between the components it connects. Edge direction is
5//! inferred from component type and pin function (e.g., triode grid = input,
6//! plate = output).
7
8pub use pedalkernel::compiler::PinDirection;
9
10use pedalkernel::compiler::Component;
11use pedalkernel::dsl::{ComponentDef, NetDef, PedalDef, Pin};
12use std::collections::{BTreeMap, BTreeSet};
13
14// ---------------------------------------------------------------------------
15// Layout graph
16// ---------------------------------------------------------------------------
17
18/// A node in the layout graph (one per component + special anchor nodes).
19#[derive(Debug, Clone)]
20pub struct LayoutNode {
21    /// Index into `LayoutGraph::nodes`.
22    pub id: usize,
23    /// Component ID (or `"__in"`, `"__out"`, `"__vcc"`, `"__gnd"` for anchors).
24    pub comp_id: String,
25    /// Component definition (None for anchor nodes).
26    pub comp: Option<ComponentDef>,
27    /// Whether this is an anchor node (`in`, `out`, `vcc`, `gnd`).
28    pub is_anchor: bool,
29}
30
31/// A directed edge between two layout nodes.
32#[derive(Debug, Clone)]
33pub struct LayoutEdge {
34    pub from: usize,
35    pub to: usize,
36    /// Net name (if any).
37    pub net_name: Option<String>,
38    /// Whether this edge is on the primary signal path.
39    pub signal_path: bool,
40    /// Whether this is a supply connection (to vcc/gnd).
41    pub is_supply: bool,
42    /// Whether this is a feedback edge (back-edge in topological sort).
43    pub is_feedback: bool,
44}
45
46/// Directed graph built from the `.pedal` netlist for layout purposes.
47#[derive(Debug)]
48pub struct LayoutGraph {
49    pub nodes: Vec<LayoutNode>,
50    pub edges: Vec<LayoutEdge>,
51    /// Map from component ID to node index.
52    pub id_to_node: BTreeMap<String, usize>,
53    /// Index of the `in` anchor node.
54    pub in_node: usize,
55    /// Index of the `out` anchor node.
56    pub out_node: usize,
57    /// Index of the `vcc` anchor node.
58    pub vcc_node: usize,
59    /// Index of the `gnd` anchor node.
60    pub gnd_node: usize,
61    /// Merged net groups: each group is a set of pins that are connected.
62    pub net_groups: Vec<NetGroup>,
63    /// Monitor component IDs from the pedal definition.
64    pub monitor_ids: Vec<String>,
65}
66
67/// A group of pins that are all connected (one electrical net).
68#[derive(Debug, Clone)]
69pub struct NetGroup {
70    /// Net index (0-based).
71    pub index: usize,
72    /// Optional net name (from reserved pins or named nodes).
73    pub name: Option<String>,
74    /// All pins in this net.
75    pub pins: Vec<Pin>,
76    /// Component IDs that this net touches.
77    pub component_ids: Vec<String>,
78}
79
80impl LayoutGraph {
81    /// Build a layout graph from a parsed pedal definition.
82    pub fn from_pedal(pedal: &PedalDef) -> Self {
83        let mut nodes = Vec::new();
84        let mut id_to_node = BTreeMap::new();
85
86        // Create anchor nodes for in, out, vcc, gnd
87        let anchors = ["__in", "__out", "__vcc", "__gnd"];
88        for name in &anchors {
89            let id = nodes.len();
90            id_to_node.insert(name.to_string(), id);
91            nodes.push(LayoutNode {
92                id,
93                comp_id: name.to_string(),
94                comp: None,
95                is_anchor: true,
96            });
97        }
98        let in_node = id_to_node["__in"];
99        let out_node = id_to_node["__out"];
100        let vcc_node = id_to_node["__vcc"];
101        let gnd_node = id_to_node["__gnd"];
102
103        // Create a node for each component
104        for comp in &pedal.components {
105            let id = nodes.len();
106            id_to_node.insert(comp.id.clone(), id);
107            nodes.push(LayoutNode {
108                id,
109                comp_id: comp.id.clone(),
110                comp: Some(comp.clone()),
111                is_anchor: false,
112            });
113        }
114
115        // Build merged net groups using union-find on pin connectivity
116        let net_groups = build_net_groups(&pedal.nets);
117
118        // Build edges from net groups
119        let edges = build_edges(&net_groups, &nodes, &id_to_node, pedal);
120
121        // Collect monitor IDs
122        let monitor_ids = pedal.monitors.iter().map(|m| m.component.clone()).collect();
123
124        LayoutGraph {
125            nodes,
126            edges,
127            id_to_node,
128            in_node,
129            out_node,
130            vcc_node,
131            gnd_node,
132            net_groups,
133            monitor_ids,
134        }
135    }
136
137    /// Get all outgoing edges from a node.
138    pub fn outgoing(&self, node_id: usize) -> Vec<&LayoutEdge> {
139        self.edges.iter().filter(|e| e.from == node_id).collect()
140    }
141
142    /// Get all incoming edges to a node.
143    pub fn incoming(&self, node_id: usize) -> Vec<&LayoutEdge> {
144        self.edges.iter().filter(|e| e.to == node_id).collect()
145    }
146
147    /// Get all neighbors (both directions) of a node.
148    pub fn neighbors(&self, node_id: usize) -> BTreeSet<usize> {
149        let mut result = BTreeSet::new();
150        for e in &self.edges {
151            if e.from == node_id {
152                result.insert(e.to);
153            }
154            if e.to == node_id {
155                result.insert(e.from);
156            }
157        }
158        result
159    }
160
161    /// Check if a node connects to vcc (directly or through supply components).
162    pub fn connects_to_vcc(&self, node_id: usize) -> bool {
163        self.edges.iter().any(|e| {
164            (e.from == node_id && e.to == self.vcc_node)
165                || (e.to == node_id && e.from == self.vcc_node)
166        })
167    }
168
169    /// Check if a node connects to gnd (directly or through ground components).
170    pub fn connects_to_gnd(&self, node_id: usize) -> bool {
171        self.edges.iter().any(|e| {
172            (e.from == node_id && e.to == self.gnd_node)
173                || (e.to == node_id && e.from == self.gnd_node)
174        })
175    }
176
177    /// Get the component kind for a node, if it has one.
178    pub fn node_kind(&self, node_id: usize) -> Option<&dyn Component> {
179        self.nodes[node_id].comp.as_ref().map(|c| c.kind.as_ref())
180    }
181
182    /// Check if a node is an active device (tube, transistor, op-amp).
183    pub fn is_active_device(&self, node_id: usize) -> bool {
184        self.node_kind(node_id)
185            .map_or(false, |k| k.is_gain_device() || k.op_amp_type().is_some())
186    }
187
188    /// Check if a component is a simple passive (R, C, L — not pot/transformer).
189    pub fn is_passive(&self, node_id: usize) -> bool {
190        self.node_kind(node_id)
191            .map_or(false, |k| k.is_simple_passive() && !k.is_pot())
192    }
193}
194
195// ---------------------------------------------------------------------------
196// Net group construction
197// ---------------------------------------------------------------------------
198
199/// Map a reserved pin name to the corresponding anchor node ID.
200fn reserved_to_anchor(name: &str) -> Option<&'static str> {
201    match name {
202        "in" => Some("__in"),
203        "out" => Some("__out"),
204        "vcc" => Some("__vcc"),
205        "gnd" => Some("__gnd"),
206        _ => None,
207    }
208}
209
210/// Extract the component ID from a pin reference.
211fn pin_component_id(pin: &Pin) -> Option<String> {
212    match pin {
213        Pin::Reserved(name) => reserved_to_anchor(name).map(|s| s.to_string()),
214        Pin::ComponentPin { component, .. } => Some(component.clone()),
215        // For Fork, use the switch component as the main reference
216        Pin::Fork { switch, .. } => Some(switch.clone()),
217        // SubcircuitPort references are resolved before layout
218        Pin::SubcircuitPort { subcircuit, .. } => Some(subcircuit.clone()),
219    }
220}
221
222fn build_net_groups(nets: &[NetDef]) -> Vec<NetGroup> {
223    // Merge nets that share pins (same logic as kicad.rs build_net_map)
224    let mut groups: Vec<Vec<Pin>> = Vec::new();
225
226    for net in nets {
227        let mut all_pins: Vec<Pin> = vec![net.from.clone()];
228        all_pins.extend(net.to.iter().cloned());
229
230        let mut merge_indices: Vec<usize> = Vec::new();
231        for (i, group) in groups.iter().enumerate() {
232            if all_pins.iter().any(|p| group.contains(p)) {
233                merge_indices.push(i);
234            }
235        }
236
237        if merge_indices.is_empty() {
238            groups.push(all_pins);
239        } else {
240            merge_indices.sort();
241            let target = merge_indices[0];
242            for &idx in merge_indices.iter().skip(1).rev() {
243                let g = groups.remove(idx);
244                groups[target].extend(g);
245            }
246            for p in all_pins {
247                if !groups[target].contains(&p) {
248                    groups[target].push(p);
249                }
250            }
251        }
252    }
253
254    groups
255        .into_iter()
256        .enumerate()
257        .map(|(i, pins)| {
258            let name = pins.iter().find_map(|p| {
259                if let Pin::Reserved(n) = p {
260                    Some(n.clone())
261                } else {
262                    None
263                }
264            });
265            let mut comp_ids: Vec<String> =
266                pins.iter().filter_map(|p| pin_component_id(p)).collect();
267            comp_ids.sort();
268            comp_ids.dedup();
269            NetGroup {
270                index: i,
271                name,
272                pins,
273                component_ids: comp_ids,
274            }
275        })
276        .collect()
277}
278
279// ---------------------------------------------------------------------------
280// Edge construction
281// ---------------------------------------------------------------------------
282
283fn build_edges(
284    net_groups: &[NetGroup],
285    nodes: &[LayoutNode],
286    id_to_node: &BTreeMap<String, usize>,
287    pedal: &PedalDef,
288) -> Vec<LayoutEdge> {
289    let mut edges = Vec::new();
290    let mut seen_pairs: BTreeSet<(usize, usize)> = BTreeSet::new();
291
292    // For each net group, create edges between all component pairs in the net
293    for ng in net_groups {
294        let node_ids: Vec<usize> = ng
295            .component_ids
296            .iter()
297            .filter_map(|cid| id_to_node.get(cid).copied())
298            .collect();
299
300        let is_supply_net = ng.name.as_deref() == Some("vcc") || ng.name.as_deref() == Some("gnd");
301        let is_signal = ng.name.as_deref() == Some("in") || ng.name.as_deref() == Some("out");
302
303        // Create directed edges based on pin direction inference
304        for &from_id in &node_ids {
305            for &to_id in &node_ids {
306                if from_id == to_id {
307                    continue;
308                }
309
310                // Determine direction from pin roles
311                let direction = infer_edge_direction(from_id, to_id, ng, nodes, pedal);
312
313                match direction {
314                    EdgeDir::Forward => {
315                        if seen_pairs.insert((from_id, to_id)) {
316                            edges.push(LayoutEdge {
317                                from: from_id,
318                                to: to_id,
319                                net_name: ng.name.clone(),
320                                signal_path: is_signal,
321                                is_supply: is_supply_net,
322                                is_feedback: false,
323                            });
324                        }
325                    }
326                    EdgeDir::Backward => {
327                        if seen_pairs.insert((to_id, from_id)) {
328                            edges.push(LayoutEdge {
329                                from: to_id,
330                                to: from_id,
331                                net_name: ng.name.clone(),
332                                signal_path: is_signal,
333                                is_supply: is_supply_net,
334                                is_feedback: false,
335                            });
336                        }
337                    }
338                    EdgeDir::Both => {
339                        // For bidirectional, pick the lower-id → higher-id direction
340                        let (a, b) = if from_id < to_id {
341                            (from_id, to_id)
342                        } else {
343                            (to_id, from_id)
344                        };
345                        if seen_pairs.insert((a, b)) {
346                            edges.push(LayoutEdge {
347                                from: a,
348                                to: b,
349                                net_name: ng.name.clone(),
350                                signal_path: is_signal,
351                                is_supply: is_supply_net,
352                                is_feedback: false,
353                            });
354                        }
355                    }
356                }
357            }
358        }
359    }
360
361    // Mark feedback edges (back-edges in topological order)
362    mark_feedback_edges(&mut edges, nodes.len(), id_to_node.get("__in").copied());
363
364    edges
365}
366
367enum EdgeDir {
368    Forward,
369    Backward,
370    Both,
371}
372
373fn infer_edge_direction(
374    from_id: usize,
375    to_id: usize,
376    ng: &NetGroup,
377    nodes: &[LayoutNode],
378    _pedal: &PedalDef,
379) -> EdgeDir {
380    let from_node = &nodes[from_id];
381    let to_node = &nodes[to_id];
382
383    // Check pin directions for from_node in this net
384    let from_dir = get_pin_direction_in_net(from_node, ng);
385    let to_dir = get_pin_direction_in_net(to_node, ng);
386
387    match (from_dir, to_dir) {
388        (PinDirection::Output, PinDirection::Input) => EdgeDir::Forward,
389        (PinDirection::Input, PinDirection::Output) => EdgeDir::Backward,
390        (PinDirection::Output, _) => EdgeDir::Forward,
391        (_, PinDirection::Input) => EdgeDir::Forward,
392        (PinDirection::Input, _) => EdgeDir::Backward,
393        (_, PinDirection::Output) => EdgeDir::Backward,
394        _ => EdgeDir::Both,
395    }
396}
397
398fn get_pin_direction_in_net(node: &LayoutNode, ng: &NetGroup) -> PinDirection {
399    if node.is_anchor {
400        return match node.comp_id.as_str() {
401            "__in" => PinDirection::Output, // 'in' anchor sends signal out
402            "__out" => PinDirection::Input, // 'out' anchor receives signal
403            "__vcc" => PinDirection::Up,
404            "__gnd" => PinDirection::Down,
405            _ => PinDirection::Bidirectional,
406        };
407    }
408
409    let comp = match &node.comp {
410        Some(c) => c,
411        None => return PinDirection::Bidirectional,
412    };
413
414    // Find which pin of this component appears in the net
415    for pin in &ng.pins {
416        if let Pin::ComponentPin {
417            component,
418            pin: pin_name,
419        } = pin
420        {
421            if *component == node.comp_id {
422                return comp.kind.pin_direction(pin_name);
423            }
424        }
425    }
426
427    PinDirection::Bidirectional
428}
429
430/// Mark back-edges in the graph (feedback loops) using DFS from the input.
431fn mark_feedback_edges(edges: &mut [LayoutEdge], num_nodes: usize, start: Option<usize>) {
432    let start = match start {
433        Some(s) => s,
434        None => return,
435    };
436
437    // Build adjacency list
438    let mut adj: Vec<Vec<usize>> = vec![Vec::new(); num_nodes];
439    for (i, e) in edges.iter().enumerate() {
440        adj[e.from].push(i);
441    }
442
443    // DFS to find back-edges
444    let mut visited = vec![false; num_nodes];
445    let mut in_stack = vec![false; num_nodes];
446    let mut stack = vec![(start, 0usize)];
447    visited[start] = true;
448    in_stack[start] = true;
449
450    while let Some((node, edge_idx)) = stack.last_mut() {
451        let node = *node;
452        if *edge_idx >= adj[node].len() {
453            in_stack[node] = false;
454            stack.pop();
455            continue;
456        }
457        let ei = adj[node][*edge_idx];
458        *edge_idx += 1;
459        let target = edges[ei].to;
460        if in_stack[target] {
461            // Back-edge found — mark as feedback
462            edges[ei].is_feedback = true;
463        } else if !visited[target] {
464            visited[target] = true;
465            in_stack[target] = true;
466            stack.push((target, 0));
467        }
468    }
469}
470
471#[cfg(test)]
472mod tests {
473    use super::*;
474    use pedalkernel::compiler::components::*;
475    use pedalkernel::dsl::*;
476
477    fn simple_pedal() -> PedalDef {
478        PedalDef {
479            name: "Test".into(),
480            supplies: vec![],
481            components: vec![
482                ComponentDef {
483                    id: "R1".into(),
484                    kind: Box::new(Resistor { value: 4700.0 }),
485                },
486                ComponentDef {
487                    id: "C1".into(),
488                    kind: Box::new(Capacitor {
489                        config: CapConfig::new(100e-9),
490                    }),
491                },
492            ],
493            nets: vec![
494                NetDef {
495                    from: Pin::Reserved("in".into()),
496                    to: vec![Pin::ComponentPin {
497                        component: "C1".into(),
498                        pin: "a".into(),
499                    }],
500                },
501                NetDef {
502                    from: Pin::ComponentPin {
503                        component: "C1".into(),
504                        pin: "b".into(),
505                    },
506                    to: vec![Pin::ComponentPin {
507                        component: "R1".into(),
508                        pin: "a".into(),
509                    }],
510                },
511                NetDef {
512                    from: Pin::ComponentPin {
513                        component: "R1".into(),
514                        pin: "b".into(),
515                    },
516                    to: vec![Pin::Reserved("out".into())],
517                },
518            ],
519            controls: vec![],
520            trims: vec![],
521            monitors: vec![],
522            sidechains: vec![],
523            mirrors: Default::default(),
524            subtitle: None,
525            calibrate: false,
526            subcircuits: vec![],
527            ports: vec![],
528            init_hints: vec![],
529            uses: vec![],
530        }
531    }
532
533    #[test]
534    fn graph_has_correct_node_count() {
535        let g = LayoutGraph::from_pedal(&simple_pedal());
536        // 4 anchors + 2 components = 6
537        assert_eq!(g.nodes.len(), 6);
538    }
539
540    #[test]
541    fn graph_has_edges() {
542        let g = LayoutGraph::from_pedal(&simple_pedal());
543        assert!(!g.edges.is_empty());
544    }
545
546    #[test]
547    fn pin_direction_triode() {
548        let triode = Triode {
549            model: "12AX7".into(),
550        };
551        assert_eq!(triode.pin_direction("grid"), PinDirection::Input);
552        assert_eq!(triode.pin_direction("plate"), PinDirection::Output);
553        assert_eq!(triode.pin_direction("cathode"), PinDirection::Down);
554    }
555
556    #[test]
557    fn pin_direction_opamp() {
558        let opamp = OpAmp {
559            op_type: OpAmpType::Jrc4558,
560        };
561        assert_eq!(opamp.pin_direction("pos"), PinDirection::Input);
562        assert_eq!(opamp.pin_direction("out"), PinDirection::Output);
563    }
564}