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