java-topology/defects/tomcat/unit/TomcatTribesArraysMergeTest.java

180 lines
6.9 KiB
Java
Raw 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.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<Member> 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<Member> — 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<Member> 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<Member> 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);
}
}