1use 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 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
117fn 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
173fn 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 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
268pub 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 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 #[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
351const 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
360fn 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}