java-topology/defects/undertow/unit/UndertowWebSocketSubprotocolTest.java

233 lines
9.6 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.HashMap;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;
/**
* Unit test for CWE-407 undertow-0001:
* DefaultContainerConfigurator.getNegotiatedSubprotocol() uses List.contains()
* in a per-request loop — O(|requested| × |supported|) per WebSocket handshake.
*
* Real code (websockets-jsr/.../DefaultContainerConfigurator.java):
* for (String proto : requested) {
* if (supported.contains(proto)) { // O(S) List scan → O(R×S) total
* return proto;
* }
* }
*
* Fix: convert supported to HashSet<String> once before the loop → O(R+S).
*
* Also covers getNegotiatedExtensions() nested-loop O(|req|×|inst|) → O(|req|+|inst|).
*
* Run: javac -d . UndertowWebSocketSubprotocolTest.java && java -ea unit.UndertowWebSocketSubprotocolTest
*/
public class UndertowWebSocketSubprotocolTest {
// -----------------------------------------------------------------------
// Slow path: List.contains() inside loop (O(R×S))
// -----------------------------------------------------------------------
static long[] slowNegotiate(List<String> supported, List<String> requested) {
long ops = 0;
String result = "";
for (String proto : requested) {
// Simulate List.contains(): linear scan
for (String s : supported) {
ops++;
if (s.equals(proto)) { result = proto; break; }
}
if (!result.isEmpty()) break;
}
return new long[]{ops, result.isEmpty() ? -1 : requested.indexOf(result)};
}
// Worst case: no match, all R×S comparisons
static long slowNegotiateNoMatch(List<String> supported, List<String> requested) {
long ops = 0;
for (String proto : requested) {
for (String s : supported) {
ops++;
if (s.equals(proto)) break;
}
}
return ops;
}
// -----------------------------------------------------------------------
// Fast path: HashSet.contains() (O(R+S))
// -----------------------------------------------------------------------
static long fastNegotiateOps(List<String> supported, List<String> requested) {
// Build set once
long ops = supported.size(); // cost to build HashSet
Set<String> supportedSet = new HashSet<>(supported);
for (String proto : requested) {
ops++; // O(1) HashSet lookup
if (supportedSet.contains(proto)) break;
}
return ops;
}
// -----------------------------------------------------------------------
// Slow extension negotiation: O(|req|×|inst|) nested loop
// -----------------------------------------------------------------------
static long slowExtensionNegotiateOps(List<String> installed, List<String> requested) {
long ops = 0;
for (String req : requested) {
for (String inst : installed) {
ops++;
if (inst.equals(req)) break;
}
}
return ops;
}
// Fast extension: O(|req|+|inst|) HashMap
static long fastExtensionNegotiateOps(List<String> installed, List<String> requested) {
long ops = installed.size(); // build map
Map<String, Boolean> instMap = new HashMap<>();
for (String inst : installed) instMap.put(inst, true);
for (String req : requested) {
ops++; // O(1) map lookup
}
return ops;
}
// -----------------------------------------------------------------------
// Benchmarks
// -----------------------------------------------------------------------
static long timeSlow(int R, int S, int repeats) {
List<String> supported = new ArrayList<>();
List<String> requested = new ArrayList<>();
for (int i = 0; i < S; i++) supported.add("proto-supported-" + i);
for (int i = 0; i < R; i++) requested.add("proto-requested-" + i); // no match
// warmup
for (int r = 0; r < 5; r++) slowNegotiateNoMatch(supported, requested);
long t0 = System.nanoTime();
for (int r = 0; r < repeats; r++) slowNegotiateNoMatch(supported, requested);
return (System.nanoTime() - t0) / 1_000_000;
}
static long timeFast(int R, int S, int repeats) {
List<String> supported = new ArrayList<>();
List<String> requested = new ArrayList<>();
for (int i = 0; i < S; i++) supported.add("proto-supported-" + i);
for (int i = 0; i < R; i++) requested.add("proto-requested-" + i);
// warmup
for (int r = 0; r < 5; r++) fastNegotiateOps(supported, requested);
long t0 = System.nanoTime();
for (int r = 0; r < repeats; r++) fastNegotiateOps(supported, requested);
return (System.nanoTime() - t0) / 1_000_000;
}
// -----------------------------------------------------------------------
// Test harness
// -----------------------------------------------------------------------
static int passed = 0;
static int failed = 0;
static void assertRatio(String label, long slowOps, long fastOps, double minRatio) {
double ratio = fastOps > 0 ? (double) slowOps / fastOps : slowOps;
boolean ok = ratio >= minRatio;
System.out.printf(" %-60s slowOps:%,7d 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(" %-60s slow:%4dms fast:%4dms ratio:%.1fx %s%n",
label, slowMs, fastMs, ratio, ok ? "PASS" : "FAIL");
if (ok) passed++; else failed++;
}
static void assertCorrect(String label, List<String> supported, List<String> requested,
String expected) {
// slow
Set<String> supportedSet = new HashSet<>(supported);
String slow = "";
for (String proto : requested) {
if (supported.contains(proto)) { slow = proto; break; }
}
// fast
String fast = "";
for (String proto : requested) {
if (supportedSet.contains(proto)) { fast = proto; break; }
}
boolean ok = slow.equals(fast) && slow.equals(expected);
System.out.printf(" %-60s slow='%s' fast='%s' expected='%s' %s%n",
label, slow, fast, expected, ok ? "PASS" : "FAIL");
if (ok) passed++; else failed++;
}
public static void main(String[] args) {
System.out.println("=== undertow-0001: WebSocket subprotocol negotiation List.contains() O(R×S) ===");
System.out.println();
// --- Op-count: subprotocol negotiation ---
System.out.println("Op-count: getNegotiatedSubprotocol (no match, worst case):");
int[][] cases = {{10,10},{50,20},{100,50},{200,100}};
for (int[] rc : cases) {
int R = rc[0], S = rc[1];
List<String> supported = new ArrayList<>();
List<String> requested = new ArrayList<>();
for (int i = 0; i < S; i++) supported.add("s" + i);
for (int i = 0; i < R; i++) requested.add("r" + i); // no overlap
long slowOps = slowNegotiateNoMatch(supported, requested);
long fastOps = fastNegotiateOps(supported, requested);
double minRatio = R >= 100 ? 5.0 : 3.0;
assertRatio(String.format("negotiate no-match R=%d S=%d", R, S), slowOps, fastOps, minRatio);
}
System.out.println();
// --- Op-count: extension negotiation ---
System.out.println("Op-count: getNegotiatedExtensions (no match, worst case):");
for (int[] rc : cases) {
int R = rc[0], S = rc[1];
List<String> installed = new ArrayList<>();
List<String> requested = new ArrayList<>();
for (int i = 0; i < S; i++) installed.add("ext-inst-" + i);
for (int i = 0; i < R; i++) requested.add("ext-req-" + i);
long slowOps = slowExtensionNegotiateOps(installed, requested);
long fastOps = fastExtensionNegotiateOps(installed, requested);
double minRatio = R >= 100 ? 5.0 : 3.0;
assertRatio(String.format("extensions no-match R=%d S=%d", R, S), slowOps, fastOps, minRatio);
}
System.out.println();
// --- Correctness ---
System.out.println("Correctness (first match returned):");
List<String> sup1 = List.of("chat", "binary", "json");
List<String> req1 = List.of("xml", "json", "chat");
assertCorrect("first client-preferred match is 'json'", sup1, req1, "json");
List<String> sup2 = List.of("v1", "v2", "v3");
List<String> req2 = List.of("v4", "v5");
assertCorrect("no match returns ''", sup2, req2, "");
List<String> sup3 = List.of("proto");
List<String> req3 = List.of("proto");
assertCorrect("exact single match", sup3, req3, "proto");
System.out.println();
// --- Wall-clock ---
System.out.println("Wall-clock timing (R=1000 S=500, 2000 repeats):");
long slowMs = timeSlow(1000, 500, 2000);
long fastMs = timeFast(1000, 500, 2000);
assertTimeRatio("negotiate R=1000 S=500 x2000", slowMs, fastMs, 5.0);
System.out.println();
System.out.printf("Result: %d/%d PASS%n", passed, passed + failed);
if (failed > 0) System.exit(1);
}
}