diff --git a/actor-tests/src/test/scala/org/apache/pekko/routing/ConsistentHashSpec.scala b/actor-tests/src/test/scala/org/apache/pekko/routing/ConsistentHashSpec.scala new file mode 100644 index 00000000000..cdd4b688677 --- /dev/null +++ b/actor-tests/src/test/scala/org/apache/pekko/routing/ConsistentHashSpec.scala @@ -0,0 +1,83 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.pekko.routing + +import org.scalatest.matchers.should.Matchers +import org.scalatest.wordspec.AnyWordSpec + +class ConsistentHashSpec extends AnyWordSpec with Matchers { + + // virtual nodes of these two nodes hash to the same ring positions with a factor of 10 + private val nodeA = "node-4230" + private val nodeB = "node-14323" + private val nodeC = "node-1" + private val factor = 10 + private val keys = (0 until 10000).map(i => s"key-$i") + + private def routing(ch: ConsistentHash[String]): Map[String, String] = + keys.iterator.map(key => key -> ch.nodeFor(key)).toMap + + "ConsistentHash" must { + + "route keys independently of the order the nodes were given" in { + routing(ConsistentHash(List(nodeA, nodeB, nodeC), factor)) should + ===(routing(ConsistentHash(List(nodeB, nodeA, nodeC), factor))) + routing(ConsistentHash(List(nodeA, nodeB, nodeC), factor)) should + ===(routing(ConsistentHash(List(nodeC, nodeB, nodeA), factor))) + } + + "route keys independently of the order the nodes were added" in { + val expected = routing(ConsistentHash(List(nodeA, nodeB, nodeC), factor)) + routing(ConsistentHash(List(nodeC), factor) :+ nodeA :+ nodeB) should ===(expected) + routing(ConsistentHash(List(nodeC), factor) :+ nodeB :+ nodeA) should ===(expected) + } + + "restore the colliding virtual nodes of the remaining node when a node is removed" in { + routing(ConsistentHash(List(nodeA, nodeB, nodeC), factor) :- nodeA) should + ===(routing(ConsistentHash(List(nodeB, nodeC), factor))) + routing(ConsistentHash(List(nodeB, nodeA, nodeC), factor) :- nodeA) should + ===(routing(ConsistentHash(List(nodeB, nodeC), factor))) + routing(ConsistentHash(List(nodeA, nodeB, nodeC), factor) :- nodeB) should + ===(routing(ConsistentHash(List(nodeA, nodeC), factor))) + routing(ConsistentHash(List(nodeB, nodeA, nodeC), factor) :- nodeB) should + ===(routing(ConsistentHash(List(nodeA, nodeC), factor))) + } + + "not remove virtual nodes of other nodes when removing a node that is not in the ring" in { + routing(ConsistentHash(List(nodeB, nodeC), factor) :- nodeA) should + ===(routing(ConsistentHash(List(nodeB, nodeC), factor))) + } + + "build the same ring with apply as by adding the nodes one by one" in { + val many = (0 until 200).map(n => s"node-host-$n") ++ List(nodeA, nodeB, nodeC, nodeA) + val incremental = many.foldLeft(ConsistentHash(Nil: Seq[String], factor))(_ :+ _) + routing(ConsistentHash(many, factor)) should ===(routing(incremental)) + routing(ConsistentHash(many.reverse, factor)) should ===(routing(incremental)) + } + + "be empty after all nodes are removed" in { + (ConsistentHash(List(nodeA, nodeB), factor) :- nodeA :- nodeB).isEmpty should ===(true) + } + + "not change the ring when a node is added again" in { + val ch = ConsistentHash(List(nodeA, nodeB, nodeC), factor) + routing(ch :+ nodeA) should ===(routing(ch)) + routing(ch :+ nodeB) should ===(routing(ch)) + } + } +} diff --git a/actor/src/main/scala/org/apache/pekko/routing/ConsistentHash.scala b/actor/src/main/scala/org/apache/pekko/routing/ConsistentHash.scala index 343970fa552..355a1e673d9 100644 --- a/actor/src/main/scala/org/apache/pekko/routing/ConsistentHash.scala +++ b/actor/src/main/scala/org/apache/pekko/routing/ConsistentHash.scala @@ -26,8 +26,17 @@ import scala.reflect.ClassTag * * Note that toString of the ring nodes are used for the node * hash, i.e. make sure it is different for different nodes. + * + * If virtual nodes of different nodes hash to the same ring position, the node + * with the lowest toString owns that position, so the ring does not depend on the + * order in which nodes were added. */ -class ConsistentHash[T: ClassTag] private (nodes: immutable.SortedMap[Int, T], val virtualNodesFactor: Int) { +class ConsistentHash[T: ClassTag] private ( + nodes: immutable.SortedMap[Int, T], + // other nodes whose virtual nodes hash to an owned ring position, kept so that + // they can take over the position if its owner is removed + collisions: immutable.Map[Int, List[T]], + val virtualNodesFactor: Int) { import ConsistentHash._ @@ -47,11 +56,8 @@ class ConsistentHash[T: ClassTag] private (nodes: immutable.SortedMap[Int, T], v * operation returns a new instance. */ def :+(node: T): ConsistentHash[T] = { - val nodeHash = hashFor(node.toString) - new ConsistentHash(nodes ++ - ((1 to virtualNodesFactor).map { r => - concatenateNodeHash(nodeHash, r) -> node - }), virtualNodesFactor) + val (newNodes, newCollisions) = claim(nodes, collisions, node, virtualNodesFactor) + new ConsistentHash(newNodes, newCollisions, virtualNodesFactor) } /** @@ -68,10 +74,13 @@ class ConsistentHash[T: ClassTag] private (nodes: immutable.SortedMap[Int, T], v */ def :-(node: T): ConsistentHash[T] = { val nodeHash = hashFor(node.toString) - new ConsistentHash(nodes -- - ((1 to virtualNodesFactor).map { r => - concatenateNodeHash(nodeHash, r) - }), virtualNodesFactor) + val (newNodes, newCollisions) = (1 to virtualNodesFactor).foldLeft((nodes, collisions)) { + case ((ns, cs), r) => + val hash = concatenateNodeHash(nodeHash, r) + val others = claimants(ns, cs, hash).filterNot(sameNode(_, node)) + setClaimants(ns, cs, hash, others) + } + new ConsistentHash(newNodes, newCollisions, virtualNodesFactor) } /** @@ -123,14 +132,52 @@ class ConsistentHash[T: ClassTag] private (nodes: immutable.SortedMap[Int, T], v object ConsistentHash { def apply[T: ClassTag](nodes: Iterable[T], virtualNodesFactor: Int): ConsistentHash[T] = { - new ConsistentHash( - immutable.SortedMap.empty[Int, T] ++ - (for { - node <- nodes - nodeHash = hashFor(node.toString) - vnode <- 1 to virtualNodesFactor - } yield concatenateNodeHash(nodeHash, vnode) -> node), - virtualNodesFactor) + if (virtualNodesFactor < 1) throw new IllegalArgumentException("virtualNodesFactor must be >= 1") + val nodeArray = nodes.toArray + val total = nodeArray.length * virtualNodesFactor + // each virtual node is encoded as (ring position << 32 | index), so that a primitive sort orders + // them by ring position, and by insertion order for the same position + val points = new Array[Long](total) + var n = 0 + while (n < nodeArray.length) { + val nodeHash = hashFor(nodeArray(n).toString) + var r = 1 + while (r <= virtualNodesFactor) { + val i = n * virtualNodesFactor + r - 1 + points(i) = (concatenateNodeHash(nodeHash, r).toLong << 32) | i + r += 1 + } + n += 1 + } + Arrays.sort(points) + + def hashAt(i: Int): Int = (points(i) >> 32).toInt + def nodeAt(i: Int): T = nodeArray((points(i) & 0xFFFFFFFFL).toInt / virtualNodesFactor) + + val ring = immutable.TreeMap.newBuilder[Int, T] + var collisions = immutable.Map.empty[Int, List[T]] + var i = 0 + while (i < total) { + val hash = hashAt(i) + var end = i + 1 + while (end < total && hashAt(end) == hash) end += 1 + if (end == i + 1) ring += hash -> nodeAt(i) + else { + // rare: several virtual nodes at the same ring position, resolve them like `:+` does + var claimants = List.empty[T] + var j = i + while (j < end) { + val node = nodeAt(j) + claimants = node :: claimants.filterNot(sameNode(_, node)) + j += 1 + } + val (owner, others) = ownerOf(claimants) + ring += hash -> owner + if (others.nonEmpty) collisions = collisions.updated(hash, others) + } + i = end + } + new ConsistentHash(ring.result(), collisions, virtualNodesFactor) } /** @@ -142,6 +189,54 @@ object ConsistentHash { apply(nodes.asScala, virtualNodesFactor) } + // nodes are identified by their toString, see the class documentation + private def sameNode[T](a: T, b: T): Boolean = a.toString == b.toString + + // all nodes with a virtual node at the given ring position, the owner first + private def claimants[T]( + nodes: immutable.SortedMap[Int, T], + collisions: immutable.Map[Int, List[T]], + hash: Int): List[T] = + nodes.get(hash) match { + case Some(owner) => owner :: collisions.getOrElse(hash, Nil) + case None => Nil + } + + private def setClaimants[T]( + nodes: immutable.SortedMap[Int, T], + collisions: immutable.Map[Int, List[T]], + hash: Int, + claimants: List[T]): (immutable.SortedMap[Int, T], immutable.Map[Int, List[T]]) = + if (claimants.isEmpty) (nodes - hash, collisions - hash) + else { + val (owner, others) = ownerOf(claimants) + (nodes.updated(hash, owner), if (others.isEmpty) collisions - hash else collisions.updated(hash, others)) + } + + // the lowest toString owns the position, independent of the order the nodes were added + private def ownerOf[T](claimants: List[T]): (T, List[T]) = + claimants match { + case single :: Nil => (single, Nil) + case _ => + val sorted = claimants.sortBy(_.toString) + (sorted.head, sorted.tail) + } + + // adds the virtual nodes of `node`, replacing any previous virtual nodes of the same node + private def claim[T]( + nodes: immutable.SortedMap[Int, T], + collisions: immutable.Map[Int, List[T]], + node: T, + virtualNodesFactor: Int): (immutable.SortedMap[Int, T], immutable.Map[Int, List[T]]) = { + val nodeHash = hashFor(node.toString) + (1 to virtualNodesFactor).foldLeft((nodes, collisions)) { + case ((ns, cs), r) => + val hash = concatenateNodeHash(nodeHash, r) + val others = claimants(ns, cs, hash).filterNot(sameNode(_, node)) + setClaimants(ns, cs, hash, node :: others) + } + } + private def concatenateNodeHash(nodeHash: Int, vnode: Int): Int = { import MurmurHash._ var h = startHash(nodeHash)