1use crate::graph::LayoutGraph;
8use crate::groups::{FunctionalGroup, GroupKind};
9use std::collections::{BTreeMap, BTreeSet, VecDeque};
10
11#[derive(Debug, Clone)]
13pub struct ColumnAssignment {
14 pub group_columns: BTreeMap<String, usize>,
16 pub num_columns: usize,
18 pub column_widths: Vec<f32>,
20}
21
22pub fn assign_columns(graph: &LayoutGraph, groups: &[FunctionalGroup]) -> ColumnAssignment {
24 let group_graph = build_group_graph(graph, groups);
26
27 let topo_order = topological_sort_groups(&group_graph, groups);
29
30 let layers = longest_path_layering(&group_graph, &topo_order, groups);
32
33 let layers = compact_layers(layers, groups);
35
36 let num_columns = layers.values().copied().max().map(|m| m + 1).unwrap_or(1);
37
38 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
48struct GroupGraph {
54 adj: Vec<BTreeSet<usize>>,
56 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 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 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 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 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 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 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 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 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
194fn compact_layers(
196 mut layers: BTreeMap<String, usize>,
197 groups: &[FunctionalGroup],
198) -> BTreeMap<String, usize> {
199 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 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 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}