diff --git a/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/ClientState.java b/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/ClientState.java index 170ec315a90ce..32599287bef20 100644 --- a/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/ClientState.java +++ b/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/ClientState.java @@ -93,12 +93,23 @@ public ClientState(final Set previousActiveTasks, final Map taskLagTotals, final Map clientTags, final int capacity) { + this(previousActiveTasks, previousStandbyTasks, taskLagTotals, clientTags, capacity, null); + } + + // For testing only + public ClientState(final Set previousActiveTasks, + final Set previousStandbyTasks, + final Map taskLagTotals, + final Map clientTags, + final int capacity, + final UUID processId) { this.previousStandbyTasks.taskIds(unmodifiableSet(new TreeSet<>(previousStandbyTasks))); this.previousActiveTasks.taskIds(unmodifiableSet(new TreeSet<>(previousActiveTasks))); taskOffsetSums = emptyMap(); this.taskLagTotals = unmodifiableMap(taskLagTotals); this.capacity = capacity; this.clientTags = unmodifiableMap(clientTags); + this.processId = processId; } int capacity() { @@ -133,6 +144,10 @@ public void assignActiveTasks(final Collection tasks) { assignedActiveTasks.taskIds().addAll(tasks); } + public void assignStandbyTasks(final Collection tasks) { + assignedStandbyTasks.taskIds().addAll(tasks); + } + public void assignActiveToConsumer(final TaskId task, final String consumer) { if (!assignedActiveTasks.taskIds().contains(task)) { throw new IllegalStateException("added not assign active task " + task + " to this client state."); @@ -206,6 +221,10 @@ boolean hasStandbyTask(final TaskId taskId) { return assignedStandbyTasks.taskIds().contains(taskId); } + boolean hasActiveTask(final TaskId taskId) { + return assignedActiveTasks.taskIds().contains(taskId); + } + int standbyTaskCount() { return assignedStandbyTasks.taskIds().size(); } diff --git a/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/ClientTagAwareStandbyTaskAssignor.java b/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/ClientTagAwareStandbyTaskAssignor.java index de5036fe809d1..07cf73304ecd2 100644 --- a/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/ClientTagAwareStandbyTaskAssignor.java +++ b/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/ClientTagAwareStandbyTaskAssignor.java @@ -17,8 +17,11 @@ package org.apache.kafka.streams.processor.internals.assignment; import java.util.List; +import java.util.function.BiConsumer; import java.util.function.BiFunction; import java.util.function.Function; +import java.util.stream.Collectors; +import org.apache.kafka.streams.KeyValue; import org.apache.kafka.streams.processor.TaskId; import org.apache.kafka.streams.processor.internals.assignment.AssignorConfiguration.AssignmentConfigs; import org.slf4j.Logger; @@ -156,6 +159,48 @@ public boolean isAllowedTaskMovement(final ClientState source, final ClientState return true; } + /** + * Whether one task can be moved from source to destination. If the number of distinct tags including active + * and standby after the movement isn't decreased, then we can move the task. Otherwise, we can not move + * the task. + * @param source Source client + * @param destination Destination client + * @param sourceTask Task to move + * @param clientStateMap All client metadata + * @return If the task can be moved + */ + @Override + public boolean isAllowedTaskMovement(final ClientState source, + final ClientState destination, + final TaskId sourceTask, + final Map clientStateMap) { + + final BiConsumer>> addTags = (cs, tagSet) -> { + final Map tags = clientTagFunction.apply(cs.processId(), cs); + if (tags != null) { + tagSet.addAll(tags.entrySet().stream() + .map(entry -> KeyValue.pair(entry.getKey(), entry.getValue())) + .collect(Collectors.toList()) + ); + } + }; + + final Set> tagsWithSource = new HashSet<>(); + final Set> tagsWithDestination = new HashSet<>(); + for (final ClientState clientState : clientStateMap.values()) { + if (clientState.hasAssignedTask(sourceTask) + && !clientState.processId().equals(source.processId()) + && !clientState.processId().equals(destination.processId())) { + addTags.accept(clientState, tagsWithSource); + addTags.accept(clientState, tagsWithDestination); + } + } + addTags.accept(source, tagsWithSource); + addTags.accept(destination, tagsWithDestination); + + return tagsWithDestination.size() >= tagsWithSource.size(); + } + // Visible for testing void fillClientsTagStatistics(final Map clientStates, final Map> tagEntryToClients, diff --git a/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/RackAwareTaskAssignor.java b/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/RackAwareTaskAssignor.java index 9eb82593062e0..cd6e2a49b388b 100644 --- a/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/RackAwareTaskAssignor.java +++ b/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/RackAwareTaskAssignor.java @@ -17,6 +17,7 @@ package org.apache.kafka.streams.processor.internals.assignment; import java.util.ArrayList; +import java.util.Collection; import java.util.Collections; import java.util.HashMap; import java.util.HashSet; @@ -28,7 +29,13 @@ import java.util.Set; import java.util.SortedMap; import java.util.SortedSet; +import java.util.TreeSet; import java.util.UUID; +import java.util.function.BiConsumer; +import java.util.function.BiFunction; +import java.util.function.BiPredicate; +import java.util.stream.Collectors; +import java.util.stream.Stream; import org.apache.kafka.common.Cluster; import org.apache.kafka.common.Node; import org.apache.kafka.common.PartitionInfo; @@ -43,26 +50,39 @@ import org.slf4j.LoggerFactory; public class RackAwareTaskAssignor { + + @FunctionalInterface + public interface MoveStandbyTaskPredicate { + boolean canMove(final ClientState source, + final ClientState destination, + final TaskId taskId, + final Map clientStateMap); + } + private static final Logger log = LoggerFactory.getLogger(RackAwareTaskAssignor.class); private static final int SOURCE_ID = -1; private final Cluster fullMetadata; private final Map> partitionsForTask; + private final Map> changelogPartitionsForTask; private final AssignmentConfigs assignmentConfigs; private final Map> racksForPartition; private final Map racksForProcess; private final InternalTopicManager internalTopicManager; private final boolean validClientRack; + private Boolean canEnable = null; public RackAwareTaskAssignor(final Cluster fullMetadata, final Map> partitionsForTask, + final Map> changelogPartitionsForTask, final Map> tasksForTopicGroup, final Map>> racksForProcessConsumer, final InternalTopicManager internalTopicManager, final AssignmentConfigs assignmentConfigs) { this.fullMetadata = fullMetadata; this.partitionsForTask = partitionsForTask; + this.changelogPartitionsForTask = changelogPartitionsForTask; this.internalTopicManager = internalTopicManager; this.assignmentConfigs = assignmentConfigs; this.racksForPartition = new HashMap<>(); @@ -78,16 +98,30 @@ public synchronized boolean canEnableRackAwareAssignor() { /* TODO: enable this after we add the config if (StreamsConfig.RACK_AWARE_ASSSIGNMENT_STRATEGY_NONE.equals(assignmentConfigs.rackAwareAssignmentStrategy)) { - canEnableForActive = false; + canEnable = false; return false; } */ - return validClientRack && validateTopicPartitionRack(); - // TODO: add changelog topic, standby task validation + if (canEnable != null) { + return canEnable; + } + canEnable = validClientRack && validateTopicPartitionRack(false); + if (assignmentConfigs.numStandbyReplicas == 0 || !canEnable) { + return canEnable; + } + + canEnable = validateTopicPartitionRack(true); + return canEnable; } // Visible for testing. This method also checks if all TopicPartitions exist in cluster - public boolean populateTopicsToDescribe(final Set topicsToDescribe) { + public boolean populateTopicsToDescribe(final Set topicsToDescribe, final boolean changelog) { + if (changelog) { + // Changelog topics are not in metadata, we need to describe them + changelogPartitionsForTask.values().stream().flatMap(Collection::stream).forEach(tp -> topicsToDescribe.add(tp.topic())); + return true; + } + // Make sure rackId exist for all TopicPartitions needed for (final Set topicPartitions : partitionsForTask.values()) { for (final TopicPartition topicPartition : topicPartitions) { @@ -114,10 +148,10 @@ public boolean populateTopicsToDescribe(final Set topicsToDescribe) { return true; } - private boolean validateTopicPartitionRack() { + private boolean validateTopicPartitionRack(final boolean changelogTopics) { // Make sure rackId exist for all TopicPartitions needed final Set topicsToDescribe = new HashSet<>(); - if (!populateTopicsToDescribe(topicsToDescribe)) { + if (!populateTopicsToDescribe(topicsToDescribe, changelogTopics)) { return false; } @@ -201,13 +235,17 @@ public Map racksForProcess() { return Collections.unmodifiableMap(racksForProcess); } - private int getCost(final TaskId taskId, final UUID processId, final boolean inCurrentAssignment, final int trafficCost, final int nonOverlapCost) { + public Map> racksForPartition() { + return Collections.unmodifiableMap(racksForPartition); + } + + private int getCost(final TaskId taskId, final UUID processId, final boolean inCurrentAssignment, final int trafficCost, final int nonOverlapCost, final boolean isStandby) { final String clientRack = racksForProcess.get(processId); if (clientRack == null) { throw new IllegalStateException("Client " + processId + " doesn't have rack configured. Maybe forgot to call canEnableRackAwareAssignor first"); } - final Set topicPartitions = partitionsForTask.get(taskId); + final Set topicPartitions = isStandby ? changelogPartitionsForTask.get(taskId) : partitionsForTask.get(taskId); if (topicPartitions == null || topicPartitions.isEmpty()) { throw new IllegalStateException("Task " + taskId + " has no TopicPartitions"); } @@ -230,10 +268,18 @@ private int getCost(final TaskId taskId, final UUID processId, final boolean inC return cost; } - private static int getSinkID(final List clientList, final List taskIdList) { + private static int getSinkNodeID(final List clientList, final List taskIdList) { return clientList.size() + taskIdList.size(); } + private static int getClientNodeId(final List taskIdList, final int clientIndex) { + return clientIndex + taskIdList.size(); + } + + private static int getClientIndex(final List taskIdList, final int clientNodeId) { + return clientNodeId - taskIdList.size(); + } + /** * Compute the cost for the provided {@code activeTasks}. The passed in active tasks must be contained in {@code clientState}. */ @@ -241,14 +287,33 @@ long activeTasksCost(final SortedSet activeTasks, final SortedMap clientStates, final int trafficCost, final int nonOverlapCost) { - if (activeTasks.isEmpty()) { + return tasksCost(activeTasks, clientStates, trafficCost, nonOverlapCost, ClientState::hasActiveTask, false, false); + } + + /** + * Compute the cost for the provided {@code standbyTasks}. The passed in standby tasks must be contained in {@code clientState}. + */ + long standByTasksCost(final SortedSet standbyTasks, + final SortedMap clientStates, + final int trafficCost, + final int nonOverlapCost) { + return tasksCost(standbyTasks, clientStates, trafficCost, nonOverlapCost, ClientState::hasStandbyTask, true, true); + } + + private long tasksCost(final SortedSet tasks, + final SortedMap clientStates, + final int trafficCost, + final int nonOverlapCost, + final BiPredicate hasAssignedTask, + final boolean hasReplica, + final boolean isStandby) { + if (tasks.isEmpty()) { return 0; } - final List clientList = new ArrayList<>(clientStates.keySet()); - final List taskIdList = new ArrayList<>(activeTasks); - final Graph graph = constructActiveTaskGraph(clientList, taskIdList, - clientStates, new HashMap<>(), new HashMap<>(), trafficCost, nonOverlapCost); + final List taskIdList = new ArrayList<>(tasks); + final Graph graph = constructTaskGraph(clientList, taskIdList, + clientStates, new HashMap<>(), new HashMap<>(), hasAssignedTask, trafficCost, nonOverlapCost, hasReplica, isStandby); return graph.totalCost(); } @@ -279,30 +344,87 @@ public long optimizeActiveTasks(final SortedSet activeTasks, final List taskIdList = new ArrayList<>(activeTasks); final Map taskClientMap = new HashMap<>(); final Map originalAssignedTaskNumber = new HashMap<>(); - final Graph graph = constructActiveTaskGraph(clientList, taskIdList, - clientStates, taskClientMap, originalAssignedTaskNumber, trafficCost, nonOverlapCost); + final Graph graph = constructTaskGraph(clientList, taskIdList, + clientStates, taskClientMap, originalAssignedTaskNumber, ClientState::hasActiveTask, trafficCost, nonOverlapCost, false, false); graph.solveMinCostFlow(); final long cost = graph.totalCost(); - assignActiveTaskFromMinCostFlow(graph, clientList, taskIdList, - clientStates, originalAssignedTaskNumber, taskClientMap); + assignTaskFromMinCostFlow(graph, clientList, taskIdList, clientStates, originalAssignedTaskNumber, + taskClientMap, ClientState::assignActive, ClientState::unassignActive, ClientState::hasActiveTask); return cost; } - private Graph constructActiveTaskGraph(final List clientList, - final List taskIdList, - final Map clientStates, - final Map taskClientMap, - final Map originalAssignedTaskNumber, - final int trafficCost, - final int nonOverlapCost) { + public long optimizeStandbyTasks(final SortedMap clientStates, + final int trafficCost, + final int nonOverlapCost, + final MoveStandbyTaskPredicate moveStandbyTask) { + final BiFunction> getMovableTasks = (source, destination) -> source.standbyTasks().stream() + .filter(task -> !destination.hasAssignedTask(task)) + .filter(task -> moveStandbyTask.canMove(source, destination, task, clientStates)) + .sorted() + .collect(Collectors.toList()); + + final List clientList = new ArrayList<>(clientStates.keySet()); + final SortedSet standbyTasks = new TreeSet<>(); + for (int i = 0; i < clientList.size(); i++) { + final ClientState clientState1 = clientStates.get(clientList.get(i)); + standbyTasks.addAll(clientState1.standbyTasks()); + for (int j = i + 1; j < clientList.size(); j++) { + final ClientState clientState2 = clientStates.get(clientList.get(j)); + + final String rack1 = racksForProcess.get(clientState1.processId()); + final String rack2 = racksForProcess.get(clientState2.processId()); + // Cross rack traffic can not be reduced if racks are the same + if (rack1.equals(rack2)) { + continue; + } + + final List movable1 = getMovableTasks.apply(clientState1, clientState2); + final List movable2 = getMovableTasks.apply(clientState2, clientState1); + + // There's no needed to optimize if one is empty because the optimization + // can only swap tasks to keep the client's load balanced + if (movable1.isEmpty() || movable2.isEmpty()) { + continue; + } + + final List taskIdList = Stream.concat(movable1.stream(), movable2.stream()) + .sorted() + .collect(Collectors.toList()); + + final Map taskClientMap = new HashMap<>(); + final List clients = Stream.of(clientList.get(i), clientList.get(j)).sorted().collect( + Collectors.toList()); + final Map originalAssignedTaskNumber = new HashMap<>(); + + final Graph graph = constructTaskGraph(clients, taskIdList, clientStates, taskClientMap, originalAssignedTaskNumber, + ClientState::hasStandbyTask, trafficCost, nonOverlapCost, true, true); + graph.solveMinCostFlow(); + + assignTaskFromMinCostFlow(graph, clients, taskIdList, clientStates, originalAssignedTaskNumber, + taskClientMap, ClientState::assignStandby, ClientState::unassignStandby, ClientState::hasStandbyTask); + } + } + return standByTasksCost(standbyTasks, clientStates, trafficCost, nonOverlapCost); + } + + private Graph constructTaskGraph(final List clientList, + final List taskIdList, + final Map clientStates, + final Map taskClientMap, + final Map originalAssignedTaskNumber, + final BiPredicate hasAssignedTask, + final int trafficCost, + final int nonOverlapCost, + final boolean hasReplica, + final boolean isStandby) { final Graph graph = new Graph<>(); for (final TaskId taskId : taskIdList) { for (final Entry clientState : clientStates.entrySet()) { - if (clientState.getValue().hasAssignedTask(taskId)) { + if (hasAssignedTask.test(clientState.getValue(), taskId)) { originalAssignedTaskNumber.merge(clientState.getKey(), 1, Integer::sum); } } @@ -312,13 +434,13 @@ private Graph constructActiveTaskGraph(final List clientList, for (int taskNodeId = 0; taskNodeId < taskIdList.size(); taskNodeId++) { final TaskId taskId = taskIdList.get(taskNodeId); for (int j = 0; j < clientList.size(); j++) { - final int clientNodeId = taskIdList.size() + j; + final int clientNodeId = getClientNodeId(taskIdList, j); final UUID processId = clientList.get(j); - final int flow = clientStates.get(processId).hasAssignedTask(taskId) ? 1 : 0; - final int cost = getCost(taskId, processId, flow == 1, trafficCost, nonOverlapCost); + final int flow = hasAssignedTask.test(clientStates.get(processId), taskId) ? 1 : 0; + final int cost = getCost(taskId, processId, flow == 1, trafficCost, nonOverlapCost, isStandby); if (flow == 1) { - if (taskClientMap.containsKey(taskId)) { + if (!hasReplica && taskClientMap.containsKey(taskId)) { throw new IllegalArgumentException("Task " + taskId + " assigned to multiple clients " + processId + ", " + taskClientMap.get(taskId)); } @@ -335,11 +457,11 @@ private Graph constructActiveTaskGraph(final List clientList, graph.addEdge(SOURCE_ID, taskNodeId, 1, 0, 1); } - final int sinkId = getSinkID(clientList, taskIdList); + final int sinkId = getSinkNodeID(clientList, taskIdList); // It's possible that some clients have 0 task assign. These clients will have 0 tasks assigned // even though it may have higher traffic cost. This is to maintain the original assigned task count for (int i = 0; i < clientList.size(); i++) { - final int clientNodeId = taskIdList.size() + i; + final int clientNodeId = getClientNodeId(taskIdList, i); final int capacity = originalAssignedTaskNumber.getOrDefault(clientList.get(i), 0); // Flow equals to capacity for edges to sink graph.addEdge(clientNodeId, sinkId, capacity, 0, capacity); @@ -351,12 +473,15 @@ private Graph constructActiveTaskGraph(final List clientList, return graph; } - private void assignActiveTaskFromMinCostFlow(final Graph graph, + private void assignTaskFromMinCostFlow(final Graph graph, final List clientList, final List taskIdList, final Map clientStates, final Map originalAssignedTaskNumber, - final Map taskClientMap) { + final Map taskClientMap, + final BiConsumer assignTask, + final BiConsumer unassignTask, + final BiPredicate hasAssignedTask) { int tasksAssigned = 0; for (int taskNodeId = 0; taskNodeId < taskIdList.size(); taskNodeId++) { final TaskId taskId = taskIdList.get(taskNodeId); @@ -364,7 +489,7 @@ private void assignActiveTaskFromMinCostFlow(final Graph graph, for (final Graph.Edge edge : edges.values()) { if (edge.flow > 0) { tasksAssigned++; - final int clientIndex = edge.destination - taskIdList.size(); + final int clientIndex = getClientIndex(taskIdList, edge.destination); final UUID processId = clientList.get(clientIndex); final UUID originalProcessId = taskClientMap.get(taskId); @@ -373,8 +498,8 @@ private void assignActiveTaskFromMinCostFlow(final Graph graph, break; } - clientStates.get(originalProcessId).unassignActive(taskId); - clientStates.get(processId).assignActive(taskId); + unassignTask.accept(clientStates.get(originalProcessId), taskId); + assignTask.accept(clientStates.get(processId), taskId); } } } @@ -389,7 +514,7 @@ private void assignActiveTaskFromMinCostFlow(final Graph graph, final Map assignedTaskNumber = new HashMap<>(); for (final TaskId taskId : taskIdList) { for (final Entry clientState : clientStates.entrySet()) { - if (clientState.getValue().hasAssignedTask(taskId)) { + if (hasAssignedTask.test(clientState.getValue(), taskId)) { assignedTaskNumber.merge(clientState.getKey(), 1, Integer::sum); } } diff --git a/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/StandbyTaskAssignor.java b/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/StandbyTaskAssignor.java index 3b2ce99d2e2b4..a5c1ca2ddb5a3 100644 --- a/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/StandbyTaskAssignor.java +++ b/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/StandbyTaskAssignor.java @@ -16,8 +16,27 @@ */ package org.apache.kafka.streams.processor.internals.assignment; +import java.util.Map; +import java.util.UUID; +import org.apache.kafka.streams.processor.TaskId; + interface StandbyTaskAssignor extends TaskAssignor { default boolean isAllowedTaskMovement(final ClientState source, final ClientState destination) { return true; } + + /** + * If a specific task can be moved from source to destination + * @param source Source client + * @param destination Destination client + * @param sourceTask Task to move + * @param clientStateMap All client metadata + * @return True if task can be moved, false otherwise + */ + default boolean isAllowedTaskMovement(final ClientState source, + final ClientState destination, + final TaskId sourceTask, + final Map clientStateMap) { + return true; + } } \ No newline at end of file diff --git a/streams/src/test/java/org/apache/kafka/streams/processor/internals/assignment/AssignmentTestUtils.java b/streams/src/test/java/org/apache/kafka/streams/processor/internals/assignment/AssignmentTestUtils.java index 5da4da5ce96f4..5ab04f4d400df 100644 --- a/streams/src/test/java/org/apache/kafka/streams/processor/internals/assignment/AssignmentTestUtils.java +++ b/streams/src/test/java/org/apache/kafka/streams/processor/internals/assignment/AssignmentTestUtils.java @@ -91,6 +91,16 @@ public final class AssignmentTestUtils { public static final String TP_0_NAME = "topic0"; public static final String TP_1_NAME = "topic1"; + public static final String CHANGELOG_TP_0_NAME = "store-0-changelog"; + public static final String CHANGELOG_TP_1_NAME = "store-1-changelog"; + + public static final TopicPartition CHANGELOG_TP_0_0 = new TopicPartition(CHANGELOG_TP_0_NAME, 0); + public static final TopicPartition CHANGELOG_TP_0_1 = new TopicPartition(CHANGELOG_TP_0_NAME, 1); + public static final TopicPartition CHANGELOG_TP_0_2 = new TopicPartition(CHANGELOG_TP_0_NAME, 2); + public static final TopicPartition CHANGELOG_TP_1_0 = new TopicPartition(CHANGELOG_TP_1_NAME, 0); + public static final TopicPartition CHANGELOG_TP_1_1 = new TopicPartition(CHANGELOG_TP_1_NAME, 1); + public static final TopicPartition CHANGELOG_TP_1_2 = new TopicPartition(CHANGELOG_TP_1_NAME, 2); + public static final TopicPartition TP_0_0 = new TopicPartition(TP_0_NAME, 0); public static final TopicPartition TP_0_1 = new TopicPartition(TP_0_NAME, 1); public static final TopicPartition TP_0_2 = new TopicPartition(TP_0_NAME, 2); @@ -100,6 +110,7 @@ public final class AssignmentTestUtils { public static final PartitionInfo PI_0_0 = new PartitionInfo(TP_0_NAME, 0, NODE_0, REPLICA_0, REPLICA_0); public static final PartitionInfo PI_0_1 = new PartitionInfo(TP_0_NAME, 1, NODE_1, REPLICA_1, REPLICA_1); + public static final PartitionInfo PI_0_2 = new PartitionInfo(TP_0_NAME, 2, NODE_1, REPLICA_1, REPLICA_1); public static final PartitionInfo PI_1_0 = new PartitionInfo(TP_1_NAME, 0, NODE_2, REPLICA_2, REPLICA_2); public static final PartitionInfo PI_1_1 = new PartitionInfo(TP_1_NAME, 1, NODE_3, REPLICA_3, REPLICA_3); public static final PartitionInfo PI_1_2 = new PartitionInfo(TP_1_NAME, 2, NODE_0, REPLICA_0, REPLICA_0); diff --git a/streams/src/test/java/org/apache/kafka/streams/processor/internals/assignment/ClientTagAwareStandbyTaskAssignorTest.java b/streams/src/test/java/org/apache/kafka/streams/processor/internals/assignment/ClientTagAwareStandbyTaskAssignorTest.java index 07da8e4f83ea8..f5bf95ac535f6 100644 --- a/streams/src/test/java/org/apache/kafka/streams/processor/internals/assignment/ClientTagAwareStandbyTaskAssignorTest.java +++ b/streams/src/test/java/org/apache/kafka/streams/processor/internals/assignment/ClientTagAwareStandbyTaskAssignorTest.java @@ -230,6 +230,58 @@ public void shouldDeclineTaskMovementWhenClientTagsDoNotMatch() { assertFalse(standbyTaskAssignor.isAllowedTaskMovement(source, destination)); } + @Test + public void shouldPermitSingleTaskMoveWhenClientTagMatch() { + final ClientState source = createClientStateWithCapacity(UUID_1, 1, mkMap(mkEntry(ZONE_TAG, ZONE_1), mkEntry(CLUSTER_TAG, CLUSTER_1))); + final ClientState destination = createClientStateWithCapacity(UUID_2, 1, mkMap(mkEntry(ZONE_TAG, ZONE_1), mkEntry(CLUSTER_TAG, CLUSTER_1))); + final ClientState clientState = createClientStateWithCapacity(UUID_3, 1, mkMap(mkEntry(ZONE_TAG, ZONE_3), mkEntry(CLUSTER_TAG, CLUSTER_2))); + final Map clientStateMap = mkMap( + mkEntry(UUID_1, source), + mkEntry(UUID_2, destination), + mkEntry(UUID_3, clientState) + ); + final TaskId taskId = new TaskId(0, 0); + clientState.assignActive(taskId); + source.assignStandby(taskId); + + assertTrue(standbyTaskAssignor.isAllowedTaskMovement(source, destination, taskId, clientStateMap)); + } + + @Test + public void shouldPermitSingleTaskMoveWhenDifferentClientTagCountNotChange() { + final ClientState source = createClientStateWithCapacity(UUID_1, 1, mkMap(mkEntry(ZONE_TAG, ZONE_1), mkEntry(CLUSTER_TAG, CLUSTER_1))); + final ClientState destination = createClientStateWithCapacity(UUID_2, 1, mkMap(mkEntry(ZONE_TAG, ZONE_2), mkEntry(CLUSTER_TAG, CLUSTER_1))); + final ClientState clientState = createClientStateWithCapacity(UUID_3, 1, mkMap(mkEntry(ZONE_TAG, ZONE_3), mkEntry(CLUSTER_TAG, CLUSTER_2))); + final Map clientStateMap = mkMap( + mkEntry(UUID_1, source), + mkEntry(UUID_2, destination), + mkEntry(UUID_3, clientState) + ); + final TaskId taskId = new TaskId(0, 0); + clientState.assignActive(taskId); + source.assignStandby(taskId); + + assertTrue(standbyTaskAssignor.isAllowedTaskMovement(source, destination, taskId, clientStateMap)); + } + + @Test + public void shouldDeclineSingleTaskMoveWhenReduceClientTagCount() { + final ClientState source = createClientStateWithCapacity(UUID_1, 1, mkMap(mkEntry(ZONE_TAG, ZONE_1), mkEntry(CLUSTER_TAG, CLUSTER_1))); + final ClientState destination = createClientStateWithCapacity(UUID_2, 1, mkMap(mkEntry(ZONE_TAG, ZONE_3), mkEntry(CLUSTER_TAG, CLUSTER_1))); + final ClientState clientState = createClientStateWithCapacity(UUID_3, 1, mkMap(mkEntry(ZONE_TAG, ZONE_3), mkEntry(CLUSTER_TAG, CLUSTER_2))); + final Map clientStateMap = mkMap( + mkEntry(UUID_1, source), + mkEntry(UUID_2, destination), + mkEntry(UUID_3, clientState) + ); + final TaskId taskId = new TaskId(0, 0); + clientState.assignActive(taskId); + source.assignStandby(taskId); + + // Because destination has ZONE_3 which is the same as active's zone + assertFalse(standbyTaskAssignor.isAllowedTaskMovement(source, destination, taskId, clientStateMap)); + } + @Test public void shouldDistributeStandbyTasksWhenActiveTasksAreLocatedOnSameZone() { final Map clientStates = mkMap( diff --git a/streams/src/test/java/org/apache/kafka/streams/processor/internals/assignment/RackAwareTaskAssignorTest.java b/streams/src/test/java/org/apache/kafka/streams/processor/internals/assignment/RackAwareTaskAssignorTest.java index 572cc1e15aaa7..63698be5e2864 100644 --- a/streams/src/test/java/org/apache/kafka/streams/processor/internals/assignment/RackAwareTaskAssignorTest.java +++ b/streams/src/test/java/org/apache/kafka/streams/processor/internals/assignment/RackAwareTaskAssignorTest.java @@ -23,6 +23,14 @@ import static org.apache.kafka.common.utils.Utils.mkMap; import static org.apache.kafka.common.utils.Utils.mkSet; import static org.apache.kafka.common.utils.Utils.mkSortedSet; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.CHANGELOG_TP_0_0; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.CHANGELOG_TP_0_1; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.CHANGELOG_TP_0_2; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.CHANGELOG_TP_0_NAME; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.CHANGELOG_TP_1_0; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.CHANGELOG_TP_1_1; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.CHANGELOG_TP_1_2; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.CHANGELOG_TP_1_NAME; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.EMPTY_CLIENT_TAGS; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.NODE_0; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.NODE_1; @@ -31,6 +39,7 @@ import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.NO_RACK_NODE; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.PI_0_0; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.PI_0_1; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.PI_0_2; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.PI_1_0; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.PI_1_1; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.PI_1_2; @@ -39,15 +48,20 @@ import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.RACK_2; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.RACK_3; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.RACK_4; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.REPLICA_0; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.REPLICA_1; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.REPLICA_2; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.REPLICA_3; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.SUBTOPOLOGY_0; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.TASK_0_0; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.TASK_0_1; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.TASK_0_2; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.TASK_1_0; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.TASK_1_1; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.TASK_1_2; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.TP_0_0; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.TP_0_1; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.TP_0_2; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.TP_0_NAME; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.TP_1_0; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.TP_1_1; @@ -57,16 +71,27 @@ import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.UUID_3; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.UUID_4; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.UUID_5; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.UUID_6; +import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.UUID_7; import static org.apache.kafka.streams.processor.internals.assignment.AssignmentTestUtils.uuidForInt; import static org.hamcrest.MatcherAssert.assertThat; +import static org.hamcrest.Matchers.greaterThanOrEqualTo; +import static org.hamcrest.Matchers.hasItems; import static org.hamcrest.Matchers.lessThanOrEqualTo; +import static org.hamcrest.Matchers.not; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertTrue; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.mockito.ArgumentMatchers.anySet; +import static org.mockito.ArgumentMatchers.eq; import static org.mockito.Mockito.doReturn; import static org.mockito.Mockito.doThrow; +import static org.mockito.Mockito.never; import static org.mockito.Mockito.spy; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; import java.util.ArrayList; import java.util.Arrays; @@ -79,11 +104,14 @@ import java.util.Map; import java.util.Map.Entry; import java.util.Optional; +import java.util.Random; import java.util.Set; import java.util.SortedMap; import java.util.SortedSet; import java.util.TreeMap; +import java.util.TreeSet; import java.util.UUID; +import java.util.function.Function; import java.util.stream.Collectors; import org.apache.kafka.common.Cluster; import org.apache.kafka.common.Node; @@ -95,7 +123,9 @@ import org.apache.kafka.streams.StreamsConfig; import org.apache.kafka.streams.StreamsConfig.InternalConfig; import org.apache.kafka.streams.processor.TaskId; +import org.apache.kafka.streams.processor.internals.InternalTopicManager; import org.apache.kafka.streams.processor.internals.TopologyMetadata.Subtopology; +import org.apache.kafka.streams.processor.internals.assignment.AssignorConfiguration.AssignmentConfigs; import org.apache.kafka.test.MockClientSupplier; import org.apache.kafka.test.MockInternalTopicManager; import org.junit.Before; @@ -104,11 +134,15 @@ import org.junit.runner.RunWith; import org.junit.runners.Parameterized; import org.junit.runners.Parameterized.Parameter; +import org.mockito.Mockito; @RunWith(Parameterized.class) public class RackAwareTaskAssignorTest { private static final String USER_END_POINT = "localhost:8080"; private static final String APPLICATION_ID = "stream-partition-assignor-test"; + private static final String TOPIC_PREFIX = "topic"; + private static final String CHANGELOG_TOPIC_PREFIX = "changelog-topic"; + private static final String RACK_PREFIX = "rack"; private final MockTime time = new MockTime(); private final StreamsConfig streamsConfig = new StreamsConfig(configProps()); @@ -147,9 +181,14 @@ public void setUp() { } private Map configProps() { + return configProps(0); + } + + private Map configProps(final int standbyNum) { final Map configurationMap = new HashMap<>(); configurationMap.put(StreamsConfig.APPLICATION_ID_CONFIG, APPLICATION_ID); configurationMap.put(StreamsConfig.BOOTSTRAP_SERVERS_CONFIG, USER_END_POINT); + configurationMap.put(StreamsConfig.NUM_STANDBY_REPLICAS_CONFIG, standbyNum); final ReferenceContainer referenceContainer = new ReferenceContainer(); /* @@ -165,34 +204,40 @@ private Map configProps() { @Test public void shouldDisableActiveWhenMissingClusterInfo() { - final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( + final RackAwareTaskAssignor assignor = spy(new RackAwareTaskAssignor( getClusterForTopic0(), getTaskTopicPartitionMapForTask0(true), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForProcess0(), mockInternalTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() - ); + )); // False since partitionWithoutInfo10 is missing in cluster metadata assertFalse(assignor.canEnableRackAwareAssignor()); - assertFalse(assignor.populateTopicsToDescribe(new HashSet<>())); + verify(assignor, times(1)).populateTopicsToDescribe(anySet(), eq(false)); + verify(assignor, never()).populateTopicsToDescribe(anySet(), eq(true)); + assertFalse(assignor.populateTopicsToDescribe(new HashSet<>(), false)); } @Test public void shouldDisableActiveWhenRackMissingInNode() { - final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( + final RackAwareTaskAssignor assignor = spy(new RackAwareTaskAssignor( getClusterWithPartitionMissingRack(), getTaskTopicPartitionMapForTask0(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForProcess0(), mockInternalTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() - ); + )); - assertFalse(assignor.populateTopicsToDescribe(new HashSet<>())); // False since nodeMissingRack has one node which doesn't have rack assertFalse(assignor.canEnableRackAwareAssignor()); + verify(assignor, times(1)).populateTopicsToDescribe(anySet(), eq(false)); + verify(assignor, never()).populateTopicsToDescribe(anySet(), eq(true)); + assertFalse(assignor.populateTopicsToDescribe(new HashSet<>(), false)); } @Test @@ -200,6 +245,7 @@ public void shouldReturnInvalidClientRackWhenRackMissingInClientConsumer() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0(), getTaskTopicPartitionMapForTask0(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForProcess0(true), mockInternalTopicManager, @@ -214,6 +260,7 @@ public void shouldReturnFalseWhenRackMissingInProcess() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0(), getTaskTopicPartitionMapForTask0(), + mkMap(), getTopologyGroupTaskMap(), getProcessWithNoConsumerRacks(), mockInternalTopicManager, @@ -229,6 +276,7 @@ public void shouldPopulateRacksForProcess() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0(), getTaskTopicPartitionMapForTask0(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForProcess0(), mockInternalTopicManager, @@ -250,6 +298,7 @@ public void shouldReturnInvalidClientRackWhenRackDiffersInSameProcess() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0(), getTaskTopicPartitionMapForTask0(), + mkMap(), getTopologyGroupTaskMap(), processRacks, mockInternalTopicManager, @@ -264,6 +313,7 @@ public void shouldEnableRackAwareAssignorWithoutDescribingTopics() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0(), getTaskTopicPartitionMapForTask0(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForProcess0(), mockInternalTopicManager, @@ -274,6 +324,28 @@ public void shouldEnableRackAwareAssignorWithoutDescribingTopics() { assertTrue(assignor.canEnableRackAwareAssignor()); } + @Test + public void shouldEnableRackAwareAssignorWithCacheResult() { + final RackAwareTaskAssignor assignor = spy(new RackAwareTaskAssignor( + getClusterForTopic0(), + getTaskTopicPartitionMapForTask0(), + mkMap(), + getTopologyGroupTaskMap(), + getProcessRacksForProcess0(), + mockInternalTopicManager, + new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() + )); + + // partitionWithoutInfo00 has rackInfo in cluster metadata + assertTrue(assignor.canEnableRackAwareAssignor()); + verify(assignor, times(1)).populateTopicsToDescribe(anySet(), eq(false)); + + // Should use cache result + Mockito.reset(assignor); + assertTrue(assignor.canEnableRackAwareAssignor()); + verify(assignor, never()).populateTopicsToDescribe(anySet(), eq(false)); + } + @Test public void shouldEnableRackAwareAssignorWithDescribingTopics() { final MockInternalTopicManager spyTopicManager = spy(mockInternalTopicManager); @@ -289,6 +361,7 @@ public void shouldEnableRackAwareAssignorWithDescribingTopics() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterWithNoNode(), getTaskTopicPartitionMapForTask0(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForProcess0(), spyTopicManager, @@ -298,23 +371,101 @@ public void shouldEnableRackAwareAssignorWithDescribingTopics() { assertTrue(assignor.canEnableRackAwareAssignor()); } + @Test + public void shouldEnableRackAwareAssignorWithStandbyDescribingTopics() { + final MockInternalTopicManager spyTopicManager = spy(mockInternalTopicManager); + doReturn( + Collections.singletonMap( + TP_0_NAME, + Collections.singletonList( + new TopicPartitionInfo(0, NODE_0, Arrays.asList(REPLICA_1), Collections.emptyList()) + ) + ) + ).when(spyTopicManager).getTopicPartitionInfo(Collections.singleton(TP_0_NAME)); + + doReturn( + Collections.singletonMap( + CHANGELOG_TP_0_NAME, + Collections.singletonList( + new TopicPartitionInfo(0, NODE_0, Arrays.asList(REPLICA_1), Collections.emptyList()) + ) + ) + ).when(spyTopicManager).getTopicPartitionInfo(Collections.singleton(CHANGELOG_TP_0_NAME)); + + final StreamsConfig streamsConfig1 = new StreamsConfig(configProps(1)); + final RackAwareTaskAssignor assignor = spy(new RackAwareTaskAssignor( + getClusterWithNoNode(), + getTaskTopicPartitionMapForTask0(), + getTaskChangeLogTopicPartitionMapForTask0(), + getTopologyGroupTaskMap(), + getProcessRacksForProcess0(), + spyTopicManager, + new AssignorConfiguration(streamsConfig1.originals()).assignmentConfigs() + )); + + assertTrue(assignor.canEnableRackAwareAssignor()); + verify(assignor, times(1)).populateTopicsToDescribe(anySet(), eq(false)); + verify(assignor, times(1)).populateTopicsToDescribe(anySet(), eq(true)); + + final Map> racksForPartition = assignor.racksForPartition(); + final Map> expected = mkMap( + mkEntry(TP_0_0, mkSet(RACK_1, RACK_2)), + mkEntry(CHANGELOG_TP_0_0, mkSet(RACK_1, RACK_2)) + ); + assertEquals(expected, racksForPartition); + } + + @Test + public void shouldDisableRackAwareAssignorWithStandbyDescribingTopicsFailure() { + final MockInternalTopicManager spyTopicManager = spy(mockInternalTopicManager); + doReturn( + Collections.singletonMap( + TP_0_NAME, + Collections.singletonList( + new TopicPartitionInfo(0, NODE_0, Arrays.asList(REPLICA_1), Collections.emptyList()) + ) + ) + ).when(spyTopicManager).getTopicPartitionInfo(Collections.singleton(TP_0_NAME)); + + doThrow(new TimeoutException("Timeout describing topic")).when(spyTopicManager).getTopicPartitionInfo(Collections.singleton( + CHANGELOG_TP_0_NAME)); + + final StreamsConfig streamsConfig1 = new StreamsConfig(configProps(1)); + final RackAwareTaskAssignor assignor = spy(new RackAwareTaskAssignor( + getClusterWithNoNode(), + getTaskTopicPartitionMapForTask0(), + getTaskChangeLogTopicPartitionMapForTask0(), + getTopologyGroupTaskMap(), + getProcessRacksForProcess0(), + spyTopicManager, + new AssignorConfiguration(streamsConfig1.originals()).assignmentConfigs() + )); + + assertFalse(assignor.canEnableRackAwareAssignor()); + verify(assignor, times(1)).populateTopicsToDescribe(anySet(), eq(false)); + verify(assignor, times(1)).populateTopicsToDescribe(anySet(), eq(true)); + } + @Test public void shouldDisableRackAwareAssignorWithDescribingTopicsFailure() { final MockInternalTopicManager spyTopicManager = spy(mockInternalTopicManager); doThrow(new TimeoutException("Timeout describing topic")).when(spyTopicManager).getTopicPartitionInfo(Collections.singleton( TP_0_NAME)); - final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( + final RackAwareTaskAssignor assignor = spy(new RackAwareTaskAssignor( getClusterWithNoNode(), getTaskTopicPartitionMapForTask0(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForProcess0(), spyTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() - ); + )); assertFalse(assignor.canEnableRackAwareAssignor()); - assertTrue(assignor.populateTopicsToDescribe(new HashSet<>())); + verify(assignor, times(1)).populateTopicsToDescribe(anySet(), eq(false)); + verify(assignor, never()).populateTopicsToDescribe(anySet(), eq(true)); + assertTrue(assignor.populateTopicsToDescribe(new HashSet<>(), false)); } @Test @@ -322,18 +473,19 @@ public void shouldOptimizeEmptyActiveTasks() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0And1(), getTaskTopicPartitionMapForAllTasks(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForAllProcess(), mockInternalTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() ); - final ClientState clientState0 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); + final ClientState clientState1 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); - clientState0.assignActiveTasks(mkSet(TASK_0_1, TASK_1_1)); + clientState1.assignActiveTasks(mkSet(TASK_0_1, TASK_1_1)); final SortedMap clientStateMap = new TreeMap<>(mkMap( - mkEntry(UUID_1, clientState0) + mkEntry(UUID_1, clientState1) )); final SortedSet taskIds = mkSortedSet(); @@ -344,7 +496,7 @@ public void shouldOptimizeEmptyActiveTasks() { final long cost = assignor.optimizeActiveTasks(taskIds, clientStateMap, trafficCost, nonOverlapCost); assertEquals(0, cost); - assertEquals(mkSet(TASK_0_1, TASK_1_1), clientState0.activeTasks()); + assertEquals(mkSet(TASK_0_1, TASK_1_1), clientState1.activeTasks()); } @Test @@ -352,19 +504,20 @@ public void shouldOptimizeActiveTasks() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0And1(), getTaskTopicPartitionMapForAllTasks(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForAllProcess(), mockInternalTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() ); - final ClientState clientState0 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); final ClientState clientState1 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); final ClientState clientState2 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); + final ClientState clientState3 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); - clientState0.assignActiveTasks(mkSet(TASK_0_1, TASK_1_1)); - clientState1.assignActive(TASK_1_0); - clientState2.assignActive(TASK_0_0); + clientState1.assignActiveTasks(mkSet(TASK_0_1, TASK_1_1)); + clientState2.assignActive(TASK_1_0); + clientState3.assignActive(TASK_0_0); // task_0_0 has same rack as UUID_1 // task_0_1 has same rack as UUID_2 and UUID_3 @@ -372,9 +525,9 @@ public void shouldOptimizeActiveTasks() { // task_1_1 has same rack as UUID_2 // Optimal assignment is UUID_1: {0_0, 1_0}, UUID_2: {1_1}, UUID_3: {0_1} which result in no cross rack traffic final SortedMap clientStateMap = new TreeMap<>(mkMap( - mkEntry(UUID_1, clientState0), - mkEntry(UUID_2, clientState1), - mkEntry(UUID_3, clientState2) + mkEntry(UUID_1, clientState1), + mkEntry(UUID_2, clientState2), + mkEntry(UUID_3, clientState3) )); final SortedSet taskIds = mkSortedSet(TASK_0_0, TASK_0_1, TASK_1_0, TASK_1_1); @@ -387,31 +540,31 @@ public void shouldOptimizeActiveTasks() { final long cost = assignor.optimizeActiveTasks(taskIds, clientStateMap, trafficCost, nonOverlapCost); assertEquals(expected, cost); - assertEquals(mkSet(TASK_0_0, TASK_1_0), clientState0.activeTasks()); - assertEquals(mkSet(TASK_1_1), clientState1.activeTasks()); - assertEquals(mkSet(TASK_0_1), clientState2.activeTasks()); + assertEquals(mkSet(TASK_0_0, TASK_1_0), clientState1.activeTasks()); + assertEquals(mkSet(TASK_1_1), clientState2.activeTasks()); + assertEquals(mkSet(TASK_0_1), clientState3.activeTasks()); } @Test - public void shouldOptimizeRandom() { + public void shouldOptimizeRandomActive() { final int nodeSize = 30; final int tpSize = 40; final int clientSize = 30; - final SortedMap> taskTopicPartitionMap = getTaskTopicPartitionMap(tpSize); + final SortedMap> taskTopicPartitionMap = getTaskTopicPartitionMap(tpSize, false); final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getRandomCluster(nodeSize, tpSize), taskTopicPartitionMap, + mkMap(), getTopologyGroupTaskMap(), getRandomProcessRacks(clientSize, nodeSize), mockInternalTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() ); - final SortedMap clientStateMap = getRandomClientState(clientSize, tpSize); + final SortedMap clientStateMap = getRandomClientState(clientSize, tpSize, 1); final SortedSet taskIds = (SortedSet) taskTopicPartitionMap.keySet(); - final Map clientTaskCount = clientStateMap.entrySet().stream() - .collect(Collectors.toMap(Map.Entry::getKey, e -> e.getValue().activeTasks().size())); + final Map clientTaskCount = clientTaskCount(clientStateMap, ClientState::activeTaskCount); assertTrue(assignor.canEnableRackAwareAssignor()); final long originalCost = assignor.activeTasksCost(taskIds, clientStateMap, trafficCost, nonOverlapCost); @@ -428,17 +581,18 @@ public void shouldMaintainOriginalAssignment() { final int nodeSize = 20; final int tpSize = 40; final int clientSize = 30; - final SortedMap> taskTopicPartitionMap = getTaskTopicPartitionMap(tpSize); + final SortedMap> taskTopicPartitionMap = getTaskTopicPartitionMap(tpSize, false); final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getRandomCluster(nodeSize, tpSize), taskTopicPartitionMap, + mkMap(), getTopologyGroupTaskMap(), getRandomProcessRacks(clientSize, nodeSize), mockInternalTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() ); - final SortedMap clientStateMap = getRandomClientState(clientSize, tpSize); + final SortedMap clientStateMap = getRandomClientState(clientSize, tpSize, 1); final SortedSet taskIds = (SortedSet) taskTopicPartitionMap.keySet(); final Map taskClientMap = new HashMap<>(); @@ -467,27 +621,28 @@ public void shouldOptimizeActiveTasksWithMoreClients() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0And1(), getTaskTopicPartitionMapForAllTasks(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForAllProcess(), mockInternalTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() ); - final ClientState clientState0 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); final ClientState clientState1 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); final ClientState clientState2 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); + final ClientState clientState3 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); - clientState1.assignActive(TASK_1_0); - clientState2.assignActive(TASK_0_0); + clientState2.assignActive(TASK_1_0); + clientState3.assignActive(TASK_0_0); // task_0_0 has same rack as UUID_1 and UUID_2 // task_1_0 has same rack as UUID_1 and UUID_3 // Optimal assignment is UUID_1: {}, UUID_2: {0_0}, UUID_3: {1_0} which result in no cross rack traffic // and keeps UUID_1 empty since it was originally empty final SortedMap clientStateMap = new TreeMap<>(mkMap( - mkEntry(UUID_1, clientState0), - mkEntry(UUID_2, clientState1), - mkEntry(UUID_3, clientState2) + mkEntry(UUID_1, clientState1), + mkEntry(UUID_2, clientState2), + mkEntry(UUID_3, clientState3) )); final SortedSet taskIds = mkSortedSet(TASK_0_0, TASK_1_0); @@ -501,9 +656,9 @@ public void shouldOptimizeActiveTasksWithMoreClients() { assertEquals(expected, cost); // UUID_1 remains empty - assertEquals(mkSet(), clientState0.activeTasks()); - assertEquals(mkSet(TASK_0_0), clientState1.activeTasks()); - assertEquals(mkSet(TASK_1_0), clientState2.activeTasks()); + assertEquals(mkSet(), clientState1.activeTasks()); + assertEquals(mkSet(TASK_0_0), clientState2.activeTasks()); + assertEquals(mkSet(TASK_1_0), clientState3.activeTasks()); } @Test @@ -511,27 +666,28 @@ public void shouldOptimizeActiveTasksWithMoreClientsWithMoreThanOneTask() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0And1(), getTaskTopicPartitionMapForAllTasks(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForAllProcess(), mockInternalTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() ); - final ClientState clientState0 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); final ClientState clientState1 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); final ClientState clientState2 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); + final ClientState clientState3 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); - clientState1.assignActiveTasks(mkSet(TASK_0_1, TASK_1_0)); - clientState2.assignActive(TASK_0_0); + clientState2.assignActiveTasks(mkSet(TASK_0_1, TASK_1_0)); + clientState3.assignActive(TASK_0_0); // task_0_0 has same rack as UUID_1 and UUID_2 // task_0_1 has same rack as UUID_2 and UUID_3 // task_1_0 has same rack as UUID_1 and UUID_3 // Optimal assignment is UUID_1: {}, UUID_2: {0_0, 0_1}, UUID_3: {1_0} which result in no cross rack traffic final SortedMap clientStateMap = new TreeMap<>(mkMap( - mkEntry(UUID_1, clientState0), - mkEntry(UUID_2, clientState1), - mkEntry(UUID_3, clientState2) + mkEntry(UUID_1, clientState1), + mkEntry(UUID_2, clientState2), + mkEntry(UUID_3, clientState3) )); final SortedSet taskIds = mkSortedSet(TASK_0_0, TASK_0_1, TASK_1_0); @@ -545,9 +701,9 @@ public void shouldOptimizeActiveTasksWithMoreClientsWithMoreThanOneTask() { assertEquals(expected, cost); // Because original assignment is not balanced (3 tasks but client 0 has no task), we maintain it - assertEquals(mkSet(), clientState0.activeTasks()); - assertEquals(mkSet(TASK_0_0, TASK_0_1), clientState1.activeTasks()); - assertEquals(mkSet(TASK_1_0), clientState2.activeTasks()); + assertEquals(mkSet(), clientState1.activeTasks()); + assertEquals(mkSet(TASK_0_0, TASK_0_1), clientState2.activeTasks()); + assertEquals(mkSet(TASK_1_0), clientState3.activeTasks()); } @Test @@ -555,25 +711,26 @@ public void shouldBalanceAssignmentWithMoreCost() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0And1(), getTaskTopicPartitionMapForAllTasks(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForAllProcess(), mockInternalTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() ); - final ClientState clientState0 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); final ClientState clientState1 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); + final ClientState clientState2 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); - clientState0.assignActiveTasks(mkSet(TASK_0_0, TASK_1_1)); - clientState1.assignActive(TASK_0_1); + clientState1.assignActiveTasks(mkSet(TASK_0_0, TASK_1_1)); + clientState2.assignActive(TASK_0_1); // task_0_0 has same rack as UUID_2 // task_0_1 has same rack as UUID_2 // task_1_1 has same rack as UUID_2 // UUID_5 is not in same rack as any task final SortedMap clientStateMap = new TreeMap<>(mkMap( - mkEntry(UUID_2, clientState0), - mkEntry(UUID_5, clientState1) + mkEntry(UUID_2, clientState1), + mkEntry(UUID_5, clientState2) )); final SortedSet taskIds = mkSortedSet(TASK_0_0, TASK_0_1, TASK_1_1); @@ -587,8 +744,8 @@ public void shouldBalanceAssignmentWithMoreCost() { // Even though assigning all tasks to UUID_2 will result in min cost, but it's not balanced // assignment. That's why TASK_0_1 is still assigned to UUID_5 - assertEquals(mkSet(TASK_0_0, TASK_1_1), clientState0.activeTasks()); - assertEquals(mkSet(TASK_0_1), clientState1.activeTasks()); + assertEquals(mkSet(TASK_0_0, TASK_1_1), clientState1.activeTasks()); + assertEquals(mkSet(TASK_0_1), clientState2.activeTasks()); } @Test @@ -596,21 +753,22 @@ public void shouldThrowIfMissingCallcanEnableRackAwareAssignor() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0And1(), getTaskTopicPartitionMapForAllTasks(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForAllProcess(), mockInternalTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() ); - final ClientState clientState0 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); final ClientState clientState1 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); + final ClientState clientState2 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); - clientState0.assignActiveTasks(mkSet(TASK_0_0, TASK_1_1)); - clientState1.assignActive(TASK_0_1); + clientState1.assignActiveTasks(mkSet(TASK_0_0, TASK_1_1)); + clientState2.assignActive(TASK_0_1); final SortedMap clientStateMap = new TreeMap<>(mkMap( - mkEntry(UUID_2, clientState0), - mkEntry(UUID_5, clientState1) + mkEntry(UUID_2, clientState1), + mkEntry(UUID_5, clientState2) )); final SortedSet taskIds = mkSortedSet(TASK_0_0, TASK_0_1, TASK_1_1); final Exception exception = assertThrows(IllegalStateException.class, @@ -624,21 +782,22 @@ public void shouldThrowIfTaskInMultipleClients() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0And1(), getTaskTopicPartitionMapForAllTasks(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForAllProcess(), mockInternalTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() ); - final ClientState clientState0 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); final ClientState clientState1 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); + final ClientState clientState2 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); - clientState0.assignActiveTasks(mkSet(TASK_0_0, TASK_1_1)); - clientState1.assignActiveTasks(mkSet(TASK_0_1, TASK_1_1)); + clientState1.assignActiveTasks(mkSet(TASK_0_0, TASK_1_1)); + clientState2.assignActiveTasks(mkSet(TASK_0_1, TASK_1_1)); final SortedMap clientStateMap = new TreeMap<>(mkMap( - mkEntry(UUID_2, clientState0), - mkEntry(UUID_5, clientState1) + mkEntry(UUID_2, clientState1), + mkEntry(UUID_5, clientState2) )); final SortedSet taskIds = mkSortedSet(TASK_0_0, TASK_0_1, TASK_1_1); assertTrue(assignor.canEnableRackAwareAssignor()); @@ -654,21 +813,22 @@ public void shouldThrowIfTaskMissingInClients() { final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( getClusterForTopic0And1(), getTaskTopicPartitionMapForAllTasks(), + mkMap(), getTopologyGroupTaskMap(), getProcessRacksForAllProcess(), mockInternalTopicManager, new AssignorConfiguration(streamsConfig.originals()).assignmentConfigs() ); - final ClientState clientState0 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); final ClientState clientState1 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); + final ClientState clientState2 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); - clientState0.assignActiveTasks(mkSet(TASK_0_0, TASK_1_1)); - clientState1.assignActive(TASK_0_1); + clientState1.assignActiveTasks(mkSet(TASK_0_0, TASK_1_1)); + clientState2.assignActive(TASK_0_1); final SortedMap clientStateMap = new TreeMap<>(mkMap( - mkEntry(UUID_2, clientState0), - mkEntry(UUID_5, clientState1) + mkEntry(UUID_2, clientState1), + mkEntry(UUID_5, clientState2) )); final SortedSet taskIds = mkSortedSet(TASK_0_0, TASK_0_1, TASK_1_0, TASK_1_1); assertTrue(assignor.canEnableRackAwareAssignor()); @@ -678,18 +838,230 @@ public void shouldThrowIfTaskMissingInClients() { "Task 1_0 not assigned to any client", exception.getMessage()); } - private Cluster getRandomCluster(final int nodeSize, final int tpSize) { + @Test + public void shouldNotCrashForEmptyStandby() { + final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( + getClusterForTopic0And1(), + getTaskTopicPartitionMapForAllTasks(), + mkMap(), + getTopologyGroupTaskMap(), + getProcessRacksForAllProcess(), + mockInternalTopicManagerForChangelog(), + new AssignorConfiguration(new StreamsConfig(configProps(1)).originals()).assignmentConfigs() + ); + + final ClientState clientState1 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_1); + final ClientState clientState2 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_2); + final ClientState clientState3 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_3); + + clientState1.assignActiveTasks(mkSet(TASK_0_1, TASK_1_1)); + clientState2.assignActive(TASK_1_0); + clientState3.assignActive(TASK_0_0); + + final SortedMap clientStateMap = new TreeMap<>(mkMap( + mkEntry(UUID_1, clientState1), + mkEntry(UUID_2, clientState2), + mkEntry(UUID_3, clientState3) + )); + + final long originalCost = assignor.standByTasksCost(new TreeSet<>(), clientStateMap, 10, 1); + assertEquals(0, originalCost); + + final long cost = assignor.optimizeStandbyTasks(clientStateMap, 10, 1, + (source, destination, task, clientStates) -> true); + assertEquals(0, cost); + } + + @Test + public void shouldOptimizeStandbyTasksWhenTasksAllMovable() { + final int replicaCount = 2; + final AssignmentConfigs assignorConfiguration = new AssignorConfiguration(new StreamsConfig(configProps(replicaCount)).originals()).assignmentConfigs(); + final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( + getClusterForTopic0And1(), + getTaskTopicPartitionMapForAllTasks(), + getTaskChangelogMapForAllTasks(), + getTopologyGroupTaskMap(), + getProcessRacksForAllProcess(), + mockInternalTopicManagerForChangelog(), + assignorConfiguration + ); + + final ClientState clientState1 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_1); + final ClientState clientState2 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_2); + final ClientState clientState3 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_3); + final ClientState clientState4 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_4); + final ClientState clientState5 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_6); + final ClientState clientState6 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_7); + + final SortedMap clientStateMap = new TreeMap<>(mkMap( + mkEntry(UUID_1, clientState1), + mkEntry(UUID_2, clientState2), + mkEntry(UUID_3, clientState3), + mkEntry(UUID_4, clientState4), + mkEntry(UUID_6, clientState5), + mkEntry(UUID_7, clientState6) + )); + + clientState1.assignActive(TASK_0_0); + clientState2.assignActive(TASK_0_1); + clientState3.assignActive(TASK_1_0); + clientState4.assignActive(TASK_1_1); + clientState5.assignActive(TASK_0_2); + clientState6.assignActive(TASK_1_2); + + clientState1.assignStandbyTasks(mkSet(TASK_0_1, TASK_1_1)); // Cost 10 + clientState2.assignStandbyTasks(mkSet(TASK_0_0, TASK_1_0)); // Cost 10 + clientState3.assignStandbyTasks(mkSet(TASK_0_0, TASK_0_2)); // Cost 20 + clientState4.assignStandbyTasks(mkSet(TASK_0_1, TASK_1_2)); // Cost 10 + clientState5.assignStandbyTasks(mkSet(TASK_1_0, TASK_1_2)); // Cost 10 + clientState6.assignStandbyTasks(mkSet(TASK_0_2, TASK_1_1)); // Cost 10 + + final SortedSet taskIds = new TreeSet<>(mkSet(TASK_0_0, TASK_0_1, TASK_0_2, TASK_1_0, TASK_1_1, TASK_1_2)); + final Map standbyTaskCount = clientTaskCount(clientStateMap, ClientState::standbyTaskCount); + + assertTrue(assignor.canEnableRackAwareAssignor()); + verifyStandbySatisfyRackReplica(taskIds, assignor.racksForProcess(), clientStateMap, replicaCount, false, null); + + final long originalCost = assignor.standByTasksCost(taskIds, clientStateMap, 10, 1); + assertEquals(60, originalCost); + + // Task can be moved anywhere so cost can be reduced to 30 compared to in shouldOptimizeStandbyTasksWithMovingConstraint it + // can only be reduced to 50 since there are moving constraints + final long cost = assignor.optimizeStandbyTasks(clientStateMap, 10, 1, + (source, destination, task, clients) -> true); + assertEquals(30, cost); + // Don't validate tasks in different racks after moving + verifyStandbySatisfyRackReplica(taskIds, assignor.racksForProcess(), clientStateMap, replicaCount, true, standbyTaskCount); + } + + @Test + public void shouldOptimizeStandbyTasksWithMovingConstraint() { + final int replicaCount = 2; + final AssignmentConfigs assignorConfiguration = new AssignorConfiguration(new StreamsConfig(configProps(replicaCount)).originals()).assignmentConfigs(); + final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( + getClusterForTopic0And1(), + getTaskTopicPartitionMapForAllTasks(), + getTaskChangelogMapForAllTasks(), + getTopologyGroupTaskMap(), + getProcessRacksForAllProcess(), + mockInternalTopicManagerForChangelog(), + assignorConfiguration + ); + + final ClientState clientState1 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_1); + final ClientState clientState2 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_2); + final ClientState clientState3 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_3); + final ClientState clientState4 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_4); + final ClientState clientState5 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_6); + final ClientState clientState6 = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1, UUID_7); + + final SortedMap clientStateMap = new TreeMap<>(mkMap( + mkEntry(UUID_1, clientState1), + mkEntry(UUID_2, clientState2), + mkEntry(UUID_3, clientState3), + mkEntry(UUID_4, clientState4), + mkEntry(UUID_6, clientState5), + mkEntry(UUID_7, clientState6) + )); + + clientState1.assignActive(TASK_0_0); + clientState2.assignActive(TASK_0_1); + clientState3.assignActive(TASK_1_0); + clientState4.assignActive(TASK_1_1); + clientState5.assignActive(TASK_0_2); + clientState6.assignActive(TASK_1_2); + + clientState1.assignStandbyTasks(mkSet(TASK_0_1, TASK_1_1)); // Cost 10 + clientState2.assignStandbyTasks(mkSet(TASK_0_0, TASK_1_0)); // Cost 10 + clientState3.assignStandbyTasks(mkSet(TASK_0_0, TASK_0_2)); // Cost 20 + clientState4.assignStandbyTasks(mkSet(TASK_0_1, TASK_1_2)); // Cost 10 + clientState5.assignStandbyTasks(mkSet(TASK_1_0, TASK_1_2)); // Cost 10 + clientState6.assignStandbyTasks(mkSet(TASK_0_2, TASK_1_1)); // Cost 10 + + final SortedSet taskIds = new TreeSet<>(mkSet(TASK_0_0, TASK_0_1, TASK_0_2, TASK_1_0, TASK_1_1, TASK_1_2)); + final Map standbyTaskCount = clientTaskCount(clientStateMap, ClientState::standbyTaskCount); + + assertTrue(assignor.canEnableRackAwareAssignor()); + verifyStandbySatisfyRackReplica(taskIds, assignor.racksForProcess(), clientStateMap, replicaCount, false, null); + + final long originalCost = assignor.standByTasksCost(taskIds, clientStateMap, 10, 1); + assertEquals(60, originalCost); + + final StandbyTaskAssignor standbyTaskAssignor = StandbyTaskAssignorFactory.create(assignorConfiguration, assignor); + assertInstanceOf(ClientTagAwareStandbyTaskAssignor.class, standbyTaskAssignor); + final long cost = assignor.optimizeStandbyTasks(clientStateMap, 10, 1, + standbyTaskAssignor::isAllowedTaskMovement); + assertEquals(50, cost); + // Validate tasks in different racks after moving + verifyStandbySatisfyRackReplica(taskIds, assignor.racksForProcess(), clientStateMap, replicaCount, false, standbyTaskCount); + } + + @Test + public void shouldOptimizeRandomStandby() { + final int nodeSize = 50; + final int tpSize = 60; + final int clientSize = 50; + final int replicaCount = 3; + final int maxCapacity = 3; + final SortedMap> taskTopicPartitionMap = getTaskTopicPartitionMap( + tpSize, false); + final AssignmentConfigs assignorConfiguration = new AssignorConfiguration( + new StreamsConfig(configProps(replicaCount)).originals()).assignmentConfigs(); + + final RackAwareTaskAssignor assignor = new RackAwareTaskAssignor( + getRandomCluster(nodeSize, tpSize), + taskTopicPartitionMap, + getTaskTopicPartitionMap(tpSize, true), + getTopologyGroupTaskMap(), + getRandomProcessRacks(clientSize, nodeSize), + mockInternalTopicManagerForRandomChangelog(nodeSize, tpSize), + assignorConfiguration + ); + + final SortedMap clientStateMap = getRandomClientState(clientSize, + tpSize, maxCapacity); + final SortedSet taskIds = (SortedSet) taskTopicPartitionMap.keySet(); + + final StandbyTaskAssignor standbyTaskAssignor = StandbyTaskAssignorFactory.create( + assignorConfiguration, assignor); + assertInstanceOf(ClientTagAwareStandbyTaskAssignor.class, standbyTaskAssignor); + // Get a standby assignment + standbyTaskAssignor.assign(clientStateMap, taskIds, taskIds, assignorConfiguration); + final Map standbyTaskCount = clientTaskCount(clientStateMap, + ClientState::standbyTaskCount); + + assertTrue(assignor.canEnableRackAwareAssignor()); + verifyStandbySatisfyRackReplica(taskIds, assignor.racksForProcess(), clientStateMap, + replicaCount, false, null); + + final long originalCost = assignor.standByTasksCost(taskIds, clientStateMap, 10, 1); + assertThat(originalCost, greaterThanOrEqualTo(0L)); + + final long cost = assignor.optimizeStandbyTasks(clientStateMap, 10, 1, + standbyTaskAssignor::isAllowedTaskMovement); + assertThat(cost, lessThanOrEqualTo(originalCost)); + // Validate tasks in different racks after moving + verifyStandbySatisfyRackReplica(taskIds, assignor.racksForProcess(), clientStateMap, + replicaCount, false, standbyTaskCount); + } + + private List getRandomNodes(final int nodeSize) { final List nodeList = new ArrayList<>(nodeSize); for (int i = 0; i < nodeSize; i++) { - nodeList.add(new Node(i, "node" + i, 1, "rack" + i)); + nodeList.add(new Node(i, "node" + i, 1, RACK_PREFIX + i)); } Collections.shuffle(nodeList); + return nodeList; + } + + private Cluster getRandomCluster(final int nodeSize, final int tpSize) { + final List nodeList = getRandomNodes(nodeSize); final Set partitionInfoSet = new HashSet<>(); for (int i = 0; i < tpSize; i++) { final Node firstNode = nodeList.get(i % nodeSize); final Node secondNode = nodeList.get((i + 1) % nodeSize); final Node[] replica = new Node[] {firstNode, secondNode}; - partitionInfoSet.add(new PartitionInfo("topic" + i, 0, firstNode, replica, replica)); + partitionInfoSet.add(new PartitionInfo(TOPIC_PREFIX + i, 0, firstNode, replica, replica)); } return new Cluster( @@ -704,7 +1076,7 @@ private Cluster getRandomCluster(final int nodeSize, final int tpSize) { private Map>> getRandomProcessRacks(final int clientSize, final int nodeSize) { final List racks = new ArrayList<>(nodeSize); for (int i = 0; i < nodeSize; i++) { - racks.add("rack" + i); + racks.add(RACK_PREFIX + i); } Collections.shuffle(racks); final Map>> processRacks = new HashMap<>(); @@ -715,23 +1087,49 @@ private Map>> getRandomProcessRacks(final int return processRacks; } - private SortedMap> getTaskTopicPartitionMap(final int tpSize) { + private SortedMap> getTaskTopicPartitionMap(final int tpSize, final boolean changelog) { final SortedMap> taskTopicPartitionMap = new TreeMap<>(); + final String topicName = changelog ? CHANGELOG_TOPIC_PREFIX : TOPIC_PREFIX; for (int i = 0; i < tpSize; i++) { - taskTopicPartitionMap.put(new TaskId(i, 0), mkSet(new TopicPartition("topic" + i, 0))); + taskTopicPartitionMap.put(new TaskId(i, 0), mkSet( + new TopicPartition(topicName + i, 0), + new TopicPartition(topicName + (i + 1) % tpSize, 0) + )); } return taskTopicPartitionMap; } - private SortedMap getRandomClientState(final int clientSize, final int tpSize) { + private InternalTopicManager mockInternalTopicManagerForRandomChangelog(final int nodeSize, final int tpSize) { + final Set changelogNames = new HashSet<>(); + final List nodeList = getRandomNodes(nodeSize); + final Map> topicPartitionInfo = new HashMap<>(); + for (int i = 0; i < tpSize; i++) { + final String topicName = CHANGELOG_TOPIC_PREFIX + i; + changelogNames.add(topicName); + + final Node firstNode = nodeList.get(i % nodeSize); + final Node secondNode = nodeList.get((i + 1) % nodeSize); + final TopicPartitionInfo info = new TopicPartitionInfo(0, firstNode, Arrays.asList(firstNode, secondNode), Collections.emptyList()); + + topicPartitionInfo.computeIfAbsent(topicName, tp -> new ArrayList<>()).add(info); + } + + final MockInternalTopicManager spyTopicManager = spy(mockInternalTopicManager); + doReturn(topicPartitionInfo).when(spyTopicManager).getTopicPartitionInfo(changelogNames); + return spyTopicManager; + } + + private SortedMap getRandomClientState(final int clientSize, final int tpSize, final int maxCapacity) { final SortedMap clientStates = new TreeMap<>(); final List taskIds = new ArrayList<>(tpSize); for (int i = 0; i < tpSize; i++) { taskIds.add(new TaskId(i, 0)); } Collections.shuffle(taskIds); + final Random random = new Random(); for (int i = 0; i < clientSize; i++) { - final ClientState clientState = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, 1); + final int capacity = random.nextInt(maxCapacity) + 1; + final ClientState clientState = new ClientState(emptySet(), emptySet(), emptyMap(), EMPTY_CLIENT_TAGS, capacity, uuidForInt(i)); clientStates.put(uuidForInt(i), clientState); } Iterator> iterator = clientStates.entrySet().iterator(); @@ -750,7 +1148,7 @@ private Cluster getClusterForTopic0And1() { return new Cluster( "cluster", mkSet(NODE_0, NODE_1, NODE_2, NODE_3), - mkSet(PI_0_0, PI_0_1, PI_1_0, PI_1_1, PI_1_2), + mkSet(PI_0_0, PI_0_1, PI_0_2, PI_1_0, PI_1_1, PI_1_2), Collections.emptySet(), Collections.emptySet() ); @@ -796,7 +1194,9 @@ private Map>> getProcessRacksForAllProcess() mkEntry(UUID_2, mkMap(mkEntry("1", Optional.of(RACK_1)))), mkEntry(UUID_3, mkMap(mkEntry("1", Optional.of(RACK_2)))), mkEntry(UUID_4, mkMap(mkEntry("1", Optional.of(RACK_3)))), - mkEntry(UUID_5, mkMap(mkEntry("1", Optional.of(RACK_4)))) + mkEntry(UUID_5, mkMap(mkEntry("1", Optional.of(RACK_4)))), + mkEntry(UUID_6, mkMap(mkEntry("1", Optional.of(RACK_0)))), + mkEntry(UUID_7, mkMap(mkEntry("1", Optional.of(RACK_1)))) ); } @@ -821,6 +1221,12 @@ private Map> getTaskTopicPartitionMapForTask0() { return getTaskTopicPartitionMapForTask0(false); } + private Map> getTaskChangeLogTopicPartitionMapForTask0() { + return mkMap( + mkEntry(TASK_0_0, mkSet(CHANGELOG_TP_0_0)) + ); + } + private Map> getTaskTopicPartitionMapForTask0(final boolean extraTopic) { final Set topicPartitions = new HashSet<>(); topicPartitions.add(TP_0_0); @@ -834,13 +1240,100 @@ private Map> getTaskTopicPartitionMapForAllTasks() { return mkMap( mkEntry(TASK_0_0, mkSet(TP_0_0)), mkEntry(TASK_0_1, mkSet(TP_0_1)), + mkEntry(TASK_0_2, mkSet(TP_0_2)), mkEntry(TASK_1_0, mkSet(TP_1_0)), mkEntry(TASK_1_1, mkSet(TP_1_1)), mkEntry(TASK_1_2, mkSet(TP_1_2)) ); } + private Map> getTaskChangelogMapForAllTasks() { + return mkMap( + mkEntry(TASK_0_0, mkSet(CHANGELOG_TP_0_0)), + mkEntry(TASK_0_1, mkSet(CHANGELOG_TP_0_1)), + mkEntry(TASK_0_2, mkSet(CHANGELOG_TP_0_2)), + mkEntry(TASK_1_0, mkSet(CHANGELOG_TP_1_0)), + mkEntry(TASK_1_1, mkSet(CHANGELOG_TP_1_1)), + mkEntry(TASK_1_2, mkSet(CHANGELOG_TP_1_2)) + ); + } + + private InternalTopicManager mockInternalTopicManagerForChangelog() { + final MockInternalTopicManager spyTopicManager = spy(mockInternalTopicManager); + doReturn( + mkMap( + mkEntry( + CHANGELOG_TP_0_NAME, Arrays.asList( + new TopicPartitionInfo(0, NODE_0, Arrays.asList(REPLICA_0), Collections.emptyList()), + new TopicPartitionInfo(1, NODE_1, Arrays.asList(REPLICA_1), Collections.emptyList()), + new TopicPartitionInfo(2, NODE_1, Arrays.asList(REPLICA_1), Collections.emptyList()) + ) + ), + mkEntry( + CHANGELOG_TP_1_NAME, Arrays.asList( + new TopicPartitionInfo(0, NODE_2, Arrays.asList(REPLICA_2), Collections.emptyList()), + new TopicPartitionInfo(1, NODE_3, Arrays.asList(REPLICA_3), Collections.emptyList()), + new TopicPartitionInfo(2, NODE_0, Arrays.asList(REPLICA_0), Collections.emptyList()) + ) + ) + ) + ).when(spyTopicManager).getTopicPartitionInfo(mkSet(CHANGELOG_TP_0_NAME, CHANGELOG_TP_1_NAME)); + return spyTopicManager; + } + private Map> getTopologyGroupTaskMap() { return Collections.singletonMap(SUBTOPOLOGY_0, Collections.singleton(new TaskId(1, 1))); } + + private void verifyStandbySatisfyRackReplica(final Set taskIds, + final Map racksForProcess, + final Map clientStateMap, + final int replica, + final boolean relaxRackCheck, + final Map standbyTaskCount) { + if (standbyTaskCount != null) { + for (final Entry entry : clientStateMap.entrySet()) { + final int expected = standbyTaskCount.get(entry.getKey()); + final int actual = entry.getValue().standbyTaskCount(); + assertEquals("StandbyTaskCount for " + entry.getKey() + " doesn't match", expected, actual); + } + } + for (final TaskId taskId : taskIds) { + int activeCount = 0; + int standbyCount = 0; + final Map racks = new HashMap<>(); + for (final Map.Entry entry : clientStateMap.entrySet()) { + final UUID processId = entry.getKey(); + final ClientState clientState = entry.getValue(); + + if (!relaxRackCheck && clientState.hasAssignedTask(taskId)) { + final String rack = racksForProcess.get(processId); + assertThat("Task " + taskId + " appears in both " + processId + " and " + racks.get(rack), racks.keySet(), not(hasItems(rack))); + racks.put(rack, processId); + } + + boolean hasActive = false; + if (clientState.hasActiveTask(taskId)) { + activeCount++; + hasActive = true; + } + + boolean hasStandby = false; + if (clientState.hasStandbyTask(taskId)) { + standbyCount++; + hasStandby = true; + } + + assertFalse(clientState + " has both active and standby task " + taskId, hasActive && hasStandby); + } + + assertEquals("Task " + taskId + " should have 1 active task", 1, activeCount); + assertEquals("Task " + taskId + " has wrong replica count", replica, standbyCount); + } + } + + private Map clientTaskCount(final Map clientStateMap, + final Function taskFunc) { + return clientStateMap.entrySet().stream().collect(Collectors.toMap(Entry::getKey, v -> taskFunc.apply(v.getValue()))); + } }