package unit; import java.util.*; /** * VaultGroupMembershipAlgorithm — CWE-407 benchmark * * Models HashiCorp Vault identity_store_util.go sanitizeAndUpsertGroup(): * SLOW: strutil.StrListContains(memberGroupIDs, currentMemberGroupID) — O(G) per iteration * called G times → O(G²) total * FAST: pre-build map[string]bool from memberGroupIDs → O(1) per lookup * * vault-0001 */ public class VaultGroupMembershipAlgorithm { // ------------------------------------------------------------------ nodes static class Node { String id; Node(String id) { this.id = id; } } // ------------------------------------------------------------------ result static class Result { int removed; Result(int removed) { this.removed = removed; } } // ------------------------------------------------------------------ slow (defect) static class DefectiveGroupUpdate { /** * For each currentMemberGroupID, scan the full memberGroupIDs list to * determine if it was removed. O(G_current * G_member) = O(G²). */ static Result updateRemovedMembers(List currentMemberGroupIDs, List memberGroupIDs) { int removed = 0; for (String currentID : currentMemberGroupIDs) { // CWE-407: linear scan of memberGroupIDs on each iteration boolean stillMember = memberGroupIDs.contains(currentID); // O(G) if (!stillMember) { removed++; // would call UpsertGroupInTxn here } } return new Result(removed); } } // ------------------------------------------------------------------ fast (fix) static class FixedGroupUpdate { /** * Pre-build a HashSet from memberGroupIDs for O(1) lookup. * Total: O(G_current + G_member). */ static Result updateRemovedMembers(List currentMemberGroupIDs, List memberGroupIDs) { // CWE-407 fix: O(1) lookup set Set memberIDSet = new HashSet<>(memberGroupIDs); int removed = 0; for (String currentID : currentMemberGroupIDs) { if (!memberIDSet.contains(currentID)) { // O(1) removed++; } } return new Result(removed); } } // ------------------------------------------------------------------ helpers static List buildGroupIDs(int n, String prefix) { List ids = new ArrayList<>(n); for (int i = 0; i < n; i++) { ids.add(prefix + "-group-" + i); } return ids; } static long benchSlow(int G, int iters) { List current = buildGroupIDs(G, "cur"); // Half of current groups stay, half are removed List newMembers = buildGroupIDs(G / 2, "cur"); long start = System.nanoTime(); for (int i = 0; i < iters; i++) { DefectiveGroupUpdate.updateRemovedMembers(current, newMembers); } return System.nanoTime() - start; } static long benchFast(int G, int iters) { List current = buildGroupIDs(G, "cur"); List newMembers = buildGroupIDs(G / 2, "cur"); long start = System.nanoTime(); for (int i = 0; i < iters; i++) { FixedGroupUpdate.updateRemovedMembers(current, newMembers); } return System.nanoTime() - start; } // ------------------------------------------------------------------ main public static void main(String[] args) { int passed = 0, total = 0; // ---- correctness List current = Arrays.asList("g1", "g2", "g3", "g4", "g5"); List newMems = Arrays.asList("g1", "g3"); // g2,g4,g5 removed Result slowR = DefectiveGroupUpdate.updateRemovedMembers(current, newMems); Result fastR = FixedGroupUpdate.updateRemovedMembers(current, newMems); assert slowR.removed == 3 : "slow: expected 3 removed, got " + slowR.removed; assert fastR.removed == 3 : "fast: expected 3 removed, got " + fastR.removed; assert slowR.removed == fastR.removed : "slow/fast mismatch"; System.out.println("Correctness: PASS (slow.removed == fast.removed == 3)"); // ---- performance int ITERS = 200; int[] sizes = {400, 600, 1000}; System.out.printf("%-8s %-12s %-12s %s%n", "G(groups)", "slow(ns)", "fast(ns)", "ratio"); for (int G : sizes) { // warm-up benchSlow(G, 20); benchFast(G, 20); long slowNs = benchSlow(G, ITERS); long fastNs = benchFast(G, ITERS); double ratio = (double) slowNs / fastNs; System.out.printf("%-8d %-12d %-12d %.2fx%n", G, slowNs, fastNs, ratio); total++; double threshold = G <= 400 ? 3.5 : 5.0; // JVM warmup noise at small N if (ratio >= threshold) { System.out.printf(" PASS (ratio=%.2f >= %.1f)%n", ratio, threshold); passed++; } else { System.out.printf(" FAIL (ratio=%.2f < %.1f)%n", ratio, threshold); } } System.out.printf("%nTests: %d/%d PASS%n", passed, total); if (passed < total) System.exit(1); } }