diff --git a/dd-java-agent/instrumentation/kafka/kafka-clients-0.11/src/main/java/datadog/trace/instrumentation/kafka_clients/TextMapExtractAdapter.java b/dd-java-agent/instrumentation/kafka/kafka-clients-0.11/src/main/java/datadog/trace/instrumentation/kafka_clients/TextMapExtractAdapter.java index 429038773ba..b7aa2936643 100644 --- a/dd-java-agent/instrumentation/kafka/kafka-clients-0.11/src/main/java/datadog/trace/instrumentation/kafka_clients/TextMapExtractAdapter.java +++ b/dd-java-agent/instrumentation/kafka/kafka-clients-0.11/src/main/java/datadog/trace/instrumentation/kafka_clients/TextMapExtractAdapter.java @@ -6,6 +6,7 @@ import static datadog.trace.instrumentation.kafka_clients.KafkaDecorator.KAFKA_PRODUCED_KEY; import datadog.trace.api.Config; +import datadog.trace.api.Functions; import datadog.trace.bootstrap.instrumentation.api.AgentPropagation; import datadog.trace.bootstrap.instrumentation.api.AgentPropagation.ContextVisitor; import java.nio.ByteBuffer; @@ -20,14 +21,21 @@ public class TextMapExtractAdapter implements ContextVisitor { private static final Logger log = LoggerFactory.getLogger(TextMapExtractAdapter.class); public static final TextMapExtractAdapter GETTER = - new TextMapExtractAdapter(Config.get().isKafkaClientBase64DecodingEnabled()); + new TextMapExtractAdapter( + Config.get().isKafkaClientBase64DecodingEnabled(), + Config.get().isKafkaClientBase64DecodingGuardEnabled()); private final Function headerValueTransformer; private final Base64.Decoder decoder; public TextMapExtractAdapter(boolean decodeBase64Headers) { + this(decodeBase64Headers, true); + } + + public TextMapExtractAdapter(boolean decodeBase64Headers, boolean guardEnabled) { if (decodeBase64Headers) { - this.headerValueTransformer = BASE64_DECODE; + this.headerValueTransformer = + guardEnabled ? new Functions.GuardedBase64Decode()::tryApply : BASE64_DECODE; this.decoder = Base64.getDecoder(); } else { this.headerValueTransformer = UTF8_BYTES_TO_STRING; diff --git a/dd-java-agent/instrumentation/kafka/kafka-clients-3.8/src/main/java17/datadog/trace/instrumentation/kafka_clients38/TextMapExtractAdapter.java b/dd-java-agent/instrumentation/kafka/kafka-clients-3.8/src/main/java17/datadog/trace/instrumentation/kafka_clients38/TextMapExtractAdapter.java index 30b35216a4b..dd9e5683367 100644 --- a/dd-java-agent/instrumentation/kafka/kafka-clients-3.8/src/main/java17/datadog/trace/instrumentation/kafka_clients38/TextMapExtractAdapter.java +++ b/dd-java-agent/instrumentation/kafka/kafka-clients-3.8/src/main/java17/datadog/trace/instrumentation/kafka_clients38/TextMapExtractAdapter.java @@ -5,6 +5,7 @@ import static datadog.trace.api.telemetry.LogCollector.EXCLUDE_TELEMETRY; import datadog.trace.api.Config; +import datadog.trace.api.Functions; import datadog.trace.bootstrap.instrumentation.api.AgentPropagation; import datadog.trace.bootstrap.instrumentation.api.AgentPropagation.ContextVisitor; import java.nio.ByteBuffer; @@ -19,14 +20,21 @@ public class TextMapExtractAdapter implements ContextVisitor { private static final Logger log = LoggerFactory.getLogger(TextMapExtractAdapter.class); public static final TextMapExtractAdapter GETTER = - new TextMapExtractAdapter(Config.get().isKafkaClientBase64DecodingEnabled()); + new TextMapExtractAdapter( + Config.get().isKafkaClientBase64DecodingEnabled(), + Config.get().isKafkaClientBase64DecodingGuardEnabled()); private final Function headerValueTransformer; private final Base64.Decoder decoder; public TextMapExtractAdapter(boolean decodeBase64Headers) { + this(decodeBase64Headers, true); + } + + public TextMapExtractAdapter(boolean decodeBase64Headers, boolean guardEnabled) { if (decodeBase64Headers) { - this.headerValueTransformer = BASE64_DECODE; + this.headerValueTransformer = + guardEnabled ? new Functions.GuardedBase64Decode()::tryApply : BASE64_DECODE; this.decoder = Base64.getDecoder(); } else { this.headerValueTransformer = UTF8_BYTES_TO_STRING; diff --git a/dd-trace-api/src/main/java/datadog/trace/api/config/TraceInstrumentationConfig.java b/dd-trace-api/src/main/java/datadog/trace/api/config/TraceInstrumentationConfig.java index 137e9805519..86ccd4b9fdd 100644 --- a/dd-trace-api/src/main/java/datadog/trace/api/config/TraceInstrumentationConfig.java +++ b/dd-trace-api/src/main/java/datadog/trace/api/config/TraceInstrumentationConfig.java @@ -113,6 +113,8 @@ public final class TraceInstrumentationConfig { "kafka.client.propagation.disabled.topics"; public static final String KAFKA_CLIENT_BASE64_DECODING_ENABLED = "kafka.client.base64.decoding.enabled"; + public static final String KAFKA_CLIENT_BASE64_DECODING_GUARD_ENABLED = + "kafka.client.base64.decoding.guard.enabled"; public static final String JMS_PROPAGATION_DISABLED_TOPICS = "jms.propagation.disabled.topics"; public static final String JMS_PROPAGATION_DISABLED_QUEUES = "jms.propagation.disabled.queues"; diff --git a/internal-api/src/jmh/java/datadog/trace/api/Base64DecodeBenchmark.java b/internal-api/src/jmh/java/datadog/trace/api/Base64DecodeBenchmark.java new file mode 100644 index 00000000000..c7197995615 --- /dev/null +++ b/internal-api/src/jmh/java/datadog/trace/api/Base64DecodeBenchmark.java @@ -0,0 +1,137 @@ +package datadog.trace.api; + +import static java.nio.charset.StandardCharsets.UTF_8; + +import datadog.trace.api.Functions.GuardedBase64Decode; +import java.util.Base64; +import java.util.Random; +import org.openjdk.jmh.annotations.Benchmark; +import org.openjdk.jmh.annotations.Fork; +import org.openjdk.jmh.annotations.Measurement; +import org.openjdk.jmh.annotations.Param; +import org.openjdk.jmh.annotations.Scope; +import org.openjdk.jmh.annotations.Setup; +import org.openjdk.jmh.annotations.State; +import org.openjdk.jmh.annotations.Threads; +import org.openjdk.jmh.annotations.Warmup; + +/** + * {@link GuardedBase64Decode#decodeOrNull}, the exception-free decoder, against {@link + * Base64#getDecoder()}. These are the two costs that set the guard's {@code closeAfter}: the + * exception-free decoder's overhead on valid input (the cost of staying engaged), and the JDK + * decoder's failure on invalid input (the cost of disengaging too early). + * + *
    + *
  • {@code size}: {@code header} is a trace-header-sized value, {@code long} is about 1.4 KB of + * Base64, where the JDK's block decoding, and any intrinsic for it, has room to pay off. + *
  • {@code depth}: the stack the JDK decoder's exception fills in; a consumer thread's stack is + * deeper than a benchmark thread's. + *
+ * + *

Run with {@code ./gradlew :internal-api:jmh -Pjmh.includes=Base64DecodeBenchmark + * -Pjmh.profilers=gc}, and with {@code -PtestJvm=17} for a newer JDK. + * + *

Results, one run each: Zulu 8.0.382 and Zulu 17.0.7 (HotSpot), MacBook M1, single thread, 2 + * forks of 5 one-second iterations, on a laptop with normal background activity. x86 is not + * measured. ns/op is derived from ops/s; B/op is from {@code -prof gc}. + * + *

+ * ns/op (B/op)                  JDK 8                  JDK 17
+ *                          header      long       header      long
+ * jdkValid,  depth 0       87.7 (160)  3631       53.4 (104)  1598
+ * exceptionFreeValid       86.3 (160)  3904       80.3 (104)  3072
+ * jdkInvalid, depth 0       889 (928)   926        921 (1024)  956
+ * jdkInvalid, depth 50     2080 (1904) 2189       2498 (2384) 2512
+ * exceptionFreeInvalid      6.4 (0)     6.4        6.2 (0)     6.3
+ * 
+ * + * On JDK 8 the exception-free decoder is about as fast as the JDK's on header-sized input, and + * about 7% slower on long input. On JDK 17 the JDK's decoder, which decodes in blocks and has an + * intrinsic on this platform, is about 27 ns faster on header-sized input and about twice as fast + * on long input; that difference is why the guard switches between the two rather than always using + * the exception-free one. Invalid input costs the exception-free decoder about 6 ns and no + * allocation, at any length, where the JDK's costs about 0.9 us, or 2.1 to 2.5 us at depth 50. + * + *

The exception-free decoder allocates its output only once the first unit is clean. Against a + * copy that allocated up front, in the same run, that cost about 7 ns (JDK 17) to 12 ns (JDK 8) on + * valid header-sized input at depth 0, and nothing measurable at depth 50 or on long input, while + * saving 2 to 4 ns and 40 B on invalid header-sized input and about 55 ns and 1 KB on invalid long + * input. It runs only while the guard is engaged. The {@code jdk*} rows are from an earlier run of + * the same day, the {@code exceptionFree*} rows from the later one. + */ +@Fork(2) +@Warmup(iterations = 3, time = 1) +@Measurement(iterations = 5, time = 1) +@Threads(1) +@State(Scope.Benchmark) +public class Base64DecodeBenchmark { + + @Param({"header", "long"}) + String size; + + @Param({"0", "50"}) + int depth; + + byte[] valid; + byte[] invalid; + + @Setup + public void setup() { + byte[] data; + if ("header".equals(size)) { + data = "1234567890123456789".getBytes(UTF_8); + } else { + data = new byte[1024]; + new Random(12672).nextBytes(data); + } + valid = Base64.getEncoder().encode(data); + // a plain-text value where Base64 was expected, bad from its fourth byte + invalid = valid.clone(); + invalid[3] = '-'; + String expected = new String(data, UTF_8); + if (!expected.equals(GuardedBase64Decode.decodeOrNull(valid)) + || !expected.equals(jdk(valid)) + || GuardedBase64Decode.decodeOrNull(invalid) != null + || jdk(invalid) != null) { + throw new IllegalStateException("the two decoders must agree on the benchmark inputs"); + } + } + + static String jdk(byte[] src) { + try { + return new String(Base64.getDecoder().decode(src), UTF_8); + } catch (IllegalArgumentException e) { + return null; + } + } + + @Benchmark + public Object jdkValid() { + return jdkAt(depth, valid); + } + + @Benchmark + public Object exceptionFreeValid() { + return exceptionFreeAt(depth, valid); + } + + @Benchmark + public Object jdkInvalid() { + return jdkAt(depth, invalid); + } + + @Benchmark + public Object exceptionFreeInvalid() { + return exceptionFreeAt(depth, invalid); + } + + private static Object jdkAt(int remaining, byte[] src) { + return remaining > 0 ? jdkAt(remaining - 1, src) : jdk(src); + } + + private static Object exceptionFreeAt(int remaining, byte[] src) { + return remaining > 0 + ? exceptionFreeAt(remaining - 1, src) + : GuardedBase64Decode.decodeOrNull(src); + } +} diff --git a/internal-api/src/jmh/java/datadog/trace/api/FunctionsBase64Benchmark.java b/internal-api/src/jmh/java/datadog/trace/api/FunctionsBase64Benchmark.java new file mode 100644 index 00000000000..c9e1ec14b9f --- /dev/null +++ b/internal-api/src/jmh/java/datadog/trace/api/FunctionsBase64Benchmark.java @@ -0,0 +1,315 @@ +package datadog.trace.api; + +import static datadog.trace.api.Functions.BASE64_DECODE; + +import datadog.trace.api.Functions.GuardedBase64Decode; +import datadog.trace.util.AdaptiveLatch; +import java.nio.charset.StandardCharsets; +import java.util.Base64; +import java.util.function.Function; +import org.openjdk.jmh.annotations.Benchmark; +import org.openjdk.jmh.annotations.Fork; +import org.openjdk.jmh.annotations.Measurement; +import org.openjdk.jmh.annotations.Param; +import org.openjdk.jmh.annotations.Scope; +import org.openjdk.jmh.annotations.Setup; +import org.openjdk.jmh.annotations.State; +import org.openjdk.jmh.annotations.Warmup; +import org.openjdk.jmh.infra.Blackhole; + +/** + * Compares {@link Functions#BASE64_DECODE} on valid input (no exception) against malformed input, + * where every call throws and catches an {@link IllegalArgumentException}. Motivated by a customer + * seeing ~180K/day of this exact throw from Kafka header extraction (non-Base64 header values from + * a mixed producer). Point is to see whether HotSpot's fast-throw stack-trace omission actually + * kicks in for this call site under sustained repeated throws, or whether the caught path pays for + * a full stack trace fill-in every time. + * + *

It also measures the reusable form of the guard, {@link AdaptiveLatch} (APMLP-1884): an + * abstract class whose subclass is the strategy, held in a {@code static final}. The arms + * use the production subclass, {@link Functions.GuardedBase64Decode}, whose cautious path is an + * exception-free decoder; {@code Base64DecodeBenchmark} compares that decoder with the JDK's + * directly. {@link AdaptiveLatch#tryApply} converts a failure to {@code null} and never builds an + * exception at all while engaged. + * + *

The numbers below were measured on an earlier, benchmark-local sketch of the latch, then named + * {@code DynamicLatch}, with {@code tryGetOrNull} for {@code tryApply}, while its engaged path was + * an alphabet pre-check in front of the JDK decoder, before {@code fallback} was added and before + * the pre-check stopped rejecting unpadded input. The sketch also had a flow-through flavor, {@code + * get}, which threw a stackless stand-in while engaged; it was dropped when the latch moved into + * production because no caller needs it, and its arms were removed. The hand-written throwing + * {@link Breaker} still measures that shape. Not re-measured since. + * + *

The hand-written {@link Breaker} is the specialized baseline the latch is compared against. + * + *

Results, one run: Zulu 17.0.7 (HotSpot), MacBook M1, single thread, 5 forks, on a laptop with + * normal background activity (load about 4 to 7). JDK 8 and x86 are not measured. Error margins are + * below 4%. + * + *

+ * ns/op                                valid input  invalid input
+ * unguarded (status quo)                      30.5          905.7
+ * always pre-check                            71.3           2.75
+ * hand-written Breaker, converting            31.2           2.73
+ * hand-written Breaker, throwing              31.4           11.6
+ * sketch, converting (tryApply)               31.1           2.78
+ * sketch, flow-through (get)                  31.2           11.5
+ *
+ * Mix, ns/op, one invalid input in every N
+ *      N  unguarded  pre-check  Breaker  throwing  latch.get      tryApply
+ *      2      467.1       36.6     47.7      49.9       51.9          47.3
+ *     10      118.4       64.9     71.0      71.3       72.1          70.0
+ *    100       41.9       70.0     46.1      46.4       47.9          45.9
+ *   1000       41.6       70.3     44.7      45.7       45.1          45.6
+ *  10000       34.3       70.2     36.2      35.3       39.6          35.5
+ * 
+ * + * The sketch matches the hand-written breaker: within 0.2 ns in the single-input arms, and within + * about 4 ns in the mixed ones (the largest gap is flow-through at one in 10,000, 39.6 ns against + * 35.3 ns). While engaged, the converting flavor costs about 2.8 ns where the status quo costs + * about 906 ns; the flow-through flavor costs about 11.5 ns, the price of building a stackless + * exception. Always pre-checking more than doubles the cost of valid input (71 ns against 30 ns), + * which is what the adaptive form avoids. + * + *

It is a tradeoff, not a free win. With one invalid input in 100 or rarer, the guard costs 1 to + * 4 ns over doing nothing and is about 24 to 34 ns cheaper than always pre-checking. With a high + * rate (one in 2 or one in 10), always pre-checking is faster than the adaptive guard (36.6 ns + * against about 47 to 52 ns at one in 2, and 65 ns against about 70 to 72 ns at one in 10). The + * cause was not investigated. + */ +@Fork(2) +@Warmup(iterations = 3, time = 1) +@Measurement(iterations = 5, time = 1) +@State(Scope.Benchmark) +public class FunctionsBase64Benchmark { + + static final byte[] VALID = + Base64.getEncoder().encode("x-datadog-trace-id=1234567890".getBytes(StandardCharsets.UTF_8)); + + static final byte[] INVALID = "not-valid-base64!@#".getBytes(StandardCharsets.UTF_8); + + private static final boolean[] BASE64_ALPHABET = buildAlphabetTable(); + + private static boolean[] buildAlphabetTable() { + boolean[] table = new boolean[256]; + String alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/="; + for (int i = 0; i < alphabet.length(); i++) { + table[alphabet.charAt(i)] = true; + } + return table; + } + + // Branch-free per byte on purpose: a data-dependent early exit helps the rare "bad byte early" + // case but hurts the common valid case, which always scans the whole buffer anyway. + private static boolean looksLikeBase64(byte[] bytes) { + if (bytes.length == 0 || (bytes.length & 3) != 0) { + return false; + } + boolean valid = true; + for (byte b : bytes) { + valid &= BASE64_ALPHABET[b & 0xFF]; + } + return valid; + } + + static final Function PRECHECK_BASE64_DECODE = + bytes -> looksLikeBase64(bytes) ? BASE64_DECODE.apply(bytes) : null; + + private static final int CLOSE_THRESHOLD = 20; + + /** + * Thrown only when the cheap alphabet-scan precheck already knows the input can't be valid + * Base64, so we never call into {@link Base64}'s decoder (or pay for its exception's stack-trace + * capture) at all. Deliberately extends {@link IllegalArgumentException} so it's a drop-in for + * existing catch sites around the real decoder; callers should not rely on a stack trace being + * present for this specific instance — a real decode failure still throws the JDK's own + * exception, unmodified, with its own message and stack trace. + * + *

A new instance is built per failure, never shared, so suppression cannot accumulate on it. A + * shared instance would also have to disable suppression and forbid {@code initCause}. + */ + static final class FastFailBase64Exception extends IllegalArgumentException { + private static final String MESSAGE = "Header value is not valid Base64"; + + FastFailBase64Exception() { + super(MESSAGE); + } + + // IllegalArgumentException has no writableStackTrace-suppressing constructor of its own, so + // skip the stack walk here instead. + @Override + public synchronized Throwable fillInStackTrace() { + return this; + } + } + + // Single-word countdown: 0 == closed (no precheck), >0 == guarded, counting down to close. + // Plain int on purpose: this is advisory hysteresis, not correctness-critical state, so a lost + // update or a stale read across threads just means one extra precheck or one extra exception. + static final class Breaker { + int state; + + String decode(byte[] bytes) { + if (state > 0 && !looksLikeBase64(bytes)) { + return null; + } + String result = BASE64_DECODE.apply(bytes); + if (result == null) { + state = CLOSE_THRESHOLD; + } else if (state > 0) { + state--; + } + return result; + } + + // Same hysteresis, but preserves throw-based failure semantics: a precheck-known failure + // throws our stack-trace-free stand-in, a real decode failure lets the JDK's own + // IllegalArgumentException (with its own message and stack trace) propagate untouched. + String decodeOrThrow(byte[] bytes) { + if (state > 0 && !looksLikeBase64(bytes)) { + throw new FastFailBase64Exception(); + } + try { + String result = new String(Base64.getDecoder().decode(bytes), StandardCharsets.UTF_8); + if (state > 0) { + state--; + } + return result; + } catch (IllegalArgumentException e) { + state = CLOSE_THRESHOLD; + throw e; + } + } + } + + final Breaker breakerForValid = new Breaker(); + final Breaker breakerForInvalid = new Breaker(); + final Breaker breakerThrowingForValid = new Breaker(); + final Breaker breakerThrowingForInvalid = new Breaker(); + + // One latch per arm, each a static final of the exact type, as a call site would hold it. + static final GuardedBase64Decode LATCH_CONVERT_VALID = new GuardedBase64Decode(); + static final GuardedBase64Decode LATCH_CONVERT_INVALID = new GuardedBase64Decode(); + static final GuardedBase64Decode LATCH_CONVERT_MIX = new GuardedBase64Decode(); + + /** Fails fast if the latch does not behave as the benchmark assumes. */ + @Setup + public void checkTheLatchBehaves() { + GuardedBase64Decode latch = new GuardedBase64Decode(); + String expected = new String(Base64.getDecoder().decode(VALID), StandardCharsets.UTF_8); + if (!expected.equals(latch.tryApply(VALID))) { + throw new IllegalStateException("a valid value must decode"); + } + if (latch.tryApply(INVALID) != null || !latch.isEngaged()) { + throw new IllegalStateException("bad input must convert to null and engage the latch"); + } + if (!expected.equals(latch.tryApply(VALID))) { + throw new IllegalStateException("good input must still decode while engaged"); + } + } + + @Benchmark + public void breakerValid(Blackhole bh) { + bh.consume(breakerForValid.decode(VALID)); + } + + @Benchmark + public void breakerInvalid(Blackhole bh) { + bh.consume(breakerForInvalid.decode(INVALID)); + } + + @Benchmark + public void breakerThrowingValid(Blackhole bh) { + bh.consume(breakerThrowingForValid.decodeOrThrow(VALID)); + } + + @Benchmark + public void breakerThrowingInvalid(Blackhole bh) { + try { + bh.consume(breakerThrowingForInvalid.decodeOrThrow(INVALID)); + } catch (IllegalArgumentException e) { + bh.consume(e); + } + } + + @Benchmark + public void latchConvertValid(Blackhole bh) { + bh.consume(LATCH_CONVERT_VALID.tryApply(VALID)); + } + + @Benchmark + public void latchConvertInvalid(Blackhole bh) { + bh.consume(LATCH_CONVERT_INVALID.tryApply(INVALID)); + } + + @Benchmark + public void valid(Blackhole bh) { + bh.consume(BASE64_DECODE.apply(VALID)); + } + + @Benchmark + public void invalid(Blackhole bh) { + bh.consume(BASE64_DECODE.apply(INVALID)); + } + + @Benchmark + public void precheckValid(Blackhole bh) { + bh.consume(PRECHECK_BASE64_DECODE.apply(VALID)); + } + + @Benchmark + public void precheckInvalid(Blackhole bh) { + bh.consume(PRECHECK_BASE64_DECODE.apply(INVALID)); + } + + // Deterministic stream: one INVALID every invalidEveryN calls, VALID otherwise. Each mix + // benchmark gets its own counter and state so the strategies see the identical pattern. + @State(Scope.Thread) + public static class Mix { + + @Param({"2", "10", "100", "1000", "10000"}) + int invalidEveryN; + + int counter; + final Breaker breaker = new Breaker(); + final Breaker throwingBreaker = new Breaker(); + + byte[] next() { + if (++counter >= invalidEveryN) { + counter = 0; + return INVALID; + } + return VALID; + } + } + + @Benchmark + public void mixUnguarded(Mix mix, Blackhole bh) { + bh.consume(BASE64_DECODE.apply(mix.next())); + } + + @Benchmark + public void mixGuarded(Mix mix, Blackhole bh) { + bh.consume(PRECHECK_BASE64_DECODE.apply(mix.next())); + } + + @Benchmark + public void mixBreaker(Mix mix, Blackhole bh) { + bh.consume(mix.breaker.decode(mix.next())); + } + + @Benchmark + public void mixBreakerThrowing(Mix mix, Blackhole bh) { + byte[] bytes = mix.next(); + try { + bh.consume(mix.throwingBreaker.decodeOrThrow(bytes)); + } catch (IllegalArgumentException e) { + bh.consume(e); + } + } + + @Benchmark + public void mixLatchConvert(Mix mix, Blackhole bh) { + bh.consume(LATCH_CONVERT_MIX.tryApply(mix.next())); + } +} diff --git a/internal-api/src/jmh/java/datadog/trace/util/AdaptiveLatchBenchmark.java b/internal-api/src/jmh/java/datadog/trace/util/AdaptiveLatchBenchmark.java new file mode 100644 index 00000000000..66d0c9e4e24 --- /dev/null +++ b/internal-api/src/jmh/java/datadog/trace/util/AdaptiveLatchBenchmark.java @@ -0,0 +1,279 @@ +package datadog.trace.util; + +import org.openjdk.jmh.annotations.Benchmark; +import org.openjdk.jmh.annotations.Fork; +import org.openjdk.jmh.annotations.Measurement; +import org.openjdk.jmh.annotations.Param; +import org.openjdk.jmh.annotations.Scope; +import org.openjdk.jmh.annotations.Setup; +import org.openjdk.jmh.annotations.State; +import org.openjdk.jmh.annotations.Threads; +import org.openjdk.jmh.annotations.Warmup; + +/** + * What {@link AdaptiveLatch#tryApply} costs and saves for an operation whose failure depends on the + * input. The operation is {@link Integer#parseInt}, whose {@link NumberFormatException} fills in a + * full stack trace; the pre-check is a digit scan, which is correct (it flags only input {@code + * parseInt} rejects) but not complete (it lets an overflowing number through). + * + *

    + *
  • {@code unguarded*}: the status quo -- parse, and catch the exception for bad input. + *
  • {@code precheck*}: the non-adaptive alternative -- always scan before parsing. + *
  • {@code latchGood}: a disengaged latch on good input. The difference from {@code + * unguardedGood} is the latch's overhead on the path that works. + *
  • {@code latchEngagedGood}: good input while engaged, so it pays for the pre-check too. The + * latch is pinned engaged ({@code closeAfter} is {@link Integer#MAX_VALUE}) so that the state + * holds for the whole run. + *
  • {@code latchBad}: bad input while engaged, the steady state for a producer that keeps + * sending garbage: turned away without parsing. + *
  • {@code mix*}: one bad input in every {@code invalidEveryN}, with the default {@code + * closeAfter}, so the latch engages, disengages and re-engages as it would in production. + *
+ * + *

The cost of a throw grows with the depth of the stack it fills in, which is why {@code depth} + * is a parameter: a benchmark thread's stack is shallow, a request thread's is not. Each arm + * descends on its own, and each has its own {@code static final} latch, so that no arm's profile is + * shaped by another's. + * + *

Run with {@code ./gradlew :internal-api:jmh -Pjmh.includes=AdaptiveLatchBenchmark + * -Pjmh.profilers=gc}. + * + *

The latch here takes the pre-check form of a cautious path: the digit scan, then {@code + * parseInt} only if it passes. The results were recorded when that was the latch's only form; it + * does the same work now. + * + *

Results, one run: Zulu 17.0.7 (HotSpot), MacBook M1, single thread, 2 forks of 5 one-second + * iterations, on a laptop with normal background activity. JDK 8 and x86 are not measured. ns/op is + * derived from ops/s; B/op is from {@code -prof gc}. Good input allocates 16 B for the boxed result + * in every arm. + * + *

+ * ns/op (B/op)          depth 0          depth 50
+ * unguardedGood         12.7  (16)        31.7  (16)
+ * precheckGood          15.8  (16)        40.1  (16)
+ * latchGood             15.7  (16)        36.2  (16)
+ * latchEngagedGood      11.3  (16)        40.3  (16)
+ * unguardedBad         908    (880)     2403    (2240)
+ * precheckBad            2.79  (0)        25.6   (0)
+ * latchBad               2.57  (0)        31.2   (0)
+ *
+ * Mix, ns/op, one bad input in every N
+ *           depth 0                         depth 50
+ *      N    unguarded  precheck  latch      unguarded  precheck  latch
+ *      2        452       9.85    31.9         1238      33.7     98.7
+ *    100       18.1      14.1     19.7         55.3      37.0     62.5
+ *  10000       13.5      13.8     11.9         31.7      38.8     38.6
+ * 
+ * + * Engaged, bad input costs about 2.6 ns where the status quo costs about 908 ns (31 ns against 2.4 + * us at depth 50), and allocates nothing where the status quo allocates 880 B (2,240 B). + * + *

Disengaged, the latch is not free on good input: about 3 ns over {@code unguardedGood} at + * depth 0 and about 4.5 ns at depth 50, more than one field read should cost. The depth-0 + * good-input arms have errors of 9 to 11%, which is also why {@code latchEngagedGood} appears + * faster than {@code latchGood} there; the depth-50 gap is outside the error (2 to 3%). The cause + * was not investigated. + * + *

The mix shows when the latch pays off. At one bad input in 2 it is about 14 times cheaper than + * the status quo, though always pre-checking is cheaper still. At one in 100 it saves nothing and + * costs a little: with the default {@code closeAfter} of 20, it disengages between bad inputs, so + * every bad input still throws. It only helps while bad input arrives more often than once per + * {@code closeAfter} calls; at one in 10,000 the three are within a few ns of each other. + */ +@Fork(2) +@Warmup(iterations = 3, time = 1) +@Measurement(iterations = 5, time = 1) +@Threads(1) +@State(Scope.Benchmark) +public class AdaptiveLatchBenchmark { + + static final String GOOD = "1234567"; + static final String BAD = "12x4567"; + + @Param({"0", "50"}) + int depth; + + /** Parses an int; while engaged, a digit scan turns away anything that is not plain digits. */ + static class Parse extends AdaptiveLatch { + Parse() { + // the value the recorded results were measured with; not tuned for this operation + this(20); + } + + Parse(int closeAfter) { + super(NumberFormatException.class, closeAfter); + } + + @Override + protected Integer apply(String input) { + return Integer.parseInt(input); + } + + @Override + protected Integer applySafely(String input) { + return isAllDigits(input) ? Integer.parseInt(input) : reject(input); + } + } + + /** Never disengages, so that the engaged state holds for a whole run of good input. */ + static final class PinnedParse extends Parse { + PinnedParse() { + super(Integer.MAX_VALUE); + } + } + + static boolean isAllDigits(String input) { + if (input.isEmpty()) { + return false; + } + // branch-free per char, as a pre-check over good input always scans it all anyway + boolean digits = true; + for (int i = 0; i < input.length(); i++) { + char c = input.charAt(i); + digits &= c >= '0' & c <= '9'; + } + return digits; + } + + static Integer unguarded(String input) { + try { + return Integer.parseInt(input); + } catch (NumberFormatException e) { + return null; + } + } + + static Integer precheck(String input) { + return isAllDigits(input) ? Integer.parseInt(input) : null; + } + + static final Parse LATCH_GOOD = new Parse(); + static final PinnedParse LATCH_ENGAGED_GOOD = new PinnedParse(); + static final Parse LATCH_BAD = new Parse(); + static final Parse LATCH_MIX = new Parse(); + + /** Puts each latch in the state its arm measures, and fails fast if one does not get there. */ + @Setup + public void setup() { + if (LATCH_GOOD.tryApply(GOOD) == null || LATCH_GOOD.isEngaged()) { + throw new IllegalStateException("good input must parse without engaging"); + } + LATCH_ENGAGED_GOOD.tryApply(BAD); + if (!LATCH_ENGAGED_GOOD.isEngaged() || LATCH_ENGAGED_GOOD.tryApply(GOOD) == null) { + throw new IllegalStateException("good input must still parse while engaged"); + } + LATCH_BAD.tryApply(BAD); + if (!LATCH_BAD.isEngaged() || LATCH_BAD.tryApply(BAD) != null) { + throw new IllegalStateException("bad input must engage the latch and be turned away"); + } + } + + @Benchmark + public Object unguardedGood() { + return unguardedGood(depth); + } + + @Benchmark + public Object precheckGood() { + return precheckGood(depth); + } + + @Benchmark + public Object latchGood() { + return latchGood(depth); + } + + @Benchmark + public Object latchEngagedGood() { + return latchEngagedGood(depth); + } + + @Benchmark + public Object unguardedBad() { + return unguardedBad(depth); + } + + @Benchmark + public Object precheckBad() { + return precheckBad(depth); + } + + @Benchmark + public Object latchBad() { + return latchBad(depth); + } + + /** + * A deterministic stream: one {@link #BAD} in every {@code invalidEveryN}, else {@link #GOOD}. + */ + @State(Scope.Thread) + public static class Mix { + @Param({"2", "100", "10000"}) + int invalidEveryN; + + int counter; + + String next() { + if (++counter >= invalidEveryN) { + counter = 0; + return BAD; + } + return GOOD; + } + } + + @Benchmark + public Object mixUnguarded(Mix mix) { + return mixUnguarded(depth, mix.next()); + } + + @Benchmark + public Object mixPrecheck(Mix mix) { + return mixPrecheck(depth, mix.next()); + } + + @Benchmark + public Object mixLatch(Mix mix) { + return mixLatch(depth, mix.next()); + } + + private Object unguardedGood(int remaining) { + return remaining > 0 ? unguardedGood(remaining - 1) : unguarded(GOOD); + } + + private Object precheckGood(int remaining) { + return remaining > 0 ? precheckGood(remaining - 1) : precheck(GOOD); + } + + private Object latchGood(int remaining) { + return remaining > 0 ? latchGood(remaining - 1) : LATCH_GOOD.tryApply(GOOD); + } + + private Object latchEngagedGood(int remaining) { + return remaining > 0 ? latchEngagedGood(remaining - 1) : LATCH_ENGAGED_GOOD.tryApply(GOOD); + } + + private Object unguardedBad(int remaining) { + return remaining > 0 ? unguardedBad(remaining - 1) : unguarded(BAD); + } + + private Object precheckBad(int remaining) { + return remaining > 0 ? precheckBad(remaining - 1) : precheck(BAD); + } + + private Object latchBad(int remaining) { + return remaining > 0 ? latchBad(remaining - 1) : LATCH_BAD.tryApply(BAD); + } + + private Object mixUnguarded(int remaining, String input) { + return remaining > 0 ? mixUnguarded(remaining - 1, input) : unguarded(input); + } + + private Object mixPrecheck(int remaining, String input) { + return remaining > 0 ? mixPrecheck(remaining - 1, input) : precheck(input); + } + + private Object mixLatch(int remaining, String input) { + return remaining > 0 ? mixLatch(remaining - 1, input) : LATCH_MIX.tryApply(input); + } +} diff --git a/internal-api/src/main/java/datadog/trace/api/Config.java b/internal-api/src/main/java/datadog/trace/api/Config.java index b1707b77a71..5194934ba10 100644 --- a/internal-api/src/main/java/datadog/trace/api/Config.java +++ b/internal-api/src/main/java/datadog/trace/api/Config.java @@ -623,6 +623,7 @@ import static datadog.trace.api.config.TraceInstrumentationConfig.JMS_PROPAGATION_DISABLED_TOPICS; import static datadog.trace.api.config.TraceInstrumentationConfig.JMS_UNACKNOWLEDGED_MAX_AGE; import static datadog.trace.api.config.TraceInstrumentationConfig.KAFKA_CLIENT_BASE64_DECODING_ENABLED; +import static datadog.trace.api.config.TraceInstrumentationConfig.KAFKA_CLIENT_BASE64_DECODING_GUARD_ENABLED; import static datadog.trace.api.config.TraceInstrumentationConfig.KAFKA_CLIENT_PROPAGATION_DISABLED_TOPICS; import static datadog.trace.api.config.TraceInstrumentationConfig.LOGS_INJECTION; import static datadog.trace.api.config.TraceInstrumentationConfig.LOGS_INJECTION_ENABLED; @@ -1321,6 +1322,7 @@ public static String getHostName() { private final boolean kafkaClientPropagationEnabled; private final Set kafkaClientPropagationDisabledTopics; private final boolean kafkaClientBase64DecodingEnabled; + private final boolean kafkaClientBase64DecodingGuardEnabled; private final boolean jmsPropagationEnabled; private final Set jmsPropagationDisabledTopics; @@ -3156,6 +3158,8 @@ PROFILING_DATADOG_PROFILER_ENABLED, isDatadogProfilerSafeInCurrentEnvironment()) tryMakeImmutableSet(configProvider.getList(KAFKA_CLIENT_PROPAGATION_DISABLED_TOPICS)); kafkaClientBase64DecodingEnabled = configProvider.getBoolean(KAFKA_CLIENT_BASE64_DECODING_ENABLED, false); + kafkaClientBase64DecodingGuardEnabled = + configProvider.getBoolean(KAFKA_CLIENT_BASE64_DECODING_GUARD_ENABLED, true); jmsPropagationEnabled = isPropagationEnabled(true, "jms"); jmsPropagationDisabledTopics = tryMakeImmutableSet(configProvider.getList(JMS_PROPAGATION_DISABLED_TOPICS)); @@ -5180,6 +5184,10 @@ public boolean isKafkaClientBase64DecodingEnabled() { return kafkaClientBase64DecodingEnabled; } + public boolean isKafkaClientBase64DecodingGuardEnabled() { + return kafkaClientBase64DecodingGuardEnabled; + } + public boolean isRabbitPropagationEnabled() { return rabbitPropagationEnabled; } @@ -6983,6 +6991,8 @@ public String toString() { + kafkaClientPropagationDisabledTopics + ", kafkaClientBase64DecodingEnabled=" + kafkaClientBase64DecodingEnabled + + ", kafkaClientBase64DecodingGuardEnabled=" + + kafkaClientBase64DecodingGuardEnabled + ", jmsPropagationEnabled=" + jmsPropagationEnabled + ", jmsPropagationDisabledTopics=" diff --git a/internal-api/src/main/java/datadog/trace/api/Functions.java b/internal-api/src/main/java/datadog/trace/api/Functions.java index 731cd711a9a..bd9aee0b6f3 100644 --- a/internal-api/src/main/java/datadog/trace/api/Functions.java +++ b/internal-api/src/main/java/datadog/trace/api/Functions.java @@ -4,13 +4,16 @@ import static java.util.function.Function.identity; import datadog.trace.bootstrap.instrumentation.api.UTF8BytesString; +import datadog.trace.util.AdaptiveLatch; import java.lang.invoke.MethodHandle; import java.lang.invoke.MethodHandles; import java.lang.invoke.MethodType; +import java.util.Arrays; import java.util.Base64; import java.util.Locale; import java.util.function.BiFunction; import java.util.function.Function; +import javax.annotation.Nullable; public final class Functions { @@ -188,4 +191,131 @@ public T apply(Object input) { return null; } }; + + /** + * Base64-decodes bytes that are occasionally not Base64 at all, for example header values from a + * misconfigured producer that mixes encodings, without paying for {@link Base64}'s decoder to + * throw, and fill in a stack trace, on every one of them. Returns {@code null} for input that + * does not decode, like {@link #BASE64_DECODE}. + * + *

Decodes with {@link Base64#getDecoder()} until it fails, then with {@link + * #decodeOrNull(byte[])}, an exception-free decoder that accepts exactly what {@link + * Base64#getDecoder()} accepts, until {@code CLOSE_AFTER} consecutive inputs decode again (see + * {@link AdaptiveLatch}). + * + *

One instance can be shared across threads: its state is advisory, so a stale read costs one + * extra cautious decode or one extra real failure, never a wrong result. + */ + public static final class GuardedBase64Decode + extends AdaptiveLatch { + /** Each byte's 6-bit value in the basic Base64 alphabet, {@code -2} for padding, else -1. */ + private static final byte[] FROM_BASE64 = buildDecodeTable(); + + private static final int PADDING = -2; + + private static byte[] buildDecodeTable() { + byte[] table = new byte[256]; + Arrays.fill(table, (byte) -1); + String alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/"; + for (int i = 0; i < alphabet.length(); i++) { + table[alphabet.charAt(i)] = (byte) i; + } + table['='] = PADDING; + return table; + } + + /** + * The rent-or-buy break-even (see {@link AdaptiveLatch}), from {@code Base64DecodeBenchmark} on + * JDK 17 for header-sized values: a JDK decoder failure costs about 913 ns more than a + * rejection here at stack depth 0 and about 2,468 ns at depth 50, and on valid input this + * decoder costs about 27 ns more than the JDK's, giving 34 to 91. A consumer's stack is deeper + * than a benchmark's. On JDK 8 the two decoders cost about the same on valid input, so the + * value barely matters there. + */ + private static final int CLOSE_AFTER = 64; + + public GuardedBase64Decode() { + super(IllegalArgumentException.class, CLOSE_AFTER); + } + + @Override + protected String apply(byte[] bytes) { + return new String(Base64.getDecoder().decode(bytes), UTF_8); + } + + @Override + protected String applySafely(byte[] bytes) { + String decoded = decodeOrNull(bytes); + return decoded != null ? decoded : reject(bytes); + } + + /** + * Decodes {@code src} as {@link Base64#getDecoder()} would, returning {@code null} where it + * would throw. Follows the same rules: padding is optional, but if present must complete the + * final unit, and nothing may follow it; a final unit of one character is rejected. + */ + @Nullable + static String decodeOrNull(byte[] src) { + final int length = src.length; + if (length == 0) { + return ""; + } + // allocated up front only if the first unit is clean, so input that is bad from the start + // costs no allocation; input that goes bad later still does, since one pass cannot know in + // advance. If the first unit is not clean, no unit can complete before the loop meets its + // bad or padding byte, so the loop never writes to a null dst. + byte[] dst = + length >= 4 + && (FROM_BASE64[src[0] & 0xFF] + | FROM_BASE64[src[1] & 0xFF] + | FROM_BASE64[src[2] & 0xFF] + | FROM_BASE64[src[3] & 0xFF]) + >= 0 + ? new byte[3 * ((length + 3) / 4)] + : null; + int dp = 0; + int bits = 0; + // the bit position of the next character within its 4-character unit + int shift = 18; + int sp = 0; + while (sp < length) { + final int b = FROM_BASE64[src[sp++] & 0xFF]; + if (b < 0) { + // padding is only legal after two or three characters of a unit, and after two it + // must be doubled + if (b != PADDING || shift == 18 || (shift == 6 && (sp == length || src[sp++] != '='))) { + return null; + } + break; + } + bits |= b << shift; + shift -= 6; + if (shift < 0) { + dst[dp++] = (byte) (bits >> 16); + dst[dp++] = (byte) (bits >> 8); + dst[dp++] = (byte) bits; + shift = 18; + bits = 0; + } + } + // a dangling single character, or anything after the padding + if (shift == 12 || sp < length) { + return null; + } + if (shift != 18) { + if (dst == null) { + // no full unit: at most two bytes + dst = new byte[2]; + } + dst[dp++] = (byte) (bits >> 16); + if (shift == 0) { + dst[dp++] = (byte) (bits >> 8); + } + } + if (dst == null) { + return ""; + } + return new String(dst, 0, dp, UTF_8); + } + } } diff --git a/internal-api/src/main/java/datadog/trace/util/AdaptiveLatch.java b/internal-api/src/main/java/datadog/trace/util/AdaptiveLatch.java new file mode 100644 index 00000000000..96e5408d957 --- /dev/null +++ b/internal-api/src/main/java/datadog/trace/util/AdaptiveLatch.java @@ -0,0 +1,124 @@ +package datadog.trace.util; + +import javax.annotation.Nullable; +import javax.annotation.concurrent.ThreadSafe; + +/** + * A self-resetting switch between two ways of performing an operation whose failure depends on the + * input, such as parsing a value from a producer that sometimes sends garbage: an optimistic path, + * {@link #apply}, that is fastest on good input but throws {@code X} on bad input, and a cautious + * path, {@link #applySafely}, that never throws {@code X}. Unlike {@code Latch} and {@link + * ClassLatch}, it never skips the operation: a bad input says nothing about the next one. + * + *

Disengaged, every call takes the optimistic path. A failure there engages the latch, and from + * then on calls take the cautious path until {@code closeAfter} consecutive inputs pass it cleanly. + * The cautious path reports bad input by returning {@link #reject}, which restarts that count; + * anything else it returns counts as clean. + * + *

The cautious path can be built in several ways: + * + *

    + *
  • an exception-free implementation of the same operation; + *
  • a cheap, correct pre-check in front of the optimistic path: {@code return isKnownBad(input) + * ? reject(input) : apply(input);} + *
  • a repair that turns common bad input into good input before the optimistic path. + *
+ * + *

Choosing {@code closeAfter} is a rent-or-buy decision. Staying engaged costs the cautious + * path's overhead on every good input; disengaging costs one optimistic failure when the next bad + * input arrives. Disengaging once the accumulated overhead would match one failure, that is after + * about {@code failureCost / cautiousOverhead} good inputs, is never worse than twice the best + * possible schedule, whatever the input. Measure both costs at a realistic stack depth: the cost of + * a throw grows with the stack it fills in. + * + *

Intended as a {@code static final} field of a named final subclass, one per call site: the + * receiver is then a constant of a known exact type, so the JIT can inline the hooks. Disengaged, + * the latch adds one plain field read to the optimistic path. + * + *

This is a hint, not a lock. The state is a plain counter, and it counts calls, not time. A + * stale or lost update only costs one more cautious call or one more optimistic failure, never a + * wrong result. + * + * @param the type of the value the operation is applied to + * @param the type of the result + * @param the failure that engages the latch; anything else propagates unchanged + */ +@ThreadSafe +public abstract class AdaptiveLatch { + private final Class failureType; + private final int closeAfter; + + /** 0 when disengaged; otherwise the clean calls left before disengaging. */ + private int remaining; + + /** + * @param failureType the failure of {@link #apply} that engages the latch + * @param closeAfter how many consecutive clean calls disengage it, at least 1; see the class + * comment for how to choose it + */ + protected AdaptiveLatch(Class failureType, int closeAfter) { + if (closeAfter < 1) { + throw new IllegalArgumentException("closeAfter must be at least 1: " + closeAfter); + } + this.failureType = failureType; + this.closeAfter = closeAfter; + } + + /** The optimistic path: fastest on good input, but may throw {@code X} for bad input. */ + @Nullable + protected abstract R apply(T input); + + /** + * The cautious path: must not throw {@code X}. Returns {@link #reject} for bad input; any other + * return counts as clean, including {@code null}. + */ + @Nullable + protected abstract R applySafely(T input); + + /** How many consecutive clean calls disengage the latch. */ + public final int closeAfter() { + return closeAfter; + } + + /** What a rejected input yields. {@code null} unless overridden. */ + @Nullable + protected R fallback(T input) { + return null; + } + + /** Reports bad input from {@link #applySafely}: restarts the count, and returns the fallback. */ + @Nullable + protected final R reject(T input) { + remaining = closeAfter; + return fallback(input); + } + + /** + * Performs the operation: optimistically while disengaged, cautiously while engaged. An + * optimistic failure engages the latch, and the same input is then retried cautiously, so a + * cautious path that can repair it still gets the chance. + */ + @Nullable + public final R tryApply(T input) { + final int remaining = this.remaining; + if (remaining > 0) { + // counted as clean up front; reject() re-arms the count if it is not + this.remaining = remaining - 1; + return applySafely(input); + } + try { + return apply(input); + } catch (RuntimeException e) { + if (!failureType.isInstance(e)) { + throw e; + } + this.remaining = closeAfter; + return applySafely(input); + } + } + + /** Returns whether calls currently take the cautious path. */ + public final boolean isEngaged() { + return remaining > 0; + } +} diff --git a/internal-api/src/test/java/datadog/trace/api/GuardedBase64DecodeTest.java b/internal-api/src/test/java/datadog/trace/api/GuardedBase64DecodeTest.java new file mode 100644 index 00000000000..d2f67b4a74e --- /dev/null +++ b/internal-api/src/test/java/datadog/trace/api/GuardedBase64DecodeTest.java @@ -0,0 +1,176 @@ +package datadog.trace.api; + +import static java.nio.charset.StandardCharsets.UTF_8; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.junit.jupiter.api.Assertions.fail; + +import datadog.trace.api.Functions.GuardedBase64Decode; +import java.util.Base64; +import java.util.Random; +import org.junit.jupiter.api.Test; +import org.tabletest.junit.TableTest; + +class GuardedBase64DecodeTest { + + private static final byte[] VALID = + Base64.getEncoder().encode("x-datadog-trace-id".getBytes(UTF_8)); + private static final byte[] INVALID = "not-valid-base64!@#".getBytes(UTF_8); + + private static byte[] bytes(String s) { + return s.getBytes(UTF_8); + } + + private static GuardedBase64Decode engaged() { + GuardedBase64Decode decode = new GuardedBase64Decode(); + assertNull(decode.tryApply(INVALID)); + assertTrue(decode.isEngaged()); + return decode; + } + + /** What {@link Base64#getDecoder()} makes of {@code src}: the decoded string, or null. */ + private static String jdk(byte[] src) { + try { + return new String(Base64.getDecoder().decode(src), UTF_8); + } catch (IllegalArgumentException e) { + return null; + } + } + + private static void assertAgreesWithTheJdk(byte[] src) { + String expected; + try { + expected = jdk(src); + } catch (RuntimeException e) { + // the guard relies on the JDK decoder failing only with IllegalArgumentException + throw new AssertionError("JDK decoder threw " + e + " for " + new String(src, UTF_8), e); + } + String actual = GuardedBase64Decode.decodeOrNull(src); + if (expected == null ? actual != null : !expected.equals(actual)) { + fail( + "for \"" + + new String(src, UTF_8) + + "\": JDK " + + (expected == null ? "rejects" : "gives \"" + expected + "\"") + + ", decodeOrNull " + + (actual == null ? "rejects" : "gives \"" + actual + "\"")); + } + } + + @Test + void decodesValidInput() { + assertEquals("x-datadog-trace-id", new GuardedBase64Decode().tryApply(VALID)); + } + + @Test + void invalidInputYieldsNullAndEngages() { + GuardedBase64Decode decode = new GuardedBase64Decode(); + + assertNull(decode.tryApply(INVALID)); + assertTrue(decode.isEngaged()); + } + + @Test + void whileEngagedItStillDecodesValidInputAndRejectsInvalidInput() { + GuardedBase64Decode decode = engaged(); + + assertEquals("x-datadog-trace-id", decode.tryApply(VALID)); + assertNull(decode.tryApply(INVALID)); + } + + @Test + void disengagesAfterEnoughValidInput() { + GuardedBase64Decode decode = engaged(); + + for (int i = 0; i < decode.closeAfter(); i++) { + decode.tryApply(VALID); + } + + assertFalse(decode.isEngaged()); + } + + @Test + void otherFailuresPropagate() { + assertThrows(NullPointerException.class, () -> new GuardedBase64Decode().tryApply(null)); + } + + @TableTest({ + "scenario | input | expected", + "empty | '' | '' ", + "two-char unit, unpadded | YQ | a ", + "three-char unit, unpadded | YWI | ab ", + "full unit | YWJj | abc ", + "two-char unit, padded | 'YQ==' | a ", + "three-char unit, padded | 'YWI=' | ab ", + "non-zero trailing bits | 'YR==' | a ", + "one char | Y | ", + "one-char final unit | YWJjZ | ", + "dangling char, padded | 'Y=' | ", + "lone padding | '=' | ", + "padding starts a unit | 'YWJj=' | ", + "two-char unit, one pad | 'YQ=' | ", + "two-char unit, pad, char | 'YQ=Y' | ", + "char after padding | 'YQ==Y' | ", + "unit after padding | 'YQ==YQ==' | ", + "space | 'YW Jj' | ", + "url-safe alphabet | 'YQ-_' | " + }) + void decodeOrNullFollowsTheJdkRules(String input, String expected) { + assertEquals(expected, GuardedBase64Decode.decodeOrNull(bytes(input))); + assertAgreesWithTheJdk(bytes(input)); + } + + @Test + void decodeOrNullAgreesWithTheJdkOnEveryShortInputOverAReducedAlphabet() { + // every placement of padding, of alphabet characters with low and high bits set, and of an + // illegal character, up to two full units + byte[] symbols = bytes("AQ/=!"); + for (int length = 0; length <= 8; length++) { + byte[] src = new byte[length]; + int[] digits = new int[length]; + while (true) { + for (int i = 0; i < length; i++) { + src[i] = symbols[digits[i]]; + } + assertAgreesWithTheJdk(src.clone()); + int i = 0; + while (i < length && ++digits[i] == symbols.length) { + digits[i++] = 0; + } + if (i == length) { + break; + } + } + } + } + + @Test + void decodeOrNullAgreesWithTheJdkOnRandomEncodingsAndTheirMutations() { + Random random = new Random(12672); + Base64.Encoder[] encoders = {Base64.getEncoder(), Base64.getEncoder().withoutPadding()}; + byte[] mutations = bytes("=!-_ \nA/"); + for (int n = 0; n < 20_000; n++) { + byte[] data = new byte[random.nextInt(48)]; + random.nextBytes(data); + byte[] encoded = encoders[n & 1].encode(data); + assertAgreesWithTheJdk(encoded); + + if (encoded.length > 0) { + byte[] mutated = encoded.clone(); + mutated[random.nextInt(mutated.length)] = mutations[random.nextInt(mutations.length)]; + assertAgreesWithTheJdk(mutated); + + byte[] truncated = new byte[random.nextInt(encoded.length)]; + System.arraycopy(encoded, 0, truncated, 0, truncated.length); + assertAgreesWithTheJdk(truncated); + } + + byte[] noise = new byte[random.nextInt(16)]; + random.nextBytes(noise); + assertAgreesWithTheJdk(noise); + } + } +} diff --git a/internal-api/src/test/java/datadog/trace/util/AdaptiveLatchTest.java b/internal-api/src/test/java/datadog/trace/util/AdaptiveLatchTest.java new file mode 100644 index 00000000000..cf8422325e0 --- /dev/null +++ b/internal-api/src/test/java/datadog/trace/util/AdaptiveLatchTest.java @@ -0,0 +1,192 @@ +package datadog.trace.util; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.Test; + +class AdaptiveLatchTest { + + /** + * Parses an int. The cautious path accepts only plain digits and rejects anything else, so + * "sneaky" input such as an overflowing number passes neither path; "1_000" is repaired by the + * cautious path, which the optimistic path rejects. + */ + private static final class Parsing extends AdaptiveLatch { + final AtomicInteger optimistic = new AtomicInteger(); + final AtomicInteger cautious = new AtomicInteger(); + + Parsing() { + super(NumberFormatException.class, 3); + } + + @Override + protected Integer apply(String input) { + optimistic.incrementAndGet(); + if ("boom".equals(input)) { + throw new IllegalStateException("boom"); + } + return Integer.parseInt(input); + } + + @Override + protected Integer applySafely(String input) { + cautious.incrementAndGet(); + String digits = input.replace("_", ""); + if (digits.isEmpty() || digits.length() > 9) { + return reject(input); + } + int value = 0; + for (int i = 0; i < digits.length(); i++) { + char c = digits.charAt(i); + if (c < '0' || c > '9') { + return reject(input); + } + value = value * 10 + (c - '0'); + } + return value; + } + } + + @Test + void disengagedItTakesOnlyTheOptimisticPath() { + Parsing latch = new Parsing(); + + assertEquals(42, latch.tryApply("42")); + assertEquals(7, latch.tryApply("7")); + + assertEquals(2, latch.optimistic.get()); + assertEquals(0, latch.cautious.get()); + assertFalse(latch.isEngaged()); + } + + @Test + void anOptimisticFailureEngagesAndRetriesTheSameInputCautiously() { + Parsing latch = new Parsing(); + + assertNull(latch.tryApply("bad")); + + assertEquals(1, latch.optimistic.get()); + assertEquals(1, latch.cautious.get()); + assertTrue(latch.isEngaged()); + } + + @Test + void theCautiousRetryCanRepairWhatTheOptimisticPathRejected() { + Parsing latch = new Parsing(); + + assertEquals(1000, latch.tryApply("1_000")); + assertTrue(latch.isEngaged()); + } + + @Test + void engagedItTakesOnlyTheCautiousPath() { + Parsing latch = new Parsing(); + latch.tryApply("bad"); + + assertEquals(42, latch.tryApply("42")); + assertNull(latch.tryApply("bad")); + + assertEquals(1, latch.optimistic.get(), "no optimistic call, and so no throw, while engaged"); + assertEquals(3, latch.cautious.get()); + } + + @Test + void disengagesAfterEnoughConsecutiveCleanCalls() { + Parsing latch = new Parsing(); + latch.tryApply("bad"); + + latch.tryApply("1"); + latch.tryApply("2"); + assertTrue(latch.isEngaged()); + latch.tryApply("3"); + assertFalse(latch.isEngaged()); + + assertEquals(4, latch.tryApply("4")); + assertEquals(2, latch.optimistic.get(), "disengaged again, so back on the optimistic path"); + } + + @Test + void aRejectionWhileEngagedRestartsTheCount() { + Parsing latch = new Parsing(); + latch.tryApply("bad"); + latch.tryApply("1"); + latch.tryApply("2"); + + assertNull(latch.tryApply("bad")); + + latch.tryApply("3"); + latch.tryApply("4"); + assertTrue(latch.isEngaged(), "the count restarted from closeAfter"); + latch.tryApply("5"); + assertFalse(latch.isEngaged()); + } + + @Test + void otherExceptionsPropagateWithoutEngaging() { + Parsing latch = new Parsing(); + + assertThrows(IllegalStateException.class, () -> latch.tryApply("boom")); + assertFalse(latch.isEngaged()); + assertEquals(0, latch.cautious.get()); + } + + @Test + void rejectYieldsTheFallbackAndARealNullCountsAsClean() { + AdaptiveLatch latch = + new AdaptiveLatch( + IllegalArgumentException.class, 1) { + @Override + protected String apply(String input) { + if (input.isEmpty()) { + throw new IllegalArgumentException(); + } + return "null".equals(input) ? null : input; + } + + @Override + protected String applySafely(String input) { + if (input.isEmpty()) { + return reject(input); + } + return "null".equals(input) ? null : input; + } + + @Override + protected String fallback(String input) { + return "fallback"; + } + }; + + // failed optimistically, then rejected cautiously + assertEquals("fallback", latch.tryApply("")); + // rejected while engaged + assertEquals("fallback", latch.tryApply("")); + // a cautious call that succeeds with null is clean, and with closeAfter 1 disengages + assertNull(latch.tryApply("null")); + assertFalse(latch.isEngaged()); + } + + @Test + void closeAfterMustBeAtLeastOne() { + assertThrows( + IllegalArgumentException.class, + () -> + new AdaptiveLatch( + IllegalArgumentException.class, 0) { + @Override + protected String apply(String input) { + return input; + } + + @Override + protected String applySafely(String input) { + return input; + } + }); + } +} diff --git a/metadata/supported-configurations.json b/metadata/supported-configurations.json index e981287521e..11e06c79f45 100644 --- a/metadata/supported-configurations.json +++ b/metadata/supported-configurations.json @@ -2340,6 +2340,14 @@ "aliases": [] } ], + "DD_KAFKA_CLIENT_BASE64_DECODING_GUARD_ENABLED": [ + { + "version": "A", + "type": "boolean", + "default": "true", + "aliases": [] + } + ], "DD_KAFKA_CLIENT_PROPAGATION_DISABLED_TOPICS": [ { "version": "A",