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