Skip to main content

helix_graph_algorithms/algorithms/
layout.rs

1use std::collections::BTreeSet;
2use std::num::NonZeroUsize;
3
4use serde::{Deserialize, Serialize};
5
6use crate::{Graph, GraphError, NodeId, PositiveFiniteF64};
7
8/// Deterministic Fruchterman-Reingold layout options.
9#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
10pub struct LayoutOptions {
11    /// Optional ideal edge distance. `None` uses `sqrt(1 / node_count)`.
12    pub k: Option<PositiveFiniteF64>,
13    /// Number of force/cooling iterations.
14    pub iterations: NonZeroUsize,
15    /// Deterministic position seed.
16    pub seed: u64,
17    /// Whether edge attraction multiplies selected edge weights.
18    pub weighted: bool,
19    /// Optional deterministic position overrides. Nodes without an override
20    /// retain their seeded initial position.
21    pub initial_positions: Vec<NodePosition>,
22}
23
24impl Default for LayoutOptions {
25    fn default() -> Self {
26        Self {
27            k: None,
28            iterations: NonZeroUsize::new(50).expect("50 is non-zero"),
29            seed: 42,
30            weighted: true,
31            initial_positions: Vec::new(),
32        }
33    }
34}
35
36/// One final two-dimensional node position.
37#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
38pub struct NodePosition {
39    /// External node identity.
40    pub node_id: NodeId,
41    /// Rescaled horizontal coordinate in `[-1, 1]`.
42    pub x: f64,
43    /// Rescaled vertical coordinate in `[-1, 1]`.
44    pub y: f64,
45}
46
47impl Graph {
48    /// Compute deterministic Fruchterman-Reingold positions.
49    ///
50    /// This exact kernel is quadratic in node count per iteration. It is
51    /// intentionally explicit rather than silently switching to an
52    /// approximate layout for larger graphs.
53    pub fn spring_layout(&self, options: LayoutOptions) -> Result<Vec<NodePosition>, GraphError> {
54        let mut positioned = BTreeSet::new();
55        for position in &options.initial_positions {
56            if !position.x.is_finite() || !position.y.is_finite() {
57                return Err(GraphError::InvalidOption(format!(
58                    "initial position for {} must be finite",
59                    position.node_id
60                )));
61            }
62            self.node_index(&position.node_id)?;
63            if !positioned.insert(position.node_id.clone()) {
64                return Err(GraphError::InvalidOption(format!(
65                    "duplicate initial position for {}",
66                    position.node_id
67                )));
68            }
69        }
70        match self.node_count() {
71            0 => return Ok(Vec::new()),
72            1 => {
73                return Ok(vec![NodePosition {
74                    node_id: self.node_id(0).clone(),
75                    x: 0.0,
76                    y: 0.0,
77                }]);
78            }
79            _ => {}
80        }
81
82        let k = options
83            .k
84            .map(PositiveFiniteF64::get)
85            .unwrap_or_else(|| (1.0 / self.node_count() as f64).sqrt());
86        let mut seed = options.seed;
87        let mut positions = (0..self.node_count())
88            .map(|_| (random_unit(&mut seed), random_unit(&mut seed)))
89            .collect::<Vec<_>>();
90        for position in &options.initial_positions {
91            let node = self.node_index(&position.node_id)?;
92            positions[node] = (position.x, position.y);
93        }
94        const MIN_DISTANCE: f64 = 1e-9;
95        for step in 0..options.iterations.get() {
96            let mut displacement = vec![(0.0, 0.0); self.node_count()];
97            for left in 0..self.node_count() {
98                for right in left + 1..self.node_count() {
99                    let mut dx = positions[left].0 - positions[right].0;
100                    let mut dy = positions[left].1 - positions[right].1;
101                    let mut distance = dx.hypot(dy);
102                    if distance < MIN_DISTANCE {
103                        let jitter = deterministic_jitter(left, right);
104                        dx = jitter.0;
105                        dy = jitter.1;
106                        distance = dx.hypot(dy);
107                    }
108                    let force = k * k / distance;
109                    let force_x = dx / distance * force;
110                    let force_y = dy / distance * force;
111                    displacement[left].0 += force_x;
112                    displacement[left].1 += force_y;
113                    displacement[right].0 -= force_x;
114                    displacement[right].1 -= force_y;
115                }
116            }
117            for edge in self.edges() {
118                let source = self.node_index(&edge.source)?;
119                let target = self.node_index(&edge.target)?;
120                if source == target {
121                    continue;
122                }
123                let dx = positions[source].0 - positions[target].0;
124                let dy = positions[source].1 - positions[target].1;
125                let distance = dx.hypot(dy).max(MIN_DISTANCE);
126                let weight = if options.weighted {
127                    edge.weight.unwrap_or(1.0)
128                } else {
129                    1.0
130                };
131                let force = distance * distance / k * weight;
132                let force_x = dx / distance * force;
133                let force_y = dy / distance * force;
134                displacement[source].0 -= force_x;
135                displacement[source].1 -= force_y;
136                displacement[target].0 += force_x;
137                displacement[target].1 += force_y;
138            }
139            let temperature = 0.1 * (1.0 - step as f64 / options.iterations.get() as f64).max(0.0);
140            for node in 0..self.node_count() {
141                let (dx, dy) = displacement[node];
142                let distance = dx.hypot(dy).max(MIN_DISTANCE);
143                positions[node].0 += dx / distance * distance.min(temperature);
144                positions[node].1 += dy / distance * distance.min(temperature);
145            }
146        }
147        rescale_positions(&mut positions);
148        Ok(positions
149            .into_iter()
150            .enumerate()
151            .map(|(node, (x, y))| NodePosition {
152                node_id: self.node_id(node).clone(),
153                x,
154                y,
155            })
156            .collect())
157    }
158}
159
160fn random_unit(seed: &mut u64) -> f64 {
161    *seed = seed
162        .wrapping_mul(6_364_136_223_846_793_005)
163        .wrapping_add(1_442_695_040_888_963_407);
164    (*seed >> 11) as f64 / ((1_u64 << 53) - 1) as f64
165}
166
167fn deterministic_jitter(left: usize, right: usize) -> (f64, f64) {
168    let angle = ((left.wrapping_mul(31) ^ right.wrapping_mul(17)) % 360) as f64
169        * std::f64::consts::PI
170        / 180.0;
171    (angle.cos() * 1e-6, angle.sin() * 1e-6)
172}
173
174fn rescale_positions(positions: &mut [(f64, f64)]) {
175    let mean_x = positions.iter().map(|position| position.0).sum::<f64>() / positions.len() as f64;
176    let mean_y = positions.iter().map(|position| position.1).sum::<f64>() / positions.len() as f64;
177    let scale = positions
178        .iter()
179        .map(|(x, y)| (x - mean_x).abs().max((y - mean_y).abs()))
180        .fold(0.0, f64::max);
181    if scale == 0.0 {
182        return;
183    }
184    for position in positions {
185        position.0 = (position.0 - mean_x) / scale;
186        position.1 = (position.1 - mean_y) / scale;
187    }
188}
189
190#[cfg(test)]
191mod tests {
192    use super::*;
193    use crate::{Edge, GraphKind, Node};
194
195    #[test]
196    fn layout_handles_empty_and_singleton_graphs() {
197        let empty = Graph::new(GraphKind::Graph, [], []).unwrap();
198        assert!(empty
199            .spring_layout(LayoutOptions::default())
200            .unwrap()
201            .is_empty());
202
203        let singleton = Graph::new(GraphKind::Graph, [Node::new("a")], []).unwrap();
204        assert_eq!(
205            singleton.spring_layout(LayoutOptions::default()).unwrap(),
206            [NodePosition {
207                node_id: "a".into(),
208                x: 0.0,
209                y: 0.0,
210            }]
211        );
212    }
213
214    #[test]
215    fn layout_is_finite_seeded_and_rescaled() {
216        let graph = Graph::new(
217            GraphKind::Graph,
218            [Node::new("a"), Node::new("b"), Node::new("c")],
219            [Edge::new("ab", "a", "b"), Edge::new("bc", "b", "c")],
220        )
221        .unwrap();
222        let first = graph.spring_layout(LayoutOptions::default()).unwrap();
223        let second = graph.spring_layout(LayoutOptions::default()).unwrap();
224        assert_eq!(first, second);
225        assert!(first.iter().all(|position| {
226            position.x.is_finite()
227                && position.y.is_finite()
228                && position.x.abs() <= 1.0
229                && position.y.abs() <= 1.0
230        }));
231    }
232
233    #[test]
234    fn layout_k_type_rejects_invalid_values() {
235        assert!(PositiveFiniteF64::new(0.0).is_err());
236        assert!(PositiveFiniteF64::new(f64::NAN).is_err());
237    }
238
239    #[test]
240    fn layout_accepts_partial_initial_positions_and_rejects_bad_ones() {
241        let graph = Graph::new(
242            GraphKind::Graph,
243            [Node::new("a"), Node::new("b")],
244            [Edge::new("ab", "a", "b")],
245        )
246        .unwrap();
247        let options = LayoutOptions {
248            initial_positions: vec![NodePosition {
249                node_id: "a".into(),
250                x: 0.25,
251                y: 0.75,
252            }],
253            ..LayoutOptions::default()
254        };
255        assert_eq!(graph.spring_layout(options.clone()).unwrap().len(), 2);
256        let invalid = LayoutOptions {
257            initial_positions: vec![NodePosition {
258                node_id: "missing".into(),
259                x: 0.0,
260                y: 0.0,
261            }],
262            ..options
263        };
264        assert!(matches!(
265            graph.spring_layout(invalid),
266            Err(GraphError::UnknownNode(_))
267        ));
268
269        let non_finite = LayoutOptions {
270            initial_positions: vec![NodePosition {
271                node_id: "a".into(),
272                x: f64::NAN,
273                y: 0.0,
274            }],
275            ..LayoutOptions::default()
276        };
277        assert!(matches!(
278            graph.spring_layout(non_finite),
279            Err(GraphError::InvalidOption(_))
280        ));
281        let duplicate = LayoutOptions {
282            initial_positions: vec![
283                NodePosition {
284                    node_id: "a".into(),
285                    x: 0.0,
286                    y: 0.0,
287                },
288                NodePosition {
289                    node_id: "a".into(),
290                    x: 1.0,
291                    y: 1.0,
292                },
293            ],
294            ..LayoutOptions::default()
295        };
296        assert!(matches!(
297            graph.spring_layout(duplicate),
298            Err(GraphError::InvalidOption(_))
299        ));
300    }
301
302    #[test]
303    fn layout_covers_coincident_unweighted_parallel_and_self_loop_forces() {
304        let graph = Graph::new(
305            GraphKind::MultiGraph,
306            [Node::new("a"), Node::new("b")],
307            [
308                Edge::new("aa", "a", "a").with_weight(10.0),
309                Edge::new("ab1", "a", "b").with_weight(2.0),
310                Edge::new("ab2", "a", "b").with_weight(3.0),
311            ],
312        )
313        .unwrap();
314        let positions = graph
315            .spring_layout(LayoutOptions {
316                weighted: false,
317                initial_positions: vec![
318                    NodePosition {
319                        node_id: "a".into(),
320                        x: 0.0,
321                        y: 0.0,
322                    },
323                    NodePosition {
324                        node_id: "b".into(),
325                        x: 0.0,
326                        y: 0.0,
327                    },
328                ],
329                ..LayoutOptions::default()
330            })
331            .unwrap();
332        assert!(positions.iter().all(|position| position.x.is_finite()));
333        let mut same = [(1.0, 1.0), (1.0, 1.0)];
334        rescale_positions(&mut same);
335        assert_eq!(same, [(1.0, 1.0), (1.0, 1.0)]);
336    }
337}