Skip to main content

pedalkernel_layout/
layering.rs

1//! Phase 3: Sugiyama layer (column) assignment.
2//!
3//! Assigns each functional group to a horizontal column based on signal flow
4//! order. Uses longest-path layering from the input anchor to ensure signal
5//! flows strictly left to right.
6
7use crate::graph::LayoutGraph;
8use crate::groups::{FunctionalGroup, GroupKind};
9use std::collections::{BTreeMap, BTreeSet, VecDeque};
10
11/// Column assignment for each functional group.
12#[derive(Debug, Clone)]
13pub struct ColumnAssignment {
14    /// Map from group name to column index (0 = leftmost).
15    pub group_columns: BTreeMap<String, usize>,
16    /// Total number of columns.
17    pub num_columns: usize,
18    /// Column widths (relative, normalized to sum=1.0 later).
19    pub column_widths: Vec<f32>,
20}
21
22/// Assign columns to functional groups using longest-path layering.
23pub fn assign_columns(graph: &LayoutGraph, groups: &[FunctionalGroup]) -> ColumnAssignment {
24    // Build a group-level graph: edges between groups based on signal flow
25    let group_graph = build_group_graph(graph, groups);
26
27    // Topological sort of groups, respecting signal flow
28    let topo_order = topological_sort_groups(&group_graph, groups);
29
30    // Longest-path layering: assign each group to its longest distance from input
31    let layers = longest_path_layering(&group_graph, &topo_order, groups);
32
33    // Merge single-component groups into adjacent layers if they're just coupling elements
34    let layers = compact_layers(layers, groups);
35
36    let num_columns = layers.values().copied().max().map(|m| m + 1).unwrap_or(1);
37
38    // Compute column widths based on group complexity
39    let column_widths = compute_column_widths(groups, &layers, num_columns);
40
41    ColumnAssignment {
42        group_columns: layers,
43        num_columns,
44        column_widths,
45    }
46}
47
48// ---------------------------------------------------------------------------
49// Group-level graph
50// ---------------------------------------------------------------------------
51
52/// Adjacency list for the group-level directed graph.
53struct GroupGraph {
54    /// Edges: group index → set of successor group indices.
55    adj: Vec<BTreeSet<usize>>,
56    /// Reverse edges: group index → set of predecessor group indices.
57    rev: Vec<BTreeSet<usize>>,
58}
59
60fn build_group_graph(graph: &LayoutGraph, groups: &[FunctionalGroup]) -> GroupGraph {
61    let n = groups.len();
62    let mut adj = vec![BTreeSet::new(); n];
63    let mut rev = vec![BTreeSet::new(); n];
64
65    // Map node ID → group index
66    let mut node_to_group: BTreeMap<usize, usize> = BTreeMap::new();
67    for (gi, g) in groups.iter().enumerate() {
68        for &m in &g.members {
69            node_to_group.insert(m, gi);
70        }
71    }
72
73    // For each edge in the layout graph, if it connects two different groups,
74    // add a group-level edge
75    for edge in &graph.edges {
76        if edge.is_feedback || edge.is_supply {
77            continue;
78        }
79        if let (Some(&from_g), Some(&to_g)) =
80            (node_to_group.get(&edge.from), node_to_group.get(&edge.to))
81        {
82            if from_g != to_g {
83                adj[from_g].insert(to_g);
84                rev[to_g].insert(from_g);
85            }
86        }
87    }
88
89    // Also add edges from anchor nodes to their connected groups
90    // Input anchor → first groups; last groups → output anchor
91    for edge in &graph.edges {
92        if edge.from == graph.in_node {
93            if let Some(&to_g) = node_to_group.get(&edge.to) {
94                // Mark input group as having no predecessors (it's the start)
95                // — already implicit from lack of reverse edges
96                let _ = to_g;
97            }
98        }
99    }
100
101    GroupGraph { adj, rev }
102}
103
104fn topological_sort_groups(gg: &GroupGraph, groups: &[FunctionalGroup]) -> Vec<usize> {
105    let n = groups.len();
106    let mut in_degree: Vec<usize> = vec![0; n];
107
108    for (gi, succs) in gg.adj.iter().enumerate() {
109        let _ = gi;
110        for &s in succs {
111            in_degree[s] += 1;
112        }
113    }
114
115    let mut queue: VecDeque<usize> = VecDeque::new();
116    for i in 0..n {
117        if in_degree[i] == 0 {
118            queue.push_back(i);
119        }
120    }
121
122    // Prioritize input sections first
123    let mut sorted_queue: Vec<usize> = queue.drain(..).collect();
124    sorted_queue.sort_by_key(|&gi| match groups[gi].kind {
125        GroupKind::InputSection => 0,
126        GroupKind::GainStage => 1,
127        GroupKind::OpAmpStage => 1,
128        GroupKind::ToneStack => 2,
129        GroupKind::PhaseInverter => 3,
130        GroupKind::PushPullOutput => 4,
131        GroupKind::OutputSection => 5,
132        GroupKind::Generic => 6,
133    });
134    for gi in sorted_queue {
135        queue.push_back(gi);
136    }
137
138    let mut order = Vec::with_capacity(n);
139    while let Some(gi) = queue.pop_front() {
140        order.push(gi);
141        for &succ in &gg.adj[gi] {
142            in_degree[succ] -= 1;
143            if in_degree[succ] == 0 {
144                queue.push_back(succ);
145            }
146        }
147    }
148
149    // If there are cycles (shouldn't happen after feedback removal), add remaining
150    for i in 0..n {
151        if !order.contains(&i) {
152            order.push(i);
153        }
154    }
155
156    order
157}
158
159fn longest_path_layering(
160    gg: &GroupGraph,
161    topo_order: &[usize],
162    groups: &[FunctionalGroup],
163) -> BTreeMap<String, usize> {
164    let n = groups.len();
165    let mut layer = vec![0usize; n];
166
167    // Process in topological order; each group's layer is max(predecessors' layers) + 1
168    for &gi in topo_order {
169        let max_pred = gg.rev[gi]
170            .iter()
171            .map(|&pred| layer[pred] + 1)
172            .max()
173            .unwrap_or(0);
174        layer[gi] = max_pred;
175    }
176
177    // Force input sections to column 0 and output sections to the last column
178    let max_layer = *layer.iter().max().unwrap_or(&0);
179    for (gi, g) in groups.iter().enumerate() {
180        if g.kind == GroupKind::InputSection {
181            layer[gi] = 0;
182        } else if g.kind == GroupKind::OutputSection {
183            layer[gi] = max_layer;
184        }
185    }
186
187    groups
188        .iter()
189        .enumerate()
190        .map(|(gi, g)| (g.name.clone(), layer[gi]))
191        .collect()
192}
193
194/// Compact layers by merging single-passive groups into adjacent columns.
195fn compact_layers(
196    mut layers: BTreeMap<String, usize>,
197    groups: &[FunctionalGroup],
198) -> BTreeMap<String, usize> {
199    // If a group has only one passive member and it bridges two columns,
200    // merge it with the earlier column.
201    for g in groups {
202        if g.members.len() == 1 && g.kind == GroupKind::Generic {
203            if let Some(&col) = layers.get(&g.name) {
204                if col > 0 {
205                    layers.insert(g.name.clone(), col.saturating_sub(1));
206                }
207            }
208        }
209    }
210
211    // Re-number columns to close gaps
212    let mut used: Vec<usize> = layers.values().copied().collect();
213    used.sort();
214    used.dedup();
215
216    let remap: BTreeMap<usize, usize> = used
217        .iter()
218        .enumerate()
219        .map(|(new, &old)| (old, new))
220        .collect();
221
222    for val in layers.values_mut() {
223        if let Some(&new) = remap.get(val) {
224            *val = new;
225        }
226    }
227
228    layers
229}
230
231fn compute_column_widths(
232    groups: &[FunctionalGroup],
233    layers: &BTreeMap<String, usize>,
234    num_columns: usize,
235) -> Vec<f32> {
236    let mut widths = vec![1.0f32; num_columns];
237
238    for g in groups {
239        if let Some(&col) = layers.get(&g.name) {
240            let complexity = match g.kind {
241                GroupKind::GainStage => 1.2,
242                GroupKind::ToneStack => 1.5 + 0.3 * (g.members.len() as f32 - 3.0).max(0.0),
243                GroupKind::PushPullOutput => 2.0,
244                GroupKind::PhaseInverter => 1.5,
245                GroupKind::OpAmpStage => 1.3,
246                GroupKind::InputSection => 0.8,
247                GroupKind::OutputSection => 0.8,
248                GroupKind::Generic => 1.0,
249            };
250            if complexity > widths[col] {
251                widths[col] = complexity;
252            }
253        }
254    }
255
256    // Normalize so widths sum to 1.0
257    let total: f32 = widths.iter().sum();
258    if total > 0.0 {
259        for w in &mut widths {
260            *w /= total;
261        }
262    }
263
264    widths
265}
266
267#[cfg(test)]
268mod tests {
269    use super::*;
270    use crate::graph::LayoutGraph;
271    use pedalkernel::compiler::components::*;
272    use pedalkernel::dsl::*;
273
274    #[test]
275    fn linear_chain_assigns_increasing_columns() {
276        let pedal = PedalDef {
277            name: "Test".into(),
278            supplies: vec![],
279            components: vec![
280                ComponentDef {
281                    id: "C1".into(),
282                    kind: Box::new(Capacitor {
283                        config: CapConfig::new(100e-9),
284                    }),
285                },
286                ComponentDef {
287                    id: "R1".into(),
288                    kind: Box::new(Resistor { value: 1e6 }),
289                },
290                ComponentDef {
291                    id: "V1".into(),
292                    kind: Box::new(Triode {
293                        model: "12AX7".into(),
294                    }),
295                },
296                ComponentDef {
297                    id: "R2".into(),
298                    kind: Box::new(Resistor { value: 100e3 }),
299                },
300            ],
301            nets: vec![
302                NetDef {
303                    from: Pin::Reserved("in".into()),
304                    to: vec![Pin::ComponentPin {
305                        component: "C1".into(),
306                        pin: "a".into(),
307                    }],
308                },
309                NetDef {
310                    from: Pin::ComponentPin {
311                        component: "C1".into(),
312                        pin: "b".into(),
313                    },
314                    to: vec![
315                        Pin::ComponentPin {
316                            component: "R1".into(),
317                            pin: "a".into(),
318                        },
319                        Pin::ComponentPin {
320                            component: "V1".into(),
321                            pin: "grid".into(),
322                        },
323                    ],
324                },
325                NetDef {
326                    from: Pin::Reserved("vcc".into()),
327                    to: vec![Pin::ComponentPin {
328                        component: "R2".into(),
329                        pin: "a".into(),
330                    }],
331                },
332                NetDef {
333                    from: Pin::ComponentPin {
334                        component: "R2".into(),
335                        pin: "b".into(),
336                    },
337                    to: vec![Pin::ComponentPin {
338                        component: "V1".into(),
339                        pin: "plate".into(),
340                    }],
341                },
342                NetDef {
343                    from: Pin::ComponentPin {
344                        component: "R1".into(),
345                        pin: "b".into(),
346                    },
347                    to: vec![Pin::Reserved("gnd".into())],
348                },
349                NetDef {
350                    from: Pin::ComponentPin {
351                        component: "V1".into(),
352                        pin: "plate".into(),
353                    },
354                    to: vec![Pin::Reserved("out".into())],
355                },
356            ],
357            controls: vec![],
358            trims: vec![],
359            monitors: vec![],
360            sidechains: vec![],
361            mirrors: Default::default(),
362            subtitle: None,
363            calibrate: false,
364            subcircuits: vec![],
365            ports: vec![],
366            init_hints: vec![],
367            uses: vec![],
368        };
369
370        let graph = LayoutGraph::from_pedal(&pedal);
371        let groups = crate::groups::detect_groups(&graph);
372        let cols = assign_columns(&graph, &groups);
373
374        assert!(cols.num_columns >= 1, "Should have at least one column");
375    }
376}