package unit; import java.util.*; /** * EjabberdTest — CWE-407 benchmark for ejabberd-0001 and ejabberd-0002 * * ejabberd-0001 (EJABBERD_MAM_PREFS): * Models should_archive_peer(): archive-preference membership test. * SLOW: lists:member on [ljid()] — O(n_prefs) per archived message * FAST: gb_sets:is_member — O(log n) / effectively O(1) for small sets * * ejabberd-0002 (EJABBERD_SHARED_ROSTER): * Models is_user_in_group() / process_subscription() membership test. * SLOW: lists:member on flat user list — O(n_group_members) per subscription * FAST: gb_sets:from_list + is_member — O(1) amortised per subscription */ public class EjabberdTest { // ------------------------------------------------------------------------- // MAM archive prefs model // ------------------------------------------------------------------------- /** * Simulate should_archive_peer() — called once per archived message. * SLOW: O(|alwaysList|) + O(|neverList|) per call. * Returns number of list comparisons performed. */ static long shouldArchivePeer_slow(List alwaysList, List neverList, String peer) { long ops = 0; for (String jid : alwaysList) { ops++; if (jid.equals(peer)) return ops; // match in always } for (String jid : neverList) { ops++; if (jid.equals(peer)) return ops; // match in never } return ops; // default } /** * FAST: O(1) HashSet lookup (models gb_sets:is_member after from_list). * Returns ops = 1 per lookup (hash probe). */ static long shouldArchivePeer_fast(Set alwaysSet, Set neverSet, String peer) { long ops = 1; if (alwaysSet.contains(peer)) return ops; ops++; if (neverSet.contains(peer)) return ops; return ops; } // ------------------------------------------------------------------------- // Shared roster model // ------------------------------------------------------------------------- /** * Models is_user_in_group(): rebuild list + linear scan. * SLOW: O(n_members) per call. */ static long isUserInGroup_slow(List groupUsers, String us) { long ops = 0; for (String u : groupUsers) { ops++; if (u.equals(us)) return ops; } return ops; } /** * FAST: O(1) set lookup (models gb_sets:from_list + is_member). */ static long isUserInGroup_fast(Set groupSet, String us) { groupSet.contains(us); // O(1) return 1; } /** * Models process_subscription(): build flat SRUsers list from all groups, * then lists:member(US1, SRUsers). Called once per subscription stanza. * The target is NOT present — full list scan every call. * SLOW: O(total_users_across_groups) per call. * Returns ops = list comparisons performed. */ static long processSubscription_slow(List srUsers, String targetUser) { long ops = 0; for (String u : srUsers) { ops++; if (u.equals(targetUser)) return ops; } return ops; } /** * FAST: pre-built HashSet passed in (built once, reused across calls). * O(1) per call. Returns ops = 1. */ static long processSubscription_fast(Set srSet, String targetUser) { srSet.contains(targetUser); return 1; } // ------------------------------------------------------------------------- // Benchmark harness // ------------------------------------------------------------------------- static void bench(String label, Runnable slow, Runnable fast, long sOps, long fOps) { slow.run(); fast.run(); long t0 = System.nanoTime(); slow.run(); long sMs = (System.nanoTime() - t0) / 1_000_000; long t1 = System.nanoTime(); fast.run(); long fMs = (System.nanoTime() - t1) / 1_000_000; double r = fOps > 0 ? (double) sOps / fOps : 0; System.out.printf(" %-52s slow:%4dms (%,d ops) fast:%4dms (%,d ops) speedup:%.0fx%n", label, sMs, sOps, fMs, fOps, r); } // ------------------------------------------------------------------------- // Setup helpers // ------------------------------------------------------------------------- static List makeJidList(String prefix, int n) { List list = new ArrayList<>(n); for (int i = 0; i < n; i++) list.add(prefix + i + "@example.org"); return list; } static Set toHashSet(List list) { return new HashSet<>(list); } // ------------------------------------------------------------------------- // Main // ------------------------------------------------------------------------- public static void main(String[] args) { System.out.println("EjabberdTest — CWE-407 benchmarks: ejabberd-0001 + ejabberd-0002"); System.out.println("=================================================================="); int passed = 0, total = 0; // ----- ejabberd-0001: MAM archive prefs ----- System.out.println("\n[ejabberd-0001] mod_mam: should_archive_peer lists:member"); { // N=200 contacts in always+never lists; target is NOT in either (worst case: full scan) int N = 200; List alwaysList = makeJidList("always", N / 2); List neverList = makeJidList("never", N / 2); Set alwaysSet = toHashSet(alwaysList); Set neverSet = toHashSet(neverList); String peer = "unknown@remote.org"; // not in either list — worst case int MSGS = 50_000; long[] sOps = {0}, fOps = {0}; Runnable slow = () -> { long ops = 0; for (int i = 0; i < MSGS; i++) ops += shouldArchivePeer_slow(alwaysList, neverList, peer); sOps[0] = ops; }; Runnable fast = () -> { long ops = 0; for (int i = 0; i < MSGS; i++) ops += shouldArchivePeer_fast(alwaysSet, neverSet, peer); fOps[0] = ops; }; slow.run(); fast.run(); bench("MAM prefs N=200 unknown-peer (50k msgs)", slow, fast, sOps[0], fOps[0]); total++; assert sOps[0] > fOps[0] * 50 : "FAIL: slow=" + sOps[0] + " fast=" + fOps[0]; passed++; } { // N=500 worst case — large block/allow lists int N = 500; List alwaysList = makeJidList("always", N / 2); List neverList = makeJidList("never", N / 2); Set alwaysSet = toHashSet(alwaysList); Set neverSet = toHashSet(neverList); String peer = "unknown2@remote.org"; int MSGS = 20_000; long[] sOps = {0}, fOps = {0}; Runnable slow = () -> { long ops = 0; for (int i = 0; i < MSGS; i++) ops += shouldArchivePeer_slow(alwaysList, neverList, peer); sOps[0] = ops; }; Runnable fast = () -> { long ops = 0; for (int i = 0; i < MSGS; i++) ops += shouldArchivePeer_fast(alwaysSet, neverSet, peer); fOps[0] = ops; }; slow.run(); fast.run(); bench("MAM prefs N=500 unknown-peer (20k msgs)", slow, fast, sOps[0], fOps[0]); total++; assert sOps[0] > fOps[0] * 100 : "FAIL: slow=" + sOps[0] + " fast=" + fOps[0]; passed++; } // ----- ejabberd-0002: shared roster ----- System.out.println("\n[ejabberd-0002] mod_shared_roster: is_user_in_group + process_subscription"); { // is_user_in_group: 1000-member group, target not present (worst case) int G = 1000; List groupUsers = makeJidList("member", G); Set groupSet = toHashSet(groupUsers); String targetUser = "newuser@example.org"; // not in group int CALLS = 20_000; long[] sOps = {0}, fOps = {0}; Runnable slow = () -> { long ops = 0; for (int i = 0; i < CALLS; i++) ops += isUserInGroup_slow(groupUsers, targetUser); sOps[0] = ops; }; Runnable fast = () -> { long ops = 0; for (int i = 0; i < CALLS; i++) ops += isUserInGroup_fast(groupSet, targetUser); fOps[0] = ops; }; slow.run(); fast.run(); bench("is_user_in_group G=1000 not-member (20k)", slow, fast, sOps[0], fOps[0]); total++; assert sOps[0] > fOps[0] * 500 : "FAIL: slow=" + sOps[0] + " fast=" + fOps[0]; passed++; } { // process_subscription: 5 displayed groups × 500 members each = 2500 total. // The SRUsers flat list is rebuilt per subscription in the slow path (lists:usort + // flatmap each call). In the fast path the set is built once and reused. // Target not in any group (worst case: full list scan every call). int GROUPS = 5, PER_GROUP = 500; List srUsers = new ArrayList<>(); for (int g = 0; g < GROUPS; g++) srUsers.addAll(makeJidList("grp" + g + "member", PER_GROUP)); Set srSet = toHashSet(srUsers); String targetUser = "outsider@other.org"; int SUBS = 20_000; long[] sOps = {0}, fOps = {0}; Runnable slow = () -> { long ops = 0; for (int i = 0; i < SUBS; i++) ops += processSubscription_slow(srUsers, targetUser); sOps[0] = ops; }; Runnable fast = () -> { long ops = 0; for (int i = 0; i < SUBS; i++) ops += processSubscription_fast(srSet, targetUser); fOps[0] = ops; }; slow.run(); fast.run(); bench("process_subscription 5grp×500 not-member (20k)", slow, fast, sOps[0], fOps[0]); total++; assert sOps[0] > fOps[0] * 1000 : "FAIL: slow=" + sOps[0] + " fast=" + fOps[0]; passed++; } System.out.println("\n" + passed + "/" + total + " PASS"); if (passed < total) System.exit(1); } }