package unit; import java.util.*; /** * kafka-0004: AbstractStickyAssignor prepopulateCurrentAssignments — * consumerSubscription.topics().contains(partition.topic()) where topics() returns List. * * Inside a double for loop (consumers x partitions), this gives O(C * P * T). * Fix: convert subscription topics to HashSet before the inner loop. */ public class KafkaStickyAssignorTopicsTest { // Simulate the slow path: subscription.topics() returns List // Returns total number of element comparisons performed static long slowPrepopulate( Map> consumerToTopics, // consumer -> topic subscription list Map> consumerToPartitions // consumer -> assigned partitions (topic names) ) { long ops = 0; for (Map.Entry> entry : consumerToPartitions.entrySet()) { String consumer = entry.getKey(); List subTopics = consumerToTopics.get(consumer); // returns List for (String partitionTopic : entry.getValue()) { // List.contains: O(T) scan for (String t : subTopics) { ops++; if (t.equals(partitionTopic)) break; } } } return ops; } // Simulate the fast path: subscription topics converted to HashSet static long fastPrepopulate( Map> consumerToTopicsSet, Map> consumerToPartitions ) { long ops = 0; for (Map.Entry> entry : consumerToPartitions.entrySet()) { String consumer = entry.getKey(); Set subTopics = consumerToTopicsSet.get(consumer); for (String partitionTopic : entry.getValue()) { ops++; // O(1) HashSet lookup counts as 1 op subTopics.contains(partitionTopic); } } return ops; } public static void main(String[] args) { int passed = 0; int total = 0; // Test 1: basic op-count comparison { total++; int C = 20; // consumers int P = 100; // partitions per consumer int T = 30; // topics per consumer (worst-case: partition topic is last in list) Map> consumerToTopics = new HashMap<>(); Map> consumerToTopicsSet = new HashMap<>(); Map> consumerToPartitions = new HashMap<>(); for (int c = 0; c < C; c++) { String consumerId = "consumer-" + c; List topics = new ArrayList<>(); Set topicsSet = new HashSet<>(); for (int t = 0; t < T; t++) { String topic = "topic-" + t; topics.add(topic); topicsSet.add(topic); } consumerToTopics.put(consumerId, topics); consumerToTopicsSet.put(consumerId, topicsSet); List partitions = new ArrayList<>(); for (int p = 0; p < P; p++) { // worst case: topic is last in subscription list partitions.add("topic-" + (T - 1)); } consumerToPartitions.put(consumerId, partitions); } long slowOps = slowPrepopulate(consumerToTopics, consumerToPartitions); long fastOps = fastPrepopulate(consumerToTopicsSet, consumerToPartitions); long expectedSlowOps = (long) C * P * T; // every contains scans all T long expectedFastOps = (long) C * P; assert slowOps == expectedSlowOps : "Expected slowOps=" + expectedSlowOps + " got " + slowOps; assert fastOps == expectedFastOps : "Expected fastOps=" + expectedFastOps + " got " + fastOps; long speedup = slowOps / fastOps; assert speedup == T : "Expected speedup=" + T + " got " + speedup; System.out.println(" Test 1 PASS: slow=" + slowOps + " fast=" + fastOps + " speedup=" + speedup + "x"); passed++; } // Test 2: correctness — both paths agree on membership { total++; List subTopics = Arrays.asList("sports", "news", "tech", "finance"); Set subTopicsSet = new HashSet<>(subTopics); String[] testPartitionTopics = {"sports", "news", "tech", "finance", "other"}; boolean[] expected = {true, true, true, true, false}; for (int i = 0; i < testPartitionTopics.length; i++) { boolean listResult = subTopics.contains(testPartitionTopics[i]); boolean setResult = subTopicsSet.contains(testPartitionTopics[i]); assert listResult == expected[i] : "List result wrong for " + testPartitionTopics[i]; assert setResult == expected[i] : "Set result wrong for " + testPartitionTopics[i]; assert listResult == setResult : "List and Set disagree for " + testPartitionTopics[i]; } System.out.println(" Test 2 PASS: List and Set agree on all " + testPartitionTopics.length + " membership checks"); passed++; } // Test 3: scaling — speedup grows linearly with T { total++; for (int T : new int[]{10, 50, 100, 200}) { List topics = new ArrayList<>(); for (int t = 0; t < T; t++) topics.add("t" + t); Set topicsSet = new HashSet<>(topics); String worstCase = "t" + (T - 1); // List: scans all T long listWork = T; // HashSet: O(1) long hashWork = 1; assert topics.contains(worstCase) == topicsSet.contains(worstCase) : "Correctness check failed at T=" + T; assert listWork / hashWork == T : "speedup should == T"; } System.out.println(" Test 3 PASS: speedup scales linearly with T for T in [10,50,100,200]"); passed++; } System.out.println(passed + "/" + total + " PASS"); if (passed != total) System.exit(1); } }