helix_graph_algorithms/algorithms/
layout.rs1use std::collections::BTreeSet;
2use std::num::NonZeroUsize;
3
4use serde::{Deserialize, Serialize};
5
6use crate::{Graph, GraphError, NodeId, PositiveFiniteF64};
7
8#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
10pub struct LayoutOptions {
11 pub k: Option<PositiveFiniteF64>,
13 pub iterations: NonZeroUsize,
15 pub seed: u64,
17 pub weighted: bool,
19 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#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
38pub struct NodePosition {
39 pub node_id: NodeId,
41 pub x: f64,
43 pub y: f64,
45}
46
47impl Graph {
48 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}