java-topology/defects/opensearch/unit/ShardRoutingWeightedContains.java

113 lines
4.7 KiB
Java
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package unit;
import java.util.ArrayList;
import java.util.HashSet;
import java.util.List;
import java.util.Set;
import java.util.stream.Collectors;
/**
* CWE-407 unit test: opensearch-003
* IndexShardRoutingTable.populateActiveShardWeightsMap() uses List.contains() inside a
* .filter() stream to compute non-weighted shard routings.
*
* Defect: O(n²) — weightedRoutings (ArrayList).contains() called per element of activeShards
* Fix: O(n) — HashSet.contains() is O(1)
*
* No JUnit. Run with: javac ShardRoutingWeightedContains.java && java -cp . unit.ShardRoutingWeightedContains
*/
public class ShardRoutingWeightedContains {
// Minimal stand-in for ShardRouting (an integer ID)
static class FakeShard {
final int id;
FakeShard(int id) { this.id = id; }
@Override public boolean equals(Object o) {
return o instanceof FakeShard && ((FakeShard) o).id == this.id;
}
@Override public int hashCode() { return Integer.hashCode(id); }
@Override public String toString() { return "S" + id; }
}
// ---- SLOW: mirrors the defective populateActiveShardWeightsMap logic ----
static List<FakeShard> computeNonWeightedSlow(List<FakeShard> allShards, List<FakeShard> weightedShards) {
return allShards.stream()
.filter(shard -> !weightedShards.contains(shard)) // O(weightedShards.size()) per shard
.collect(Collectors.toList());
}
// ---- FAST: convert weightedShards to a HashSet first ----
static List<FakeShard> computeNonWeightedFast(List<FakeShard> allShards, List<FakeShard> weightedShards) {
Set<FakeShard> weightedSet = new HashSet<>(weightedShards); // O(n) to build
return allShards.stream()
.filter(shard -> !weightedSet.contains(shard)) // O(1) per shard
.collect(Collectors.toList());
}
// ---- count .contains() calls for slow path ----
static long countSlowContainsCalls(int allShardsSize) {
// For each of allShardsSize shards, contains() scans weightedShards (same size in worst case)
return (long) allShardsSize * allShardsSize;
}
// ---- count .contains() calls for fast path ----
static long countFastContainsCalls(int allShardsSize) {
return allShardsSize; // one O(1) HashSet lookup per shard
}
public static void main(String[] args) {
System.out.println("=== opensearch-003: IndexShardRoutingTable weighted routing CWE-407 ===\n");
// --- Correctness check ---
// 10 shards total; first 6 are "weighted", last 4 are "non-weighted"
List<FakeShard> allShards = new ArrayList<>();
for (int i = 0; i < 10; i++) allShards.add(new FakeShard(i));
List<FakeShard> weighted = allShards.subList(0, 6);
List<FakeShard> slowResult = computeNonWeightedSlow(allShards, weighted);
List<FakeShard> fastResult = computeNonWeightedFast(allShards, weighted);
if (!slowResult.equals(fastResult)) {
System.out.println("FAIL correctness: slow=" + slowResult + " fast=" + fastResult);
System.exit(1);
}
if (slowResult.size() != 4) {
System.out.println("FAIL: expected 4 non-weighted shards, got " + slowResult.size());
System.exit(1);
}
System.out.println("correctness OK nonWeighted=" + slowResult);
// --- Op-count comparison at scale ---
System.out.println("\n=== Op-count: worst case (all shards weighted, allShards == weightedShards) ===");
System.out.printf("%-8s %-16s %-14s %-10s%n", "n", "slow_contains", "fast_contains", "ratio");
System.out.println("-".repeat(56));
int[] sizes = {10, 50, 100, 200, 500};
for (int n : sizes) {
long slowOps = countSlowContainsCalls(n);
long fastOps = countFastContainsCalls(n);
double ratio = (double) slowOps / fastOps;
System.out.printf("%-8d %-16d %-14d %-10.1f%n", n, slowOps, fastOps, ratio);
if (slowOps <= fastOps && n > 1) {
System.out.println("FAIL: slow was not worse than fast at n=" + n);
System.exit(1);
}
}
// --- Verify at n=200: should be ≥ 100× worse ----
int n = 200;
long slowOps = countSlowContainsCalls(n);
long fastOps = countFastContainsCalls(n);
double ratio = (double) slowOps / fastOps;
if (ratio < 100.0) {
System.out.printf("FAIL: expected ratio >= 100x at n=200, got %.1fx%n", ratio);
System.exit(1);
}
System.out.printf("%nspeedup at n=200 shards: %.1fx PASS%n", ratio);
System.out.println("\nALL PASS");
}
}