package unit; import java.util.ArrayList; import java.util.LinkedHashSet; /** * Unit test for CWE-407 tomcat-0002: * Arrays.merge() uses ArrayList.contains() for Member deduplication — O(M×N). * * Real code (java/org/apache/catalina/tribes/util/Arrays.java): * ArrayList list = new ArrayList<>(Arrays.asList(m1)); * for (Member member : m2) { * if (!list.contains(member)) { // O(N) scan → O(M×N) total * list.add(member); * } * } * * Fix: LinkedHashSet — O(1) add, idempotent deduplication. * * Run: javac -d . TomcatTribesArraysMergeTest.java && java -ea unit.TomcatTribesArraysMergeTest */ public class TomcatTribesArraysMergeTest { // ----------------------------------------------------------------------- // Minimal Member stand-in — identity = id (integer) // hashCode/equals match MemberImpl contract (host-based identity) // ----------------------------------------------------------------------- static final class Member { final int id; Member(int id) { this.id = id; } @Override public boolean equals(Object o) { return o instanceof Member && ((Member) o).id == this.id; } @Override public int hashCode() { return id; } } // ----------------------------------------------------------------------- // Slow path: ArrayList.contains() inside loop (O(M×N)) // ----------------------------------------------------------------------- static long[] slowMerge(Member[] m1, Member[] m2) { long ops = 0; ArrayList list = new ArrayList<>(); for (Member m : m1) list.add(m); for (Member member : m2) { // Simulate ArrayList.contains() cost: count equals() calls boolean found = false; for (Member existing : list) { ops++; if (existing.equals(member)) { found = true; break; } } if (!found) list.add(member); } return new long[]{list.size(), ops}; } // ----------------------------------------------------------------------- // Fast path: LinkedHashSet.add() (O(M+N)) // ----------------------------------------------------------------------- static long[] fastMerge(Member[] m1, Member[] m2) { LinkedHashSet set = new LinkedHashSet<>(); for (Member m : m1) set.add(m); long ops = 0; for (Member member : m2) { ops++; // one hash + equals in best case set.add(member); } return new long[]{set.size(), ops}; } // ----------------------------------------------------------------------- // Benchmark helpers // ----------------------------------------------------------------------- static long timeSlowMerge(int n, int repeats) { Member[] m1 = new Member[n]; Member[] m2 = new Member[n]; // fully overlapping → worst case for slow path for (int i = 0; i < n; i++) { m1[i] = new Member(i); m2[i] = new Member(i); } // warmup for (int r = 0; r < 3; r++) slowMerge(m1, m2); long t0 = System.nanoTime(); for (int r = 0; r < repeats; r++) slowMerge(m1, m2); return (System.nanoTime() - t0) / 1_000_000; } static long timeFastMerge(int n, int repeats) { Member[] m1 = new Member[n]; Member[] m2 = new Member[n]; for (int i = 0; i < n; i++) { m1[i] = new Member(i); m2[i] = new Member(i); } for (int r = 0; r < 3; r++) fastMerge(m1, m2); long t0 = System.nanoTime(); for (int r = 0; r < repeats; r++) fastMerge(m1, m2); return (System.nanoTime() - t0) / 1_000_000; } // ----------------------------------------------------------------------- // Tests // ----------------------------------------------------------------------- static int passed = 0; static int failed = 0; static void assertOpsRatio(String label, long slowOps, long fastOps, double minRatio) { double ratio = fastOps > 0 ? (double) slowOps / fastOps : slowOps; boolean ok = ratio >= minRatio; System.out.printf(" %-55s slowOps:%,6d fastOps:%,6d ratio:%.1fx %s%n", label, slowOps, fastOps, ratio, ok ? "PASS" : "FAIL"); if (ok) passed++; else failed++; } static void assertTimeRatio(String label, long slowMs, long fastMs, double minRatio) { double ratio = fastMs > 0 ? (double) slowMs / fastMs : (slowMs > 0 ? 100.0 : 1.0); boolean ok = ratio >= minRatio; System.out.printf(" %-55s slow:%4dms fast:%4dms ratio:%.1fx %s%n", label, slowMs, fastMs, ratio, ok ? "PASS" : "FAIL"); if (ok) passed++; else failed++; } static void assertCorrectness(String label, int n, boolean expectedSize) { Member[] m1 = new Member[n]; Member[] m2 = new Member[n / 2]; for (int i = 0; i < n; i++) m1[i] = new Member(i); for (int i = 0; i < n / 2; i++) m2[i] = new Member(i + n / 2); // half overlap long[] slow = slowMerge(m1, m2); long[] fast = fastMerge(m1, m2); boolean ok = slow[0] == fast[0]; System.out.printf(" %-55s slowSize:%d fastSize:%d %s%n", label, slow[0], fast[0], ok ? "PASS" : "FAIL"); if (ok) passed++; else failed++; } public static void main(String[] args) { System.out.println("=== tomcat-0002: Arrays.merge() ArrayList.contains() O(M×N) ==="); System.out.println(); // --- Op-count tests --- System.out.println("Op-count comparison (equals() calls, worst-case all-overlap):"); int[] sizes = {50, 100, 200, 500}; for (int n : sizes) { Member[] m1 = new Member[n]; Member[] m2 = new Member[n]; for (int i = 0; i < n; i++) { m1[i] = new Member(i); m2[i] = new Member(i); } long slowOps = slowMerge(m1, m2)[1]; long fastOps = fastMerge(m1, m2)[1]; // At n=50: slow does ~50*50/2=1250 ops, fast does 50 ops → ratio ~25x double minRatio = n >= 200 ? 50.0 : 10.0; assertOpsRatio(String.format("merge all-overlap N=%d", n), slowOps, fastOps, minRatio); } System.out.println(); // --- Correctness tests --- System.out.println("Correctness (result size matches between slow and fast):"); assertCorrectness("merge half-overlap N=100", 100, true); assertCorrectness("merge half-overlap N=500", 500, true); System.out.println(); // --- Wall-clock timing --- System.out.println("Wall-clock timing (N=2000, 500 repeats):"); long slowMs = timeSlowMerge(2000, 500); long fastMs = timeFastMerge(2000, 500); assertTimeRatio("merge N=2000 all-overlap x500 repeats", slowMs, fastMs, 5.0); System.out.println(); System.out.printf("Result: %d/%d PASS%n", passed, passed + failed); if (failed > 0) System.exit(1); } }