Skip to main content

piw/
layout.rs

1//! Faithful port of `src/render/graph.ts`: expand switch edges, classify
2//! back edges via DFS, longest-path layering with terminal tail pull-down,
3//! virtual pass-through cells, and barycenter ordering. Behavior is pinned
4//! against the golden fixtures in `fixtures/layout/` — any divergence is a
5//! bug in this port, not a stylistic choice.
6
7use crate::state::types::{DefinitionSnapshot, EdgeDef};
8use serde::{Deserialize, Serialize};
9use std::collections::{HashMap, HashSet};
10
11#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
12pub struct GraphEdge {
13    #[serde(rename = "edgeId")]
14    pub edge_id: String,
15    pub from: String,
16    pub to: String,
17    #[serde(skip_serializing_if = "Option::is_none")]
18    pub label: Option<String>,
19    #[serde(rename = "isBackEdge")]
20    pub is_back_edge: bool,
21}
22
23#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
24#[serde(tag = "kind", rename_all = "lowercase")]
25pub enum GraphCell {
26    Node {
27        #[serde(rename = "nodeId")]
28        node_id: String,
29    },
30    Virtual {
31        #[serde(rename = "edgeId")]
32        edge_id: String,
33    },
34}
35
36impl GraphCell {
37    pub fn is_node(&self) -> bool {
38        matches!(self, GraphCell::Node { .. })
39    }
40}
41
42#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
43pub struct GraphSegment {
44    #[serde(rename = "edgeId")]
45    pub edge_id: String,
46    pub rank: usize,
47    #[serde(rename = "fromCell")]
48    pub from_cell: usize,
49    #[serde(rename = "toCell")]
50    pub to_cell: usize,
51    #[serde(skip_serializing_if = "Option::is_none")]
52    pub label: Option<String>,
53}
54
55#[derive(Debug, Clone, PartialEq)]
56pub struct GraphLayout {
57    pub ranks: Vec<Vec<GraphCell>>,
58    pub edges: Vec<GraphEdge>,
59    pub segments: Vec<GraphSegment>,
60    pub rank_of_node: HashMap<String, usize>,
61}
62
63pub fn expand_edges(snapshot: &DefinitionSnapshot) -> Vec<GraphEdge> {
64    let mut edges = Vec::new();
65    for (index, edge) in snapshot.edges.iter().enumerate() {
66        match edge {
67            EdgeDef::Simple { from, to } => edges.push(GraphEdge {
68                edge_id: format!("{from}->{to}#{index}.0"),
69                from: from.clone(),
70                to: to.clone(),
71                label: None,
72                is_back_edge: false,
73            }),
74            EdgeDef::Switch { from, switch } => {
75                for (branch, (case_key, target)) in switch.cases.iter().enumerate() {
76                    let target = target.as_str().unwrap_or_default().to_string();
77                    edges.push(GraphEdge {
78                        edge_id: format!("{from}->{target}#{index}.{branch}"),
79                        from: from.clone(),
80                        to: target,
81                        // Case keys are author-controlled text from the run;
82                        // scrub them so drawing a label can't emit escapes.
83                        label: Some(crate::format::sanitize_text(case_key)),
84                        is_back_edge: false,
85                    });
86                }
87            }
88        }
89    }
90    edges
91}
92
93fn bfs_order(snapshot: &DefinitionSnapshot, edges: &[GraphEdge]) -> Vec<String> {
94    let mut queue = std::collections::VecDeque::from([snapshot.start_at.clone()]);
95    let mut visited = HashSet::new();
96    let mut ordered = Vec::new();
97    while let Some(node_id) = queue.pop_front() {
98        if !visited.insert(node_id.clone()) {
99            continue;
100        }
101        ordered.push(node_id.clone());
102        for edge in edges {
103            if edge.from == node_id {
104                queue.push_back(edge.to.clone());
105            }
106        }
107    }
108    for node_id in snapshot.node_ids() {
109        if visited.insert(node_id.to_string()) {
110            ordered.push(node_id.to_string());
111        }
112    }
113    ordered
114}
115
116/// DFS cycle detection: an edge is a back edge only when it closes a real
117/// cycle (its target is an ancestor on the DFS stack).
118fn mark_back_edges(edges: &mut [GraphEdge], ordered_node_ids: &[String]) {
119    #[derive(Clone, Copy, PartialEq)]
120    enum Color {
121        Gray,
122        Black,
123    }
124    fn visit(node_id: &str, edges: &mut [GraphEdge], color: &mut HashMap<String, Color>) {
125        color.insert(node_id.to_string(), Color::Gray);
126        for index in 0..edges.len() {
127            if edges[index].from != node_id {
128                continue;
129            }
130            let target = edges[index].to.clone();
131            match color.get(&target) {
132                Some(Color::Gray) => edges[index].is_back_edge = true,
133                None => visit(&target, edges, color),
134                Some(Color::Black) => {}
135            }
136        }
137        color.insert(node_id.to_string(), Color::Black);
138    }
139    let mut color = HashMap::new();
140    for node_id in ordered_node_ids {
141        if !color.contains_key(node_id) {
142            visit(node_id, edges, &mut color);
143        }
144    }
145}
146
147fn compute_longest_levels(
148    start_at: &str,
149    ordered_node_ids: &[String],
150    forward_edges: &[&GraphEdge],
151) -> HashMap<String, i64> {
152    let mut levels = HashMap::from([(start_at.to_string(), 0_i64)]);
153    for _pass in 0..=ordered_node_ids.len() {
154        let mut changed = false;
155        for edge in forward_edges {
156            let Some(&from_level) = levels.get(&edge.from) else {
157                continue;
158            };
159            let proposed = from_level + 1;
160            if proposed > levels.get(&edge.to).copied().unwrap_or(-1) {
161                levels.insert(edge.to.clone(), proposed);
162                changed = true;
163            }
164        }
165        if !changed {
166            break;
167        }
168    }
169    levels
170}
171
172/// Nodes on a straight single-path tail into a terminal node sink to the
173/// bottom ranks.
174fn compute_tail_depths(
175    ordered_node_ids: &[String],
176    forward_edges: &[&GraphEdge],
177    terminal_node_ids: &HashSet<String>,
178) -> HashMap<String, i64> {
179    let mut outgoing: HashMap<&str, Vec<&str>> = HashMap::new();
180    for edge in forward_edges {
181        outgoing.entry(&edge.from).or_default().push(&edge.to);
182    }
183    fn visit(
184        node_id: &str,
185        outgoing: &HashMap<&str, Vec<&str>>,
186        terminal: &HashSet<String>,
187        memo: &mut HashMap<String, Option<i64>>,
188    ) -> Option<i64> {
189        if let Some(existing) = memo.get(node_id) {
190            return *existing;
191        }
192        if terminal.contains(node_id) {
193            memo.insert(node_id.to_string(), Some(0));
194            return Some(0);
195        }
196        let targets = outgoing.get(node_id).map(Vec::as_slice).unwrap_or(&[]);
197        if targets.len() != 1 {
198            memo.insert(node_id.to_string(), None);
199            return None;
200        }
201        // Break potential cycles before recursing.
202        memo.insert(node_id.to_string(), None);
203        let child = targets[0].to_string();
204        let depth = visit(&child, outgoing, terminal, memo).map(|value| value + 1);
205        memo.insert(node_id.to_string(), depth);
206        depth
207    }
208    let mut memo = HashMap::new();
209    for node_id in ordered_node_ids {
210        visit(node_id, &outgoing, terminal_node_ids, &mut memo);
211    }
212    memo.into_iter()
213        .filter_map(|(node_id, depth)| depth.map(|value| (node_id, value)))
214        .collect()
215}
216
217fn compute_node_ranks(
218    snapshot: &DefinitionSnapshot,
219    ordered_node_ids: &[String],
220    edges: &[GraphEdge],
221) -> HashMap<String, usize> {
222    let forward_edges: Vec<&GraphEdge> = edges.iter().filter(|edge| !edge.is_back_edge).collect();
223    let longest = compute_longest_levels(&snapshot.start_at, ordered_node_ids, &forward_edges);
224    let mut outgoing_counts: HashMap<&str, usize> = HashMap::new();
225    for edge in &forward_edges {
226        *outgoing_counts.entry(&edge.from).or_default() += 1;
227    }
228    let terminal_node_ids: HashSet<String> = ordered_node_ids
229        .iter()
230        .filter(|node_id| outgoing_counts.get(node_id.as_str()).copied().unwrap_or(0) == 0)
231        .cloned()
232        .collect();
233    let tail_depths = compute_tail_depths(ordered_node_ids, &forward_edges, &terminal_node_ids);
234
235    let mut rank_of_node: HashMap<String, i64> = HashMap::new();
236    let mut fallback = longest.values().copied().max().unwrap_or(0).max(0);
237    for node_id in ordered_node_ids {
238        match longest.get(node_id) {
239            Some(&base) => {
240                rank_of_node.insert(node_id.clone(), base);
241            }
242            None => {
243                fallback += 1;
244                rank_of_node.insert(node_id.clone(), fallback);
245            }
246        }
247    }
248    let max_rank = rank_of_node.values().copied().max().unwrap_or(0).max(0);
249    for node_id in ordered_node_ids {
250        if let Some(&tail_depth) = tail_depths.get(node_id) {
251            let current = rank_of_node.get(node_id).copied().unwrap_or(0);
252            rank_of_node.insert(node_id.clone(), current.max(max_rank - tail_depth));
253        }
254    }
255    rank_of_node
256        .into_iter()
257        .map(|(node_id, rank)| (node_id, rank.max(0) as usize))
258        .collect()
259}
260
261#[derive(Clone, Copy)]
262struct CellRef {
263    rank: usize,
264    index: usize,
265}
266
267/// Build ranks with virtual pass-through cells so every forward edge connects
268/// adjacent ranks, then order cells within ranks by neighbor barycenter.
269pub fn layout_graph(snapshot: &DefinitionSnapshot) -> GraphLayout {
270    let mut edges = expand_edges(snapshot);
271    let ordered_node_ids = bfs_order(snapshot, &edges);
272    mark_back_edges(&mut edges, &ordered_node_ids);
273    let rank_of_node = compute_node_ranks(snapshot, &ordered_node_ids, &edges);
274
275    let rank_count = rank_of_node.values().copied().max().unwrap_or(0) + 1;
276    let mut ranks: Vec<Vec<GraphCell>> = vec![Vec::new(); rank_count];
277    let mut cell_ref: HashMap<String, CellRef> = HashMap::new();
278    for node_id in &ordered_node_ids {
279        let rank = rank_of_node.get(node_id).copied().unwrap_or(0);
280        cell_ref.insert(
281            node_id.clone(),
282            CellRef {
283                rank,
284                index: ranks[rank].len(),
285            },
286        );
287        ranks[rank].push(GraphCell::Node {
288            node_id: node_id.clone(),
289        });
290    }
291
292    // Chain each long forward edge through virtual cells in intermediate ranks.
293    let mut segments: Vec<GraphSegment> = Vec::new();
294    for edge in &edges {
295        if edge.is_back_edge {
296            continue;
297        }
298        let (Some(&from_rank), Some(&to_rank)) =
299            (rank_of_node.get(&edge.from), rank_of_node.get(&edge.to))
300        else {
301            continue;
302        };
303        if to_rank <= from_rank {
304            continue;
305        }
306        let mut previous = cell_ref[&edge.from];
307        // Index loop kept deliberately: `rank` also feeds the CellRef chain.
308        #[allow(clippy::needless_range_loop)]
309        for rank in (from_rank + 1)..to_rank {
310            let index = ranks[rank].len();
311            ranks[rank].push(GraphCell::Virtual {
312                edge_id: edge.edge_id.clone(),
313            });
314            segments.push(GraphSegment {
315                edge_id: edge.edge_id.clone(),
316                rank: previous.rank,
317                from_cell: previous.index,
318                to_cell: index,
319                label: if previous.rank == from_rank {
320                    edge.label.clone()
321                } else {
322                    None
323                },
324            });
325            previous = CellRef { rank, index };
326        }
327        let target = cell_ref[&edge.to];
328        segments.push(GraphSegment {
329            edge_id: edge.edge_id.clone(),
330            rank: previous.rank,
331            from_cell: previous.index,
332            to_cell: target.index,
333            label: if previous.rank == from_rank {
334                edge.label.clone()
335            } else {
336                None
337            },
338        });
339    }
340
341    order_ranks_by_barycenter(&mut ranks, &mut segments);
342    GraphLayout {
343        ranks,
344        edges,
345        segments,
346        rank_of_node,
347    }
348}
349
350/// JS `Number.MAX_SAFE_INTEGER`, used as the "no neighbors" score.
351const MAX_SAFE_INTEGER: f64 = 9_007_199_254_740_991.0;
352
353#[derive(Clone, Copy, PartialEq)]
354enum Direction {
355    Down,
356    Up,
357}
358
359/// Sweep up and down, sorting each rank by the mean position of neighbors.
360fn order_ranks_by_barycenter(ranks: &mut [Vec<GraphCell>], segments: &mut [GraphSegment]) {
361    fn reindex(
362        ranks: &mut [Vec<GraphCell>],
363        segments: &mut [GraphSegment],
364        rank: usize,
365        order: &[usize],
366    ) {
367        let cells = std::mem::take(&mut ranks[rank]);
368        let mut inverse = vec![0usize; order.len()];
369        for (new_index, &old_index) in order.iter().enumerate() {
370            inverse[old_index] = new_index;
371        }
372        ranks[rank] = order.iter().map(|&old| cells[old].clone()).collect();
373        for segment in segments.iter_mut() {
374            if segment.rank == rank {
375                segment.from_cell = inverse[segment.from_cell];
376            }
377            if rank > 0 && segment.rank == rank - 1 {
378                segment.to_cell = inverse[segment.to_cell];
379            }
380        }
381    }
382
383    fn sort_rank(
384        ranks: &mut [Vec<GraphCell>],
385        segments: &mut [GraphSegment],
386        rank: usize,
387        direction: Direction,
388    ) {
389        let len = ranks[rank].len();
390        if len < 2 {
391            return;
392        }
393        let mut scores: Vec<(usize, f64)> = Vec::with_capacity(len);
394        for index in 0..len {
395            let neighbors: Vec<usize> = segments
396                .iter()
397                .filter(|segment| match direction {
398                    Direction::Down => {
399                        rank > 0 && segment.rank == rank - 1 && segment.to_cell == index
400                    }
401                    Direction::Up => segment.rank == rank && segment.from_cell == index,
402                })
403                .map(|segment| match direction {
404                    Direction::Down => segment.from_cell,
405                    Direction::Up => segment.to_cell,
406                })
407                .collect();
408            let score = if neighbors.is_empty() {
409                MAX_SAFE_INTEGER
410            } else {
411                neighbors.iter().sum::<usize>() as f64 / neighbors.len() as f64
412            };
413            scores.push((index, score));
414        }
415        let mut order: Vec<(usize, f64)> = scores.clone();
416        order.sort_by(|left, right| {
417            left.1
418                .partial_cmp(&right.1)
419                .unwrap_or(std::cmp::Ordering::Equal)
420                .then(left.0.cmp(&right.0))
421        });
422        let order: Vec<usize> = order.into_iter().map(|(index, _)| index).collect();
423        if order.iter().enumerate().any(|(new, &old)| new != old) {
424            reindex(ranks, segments, rank, &order);
425        }
426    }
427
428    for _pass in 0..4 {
429        for rank in 1..ranks.len() {
430            sort_rank(ranks, segments, rank, Direction::Down);
431        }
432        for rank in (0..ranks.len().saturating_sub(1)).rev() {
433            sort_rank(ranks, segments, rank, Direction::Up);
434        }
435    }
436}
437
438#[cfg(test)]
439mod tests {
440    use super::*;
441
442    #[test]
443    fn switch_case_labels_are_scrubbed_of_escapes() {
444        let snapshot: DefinitionSnapshot = serde_json::from_value(serde_json::json!({
445            "schema": "pi-workflows.workflow.v1",
446            "name": "test",
447            "startAt": "a",
448            "nodes": { "a": { "nodeType": "agent" }, "b": { "nodeType": "agent" } },
449            "edges": [
450                { "from": "a", "switch": { "on": "x", "cases": { "\u{1b}]52;c;evil\u{7}ok": "b" } } },
451            ],
452        }))
453        .unwrap();
454        let edges = expand_edges(&snapshot);
455        assert_eq!(edges[0].label.as_deref(), Some("]52;c;evilok"));
456    }
457}