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)]
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 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
116fn 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
172fn 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 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
267pub 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 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 #[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
350const 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
359fn 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}