1use crate::util;
4use std::collections::hash_map::RandomState;
5use std::f64;
6use std::hash::{BuildHasher, Hash};
7
8pub struct Node<'a, T> {
14 id: &'a T,
15 hash: u64,
16 weight: f64,
17 relative_weight: f64,
18}
19
20impl<'a, T> Node<'a, T> {
21 pub fn new(id: &'a T, weight: f64) -> Self {
31 Node {
32 id,
33 hash: 0,
34 weight,
35 relative_weight: 0f64,
36 }
37 }
38}
39
40pub struct Ring<'a, T, H = RandomState> {
68 nodes: Vec<Node<'a, T>>,
69 hash_builder: H,
70}
71
72impl<'a, T> Ring<'a, T, RandomState> {
73 pub fn new(nodes: Vec<Node<'a, T>>) -> Self
83 where
84 T: Hash + Ord,
85 {
86 Self::with_hasher(Default::default(), nodes)
87 }
88}
89
90impl<'a, T, H> Ring<'a, T, H> {
91 fn rebalance(&mut self) {
92 let mut product = 1f64;
93 let len = self.nodes.len() as f64;
94 for i in 0..self.nodes.len() {
95 let index = i as f64;
96 let mut res;
97 if i == 0 {
98 res = (len * self.nodes[i].weight).powf(1f64 / len);
99 } else {
100 res = (len - index) * (self.nodes[i].weight - self.nodes[i - 1].weight) / product;
101 res += self.nodes[i - 1].relative_weight.powf(len - index);
102 res = res.powf(1f64 / (len - index));
103 }
104
105 product *= res;
106 self.nodes[i].relative_weight = res;
107 }
108 if let Some(max_relative_weight) = self.nodes.last().map(|node| node.relative_weight) {
109 for node in &mut self.nodes {
110 node.relative_weight /= max_relative_weight
111 }
112 }
113 }
114
115 pub fn with_hasher(hash_builder: H, mut nodes: Vec<Node<'a, T>>) -> Self
129 where
130 T: Hash + Ord,
131 H: BuildHasher + Default,
132 {
133 for node in &mut nodes {
134 node.hash = util::gen_hash(&hash_builder, node.id);
135 }
136 nodes.reverse();
137 nodes.sort_by_key(|node| node.id);
138 nodes.dedup_by_key(|node| node.id);
139 nodes.sort_by(|n, m| {
140 if (n.weight - m.weight).abs() < f64::EPSILON {
141 n.id.cmp(m.id)
142 } else {
143 n.weight
144 .partial_cmp(&m.weight)
145 .expect("Expected all non-NaN floats.")
146 }
147 });
148 let mut ret = Self {
149 nodes,
150 hash_builder,
151 };
152 ret.rebalance();
153 ret
154 }
155
156 pub fn insert_node(&mut self, mut new_node: Node<'a, T>)
172 where
173 T: Hash + Ord,
174 H: BuildHasher,
175 {
176 new_node.hash = util::gen_hash(&self.hash_builder, new_node.id);
177 if let Some(index) = self.nodes.iter().position(|node| node.id == new_node.id) {
178 self.nodes[index] = new_node;
179 } else {
180 self.nodes.push(new_node);
181 }
182 self.nodes.sort_by(|n, m| {
183 if (n.weight - m.weight).abs() < f64::EPSILON {
184 n.id.cmp(m.id)
185 } else {
186 n.weight
187 .partial_cmp(&m.weight)
188 .expect("Expected all non-NaN floats.")
189 }
190 });
191 self.rebalance();
192 }
193
194 pub fn remove_node(&mut self, id: &T)
206 where
207 T: Eq,
208 {
209 if let Some(index) = self.nodes.iter().position(|node| node.id == id) {
210 self.nodes.remove(index);
211 self.rebalance();
212 }
213 }
214
215 pub fn get_node<U>(&self, point: &U) -> &'a T
231 where
232 T: Ord,
233 U: Hash,
234 H: BuildHasher,
235 {
236 let point_hash = util::gen_hash(&self.hash_builder, point);
237 self.nodes
238 .iter()
239 .map(|node| {
240 (
241 util::combine_hash(&self.hash_builder, node.hash, point_hash) as f64
242 * node.relative_weight,
243 node.id,
244 )
245 })
246 .max_by(|n, m| {
247 if n == m {
248 n.1.cmp(m.1)
249 } else {
250 n.0.partial_cmp(&m.0).expect("Expected all non-NaN floats.")
251 }
252 })
253 .expect("Expected non-empty ring.")
254 .1
255 }
256
257 pub fn len(&self) -> usize {
269 self.nodes.len()
270 }
271
272 pub fn is_empty(&self) -> bool {
284 self.nodes.is_empty()
285 }
286
287 pub fn iter(&'a self) -> impl Iterator<Item = (&'a T, f64)> {
304 self.nodes.iter().map(|node| (&*node.id, node.weight))
305 }
306}
307
308impl<'a, T, H> IntoIterator for &'a Ring<'a, T, H> {
309 type IntoIter = Box<dyn Iterator<Item = (&'a T, f64)> + 'a>;
310 type Item = (&'a T, f64);
311
312 fn into_iter(self) -> Self::IntoIter {
313 Box::new(self.iter())
314 }
315}
316
317#[cfg(test)]
318mod tests {
319 use super::{Node, Ring};
320 use crate::test_util::BuildDefaultHasher;
321
322 macro_rules! assert_approx_eq {
323 ($a:expr, $b:expr) => {{
324 let (a, b) = (&$a, &$b);
325 assert!(
326 (*a - *b).abs() < 1.0e-6,
327 "{} is not approximately equal to {}",
328 *a,
329 *b
330 );
331 }};
332 }
333
334 #[test]
335 fn test_size_empty() {
336 let ring: Ring<'_, u32, _> = Ring::with_hasher(BuildDefaultHasher::default(), vec![]);
337 assert!(ring.is_empty());
338 assert_eq!(ring.len(), 0);
339 }
340
341 #[test]
342 fn test_correct_weights() {
343 let ring = Ring::with_hasher(
344 BuildDefaultHasher::default(),
345 vec![Node::new(&0, 0.4), Node::new(&1, 0.4), Node::new(&2, 0.2)],
346 );
347 assert_eq!(ring.nodes[0].id, &2);
348 assert_eq!(ring.nodes[1].id, &0);
349 assert_eq!(ring.nodes[2].id, &1);
350 assert_approx_eq!(ring.nodes[0].relative_weight, 0.774_596);
351 assert_approx_eq!(ring.nodes[1].relative_weight, 1.000_000);
352 assert_approx_eq!(ring.nodes[2].relative_weight, 1.000_000);
353 }
354
355 #[test]
356 fn test_new_replace() {
357 let ring = Ring::with_hasher(
358 BuildDefaultHasher::default(),
359 vec![Node::new(&0, 0.5), Node::new(&1, 0.1), Node::new(&1, 0.5)],
360 );
361
362 assert_eq!(ring.nodes[0].id, &0);
363 assert_eq!(ring.nodes[1].id, &1);
364 assert_approx_eq!(ring.nodes[0].relative_weight, 1.000_000);
365 assert_approx_eq!(ring.nodes[1].relative_weight, 1.000_000);
366 }
367
368 #[test]
369 fn test_insert_node() {
370 let mut ring = Ring::with_hasher(BuildDefaultHasher::default(), vec![Node::new(&0, 0.5)]);
371 ring.insert_node(Node::new(&1, 0.5));
372
373 assert_eq!(ring.nodes[0].id, &0);
374 assert_eq!(ring.nodes[1].id, &1);
375 assert_approx_eq!(ring.nodes[0].relative_weight, 1.000_000);
376 assert_approx_eq!(ring.nodes[1].relative_weight, 1.000_000);
377 }
378
379 #[test]
380 fn test_insert_node_replace() {
381 let mut ring = Ring::with_hasher(
382 BuildDefaultHasher::default(),
383 vec![Node::new(&0, 0.5), Node::new(&1, 0.1)],
384 );
385 ring.insert_node(Node::new(&1, 0.5));
386
387 assert_eq!(ring.nodes[0].id, &0);
388 assert_eq!(ring.nodes[1].id, &1);
389 assert_approx_eq!(ring.nodes[0].relative_weight, 1.000_000);
390 assert_approx_eq!(ring.nodes[1].relative_weight, 1.000_000);
391 }
392
393 #[test]
394 fn test_remove_node() {
395 let mut ring = Ring::with_hasher(
396 BuildDefaultHasher::default(),
397 vec![Node::new(&0, 0.5), Node::new(&1, 0.5), Node::new(&2, 0.1)],
398 );
399 ring.remove_node(&2);
400
401 assert_eq!(ring.nodes[0].id, &0);
402 assert_eq!(ring.nodes[1].id, &1);
403 assert_approx_eq!(ring.nodes[0].relative_weight, 1.000_000);
404 assert_approx_eq!(ring.nodes[1].relative_weight, 1.000_000);
405 }
406
407 #[test]
408 fn test_get_node() {
409 let ring = Ring::with_hasher(
410 BuildDefaultHasher::default(),
411 vec![Node::new(&0, 1.0), Node::new(&1, 1.0)],
412 );
413
414 assert_eq!(ring.get_node(&0), &0);
415 assert_eq!(ring.get_node(&1), &0);
416 assert_eq!(ring.get_node(&2), &0);
417 assert_eq!(ring.get_node(&3), &1);
418 assert_eq!(ring.get_node(&4), &1);
419 assert_eq!(ring.get_node(&5), &1);
420 }
421
422 #[test]
423 fn test_iter() {
424 let ring = Ring::with_hasher(
425 BuildDefaultHasher::default(),
426 vec![Node::new(&0, 0.4), Node::new(&1, 0.4), Node::new(&2, 0.2)],
427 );
428
429 let mut iterator = ring.iter();
430 let mut node;
431
432 node = iterator.next().unwrap();
433 assert_eq!(node.0, &2);
434 assert_approx_eq!(node.1, 0.2);
435
436 node = iterator.next().unwrap();
437 assert_eq!(node.0, &0);
438 assert_approx_eq!(node.1, 0.4);
439
440 node = iterator.next().unwrap();
441 assert_eq!(node.0, &1);
442 assert_approx_eq!(node.1, 0.4);
443
444 assert_eq!(iterator.next(), None);
445 }
446}